File size: 7,780 Bytes
619344d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
# Copyright (c) 2025 Hansheng Chen

import gc
import contextlib
import functools
import torch
import torch.distributed as dist

from torch.distributed.tensor import DTensor
try:
    from torch.distributed.fsdp import FSDPModule
except:
    pass
from functools import partial
from six.moves import map, zip
from mmcv.parallel import is_module_wrapper
from peft.tuners.lora import LoraLayer


def multi_apply(func, *args, **kwargs):
    """Apply function to a list of arguments.

    Note:
        This function applies the ``func`` to multiple inputs and
        map the multiple outputs of the ``func`` into different
        list. Each list contains the same type of outputs corresponding
        to different inputs.

    Args:
        func (Function): A function that will be applied to a list of
            arguments

    Returns:
        tuple(list): A tuple containing multiple list, each list contains \
            a kind of returned results by the function
    """
    pfunc = partial(func, **kwargs) if kwargs else func
    map_results = map(pfunc, *args)
    return tuple(map(list, zip(*map_results)))


def reduce_mean(tensor):
    """Obtain the mean of tensor on different GPUs."""
    if not (dist.is_available() and dist.is_initialized()):
        return tensor
    tensor = tensor.clone()
    dist.all_reduce(tensor.div_(dist.get_world_size()), op=dist.ReduceOp.SUM)
    return tensor


def rgetattr(obj, attr, *args):
    def _getattr(obj, attr):
        if is_module_wrapper(obj):
            obj = obj.module
        return getattr(obj, attr, *args)
    return functools.reduce(_getattr, [obj] + attr.split('.'))


def rsetattr(obj, attr, val):
    pre, _, post = attr.rpartition('.')
    pre = rgetattr(obj, pre) if pre else obj
    if is_module_wrapper(pre):
        pre = pre.module
    return setattr(pre, post, val)


def rhasattr(obj, attr):
    return rgetattr(obj, attr, None) is not None


def rdelattr(obj, attr):
    pre, _, post = attr.rpartition('.')
    pre = rgetattr(obj, pre) if pre else obj
    if is_module_wrapper(pre):
        pre = pre.module
    return delattr(pre, post)


class module_requires_grad:
    def __init__(self, module, requires_grad=True):
        self.module = module
        self.requires_grad = requires_grad
        self.prev = []

    def __enter__(self):
        for p in self.module.parameters():
            self.prev.append(p.requires_grad)
            p.requires_grad = self.requires_grad

    def __exit__(self, exc_type, exc_value, traceback):
        for p, r in zip(self.module.parameters(), self.prev):
            p.requires_grad = r


class module_eval:
    def __init__(self, module):
        self.module = module
        self.prev = None

    def __enter__(self):
        self.prev = self.module.training
        self.module.train(False)

    def __exit__(self, exc_type, exc_value, traceback):
        self.module.train(self.prev)


def all_frozen(modules):
    for module in modules:
        for p in module.parameters():
            if p.requires_grad:
                return False
    return True


def tie_untrained_submodules(tgt_module, src_module, tie_tgt_lora_base_layer=False):
    for key, src_submodule in src_module._modules.items():
        if key in tgt_module._modules:
            if (tie_tgt_lora_base_layer
                    and isinstance(tgt_module._modules[key], LoraLayer)
                    and not isinstance(src_submodule, LoraLayer)):
                if all_frozen((tgt_module._modules[key]._modules['base_layer'], src_submodule)):
                    tgt_module._modules[key]._modules['base_layer'] = src_submodule
                else:
                    tie_untrained_submodules(
                        tgt_module._modules[key]._modules['base_layer'], src_submodule, tie_tgt_lora_base_layer)
            else:
                if all_frozen((tgt_module._modules[key], src_submodule)):
                    tgt_module._modules[key] = src_submodule
                else:
                    tie_untrained_submodules(
                        tgt_module._modules[key], src_submodule, tie_tgt_lora_base_layer)


def clone_params(tgt_module, src_module, recursive=True):
    """Clone parameters and buffers from src_module to tgt_module (sharing the same structure).
    Tied parameters/buffers are not cloned. Used for EMA model initialization.
    """
    for key, val in src_module._parameters.items():
        if (val is not None) \
                and (val is not tgt_module._parameters[key]):
            tgt_module._parameters[key] = val.clone()
    for key, val in src_module._buffers.items():
        if val is not tgt_module._buffers[key]:
            tgt_module._buffers[key] = val.clone()
    if recursive:
        for key, val in src_module._modules.items():
            clone_params(
                tgt_module._modules[key], val, recursive)


@torch.no_grad()
def kai_zhang_clip_grad(model, max_norm: float) -> float:
    """
    https://github.com/Kai-46/minFM/blob/385568691b021aa4888269dafa671900daf17cf6/utils/clip_grad.py#L9
    """
    shard_size, replicate_factor = 1, dist.get_world_size()
    if isinstance(model, FSDPModule):
        shard_size = model._get_fsdp_state()._fsdp_param_group.mesh_info.shard_mesh_size
        replicate_factor = dist.get_world_size() // shard_size

    # Separate DTensor and non-DTensor parameters
    all_param_grads = []
    dtensor_param_grads = []
    regular_param_grads = []

    for p in model.parameters():
        if (not p.requires_grad) or (p.grad is None):
            continue

        if isinstance(p.grad.data, DTensor):
            local_p_grad = p.grad.data.to_local()
            dtensor_param_grads.append(local_p_grad.ravel())
        else:
            local_p_grad = p.grad.data
            regular_param_grads.append(local_p_grad.ravel())

        all_param_grads.append(local_p_grad)

    # Compute local square sum for each group separately
    local_sq_sum = torch.tensor(0.0, device=all_param_grads[0].device)

    if dtensor_param_grads:
        dtensor_sq_sum = (torch.cat(dtensor_param_grads, dim=0) ** 2).float().sum()
        local_sq_sum = local_sq_sum + dtensor_sq_sum

    if regular_param_grads:
        regular_sq_sum = (torch.cat(regular_param_grads, dim=0) ** 2).float().sum()
        local_sq_sum = local_sq_sum + regular_sq_sum / shard_size

    # Single all-reduce operation
    global_sq_sum = local_sq_sum.clone()
    dist.all_reduce(global_sq_sum, op=dist.ReduceOp.SUM)
    global_sq_sum = global_sq_sum / replicate_factor

    total_norm = global_sq_sum.sqrt().item()

    # Only apply clipping when exceeding threshold
    if total_norm > max_norm:
        clip_factor = max_norm / total_norm
        torch._foreach_mul_(all_param_grads, clip_factor)

    return total_norm


def materialize_meta_states(module, device=None):
    if device is None:
        device = torch.get_default_device()
    assert device != torch.device('meta'), 'Please specify a non-meta device.'
    for mod in module.modules():
        for name, p in list(mod._parameters.items()):
            if p is not None and p.is_meta:
                new = torch.empty_like(p, device=device)
                new = torch.nn.Parameter(new, requires_grad=p.requires_grad)
                mod._parameters[name] = new
        for name, b in list(mod._buffers.items()):
            if b is not None and b.is_meta:
                new = torch.empty_like(b, device=device)
                mod._buffers[name] = new
    return module


@contextlib.contextmanager
def gc_context(enable=False):
    prev_enabled = gc.isenabled()
    if enable:
        gc.enable()
    else:
        gc.disable()
    try:
        yield
    finally:
        if prev_enabled:
            gc.enable()
        else:
            gc.disable()