Octopus1 commited on
Commit
4fa7707
·
verified ·
1 Parent(s): 96974a1

Add load-time DINOv3 key remap hook (compat transformers 4.56 & 5.6.x)

Browse files
Files changed (1) hide show
  1. modeling_page.py +38 -0
modeling_page.py CHANGED
@@ -695,6 +695,44 @@ class PaGEBackbone(nn.Module):
695
  self.embed_dim = int(dinov3_config.hidden_size)
696
  # CLS(1) + num_register_tokens
697
  self._num_front = 1 + int(getattr(dinov3_config, "num_register_tokens", 0) or 0)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
698
 
699
  def _get_patch_tokens(self, x: torch.Tensor) -> torch.Tensor:
700
  out = self.model(pixel_values=x, return_dict=True)
 
695
  self.embed_dim = int(dinov3_config.hidden_size)
696
  # CLS(1) + num_register_tokens
697
  self._num_front = 1 + int(getattr(dinov3_config, "num_register_tokens", 0) or 0)
698
+ # DINOv3ViTModel's internal naming differs across transformers versions:
699
+ # 4.56.x -> layer stack flattened: model.layer.N.* (model.<dinov3 top-level>)
700
+ # 5.6.x -> layer stack nested: model.model.layer.N.* (extra `.model` wrapper)
701
+ # A single safetensors file must load under both, so remap incoming keys at load time.
702
+ self._register_load_state_dict_pre_hook(self._remap_dinov3_keys)
703
+
704
+ @staticmethod
705
+ def _dinov3_has_nested_layer(dinov3_module) -> bool:
706
+ """True if this transformers version nests the layer stack under an inner `.model`."""
707
+ inner = getattr(dinov3_module, "model", None)
708
+ if not isinstance(inner, nn.Module):
709
+ return False
710
+ return any(k.startswith("model.layer.") for k in inner.state_dict().keys())
711
+
712
+ def _remap_dinov3_keys(self, state_dict, prefix, *args, **kwargs):
713
+ """Normalize DINOv3 backbone keys (embeddings / layer / norm / rope_embeddings)
714
+ from whichever convention the checkpoint uses into the one this transformers
715
+ version expects."""
716
+ nested = self._dinov3_has_nested_layer(self.model)
717
+ model_pref = prefix + "model."
718
+ new = {}
719
+ for k in list(state_dict.keys()):
720
+ if not k.startswith(model_pref):
721
+ continue
722
+ rest = k[len(model_pref):] # after "<prefix>model."
723
+ # rest is one of: "embeddings...", "norm...", "rope_embeddings...", "layer...",
724
+ # or "model.layer..." (nested-conv checkpoint under a flat version, etc.)
725
+ if rest.startswith("model.layer."):
726
+ core = rest[len("model."):] # -> "layer..."
727
+ else:
728
+ core = rest # "layer..." / "embeddings..." / "norm..." / "rope_embeddings..."
729
+ if nested and core.startswith("layer."):
730
+ target = model_pref + "model." + core
731
+ else:
732
+ target = model_pref + core
733
+ if target != k:
734
+ new[target] = state_dict.pop(k)
735
+ state_dict.update(new)
736
 
737
  def _get_patch_tokens(self, x: torch.Tensor) -> torch.Tensor:
738
  out = self.model(pixel_values=x, return_dict=True)