Spaces:
Running on Zero
Running on Zero
Download models/encoder.py from akashch1512/SingleViewHeigthEstimation: direct link, hf CLI and curl.
- Browser
- Download file 17 kB
-
https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/models/encoder.py
- Command line
-
hf download hf://spaces/akashch1512/SingleViewHeigthEstimation/models/encoder.py
-
curl -L -o encoder.py https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/models/encoder.py
17 kB
| """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)] | |