From dbe72263a0dd0a678f0227edd7d6dcd050915640 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Feb 2025 07:25:05 -0800 Subject: [PATCH] Update rl.py --- unsloth/models/rl.py | 33 ++++++++++++++++++++++----------- 1 file changed, 22 insertions(+), 11 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 06634ae3cd..88db94bdf8 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -33,27 +33,38 @@ from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union def PatchRL(FastLanguageModel): - from trl.models import unwrap_model_for_generation + from trl.models.utils 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) + def unsloth_unwrap_model_for_generation(model, *args, **kwargs): with unwrap_model_for_generation(model, accelerator) as unwrapped_model: + # Put the model in inference mode. + FastLanguageModel.for_inference(unwrapped_model) + + # Monkey-patch the generate method so it clones its output. + original_generate = unwrapped_model.generate + + def generate_with_clone(*args, **kwargs): + out = original_generate(*args, **kwargs) + # If the output is a tensor (i.e. an inference tensor), clone it. + if isinstance(out, torch.Tensor): + return out.clone() + # Optionally, if out is a tuple or dict containing tensors, you + # might want to iterate over it and clone all tensors. + return out + + # Replace the generate method. + unwrapped_model.generate = generate_with_clone + try: yield unwrapped_model finally: - # Finally return back training + # Restore the original generate method and reset the model mode. + unwrapped_model.generate = original_generate FastLanguageModel.for_training(model) - pass - 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")]