From aebeeb4901a61434d1630ccf4e8ea13ef3eaf9d3 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Feb 2025 06:58:04 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 16 +++++++--------- 1 file changed, 7 insertions(+), 9 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 72b568790d..2431e5a70f 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -39,15 +39,13 @@ def PatchRL(FastLanguageModel): @contextmanager def unsloth_unwrap_model_for_generation(model, accelerator): # Must use for_inference to allow inference in Unsloth - with torch.inference_mode(): - with unwrap_model_for_generation(model, accelerator) as unwrapped_model: - FastLanguageModel.for_inference(unwrapped_model) - try: - yield unwrapped_model - finally: - # Finally return back training - FastLanguageModel.for_training(model) - pass + with unwrap_model_for_generation(model, accelerator) as unwrapped_model: + FastLanguageModel.for_inference(unwrapped_model) + try: + yield unwrapped_model.eval() + finally: + # Finally return back training + FastLanguageModel.for_training(model) pass pass pass