Update rl_replacements.py

This commit is contained in:
Daniel Han 2025-02-11 22:02:41 -08:00
commit c643d950f7

View file

@ -110,7 +110,7 @@ def sft_trainer_compute_loss(function_name, function):
)
if len(replacer) != 0:
replacer = replacer[0]
returner = " "*8 + "return (loss, outputs) if return_outputs else loss"
returner = "\n" + " "*8 + "return (loss, outputs) if return_outputs else loss"
function = function.replace(replacer, replacer + returner)
pass
return function