Instructions to use nvidia/Cosmos-Embed1-336p with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Cosmos
How to use nvidia/Cosmos-Embed1-336p with Cosmos:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- NeMo
How to use nvidia/Cosmos-Embed1-336p with NeMo:
# tag did not correspond to a valid NeMo domain.
- Notebooks
- Google Colab
- Kaggle
Support Transformers 4 and 5 in Cosmos-Embed1-336p remote code
#2
by manuel-tulip - opened
- README.md +39 -16
- examples/test_transformers_compat.py +147 -0
- modeling_embed1.py +31 -13
- modeling_qformer.py +54 -3
- modeling_vit.py +122 -13
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 |
-
|
| 212 |
-
|
| 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 |
-
|
| 222 |
-
|
| 223 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
#
|
| 242 |
-
video_inputs = preprocess(videos=batch).to(
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 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
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 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 |
-
|
|
|
|
|
|
|
| 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.
|
| 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.
|
| 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]
|
| 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 =
|
| 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
|
| 593 |
logger.info("Positional interpolation from %dx%d to %dx%d" % (orig_size, orig_size, new_h, new_w))
|
| 594 |
-
|
| 595 |
-
|
| 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,
|