From 635cfdbbb0664776c2907d455de06a69ebd4ebd2 Mon Sep 17 00:00:00 2001 From: Datta Nimmaturi Date: Thu, 23 Oct 2025 14:13:55 +0530 Subject: [PATCH] Sleep trl patch (#3494) * Patch sleep mode properly for trl * empty cache after sleep/wakeup * no extra wakeups * Do not redo wakeups * cleanup --- unsloth/models/rl_replacements.py | 60 +++++-------------------------- 1 file changed, 8 insertions(+), 52 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 996c4f62b0..87287947a0 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -181,40 +181,6 @@ RL_FUNCTIONS["sft_trainer"].append(sft_trainer_compute_loss) def grpo_trainer__prepare_inputs(function_name, function): if function_name != "_prepare_inputs": return function - import re - # This matches the function signature, decorators and any comments immediately following - pattern = r"(\s*@profiling_decorator\s*\n\s*def _prepare_inputs\s*\([^\)]*\)\s*(->\s*[^:]+)?\s*:\s*\n(?:[ ]*#[^\n]*\n)*)" - - match = re.search(pattern, function) - insert = ( - " if hasattr(self, 'llm'):\n" - " if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n" - " self.llm.wake_up()\n" - ) - if match: - header_and_comments = match.group(1) - # Find where the code block starts after comments - code_start_index = match.end(1) - rest_of_function = function[code_start_index:] - - # Remove any old wake_up call that might be at the start of the function body - rest_of_function = re.sub( - r"^\s*if hasattr\(self, 'llm'\):.*?self\.llm\.wake_up\(\).*?\n", - "", - rest_of_function, - flags=re.DOTALL | re.MULTILINE - ) - - # We also need to remove the old wake up call from the beginning of the function - # since it's injected before the comments. - header_and_comments = re.sub( - r"(:\s*\n)\s*if hasattr\(self, 'llm'\):.*?self\.llm\.wake_up\(\).*?\n", - r"\1", - header_and_comments, - flags=re.DOTALL | re.MULTILINE - ) - - function = header_and_comments + insert + rest_of_function # Add mixed precision training function = function.replace( "with torch.inference_mode():", @@ -228,16 +194,6 @@ def grpo_trainer__prepare_inputs(function_name, function): "self.accelerator.unwrap_model(self.model)", "self.accelerator.unwrap_model(self.model, keep_fp32_wrapper = False)", ) - sleep_and_cache = ( - "if hasattr(self, 'llm'):\n" - " if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n" - " self.llm.sleep(os.environ.get('VLLM_SLEEP_MODE', 1))\n" - " " - ) - if re.search(r"\n\s*return ", function): - function = re.sub(r"(\n\s*)return ", f"\\1{sleep_and_cache}return ", function, count=1) - else: - function = function.rstrip() + "\n " + sleep_and_cache return function pass RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__prepare_inputs) @@ -320,11 +276,11 @@ def grpo_trainer__generate_and_score_completions(function_name, function): wrapped = ( f"{indent}if hasattr(self, 'llm'):\n" f"{indent} if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n" - f"{indent} self.llm.wake_up()\n\n" + f"{indent} self.llm.wake_up()\n" f"{line}\n\n" f"{indent}if hasattr(self, 'llm'):\n" f"{indent} if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n" - f"{indent} self.llm.sleep(os.environ.get('VLLM_SLEEP_MODE', 1))" + f"{indent} self.llm.sleep(os.environ.get('VLLM_SLEEP_MODE', 1))\n" ) patched = patched[:match.start()] + wrapped + patched[match.end():] @@ -446,13 +402,13 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): if function_name != "_get_per_token_logps_and_entropies": return function # Just copy over from _get_per_token_logps replacement function above. For now this returns None anyway - def _get_per_token_logps_and_entropies(self, model, input_ids, attention_mask, logits_to_keep, batch_size = None, + def _get_per_token_logps_and_entropies(self, model, input_ids, attention_mask, logits_to_keep, batch_size = None, compute_entropy = False, compute_efficient = False, *args, **kwargs): # if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0': # return None, None # logps, entropies Unsloth efficient GRPO if compute_efficient: return None, None - else: + else: # Otherwise, calculate normally: if not hasattr(self, '_autocast_dtype'): self._autocast_dtype = torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16 @@ -462,12 +418,12 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): pixel_attention_mask, image_sizes = kwargs.get('pixel_attention_mask',None), kwargs.get('image_sizes',None) os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" - + unwrapped_model = self.accelerator.unwrap_model(model, keep_fp32_wrapper=False) with torch.amp.autocast(device_type = 'cuda', dtype = self._autocast_dtype): with torch.inference_mode(): - if pixel_values is None: + if pixel_values is None: attention_mask = input_ids != self.processing_class.pad_token_id attention_mask = attention_mask.to(attention_mask.dtype) # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded @@ -490,7 +446,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): image_sizes = image_sizes, logits_to_keep = logits_to_keep + 1, ).logits - + entropies = None if compute_entropy: @@ -580,7 +536,7 @@ def grpo_trainer_compute_loss(function_name, function): # per_token_loss = -(per_token_loss - self.beta * per_token_kl) # loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean() old_hidden_states = inputs.get("old_per_token_logps", None) - + input_ids = input_ids[:, -logits_to_keep:] # Get logit softcapping and logit scale