From d34f2ff23f097edf2b77feaf0dc74dd4784db979 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Feb 2025 05:45:05 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 40979ec76f..0e9e28b48c 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -39,10 +39,9 @@ def PatchRL(FastLanguageModel): @contextmanager def unsloth_unwrap_model_for_generation(model, *args, **kwargs): # Must use for_inference to allow inference in Unsloth - FastLanguageModel.for_inference(model) - with torch.inference_mode(): - with unwrap_model_for_generation(model, *args, **kwargs) as unwrapped_model: - yield unwrapped_model + with unwrap_model_for_generation(model, *args, **kwargs) as unwrapped_model: + FastLanguageModel.for_inference(unwrapped_model) + yield unwrapped_model # Return back to training mode FastLanguageModel.for_training (model) pass