From 74237834620640e32fde96bdf232f3398e5de083 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Feb 2025 06:37:32 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 2282e8b313..b51be3b7fa 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -39,11 +39,14 @@ def PatchRL(FastLanguageModel): @contextmanager def unsloth_unwrap_model_for_generation(model, *args, **kwargs): # Must use for_inference to allow inference in Unsloth - with unwrap_model_for_generation(model, *args, **kwargs) as unwrapped_model: - FastLanguageModel.for_inference(unwrapped_model) + FastLanguageModel.for_inference(model) + try: + unwrapped_model = unwrap_model_for_generation(model, *args, **kwargs) yield unwrapped_model - # Return back to training mode - FastLanguageModel.for_training(model) + finally: + # Finally return back training + FastLanguageModel.for_training(model) + pass pass import trl.trainer