diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index ac1b836673..b13e6f9c78 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -548,6 +548,7 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import changed = {"__init__" : (old_init, init,)} edit_functions = RL_FUNCTIONS.get(trainer_file, []) + remover = [] for function in functions: if not hasattr(RLTrainer, function): continue @@ -591,7 +592,9 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import ) # Skip if no changes done - if source == original_source: continue + if source == original_source: + remover.append(original_source) + continue # Find all imports imports += [x for x in all_imports if not x.startswith("_") and x in source] @@ -607,9 +610,23 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import old, new = changed[function] RLTrainer_source = RLTrainer_source.replace(old, new) pass + + # Remove non editted functions + for remove in remover: + RLTrainer_source = RLTrainer_source.replace(remove, "\n") + pass + RLTrainer_source = RLTrainer_source.replace( f"class {RLTrainer_name}", f"class _Unsloth{RLTrainer_name}", 1 ) + + # Get rid of docs since we repeated it + RLTrainer_source = re.sub( + rf"class _Unsloth{RLTrainer_name}:.+?def __init__\(", + rf"class _Unsloth{RLTrainer_name}:\n def __init__(", + RLTrainer_source, + flags = re.MULTILINE | re.DOTALL, + ) return RLTrainer_source pass diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index b9ba34726a..46d44b92f6 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -40,7 +40,7 @@ torch_compile_options = { } # Check untrained tokens -def sft_trainer_fix_untraiend_tokens(call_args, extra_args): +def sft_trainer_fix_untrained_tokens(call_args, extra_args): if "model" in call_args and "train_dataset" in call_args: fix_tokenizer = \ "IGNORED_TOKENIZER_NAMES = os.environ.get('UNSLOTH_IGNORED_TOKENIZER_NAMES', '').split('\\n')\n"\ @@ -52,7 +52,7 @@ def sft_trainer_fix_untraiend_tokens(call_args, extra_args): return fix_tokenizer return "" pass -RL_EXTRA_ARGS["sft_trainer"].append(sft_trainer_fix_untraiend_tokens) +RL_EXTRA_ARGS["sft_trainer"].append(sft_trainer_fix_untrained_tokens) # Remove DPO columns which might randomnly be tokenized