Spaces:
Running on Zero
Running on Zero
File size: 8,585 Bytes
b347817 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 | """
Layer offloading for training — hook-based CPU offload.
Hook lifecycle per module:
No checkpointing:
forward_pre(load→GPU) → forward() → forward_hook(offload)
backward_pre(load→GPU) → backward() → backward_hook(offload)
With checkpointing (use_reentrant=False):
First forward:
forward_pre(load→GPU) → forward() → forward_hook(offload, mark in_recompute)
Recomputing forward (during backward):
forward_pre(load→GPU) → forward() → forward_hook(in_recompute=True → keep GPU)
Backward:
backward_pre(load→GPU) → backward() → backward_hook(offload)
"""
import torch
import torch.nn as nn
import warnings
from .offloader import StaticParamOffloader, TrainableParamOffloader, AlwaysOnGPUParamOffloader, BufferOffloader
from .memory_buffer import PinnedArenaPool, BaseBufferPool
warnings.filterwarnings("ignore", message="Full backward hook is firing when gradients are computed with respect to module outputs")
def has_parameters(module: nn.Module) -> bool:
return len(list(module.parameters())) > 0
def count_parameters(module: nn.Module) -> int:
return sum(p.numel() for p in module.parameters())
def is_leaf_module(module: nn.Module) -> bool:
return len(list(module.children())) == 0
class UnitWiseParamManager:
def __init__(self, model: nn.Module, target_device: torch.device, enable_optimizer_cpu_offload: bool = False, params: list = None, buffers: list = None, memory_buffer: BaseBufferPool = None):
self.model = model
self.target_device = target_device
self.param_offloaders = {}
for param in (model.parameters() if params is None else params):
if not param.requires_grad:
self.param_offloaders[id(param)] = StaticParamOffloader(param, target_device, memory_buffer=memory_buffer)
else:
if enable_optimizer_cpu_offload:
self.param_offloaders[id(param)] = TrainableParamOffloader(param, target_device)
else:
self.param_offloaders[id(param)] = AlwaysOnGPUParamOffloader(param, target_device)
if buffers is not None and len(buffers) > 0:
for mod, buf_name, buf in buffers:
self.param_offloaders[id(buf)] = BufferOffloader(mod, buf_name, buf, target_device, memory_buffer=memory_buffer)
def move_gradients_to_cpu(self):
for offloader in self.param_offloaders.values():
offloader.offload_grad()
def onload_module(self, module: nn.Module):
for param in module.parameters(recurse=False):
if id(param) in self.param_offloaders:
self.param_offloaders[id(param)].onload()
for name, buf in module.named_buffers(recurse=False):
if id(buf) in self.param_offloaders:
self.param_offloaders[id(buf)].onload()
def offload_module(self, module: nn.Module):
for param in module.parameters(recurse=False):
if id(param) in self.param_offloaders:
self.param_offloaders[id(param)].offload()
for name, buf in module.named_buffers(recurse=False):
if id(buf) in self.param_offloaders:
self.param_offloaders[id(buf)].offload()
class UnitWiseHookManager:
def __init__(self, model: nn.Module, target_device: torch.device, enable_optimizer_cpu_offload: bool = False,
params: list = None, buffers: list = None, memory_buffer: BaseBufferPool = None):
self.param_manager = UnitWiseParamManager(model, target_device, enable_optimizer_cpu_offload, params=params, buffers=buffers, memory_buffer=memory_buffer)
self._in_recompute: set = set()
self._register_hooks(model)
def _register_hooks(self, module: nn.Module):
def forward_pre_hook(mod, args):
self.param_manager.onload_module(mod)
def forward_hook(mod, args, output):
if mod in self._in_recompute:
return
self._in_recompute.add(mod)
self.param_manager.offload_module(mod)
def backward_pre_hook(mod, grad_output):
self.param_manager.onload_module(mod)
def backward_hook(mod, grad_input, grad_output):
self.param_manager.offload_module(mod)
module.register_forward_pre_hook(forward_pre_hook)
module.register_forward_hook(forward_hook)
module.register_full_backward_pre_hook(backward_pre_hook)
if is_leaf_module(module):
module.register_full_backward_hook(backward_hook)
else:
# Parent module backward_hook fires before child backward completes.
# Register on leaf children instead.
sub_modules = [m for m in module.modules() if is_leaf_module(m) and has_parameters(m)]
for sub_mod in sub_modules:
sub_mod.register_full_backward_hook(backward_hook)
def after_backward(self):
self._in_recompute.clear()
self.param_manager.move_gradients_to_cpu()
@property
def managed_param_ids(self):
return set(self.param_manager.param_offloaders.keys())
class OffloadTrainingManager:
def __init__(self, model: nn.Module, target_device: torch.device, enable_optimizer_cpu_offload: bool = False, cpu_offload_split_threshold: int = None):
self.model = model
self.target_device = target_device
self.enable_optimizer_cpu_offload = enable_optimizer_cpu_offload
cpu_offload_split_threshold = cpu_offload_split_threshold * 1024 * 1024 if cpu_offload_split_threshold is not None else None
self._register_units(model, target_device, enable_optimizer_cpu_offload, cpu_offload_split_threshold)
def _register_units(self, model: nn.Module, target_device: torch.device, enable_optimizer_cpu_offload: bool, cpu_offload_split_threshold: int = None):
self.memory_buffer = PinnedArenaPool.from_model(model)
units = self._find_units_recursive(model, cpu_offload_split_threshold)
self.units = [UnitWiseHookManager(u, target_device, enable_optimizer_cpu_offload, memory_buffer=self.memory_buffer) for u in units]
managed_param_ids = set().union(*[unit.managed_param_ids for unit in self.units])
orphan_params, orphan_buffers = self._find_orphan_params_and_buffers(model, managed_param_ids)
for orphan_module in set(orphan_params.keys()) | set(orphan_buffers.keys()):
params = orphan_params.get(orphan_module, [])
buffers = orphan_buffers.get(orphan_module, [])
self.units.append(UnitWiseHookManager(orphan_module, target_device, enable_optimizer_cpu_offload, params=params, buffers=buffers, memory_buffer=self.memory_buffer))
def _find_orphan_params_and_buffers(self, model: nn.Module, managed_param_ids: set):
orphan_params_by_module = {}
for _, mod in model.named_modules():
for param in mod.parameters(recurse=False):
if id(param) not in managed_param_ids:
orphan_params_by_module.setdefault(mod, []).append(param)
# Collect orphan buffers grouped by owner module
orphan_buffers_by_module = {}
for _, mod in model.named_modules():
for name, buf in mod.named_buffers(recurse=False):
orphan_buffers_by_module.setdefault(mod, []).append((mod, name, buf))
return orphan_params_by_module, orphan_buffers_by_module
def _find_units_recursive(self, module: nn.Module, cpu_offload_split_threshold: int = None) -> list:
if cpu_offload_split_threshold is None:
return [m for m in module.modules() if is_leaf_module(m) and has_parameters(m)]
if self._should_force_recurse(module, cpu_offload_split_threshold):
units = []
for child in module.children():
units.extend(self._find_units_recursive(child, cpu_offload_split_threshold))
return units
return [module]
def _should_force_recurse(self, module: nn.Module, cpu_offload_split_threshold: int = None) -> bool:
if is_leaf_module(module):
return False
if (
count_parameters(module) > cpu_offload_split_threshold
or ('forward' not in type(module).__dict__)
or (hasattr(module, 'encode') and hasattr(module, 'decode'))
):
return True
return False
# run after backward() and before optimizer.step()
def after_backward(self):
for unit in self.units:
unit.after_backward()
torch.cuda.synchronize()
|