Spaces:
Running on Zero
Running on Zero
File size: 10,094 Bytes
36cdb93 | 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 | """FSDP2 sharding for the block-causal student, and why it needs help.
Both trainers here are data-parallel today: Stage 1 wraps its step in a module
so DDP's hooks fire, and DMD all-reduces gradients by hand. Both keep a full
fp32 copy of the model, its gradients and its AdamW moments on every GPU --
16 bytes per parameter, 21 GB at 1.3B, and 224 GB at 14B, which is why nothing
larger than 1.3B trains on this box at all. Sharding those three tensor sets
across N ranks is the only thing that changes that, so it has to work before a
larger base is worth discussing.
THE TRAP THIS MODULE EXISTS FOR. FSDP2 all-gathers a sharded module's
parameters in a pre-forward hook installed on *that module*. `blockcausal._run`
never calls a `WanAttentionBlock` -- it passes the block to `_layer`, which
reaches into `blk.self_attn.q`, `blk.norm1`, `blk.ffn` and so on. Shard at the
block and that hook never fires, so the layer runs against parameters that are
still 1/N shards. It is structurally the same mistake DDP already invites here
("DDP does not sync gradients if you call the bare model
through helper functions"), and it is why this port is worth debugging at 1.3B.
Measured, on this codebase, the mis-shard is LOUD rather than silent: the first
parameter touched outside a forward is `blk.modulation`, and DTensor refuses
the mixed operand
RuntimeError: aten.add.Tensor got mixed torch.Tensor and DTensor
(scripts/verify_fsdp.py check 3 asserts exactly this). That guard is DTensor's,
not ours, and it only holds while every parameter stays a DTensor; it is not a
reason to leave the shard boundary in the wrong place. The `stragglers` check
in `shard_model` is the part that does not depend on someone else's invariant.
`CausalBlock` and `CausalHead` fix it by being real modules whose forward *is*
the computation, so the shard boundary and the call boundary coincide.
`scripts/verify_fsdp.py` is the control: it asserts a sharded forward matches
the single-GPU forward, and that a deliberately mis-sharded one does not.
What is sharded: the 30 transformer blocks (98% of parameters) plus, when they
are big enough to matter, the head and the embeddings -- all of which are
entered through their own `__call__` and so need no wrapper.
"""
import torch
from torch.distributed.checkpoint.state_dict import (
get_model_state_dict, StateDictOptions
)
from . import blockcausal as bc
class CausalBlock(torch.nn.Module):
"""One WanAttentionBlock as an FSDP shard unit.
Holds the block as a child so `fully_shard(self)` shards the block's
parameters, and runs `_layer` in its own forward so entering the shard is
the same act as entering the computation.
"""
def __init__(self, blk):
super().__init__()
self.blk = blk
def forward(self, x, e0, tbl, kv_ctx, ctx, ctx_lens, dtype, per_frame,
n_frames):
with torch.amp.autocast('cuda', enabled=False):
ec = (self.blk.modulation + e0).chunk(6, dim=1)
return bc._layer(self.blk, x, ec, tbl, kv_ctx, ctx, ctx_lens, dtype,
per_frame, n_frames)
class CausalHead(torch.nn.Module):
"""The output head as an FSDP shard unit, for the same reason: `_head`
reaches through `model.head` to `model.head.head` and reads
`model.head.modulation` directly."""
def __init__(self, head):
super().__init__()
self.head = head
def forward(self, x, e, per_frame, n_frames):
return bc._head_from(self.head, x, e, per_frame, n_frames)
def attach_causal_modules(model):
"""Install the per-layer / head modules `blockcausal._run` will enter, and
make each of them the *only* registered path to its parameters.
Re-registering is the subtle half. `CausalBlock(blk)` holds `blk` as a
child, but `model.blocks[i]` still holds it too, so every block parameter
now has two names in the module tree. FSDP2 shards by walking that tree and
skipping what a child FSDP module already owns; reached by its second name
a parameter looks unmanaged, and sharding it again dies with "Cannot
concatenate overlapping meshes". So `blocks` and `head` are demoted to
plain attributes -- still there for `len(model.blocks)` and for the
unsharded code path, invisible to `named_parameters()`.
Separated from `shard_model` so the correctness control can run the wrapped
path *without* FSDP and show that wrapping alone changes nothing.
"""
if getattr(model, 'causal_layers', None) is not None:
return model
blocks = list(model.blocks)
head = model.head
layers = torch.nn.ModuleList(CausalBlock(b) for b in blocks)
chead = CausalHead(head)
del model.blocks, model.head # drop from _modules
model.causal_layers = layers
model.causal_head = chead
# object.__setattr__, because nn.Module.__setattr__ would re-register an
# nn.Module value and undo exactly what the `del` above achieved.
object.__setattr__(model, 'blocks', blocks)
object.__setattr__(model, 'head', head)
return model
def shard_model(model, mesh=None, reshard_after_forward=True, mp_policy=None,
ignored_params=None):
"""Shard a WanModel across `mesh` with FSDP2, in place. Returns the model.
Call BEFORE moving to device and before building the optimizer: FSDP2
replaces `.weight` with a DTensor, and an optimizer built over the
unsharded parameters would hold stale references.
`reshard_after_forward=False` keeps parameters gathered between forward and
backward. That trades memory for collectives, and it matters here more than
in ordinary training: one DMD iteration runs the student a few dozen times
inside its rollout, so resharding after each one pays 30 all gathers per
forward to reclaim memory the (no_grad) rollout never needed back.
There is deliberately no root `fully_shard(model)`. The root's hook fires
on `model.forward`, and nothing here ever calls it. `block_forward` drives
the model from outside. A parameter caught only by the root group would
therefore stay sharded through every forward, silently. Everything is
sharded at a module that is genuinely entered instead, and the assertion at
the end refuses to return a model where that did not hold.
"""
from torch.distributed.fsdp import fully_shard
# An explicit, NAMED mesh. `fully_shard(mesh=None)` synthesises one per
# call with `mesh_dim_names=None`, and composing shards over children then
# dies in `_init_sharded_param` trying to concatenate those names.
if mesh is None:
from torch.distributed.device_mesh import init_device_mesh
import torch.distributed as dist
mesh = init_device_mesh('cuda', (dist.get_world_size(),),
mesh_dim_names=('dp',))
kw = {'mesh': mesh, 'reshard_after_forward': reshard_after_forward}
if mp_policy is not None:
kw['mp_policy'] = mp_policy
# `ignored_params` is how the LoRA critic survives sharding its own base.
# The frozen teacher is bf16 with no gradients; the adapter living inside
# the same blocks is fp32 and trainable, and it is 32 M parameters (small
# enough that replicating it and all reducing by hand is cheaper than
# sharding, and it keeps mixed dtypes out of a single FSDP parameter group.
if ignored_params:
kw['ignored_params'] = set(ignored_params)
attach_causal_modules(model)
units = list(model.causal_layers) + [model.causal_head]
# These four are invoked through their own __call__ already (patch_embedding
# in `_run`, text_embedding in the trainers, time_embedding and
# time_projection in `time_embed`), so they need no wrapper.
units += [model.patch_embedding, model.text_embedding,
model.time_embedding, model.time_projection]
for m in units:
fully_shard(m, **kw)
ignore = set(id(p) for p in (ignored_params or ()))
stragglers = [n for n, p in model.named_parameters()
if id(p) not in ignore
and not isinstance(p, torch.distributed.tensor.DTensor)]
if stragglers:
raise RuntimeError(
f'{len(stragglers)} parameter(s) were not sharded and are not '
f'reachable from any entered module -- they would train '
f'unsynchronised: {stragglers[:8]}')
return model
def set_grad_sync(model, flag):
"""FSDP2's equivalent of `DDP.no_sync()`, over every shard unit.
With gradient accumulation, leaving this on costs one reduce-scatter per
micro-step instead of one per optimizer step. It is not a correctness issue
-- reducing each micro-step and summing equals summing then reducing but
at `--accum 4` it is four times the collectives for the same gradient.
"""
for m in model.modules():
if hasattr(m, 'set_requires_gradient_sync'):
m.set_requires_gradient_sync(flag)
def full_state_dict(model):
"""Unsharded fp32 state dict on CPU, in the same key layout the existing
checkpoints use, so `demo.py --weights` keeps working unchanged.
Under FSDP2 `model.state_dict()` returns DTensors -- saving those would
produce a checkpoint that only reloads onto an identical mesh, and every
evaluation script here loads onto one GPU.
"""
sd = get_model_state_dict(
model, options=StateDictOptions(full_state_dict=True, cpu_offload=True))
# `attach_causal_modules` renamed the shard units, so the keys now read
# `causal_layers.7.blk.*` and `causal_head.head.*`. Put them back under the
# names WanModel.load_state_dict expects, so demo.py and every other
# evaluation script keep loading these checkpoints on a single GPU.
out = {}
for k, v in sd.items():
if k.startswith('causal_layers.'):
i, rest = k[len('causal_layers.'):].split('.blk.', 1)
k = f'blocks.{i}.{rest}'
elif k.startswith('causal_head.head.'):
k = 'head.' + k[len('causal_head.head.'):]
out[k] = v
return out
|