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()