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