Remove docs

This commit is contained in:
Daniel Han 2025-02-15 17:36:18 -08:00
commit e1130a41e3
2 changed files with 20 additions and 3 deletions

View file

@ -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

View file

@ -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