Cosmos
Safetensors
NeMo
cosmos-embed1
nvidia
custom_code

Support Transformers 4 and 5 in Cosmos-Embed1-336p remote code

#2
README.md CHANGED
@@ -206,21 +206,45 @@ One can optionally install Transformer Engine for faster inference:
206
  pip install --no-build-isolation transformer_engine[pytorch]
207
  ```
208
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
209
  ## Example inference
210
 
211
- A code snippet for video and text inference is shown below.
212
- For a step-by-step guide, please refer to the Juypter notebook [here](https://huggingface.co/nvidia/Cosmos-Embed1-336p/blob/main/examples/example.ipynb).
213
- ```
 
214
  import decord
215
  import numpy as np
216
  import torch
217
  from transformers import AutoProcessor, AutoModel
218
  import subprocess
219
- import io
220
 
221
- # load model and pre-processor
222
- model = AutoModel.from_pretrained("nvidia/Cosmos-Embed1-336p", trust_remote_code=True).to("cuda", dtype=torch.bfloat16)
223
- preprocess = AutoProcessor.from_pretrained("nvidia/Cosmos-Embed1-336p", trust_remote_code=True)
 
 
 
224
 
225
  # load mock data
226
  video_url = "https://upload.wikimedia.org/wikipedia/commons/3/3d/Branko_Paukovic%2C_javelin_throw.webm"
@@ -238,15 +262,14 @@ captions = [
238
  "a man throwing a javelin with both hands", # distractor
239
  ]
240
 
241
- # video and text processing
242
- video_inputs = preprocess(videos=batch).to("cuda", dtype=torch.bfloat16)
243
- video_out = model.get_video_embeddings(**video_inputs)
244
- text_inputs = preprocess(text=captions).to("cuda", dtype=torch.bfloat16)
245
- text_out = model.get_text_embeddings(**text_inputs)
246
-
247
- # ranking and argmax
248
- probs = (torch.softmax(model.logit_scale.exp() * video_out.visual_proj @ text_out.text_proj.T, dim=-1))[0]
249
- print(captions[probs.argmax()])
250
  ```
251
 
252
  # Training and Evaluation
 
206
  pip install --no-build-isolation transformer_engine[pytorch]
207
  ```
208
 
209
+ The example below also uses `decord` for video decoding and `wget` to fetch a
210
+ sample clip. Install `decord` separately and ensure `wget` is available.
211
+
212
+ ### Transformers compatibility
213
+
214
+ The remote model code has been validated with Transformers 4.57.6 and 5.17.0.
215
+ Transformers 4 loads positional embeddings through a state-dict pre-hook;
216
+ Transformers 5 uses a registered weight conversion operation instead. This
217
+ matters when loading the 336p and 448p checkpoints: they store positional
218
+ embeddings on a 224p patch grid, which must be interpolated to the model's
219
+ target resolution. Do not use `ignore_mismatched_sizes=True` to work around a
220
+ shape mismatch; that reinitializes the positional embeddings.
221
+
222
+ Community validation with the same checkpoints and fixed inputs found exact
223
+ text/video-vector parity with the original Transformers 4.57.6 code on CPU at
224
+ all three resolutions. Transformers 5.17.0 was also compared against that
225
+ baseline on CPU, Apple MPS, and NVIDIA CUDA. These comparisons establish
226
+ numerical parity for the tested inputs; they do not change the supported
227
+ platform statement above or establish a performance or retrieval-accuracy
228
+ guarantee.
229
+
230
  ## Example inference
231
 
232
+ This example keeps model weights in FP32 and runs on CUDA. For a step-by-step
233
+ guide, see the [example notebook](https://huggingface.co/nvidia/Cosmos-Embed1-336p/blob/main/examples/example.ipynb).
234
+
235
+ ```python
236
  import decord
237
  import numpy as np
238
  import torch
239
  from transformers import AutoProcessor, AutoModel
240
  import subprocess
 
241
 
242
+ model_id = "nvidia/Cosmos-Embed1-336p"
243
+ device = "cuda"
244
+ model = AutoModel.from_pretrained(
245
+ model_id, trust_remote_code=True, dtype=torch.float32,
246
+ ).to(device).eval()
247
+ preprocess = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)
248
 
249
  # load mock data
250
  video_url = "https://upload.wikimedia.org/wikipedia/commons/3/3d/Branko_Paukovic%2C_javelin_throw.webm"
 
262
  "a man throwing a javelin with both hands", # distractor
263
  ]
264
 
265
+ # Video and text embeddings are L2-normalized by the model.
266
+ video_inputs = preprocess(videos=batch).to(device)
267
+ text_inputs = preprocess(text=captions).to(device)
268
+ with torch.inference_mode():
269
+ video_out = model.get_video_embeddings(**video_inputs)
270
+ text_out = model.get_text_embeddings(**text_inputs)
271
+ scores = (video_out.visual_proj @ text_out.text_proj.T)[0]
272
+ print(captions[scores.argmax().item()])
 
273
  ```
274
 
275
  # Training and Evaluation
examples/test_transformers_compat.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+
16
+ """Weight-free compatibility tests for Transformers 4 and 5.
17
+
18
+ Load standalone remote-code modules without model weights. V5-specific
19
+ conversion tests are skipped when the conversion API is unavailable.
20
+ """
21
+
22
+ import importlib.util
23
+ import sys
24
+ from pathlib import Path
25
+
26
+ import pytest
27
+ import torch
28
+
29
+ ROOT = Path(__file__).resolve().parent.parent
30
+
31
+
32
+ def _load(name: str):
33
+ spec = importlib.util.spec_from_file_location(name, ROOT / f"{name}.py")
34
+ module = importlib.util.module_from_spec(spec)
35
+ sys.modules[name] = module
36
+ spec.loader.exec_module(module)
37
+ return module
38
+
39
+
40
+ modeling_vit = _load("modeling_vit")
41
+ modeling_qformer = _load("modeling_qformer")
42
+
43
+
44
+ def test_qformer_initialization_and_attention_masks() -> None:
45
+ from transformers import BertConfig
46
+
47
+ config = BertConfig(
48
+ hidden_size=16, num_hidden_layers=1, num_attention_heads=2,
49
+ intermediate_size=32, vocab_size=64,
50
+ )
51
+ model = modeling_qformer.BertModel(config).eval()
52
+ if modeling_vit._HAS_TRANSFORMERS5_CONVERSION_API:
53
+ assert hasattr(model, "all_tied_weights_keys")
54
+ assert model.get_head_mask(None, 1) == [None]
55
+ mask = model.invert_attention_mask(torch.tensor([[1, 0]]))
56
+ assert mask.shape == (1, 1, 1, 2)
57
+ assert mask[0, 0, 0, 0] == 0
58
+ assert mask[0, 0, 0, 1] < 0
59
+
60
+
61
+ def test_pruning_helper_is_vendored() -> None:
62
+ heads, index = modeling_qformer.find_pruneable_heads_and_indices([0, 1], 4, 2, set())
63
+ assert heads == {0, 1}
64
+ assert index.tolist() == [4, 5, 6, 7]
65
+
66
+
67
+ def test_bicubic_interpolation_matches_reference() -> None:
68
+ pos_embed = torch.arange(1 * 257 * 8, dtype=torch.float32).reshape(1, 257, 8)
69
+ out = modeling_vit._bicubic_interpolate_pos_embed(pos_embed, 1, 24, 24)
70
+ assert out.shape == (1, 577, 8)
71
+
72
+ # Independent reference: same math as the historical interpolate_pos_embed.
73
+ extra = pos_embed[:, :1]
74
+ tokens = pos_embed[:, 1:].reshape(-1, 16, 16, 8).permute(0, 3, 1, 2)
75
+ tokens = torch.nn.functional.interpolate(tokens, size=(24, 24), mode="bicubic", align_corners=False)
76
+ tokens = tokens.permute(0, 2, 3, 1).flatten(1, 2)
77
+ reference = torch.cat((extra, tokens), dim=1)
78
+ assert torch.equal(out, reference)
79
+ assert torch.equal(out[:, :1], pos_embed[:, :1]) # cls token untouched
80
+
81
+
82
+ def test_bicubic_interpolation_is_noop_on_matching_grid() -> None:
83
+ pos_embed = torch.rand(1, 577, 8)
84
+ out = modeling_vit._bicubic_interpolate_pos_embed(pos_embed, 1, 24, 24)
85
+ assert torch.equal(out, pos_embed)
86
+
87
+
88
+ def test_interpolation_rejects_non_square_source_and_resizes_equal_area_rectangle() -> None:
89
+ with pytest.raises(ValueError, match="square patch grid"):
90
+ modeling_vit._bicubic_interpolate_pos_embed(torch.rand(1, 258, 8), 1, 24, 24)
91
+ source = torch.arange(1 * 257 * 8, dtype=torch.float32).reshape(1, 257, 8)
92
+ resized = modeling_vit._bicubic_interpolate_pos_embed(source, 1, 8, 32)
93
+ assert resized.shape == source.shape
94
+ assert not torch.equal(resized[:, 1:], source[:, 1:])
95
+ state_dict = {"visual_encoder.pos_embed": source.clone()}
96
+ modeling_vit.interpolate_pos_embed(
97
+ "visual_encoder.pos_embed", 256, 257, state_dict, target_h=8, target_w=32
98
+ )
99
+ assert torch.equal(state_dict["visual_encoder.pos_embed"], resized)
100
+
101
+
102
+ def test_interpolate_pos_embed_writes_only_on_mismatch() -> None:
103
+ state_dict = {"visual_encoder.pos_embed": torch.rand(1, 257, 8)}
104
+ modeling_vit.interpolate_pos_embed(
105
+ "visual_encoder.pos_embed", 256, 257, state_dict, target_h=16, target_w=16
106
+ )
107
+ assert state_dict["visual_encoder.pos_embed"].shape == (1, 257, 8) # no-op
108
+
109
+
110
+ @pytest.mark.skipif(
111
+ not modeling_vit._HAS_TRANSFORMERS5_CONVERSION_API,
112
+ reason="Transformers 5 conversion API not available",
113
+ )
114
+ def test_v5_conversion_mapping_registered() -> None:
115
+ from transformers.conversion_mapping import get_checkpoint_conversion_mapping
116
+
117
+ modeling_vit.register_transformers5_pos_embed_conversion()
118
+ converters = get_checkpoint_conversion_mapping("cosmos-embed1")
119
+ assert converters and len(converters) == 1
120
+ assert converters[0].source_patterns == ["visual_encoder.pos_embed"]
121
+
122
+
123
+ @pytest.mark.skipif(
124
+ not modeling_vit._HAS_TRANSFORMERS5_CONVERSION_API,
125
+ reason="Transformers 5 conversion API not available",
126
+ )
127
+ def test_v5_operation_interpolates_257_to_577() -> None:
128
+ class _PatchEmbed:
129
+ patch_shape = (24, 24)
130
+
131
+ class _Visual:
132
+ patch_embed = _PatchEmbed()
133
+
134
+ class _Model:
135
+ visual_encoder = _Visual()
136
+
137
+ op = modeling_vit.InterpolatePosEmbed()
138
+ tensor = torch.arange(1 * 257 * 4, dtype=torch.bfloat16).reshape(1, 257, 4)
139
+ out = op.convert(
140
+ {"visual_encoder.pos_embed": [tensor]},
141
+ ["visual_encoder.pos_embed"],
142
+ ["visual_encoder.pos_embed"],
143
+ model=_Model(),
144
+ )
145
+ result = out["visual_encoder.pos_embed"][0]
146
+ assert result.shape == (1, 577, 4)
147
+ assert result.dtype == torch.bfloat16
modeling_embed1.py CHANGED
@@ -28,7 +28,18 @@ from .configuration_embed1 import CosmosEmbed1Config
28
  from .modeling_outputs import TextEmbedderOutput, TextVideoEmbedderOutput, VideoEmbedderOutput
29
  from .modeling_qformer import BertLMHeadModel, load_qformer
30
  from .modeling_utils import EncodingFactory, rank0_first
31
- from .modeling_vit import EvaViTG
 
 
 
 
 
 
 
 
 
 
 
32
 
33
 
34
  class CosmosEmbed1(PreTrainedModel):
@@ -51,17 +62,15 @@ class CosmosEmbed1(PreTrainedModel):
51
  self.transformer_engine = config.transformer_engine
52
  self.use_fp8 = config.use_fp8
53
 
54
- # visual encoder initialization
55
- self.register_buffer(
56
- "normalization_mean",
57
- torch.tensor([0.485, 0.456, 0.406]).view(1, 1, 3, 1, 1),
58
- persistent=False,
59
- )
60
- self.register_buffer(
61
- "normalization_std",
62
- torch.tensor([0.229, 0.224, 0.225]).view(1, 1, 3, 1, 1),
63
- persistent=False,
64
- )
65
  self.visual_encoder = EvaViTG(
66
  img_size=self.resolution,
67
  transformer_engine=self.transformer_engine,
@@ -97,6 +106,13 @@ class CosmosEmbed1(PreTrainedModel):
97
  self.logit_scale = nn.Parameter(torch.tensor(math.log(10.0)))
98
  self.logit_bias = nn.Parameter(torch.tensor(-10.0))
99
 
 
 
 
 
 
 
 
100
  @property
101
  def hidden_dim(self) -> int:
102
  return self.visual_encoder.embed_dim
@@ -131,7 +147,9 @@ class CosmosEmbed1(PreTrainedModel):
131
  return TextVideoEmbedderOutput(**video_output, **text_output)
132
 
133
  def get_video_embeddings(self, videos: torch.Tensor) -> VideoEmbedderOutput:
134
- videos = (videos - self.normalization_mean) / self.normalization_std
 
 
135
  batch_size, num_frames, _, H, W = videos.shape
136
  frame_batch = rearrange(videos, "b t c h w -> (b t) c h w")
137
 
 
28
  from .modeling_outputs import TextEmbedderOutput, TextVideoEmbedderOutput, VideoEmbedderOutput
29
  from .modeling_qformer import BertLMHeadModel, load_qformer
30
  from .modeling_utils import EncodingFactory, rank0_first
31
+ from .modeling_vit import EvaViTG, register_transformers5_pos_embed_conversion
32
+
33
+ # Register the Transformers>=5 pos_embed interpolation converter. This is a
34
+ # no-op under Transformers<5, where the EvaViTG load-state-dict pre-hook does
35
+ # the interpolation. Transformers 5 bypasses that pre-hook entirely.
36
+ register_transformers5_pos_embed_conversion()
37
+
38
+ # ImageNet normalization constants, created at import time so they exist on a
39
+ # real device even when the model is later constructed under Transformers 5's
40
+ # "init empty weights" (meta device) context.
41
+ _NORMALIZATION_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 1, 3, 1, 1)
42
+ _NORMALIZATION_STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 1, 3, 1, 1)
43
 
44
 
45
  class CosmosEmbed1(PreTrainedModel):
 
62
  self.transformer_engine = config.transformer_engine
63
  self.use_fp8 = config.use_fp8
64
 
65
+ # visual encoder normalization constants. These reference module-level
66
+ # CPU tensors created at import time, not tensors allocated here:
67
+ # Transformers >= 5 constructs the model under an "init empty weights"
68
+ # context, so any tensor allocated in __init__ lands on the meta device
69
+ # with no data. Plain attributes are also not touched by the loader
70
+ # (unlike non-persistent buffers, which v5 replaces with uninitialized
71
+ # storage); get_video_embeddings moves them to the input per call.
72
+ self.normalization_mean = _NORMALIZATION_MEAN
73
+ self.normalization_std = _NORMALIZATION_STD
 
 
74
  self.visual_encoder = EvaViTG(
75
  img_size=self.resolution,
76
  transformer_engine=self.transformer_engine,
 
106
  self.logit_scale = nn.Parameter(torch.tensor(math.log(10.0)))
107
  self.logit_bias = nn.Parameter(torch.tensor(-10.0))
108
 
109
+ # Transformers >= 5 requires models to run post_init explicitly; it
110
+ # populates tied-weight/parallelism bookkeeping (e.g.
111
+ # ``all_tied_weights_keys``) that the dynamic loader relies on. It is
112
+ # also the correct call under Transformers < 5 and replaces the
113
+ # historical ``init_weights()`` usage in the Q-Former classes.
114
+ self.post_init()
115
+
116
  @property
117
  def hidden_dim(self) -> int:
118
  return self.visual_encoder.embed_dim
 
147
  return TextVideoEmbedderOutput(**video_output, **text_output)
148
 
149
  def get_video_embeddings(self, videos: torch.Tensor) -> VideoEmbedderOutput:
150
+ mean = self.normalization_mean.to(device=videos.device, dtype=videos.dtype)
151
+ std = self.normalization_std.to(device=videos.device, dtype=videos.dtype)
152
+ videos = (videos - mean) / std
153
  batch_size, num_frames, _, H, W = videos.shape
154
  frame_batch = rearrange(videos, "b t c h w -> (b t) c h w")
155
 
modeling_qformer.py CHANGED
@@ -44,7 +44,6 @@ from transformers.modeling_utils import (
44
  )
45
  from transformers.pytorch_utils import (
46
  apply_chunking_to_forward,
47
- find_pruneable_heads_and_indices,
48
  prune_linear_layer,
49
  )
50
  from transformers.models.bert.configuration_bert import BertConfig
@@ -52,6 +51,24 @@ from transformers.models.bert.configuration_bert import BertConfig
52
  logger = getLogger(__file__)
53
 
54
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
55
  class BertEmbeddings(nn.Module):
56
  """Construct the embeddings from word and position embeddings."""
57
 
@@ -615,6 +632,40 @@ class BertPreTrainedModel(PreTrainedModel):
615
  if isinstance(module, nn.Linear) and module.bias is not None:
616
  module.bias.data.zero_()
617
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
618
 
619
  class BertModel(BertPreTrainedModel):
620
  """
@@ -636,7 +687,7 @@ class BertModel(BertPreTrainedModel):
636
 
637
  self.pooler = BertPooler(config) if add_pooling_layer else None
638
 
639
- self.init_weights()
640
 
641
  def get_input_embeddings(self):
642
  return self.embeddings.word_embeddings
@@ -886,7 +937,7 @@ class BertLMHeadModel(BertPreTrainedModel, GenerationMixin):
886
  self.bert = BertModel(config, add_pooling_layer=False)
887
  self.cls = BertOnlyMLMHead(config)
888
 
889
- self.init_weights()
890
 
891
  def get_output_embeddings(self):
892
  return self.cls.predictions.decoder
 
44
  )
45
  from transformers.pytorch_utils import (
46
  apply_chunking_to_forward,
 
47
  prune_linear_layer,
48
  )
49
  from transformers.models.bert.configuration_bert import BertConfig
 
51
  logger = getLogger(__file__)
52
 
53
 
54
+ def find_pruneable_heads_and_indices(heads, n_heads, head_size, already_pruned_heads):
55
+ """Vendored copy of the Transformers<5 helper removed in Transformers 5.
56
+
57
+ ``BertAttention.prune_heads`` calls this on the pruning path only (never
58
+ during inference), but the module imported it unconditionally, which broke
59
+ the import under Transformers 5. Defining it locally keeps the remote code
60
+ importable; the historical implementation is preserved verbatim.
61
+ """
62
+ mask = torch.ones(n_heads, head_size)
63
+ heads = set(heads) - already_pruned_heads
64
+ for head in heads:
65
+ head = head - sum(1 if h < head else 0 for h in already_pruned_heads)
66
+ mask[head] = 0
67
+ mask = mask.view(-1).contiguous().eq(1)
68
+ index = torch.arange(len(mask))[mask].long()
69
+ return heads, index
70
+
71
+
72
  class BertEmbeddings(nn.Module):
73
  """Construct the embeddings from word and position embeddings."""
74
 
 
632
  if isinstance(module, nn.Linear) and module.bias is not None:
633
  module.bias.data.zero_()
634
 
635
+ def invert_attention_mask(self, encoder_attention_mask: Tensor) -> Tensor:
636
+ """Vendored from Transformers<5 ``ModuleUtilsMixin`` (removed in v5)."""
637
+ if encoder_attention_mask.dim() == 3:
638
+ encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :]
639
+ if encoder_attention_mask.dim() == 2:
640
+ encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :]
641
+ # T5 has a mask that can compare sequence ids, we can simulate this here with this transposition
642
+ # Cf. https://github.com/tensorflow/mesh/blob/8d2465e9bc93129b913b5ccc6a59aa97abd96ec6/mesh_tensorflow
643
+ # /transformer/transformer_layers.py#L270
644
+ encoder_extended_attention_mask = encoder_extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility
645
+ encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * torch.finfo(self.dtype).min
646
+ return encoder_extended_attention_mask
647
+
648
+ def get_head_mask(self, head_mask, num_hidden_layers: int, is_attention_chunked: bool = False):
649
+ """Vendored from Transformers<5 ``ModuleUtilsMixin`` (removed in v5)."""
650
+ if head_mask is not None:
651
+ head_mask = self._convert_head_mask_to_5d(head_mask, num_hidden_layers)
652
+ if is_attention_chunked is True:
653
+ head_mask = head_mask.unsqueeze(-1)
654
+ else:
655
+ head_mask = [None] * num_hidden_layers
656
+ return head_mask
657
+
658
+ def _convert_head_mask_to_5d(self, head_mask, num_hidden_layers):
659
+ """Vendored from Transformers<5 ``ModuleUtilsMixin`` (removed in v5)."""
660
+ if head_mask.dim() == 1:
661
+ head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
662
+ head_mask = head_mask.expand(num_hidden_layers, -1, -1, -1, -1)
663
+ elif head_mask.dim() == 2:
664
+ head_mask = head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1)
665
+ assert head_mask.dim() == 5, f"head_mask.dim != 5, instead {head_mask.dim()}"
666
+ head_mask = head_mask.to(dtype=self.dtype)
667
+ return head_mask
668
+
669
 
670
  class BertModel(BertPreTrainedModel):
671
  """
 
687
 
688
  self.pooler = BertPooler(config) if add_pooling_layer else None
689
 
690
+ self.post_init()
691
 
692
  def get_input_embeddings(self):
693
  return self.embeddings.word_embeddings
 
937
  self.bert = BertModel(config, add_pooling_layer=False)
938
  self.cls = BertOnlyMLMHead(config)
939
 
940
+ self.post_init()
941
 
942
  def get_output_embeddings(self):
943
  return self.cls.predictions.decoder
modeling_vit.py CHANGED
@@ -48,6 +48,23 @@ try:
48
  except ImportError:
49
  pass
50
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
51
 
52
  def drop_path(x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True):
53
  """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
@@ -565,6 +582,32 @@ class VisionTransformer(nn.Module):
565
  return len(self.blocks)
566
 
567
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
568
  def interpolate_pos_embed(
569
  pos_embed_key: str,
570
  num_patches: int,
@@ -574,11 +617,10 @@ def interpolate_pos_embed(
574
  target_w: int = None,
575
  ) -> None:
576
  if pos_embed_key in checkpoint_model:
577
- pos_embed_checkpoint = checkpoint_model[pos_embed_key].float()
578
- embedding_size = pos_embed_checkpoint.shape[-1]
579
  num_extra_tokens = patch_embed_shape - num_patches
580
  # height (== width) for the checkpoint position embedding
581
- orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
582
 
583
  # If target dimensions are provided, use them; otherwise assume square
584
  if target_h is not None and target_w is not None:
@@ -589,18 +631,11 @@ def interpolate_pos_embed(
589
  new_h, new_w = new_size, new_size
590
 
591
  # class_token and dist_token are kept unchanged
592
- if orig_size * orig_size != new_h * new_w:
593
  logger.info("Positional interpolation from %dx%d to %dx%d" % (orig_size, orig_size, new_h, new_w))
594
- extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
595
- # only the position tokens are interpolated
596
- pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
597
- pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
598
- pos_tokens = torch.nn.functional.interpolate(
599
- pos_tokens, size=(new_h, new_w), mode="bicubic", align_corners=False
600
  )
601
- pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
602
- new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
603
- checkpoint_model[pos_embed_key] = new_pos_embed
604
 
605
 
606
  class PositionalEmbeddingHook:
@@ -619,6 +654,80 @@ class PositionalEmbeddingHook:
619
  )
620
 
621
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
622
  class EvaViTG(VisionTransformer):
623
  def __init__(
624
  self,
 
48
  except ImportError:
49
  pass
50
 
51
+ # Transformers >= 5 uses a new dynamic weight loader that no longer calls
52
+ # ``nn.Module.load_state_dict``. The load-state-dict pre-hook registered by
53
+ # ``EvaViTG`` therefore never runs and ``pos_embed`` is not interpolated. When
54
+ # the new conversion API is available we register an equivalent
55
+ # ``WeightConverter`` operation; under Transformers < 5 the imports fail and the
56
+ # pre-hook keeps doing the work.
57
+ try: # pragma: no cover - branch is selected by the installed version
58
+ from transformers.conversion_mapping import register_checkpoint_conversion_mapping
59
+ from transformers.core_model_loading import ConversionOps, WeightConverter
60
+
61
+ _HAS_TRANSFORMERS5_CONVERSION_API = True
62
+ except ImportError:
63
+ register_checkpoint_conversion_mapping = None
64
+ ConversionOps = None
65
+ WeightConverter = None
66
+ _HAS_TRANSFORMERS5_CONVERSION_API = False
67
+
68
 
69
  def drop_path(x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True):
70
  """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
 
582
  return len(self.blocks)
583
 
584
 
585
+ def _bicubic_interpolate_pos_embed(pos_embed, num_extra_tokens, target_h, target_w):
586
+ """Bicubically resize position tokens to ``(target_h, target_w)``.
587
+
588
+ Extra tokens (e.g. the cls token) are preserved unchanged. Shared by the
589
+ Transformers<5 pre-hook and the Transformers>=5 conversion operation so
590
+ both loading paths are numerically identical.
591
+ """
592
+ if pos_embed.ndim != 3 or num_extra_tokens < 0 or target_h <= 0 or target_w <= 0:
593
+ raise ValueError("pos_embed requires (batch, tokens, channels) and a positive target grid")
594
+ source_tokens = pos_embed.shape[1] - num_extra_tokens
595
+ orig_size = math.isqrt(max(source_tokens, 0))
596
+ if source_tokens <= 0 or orig_size * orig_size != source_tokens:
597
+ raise ValueError("pos_embed source tokens must form a square patch grid")
598
+ x = pos_embed.float()
599
+ if orig_size == target_h == target_w:
600
+ return x
601
+ extra_tokens = x[:, :num_extra_tokens]
602
+ pos_tokens = x[:, num_extra_tokens:]
603
+ pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, x.shape[-1]).permute(0, 3, 1, 2)
604
+ pos_tokens = torch.nn.functional.interpolate(
605
+ pos_tokens, size=(target_h, target_w), mode="bicubic", align_corners=False
606
+ )
607
+ pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
608
+ return torch.cat((extra_tokens, pos_tokens), dim=1)
609
+
610
+
611
  def interpolate_pos_embed(
612
  pos_embed_key: str,
613
  num_patches: int,
 
617
  target_w: int = None,
618
  ) -> None:
619
  if pos_embed_key in checkpoint_model:
620
+ pos_embed_checkpoint = checkpoint_model[pos_embed_key]
 
621
  num_extra_tokens = patch_embed_shape - num_patches
622
  # height (== width) for the checkpoint position embedding
623
+ orig_size = math.isqrt(max(pos_embed_checkpoint.shape[-2] - num_extra_tokens, 0))
624
 
625
  # If target dimensions are provided, use them; otherwise assume square
626
  if target_h is not None and target_w is not None:
 
631
  new_h, new_w = new_size, new_size
632
 
633
  # class_token and dist_token are kept unchanged
634
+ if orig_size != new_h or orig_size != new_w:
635
  logger.info("Positional interpolation from %dx%d to %dx%d" % (orig_size, orig_size, new_h, new_w))
636
+ checkpoint_model[pos_embed_key] = _bicubic_interpolate_pos_embed(
637
+ pos_embed_checkpoint, num_extra_tokens, new_h, new_w
 
 
 
 
638
  )
 
 
 
639
 
640
 
641
  class PositionalEmbeddingHook:
 
654
  )
655
 
656
 
657
+ if _HAS_TRANSFORMERS5_CONVERSION_API:
658
+
659
+ class _PosEmbedInterpolationNoOp(ConversionOps):
660
+ """Reverse-op placeholder: bicubic interpolation is not inverted."""
661
+
662
+ def convert(self, input_dict, source_patterns=None, target_patterns=None, **kwargs):
663
+ return input_dict
664
+
665
+ @property
666
+ def reverse_op(self):
667
+ return self
668
+
669
+ class InterpolatePosEmbed(ConversionOps):
670
+ """Bicubic ``pos_embed`` interpolation for the Transformers>=5 loader.
671
+
672
+ Transformers 5 assigns weights with ``setattr`` and never calls
673
+ ``nn.Module.load_state_dict``, so ``PositionalEmbeddingHook`` is dead
674
+ code there. The loader passes the instantiated model to ``convert``,
675
+ from which the target patch grid is read. The math is shared with the
676
+ pre-hook path via :func:`_bicubic_interpolate_pos_embed`.
677
+ """
678
+
679
+ def __init__(self, num_extra_tokens: int = 1):
680
+ self.num_extra_tokens = num_extra_tokens
681
+
682
+ @torch.no_grad()
683
+ def convert(self, input_dict, source_patterns, target_patterns, **kwargs):
684
+ tensor = next(iter(input_dict.values()))[0]
685
+ patch_h, patch_w = _model_patch_shape(kwargs.get("model"))
686
+ interpolated = _bicubic_interpolate_pos_embed(
687
+ tensor, self.num_extra_tokens, int(patch_h), int(patch_w)
688
+ )
689
+ return {target_patterns[0]: [interpolated.to(tensor.dtype)]}
690
+
691
+ @property
692
+ def reverse_op(self):
693
+ return _PosEmbedInterpolationNoOp()
694
+
695
+ def _model_patch_shape(model):
696
+ visual = getattr(model, "visual_encoder", None)
697
+ patch_embed = getattr(visual, "patch_embed", None)
698
+ shape = getattr(patch_embed, "patch_shape", None)
699
+ if shape is None:
700
+ raise RuntimeError(
701
+ "InterpolatePosEmbed: model has no visual_encoder.patch_embed.patch_shape"
702
+ )
703
+ return shape[0], shape[1]
704
+
705
+ def register_transformers5_pos_embed_conversion() -> None:
706
+ """Register the pos_embed converter for ``model_type='cosmos-embed1'``.
707
+
708
+ ``get_model_conversion_mapping`` skips custom-code submodules unless
709
+ their class name or model_type is user-registered, so this call is what
710
+ makes the converter apply to the remote model. Safe to call repeatedly.
711
+ """
712
+ register_checkpoint_conversion_mapping(
713
+ "cosmos-embed1",
714
+ [
715
+ WeightConverter(
716
+ source_patterns=["visual_encoder.pos_embed"],
717
+ target_patterns=["visual_encoder.pos_embed"],
718
+ operations=[InterpolatePosEmbed()],
719
+ )
720
+ ],
721
+ overwrite=True,
722
+ )
723
+
724
+ else:
725
+
726
+ def register_transformers5_pos_embed_conversion() -> None:
727
+ """No-op under Transformers<5; the load-state-dict pre-hook handles it."""
728
+ return None
729
+
730
+
731
  class EvaViTG(VisionTransformer):
732
  def __init__(
733
  self,