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")]