From 8aceac20712a28932666177e596a92f3b130aeff Mon Sep 17 00:00:00 2001 From: Kaitao Yang <21039614+ykaitao@users.noreply.github.com> Date: Tue, 3 Feb 2026 00:27:49 -0800 Subject: [PATCH] reduce code duplication (#3877) * reduce code duplication * address reviewer feedback: keep original function name - Keep original function name `_offload_frozen_module_for_training` - Make `offload_device` parameter Optional (can be None) - Keep original error handling (return None for missing modules_to_save) - Maintain code deduplication by reusing the helper function --------- Co-authored-by: Daniel Han --- unsloth/models/llama.py | 38 +++++++++++--------------------------- 1 file changed, 11 insertions(+), 27 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index fb3b8133a5..f4c057deee 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -152,21 +152,22 @@ from peft.utils.other import ModulesToSaveWrapper def _offload_frozen_module_for_training( module: ModulesToSaveWrapper, device_type: str, - offload_device: str = "cpu", + offload_device: Optional[str] = "cpu", ) -> None: """ Offload frozen module to CPU and configure trainable copy for mixed precision training. This function optimizes memory usage by: 1. Moving the trainable copy to the target device with appropriate precision - 2. Offloading the original frozen module to CPU/disk to free VRAM + 2. Optionally offloading the original frozen module to CPU/disk to free VRAM 3. Converting float16 to float32 for compatibility with certain GPUs (e.g., Tesla T4) Args: module: The module to configure. Must be a ModulesToSaveWrapper with a `modules_to_save` attribute containing trainable and original modules. device_type: Target device string for training (e.g., "cuda:0", "xpu:0") - offload_device: Device to offload frozen parameters (default: "cpu") + offload_device: Device to offload frozen parameters (default: "cpu"). + If None, the original frozen module remains on its current device. Note: Currently only "cpu" is supported; disk offloading is planned. Returns: @@ -174,7 +175,7 @@ def _offload_frozen_module_for_training( Note: - Float16 weights are automatically promoted to float32 for GPU compatibility - - Original frozen parameters are moved to CPU to reduce active VRAM usage + - When offload_device is specified, frozen parameters are moved to free VRAM - Future versions will support disk-based offloading for even larger models See Also: @@ -196,7 +197,8 @@ def _offload_frozen_module_for_training( module.modules_to_save.default.requires_grad_(True) # [TODO] Move old module to CPU - should be disk! - module.original_module.to(device = offload_device, non_blocking = True) + if offload_device is not None: + module.original_module.to(device = offload_device, non_blocking = True) module.original_module.requires_grad_(False) @@ -3083,35 +3085,17 @@ class FastLlamaModel: print("Unsloth: Training embed_tokens in mixed precision to save VRAM") assert hasattr(model.get_input_embeddings(), "modules_to_save") - new_dtype = ( - model.get_input_embeddings().modules_to_save.default.weight.dtype + _offload_frozen_module_for_training( + model.get_input_embeddings(), DEVICE_TYPE_TORCH, offload_device = None ) - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_input_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True - ) - model.get_input_embeddings().modules_to_save.default.requires_grad_(True) if train_lm_head: print("Unsloth: Training lm_head in mixed precision to save VRAM") assert hasattr(model.get_output_embeddings(), "modules_to_save") - new_dtype = ( - model.get_output_embeddings().modules_to_save.default.weight.dtype + _offload_frozen_module_for_training( + model.get_output_embeddings(), DEVICE_TYPE_TORCH, offload_device = None ) - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_output_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True - ) - model.get_output_embeddings().modules_to_save.default.requires_grad_(True) # Patch tokenizer to pad to the right internal_model = model