Spaces:
Running on Zero
Running on Zero
File size: 16,983 Bytes
40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 5878871 40d7c69 | 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 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 | """DINOv3 ViT encoder with a *robust* block finder, real unfreezing and LLRD.
v2 died here. Its `_blocks()` tried `model.layer / .layers / .blocks` on the
model and on `model.encoder`, found none of them on the HF DINOv3 module, and
raised `AttributeError: cannot locate transformer blocks on the encoder` at
epoch 21 of 30 β the exact epoch the late unfreeze was scheduled. Consequences
in the real run: the encoder was **never** unfrozen, the process died before the
final evaluation, and no TTA / viewer / qualitative artefacts were ever written.
The whole run therefore trained 11 M decoder parameters on top of a frozen
backbone, which is v1's architecture with extra steps.
v3 finds the blocks structurally instead of by name β walk the module tree for
the `nn.ModuleList` whose length equals `config.num_hidden_layers` β so it cannot
break on a naming change. `--freeze_epochs` warm-starts the decoder, then the
**whole** encoder is unfrozen with layer-wise LR decay, which is the standard
ViT dense-prediction recipe and the largest single lever v2 left on the table.
"""
from __future__ import annotations
import torch
from torch import nn
class DINOv3Encoder(nn.Module):
# Overridable so tests (and any future backbone swap) can inject a module
# without reaching the hub. Must expose `.config` and return
# `hidden_states` from forward.
BUILDER = None
def __init__(self, cfg):
super().__init__()
if type(self).BUILDER is not None:
self.model = type(self).BUILDER(cfg)
else:
from transformers import AutoModel
kw = dict(token=cfg.hf_token or None, output_hidden_states=True)
try:
self.model = AutoModel.from_pretrained(
cfg.encoder_model_id, attn_implementation="sdpa", **kw)
except (ValueError, TypeError, ImportError, KeyError):
self.model = AutoModel.from_pretrained(cfg.encoder_model_id, **kw)
c = self.model.config
self.patch = int(getattr(c, "patch_size", 16))
self.hidden = int(c.hidden_size)
self.n_layers = int(getattr(c, "num_hidden_layers", 24))
idx = tuple(cfg.encoder_feature_indices)
if max(idx) > self.n_layers:
idx = tuple(max(1, round(self.n_layers * q)) for q in (0.25, 0.5, 0.75, 1.0))
print(f"[model] taps exceed {self.n_layers} blocks -> {idx}")
self.idx = idx
self.blocks = self._find_blocks()
self._ckpt = bool(cfg.grad_checkpoint_encoder)
self._cl = bool(cfg.channels_last)
self.frozen = True
self.set_frozen(True)
print(f"[model] encoder={cfg.encoder_model_id} hidden={self.hidden} "
f"patch={self.patch} blocks={len(self.blocks) if self.blocks else '?'} "
f"taps={self.idx}")
# -- block discovery ------------------------------------------------
def _find_blocks(self) -> nn.ModuleList | None:
"""The ModuleList of transformer blocks, located by shape not by name."""
exact, any_list = None, None
for _name, mod in self.model.named_modules():
if isinstance(mod, nn.ModuleList) and len(mod) > 0:
if len(mod) == self.n_layers and exact is None:
exact = mod
if any_list is None or len(mod) > len(any_list):
any_list = mod
found = exact if exact is not None else any_list
if found is None:
print("[model] WARNING: no transformer block list found β "
"unfreeze and LLRD will fall back to whole-encoder groups")
return found
# -- freezing -------------------------------------------------------
def set_frozen(self, frozen: bool, top_blocks: int = 0) -> None:
"""Freeze the backbone, or unfreeze `top_blocks` of it (0 = all of it).
A partial unfreeze leaves the patch embedding and the lower blocks with
`requires_grad=False`, which is all three of the places that matters:
`llrd_param_groups` already skips them, DDP does not bucket them, and
AdamW never allocates their two moment tensors.
"""
self.frozen = bool(frozen)
for p in self.model.parameters():
p.requires_grad_(not self.frozen)
if not self.frozen and top_blocks > 0 and self.blocks is not None:
keep = set()
for b in list(self.blocks)[-int(top_blocks):]:
keep.update(id(p) for p in b.parameters())
n_kept = 0
for p in self.model.parameters():
if id(p) in keep:
n_kept += 1
else:
p.requires_grad_(False)
print(f"[model] partial unfreeze: top {min(int(top_blocks), len(self.blocks))}"
f"/{len(self.blocks)} blocks trainable ({n_kept} tensors); "
f"patch embedding and lower blocks stay frozen")
if self.frozen:
self.model.eval()
else:
# Toggled here, once, instead of inside forward(): HF's
# `gradient_checkpointing_enable` walks every submodule of the ViT
# and rebuilds its forwards, and v4 was calling it on every
# training step. It is also off by default on a big card β see
# `grad_checkpoint_encoder`, which trades ~35 % throughput for VRAM
# there is no shortage of on an 80 GB H100.
self._set_grad_checkpointing(self._ckpt)
self._freeze_unreachable()
n = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
print(f"[model] encoder {'FROZEN' if frozen else 'TRAINABLE'} "
f"({n / 1e6:.1f}M grad params)")
# -- parameters the tapped forward cannot reach ---------------------
def _freeze_unreachable(self) -> list[str]:
"""Clear `requires_grad` on backbone parameters no gradient reaches.
Unfreezing hands DDP ~415 encoder parameter tensors and three of them
are not on any path from `pixel_values` to the four taps:
* `embeddings.mask_token` β masked-image-modelling only. It is read
solely when `bool_masked_pos` is passed to the backbone, and
`_hidden_states` never passes it.
* `norm.weight` / `norm.bias` β on some `transformers` releases
`DINOv3ViTModel` applies its final LayerNorm to `last_hidden_state`
on the way out, while the taps come from `hidden_states`, the
*pre-norm* block outputs, so that LayerNorm's result is thrown
away. On others the last tapped hidden state is post-norm and the
gradient does reach it. Which one you are on is not knowable from
here, and it is exactly why the probe below decides rather than the
name list β measured both ways on transformers >= 4.56.
On one GPU those are three tensors AdamW steps with a `None` gradient,
i.e. nothing. Under DDP with `find_unused_parameters=False` they are a
deadlock: the reducer registers a bucket for every `requires_grad`
parameter at construction and then waits for gradients the graph never
produces, the reduction never finishes, and the *next* forward dies in
`_rebuild_buckets` with "Expected to have finished reduction in the
prior iteration". That is how the 2xT4 run died one step after the
unfreeze, reporting "did not receive grad: 1 413 414" β which is
exactly `embeddings.mask_token`, `norm.weight`, `norm.bias` in
`named_parameters()` order (5 embedding tensors, then 24 blocks x 17,
then the final norm). Same failure mode as `fuse[3].rcu1` (indices
60-65) one level down; see `models/dpt.py`.
Found by probing, not by name: one tiny forward/backward through this
very module, and whatever comes back with `grad is None` is by
definition unreachable. That survives a backbone swap or an HF
refactor, which a hard-coded name list does not β the list is only the
fallback for when the probe itself cannot run.
"""
dead = self._probe_unreachable()
if dead is None:
dead = self._named_unreachable()
by_name = dict(self.model.named_parameters())
got = []
for nm in dead:
p = by_name.get(nm)
if p is not None and p.requires_grad:
p.requires_grad_(False)
got.append(nm)
if got:
print(f"[model] {len(got)} unreachable encoder parameter(s) left "
f"frozen (never receive gradients, would hang DDP): "
f"{', '.join(got)}")
else:
print("[model] encoder probe: every trainable parameter reaches "
"the taps")
return got
def _probe_unreachable(self) -> list[str] | None:
"""Names of trainable backbone params that get no grad from `forward`.
`None` means the probe could not be run, not that nothing is dead.
The input is a deterministic ramp rather than `torch.randn`: it draws no
numbers from the global RNG, so a resumed run's RNG state stays exactly
where `_restore_rng` put it. Values do not matter β only whether an
autograd edge exists β but a *constant* input would make every
LayerNorm see zero variance, so the ramp avoids the degenerate case.
The probe runs in train mode so it exercises the same graph training
will, gradient checkpointing included.
"""
ref = next(self.model.parameters(), None)
if ref is None:
return []
s = max(self.patch * 4, self.patch)
n = 3 * s * s
x = (torch.arange(n, device=ref.device, dtype=torch.float32)
.reshape(1, 3, s, s) / n).to(ref.dtype)
was_training = self.model.training
stash = [(p, p.grad) for p in self.model.parameters()]
try:
for p, _ in stash:
p.grad = None
self.model.train()
with torch.enable_grad():
outs = self.forward(x)
total = outs[0].float().sum()
for o in outs[1:]:
total = total + o.float().sum()
total.backward()
return [nm for nm, p in self.model.named_parameters()
if p.requires_grad and p.grad is None]
except Exception as e: # noqa: BLE001
print(f"[model] unreachable-parameter probe failed ({e}) β "
"falling back to the known dead set")
return None
finally:
# The probe's gradients are not training signal; drop them and put
# back whatever was there (nothing, at every call site).
for p, g in stash:
p.grad = g
self.model.train(was_training)
def _named_unreachable(self) -> list[str]:
"""The dead set for this backbone family, by name.
Only consulted when `_probe_unreachable` raised. Each name is checked
against the live module, so a backbone without it is simply unaffected.
"""
out = []
emb = getattr(self.model, "embeddings", None)
if isinstance(getattr(emb, "mask_token", None), nn.Parameter):
out.append("embeddings.mask_token")
# The ViT's own trailing LayerNorm, which sits after the last tapped
# hidden state. Matched as a direct attribute of the backbone so the
# per-block `norm1` / `norm2` cannot be caught by accident.
final_norm = getattr(self.model, "norm", None)
if isinstance(final_norm, nn.Module):
out += [f"norm.{nm}" for nm, _ in final_norm.named_parameters()]
return out
def _set_grad_checkpointing(self, on: bool) -> None:
fn = ("gradient_checkpointing_enable" if on
else "gradient_checkpointing_disable")
try:
if on:
self.model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False})
else:
self.model.gradient_checkpointing_disable()
print(f"[model] encoder gradient checkpointing {'ON' if on else 'OFF'}")
except Exception as e: # noqa: BLE001
print(f"[model] {fn} unavailable: {e}")
def train(self, mode: bool = True):
super().train(mode)
if self.frozen:
self.model.eval()
return self
# -- layer-wise LR decay -------------------------------------------
def llrd_param_groups(self, base_lr: float, decay: float, weight_decay: float):
"""One param group per block, LR decaying towards the input.
Depth 0 = embeddings, depth L = the last block; group LR is
`base_lr * decay ** (L - depth)`. Norm/bias params get no weight decay.
"""
if self.blocks is None:
return [{"params": [p for p in self.model.parameters() if p.requires_grad],
"lr": base_lr, "weight_decay": weight_decay, "name": "encoder"}]
block_ids = {id(p) for b in self.blocks for p in b.parameters()}
depth_of: dict[int, int] = {}
for d, b in enumerate(self.blocks, start=1):
for p in b.parameters():
depth_of[id(p)] = d
L = len(self.blocks)
buckets: dict[tuple[int, bool], list] = {}
for name, p in self.model.named_parameters():
if not p.requires_grad:
continue
depth = depth_of.get(id(p), 0 if id(p) not in block_ids else L)
no_decay = p.ndim <= 1 or name.endswith(".bias")
buckets.setdefault((depth, no_decay), []).append(p)
groups = []
for (depth, no_decay), params in sorted(buckets.items()):
groups.append({
"params": params,
"lr": base_lr * (decay ** (L - depth)),
"weight_decay": 0.0 if no_decay else weight_decay,
"name": f"enc.d{depth}{'.nd' if no_decay else ''}",
})
return groups
# -- forward --------------------------------------------------------
def _tokens_to_map(self, tok: torch.Tensor, hp: int, wp: int) -> torch.Tensor:
"""(B, N, C) tokens -> (B, C, hp, wp) feature map.
`transpose(1, 2).reshape(...)` cannot be a view, so it materialised a
full copy of every tap. Reshaping to (B, hp, wp, C) *is* a view, and the
permute that follows leaves a tensor already laid out as channels_last β
which is the format the DPT convs want, so the copy and the cudnn
permute both disappear.
"""
# Drop CLS + register/prefix tokens: the patch tokens are always the last
# hp*wp entries regardless of how many prefix tokens the config uses.
patches = tok[:, -(hp * wp):, :]
m = patches.reshape(tok.shape[0], hp, wp, -1).permute(0, 3, 1, 2)
return m if self._cl else m.contiguous()
def _hidden_states(self, pixel_values: torch.Tensor) -> list[torch.Tensor]:
"""The tapped hidden states, robust to how the backbone is configured.
`output_hidden_states` is passed per call rather than trusted from the
config: setting it only via `from_pretrained` leaves `hidden_states=None`
on transformers 5.x, which turns the whole decoder into a TypeError at
the first forward.
"""
out = self.model(pixel_values=pixel_values, output_hidden_states=True)
hs = getattr(out, "hidden_states", None)
if hs is None and isinstance(out, (tuple, list)):
hs = out[-1]
if hs is None:
raise RuntimeError(
f"{type(self.model).__name__} returned no hidden_states; "
"the encoder taps cannot be built")
return [hs[i] for i in self.idx]
def forward(self, pixel_values: torch.Tensor) -> list[torch.Tensor]:
hp = pixel_values.shape[-2] // self.patch
wp = pixel_values.shape[-1] // self.patch
# No `.float()` here. Under autocast the taps come out in bf16 and the
# trunk's first op is a Conv2d that autocast casts back to bf16 anyway β
# so the upcast was a pure round-trip, ~540 MB of fp32 allocated and
# re-read per step at batch 32. Without autocast the taps are already
# fp32 and nothing changes. Gradient checkpointing is toggled once in
# `set_frozen`, not here: it used to be re-applied every single step.
if self.frozen:
self.model.eval()
with torch.no_grad():
outs = [self._tokens_to_map(h, hp, wp)
for h in self._hidden_states(pixel_values)]
return [o.detach() for o in outs]
return [self._tokens_to_map(h, hp, wp)
for h in self._hidden_states(pixel_values)]
|