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 <danielhanchen@gmail.com>
This commit is contained in:
parent
4a8edd5776
commit
8aceac2071
1 changed files with 11 additions and 27 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue