From 1022dcfa6ec2ed35007dba8912a893be16eeedcc Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Feb 2025 07:13:40 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 2431e5a70f..06634ae3cd 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -33,16 +33,16 @@ from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union def PatchRL(FastLanguageModel): - from trl.models.utils import unwrap_model_for_generation + from trl.models import unwrap_model_for_generation from contextlib import contextmanager @contextmanager def unsloth_unwrap_model_for_generation(model, accelerator): # Must use for_inference to allow inference in Unsloth + FastLanguageModel.for_inference(model) with unwrap_model_for_generation(model, accelerator) as unwrapped_model: - FastLanguageModel.for_inference(unwrapped_model) try: - yield unwrapped_model.eval() + yield unwrapped_model finally: # Finally return back training FastLanguageModel.for_training(model) @@ -50,6 +50,10 @@ def PatchRL(FastLanguageModel): pass pass + import trl.models + trl.models.utils.unwrap_model_for_generation = unwrap_model_for_generation + trl.models.unwrap_model_for_generation = unwrap_model_for_generation + import trl.trainer trainers = dir(trl.trainer) trainers = [x for x in trainers if x.endswith("_trainer")]