taewhan commited on
Commit
cfeb2ba
·
verified ·
1 Parent(s): 9117292

Initial release: MOTIF vision encoder 7B (bf16 safetensors + modeling_motif.py)

Browse files
LICENSE ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ Apache License
3
+ Version 2.0, January 2004
4
+ http://www.apache.org/licenses/
5
+
6
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
7
+
8
+ 1. Definitions.
9
+
10
+ "License" shall mean the terms and conditions for use, reproduction,
11
+ and distribution as defined by Sections 1 through 9 of this document.
12
+
13
+ "Licensor" shall mean the copyright owner or entity authorized by
14
+ the copyright owner that is granting the License.
15
+
16
+ "Legal Entity" shall mean the union of the acting entity and all
17
+ other entities that control, are controlled by, or are under common
18
+ control with that entity. For the purposes of this definition,
19
+ "control" means (i) the power, direct or indirect, to cause the
20
+ direction or management of such entity, whether by contract or
21
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
22
+ outstanding shares, or (iii) beneficial ownership of such entity.
23
+
24
+ "You" (or "Your") shall mean an individual or Legal Entity
25
+ exercising permissions granted by this License.
26
+
27
+ "Source" form shall mean the preferred form for making modifications,
28
+ including but not limited to software source code, documentation
29
+ source, and configuration files.
30
+
31
+ "Object" form shall mean any form resulting from mechanical
32
+ transformation or translation of a Source form, including but
33
+ not limited to compiled object code, generated documentation,
34
+ and conversions to other media types.
35
+
36
+ "Work" shall mean the work of authorship, whether in Source or
37
+ Object form, made available under the License, as indicated by a
38
+ copyright notice that is included in or attached to the work
39
+ (an example is provided in the Appendix below).
40
+
41
+ "Derivative Works" shall mean any work, whether in Source or Object
42
+ form, that is based on (or derived from) the Work and for which the
43
+ editorial revisions, annotations, elaborations, or other modifications
44
+ represent, as a whole, an original work of authorship. For the purposes
45
+ of this License, Derivative Works shall not include works that remain
46
+ separable from, or merely link (or bind by name) to the interfaces of,
47
+ the Work and Derivative Works thereof.
48
+
49
+ "Contribution" shall mean any work of authorship, including
50
+ the original version of the Work and any modifications or additions
51
+ to that Work or Derivative Works thereof, that is intentionally
52
+ submitted to Licensor for inclusion in the Work by the copyright owner
53
+ or by an individual or Legal Entity authorized to submit on behalf of
54
+ the copyright owner. For the purposes of this definition, "submitted"
55
+ means any form of electronic, verbal, or written communication sent
56
+ to the Licensor or its representatives, including but not limited to
57
+ communication on electronic mailing lists, source code control systems,
58
+ and issue tracking systems that are managed by, or on behalf of, the
59
+ Licensor for the purpose of discussing and improving the Work, but
60
+ excluding communication that is conspicuously marked or otherwise
61
+ designated in writing by the copyright owner as "Not a Contribution."
62
+
63
+ "Contributor" shall mean Licensor and any individual or Legal Entity
64
+ on behalf of whom a Contribution has been received by Licensor and
65
+ subsequently incorporated within the Work.
66
+
67
+ 2. Grant of Copyright License. Subject to the terms and conditions of
68
+ this License, each Contributor hereby grants to You a perpetual,
69
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
70
+ copyright license to reproduce, prepare Derivative Works of,
71
+ publicly display, publicly perform, sublicense, and distribute the
72
+ Work and such Derivative Works in Source or Object form.
73
+
74
+ 3. Grant of Patent License. Subject to the terms and conditions of
75
+ this License, each Contributor hereby grants to You a perpetual,
76
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
77
+ (except as stated in this section) patent license to make, have made,
78
+ use, offer to sell, sell, import, and otherwise transfer the Work,
79
+ where such license applies only to those patent claims licensable
80
+ by such Contributor that are necessarily infringed by their
81
+ Contribution(s) alone or by combination of their Contribution(s)
82
+ with the Work to which such Contribution(s) was submitted. If You
83
+ institute patent litigation against any entity (including a
84
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
85
+ or a Contribution incorporated within the Work constitutes direct
86
+ or contributory patent infringement, then any patent licenses
87
+ granted to You under this License for that Work shall terminate
88
+ as of the date such litigation is filed.
89
+
90
+ 4. Redistribution. You may reproduce and distribute copies of the
91
+ Work or Derivative Works thereof in any medium, with or without
92
+ modifications, and in Source or Object form, provided that You
93
+ meet the following conditions:
94
+
95
+ (a) You must give any other recipients of the Work or
96
+ Derivative Works a copy of this License; and
97
+
98
+ (b) You must cause any modified files to carry prominent notices
99
+ stating that You changed the files; and
100
+
101
+ (c) You must retain, in the Source form of any Derivative Works
102
+ that You distribute, all copyright, patent, trademark, and
103
+ attribution notices from the Source form of the Work,
104
+ excluding those notices that do not pertain to any part of
105
+ the Derivative Works; and
106
+
107
+ (d) If the Work includes a "NOTICE" text file as part of its
108
+ distribution, then any Derivative Works that You distribute must
109
+ include a readable copy of the attribution notices contained
110
+ within such NOTICE file, excluding those notices that do not
111
+ pertain to any part of the Derivative Works, in at least one
112
+ of the following places: within a NOTICE text file distributed
113
+ as part of the Derivative Works; within the Source form or
114
+ documentation, if provided along with the Derivative Works; or,
115
+ within a display generated by the Derivative Works, if and
116
+ wherever such third-party notices normally appear. The contents
117
+ of the NOTICE file are for informational purposes only and
118
+ do not modify the License. You may add Your own attribution
119
+ notices within Derivative Works that You distribute, alongside
120
+ or as an addendum to the NOTICE text from the Work, provided
121
+ that such additional attribution notices cannot be construed
122
+ as modifying the License.
123
+
124
+ You may add Your own copyright statement to Your modifications and
125
+ may provide additional or different license terms and conditions
126
+ for use, reproduction, or distribution of Your modifications, or
127
+ for any such Derivative Works as a whole, provided Your use,
128
+ reproduction, and distribution of the Work otherwise complies with
129
+ the conditions stated in this License.
130
+
131
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
132
+ any Contribution intentionally submitted for inclusion in the Work
133
+ by You to the Licensor shall be under the terms and conditions of
134
+ this License, without any additional terms or conditions.
135
+ Notwithstanding the above, nothing herein shall supersede or modify
136
+ the terms of any separate license agreement you may have executed
137
+ with Licensor regarding such Contributions.
138
+
139
+ 6. Trademarks. This License does not grant permission to use the trade
140
+ names, trademarks, service marks, or product names of the Licensor,
141
+ except as required for reasonable and customary use in describing the
142
+ origin of the Work and reproducing the content of the NOTICE file.
143
+
144
+ 7. Disclaimer of Warranty. Unless required by applicable law or
145
+ agreed to in writing, Licensor provides the Work (and each
146
+ Contributor provides its Contributions) on an "AS IS" BASIS,
147
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
148
+ implied, including, without limitation, any warranties or conditions
149
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
150
+ PARTICULAR PURPOSE. You are solely responsible for determining the
151
+ appropriateness of using or redistributing the Work and assume any
152
+ risks associated with Your exercise of permissions under this License.
153
+
154
+ 8. Limitation of Liability. In no event and under no legal theory,
155
+ whether in tort (including negligence), contract, or otherwise,
156
+ unless required by applicable law (such as deliberate and grossly
157
+ negligent acts) or agreed to in writing, shall any Contributor be
158
+ liable to You for damages, including any direct, indirect, special,
159
+ incidental, or consequential damages of any character arising as a
160
+ result of this License or out of the use or inability to use the
161
+ Work (including but not limited to damages for loss of goodwill,
162
+ work stoppage, computer failure or malfunction, or any and all
163
+ other commercial damages or losses), even if such Contributor
164
+ has been advised of the possibility of such damages.
165
+
166
+ 9. Accepting Warranty or Additional Liability. While redistributing
167
+ the Work or Derivative Works thereof, You may choose to offer,
168
+ and charge a fee for, acceptance of support, warranty, indemnity,
169
+ or other liability obligations and/or rights consistent with this
170
+ License. However, in accepting such obligations, You may act only
171
+ on Your own behalf and on Your sole responsibility, not on behalf
172
+ of any other Contributor, and only if You agree to indemnify,
173
+ defend, and hold each Contributor harmless for any liability
174
+ incurred by, or claims asserted against, such Contributor by reason
175
+ of your accepting any such warranty or additional liability.
176
+
177
+ END OF TERMS AND CONDITIONS
178
+
179
+ APPENDIX: How to apply the Apache License to your work.
180
+
181
+ To apply the Apache License to your work, attach the following
182
+ boilerplate notice, with the fields enclosed by brackets "[]"
183
+ replaced with your own identifying information. (Don't include
184
+ the brackets!) The text should be enclosed in the appropriate
185
+ comment syntax for the file format. We also recommend that a
186
+ file or class name and description of purpose be included on the
187
+ same "printed page" as the copyright notice for easier
188
+ identification within third-party archives.
189
+
190
+ Copyright [yyyy] [name of copyright owner]
191
+
192
+ Licensed under the Apache License, Version 2.0 (the "License");
193
+ you may not use this file except in compliance with the License.
194
+ You may obtain a copy of the License at
195
+
196
+ http://www.apache.org/licenses/LICENSE-2.0
197
+
198
+ Unless required by applicable law or agreed to in writing, software
199
+ distributed under the License is distributed on an "AS IS" BASIS,
200
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
201
+ See the License for the specific language governing permissions and
202
+ limitations under the License.
README.md ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: transformers
4
+ pipeline_tag: image-feature-extraction
5
+ tags:
6
+ - motif
7
+ - vision-transformer
8
+ - self-supervised
9
+ - image-feature-extraction
10
+ - video
11
+ - custom_code
12
+ ---
13
+
14
+ # MOTIF Vision Encoder
15
+
16
+ MOTIF Vision Encoder is a **unified image + video** self-supervised vision encoder on a ViT
17
+ backbone. A single **3D-convolutional tokenizer** ingests both modalities — an image is a
18
+ 1-frame clip (`T=1`), a video is `T>1` — so the same weights produce **dense patch-level
19
+ features** and a **language-aligned global (CLS) representation**.
20
+
21
+ - **Architecture**: ViT-7B (embed 4096 / depth 40 / heads 32), patch 16, 3D axial RoPE
22
+ (`base=100`), SwiGLU FFN, LayerScale, per-head QK-norm, gated attention, 4 register tokens.
23
+ - **Tokenizer**: `Conv3d(kernel=stride=(tubelet, patch, patch))` — image `(B,3,H,W)` → `T=1`,
24
+ video `(B,T,3,H,W)`. Token layout `[CLS] + [register × 4] + [patch × N]`.
25
+ - **This repo**: inference-only. Weights + self-contained `modeling_motif.py` (loaded via
26
+ `trust_remote_code`). Training/eval code lives in the GitHub repo.
27
+
28
+ ## Usage
29
+
30
+ The model ships a self-contained `modeling_motif.py`, so it loads with `trust_remote_code=True`.
31
+
32
+ ### Image
33
+
34
+ ```python
35
+ import torch
36
+ from transformers import AutoImageProcessor, AutoModel
37
+ from transformers.image_utils import load_image
38
+
39
+ url = "http://images.cocodataset.org/val2017/000000039769.jpg"
40
+ image = load_image(url)
41
+
42
+ repo = "Motif-Technologies/motif-vision-encoder-7B"
43
+ processor = AutoImageProcessor.from_pretrained(repo)
44
+ model = AutoModel.from_pretrained(repo, trust_remote_code=True, dtype=torch.bfloat16).to("cuda").eval()
45
+
46
+ inputs = processor(images=image, return_tensors="pt").to(model.device, torch.bfloat16)
47
+ with torch.inference_mode():
48
+ outputs = model(**inputs)
49
+
50
+ outputs.last_hidden_state # (1, 1 + 4 + N, 4096) CLS + registers + patch tokens
51
+ outputs.pooler_output # (1, 4096) global (CLS) representation
52
+
53
+ patch_tokens = outputs.last_hidden_state[:, 5:, :] # (1, N, 4096), N = (H/16)*(W/16)
54
+ ```
55
+
56
+ The processor resizes the shorter side to 512, center-crops to 512×512, and normalizes with
57
+ ImageNet mean/std (BICUBIC). `H`/`W` must be multiples of 16.
58
+
59
+ ### Video
60
+
61
+ An image is a 1-frame clip; a video is the same call with a `(B, T, 3, H, W)` tensor. Apply the
62
+ same per-frame transform (resize → center-crop → ImageNet norm) and stack over time:
63
+
64
+ ```python
65
+ import torch
66
+
67
+ video = torch.randn(1, 8, 3, 256, 256, device="cuda", dtype=torch.bfloat16) # (B, T, 3, H, W)
68
+ with torch.inference_mode():
69
+ outputs = model(pixel_values=video)
70
+ ```
71
+
72
+ ## Model details
73
+
74
+ | | |
75
+ |---|---|
76
+ | Backbone | ViT-7B, patch 16, embed 4096, depth 40, heads 32, SwiGLU |
77
+ | Register tokens | 4 |
78
+ | Position encoding | 3D axial RoPE (T,H,W), `base=100.0` |
79
+ | Video tokenizer | 3D Conv, tubelet size 2 |
80
+ | Precision | bf16 weights |
81
+ | Training | DINO + iBOT + KoLeo self-distillation, Gram anchoring, contrastive caption alignment |
82
+
83
+ Outputs (`BaseModelOutputWithPooling`): `last_hidden_state` `(B, 1+4+N, 4096)`,
84
+ `pooler_output` `(B, 4096)`.
85
+
86
+ ## Evaluation
87
+
88
+ Compared against the strongest publicly reported self-supervised / vision backbones. Higher is
89
+ better for every column **except KITTI depth MSE** (lower is better). Best value per column in
90
+ **bold**.
91
+
92
+ | Model | Params | ImageNet-1K<br>lin. probe ↑ | ADE20K<br>mIoU ↑ | DAVIS<br>J&F ↑ | K400 ↑ | KITTI<br>depth MSE ↓ |
93
+ |---|---|---|---|---|---|---|
94
+ | **Motif Vision Encoder** | **7B** | 87.2 | 52.0 | **71.7** | *in progress* | *in progress* |
95
+ | DINOv3 | 7B | **88.2** | **55.9** | 71.1 | **87.8** | **2.3** |
96
+ | DINOv2 | 1.1B | 86.5 | 49.0 | 63.9 | 84.4 | – |
97
+ | V-JEPA 2.1 | 2B | 85.5 | 47.9 | 69.0 | 87.7 | 3.x |
98
+ | SigLIP2 | 2B | 84.5 | 45.4 | 56.1 | 86.9 | – |
99
+
100
+ - **DAVIS video segmentation (J&F 71.7)** — best in the table, surpassing DINOv3 7B (71.1),
101
+ reflecting the encoder's dense, temporally-coherent patch features on video.
102
+ - **ImageNet-1K linear probe (87.2)** and **ADE20K semantic segmentation (52.0 mIoU)** are
103
+ second only to DINOv3 7B while ahead of DINOv2, V-JEPA 2.1, and SigLIP2.
104
+ - K400 action recognition and KITTI depth estimation are still being evaluated and will be added
105
+ as the runs complete.
106
+
107
+ Protocol: DINOv3-style linear/attentive probes for image tasks; V-JEPA 2-style protocol for
108
+ video. Comparison numbers are the best figures reported by each model's authors.
109
+
110
+ ## License
111
+
112
+ Released under **Apache-2.0** (see `LICENSE`). The model was trained on data governed by the
113
+ respective dataset licenses; downstream users are responsible for compliance with those terms.
114
+
115
+ ## Citation
116
+
117
+ ```bibtex
118
+ @misc{motif_vision_encoder,
119
+ title = {MOTIF Vision Encoder: a unified image/video self-supervised vision encoder},
120
+ author = {Motif Technologies},
121
+ year = {2026},
122
+ howpublished = {\url{https://huggingface.co/Motif-Technologies/motif-vision-encoder-7B}}
123
+ }
124
+ ```
config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "MotifVisionModel"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "modeling_motif.MotifVisionConfig",
7
+ "AutoModel": "modeling_motif.MotifVisionModel"
8
+ },
9
+ "depth": 40,
10
+ "drop_path_rate": 0.0,
11
+ "dtype": "bfloat16",
12
+ "embed_dim": 4096,
13
+ "ffn_bias": true,
14
+ "ffn_layer": "swiglu64",
15
+ "ffn_ratio": 3.0,
16
+ "gated_attention": "elementwise",
17
+ "img_size": 512,
18
+ "in_chans": 3,
19
+ "layerscale_init": 1e-05,
20
+ "mask_k_bias": true,
21
+ "model_type": "motif_vision",
22
+ "n_storage_tokens": 4,
23
+ "norm_layer": "layernormbf16",
24
+ "num_frames": 1,
25
+ "num_heads": 32,
26
+ "patch_size": 16,
27
+ "pos_embed_rope_base": 100.0,
28
+ "pos_embed_rope_rescale_coords": 2.0,
29
+ "proj_bias": true,
30
+ "qk_norm": true,
31
+ "qkv_bias": false,
32
+ "transformers_version": "5.8.1",
33
+ "tubelet_size": 2,
34
+ "untie_cls_and_patch_norms": false,
35
+ "untie_global_and_local_cls_norm": true
36
+ }
model-00001-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7199e2a02cd8547f6e2d4271f51aab282eb3686d938e8eb54bb2af242e81cc8a
3
+ size 3939675808
model-00002-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ae96352d3961197861bc7b8147d46d589b2d08557f6a2090157f576895a8e202
3
+ size 3994160376
model-00003-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:02a2a4bab1dc9dacc740b8c25aa90a21028a2e0feadfa077a1cd1268e4a7f09c
3
+ size 3994134520
model-00004-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:04f1a7bf84c7c84cff9687f30c07e3982f9411618083d9e4d58c51dac78d202f
3
+ size 2853015560
model.safetensors.index.json ADDED
@@ -0,0 +1,780 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "metadata": {
3
+ "total_parameters": 7390451712,
4
+ "total_size": 14780903552
5
+ },
6
+ "weight_map": {
7
+ "backbone.blocks.0.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
8
+ "backbone.blocks.0.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
9
+ "backbone.blocks.0.attn.k_norm.weight": "model-00001-of-00004.safetensors",
10
+ "backbone.blocks.0.attn.proj.bias": "model-00001-of-00004.safetensors",
11
+ "backbone.blocks.0.attn.proj.weight": "model-00001-of-00004.safetensors",
12
+ "backbone.blocks.0.attn.q_norm.weight": "model-00001-of-00004.safetensors",
13
+ "backbone.blocks.0.attn.qkv.weight": "model-00001-of-00004.safetensors",
14
+ "backbone.blocks.0.ls1.gamma": "model-00001-of-00004.safetensors",
15
+ "backbone.blocks.0.ls2.gamma": "model-00001-of-00004.safetensors",
16
+ "backbone.blocks.0.mlp.w1.bias": "model-00001-of-00004.safetensors",
17
+ "backbone.blocks.0.mlp.w1.weight": "model-00001-of-00004.safetensors",
18
+ "backbone.blocks.0.mlp.w2.bias": "model-00001-of-00004.safetensors",
19
+ "backbone.blocks.0.mlp.w2.weight": "model-00001-of-00004.safetensors",
20
+ "backbone.blocks.0.mlp.w3.bias": "model-00001-of-00004.safetensors",
21
+ "backbone.blocks.0.mlp.w3.weight": "model-00001-of-00004.safetensors",
22
+ "backbone.blocks.0.norm1.bias": "model-00001-of-00004.safetensors",
23
+ "backbone.blocks.0.norm1.weight": "model-00001-of-00004.safetensors",
24
+ "backbone.blocks.0.norm2.bias": "model-00001-of-00004.safetensors",
25
+ "backbone.blocks.0.norm2.weight": "model-00001-of-00004.safetensors",
26
+ "backbone.blocks.1.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
27
+ "backbone.blocks.1.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
28
+ "backbone.blocks.1.attn.k_norm.weight": "model-00001-of-00004.safetensors",
29
+ "backbone.blocks.1.attn.proj.bias": "model-00001-of-00004.safetensors",
30
+ "backbone.blocks.1.attn.proj.weight": "model-00001-of-00004.safetensors",
31
+ "backbone.blocks.1.attn.q_norm.weight": "model-00001-of-00004.safetensors",
32
+ "backbone.blocks.1.attn.qkv.weight": "model-00001-of-00004.safetensors",
33
+ "backbone.blocks.1.ls1.gamma": "model-00001-of-00004.safetensors",
34
+ "backbone.blocks.1.ls2.gamma": "model-00001-of-00004.safetensors",
35
+ "backbone.blocks.1.mlp.w1.bias": "model-00001-of-00004.safetensors",
36
+ "backbone.blocks.1.mlp.w1.weight": "model-00001-of-00004.safetensors",
37
+ "backbone.blocks.1.mlp.w2.bias": "model-00001-of-00004.safetensors",
38
+ "backbone.blocks.1.mlp.w2.weight": "model-00001-of-00004.safetensors",
39
+ "backbone.blocks.1.mlp.w3.bias": "model-00001-of-00004.safetensors",
40
+ "backbone.blocks.1.mlp.w3.weight": "model-00001-of-00004.safetensors",
41
+ "backbone.blocks.1.norm1.bias": "model-00001-of-00004.safetensors",
42
+ "backbone.blocks.1.norm1.weight": "model-00001-of-00004.safetensors",
43
+ "backbone.blocks.1.norm2.bias": "model-00001-of-00004.safetensors",
44
+ "backbone.blocks.1.norm2.weight": "model-00001-of-00004.safetensors",
45
+ "backbone.blocks.10.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
46
+ "backbone.blocks.10.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
47
+ "backbone.blocks.10.attn.k_norm.weight": "model-00001-of-00004.safetensors",
48
+ "backbone.blocks.10.attn.proj.bias": "model-00001-of-00004.safetensors",
49
+ "backbone.blocks.10.attn.proj.weight": "model-00001-of-00004.safetensors",
50
+ "backbone.blocks.10.attn.q_norm.weight": "model-00001-of-00004.safetensors",
51
+ "backbone.blocks.10.attn.qkv.weight": "model-00001-of-00004.safetensors",
52
+ "backbone.blocks.10.ls1.gamma": "model-00001-of-00004.safetensors",
53
+ "backbone.blocks.10.ls2.gamma": "model-00002-of-00004.safetensors",
54
+ "backbone.blocks.10.mlp.w1.bias": "model-00001-of-00004.safetensors",
55
+ "backbone.blocks.10.mlp.w1.weight": "model-00001-of-00004.safetensors",
56
+ "backbone.blocks.10.mlp.w2.bias": "model-00002-of-00004.safetensors",
57
+ "backbone.blocks.10.mlp.w2.weight": "model-00002-of-00004.safetensors",
58
+ "backbone.blocks.10.mlp.w3.bias": "model-00002-of-00004.safetensors",
59
+ "backbone.blocks.10.mlp.w3.weight": "model-00002-of-00004.safetensors",
60
+ "backbone.blocks.10.norm1.bias": "model-00001-of-00004.safetensors",
61
+ "backbone.blocks.10.norm1.weight": "model-00001-of-00004.safetensors",
62
+ "backbone.blocks.10.norm2.bias": "model-00001-of-00004.safetensors",
63
+ "backbone.blocks.10.norm2.weight": "model-00001-of-00004.safetensors",
64
+ "backbone.blocks.11.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
65
+ "backbone.blocks.11.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
66
+ "backbone.blocks.11.attn.k_norm.weight": "model-00002-of-00004.safetensors",
67
+ "backbone.blocks.11.attn.proj.bias": "model-00002-of-00004.safetensors",
68
+ "backbone.blocks.11.attn.proj.weight": "model-00002-of-00004.safetensors",
69
+ "backbone.blocks.11.attn.q_norm.weight": "model-00002-of-00004.safetensors",
70
+ "backbone.blocks.11.attn.qkv.weight": "model-00002-of-00004.safetensors",
71
+ "backbone.blocks.11.ls1.gamma": "model-00002-of-00004.safetensors",
72
+ "backbone.blocks.11.ls2.gamma": "model-00002-of-00004.safetensors",
73
+ "backbone.blocks.11.mlp.w1.bias": "model-00002-of-00004.safetensors",
74
+ "backbone.blocks.11.mlp.w1.weight": "model-00002-of-00004.safetensors",
75
+ "backbone.blocks.11.mlp.w2.bias": "model-00002-of-00004.safetensors",
76
+ "backbone.blocks.11.mlp.w2.weight": "model-00002-of-00004.safetensors",
77
+ "backbone.blocks.11.mlp.w3.bias": "model-00002-of-00004.safetensors",
78
+ "backbone.blocks.11.mlp.w3.weight": "model-00002-of-00004.safetensors",
79
+ "backbone.blocks.11.norm1.bias": "model-00002-of-00004.safetensors",
80
+ "backbone.blocks.11.norm1.weight": "model-00002-of-00004.safetensors",
81
+ "backbone.blocks.11.norm2.bias": "model-00002-of-00004.safetensors",
82
+ "backbone.blocks.11.norm2.weight": "model-00002-of-00004.safetensors",
83
+ "backbone.blocks.12.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
84
+ "backbone.blocks.12.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
85
+ "backbone.blocks.12.attn.k_norm.weight": "model-00002-of-00004.safetensors",
86
+ "backbone.blocks.12.attn.proj.bias": "model-00002-of-00004.safetensors",
87
+ "backbone.blocks.12.attn.proj.weight": "model-00002-of-00004.safetensors",
88
+ "backbone.blocks.12.attn.q_norm.weight": "model-00002-of-00004.safetensors",
89
+ "backbone.blocks.12.attn.qkv.weight": "model-00002-of-00004.safetensors",
90
+ "backbone.blocks.12.ls1.gamma": "model-00002-of-00004.safetensors",
91
+ "backbone.blocks.12.ls2.gamma": "model-00002-of-00004.safetensors",
92
+ "backbone.blocks.12.mlp.w1.bias": "model-00002-of-00004.safetensors",
93
+ "backbone.blocks.12.mlp.w1.weight": "model-00002-of-00004.safetensors",
94
+ "backbone.blocks.12.mlp.w2.bias": "model-00002-of-00004.safetensors",
95
+ "backbone.blocks.12.mlp.w2.weight": "model-00002-of-00004.safetensors",
96
+ "backbone.blocks.12.mlp.w3.bias": "model-00002-of-00004.safetensors",
97
+ "backbone.blocks.12.mlp.w3.weight": "model-00002-of-00004.safetensors",
98
+ "backbone.blocks.12.norm1.bias": "model-00002-of-00004.safetensors",
99
+ "backbone.blocks.12.norm1.weight": "model-00002-of-00004.safetensors",
100
+ "backbone.blocks.12.norm2.bias": "model-00002-of-00004.safetensors",
101
+ "backbone.blocks.12.norm2.weight": "model-00002-of-00004.safetensors",
102
+ "backbone.blocks.13.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
103
+ "backbone.blocks.13.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
104
+ "backbone.blocks.13.attn.k_norm.weight": "model-00002-of-00004.safetensors",
105
+ "backbone.blocks.13.attn.proj.bias": "model-00002-of-00004.safetensors",
106
+ "backbone.blocks.13.attn.proj.weight": "model-00002-of-00004.safetensors",
107
+ "backbone.blocks.13.attn.q_norm.weight": "model-00002-of-00004.safetensors",
108
+ "backbone.blocks.13.attn.qkv.weight": "model-00002-of-00004.safetensors",
109
+ "backbone.blocks.13.ls1.gamma": "model-00002-of-00004.safetensors",
110
+ "backbone.blocks.13.ls2.gamma": "model-00002-of-00004.safetensors",
111
+ "backbone.blocks.13.mlp.w1.bias": "model-00002-of-00004.safetensors",
112
+ "backbone.blocks.13.mlp.w1.weight": "model-00002-of-00004.safetensors",
113
+ "backbone.blocks.13.mlp.w2.bias": "model-00002-of-00004.safetensors",
114
+ "backbone.blocks.13.mlp.w2.weight": "model-00002-of-00004.safetensors",
115
+ "backbone.blocks.13.mlp.w3.bias": "model-00002-of-00004.safetensors",
116
+ "backbone.blocks.13.mlp.w3.weight": "model-00002-of-00004.safetensors",
117
+ "backbone.blocks.13.norm1.bias": "model-00002-of-00004.safetensors",
118
+ "backbone.blocks.13.norm1.weight": "model-00002-of-00004.safetensors",
119
+ "backbone.blocks.13.norm2.bias": "model-00002-of-00004.safetensors",
120
+ "backbone.blocks.13.norm2.weight": "model-00002-of-00004.safetensors",
121
+ "backbone.blocks.14.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
122
+ "backbone.blocks.14.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
123
+ "backbone.blocks.14.attn.k_norm.weight": "model-00002-of-00004.safetensors",
124
+ "backbone.blocks.14.attn.proj.bias": "model-00002-of-00004.safetensors",
125
+ "backbone.blocks.14.attn.proj.weight": "model-00002-of-00004.safetensors",
126
+ "backbone.blocks.14.attn.q_norm.weight": "model-00002-of-00004.safetensors",
127
+ "backbone.blocks.14.attn.qkv.weight": "model-00002-of-00004.safetensors",
128
+ "backbone.blocks.14.ls1.gamma": "model-00002-of-00004.safetensors",
129
+ "backbone.blocks.14.ls2.gamma": "model-00002-of-00004.safetensors",
130
+ "backbone.blocks.14.mlp.w1.bias": "model-00002-of-00004.safetensors",
131
+ "backbone.blocks.14.mlp.w1.weight": "model-00002-of-00004.safetensors",
132
+ "backbone.blocks.14.mlp.w2.bias": "model-00002-of-00004.safetensors",
133
+ "backbone.blocks.14.mlp.w2.weight": "model-00002-of-00004.safetensors",
134
+ "backbone.blocks.14.mlp.w3.bias": "model-00002-of-00004.safetensors",
135
+ "backbone.blocks.14.mlp.w3.weight": "model-00002-of-00004.safetensors",
136
+ "backbone.blocks.14.norm1.bias": "model-00002-of-00004.safetensors",
137
+ "backbone.blocks.14.norm1.weight": "model-00002-of-00004.safetensors",
138
+ "backbone.blocks.14.norm2.bias": "model-00002-of-00004.safetensors",
139
+ "backbone.blocks.14.norm2.weight": "model-00002-of-00004.safetensors",
140
+ "backbone.blocks.15.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
141
+ "backbone.blocks.15.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
142
+ "backbone.blocks.15.attn.k_norm.weight": "model-00002-of-00004.safetensors",
143
+ "backbone.blocks.15.attn.proj.bias": "model-00002-of-00004.safetensors",
144
+ "backbone.blocks.15.attn.proj.weight": "model-00002-of-00004.safetensors",
145
+ "backbone.blocks.15.attn.q_norm.weight": "model-00002-of-00004.safetensors",
146
+ "backbone.blocks.15.attn.qkv.weight": "model-00002-of-00004.safetensors",
147
+ "backbone.blocks.15.ls1.gamma": "model-00002-of-00004.safetensors",
148
+ "backbone.blocks.15.ls2.gamma": "model-00002-of-00004.safetensors",
149
+ "backbone.blocks.15.mlp.w1.bias": "model-00002-of-00004.safetensors",
150
+ "backbone.blocks.15.mlp.w1.weight": "model-00002-of-00004.safetensors",
151
+ "backbone.blocks.15.mlp.w2.bias": "model-00002-of-00004.safetensors",
152
+ "backbone.blocks.15.mlp.w2.weight": "model-00002-of-00004.safetensors",
153
+ "backbone.blocks.15.mlp.w3.bias": "model-00002-of-00004.safetensors",
154
+ "backbone.blocks.15.mlp.w3.weight": "model-00002-of-00004.safetensors",
155
+ "backbone.blocks.15.norm1.bias": "model-00002-of-00004.safetensors",
156
+ "backbone.blocks.15.norm1.weight": "model-00002-of-00004.safetensors",
157
+ "backbone.blocks.15.norm2.bias": "model-00002-of-00004.safetensors",
158
+ "backbone.blocks.15.norm2.weight": "model-00002-of-00004.safetensors",
159
+ "backbone.blocks.16.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
160
+ "backbone.blocks.16.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
161
+ "backbone.blocks.16.attn.k_norm.weight": "model-00002-of-00004.safetensors",
162
+ "backbone.blocks.16.attn.proj.bias": "model-00002-of-00004.safetensors",
163
+ "backbone.blocks.16.attn.proj.weight": "model-00002-of-00004.safetensors",
164
+ "backbone.blocks.16.attn.q_norm.weight": "model-00002-of-00004.safetensors",
165
+ "backbone.blocks.16.attn.qkv.weight": "model-00002-of-00004.safetensors",
166
+ "backbone.blocks.16.ls1.gamma": "model-00002-of-00004.safetensors",
167
+ "backbone.blocks.16.ls2.gamma": "model-00002-of-00004.safetensors",
168
+ "backbone.blocks.16.mlp.w1.bias": "model-00002-of-00004.safetensors",
169
+ "backbone.blocks.16.mlp.w1.weight": "model-00002-of-00004.safetensors",
170
+ "backbone.blocks.16.mlp.w2.bias": "model-00002-of-00004.safetensors",
171
+ "backbone.blocks.16.mlp.w2.weight": "model-00002-of-00004.safetensors",
172
+ "backbone.blocks.16.mlp.w3.bias": "model-00002-of-00004.safetensors",
173
+ "backbone.blocks.16.mlp.w3.weight": "model-00002-of-00004.safetensors",
174
+ "backbone.blocks.16.norm1.bias": "model-00002-of-00004.safetensors",
175
+ "backbone.blocks.16.norm1.weight": "model-00002-of-00004.safetensors",
176
+ "backbone.blocks.16.norm2.bias": "model-00002-of-00004.safetensors",
177
+ "backbone.blocks.16.norm2.weight": "model-00002-of-00004.safetensors",
178
+ "backbone.blocks.17.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
179
+ "backbone.blocks.17.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
180
+ "backbone.blocks.17.attn.k_norm.weight": "model-00002-of-00004.safetensors",
181
+ "backbone.blocks.17.attn.proj.bias": "model-00002-of-00004.safetensors",
182
+ "backbone.blocks.17.attn.proj.weight": "model-00002-of-00004.safetensors",
183
+ "backbone.blocks.17.attn.q_norm.weight": "model-00002-of-00004.safetensors",
184
+ "backbone.blocks.17.attn.qkv.weight": "model-00002-of-00004.safetensors",
185
+ "backbone.blocks.17.ls1.gamma": "model-00002-of-00004.safetensors",
186
+ "backbone.blocks.17.ls2.gamma": "model-00002-of-00004.safetensors",
187
+ "backbone.blocks.17.mlp.w1.bias": "model-00002-of-00004.safetensors",
188
+ "backbone.blocks.17.mlp.w1.weight": "model-00002-of-00004.safetensors",
189
+ "backbone.blocks.17.mlp.w2.bias": "model-00002-of-00004.safetensors",
190
+ "backbone.blocks.17.mlp.w2.weight": "model-00002-of-00004.safetensors",
191
+ "backbone.blocks.17.mlp.w3.bias": "model-00002-of-00004.safetensors",
192
+ "backbone.blocks.17.mlp.w3.weight": "model-00002-of-00004.safetensors",
193
+ "backbone.blocks.17.norm1.bias": "model-00002-of-00004.safetensors",
194
+ "backbone.blocks.17.norm1.weight": "model-00002-of-00004.safetensors",
195
+ "backbone.blocks.17.norm2.bias": "model-00002-of-00004.safetensors",
196
+ "backbone.blocks.17.norm2.weight": "model-00002-of-00004.safetensors",
197
+ "backbone.blocks.18.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
198
+ "backbone.blocks.18.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
199
+ "backbone.blocks.18.attn.k_norm.weight": "model-00002-of-00004.safetensors",
200
+ "backbone.blocks.18.attn.proj.bias": "model-00002-of-00004.safetensors",
201
+ "backbone.blocks.18.attn.proj.weight": "model-00002-of-00004.safetensors",
202
+ "backbone.blocks.18.attn.q_norm.weight": "model-00002-of-00004.safetensors",
203
+ "backbone.blocks.18.attn.qkv.weight": "model-00002-of-00004.safetensors",
204
+ "backbone.blocks.18.ls1.gamma": "model-00002-of-00004.safetensors",
205
+ "backbone.blocks.18.ls2.gamma": "model-00002-of-00004.safetensors",
206
+ "backbone.blocks.18.mlp.w1.bias": "model-00002-of-00004.safetensors",
207
+ "backbone.blocks.18.mlp.w1.weight": "model-00002-of-00004.safetensors",
208
+ "backbone.blocks.18.mlp.w2.bias": "model-00002-of-00004.safetensors",
209
+ "backbone.blocks.18.mlp.w2.weight": "model-00002-of-00004.safetensors",
210
+ "backbone.blocks.18.mlp.w3.bias": "model-00002-of-00004.safetensors",
211
+ "backbone.blocks.18.mlp.w3.weight": "model-00002-of-00004.safetensors",
212
+ "backbone.blocks.18.norm1.bias": "model-00002-of-00004.safetensors",
213
+ "backbone.blocks.18.norm1.weight": "model-00002-of-00004.safetensors",
214
+ "backbone.blocks.18.norm2.bias": "model-00002-of-00004.safetensors",
215
+ "backbone.blocks.18.norm2.weight": "model-00002-of-00004.safetensors",
216
+ "backbone.blocks.19.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
217
+ "backbone.blocks.19.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
218
+ "backbone.blocks.19.attn.k_norm.weight": "model-00002-of-00004.safetensors",
219
+ "backbone.blocks.19.attn.proj.bias": "model-00002-of-00004.safetensors",
220
+ "backbone.blocks.19.attn.proj.weight": "model-00002-of-00004.safetensors",
221
+ "backbone.blocks.19.attn.q_norm.weight": "model-00002-of-00004.safetensors",
222
+ "backbone.blocks.19.attn.qkv.weight": "model-00002-of-00004.safetensors",
223
+ "backbone.blocks.19.ls1.gamma": "model-00002-of-00004.safetensors",
224
+ "backbone.blocks.19.ls2.gamma": "model-00002-of-00004.safetensors",
225
+ "backbone.blocks.19.mlp.w1.bias": "model-00002-of-00004.safetensors",
226
+ "backbone.blocks.19.mlp.w1.weight": "model-00002-of-00004.safetensors",
227
+ "backbone.blocks.19.mlp.w2.bias": "model-00002-of-00004.safetensors",
228
+ "backbone.blocks.19.mlp.w2.weight": "model-00002-of-00004.safetensors",
229
+ "backbone.blocks.19.mlp.w3.bias": "model-00002-of-00004.safetensors",
230
+ "backbone.blocks.19.mlp.w3.weight": "model-00002-of-00004.safetensors",
231
+ "backbone.blocks.19.norm1.bias": "model-00002-of-00004.safetensors",
232
+ "backbone.blocks.19.norm1.weight": "model-00002-of-00004.safetensors",
233
+ "backbone.blocks.19.norm2.bias": "model-00002-of-00004.safetensors",
234
+ "backbone.blocks.19.norm2.weight": "model-00002-of-00004.safetensors",
235
+ "backbone.blocks.2.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
236
+ "backbone.blocks.2.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
237
+ "backbone.blocks.2.attn.k_norm.weight": "model-00001-of-00004.safetensors",
238
+ "backbone.blocks.2.attn.proj.bias": "model-00001-of-00004.safetensors",
239
+ "backbone.blocks.2.attn.proj.weight": "model-00001-of-00004.safetensors",
240
+ "backbone.blocks.2.attn.q_norm.weight": "model-00001-of-00004.safetensors",
241
+ "backbone.blocks.2.attn.qkv.weight": "model-00001-of-00004.safetensors",
242
+ "backbone.blocks.2.ls1.gamma": "model-00001-of-00004.safetensors",
243
+ "backbone.blocks.2.ls2.gamma": "model-00001-of-00004.safetensors",
244
+ "backbone.blocks.2.mlp.w1.bias": "model-00001-of-00004.safetensors",
245
+ "backbone.blocks.2.mlp.w1.weight": "model-00001-of-00004.safetensors",
246
+ "backbone.blocks.2.mlp.w2.bias": "model-00001-of-00004.safetensors",
247
+ "backbone.blocks.2.mlp.w2.weight": "model-00001-of-00004.safetensors",
248
+ "backbone.blocks.2.mlp.w3.bias": "model-00001-of-00004.safetensors",
249
+ "backbone.blocks.2.mlp.w3.weight": "model-00001-of-00004.safetensors",
250
+ "backbone.blocks.2.norm1.bias": "model-00001-of-00004.safetensors",
251
+ "backbone.blocks.2.norm1.weight": "model-00001-of-00004.safetensors",
252
+ "backbone.blocks.2.norm2.bias": "model-00001-of-00004.safetensors",
253
+ "backbone.blocks.2.norm2.weight": "model-00001-of-00004.safetensors",
254
+ "backbone.blocks.20.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
255
+ "backbone.blocks.20.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
256
+ "backbone.blocks.20.attn.k_norm.weight": "model-00002-of-00004.safetensors",
257
+ "backbone.blocks.20.attn.proj.bias": "model-00002-of-00004.safetensors",
258
+ "backbone.blocks.20.attn.proj.weight": "model-00002-of-00004.safetensors",
259
+ "backbone.blocks.20.attn.q_norm.weight": "model-00002-of-00004.safetensors",
260
+ "backbone.blocks.20.attn.qkv.weight": "model-00002-of-00004.safetensors",
261
+ "backbone.blocks.20.ls1.gamma": "model-00002-of-00004.safetensors",
262
+ "backbone.blocks.20.ls2.gamma": "model-00002-of-00004.safetensors",
263
+ "backbone.blocks.20.mlp.w1.bias": "model-00002-of-00004.safetensors",
264
+ "backbone.blocks.20.mlp.w1.weight": "model-00002-of-00004.safetensors",
265
+ "backbone.blocks.20.mlp.w2.bias": "model-00002-of-00004.safetensors",
266
+ "backbone.blocks.20.mlp.w2.weight": "model-00002-of-00004.safetensors",
267
+ "backbone.blocks.20.mlp.w3.bias": "model-00002-of-00004.safetensors",
268
+ "backbone.blocks.20.mlp.w3.weight": "model-00002-of-00004.safetensors",
269
+ "backbone.blocks.20.norm1.bias": "model-00002-of-00004.safetensors",
270
+ "backbone.blocks.20.norm1.weight": "model-00002-of-00004.safetensors",
271
+ "backbone.blocks.20.norm2.bias": "model-00002-of-00004.safetensors",
272
+ "backbone.blocks.20.norm2.weight": "model-00002-of-00004.safetensors",
273
+ "backbone.blocks.21.attn.gate_proj.bias": "model-00002-of-00004.safetensors",
274
+ "backbone.blocks.21.attn.gate_proj.weight": "model-00002-of-00004.safetensors",
275
+ "backbone.blocks.21.attn.k_norm.weight": "model-00002-of-00004.safetensors",
276
+ "backbone.blocks.21.attn.proj.bias": "model-00002-of-00004.safetensors",
277
+ "backbone.blocks.21.attn.proj.weight": "model-00002-of-00004.safetensors",
278
+ "backbone.blocks.21.attn.q_norm.weight": "model-00002-of-00004.safetensors",
279
+ "backbone.blocks.21.attn.qkv.weight": "model-00002-of-00004.safetensors",
280
+ "backbone.blocks.21.ls1.gamma": "model-00002-of-00004.safetensors",
281
+ "backbone.blocks.21.ls2.gamma": "model-00003-of-00004.safetensors",
282
+ "backbone.blocks.21.mlp.w1.bias": "model-00003-of-00004.safetensors",
283
+ "backbone.blocks.21.mlp.w1.weight": "model-00003-of-00004.safetensors",
284
+ "backbone.blocks.21.mlp.w2.bias": "model-00003-of-00004.safetensors",
285
+ "backbone.blocks.21.mlp.w2.weight": "model-00003-of-00004.safetensors",
286
+ "backbone.blocks.21.mlp.w3.bias": "model-00003-of-00004.safetensors",
287
+ "backbone.blocks.21.mlp.w3.weight": "model-00003-of-00004.safetensors",
288
+ "backbone.blocks.21.norm1.bias": "model-00002-of-00004.safetensors",
289
+ "backbone.blocks.21.norm1.weight": "model-00002-of-00004.safetensors",
290
+ "backbone.blocks.21.norm2.bias": "model-00002-of-00004.safetensors",
291
+ "backbone.blocks.21.norm2.weight": "model-00002-of-00004.safetensors",
292
+ "backbone.blocks.22.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
293
+ "backbone.blocks.22.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
294
+ "backbone.blocks.22.attn.k_norm.weight": "model-00003-of-00004.safetensors",
295
+ "backbone.blocks.22.attn.proj.bias": "model-00003-of-00004.safetensors",
296
+ "backbone.blocks.22.attn.proj.weight": "model-00003-of-00004.safetensors",
297
+ "backbone.blocks.22.attn.q_norm.weight": "model-00003-of-00004.safetensors",
298
+ "backbone.blocks.22.attn.qkv.weight": "model-00003-of-00004.safetensors",
299
+ "backbone.blocks.22.ls1.gamma": "model-00003-of-00004.safetensors",
300
+ "backbone.blocks.22.ls2.gamma": "model-00003-of-00004.safetensors",
301
+ "backbone.blocks.22.mlp.w1.bias": "model-00003-of-00004.safetensors",
302
+ "backbone.blocks.22.mlp.w1.weight": "model-00003-of-00004.safetensors",
303
+ "backbone.blocks.22.mlp.w2.bias": "model-00003-of-00004.safetensors",
304
+ "backbone.blocks.22.mlp.w2.weight": "model-00003-of-00004.safetensors",
305
+ "backbone.blocks.22.mlp.w3.bias": "model-00003-of-00004.safetensors",
306
+ "backbone.blocks.22.mlp.w3.weight": "model-00003-of-00004.safetensors",
307
+ "backbone.blocks.22.norm1.bias": "model-00003-of-00004.safetensors",
308
+ "backbone.blocks.22.norm1.weight": "model-00003-of-00004.safetensors",
309
+ "backbone.blocks.22.norm2.bias": "model-00003-of-00004.safetensors",
310
+ "backbone.blocks.22.norm2.weight": "model-00003-of-00004.safetensors",
311
+ "backbone.blocks.23.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
312
+ "backbone.blocks.23.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
313
+ "backbone.blocks.23.attn.k_norm.weight": "model-00003-of-00004.safetensors",
314
+ "backbone.blocks.23.attn.proj.bias": "model-00003-of-00004.safetensors",
315
+ "backbone.blocks.23.attn.proj.weight": "model-00003-of-00004.safetensors",
316
+ "backbone.blocks.23.attn.q_norm.weight": "model-00003-of-00004.safetensors",
317
+ "backbone.blocks.23.attn.qkv.weight": "model-00003-of-00004.safetensors",
318
+ "backbone.blocks.23.ls1.gamma": "model-00003-of-00004.safetensors",
319
+ "backbone.blocks.23.ls2.gamma": "model-00003-of-00004.safetensors",
320
+ "backbone.blocks.23.mlp.w1.bias": "model-00003-of-00004.safetensors",
321
+ "backbone.blocks.23.mlp.w1.weight": "model-00003-of-00004.safetensors",
322
+ "backbone.blocks.23.mlp.w2.bias": "model-00003-of-00004.safetensors",
323
+ "backbone.blocks.23.mlp.w2.weight": "model-00003-of-00004.safetensors",
324
+ "backbone.blocks.23.mlp.w3.bias": "model-00003-of-00004.safetensors",
325
+ "backbone.blocks.23.mlp.w3.weight": "model-00003-of-00004.safetensors",
326
+ "backbone.blocks.23.norm1.bias": "model-00003-of-00004.safetensors",
327
+ "backbone.blocks.23.norm1.weight": "model-00003-of-00004.safetensors",
328
+ "backbone.blocks.23.norm2.bias": "model-00003-of-00004.safetensors",
329
+ "backbone.blocks.23.norm2.weight": "model-00003-of-00004.safetensors",
330
+ "backbone.blocks.24.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
331
+ "backbone.blocks.24.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
332
+ "backbone.blocks.24.attn.k_norm.weight": "model-00003-of-00004.safetensors",
333
+ "backbone.blocks.24.attn.proj.bias": "model-00003-of-00004.safetensors",
334
+ "backbone.blocks.24.attn.proj.weight": "model-00003-of-00004.safetensors",
335
+ "backbone.blocks.24.attn.q_norm.weight": "model-00003-of-00004.safetensors",
336
+ "backbone.blocks.24.attn.qkv.weight": "model-00003-of-00004.safetensors",
337
+ "backbone.blocks.24.ls1.gamma": "model-00003-of-00004.safetensors",
338
+ "backbone.blocks.24.ls2.gamma": "model-00003-of-00004.safetensors",
339
+ "backbone.blocks.24.mlp.w1.bias": "model-00003-of-00004.safetensors",
340
+ "backbone.blocks.24.mlp.w1.weight": "model-00003-of-00004.safetensors",
341
+ "backbone.blocks.24.mlp.w2.bias": "model-00003-of-00004.safetensors",
342
+ "backbone.blocks.24.mlp.w2.weight": "model-00003-of-00004.safetensors",
343
+ "backbone.blocks.24.mlp.w3.bias": "model-00003-of-00004.safetensors",
344
+ "backbone.blocks.24.mlp.w3.weight": "model-00003-of-00004.safetensors",
345
+ "backbone.blocks.24.norm1.bias": "model-00003-of-00004.safetensors",
346
+ "backbone.blocks.24.norm1.weight": "model-00003-of-00004.safetensors",
347
+ "backbone.blocks.24.norm2.bias": "model-00003-of-00004.safetensors",
348
+ "backbone.blocks.24.norm2.weight": "model-00003-of-00004.safetensors",
349
+ "backbone.blocks.25.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
350
+ "backbone.blocks.25.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
351
+ "backbone.blocks.25.attn.k_norm.weight": "model-00003-of-00004.safetensors",
352
+ "backbone.blocks.25.attn.proj.bias": "model-00003-of-00004.safetensors",
353
+ "backbone.blocks.25.attn.proj.weight": "model-00003-of-00004.safetensors",
354
+ "backbone.blocks.25.attn.q_norm.weight": "model-00003-of-00004.safetensors",
355
+ "backbone.blocks.25.attn.qkv.weight": "model-00003-of-00004.safetensors",
356
+ "backbone.blocks.25.ls1.gamma": "model-00003-of-00004.safetensors",
357
+ "backbone.blocks.25.ls2.gamma": "model-00003-of-00004.safetensors",
358
+ "backbone.blocks.25.mlp.w1.bias": "model-00003-of-00004.safetensors",
359
+ "backbone.blocks.25.mlp.w1.weight": "model-00003-of-00004.safetensors",
360
+ "backbone.blocks.25.mlp.w2.bias": "model-00003-of-00004.safetensors",
361
+ "backbone.blocks.25.mlp.w2.weight": "model-00003-of-00004.safetensors",
362
+ "backbone.blocks.25.mlp.w3.bias": "model-00003-of-00004.safetensors",
363
+ "backbone.blocks.25.mlp.w3.weight": "model-00003-of-00004.safetensors",
364
+ "backbone.blocks.25.norm1.bias": "model-00003-of-00004.safetensors",
365
+ "backbone.blocks.25.norm1.weight": "model-00003-of-00004.safetensors",
366
+ "backbone.blocks.25.norm2.bias": "model-00003-of-00004.safetensors",
367
+ "backbone.blocks.25.norm2.weight": "model-00003-of-00004.safetensors",
368
+ "backbone.blocks.26.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
369
+ "backbone.blocks.26.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
370
+ "backbone.blocks.26.attn.k_norm.weight": "model-00003-of-00004.safetensors",
371
+ "backbone.blocks.26.attn.proj.bias": "model-00003-of-00004.safetensors",
372
+ "backbone.blocks.26.attn.proj.weight": "model-00003-of-00004.safetensors",
373
+ "backbone.blocks.26.attn.q_norm.weight": "model-00003-of-00004.safetensors",
374
+ "backbone.blocks.26.attn.qkv.weight": "model-00003-of-00004.safetensors",
375
+ "backbone.blocks.26.ls1.gamma": "model-00003-of-00004.safetensors",
376
+ "backbone.blocks.26.ls2.gamma": "model-00003-of-00004.safetensors",
377
+ "backbone.blocks.26.mlp.w1.bias": "model-00003-of-00004.safetensors",
378
+ "backbone.blocks.26.mlp.w1.weight": "model-00003-of-00004.safetensors",
379
+ "backbone.blocks.26.mlp.w2.bias": "model-00003-of-00004.safetensors",
380
+ "backbone.blocks.26.mlp.w2.weight": "model-00003-of-00004.safetensors",
381
+ "backbone.blocks.26.mlp.w3.bias": "model-00003-of-00004.safetensors",
382
+ "backbone.blocks.26.mlp.w3.weight": "model-00003-of-00004.safetensors",
383
+ "backbone.blocks.26.norm1.bias": "model-00003-of-00004.safetensors",
384
+ "backbone.blocks.26.norm1.weight": "model-00003-of-00004.safetensors",
385
+ "backbone.blocks.26.norm2.bias": "model-00003-of-00004.safetensors",
386
+ "backbone.blocks.26.norm2.weight": "model-00003-of-00004.safetensors",
387
+ "backbone.blocks.27.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
388
+ "backbone.blocks.27.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
389
+ "backbone.blocks.27.attn.k_norm.weight": "model-00003-of-00004.safetensors",
390
+ "backbone.blocks.27.attn.proj.bias": "model-00003-of-00004.safetensors",
391
+ "backbone.blocks.27.attn.proj.weight": "model-00003-of-00004.safetensors",
392
+ "backbone.blocks.27.attn.q_norm.weight": "model-00003-of-00004.safetensors",
393
+ "backbone.blocks.27.attn.qkv.weight": "model-00003-of-00004.safetensors",
394
+ "backbone.blocks.27.ls1.gamma": "model-00003-of-00004.safetensors",
395
+ "backbone.blocks.27.ls2.gamma": "model-00003-of-00004.safetensors",
396
+ "backbone.blocks.27.mlp.w1.bias": "model-00003-of-00004.safetensors",
397
+ "backbone.blocks.27.mlp.w1.weight": "model-00003-of-00004.safetensors",
398
+ "backbone.blocks.27.mlp.w2.bias": "model-00003-of-00004.safetensors",
399
+ "backbone.blocks.27.mlp.w2.weight": "model-00003-of-00004.safetensors",
400
+ "backbone.blocks.27.mlp.w3.bias": "model-00003-of-00004.safetensors",
401
+ "backbone.blocks.27.mlp.w3.weight": "model-00003-of-00004.safetensors",
402
+ "backbone.blocks.27.norm1.bias": "model-00003-of-00004.safetensors",
403
+ "backbone.blocks.27.norm1.weight": "model-00003-of-00004.safetensors",
404
+ "backbone.blocks.27.norm2.bias": "model-00003-of-00004.safetensors",
405
+ "backbone.blocks.27.norm2.weight": "model-00003-of-00004.safetensors",
406
+ "backbone.blocks.28.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
407
+ "backbone.blocks.28.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
408
+ "backbone.blocks.28.attn.k_norm.weight": "model-00003-of-00004.safetensors",
409
+ "backbone.blocks.28.attn.proj.bias": "model-00003-of-00004.safetensors",
410
+ "backbone.blocks.28.attn.proj.weight": "model-00003-of-00004.safetensors",
411
+ "backbone.blocks.28.attn.q_norm.weight": "model-00003-of-00004.safetensors",
412
+ "backbone.blocks.28.attn.qkv.weight": "model-00003-of-00004.safetensors",
413
+ "backbone.blocks.28.ls1.gamma": "model-00003-of-00004.safetensors",
414
+ "backbone.blocks.28.ls2.gamma": "model-00003-of-00004.safetensors",
415
+ "backbone.blocks.28.mlp.w1.bias": "model-00003-of-00004.safetensors",
416
+ "backbone.blocks.28.mlp.w1.weight": "model-00003-of-00004.safetensors",
417
+ "backbone.blocks.28.mlp.w2.bias": "model-00003-of-00004.safetensors",
418
+ "backbone.blocks.28.mlp.w2.weight": "model-00003-of-00004.safetensors",
419
+ "backbone.blocks.28.mlp.w3.bias": "model-00003-of-00004.safetensors",
420
+ "backbone.blocks.28.mlp.w3.weight": "model-00003-of-00004.safetensors",
421
+ "backbone.blocks.28.norm1.bias": "model-00003-of-00004.safetensors",
422
+ "backbone.blocks.28.norm1.weight": "model-00003-of-00004.safetensors",
423
+ "backbone.blocks.28.norm2.bias": "model-00003-of-00004.safetensors",
424
+ "backbone.blocks.28.norm2.weight": "model-00003-of-00004.safetensors",
425
+ "backbone.blocks.29.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
426
+ "backbone.blocks.29.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
427
+ "backbone.blocks.29.attn.k_norm.weight": "model-00003-of-00004.safetensors",
428
+ "backbone.blocks.29.attn.proj.bias": "model-00003-of-00004.safetensors",
429
+ "backbone.blocks.29.attn.proj.weight": "model-00003-of-00004.safetensors",
430
+ "backbone.blocks.29.attn.q_norm.weight": "model-00003-of-00004.safetensors",
431
+ "backbone.blocks.29.attn.qkv.weight": "model-00003-of-00004.safetensors",
432
+ "backbone.blocks.29.ls1.gamma": "model-00003-of-00004.safetensors",
433
+ "backbone.blocks.29.ls2.gamma": "model-00003-of-00004.safetensors",
434
+ "backbone.blocks.29.mlp.w1.bias": "model-00003-of-00004.safetensors",
435
+ "backbone.blocks.29.mlp.w1.weight": "model-00003-of-00004.safetensors",
436
+ "backbone.blocks.29.mlp.w2.bias": "model-00003-of-00004.safetensors",
437
+ "backbone.blocks.29.mlp.w2.weight": "model-00003-of-00004.safetensors",
438
+ "backbone.blocks.29.mlp.w3.bias": "model-00003-of-00004.safetensors",
439
+ "backbone.blocks.29.mlp.w3.weight": "model-00003-of-00004.safetensors",
440
+ "backbone.blocks.29.norm1.bias": "model-00003-of-00004.safetensors",
441
+ "backbone.blocks.29.norm1.weight": "model-00003-of-00004.safetensors",
442
+ "backbone.blocks.29.norm2.bias": "model-00003-of-00004.safetensors",
443
+ "backbone.blocks.29.norm2.weight": "model-00003-of-00004.safetensors",
444
+ "backbone.blocks.3.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
445
+ "backbone.blocks.3.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
446
+ "backbone.blocks.3.attn.k_norm.weight": "model-00001-of-00004.safetensors",
447
+ "backbone.blocks.3.attn.proj.bias": "model-00001-of-00004.safetensors",
448
+ "backbone.blocks.3.attn.proj.weight": "model-00001-of-00004.safetensors",
449
+ "backbone.blocks.3.attn.q_norm.weight": "model-00001-of-00004.safetensors",
450
+ "backbone.blocks.3.attn.qkv.weight": "model-00001-of-00004.safetensors",
451
+ "backbone.blocks.3.ls1.gamma": "model-00001-of-00004.safetensors",
452
+ "backbone.blocks.3.ls2.gamma": "model-00001-of-00004.safetensors",
453
+ "backbone.blocks.3.mlp.w1.bias": "model-00001-of-00004.safetensors",
454
+ "backbone.blocks.3.mlp.w1.weight": "model-00001-of-00004.safetensors",
455
+ "backbone.blocks.3.mlp.w2.bias": "model-00001-of-00004.safetensors",
456
+ "backbone.blocks.3.mlp.w2.weight": "model-00001-of-00004.safetensors",
457
+ "backbone.blocks.3.mlp.w3.bias": "model-00001-of-00004.safetensors",
458
+ "backbone.blocks.3.mlp.w3.weight": "model-00001-of-00004.safetensors",
459
+ "backbone.blocks.3.norm1.bias": "model-00001-of-00004.safetensors",
460
+ "backbone.blocks.3.norm1.weight": "model-00001-of-00004.safetensors",
461
+ "backbone.blocks.3.norm2.bias": "model-00001-of-00004.safetensors",
462
+ "backbone.blocks.3.norm2.weight": "model-00001-of-00004.safetensors",
463
+ "backbone.blocks.30.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
464
+ "backbone.blocks.30.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
465
+ "backbone.blocks.30.attn.k_norm.weight": "model-00003-of-00004.safetensors",
466
+ "backbone.blocks.30.attn.proj.bias": "model-00003-of-00004.safetensors",
467
+ "backbone.blocks.30.attn.proj.weight": "model-00003-of-00004.safetensors",
468
+ "backbone.blocks.30.attn.q_norm.weight": "model-00003-of-00004.safetensors",
469
+ "backbone.blocks.30.attn.qkv.weight": "model-00003-of-00004.safetensors",
470
+ "backbone.blocks.30.ls1.gamma": "model-00003-of-00004.safetensors",
471
+ "backbone.blocks.30.ls2.gamma": "model-00003-of-00004.safetensors",
472
+ "backbone.blocks.30.mlp.w1.bias": "model-00003-of-00004.safetensors",
473
+ "backbone.blocks.30.mlp.w1.weight": "model-00003-of-00004.safetensors",
474
+ "backbone.blocks.30.mlp.w2.bias": "model-00003-of-00004.safetensors",
475
+ "backbone.blocks.30.mlp.w2.weight": "model-00003-of-00004.safetensors",
476
+ "backbone.blocks.30.mlp.w3.bias": "model-00003-of-00004.safetensors",
477
+ "backbone.blocks.30.mlp.w3.weight": "model-00003-of-00004.safetensors",
478
+ "backbone.blocks.30.norm1.bias": "model-00003-of-00004.safetensors",
479
+ "backbone.blocks.30.norm1.weight": "model-00003-of-00004.safetensors",
480
+ "backbone.blocks.30.norm2.bias": "model-00003-of-00004.safetensors",
481
+ "backbone.blocks.30.norm2.weight": "model-00003-of-00004.safetensors",
482
+ "backbone.blocks.31.attn.gate_proj.bias": "model-00003-of-00004.safetensors",
483
+ "backbone.blocks.31.attn.gate_proj.weight": "model-00003-of-00004.safetensors",
484
+ "backbone.blocks.31.attn.k_norm.weight": "model-00003-of-00004.safetensors",
485
+ "backbone.blocks.31.attn.proj.bias": "model-00003-of-00004.safetensors",
486
+ "backbone.blocks.31.attn.proj.weight": "model-00003-of-00004.safetensors",
487
+ "backbone.blocks.31.attn.q_norm.weight": "model-00003-of-00004.safetensors",
488
+ "backbone.blocks.31.attn.qkv.weight": "model-00003-of-00004.safetensors",
489
+ "backbone.blocks.31.ls1.gamma": "model-00003-of-00004.safetensors",
490
+ "backbone.blocks.31.ls2.gamma": "model-00003-of-00004.safetensors",
491
+ "backbone.blocks.31.mlp.w1.bias": "model-00003-of-00004.safetensors",
492
+ "backbone.blocks.31.mlp.w1.weight": "model-00003-of-00004.safetensors",
493
+ "backbone.blocks.31.mlp.w2.bias": "model-00003-of-00004.safetensors",
494
+ "backbone.blocks.31.mlp.w2.weight": "model-00003-of-00004.safetensors",
495
+ "backbone.blocks.31.mlp.w3.bias": "model-00003-of-00004.safetensors",
496
+ "backbone.blocks.31.mlp.w3.weight": "model-00003-of-00004.safetensors",
497
+ "backbone.blocks.31.norm1.bias": "model-00003-of-00004.safetensors",
498
+ "backbone.blocks.31.norm1.weight": "model-00003-of-00004.safetensors",
499
+ "backbone.blocks.31.norm2.bias": "model-00003-of-00004.safetensors",
500
+ "backbone.blocks.31.norm2.weight": "model-00003-of-00004.safetensors",
501
+ "backbone.blocks.32.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
502
+ "backbone.blocks.32.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
503
+ "backbone.blocks.32.attn.k_norm.weight": "model-00004-of-00004.safetensors",
504
+ "backbone.blocks.32.attn.proj.bias": "model-00004-of-00004.safetensors",
505
+ "backbone.blocks.32.attn.proj.weight": "model-00004-of-00004.safetensors",
506
+ "backbone.blocks.32.attn.q_norm.weight": "model-00004-of-00004.safetensors",
507
+ "backbone.blocks.32.attn.qkv.weight": "model-00003-of-00004.safetensors",
508
+ "backbone.blocks.32.ls1.gamma": "model-00004-of-00004.safetensors",
509
+ "backbone.blocks.32.ls2.gamma": "model-00004-of-00004.safetensors",
510
+ "backbone.blocks.32.mlp.w1.bias": "model-00004-of-00004.safetensors",
511
+ "backbone.blocks.32.mlp.w1.weight": "model-00004-of-00004.safetensors",
512
+ "backbone.blocks.32.mlp.w2.bias": "model-00004-of-00004.safetensors",
513
+ "backbone.blocks.32.mlp.w2.weight": "model-00004-of-00004.safetensors",
514
+ "backbone.blocks.32.mlp.w3.bias": "model-00004-of-00004.safetensors",
515
+ "backbone.blocks.32.mlp.w3.weight": "model-00004-of-00004.safetensors",
516
+ "backbone.blocks.32.norm1.bias": "model-00003-of-00004.safetensors",
517
+ "backbone.blocks.32.norm1.weight": "model-00003-of-00004.safetensors",
518
+ "backbone.blocks.32.norm2.bias": "model-00004-of-00004.safetensors",
519
+ "backbone.blocks.32.norm2.weight": "model-00004-of-00004.safetensors",
520
+ "backbone.blocks.33.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
521
+ "backbone.blocks.33.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
522
+ "backbone.blocks.33.attn.k_norm.weight": "model-00004-of-00004.safetensors",
523
+ "backbone.blocks.33.attn.proj.bias": "model-00004-of-00004.safetensors",
524
+ "backbone.blocks.33.attn.proj.weight": "model-00004-of-00004.safetensors",
525
+ "backbone.blocks.33.attn.q_norm.weight": "model-00004-of-00004.safetensors",
526
+ "backbone.blocks.33.attn.qkv.weight": "model-00004-of-00004.safetensors",
527
+ "backbone.blocks.33.ls1.gamma": "model-00004-of-00004.safetensors",
528
+ "backbone.blocks.33.ls2.gamma": "model-00004-of-00004.safetensors",
529
+ "backbone.blocks.33.mlp.w1.bias": "model-00004-of-00004.safetensors",
530
+ "backbone.blocks.33.mlp.w1.weight": "model-00004-of-00004.safetensors",
531
+ "backbone.blocks.33.mlp.w2.bias": "model-00004-of-00004.safetensors",
532
+ "backbone.blocks.33.mlp.w2.weight": "model-00004-of-00004.safetensors",
533
+ "backbone.blocks.33.mlp.w3.bias": "model-00004-of-00004.safetensors",
534
+ "backbone.blocks.33.mlp.w3.weight": "model-00004-of-00004.safetensors",
535
+ "backbone.blocks.33.norm1.bias": "model-00004-of-00004.safetensors",
536
+ "backbone.blocks.33.norm1.weight": "model-00004-of-00004.safetensors",
537
+ "backbone.blocks.33.norm2.bias": "model-00004-of-00004.safetensors",
538
+ "backbone.blocks.33.norm2.weight": "model-00004-of-00004.safetensors",
539
+ "backbone.blocks.34.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
540
+ "backbone.blocks.34.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
541
+ "backbone.blocks.34.attn.k_norm.weight": "model-00004-of-00004.safetensors",
542
+ "backbone.blocks.34.attn.proj.bias": "model-00004-of-00004.safetensors",
543
+ "backbone.blocks.34.attn.proj.weight": "model-00004-of-00004.safetensors",
544
+ "backbone.blocks.34.attn.q_norm.weight": "model-00004-of-00004.safetensors",
545
+ "backbone.blocks.34.attn.qkv.weight": "model-00004-of-00004.safetensors",
546
+ "backbone.blocks.34.ls1.gamma": "model-00004-of-00004.safetensors",
547
+ "backbone.blocks.34.ls2.gamma": "model-00004-of-00004.safetensors",
548
+ "backbone.blocks.34.mlp.w1.bias": "model-00004-of-00004.safetensors",
549
+ "backbone.blocks.34.mlp.w1.weight": "model-00004-of-00004.safetensors",
550
+ "backbone.blocks.34.mlp.w2.bias": "model-00004-of-00004.safetensors",
551
+ "backbone.blocks.34.mlp.w2.weight": "model-00004-of-00004.safetensors",
552
+ "backbone.blocks.34.mlp.w3.bias": "model-00004-of-00004.safetensors",
553
+ "backbone.blocks.34.mlp.w3.weight": "model-00004-of-00004.safetensors",
554
+ "backbone.blocks.34.norm1.bias": "model-00004-of-00004.safetensors",
555
+ "backbone.blocks.34.norm1.weight": "model-00004-of-00004.safetensors",
556
+ "backbone.blocks.34.norm2.bias": "model-00004-of-00004.safetensors",
557
+ "backbone.blocks.34.norm2.weight": "model-00004-of-00004.safetensors",
558
+ "backbone.blocks.35.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
559
+ "backbone.blocks.35.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
560
+ "backbone.blocks.35.attn.k_norm.weight": "model-00004-of-00004.safetensors",
561
+ "backbone.blocks.35.attn.proj.bias": "model-00004-of-00004.safetensors",
562
+ "backbone.blocks.35.attn.proj.weight": "model-00004-of-00004.safetensors",
563
+ "backbone.blocks.35.attn.q_norm.weight": "model-00004-of-00004.safetensors",
564
+ "backbone.blocks.35.attn.qkv.weight": "model-00004-of-00004.safetensors",
565
+ "backbone.blocks.35.ls1.gamma": "model-00004-of-00004.safetensors",
566
+ "backbone.blocks.35.ls2.gamma": "model-00004-of-00004.safetensors",
567
+ "backbone.blocks.35.mlp.w1.bias": "model-00004-of-00004.safetensors",
568
+ "backbone.blocks.35.mlp.w1.weight": "model-00004-of-00004.safetensors",
569
+ "backbone.blocks.35.mlp.w2.bias": "model-00004-of-00004.safetensors",
570
+ "backbone.blocks.35.mlp.w2.weight": "model-00004-of-00004.safetensors",
571
+ "backbone.blocks.35.mlp.w3.bias": "model-00004-of-00004.safetensors",
572
+ "backbone.blocks.35.mlp.w3.weight": "model-00004-of-00004.safetensors",
573
+ "backbone.blocks.35.norm1.bias": "model-00004-of-00004.safetensors",
574
+ "backbone.blocks.35.norm1.weight": "model-00004-of-00004.safetensors",
575
+ "backbone.blocks.35.norm2.bias": "model-00004-of-00004.safetensors",
576
+ "backbone.blocks.35.norm2.weight": "model-00004-of-00004.safetensors",
577
+ "backbone.blocks.36.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
578
+ "backbone.blocks.36.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
579
+ "backbone.blocks.36.attn.k_norm.weight": "model-00004-of-00004.safetensors",
580
+ "backbone.blocks.36.attn.proj.bias": "model-00004-of-00004.safetensors",
581
+ "backbone.blocks.36.attn.proj.weight": "model-00004-of-00004.safetensors",
582
+ "backbone.blocks.36.attn.q_norm.weight": "model-00004-of-00004.safetensors",
583
+ "backbone.blocks.36.attn.qkv.weight": "model-00004-of-00004.safetensors",
584
+ "backbone.blocks.36.ls1.gamma": "model-00004-of-00004.safetensors",
585
+ "backbone.blocks.36.ls2.gamma": "model-00004-of-00004.safetensors",
586
+ "backbone.blocks.36.mlp.w1.bias": "model-00004-of-00004.safetensors",
587
+ "backbone.blocks.36.mlp.w1.weight": "model-00004-of-00004.safetensors",
588
+ "backbone.blocks.36.mlp.w2.bias": "model-00004-of-00004.safetensors",
589
+ "backbone.blocks.36.mlp.w2.weight": "model-00004-of-00004.safetensors",
590
+ "backbone.blocks.36.mlp.w3.bias": "model-00004-of-00004.safetensors",
591
+ "backbone.blocks.36.mlp.w3.weight": "model-00004-of-00004.safetensors",
592
+ "backbone.blocks.36.norm1.bias": "model-00004-of-00004.safetensors",
593
+ "backbone.blocks.36.norm1.weight": "model-00004-of-00004.safetensors",
594
+ "backbone.blocks.36.norm2.bias": "model-00004-of-00004.safetensors",
595
+ "backbone.blocks.36.norm2.weight": "model-00004-of-00004.safetensors",
596
+ "backbone.blocks.37.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
597
+ "backbone.blocks.37.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
598
+ "backbone.blocks.37.attn.k_norm.weight": "model-00004-of-00004.safetensors",
599
+ "backbone.blocks.37.attn.proj.bias": "model-00004-of-00004.safetensors",
600
+ "backbone.blocks.37.attn.proj.weight": "model-00004-of-00004.safetensors",
601
+ "backbone.blocks.37.attn.q_norm.weight": "model-00004-of-00004.safetensors",
602
+ "backbone.blocks.37.attn.qkv.weight": "model-00004-of-00004.safetensors",
603
+ "backbone.blocks.37.ls1.gamma": "model-00004-of-00004.safetensors",
604
+ "backbone.blocks.37.ls2.gamma": "model-00004-of-00004.safetensors",
605
+ "backbone.blocks.37.mlp.w1.bias": "model-00004-of-00004.safetensors",
606
+ "backbone.blocks.37.mlp.w1.weight": "model-00004-of-00004.safetensors",
607
+ "backbone.blocks.37.mlp.w2.bias": "model-00004-of-00004.safetensors",
608
+ "backbone.blocks.37.mlp.w2.weight": "model-00004-of-00004.safetensors",
609
+ "backbone.blocks.37.mlp.w3.bias": "model-00004-of-00004.safetensors",
610
+ "backbone.blocks.37.mlp.w3.weight": "model-00004-of-00004.safetensors",
611
+ "backbone.blocks.37.norm1.bias": "model-00004-of-00004.safetensors",
612
+ "backbone.blocks.37.norm1.weight": "model-00004-of-00004.safetensors",
613
+ "backbone.blocks.37.norm2.bias": "model-00004-of-00004.safetensors",
614
+ "backbone.blocks.37.norm2.weight": "model-00004-of-00004.safetensors",
615
+ "backbone.blocks.38.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
616
+ "backbone.blocks.38.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
617
+ "backbone.blocks.38.attn.k_norm.weight": "model-00004-of-00004.safetensors",
618
+ "backbone.blocks.38.attn.proj.bias": "model-00004-of-00004.safetensors",
619
+ "backbone.blocks.38.attn.proj.weight": "model-00004-of-00004.safetensors",
620
+ "backbone.blocks.38.attn.q_norm.weight": "model-00004-of-00004.safetensors",
621
+ "backbone.blocks.38.attn.qkv.weight": "model-00004-of-00004.safetensors",
622
+ "backbone.blocks.38.ls1.gamma": "model-00004-of-00004.safetensors",
623
+ "backbone.blocks.38.ls2.gamma": "model-00004-of-00004.safetensors",
624
+ "backbone.blocks.38.mlp.w1.bias": "model-00004-of-00004.safetensors",
625
+ "backbone.blocks.38.mlp.w1.weight": "model-00004-of-00004.safetensors",
626
+ "backbone.blocks.38.mlp.w2.bias": "model-00004-of-00004.safetensors",
627
+ "backbone.blocks.38.mlp.w2.weight": "model-00004-of-00004.safetensors",
628
+ "backbone.blocks.38.mlp.w3.bias": "model-00004-of-00004.safetensors",
629
+ "backbone.blocks.38.mlp.w3.weight": "model-00004-of-00004.safetensors",
630
+ "backbone.blocks.38.norm1.bias": "model-00004-of-00004.safetensors",
631
+ "backbone.blocks.38.norm1.weight": "model-00004-of-00004.safetensors",
632
+ "backbone.blocks.38.norm2.bias": "model-00004-of-00004.safetensors",
633
+ "backbone.blocks.38.norm2.weight": "model-00004-of-00004.safetensors",
634
+ "backbone.blocks.39.attn.gate_proj.bias": "model-00004-of-00004.safetensors",
635
+ "backbone.blocks.39.attn.gate_proj.weight": "model-00004-of-00004.safetensors",
636
+ "backbone.blocks.39.attn.k_norm.weight": "model-00004-of-00004.safetensors",
637
+ "backbone.blocks.39.attn.proj.bias": "model-00004-of-00004.safetensors",
638
+ "backbone.blocks.39.attn.proj.weight": "model-00004-of-00004.safetensors",
639
+ "backbone.blocks.39.attn.q_norm.weight": "model-00004-of-00004.safetensors",
640
+ "backbone.blocks.39.attn.qkv.weight": "model-00004-of-00004.safetensors",
641
+ "backbone.blocks.39.ls1.gamma": "model-00004-of-00004.safetensors",
642
+ "backbone.blocks.39.ls2.gamma": "model-00004-of-00004.safetensors",
643
+ "backbone.blocks.39.mlp.w1.bias": "model-00004-of-00004.safetensors",
644
+ "backbone.blocks.39.mlp.w1.weight": "model-00004-of-00004.safetensors",
645
+ "backbone.blocks.39.mlp.w2.bias": "model-00004-of-00004.safetensors",
646
+ "backbone.blocks.39.mlp.w2.weight": "model-00004-of-00004.safetensors",
647
+ "backbone.blocks.39.mlp.w3.bias": "model-00004-of-00004.safetensors",
648
+ "backbone.blocks.39.mlp.w3.weight": "model-00004-of-00004.safetensors",
649
+ "backbone.blocks.39.norm1.bias": "model-00004-of-00004.safetensors",
650
+ "backbone.blocks.39.norm1.weight": "model-00004-of-00004.safetensors",
651
+ "backbone.blocks.39.norm2.bias": "model-00004-of-00004.safetensors",
652
+ "backbone.blocks.39.norm2.weight": "model-00004-of-00004.safetensors",
653
+ "backbone.blocks.4.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
654
+ "backbone.blocks.4.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
655
+ "backbone.blocks.4.attn.k_norm.weight": "model-00001-of-00004.safetensors",
656
+ "backbone.blocks.4.attn.proj.bias": "model-00001-of-00004.safetensors",
657
+ "backbone.blocks.4.attn.proj.weight": "model-00001-of-00004.safetensors",
658
+ "backbone.blocks.4.attn.q_norm.weight": "model-00001-of-00004.safetensors",
659
+ "backbone.blocks.4.attn.qkv.weight": "model-00001-of-00004.safetensors",
660
+ "backbone.blocks.4.ls1.gamma": "model-00001-of-00004.safetensors",
661
+ "backbone.blocks.4.ls2.gamma": "model-00001-of-00004.safetensors",
662
+ "backbone.blocks.4.mlp.w1.bias": "model-00001-of-00004.safetensors",
663
+ "backbone.blocks.4.mlp.w1.weight": "model-00001-of-00004.safetensors",
664
+ "backbone.blocks.4.mlp.w2.bias": "model-00001-of-00004.safetensors",
665
+ "backbone.blocks.4.mlp.w2.weight": "model-00001-of-00004.safetensors",
666
+ "backbone.blocks.4.mlp.w3.bias": "model-00001-of-00004.safetensors",
667
+ "backbone.blocks.4.mlp.w3.weight": "model-00001-of-00004.safetensors",
668
+ "backbone.blocks.4.norm1.bias": "model-00001-of-00004.safetensors",
669
+ "backbone.blocks.4.norm1.weight": "model-00001-of-00004.safetensors",
670
+ "backbone.blocks.4.norm2.bias": "model-00001-of-00004.safetensors",
671
+ "backbone.blocks.4.norm2.weight": "model-00001-of-00004.safetensors",
672
+ "backbone.blocks.5.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
673
+ "backbone.blocks.5.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
674
+ "backbone.blocks.5.attn.k_norm.weight": "model-00001-of-00004.safetensors",
675
+ "backbone.blocks.5.attn.proj.bias": "model-00001-of-00004.safetensors",
676
+ "backbone.blocks.5.attn.proj.weight": "model-00001-of-00004.safetensors",
677
+ "backbone.blocks.5.attn.q_norm.weight": "model-00001-of-00004.safetensors",
678
+ "backbone.blocks.5.attn.qkv.weight": "model-00001-of-00004.safetensors",
679
+ "backbone.blocks.5.ls1.gamma": "model-00001-of-00004.safetensors",
680
+ "backbone.blocks.5.ls2.gamma": "model-00001-of-00004.safetensors",
681
+ "backbone.blocks.5.mlp.w1.bias": "model-00001-of-00004.safetensors",
682
+ "backbone.blocks.5.mlp.w1.weight": "model-00001-of-00004.safetensors",
683
+ "backbone.blocks.5.mlp.w2.bias": "model-00001-of-00004.safetensors",
684
+ "backbone.blocks.5.mlp.w2.weight": "model-00001-of-00004.safetensors",
685
+ "backbone.blocks.5.mlp.w3.bias": "model-00001-of-00004.safetensors",
686
+ "backbone.blocks.5.mlp.w3.weight": "model-00001-of-00004.safetensors",
687
+ "backbone.blocks.5.norm1.bias": "model-00001-of-00004.safetensors",
688
+ "backbone.blocks.5.norm1.weight": "model-00001-of-00004.safetensors",
689
+ "backbone.blocks.5.norm2.bias": "model-00001-of-00004.safetensors",
690
+ "backbone.blocks.5.norm2.weight": "model-00001-of-00004.safetensors",
691
+ "backbone.blocks.6.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
692
+ "backbone.blocks.6.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
693
+ "backbone.blocks.6.attn.k_norm.weight": "model-00001-of-00004.safetensors",
694
+ "backbone.blocks.6.attn.proj.bias": "model-00001-of-00004.safetensors",
695
+ "backbone.blocks.6.attn.proj.weight": "model-00001-of-00004.safetensors",
696
+ "backbone.blocks.6.attn.q_norm.weight": "model-00001-of-00004.safetensors",
697
+ "backbone.blocks.6.attn.qkv.weight": "model-00001-of-00004.safetensors",
698
+ "backbone.blocks.6.ls1.gamma": "model-00001-of-00004.safetensors",
699
+ "backbone.blocks.6.ls2.gamma": "model-00001-of-00004.safetensors",
700
+ "backbone.blocks.6.mlp.w1.bias": "model-00001-of-00004.safetensors",
701
+ "backbone.blocks.6.mlp.w1.weight": "model-00001-of-00004.safetensors",
702
+ "backbone.blocks.6.mlp.w2.bias": "model-00001-of-00004.safetensors",
703
+ "backbone.blocks.6.mlp.w2.weight": "model-00001-of-00004.safetensors",
704
+ "backbone.blocks.6.mlp.w3.bias": "model-00001-of-00004.safetensors",
705
+ "backbone.blocks.6.mlp.w3.weight": "model-00001-of-00004.safetensors",
706
+ "backbone.blocks.6.norm1.bias": "model-00001-of-00004.safetensors",
707
+ "backbone.blocks.6.norm1.weight": "model-00001-of-00004.safetensors",
708
+ "backbone.blocks.6.norm2.bias": "model-00001-of-00004.safetensors",
709
+ "backbone.blocks.6.norm2.weight": "model-00001-of-00004.safetensors",
710
+ "backbone.blocks.7.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
711
+ "backbone.blocks.7.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
712
+ "backbone.blocks.7.attn.k_norm.weight": "model-00001-of-00004.safetensors",
713
+ "backbone.blocks.7.attn.proj.bias": "model-00001-of-00004.safetensors",
714
+ "backbone.blocks.7.attn.proj.weight": "model-00001-of-00004.safetensors",
715
+ "backbone.blocks.7.attn.q_norm.weight": "model-00001-of-00004.safetensors",
716
+ "backbone.blocks.7.attn.qkv.weight": "model-00001-of-00004.safetensors",
717
+ "backbone.blocks.7.ls1.gamma": "model-00001-of-00004.safetensors",
718
+ "backbone.blocks.7.ls2.gamma": "model-00001-of-00004.safetensors",
719
+ "backbone.blocks.7.mlp.w1.bias": "model-00001-of-00004.safetensors",
720
+ "backbone.blocks.7.mlp.w1.weight": "model-00001-of-00004.safetensors",
721
+ "backbone.blocks.7.mlp.w2.bias": "model-00001-of-00004.safetensors",
722
+ "backbone.blocks.7.mlp.w2.weight": "model-00001-of-00004.safetensors",
723
+ "backbone.blocks.7.mlp.w3.bias": "model-00001-of-00004.safetensors",
724
+ "backbone.blocks.7.mlp.w3.weight": "model-00001-of-00004.safetensors",
725
+ "backbone.blocks.7.norm1.bias": "model-00001-of-00004.safetensors",
726
+ "backbone.blocks.7.norm1.weight": "model-00001-of-00004.safetensors",
727
+ "backbone.blocks.7.norm2.bias": "model-00001-of-00004.safetensors",
728
+ "backbone.blocks.7.norm2.weight": "model-00001-of-00004.safetensors",
729
+ "backbone.blocks.8.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
730
+ "backbone.blocks.8.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
731
+ "backbone.blocks.8.attn.k_norm.weight": "model-00001-of-00004.safetensors",
732
+ "backbone.blocks.8.attn.proj.bias": "model-00001-of-00004.safetensors",
733
+ "backbone.blocks.8.attn.proj.weight": "model-00001-of-00004.safetensors",
734
+ "backbone.blocks.8.attn.q_norm.weight": "model-00001-of-00004.safetensors",
735
+ "backbone.blocks.8.attn.qkv.weight": "model-00001-of-00004.safetensors",
736
+ "backbone.blocks.8.ls1.gamma": "model-00001-of-00004.safetensors",
737
+ "backbone.blocks.8.ls2.gamma": "model-00001-of-00004.safetensors",
738
+ "backbone.blocks.8.mlp.w1.bias": "model-00001-of-00004.safetensors",
739
+ "backbone.blocks.8.mlp.w1.weight": "model-00001-of-00004.safetensors",
740
+ "backbone.blocks.8.mlp.w2.bias": "model-00001-of-00004.safetensors",
741
+ "backbone.blocks.8.mlp.w2.weight": "model-00001-of-00004.safetensors",
742
+ "backbone.blocks.8.mlp.w3.bias": "model-00001-of-00004.safetensors",
743
+ "backbone.blocks.8.mlp.w3.weight": "model-00001-of-00004.safetensors",
744
+ "backbone.blocks.8.norm1.bias": "model-00001-of-00004.safetensors",
745
+ "backbone.blocks.8.norm1.weight": "model-00001-of-00004.safetensors",
746
+ "backbone.blocks.8.norm2.bias": "model-00001-of-00004.safetensors",
747
+ "backbone.blocks.8.norm2.weight": "model-00001-of-00004.safetensors",
748
+ "backbone.blocks.9.attn.gate_proj.bias": "model-00001-of-00004.safetensors",
749
+ "backbone.blocks.9.attn.gate_proj.weight": "model-00001-of-00004.safetensors",
750
+ "backbone.blocks.9.attn.k_norm.weight": "model-00001-of-00004.safetensors",
751
+ "backbone.blocks.9.attn.proj.bias": "model-00001-of-00004.safetensors",
752
+ "backbone.blocks.9.attn.proj.weight": "model-00001-of-00004.safetensors",
753
+ "backbone.blocks.9.attn.q_norm.weight": "model-00001-of-00004.safetensors",
754
+ "backbone.blocks.9.attn.qkv.weight": "model-00001-of-00004.safetensors",
755
+ "backbone.blocks.9.ls1.gamma": "model-00001-of-00004.safetensors",
756
+ "backbone.blocks.9.ls2.gamma": "model-00001-of-00004.safetensors",
757
+ "backbone.blocks.9.mlp.w1.bias": "model-00001-of-00004.safetensors",
758
+ "backbone.blocks.9.mlp.w1.weight": "model-00001-of-00004.safetensors",
759
+ "backbone.blocks.9.mlp.w2.bias": "model-00001-of-00004.safetensors",
760
+ "backbone.blocks.9.mlp.w2.weight": "model-00001-of-00004.safetensors",
761
+ "backbone.blocks.9.mlp.w3.bias": "model-00001-of-00004.safetensors",
762
+ "backbone.blocks.9.mlp.w3.weight": "model-00001-of-00004.safetensors",
763
+ "backbone.blocks.9.norm1.bias": "model-00001-of-00004.safetensors",
764
+ "backbone.blocks.9.norm1.weight": "model-00001-of-00004.safetensors",
765
+ "backbone.blocks.9.norm2.bias": "model-00001-of-00004.safetensors",
766
+ "backbone.blocks.9.norm2.weight": "model-00001-of-00004.safetensors",
767
+ "backbone.cls_token": "model-00001-of-00004.safetensors",
768
+ "backbone.local_cls_norm.bias": "model-00004-of-00004.safetensors",
769
+ "backbone.local_cls_norm.weight": "model-00004-of-00004.safetensors",
770
+ "backbone.mask_token": "model-00001-of-00004.safetensors",
771
+ "backbone.norm.bias": "model-00004-of-00004.safetensors",
772
+ "backbone.norm.weight": "model-00004-of-00004.safetensors",
773
+ "backbone.patch_embed.proj.bias": "model-00001-of-00004.safetensors",
774
+ "backbone.patch_embed.proj.weight": "model-00001-of-00004.safetensors",
775
+ "backbone.rope_embed.periods_h": "model-00001-of-00004.safetensors",
776
+ "backbone.rope_embed.periods_t": "model-00001-of-00004.safetensors",
777
+ "backbone.rope_embed.periods_w": "model-00001-of-00004.safetensors",
778
+ "backbone.storage_tokens": "model-00001-of-00004.safetensors"
779
+ }
780
+ }
modeling_motif.py ADDED
@@ -0,0 +1,1671 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Motif Technologies.
2
+ # Self-contained inference model for MOTIF Vision Encoder (image + video).
3
+ # Auto-assembled from the training repo's inference path; NO training code.
4
+ """MOTIF Vision Encoder — unified image/video ViT backbone (inference-only).
5
+
6
+ Usage:
7
+ from transformers import AutoModel
8
+ import torch
9
+ model = AutoModel.from_pretrained("Motif-Technologies/motif-vision-encoder",
10
+ trust_remote_code=True).eval()
11
+ # image: (B, 3, H, W) video: (B, T, 3, H, W) (H,W multiples of 16)
12
+ out = model(pixel_values=torch.randn(1, 3, 224, 224))
13
+ out.last_hidden_state # (B, 1+num_register+N, D)
14
+ out.pooler_output # (B, D) CLS token
15
+ """
16
+ import math
17
+ from typing import Callable, Literal
18
+
19
+ import torch
20
+ import torch.nn as nn
21
+ import torch.nn.functional as F
22
+ from torch import Tensor
23
+
24
+ from transformers import PreTrainedModel, PretrainedConfig
25
+ from transformers.modeling_outputs import BaseModelOutputWithPooling
26
+
27
+
28
+
29
+ # ---- utils ----
30
+
31
+
32
+ def cat_keep_shapes(x_list: list[Tensor]) -> tuple[Tensor, list[tuple[int]], list[int]]:
33
+ """Concatenate list of tensors while preserving their shapes for later reconstruction."""
34
+ shapes = [x.shape for x in x_list]
35
+ num_tokens = [x.select(dim=-1, index=0).numel() for x in x_list]
36
+ flattened = torch.cat([x.flatten(0, -2) for x in x_list])
37
+ return flattened, shapes, num_tokens
38
+
39
+
40
+
41
+ def uncat_with_shapes(flattened: Tensor, shapes: list[tuple[int]], num_tokens: list[int]) -> list[Tensor]:
42
+ """Reverse of cat_keep_shapes: split and reshape flattened tensor back to original shapes."""
43
+ outputs_splitted = torch.split_with_sizes(flattened, num_tokens, dim=0)
44
+ shapes_adjusted = [shape[:-1] + torch.Size([flattened.shape[-1]]) for shape in shapes]
45
+ outputs_reshaped = [o.reshape(shape) for o, shape in zip(outputs_splitted, shapes_adjusted)]
46
+ return outputs_reshaped
47
+
48
+
49
+
50
+ def named_apply(
51
+ fn: Callable,
52
+ module: nn.Module,
53
+ name: str = "",
54
+ depth_first: bool = True,
55
+ include_root: bool = False,
56
+ ) -> nn.Module:
57
+ """Apply fn to all submodules (in-place, no replacement)."""
58
+ if not depth_first and include_root:
59
+ fn(module=module, name=name)
60
+ for child_name, child_module in module.named_children():
61
+ child_name = ".".join((name, child_name)) if name else child_name
62
+ named_apply(
63
+ fn=fn,
64
+ module=child_module,
65
+ name=child_name,
66
+ depth_first=depth_first,
67
+ include_root=True,
68
+ )
69
+ if depth_first and include_root:
70
+ fn(module=module, name=name)
71
+ return module
72
+
73
+
74
+ # ---- rms_norm ----
75
+
76
+
77
+ import torch
78
+ from torch import Tensor, nn
79
+
80
+
81
+ class RMSNorm(nn.Module):
82
+ """Root Mean Square Layer Normalization.
83
+
84
+ A simpler alternative to LayerNorm that normalizes by RMS without centering.
85
+
86
+ Args:
87
+ dim: Number of features.
88
+ eps: Small constant for numerical stability.
89
+ """
90
+
91
+ def __init__(
92
+ self,
93
+ dim: int,
94
+ eps: float = 1e-6,
95
+ device: torch.device | str | None = None,
96
+ ) -> None:
97
+ super().__init__()
98
+ self.eps = eps
99
+ self.weight = nn.Parameter(torch.ones(dim, device=device))
100
+
101
+ def reset_parameters(self) -> None:
102
+ """Reset weight to ones."""
103
+ nn.init.ones_(self.weight)
104
+
105
+ def forward(self, x: Tensor) -> Tensor:
106
+ """Apply RMS normalization."""
107
+ rms = torch.sqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
108
+ return x / rms * self.weight
109
+
110
+
111
+ # ---- layer_scale ----
112
+
113
+
114
+ import torch
115
+ from torch import Tensor, nn
116
+
117
+
118
+ class LayerScale(nn.Module):
119
+ """Per-channel scaling that allows gradual incorporation of each layer's contribution.
120
+
121
+ Initializes to a small value (e.g., 1e-5) so that early in training, each layer's
122
+ contribution is nearly zero, stabilizing deep network training.
123
+
124
+ Args:
125
+ dim: Number of channels.
126
+ init_values: Initial value for all channels.
127
+ inplace: Whether to apply scaling in-place.
128
+ device: Device for parameter allocation.
129
+ """
130
+
131
+ def __init__(
132
+ self,
133
+ dim: int,
134
+ init_values: float | Tensor = 1e-5,
135
+ inplace: bool = False,
136
+ device: torch.device | None = None,
137
+ ) -> None:
138
+ super().__init__()
139
+ self.inplace = inplace
140
+ self.gamma = nn.Parameter(torch.empty(dim, device=device))
141
+ self.init_values = init_values
142
+
143
+ def reset_parameters(self) -> None:
144
+ """Reset gamma to initial values."""
145
+ nn.init.constant_(self.gamma, self.init_values)
146
+
147
+ def forward(self, x: Tensor) -> Tensor:
148
+ """Apply per-channel scaling."""
149
+ return x.mul_(self.gamma) if self.inplace else x * self.gamma
150
+
151
+
152
+ # ---- patch_embed ----
153
+
154
+
155
+ import math
156
+
157
+ import torch.nn as nn
158
+ from torch import Tensor
159
+
160
+
161
+ class PatchEmbed(nn.Module):
162
+ """Video (5D) or Image (4D) to Patch Embedding via 3D Convolution.
163
+
164
+ Handles both modalities through a single Conv3d projection:
165
+ - Image (B, C, H, W): unsqueeze temporal dim -> (B, C, 1, H, W) -> Conv3d
166
+ - Video (B, T, C, H, W): transpose -> (B, C, T, H, W) -> Conv3d
167
+
168
+ Output: (B, N_total, embed_dim)
169
+ N_total = (T // tubelet_size) * (H // patch_size) * (W // patch_size)
170
+
171
+ Args:
172
+ img_size: Input image size (used for reference only).
173
+ patch_size: Spatial patch size in pixels.
174
+ in_chans: Number of input channels.
175
+ embed_dim: Output embedding dimension.
176
+ tubelet_size: Temporal patch size (number of frames per temporal token).
177
+ flatten_embedding: Whether to flatten spatial dimensions.
178
+ """
179
+
180
+ def __init__(
181
+ self,
182
+ img_size: int = 224,
183
+ patch_size: int = 16,
184
+ in_chans: int = 3,
185
+ embed_dim: int = 768,
186
+ tubelet_size: int = 1,
187
+ flatten_embedding: bool = True,
188
+ ) -> None:
189
+ super().__init__()
190
+ self.img_size = img_size
191
+ self.patch_size = (patch_size, patch_size) if isinstance(patch_size, int) else patch_size
192
+ self.tubelet_size = tubelet_size
193
+ self.flatten_embedding = flatten_embedding
194
+ self.in_chans = in_chans
195
+
196
+ # 3D Convolution: kernel and stride = (tubelet_size, patch_h, patch_w)
197
+ self.proj = nn.Conv3d(
198
+ in_chans,
199
+ embed_dim,
200
+ kernel_size=(tubelet_size, self.patch_size[0], self.patch_size[1]),
201
+ stride=(tubelet_size, self.patch_size[0], self.patch_size[1]),
202
+ )
203
+
204
+ def forward(self, x: Tensor) -> Tensor:
205
+ """Tokenize input images or videos.
206
+
207
+ Args:
208
+ x: Input tensor.
209
+ Image: (B, C, H, W) or Video: (B, T, C, H, W)
210
+
211
+ Returns:
212
+ Patch tokens of shape (B, N_total, embed_dim).
213
+ """
214
+ if x.ndim == 4:
215
+ # Image: (B, C, H, W) -> (B, C, 1, H, W)
216
+ x = x.unsqueeze(2)
217
+ # If tubelet_size > 1, repeat the single frame to match kernel size
218
+ if self.tubelet_size > 1:
219
+ x = x.expand(-1, -1, self.tubelet_size, -1, -1)
220
+ elif x.ndim == 5:
221
+ # Video: (B, T, C, H, W) -> (B, C, T, H, W)
222
+ x = x.transpose(1, 2)
223
+
224
+ # Conv3d Projection -> (B, embed_dim, T', H', W')
225
+ x = self.proj(x)
226
+
227
+ if self.flatten_embedding:
228
+ # Flatten spatial+temporal: (B, embed_dim, N) -> (B, N, embed_dim)
229
+ x = x.flatten(2).transpose(1, 2)
230
+
231
+ return x
232
+
233
+ def reset_parameters(self) -> None:
234
+ """Reset Conv3d parameters with uniform initialization."""
235
+ k = 1 / (self.in_chans * self.tubelet_size * self.patch_size[0] * self.patch_size[1])
236
+ nn.init.uniform_(self.proj.weight, -math.sqrt(k), math.sqrt(k))
237
+ if self.proj.bias is not None:
238
+ nn.init.uniform_(self.proj.bias, -math.sqrt(k), math.sqrt(k))
239
+
240
+
241
+ # ---- rope ----
242
+
243
+
244
+ import math
245
+ from typing import Literal
246
+
247
+ import numpy as np
248
+ import torch
249
+ from torch import Tensor, nn
250
+
251
+
252
+ class RopePositionEmbedding3D(nn.Module):
253
+ """Full 3D axial RoPE with independent T/H/W frequency bands.
254
+
255
+ Unlike the original MOTIF implementation which simply repeats 2D spatial angles
256
+ across temporal frames (making temporal positions indistinguishable), this
257
+ implementation partitions the head dimension into three axis groups:
258
+
259
+ D_head = D_T + D_H + D_W (no spare dimensions)
260
+
261
+ By default, the split is spatial-heavy for SSL (spatial quality is priority):
262
+ D_T = D_head // 4 (25% temporal)
263
+ D_H = (D_head - D_T) // 2 (37.5% height)
264
+ D_W = D_head - D_T - D_H (37.5% width)
265
+ e.g., D_head=64 → T=16, H=24, W=24
266
+
267
+ This can be overridden via ``fhw_dim=(D_T, D_H, D_W)`` for full control.
268
+
269
+ For images (T=1): t=0 for all tokens, making temporal angles constant
270
+ and the output is equivalent to spatial-only RoPE.
271
+
272
+ Args:
273
+ embed_dim: Total embedding dimension.
274
+ num_heads: Number of attention heads.
275
+ fhw_dim: Optional explicit (D_T, D_H, D_W) partition. Each must be even.
276
+ If None, uses the spatial-heavy default described above.
277
+ base: Frequency base (100.0 for spatial vision convention).
278
+ min_period: Minimum period (alternative to base).
279
+ max_period: Maximum period (alternative to base).
280
+ normalize_coords: How to normalize coordinates.
281
+ shift_coords: Random shift range during training.
282
+ jitter_coords: Random jitter multiplier during training.
283
+ rescale_coords: Random rescale multiplier during training.
284
+ dtype: Data type for computation.
285
+ device: Device for parameter allocation.
286
+ """
287
+
288
+ def __init__(
289
+ self,
290
+ embed_dim: int,
291
+ *,
292
+ num_heads: int,
293
+ fhw_dim: tuple[int, int, int] | None = None,
294
+ base: float | None = 100.0,
295
+ min_period: float | None = None,
296
+ max_period: float | None = None,
297
+ normalize_coords: Literal["min", "max", "separate"] = "separate",
298
+ shift_coords: float | None = None,
299
+ jitter_coords: float | None = None,
300
+ rescale_coords: float | None = None,
301
+ dtype: torch.dtype | None = None,
302
+ device: torch.device | None = None,
303
+ ) -> None:
304
+ super().__init__()
305
+ both_periods = min_period is not None and max_period is not None
306
+ if (base is None and not both_periods) or (base is not None and both_periods):
307
+ raise ValueError("Either `base` or `min_period`+`max_period` must be provided.")
308
+
309
+ D_head = embed_dim // num_heads
310
+ self.base = base
311
+ self.min_period = min_period
312
+ self.max_period = max_period
313
+ self.D_head = D_head
314
+ self.normalize_coords = normalize_coords
315
+ self.shift_coords = shift_coords
316
+ self.jitter_coords = jitter_coords
317
+ self.rescale_coords = rescale_coords
318
+
319
+ # Partition head dimension into 3 groups: T, H, W (no spare)
320
+ if fhw_dim is not None:
321
+ self.D_T, self.D_H, self.D_W = fhw_dim
322
+ assert self.D_T + self.D_H + self.D_W == D_head, (
323
+ f"fhw_dim must sum to D_head={D_head}, got {sum(fhw_dim)}"
324
+ )
325
+ else:
326
+ # Default: spatial-heavy split (SSL prioritizes spatial quality)
327
+ self.D_T = D_head // 4 # 25% temporal
328
+ self.D_H = (D_head - self.D_T) // 2 # 37.5% height
329
+ self.D_W = D_head - self.D_T - self.D_H # 37.5% width
330
+ assert self.D_T % 2 == 0 and self.D_H % 2 == 0 and self.D_W % 2 == 0, (
331
+ f"All axis dims must be even, got T={self.D_T}, H={self.D_H}, W={self.D_W}"
332
+ )
333
+
334
+ self.dtype = dtype
335
+ # Separate period buffers for each axis (n_freqs = D_axis // 2)
336
+ self.register_buffer(
337
+ "periods_t",
338
+ torch.empty(self.D_T // 2, device=device, dtype=dtype),
339
+ persistent=True,
340
+ )
341
+ self.register_buffer(
342
+ "periods_h",
343
+ torch.empty(self.D_H // 2, device=device, dtype=dtype),
344
+ persistent=True,
345
+ )
346
+ self.register_buffer(
347
+ "periods_w",
348
+ torch.empty(self.D_W // 2, device=device, dtype=dtype),
349
+ persistent=True,
350
+ )
351
+ self._init_weights()
352
+
353
+ def forward(self, *, T: int = 1, H: int, W: int) -> tuple[Tensor, Tensor]:
354
+ """Compute 3D axial RoPE sin/cos for (T, H, W) grid.
355
+
356
+ The head dimension is partitioned as [D_T | D_H | D_W]:
357
+ - D_T: temporal frequency bands (angles vary with t)
358
+ - D_H: height frequency bands (angles vary with h)
359
+ - D_W: width frequency bands (angles vary with w)
360
+
361
+ For images (T=1), all tokens get t=0, so temporal angles are constant
362
+ and the output is equivalent to spatial-only RoPE.
363
+
364
+ Args:
365
+ T: Number of temporal positions (T_grid = num_frames // tubelet_size).
366
+ H: Height in patches.
367
+ W: Width in patches.
368
+
369
+ Returns:
370
+ Tuple of (sin, cos), each of shape (T*H*W, D_head).
371
+ """
372
+ device = self.periods_t.device
373
+ dtype = self.dtype
374
+ dd = {"device": device, "dtype": dtype}
375
+
376
+ # 1. Compute normalized coordinates for each axis
377
+ if T > 1:
378
+ coords_t = torch.arange(0.5, T, **dd) / T # [T]
379
+ else:
380
+ coords_t = torch.tensor([0.5], **dd) # [1] - constant for images
381
+
382
+ coords_h, coords_w = self._compute_spatial_coords(H, W, **dd)
383
+
384
+ # Shift to [-1, +1] range
385
+ coords_t = 2.0 * coords_t - 1.0 # [T]
386
+ coords_h = 2.0 * coords_h - 1.0 # [H]
387
+ coords_w = 2.0 * coords_w - 1.0 # [W]
388
+
389
+ # Apply training-time augmentations to spatial coords only
390
+ if self.training:
391
+ coords_h, coords_w = self._augment_spatial_coords(coords_h, coords_w, dd)
392
+
393
+ # 2. Compute raw angles for each axis (n_freqs = D_axis // 2)
394
+ angles_t = 2 * math.pi * coords_t[:, None] / self.periods_t[None, :] # [T, D_T//2]
395
+ angles_h = 2 * math.pi * coords_h[:, None] / self.periods_h[None, :] # [H, D_H//2]
396
+ angles_w = 2 * math.pi * coords_w[:, None] / self.periods_w[None, :] # [W, D_W//2]
397
+
398
+ # 3. Build full 3D grid: create (T*H*W, D_head) angle tensor
399
+ t_idx, h_idx, w_idx = torch.meshgrid(
400
+ torch.arange(T, device=device),
401
+ torch.arange(H, device=device),
402
+ torch.arange(W, device=device),
403
+ indexing="ij",
404
+ )
405
+ t_idx = t_idx.flatten() # [T*H*W]
406
+ h_idx = h_idx.flatten() # [T*H*W]
407
+ w_idx = w_idx.flatten() # [T*H*W]
408
+
409
+ # Gather per-token raw angles and concatenate to D_head//2
410
+ token_angles_t = angles_t[t_idx] # [T*H*W, D_T//2]
411
+ token_angles_h = angles_h[h_idx] # [T*H*W, D_H//2]
412
+ token_angles_w = angles_w[w_idx] # [T*H*W, D_W//2]
413
+ angles_half = torch.cat([token_angles_t, token_angles_h, token_angles_w], dim=-1) # [T*H*W, D_head//2]
414
+
415
+ # tile(2) on full concat — matches MOTIF 2D RoPE pattern
416
+ # This ensures rotate_half pairs (dim i ↔ dim i+D//2) have identical angles,
417
+ # making the rotation orthogonal (preserves dot products in attention).
418
+ angles = angles_half.tile(2) # [T*H*W, D_head]
419
+
420
+ cos = torch.cos(angles)
421
+ sin = torch.sin(angles)
422
+
423
+ return (sin, cos)
424
+
425
+ def _compute_spatial_coords(self, H: int, W: int, **dd) -> tuple[Tensor, Tensor]:
426
+ """Compute normalized spatial coordinates."""
427
+ if self.normalize_coords == "max":
428
+ max_HW = max(H, W)
429
+ coords_h = torch.arange(0.5, H, **dd) / max_HW
430
+ coords_w = torch.arange(0.5, W, **dd) / max_HW
431
+ elif self.normalize_coords == "min":
432
+ min_HW = min(H, W)
433
+ coords_h = torch.arange(0.5, H, **dd) / min_HW
434
+ coords_w = torch.arange(0.5, W, **dd) / min_HW
435
+ elif self.normalize_coords == "separate":
436
+ coords_h = torch.arange(0.5, H, **dd) / H
437
+ coords_w = torch.arange(0.5, W, **dd) / W
438
+ else:
439
+ raise ValueError(f"Unknown normalize_coords: {self.normalize_coords}")
440
+ return coords_h, coords_w
441
+
442
+ def _augment_spatial_coords(
443
+ self,
444
+ coords_h: Tensor,
445
+ coords_w: Tensor,
446
+ dd: dict,
447
+ ) -> tuple[Tensor, Tensor]:
448
+ """Apply training-time coordinate augmentations to spatial coords."""
449
+ if self.shift_coords is not None:
450
+ shift = torch.empty(2, **dd).uniform_(-self.shift_coords, self.shift_coords)
451
+ coords_h = coords_h + shift[0]
452
+ coords_w = coords_w + shift[1]
453
+ if self.jitter_coords is not None:
454
+ jitter_max = np.log(self.jitter_coords)
455
+ jitter = torch.empty(2, **dd).uniform_(-jitter_max, jitter_max).exp()
456
+ coords_h = coords_h * jitter[0]
457
+ coords_w = coords_w * jitter[1]
458
+ if self.rescale_coords is not None:
459
+ rescale_max = np.log(self.rescale_coords)
460
+ rescale = torch.empty(1, **dd).uniform_(-rescale_max, rescale_max).exp()
461
+ coords_h = coords_h * rescale
462
+ coords_w = coords_w * rescale
463
+ return coords_h, coords_w
464
+
465
+ def _compute_periods(self, n_freqs: int, device: torch.device, dtype: torch.dtype | None) -> Tensor:
466
+ """Compute frequency periods for a single axis.
467
+
468
+ Args:
469
+ n_freqs: Number of frequency bands (D_axis // 2).
470
+ device: Device for tensor allocation.
471
+ dtype: Data type for computation.
472
+
473
+ Returns:
474
+ Tensor of shape (n_freqs,) with logarithmically spaced periods.
475
+ """
476
+ if self.base is not None:
477
+ return self.base ** (
478
+ 2 * torch.arange(n_freqs, device=device, dtype=dtype) / (2 * n_freqs)
479
+ )
480
+ else:
481
+ base = self.max_period / self.min_period
482
+ exponents = torch.linspace(0, 1, n_freqs, device=device, dtype=dtype)
483
+ periods = base**exponents
484
+ periods = periods / base
485
+ return periods * self.max_period
486
+
487
+ def _init_weights(self) -> None:
488
+ """Initialize frequency periods for all three axes.
489
+
490
+ Each axis gets its own frequency schedule based on its dimension size:
491
+ periods[i] = base^(2i / D_axis)
492
+
493
+ This produces logarithmically spaced periods from 1.0 to base,
494
+ with more frequencies for axes with more allocated dimensions.
495
+ """
496
+ device = self.periods_t.device
497
+ dtype = self.dtype
498
+
499
+ self.periods_t.data = self._compute_periods(self.D_T // 2, device, dtype)
500
+ self.periods_h.data = self._compute_periods(self.D_H // 2, device, dtype)
501
+ self.periods_w.data = self._compute_periods(self.D_W // 2, device, dtype)
502
+
503
+
504
+ # ---- attention ----
505
+
506
+
507
+ import math
508
+
509
+ import torch
510
+ import torch.nn.functional as F
511
+ from torch import Tensor, nn
512
+
513
+
514
+
515
+ def rope_rotate_half(x: Tensor) -> Tensor:
516
+ """Rotate half of the dimensions: [-x2, x1] from [x1, x2].
517
+
518
+ Args:
519
+ x: Input tensor of shape (..., D).
520
+
521
+ Returns:
522
+ Rotated tensor of shape (..., D).
523
+ """
524
+ x1, x2 = x.chunk(2, dim=-1)
525
+ return torch.cat([-x2, x1], dim=-1)
526
+
527
+
528
+ def rope_apply(x: Tensor, sin: Tensor, cos: Tensor) -> Tensor:
529
+ """Apply rotary position embedding to input tensor.
530
+
531
+ Args:
532
+ x: Input tensor of shape (..., D).
533
+ sin: Sine angles of shape (..., D).
534
+ cos: Cosine angles of shape (..., D).
535
+
536
+ Returns:
537
+ Rotated tensor of shape (..., D).
538
+ """
539
+ return (x * cos) + (rope_rotate_half(x) * sin)
540
+
541
+
542
+ class LinearKMaskedBias(nn.Linear):
543
+ """Linear layer with masked bias for the K component of QKV.
544
+
545
+ Zeroes out the bias for the K component (middle third of output)
546
+ to avoid interference with RoPE positional encoding.
547
+ """
548
+
549
+ def __init__(self, *args, **kwargs) -> None:
550
+ super().__init__(*args, **kwargs)
551
+ o = self.out_features
552
+ assert o % 3 == 0
553
+ if self.bias is not None:
554
+ self.register_buffer("bias_mask", torch.full_like(self.bias, fill_value=math.nan))
555
+
556
+ def forward(self, input: Tensor) -> Tensor:
557
+ """Forward pass with masked bias."""
558
+ masked_bias = self.bias * self.bias_mask.to(self.bias.dtype) if self.bias is not None else None
559
+ return F.linear(input, self.weight, masked_bias)
560
+
561
+
562
+ class SelfAttention(nn.Module):
563
+ """Multi-head self-attention with RoPE support.
564
+
565
+ Uses torch.nn.functional.scaled_dot_product_attention for FlashAttention
566
+ compatibility. RoPE is applied to Q and K on patch tokens only (not CLS/register).
567
+
568
+ Args:
569
+ dim: Model dimension.
570
+ num_heads: Number of attention heads.
571
+ qkv_bias: Whether to use bias in QKV projection.
572
+ proj_bias: Whether to use bias in output projection.
573
+ attn_drop: Attention dropout probability.
574
+ proj_drop: Output projection dropout probability.
575
+ mask_k_bias: Whether to mask K bias (for RoPE compatibility).
576
+ device: Device for parameter allocation.
577
+ gated_attention: Gated attention variant. None disables gating,
578
+ "headwise" applies a per-head scalar gate, "elementwise" applies
579
+ a per-element gate. Gate scores are query-dependent (derived from
580
+ input) and applied as sigmoid after SDPA.
581
+ Reference: https://arxiv.org/abs/2505.06708
582
+ """
583
+
584
+ def __init__(
585
+ self,
586
+ dim: int,
587
+ num_heads: int = 8,
588
+ qkv_bias: bool = False,
589
+ proj_bias: bool = True,
590
+ attn_drop: float = 0.0,
591
+ proj_drop: float = 0.0,
592
+ mask_k_bias: bool = False,
593
+ device: str | None = None,
594
+ gated_attention: str | None = None,
595
+ qk_norm: bool = False,
596
+ ) -> None:
597
+ super().__init__()
598
+ self.num_heads = num_heads
599
+ self.head_dim = dim // num_heads
600
+ self.scale = self.head_dim**-0.5
601
+
602
+ linear_class = LinearKMaskedBias if mask_k_bias else nn.Linear
603
+ self.qkv = linear_class(dim, dim * 3, bias=qkv_bias, device=device)
604
+ self.attn_drop = nn.Dropout(attn_drop)
605
+ self.proj = nn.Linear(dim, dim, bias=proj_bias, device=device)
606
+ self.proj_drop = nn.Dropout(proj_drop)
607
+
608
+ self.qk_norm = qk_norm
609
+ if qk_norm:
610
+ self.q_norm = RMSNorm(self.head_dim, device=device)
611
+ self.k_norm = RMSNorm(self.head_dim, device=device)
612
+
613
+ self.gated_attention = gated_attention
614
+ if gated_attention == "headwise":
615
+ self.gate_proj = nn.Linear(dim, num_heads, bias=True, device=device)
616
+ elif gated_attention == "elementwise":
617
+ self.gate_proj = nn.Linear(dim, dim, bias=True, device=device)
618
+ elif gated_attention is not None:
619
+ raise ValueError(f"Unknown gated_attention mode: {gated_attention!r}. Use 'headwise' or 'elementwise'.")
620
+
621
+ def apply_rope(
622
+ self,
623
+ q: Tensor,
624
+ k: Tensor,
625
+ rope: tuple[Tensor, Tensor],
626
+ ) -> tuple[Tensor, Tensor]:
627
+ """Apply RoPE to query and key tensors.
628
+
629
+ RoPE is applied only to patch tokens (prefix tokens like CLS and register
630
+ are excluded based on the difference between sequence length and rope length).
631
+
632
+ Args:
633
+ q: Query tensor of shape (B, heads, N, D_head).
634
+ k: Key tensor of shape (B, heads, N, D_head).
635
+ rope: Tuple of (sin, cos), each of shape (N_patches, D_head).
636
+
637
+ Returns:
638
+ Tuple of rotated (q, k) tensors.
639
+ """
640
+ q_dtype = q.dtype
641
+ k_dtype = k.dtype
642
+ sin, cos = rope
643
+ rope_dtype = sin.dtype
644
+ q = q.to(dtype=rope_dtype)
645
+ k = k.to(dtype=rope_dtype)
646
+ N = q.shape[-2]
647
+ prefix = N - sin.shape[-2]
648
+ assert prefix >= 0
649
+ q_prefix = q[:, :, :prefix, :]
650
+ q = rope_apply(q[:, :, prefix:, :], sin, cos)
651
+ q = torch.cat((q_prefix, q), dim=-2)
652
+ k_prefix = k[:, :, :prefix, :]
653
+ k = rope_apply(k[:, :, prefix:, :], sin, cos)
654
+ k = torch.cat((k_prefix, k), dim=-2)
655
+ q = q.to(dtype=q_dtype)
656
+ k = k.to(dtype=k_dtype)
657
+ return q, k
658
+
659
+ def forward(self, x: Tensor, attn_bias: Tensor | None = None, rope: Tensor | None = None) -> Tensor:
660
+ """Forward pass for single tensor input.
661
+
662
+ Args:
663
+ x: Input tensor of shape (B, N, D).
664
+ attn_bias: Unused (kept for interface compatibility).
665
+ rope: Optional RoPE (sin, cos) tuple.
666
+
667
+ Returns:
668
+ Output tensor of shape (B, N, D).
669
+ """
670
+ gate_score = self._compute_gate(x) if self.gated_attention else None
671
+ qkv = self.qkv(x)
672
+ attn_v = self.compute_attention(qkv=qkv, attn_bias=attn_bias, rope=rope, gate_score=gate_score)
673
+ x = self.proj(attn_v)
674
+ x = self.proj_drop(x)
675
+ return x
676
+
677
+ def forward_list(
678
+ self,
679
+ x_list: list[Tensor],
680
+ attn_bias: Tensor | None = None,
681
+ rope_list: list[tuple[Tensor, Tensor]] | None = None,
682
+ ) -> list[Tensor]:
683
+ """Forward pass for list of tensors (multi-crop efficiency).
684
+
685
+ Concatenates inputs for a single QKV projection, then splits for per-crop
686
+ attention computation (needed because different crops have different RoPE).
687
+
688
+ Args:
689
+ x_list: List of input tensors.
690
+ attn_bias: Unused.
691
+ rope_list: List of RoPE (sin, cos) tuples, one per input.
692
+
693
+ Returns:
694
+ List of output tensors.
695
+ """
696
+ assert len(x_list) == len(rope_list)
697
+ x_flat, shapes, num_tokens = cat_keep_shapes(x_list)
698
+ qkv_flat = self.qkv(x_flat)
699
+ qkv_list = uncat_with_shapes(qkv_flat, shapes, num_tokens)
700
+
701
+ if self.gated_attention:
702
+ gate_flat = self._compute_gate(x_flat)
703
+ gate_list = uncat_with_shapes(gate_flat, shapes, num_tokens)
704
+ else:
705
+ gate_list = [None] * len(x_list)
706
+
707
+ att_out = []
708
+ for qkv, _, rope, gate_score in zip(qkv_list, shapes, rope_list, gate_list):
709
+ att_out.append(self.compute_attention(qkv, attn_bias=attn_bias, rope=rope, gate_score=gate_score))
710
+ x_flat, shapes, num_tokens = cat_keep_shapes(att_out)
711
+ x_flat = self.proj(x_flat)
712
+ return uncat_with_shapes(x_flat, shapes, num_tokens)
713
+
714
+ def _compute_gate(self, x: Tensor) -> Tensor:
715
+ """Compute raw gate scores from input.
716
+
717
+ Returns the raw projection without reshaping so that the output keeps
718
+ the same number of leading dimensions as ``x``. This is critical for
719
+ ``forward_list`` where ``uncat_with_shapes`` must split a 2-D flat
720
+ tensor back to per-crop 3-D tensors — adding extra dims here would
721
+ break that reshape. The per-head unflatten happens later inside
722
+ ``compute_attention`` where B and N are known.
723
+
724
+ Args:
725
+ x: Input tensor of shape (..., D). Supports both 2D (flat) and 3D (batched).
726
+
727
+ Returns:
728
+ Raw gate projection. Headwise: (..., num_heads). Elementwise: (..., D).
729
+ """
730
+ return self.gate_proj(x)
731
+
732
+ def compute_attention(
733
+ self,
734
+ qkv: Tensor,
735
+ attn_bias: Tensor | None = None,
736
+ rope: tuple[Tensor, Tensor] | None = None,
737
+ gate_score: Tensor | None = None,
738
+ ) -> Tensor:
739
+ """Compute scaled dot-product attention.
740
+
741
+ Args:
742
+ qkv: Combined QKV tensor of shape (B, N, 3*D).
743
+ attn_bias: Unused.
744
+ rope: Optional RoPE (sin, cos) tuple.
745
+ gate_score: Optional gate tensor from _compute_gate.
746
+
747
+ Returns:
748
+ Attention output of shape (B, N, D).
749
+ """
750
+ assert attn_bias is None
751
+ B, N, _ = qkv.shape
752
+ C = self.qkv.in_features
753
+
754
+ qkv = qkv.reshape(B, N, 3, self.num_heads, self.head_dim)
755
+ q, k, v = torch.unbind(qkv, 2)
756
+ q, k, v = [t.transpose(1, 2) for t in [q, k, v]]
757
+ if self.qk_norm:
758
+ q = self.q_norm(q)
759
+ k = self.k_norm(k)
760
+ if rope is not None:
761
+ q, k = self.apply_rope(q, k, rope)
762
+ x = torch.nn.functional.scaled_dot_product_attention(q, k, v)
763
+ x = x.transpose(1, 2) # (B, N, num_heads, head_dim)
764
+ if gate_score is not None:
765
+ # _compute_gate returns raw projection: (..., num_heads) or (..., D).
766
+ # Reshape to (B, N, num_heads, 1) or (B, N, num_heads, head_dim) here.
767
+ if self.gated_attention == "headwise":
768
+ gate_score = gate_score.unflatten(-1, (self.num_heads, 1))
769
+ else: # elementwise
770
+ gate_score = gate_score.unflatten(-1, (self.num_heads, self.head_dim))
771
+ x = x * torch.sigmoid(gate_score)
772
+ return x.reshape([B, N, C])
773
+
774
+
775
+ # ---- ffn ----
776
+
777
+
778
+ from typing import Callable
779
+
780
+ import torch.nn.functional as F
781
+ from torch import Tensor, nn
782
+
783
+
784
+
785
+ class ListForwardMixin:
786
+ """Mixin providing forward_list for efficient multi-crop processing."""
787
+
788
+ def forward(self, x: Tensor) -> Tensor:
789
+ """Forward pass for a single tensor."""
790
+ raise NotImplementedError
791
+
792
+ def forward_list(self, x_list: list[Tensor]) -> list[Tensor]:
793
+ """Forward pass for a list of tensors, concatenated for efficiency."""
794
+ x_flat, shapes, num_tokens = cat_keep_shapes(x_list)
795
+ x_flat = self.forward(x_flat)
796
+ return uncat_with_shapes(x_flat, shapes, num_tokens)
797
+
798
+
799
+ class SwiGLUFFN(nn.Module, ListForwardMixin):
800
+ """SwiGLU Feed-Forward Network: w3(silu(w1(x)) * w2(x)).
801
+
802
+ Used for larger ViT models (SO400M+) due to better gradient flow.
803
+ Hidden dimension is aligned to a multiple of `align_to` for GPU efficiency.
804
+
805
+ Args:
806
+ in_features: Input dimension.
807
+ hidden_features: Hidden dimension before alignment.
808
+ out_features: Output dimension (default: same as in_features).
809
+ act_layer: Unused (SwiGLU has built-in SiLU activation).
810
+ drop: Unused (no dropout in SwiGLU).
811
+ bias: Whether to use bias in linear layers.
812
+ align_to: Align hidden dimension to this multiple.
813
+ device: Device for parameter allocation.
814
+ """
815
+
816
+ def __init__(
817
+ self,
818
+ in_features: int,
819
+ hidden_features: int | None = None,
820
+ out_features: int | None = None,
821
+ act_layer: Callable[..., nn.Module] | None = None,
822
+ drop: float = 0.0,
823
+ bias: bool = True,
824
+ align_to: int = 8,
825
+ device: str | None = None,
826
+ ) -> None:
827
+ super().__init__()
828
+ out_features = out_features or in_features
829
+ hidden_features = hidden_features or in_features
830
+ d = int(hidden_features * 2 / 3)
831
+ swiglu_hidden_features = d + (-d % align_to)
832
+ self.w1 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device)
833
+ self.w2 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device)
834
+ self.w3 = nn.Linear(swiglu_hidden_features, out_features, bias=bias, device=device)
835
+
836
+ def forward(self, x: Tensor) -> Tensor:
837
+ """Forward pass: w3(silu(w1(x)) * w2(x))."""
838
+ x1 = self.w1(x)
839
+ x2 = self.w2(x)
840
+ hidden = F.silu(x1) * x2
841
+ return self.w3(hidden)
842
+
843
+
844
+ # ---- block ----
845
+
846
+
847
+ from typing import Callable
848
+
849
+ import torch
850
+ from torch import Tensor, nn
851
+
852
+
853
+ torch._dynamo.config.automatic_dynamic_shapes = False
854
+ torch._dynamo.config.accumulated_cache_size_limit = 1024
855
+ # Per-code-object recompile cap (default 8). With dynamic=False + fullgraph=True,
856
+ # every distinct input shape forces a static recompile of the block. Gram adds the
857
+ # Gram-teacher forward at gram_teacher_crops_size (e.g. 768/1152) on top of the
858
+ # 512/768 image + video + local-crop shapes, pushing the distinct-shape count past
859
+ # 8 → FailOnRecompileLimitHit (hard crash, not the eager-fallback warning seen when
860
+ # fullgraph=False). Raise the cap so all Gram shapes compile; the accumulated cap
861
+ # (1024) still bounds total graphs. Only the shapes that actually occur are compiled,
862
+ # so a high cap costs nothing when Gram is off.
863
+ torch._dynamo.config.recompile_limit = 64
864
+
865
+
866
+
867
+ class SelfAttentionBlock(nn.Module):
868
+ """Pre-norm transformer block: Norm -> Attention -> LayerScale -> DropPath -> Residual.
869
+
870
+ Supports both single-tensor and list-of-tensors forward for efficient multi-crop
871
+ processing during MOTIF training.
872
+
873
+ Args:
874
+ dim: Model dimension.
875
+ num_heads: Number of attention heads.
876
+ ffn_ratio: FFN hidden dimension ratio.
877
+ qkv_bias: Whether to use bias in QKV projection.
878
+ proj_bias: Whether to use bias in output projection.
879
+ ffn_bias: Whether to use bias in FFN layers.
880
+ drop: Dropout probability.
881
+ attn_drop: Attention dropout probability.
882
+ init_values: LayerScale initial values (None disables LayerScale).
883
+ drop_path: Stochastic depth drop probability.
884
+ act_layer: Activation function class.
885
+ norm_layer: Normalization layer class.
886
+ attn_class: Attention class.
887
+ ffn_layer: FFN class.
888
+ mask_k_bias: Whether to mask K bias.
889
+ device: Device for parameter allocation.
890
+ """
891
+
892
+ def __init__(
893
+ self,
894
+ dim: int,
895
+ num_heads: int,
896
+ ffn_ratio: float = 4.0,
897
+ qkv_bias: bool = False,
898
+ proj_bias: bool = True,
899
+ ffn_bias: bool = True,
900
+ drop: float = 0.0,
901
+ attn_drop: float = 0.0,
902
+ init_values: float | None = None,
903
+ drop_path: float = 0.0,
904
+ act_layer: Callable[..., nn.Module] = nn.GELU,
905
+ norm_layer: Callable[..., nn.Module] = nn.LayerNorm,
906
+ attn_class: Callable[..., nn.Module] = SelfAttention,
907
+ ffn_layer: Callable[..., nn.Module] = SwiGLUFFN,
908
+ mask_k_bias: bool = False,
909
+ device: str | None = None,
910
+ gated_attention: str | None = None,
911
+ qk_norm: bool = False,
912
+ ) -> None:
913
+ super().__init__()
914
+ self.norm1 = norm_layer(dim)
915
+ self.attn = attn_class(
916
+ dim,
917
+ num_heads=num_heads,
918
+ qkv_bias=qkv_bias,
919
+ proj_bias=proj_bias,
920
+ attn_drop=attn_drop,
921
+ proj_drop=drop,
922
+ mask_k_bias=mask_k_bias,
923
+ device=device,
924
+ gated_attention=gated_attention,
925
+ qk_norm=qk_norm,
926
+ )
927
+ self.ls1 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity()
928
+
929
+ self.norm2 = norm_layer(dim)
930
+ mlp_hidden_dim = int(dim * ffn_ratio)
931
+ self.mlp = ffn_layer(
932
+ in_features=dim,
933
+ hidden_features=mlp_hidden_dim,
934
+ act_layer=act_layer,
935
+ drop=drop,
936
+ bias=ffn_bias,
937
+ device=device,
938
+ )
939
+ self.ls2 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity()
940
+
941
+ self.sample_drop_ratio = drop_path
942
+
943
+ @staticmethod
944
+ def _maybe_index_rope(
945
+ rope: tuple[Tensor, Tensor] | None,
946
+ indices: Tensor,
947
+ ) -> tuple[Tensor, Tensor] | None:
948
+ """Optionally index into RoPE embeddings for stochastic depth."""
949
+ if rope is None:
950
+ return None
951
+ sin, cos = rope
952
+ assert sin.ndim == cos.ndim
953
+ if sin.ndim == 4:
954
+ return sin[indices], cos[indices]
955
+ else:
956
+ return sin, cos
957
+
958
+ def _forward_list(self, x_list: list[Tensor], rope_list: list | None = None) -> list[Tensor]:
959
+ """Forward pass for list of tensors with stochastic depth support.
960
+
961
+ Concatenates tokens from multiple inputs for efficient elementwise operations,
962
+ then splits for per-input attention computation (different RoPE per crop).
963
+ """
964
+ b_list = [x.shape[0] for x in x_list]
965
+ sample_subset_sizes = [max(int(b * (1 - self.sample_drop_ratio)), 1) for b in b_list]
966
+ residual_scale_factors = [b / s for b, s in zip(b_list, sample_subset_sizes)]
967
+
968
+ if self.training and self.sample_drop_ratio > 0.0:
969
+ indices_1_list = [
970
+ (torch.randperm(b, device=x.device))[:s]
971
+ for x, b, s in zip(x_list, b_list, sample_subset_sizes)
972
+ ]
973
+ x_subset_1_list = [x[idx] for x, idx in zip(x_list, indices_1_list)]
974
+
975
+ if rope_list is not None:
976
+ rope_subset_list = [
977
+ self._maybe_index_rope(rope, idx) for rope, idx in zip(rope_list, indices_1_list)
978
+ ]
979
+ else:
980
+ rope_subset_list = rope_list
981
+
982
+ norm1 = [self.norm1(x) for x in x_subset_1_list]
983
+ residual_1_list = self.attn.forward_list(norm1, rope_list=rope_subset_list)
984
+
985
+ x_attn_list = [
986
+ torch.index_add(
987
+ x, dim=0, source=self.ls1(r), index=idx, alpha=scale,
988
+ )
989
+ for x, r, idx, scale in zip(x_list, residual_1_list, indices_1_list, residual_scale_factors)
990
+ ]
991
+
992
+ indices_2_list = [
993
+ (torch.randperm(b, device=x.device))[:s]
994
+ for x, b, s in zip(x_list, b_list, sample_subset_sizes)
995
+ ]
996
+ x_subset_2_list = [x[idx] for x, idx in zip(x_attn_list, indices_2_list)]
997
+ norm2_list = [self.norm2(x) for x in x_subset_2_list]
998
+
999
+ residual_2_list = self.mlp.forward_list(norm2_list)
1000
+
1001
+ x_ffn = [
1002
+ torch.index_add(
1003
+ xa, dim=0, source=self.ls2(r), index=idx, alpha=scale,
1004
+ )
1005
+ for xa, r, idx, scale in zip(x_attn_list, residual_2_list, indices_2_list, residual_scale_factors)
1006
+ ]
1007
+ else:
1008
+ x_out = []
1009
+ for x, rope in zip(x_list, rope_list):
1010
+ x_attn = x + self.ls1(self.attn(self.norm1(x), rope=rope))
1011
+ x_ffn_item = x_attn + self.ls2(self.mlp(self.norm2(x_attn)))
1012
+ x_out.append(x_ffn_item)
1013
+ x_ffn = x_out
1014
+
1015
+ return x_ffn
1016
+
1017
+ def forward(
1018
+ self,
1019
+ x_or_x_list: Tensor | list[Tensor],
1020
+ rope_or_rope_list: tuple | list | None = None,
1021
+ ) -> Tensor | list[Tensor]:
1022
+ """Forward pass accepting either a single tensor or list of tensors.
1023
+
1024
+ Args:
1025
+ x_or_x_list: Single tensor (B, N, D) or list of tensors.
1026
+ rope_or_rope_list: Single RoPE tuple or list of RoPE tuples.
1027
+
1028
+ Returns:
1029
+ Output tensor(s) matching input format.
1030
+ """
1031
+ if isinstance(x_or_x_list, Tensor):
1032
+ return self._forward_list([x_or_x_list], rope_list=[rope_or_rope_list])[0]
1033
+ elif isinstance(x_or_x_list, list):
1034
+ if rope_or_rope_list is None:
1035
+ rope_or_rope_list = [None for _ in x_or_x_list]
1036
+ return self._forward_list(x_or_x_list, rope_list=rope_or_rope_list)
1037
+ else:
1038
+ raise AssertionError(f"Unexpected input type: {type(x_or_x_list)}")
1039
+
1040
+
1041
+ # ---- vision_transformer ----
1042
+
1043
+
1044
+ import logging
1045
+ from functools import partial
1046
+ from typing import Any
1047
+
1048
+ import torch
1049
+ import torch.nn.init
1050
+ from torch import Tensor, nn
1051
+
1052
+
1053
+ logger = logging.getLogger("motif")
1054
+
1055
+ ffn_layer_dict: dict[str, type] = {
1056
+ "swiglu": SwiGLUFFN,
1057
+ "swiglu32": partial(SwiGLUFFN, align_to=32),
1058
+ "swiglu64": partial(SwiGLUFFN, align_to=64),
1059
+ "swiglu128": partial(SwiGLUFFN, align_to=128),
1060
+ }
1061
+
1062
+ norm_layer_dict: dict[str, type] = {
1063
+ "layernorm": partial(nn.LayerNorm, eps=1e-6),
1064
+ "layernormbf16": partial(nn.LayerNorm, eps=1e-5),
1065
+ "rmsnorm": RMSNorm,
1066
+ }
1067
+
1068
+ dtype_dict: dict[str, torch.dtype] = {
1069
+ "fp32": torch.float32,
1070
+ "fp16": torch.float16,
1071
+ "bf16": torch.bfloat16,
1072
+ }
1073
+
1074
+
1075
+ def init_weights_vit(module: nn.Module, name: str = "") -> None:
1076
+ """Initialize weights for ViT components.
1077
+
1078
+ Applied recursively via named_apply to all submodules.
1079
+
1080
+ Args:
1081
+ module: Module to initialize.
1082
+ name: Module name (for logging, unused).
1083
+ """
1084
+ if isinstance(module, nn.Linear):
1085
+ torch.nn.init.trunc_normal_(module.weight, std=0.02)
1086
+ if module.bias is not None:
1087
+ nn.init.zeros_(module.bias)
1088
+ if hasattr(module, "bias_mask") and module.bias_mask is not None:
1089
+ o = module.out_features
1090
+ module.bias_mask.fill_(1)
1091
+ module.bias_mask[o // 3 : 2 * o // 3].fill_(0)
1092
+ if isinstance(module, nn.LayerNorm):
1093
+ module.reset_parameters()
1094
+ if isinstance(module, LayerScale):
1095
+ module.reset_parameters()
1096
+ if isinstance(module, PatchEmbed):
1097
+ module.reset_parameters()
1098
+ if isinstance(module, RMSNorm):
1099
+ module.reset_parameters()
1100
+
1101
+
1102
+ class MotifVisionTransformer(nn.Module):
1103
+ """Vision Transformer backbone with 3D RoPE for unified image/video processing.
1104
+
1105
+ Key features:
1106
+ - PatchEmbed (Conv3d) for unified image/video tokenization
1107
+ - Full 3D axial RoPE positional encoding (T/H/W)
1108
+ - CLS token + register (storage) tokens
1109
+ - LayerScale and stochastic depth (DropPath)
1110
+ - MLP or SwiGLU FFN variants
1111
+ - FSDP-compatible (SHARD_GRAD_OP strategy)
1112
+
1113
+ Token sequence layout: [CLS] + [Register x n_storage_tokens] + [Patch x N_total]
1114
+
1115
+ Args:
1116
+ img_size: Input image size.
1117
+ patch_size: Spatial patch size.
1118
+ in_chans: Number of input channels.
1119
+ embed_dim: Embedding dimension.
1120
+ depth: Number of transformer blocks.
1121
+ num_heads: Number of attention heads.
1122
+ ffn_ratio: FFN hidden dimension ratio.
1123
+ qkv_bias: Whether to use bias in QKV.
1124
+ drop_path_rate: Stochastic depth rate.
1125
+ layerscale_init: LayerScale initial value (None to disable).
1126
+ norm_layer: Normalization layer name.
1127
+ ffn_layer: FFN layer name.
1128
+ ffn_bias: Whether to use bias in FFN.
1129
+ proj_bias: Whether to use bias in attention output projection.
1130
+ n_storage_tokens: Number of register tokens.
1131
+ mask_k_bias: Whether to mask K bias in attention.
1132
+ untie_cls_and_patch_norms: Use separate norms for CLS and patch tokens.
1133
+ untie_global_and_local_cls_norm: Use separate norm for local CLS tokens.
1134
+ device: Device for parameter allocation.
1135
+ num_frames: Number of input video frames.
1136
+ tubelet_size: Temporal patch size for Conv3d.
1137
+ pos_embed_rope_base: RoPE frequency base (100.0 for vision).
1138
+ gated_attention: Gated attention variant (None, "headwise", "elementwise").
1139
+ See https://arxiv.org/abs/2505.06708.
1140
+ """
1141
+
1142
+ def __init__(
1143
+ self,
1144
+ *,
1145
+ img_size: int = 224,
1146
+ patch_size: int = 16,
1147
+ in_chans: int = 3,
1148
+ embed_dim: int = 768,
1149
+ depth: int = 12,
1150
+ num_heads: int = 12,
1151
+ ffn_ratio: float = 4.0,
1152
+ qkv_bias: bool = True,
1153
+ drop_path_rate: float = 0.0,
1154
+ layerscale_init: float | None = None,
1155
+ norm_layer: str = "layernorm",
1156
+ ffn_layer: str = "mlp",
1157
+ ffn_bias: bool = True,
1158
+ proj_bias: bool = True,
1159
+ n_storage_tokens: int = 0,
1160
+ mask_k_bias: bool = False,
1161
+ untie_cls_and_patch_norms: bool = False,
1162
+ untie_global_and_local_cls_norm: bool = False,
1163
+ device: Any | None = None,
1164
+ num_frames: int = 1,
1165
+ tubelet_size: int = 1,
1166
+ pos_embed_rope_base: float = 100.0,
1167
+ pos_embed_rope_rescale_coords: float | None = None,
1168
+ pos_embed_rope_shift_coords: float | None = None,
1169
+ pos_embed_rope_jitter_coords: float | None = None,
1170
+ pos_embed_rope_fhw_dim: tuple[int, int, int] | None = None,
1171
+ gated_attention: str | None = None,
1172
+ qk_norm: bool = False,
1173
+ **ignored_kwargs,
1174
+ ) -> None:
1175
+ super().__init__()
1176
+ if len(ignored_kwargs) > 0:
1177
+ logger.warning(f"Ignored kwargs: {ignored_kwargs}")
1178
+
1179
+ norm_layer_cls = norm_layer_dict[norm_layer]
1180
+
1181
+ self.num_features = self.embed_dim = embed_dim
1182
+ self.n_blocks = depth
1183
+ self.num_heads = num_heads
1184
+ self.patch_size = patch_size
1185
+
1186
+ self.patch_embed = PatchEmbed(
1187
+ img_size=img_size,
1188
+ patch_size=patch_size,
1189
+ in_chans=in_chans,
1190
+ embed_dim=embed_dim,
1191
+ tubelet_size=tubelet_size,
1192
+ flatten_embedding=True,
1193
+ )
1194
+
1195
+ self.cls_token = nn.Parameter(torch.empty(1, 1, embed_dim, device=device))
1196
+ self.n_storage_tokens = n_storage_tokens
1197
+ if self.n_storage_tokens > 0:
1198
+ self.storage_tokens = nn.Parameter(torch.empty(1, n_storage_tokens, embed_dim, device=device))
1199
+
1200
+ # Convert 0.0 to None for backward compat (0.0 means disabled)
1201
+ _rescale = pos_embed_rope_rescale_coords if pos_embed_rope_rescale_coords else None
1202
+ _shift = pos_embed_rope_shift_coords if pos_embed_rope_shift_coords else None
1203
+ _jitter = pos_embed_rope_jitter_coords if pos_embed_rope_jitter_coords else None
1204
+ self.rope_embed = RopePositionEmbedding3D(
1205
+ embed_dim=embed_dim,
1206
+ num_heads=num_heads,
1207
+ fhw_dim=pos_embed_rope_fhw_dim,
1208
+ base=pos_embed_rope_base,
1209
+ rescale_coords=_rescale,
1210
+ shift_coords=_shift,
1211
+ jitter_coords=_jitter,
1212
+ )
1213
+
1214
+ logger.info(f"using {ffn_layer} layer as FFN")
1215
+ ffn_layer_cls = ffn_layer_dict[ffn_layer]
1216
+ ffn_ratio_sequence = [ffn_ratio] * depth
1217
+
1218
+ blocks_list = [
1219
+ SelfAttentionBlock(
1220
+ dim=embed_dim,
1221
+ num_heads=num_heads,
1222
+ ffn_ratio=ffn_ratio_sequence[i],
1223
+ qkv_bias=qkv_bias,
1224
+ proj_bias=proj_bias,
1225
+ ffn_bias=ffn_bias,
1226
+ drop_path=drop_path_rate,
1227
+ norm_layer=norm_layer_cls,
1228
+ act_layer=nn.GELU,
1229
+ ffn_layer=ffn_layer_cls,
1230
+ init_values=layerscale_init,
1231
+ mask_k_bias=mask_k_bias,
1232
+ device=device,
1233
+ gated_attention=gated_attention,
1234
+ qk_norm=qk_norm,
1235
+ )
1236
+ for i in range(depth)
1237
+ ]
1238
+
1239
+ self.chunked_blocks = False
1240
+ self.blocks = nn.ModuleList(blocks_list)
1241
+
1242
+ self.norm = norm_layer_cls(embed_dim)
1243
+
1244
+ self.untie_cls_and_patch_norms = untie_cls_and_patch_norms
1245
+ if untie_cls_and_patch_norms:
1246
+ self.cls_norm = norm_layer_cls(embed_dim)
1247
+ else:
1248
+ self.cls_norm = None
1249
+
1250
+ self.untie_global_and_local_cls_norm = untie_global_and_local_cls_norm
1251
+ if untie_global_and_local_cls_norm:
1252
+ self.local_cls_norm = norm_layer_cls(embed_dim)
1253
+ else:
1254
+ self.local_cls_norm = None
1255
+ self.head = nn.Identity()
1256
+ self.mask_token = nn.Parameter(torch.empty(1, embed_dim, device=device))
1257
+
1258
+ def init_weights(self) -> None:
1259
+ """Initialize all model weights."""
1260
+ self.rope_embed._init_weights()
1261
+ nn.init.normal_(self.cls_token, std=0.02)
1262
+ if self.n_storage_tokens > 0:
1263
+ nn.init.normal_(self.storage_tokens, std=0.02)
1264
+ nn.init.zeros_(self.mask_token)
1265
+ named_apply(init_weights_vit, self)
1266
+
1267
+ def prepare_tokens_with_masks(
1268
+ self,
1269
+ x: Tensor,
1270
+ masks: Tensor | None = None,
1271
+ ) -> tuple[Tensor, tuple[int, int, int]]:
1272
+ """Tokenize input and assemble token sequence with CLS + register + patches.
1273
+
1274
+ Args:
1275
+ x: Input tensor. Image: (B, C, H, W) or Video: (B, T, C, H, W).
1276
+ masks: Boolean mask of shape (B, N_spatial) indicating which patches to mask.
1277
+
1278
+ Returns:
1279
+ Tuple of:
1280
+ - Token sequence: (B, 1 + n_storage + N_total, embed_dim)
1281
+ - Grid dimensions: (T_grid, H_grid, W_grid)
1282
+ """
1283
+ if x.ndim == 5:
1284
+ B, T, C, H, W = x.shape
1285
+ # Video: Conv3d kernel=stride=tubelet downsamples raw T to T // tubelet.
1286
+ T_grid = T // self.patch_embed.tubelet_size
1287
+ else:
1288
+ B, C, H, W = x.shape
1289
+ # Image (4D): PatchEmbed expands raw T=1 to tubelet then Conv3d(stride=tubelet)
1290
+ # produces a single temporal token (output T_out = (tubelet - tubelet)/tubelet + 1 = 1).
1291
+ # vjepa2 vision_transformer.py:171-173 passes T=1 (no division) for the same reason.
1292
+ T_grid = 1
1293
+
1294
+ x = self.patch_embed(x) # (B, N_total, D)
1295
+
1296
+ # Grid dimensions for RoPE computation
1297
+ H_grid = H // self.patch_embed.patch_size[0]
1298
+ W_grid = W // self.patch_embed.patch_size[1]
1299
+
1300
+ if masks is not None:
1301
+ # Expand spatial mask to spatio-temporal if needed (tube masking)
1302
+ if masks.shape[1] != x.shape[1]:
1303
+ ratio = x.shape[1] // masks.shape[1]
1304
+ masks = masks.unsqueeze(1).repeat(1, ratio, 1).flatten(1)
1305
+
1306
+ # Replace masked positions with mask_token
1307
+ x = torch.where(masks.unsqueeze(-1), self.mask_token.to(x.dtype).unsqueeze(0), x)
1308
+ cls_token = self.cls_token
1309
+ else:
1310
+ # Include mask_token in computation graph even when not masking
1311
+ cls_token = self.cls_token + 0 * self.mask_token
1312
+
1313
+ if self.n_storage_tokens > 0:
1314
+ storage_tokens = self.storage_tokens
1315
+ else:
1316
+ storage_tokens = torch.empty(
1317
+ 1, 0, cls_token.shape[-1],
1318
+ dtype=cls_token.dtype, device=cls_token.device,
1319
+ )
1320
+
1321
+ x = torch.cat(
1322
+ [
1323
+ cls_token.expand(B, -1, -1),
1324
+ storage_tokens.expand(B, -1, -1),
1325
+ x,
1326
+ ],
1327
+ dim=1,
1328
+ )
1329
+
1330
+ return x, (T_grid, H_grid, W_grid)
1331
+
1332
+ def forward_features_list(
1333
+ self,
1334
+ x_list: list[Tensor],
1335
+ masks_list: list[Tensor | None],
1336
+ ) -> list[dict[str, Tensor]]:
1337
+ """Forward pass for a list of inputs (multi-crop).
1338
+
1339
+ Args:
1340
+ x_list: List of input tensors (global crops, local crops).
1341
+ masks_list: List of corresponding masks (None for unmasked).
1342
+
1343
+ Returns:
1344
+ List of output dictionaries, one per input, containing:
1345
+ - x_norm_clstoken: Normalized CLS token (B, D)
1346
+ - x_storage_tokens: Normalized register tokens (B, n_storage, D)
1347
+ - x_norm_patchtokens: Normalized patch tokens (B, N, D)
1348
+ - x_prenorm: Pre-normalization features (B, 1+n_storage+N, D)
1349
+ - masks: Original masks
1350
+ """
1351
+ x = []
1352
+ rope_params = []
1353
+ for t_x, t_masks in zip(x_list, masks_list):
1354
+ t2_x, grid_tuple = self.prepare_tokens_with_masks(t_x, t_masks)
1355
+ x.append(t2_x)
1356
+ rope_params.append(grid_tuple)
1357
+
1358
+ # Pre-compute RoPE sin/cos once — identical across all blocks.
1359
+ # Hoisting this out of the loop avoids breaking FSDP2's forward prefetch
1360
+ # chain (rope_embed is part of the outer FSDP unit, calling it between
1361
+ # block forwards disrupts the prefetch scheduling).
1362
+ if self.rope_embed is not None:
1363
+ rope_sincos = [self.rope_embed(T=t, H=h, W=w) for t, h, w in rope_params]
1364
+ else:
1365
+ rope_sincos = [None for _ in rope_params]
1366
+
1367
+ for _, blk in enumerate(self.blocks):
1368
+ x = blk(x, rope_sincos)
1369
+
1370
+ all_x = x
1371
+ output = []
1372
+ for idx, (x, masks) in enumerate(zip(all_x, masks_list)):
1373
+ if self.untie_cls_and_patch_norms or self.untie_global_and_local_cls_norm:
1374
+ if self.untie_global_and_local_cls_norm and self.training and idx == 1:
1375
+ x_norm_cls_reg = self.local_cls_norm(x[:, : self.n_storage_tokens + 1])
1376
+ elif self.untie_cls_and_patch_norms:
1377
+ x_norm_cls_reg = self.cls_norm(x[:, : self.n_storage_tokens + 1])
1378
+ else:
1379
+ x_norm_cls_reg = self.norm(x[:, : self.n_storage_tokens + 1])
1380
+ x_norm_patch = self.norm(x[:, self.n_storage_tokens + 1 :])
1381
+ else:
1382
+ x_norm = self.norm(x)
1383
+ x_norm_cls_reg = x_norm[:, : self.n_storage_tokens + 1]
1384
+ x_norm_patch = x_norm[:, self.n_storage_tokens + 1 :]
1385
+ output.append(
1386
+ {
1387
+ "x_norm_clstoken": x_norm_cls_reg[:, 0],
1388
+ "x_storage_tokens": x_norm_cls_reg[:, 1:],
1389
+ "x_norm_patchtokens": x_norm_patch,
1390
+ "x_prenorm": x,
1391
+ "masks": masks,
1392
+ }
1393
+ )
1394
+ return output
1395
+
1396
+ def forward_features(
1397
+ self,
1398
+ x: Tensor | list[Tensor],
1399
+ masks: Tensor | list[Tensor | None] | None = None,
1400
+ ) -> dict[str, Tensor] | list[dict[str, Tensor]]:
1401
+ """Forward pass for single or multiple inputs.
1402
+
1403
+ Args:
1404
+ x: Single tensor or list of tensors.
1405
+ masks: Single mask or list of masks.
1406
+
1407
+ Returns:
1408
+ Output dict (single input) or list of output dicts (multiple inputs).
1409
+ """
1410
+ if isinstance(x, torch.Tensor):
1411
+ return self.forward_features_list([x], [masks])[0]
1412
+ else:
1413
+ return self.forward_features_list(x, masks)
1414
+
1415
+ def _get_intermediate_layers_not_chunked(
1416
+ self,
1417
+ x: Tensor,
1418
+ n: int | list[int] = 1,
1419
+ ) -> list[Tensor]:
1420
+ """Run forward pass and collect intermediate block outputs.
1421
+
1422
+ Args:
1423
+ x: Input tensor (B, C, H, W) or (B, T, C, H, W).
1424
+ n: If int, return last n layers. If list, return specific layer indices.
1425
+
1426
+ Returns:
1427
+ List of intermediate outputs, each (B, 1+n_storage+N, D).
1428
+ """
1429
+ x, grid_tuple = self.prepare_tokens_with_masks(x, masks=None)
1430
+ T, H, W = grid_tuple
1431
+
1432
+ output, total_block_len = [], len(self.blocks)
1433
+ blocks_to_take = range(total_block_len - n, total_block_len) if isinstance(n, int) else n
1434
+
1435
+ if self.rope_embed is not None:
1436
+ rope_sincos = self.rope_embed(T=T, H=H, W=W)
1437
+ else:
1438
+ rope_sincos = None
1439
+
1440
+ for i, blk in enumerate(self.blocks):
1441
+ x = blk([x], [rope_sincos])[0]
1442
+ if i in blocks_to_take:
1443
+ output.append(x)
1444
+
1445
+ assert len(output) == len(blocks_to_take), (
1446
+ f"only {len(output)} / {len(blocks_to_take)} blocks found"
1447
+ )
1448
+ return output
1449
+
1450
+ def get_intermediate_layers(
1451
+ self,
1452
+ x: Tensor,
1453
+ n: int | list[int] = 1,
1454
+ reshape: bool = False,
1455
+ return_class_token: bool = False,
1456
+ norm: bool = True,
1457
+ ) -> tuple[Tensor, ...]:
1458
+ """Extract intermediate layer outputs for downstream evaluation.
1459
+
1460
+ This method is critical for dense prediction tasks (segmentation, depth)
1461
+ that need multi-scale features from different transformer blocks.
1462
+
1463
+ Args:
1464
+ x: Input image (B, C, H, W) or video (B, T, C, H, W).
1465
+ n: If int, return outputs from last n layers.
1466
+ If list[int], return outputs from specific layer indices.
1467
+ reshape: If True, reshape patch tokens to spatial form (B, D, H_grid, W_grid).
1468
+ return_class_token: If True, return (patch_tokens, cls_token) tuples.
1469
+ norm: If True, apply final LayerNorm to outputs.
1470
+
1471
+ Returns:
1472
+ If return_class_token is False:
1473
+ Tuple of patch token tensors, one per requested layer.
1474
+ Each tensor is (B, N, D) or (B, D, H_grid, W_grid) if reshape=True.
1475
+ If return_class_token is True:
1476
+ Tuple of (patch_tokens, cls_token) pairs.
1477
+ """
1478
+ # Determine spatial dims for reshape
1479
+ if x.ndim == 5:
1480
+ B, T_in, C, H, W = x.shape
1481
+ else:
1482
+ B, C, H, W = x.shape
1483
+ T_in = 1
1484
+
1485
+ outputs = self._get_intermediate_layers_not_chunked(x, n)
1486
+
1487
+ if norm:
1488
+ outputs_normed = []
1489
+ for out in outputs:
1490
+ if self.untie_cls_and_patch_norms:
1491
+ x_norm_cls_reg = self.cls_norm(out[:, : self.n_storage_tokens + 1])
1492
+ x_norm_patch = self.norm(out[:, self.n_storage_tokens + 1 :])
1493
+ outputs_normed.append(torch.cat((x_norm_cls_reg, x_norm_patch), dim=1))
1494
+ else:
1495
+ outputs_normed.append(self.norm(out))
1496
+ outputs = outputs_normed
1497
+
1498
+ class_tokens = [out[:, 0] for out in outputs]
1499
+ outputs = [out[:, self.n_storage_tokens + 1 :] for out in outputs]
1500
+
1501
+ if reshape:
1502
+ # Image (T_in=1): PatchEmbed expands to tubelet then Conv3d(stride=tubelet) → T_out=1.
1503
+ # Video: Conv3d downsamples T_in → T_in // tubelet. Matches prepare_tokens_with_masks
1504
+ # and vjepa2 vision_transformer.py:171-177.
1505
+ if x.ndim == 5:
1506
+ T_grid = T_in // self.patch_embed.tubelet_size
1507
+ else:
1508
+ T_grid = 1
1509
+ H_grid = H // self.patch_size
1510
+ W_grid = W // self.patch_size
1511
+ if T_grid > 1:
1512
+ # Video: reshape to (B, D, T_grid, H_grid, W_grid)
1513
+ outputs = [
1514
+ out.reshape(B, T_grid, H_grid, W_grid, -1).permute(0, 4, 1, 2, 3).contiguous()
1515
+ for out in outputs
1516
+ ]
1517
+ else:
1518
+ # Image: reshape to (B, D, H_grid, W_grid)
1519
+ outputs = [
1520
+ out.reshape(B, H_grid, W_grid, -1).permute(0, 3, 1, 2).contiguous()
1521
+ for out in outputs
1522
+ ]
1523
+
1524
+ if return_class_token:
1525
+ return tuple(zip(outputs, class_tokens))
1526
+ return tuple(outputs)
1527
+
1528
+ def forward(
1529
+ self,
1530
+ *args,
1531
+ is_training: bool = False,
1532
+ **kwargs,
1533
+ ) -> dict[str, Tensor] | list[dict[str, Tensor]] | Tensor:
1534
+ """High-level forward: training returns feature dict, inference returns CLS logits.
1535
+
1536
+ Args:
1537
+ is_training: If True, return full feature dictionary.
1538
+
1539
+ Returns:
1540
+ Feature dict(s) if training, CLS token logits if inference.
1541
+ """
1542
+ ret = self.forward_features(*args, **kwargs)
1543
+ if is_training:
1544
+ return ret
1545
+ else:
1546
+ return self.head(ret["x_norm_clstoken"])
1547
+
1548
+
1549
+ # ============================================================================
1550
+ # HuggingFace transformers wrapper (inference-only)
1551
+ # ============================================================================
1552
+ class MotifVisionConfig(PretrainedConfig):
1553
+ """Config for the MOTIF Vision Encoder backbone (image + video)."""
1554
+
1555
+ model_type = "motif_vision"
1556
+
1557
+ def __init__(
1558
+ self,
1559
+ img_size: int = 224,
1560
+ patch_size: int = 16,
1561
+ in_chans: int = 3,
1562
+ embed_dim: int = 4096,
1563
+ depth: int = 40,
1564
+ num_heads: int = 32,
1565
+ ffn_ratio: float = 3.0,
1566
+ qkv_bias: bool = False,
1567
+ drop_path_rate: float = 0.0,
1568
+ layerscale_init: float | None = 1.0e-5,
1569
+ norm_layer: str = "layernormbf16",
1570
+ ffn_layer: str = "swiglu64",
1571
+ ffn_bias: bool = True,
1572
+ proj_bias: bool = True,
1573
+ n_storage_tokens: int = 4,
1574
+ mask_k_bias: bool = True,
1575
+ untie_cls_and_patch_norms: bool = False,
1576
+ untie_global_and_local_cls_norm: bool = True,
1577
+ num_frames: int = 1,
1578
+ tubelet_size: int = 2,
1579
+ pos_embed_rope_base: float = 100.0,
1580
+ pos_embed_rope_rescale_coords: float | None = 2.0,
1581
+ gated_attention: str | None = "elementwise",
1582
+ qk_norm: bool = True,
1583
+ **kwargs,
1584
+ ):
1585
+ self.img_size = img_size
1586
+ self.patch_size = patch_size
1587
+ self.in_chans = in_chans
1588
+ self.embed_dim = embed_dim
1589
+ self.depth = depth
1590
+ self.num_heads = num_heads
1591
+ self.ffn_ratio = ffn_ratio
1592
+ self.qkv_bias = qkv_bias
1593
+ self.drop_path_rate = drop_path_rate
1594
+ self.layerscale_init = layerscale_init
1595
+ self.norm_layer = norm_layer
1596
+ self.ffn_layer = ffn_layer
1597
+ self.ffn_bias = ffn_bias
1598
+ self.proj_bias = proj_bias
1599
+ self.n_storage_tokens = n_storage_tokens
1600
+ self.mask_k_bias = mask_k_bias
1601
+ self.untie_cls_and_patch_norms = untie_cls_and_patch_norms
1602
+ self.untie_global_and_local_cls_norm = untie_global_and_local_cls_norm
1603
+ self.num_frames = num_frames
1604
+ self.tubelet_size = tubelet_size
1605
+ self.pos_embed_rope_base = pos_embed_rope_base
1606
+ self.pos_embed_rope_rescale_coords = pos_embed_rope_rescale_coords
1607
+ self.gated_attention = gated_attention
1608
+ self.qk_norm = qk_norm
1609
+ super().__init__(**kwargs)
1610
+
1611
+
1612
+ class MotifVisionModel(PreTrainedModel):
1613
+ """MOTIF Vision Encoder for HF `AutoModel` (inference). Returns dense + CLS features."""
1614
+
1615
+ config_class = MotifVisionConfig
1616
+ base_model_prefix = "motif"
1617
+ main_input_name = "pixel_values"
1618
+ _no_split_modules = ["SelfAttentionBlock"]
1619
+ supports_gradient_checkpointing = False
1620
+
1621
+ def __init__(self, config: MotifVisionConfig):
1622
+ super().__init__(config)
1623
+ self.backbone = MotifVisionTransformer(
1624
+ img_size=config.img_size,
1625
+ patch_size=config.patch_size,
1626
+ in_chans=config.in_chans,
1627
+ embed_dim=config.embed_dim,
1628
+ depth=config.depth,
1629
+ num_heads=config.num_heads,
1630
+ ffn_ratio=config.ffn_ratio,
1631
+ qkv_bias=config.qkv_bias,
1632
+ drop_path_rate=config.drop_path_rate,
1633
+ layerscale_init=config.layerscale_init,
1634
+ norm_layer=config.norm_layer,
1635
+ ffn_layer=config.ffn_layer,
1636
+ ffn_bias=config.ffn_bias,
1637
+ proj_bias=config.proj_bias,
1638
+ n_storage_tokens=config.n_storage_tokens,
1639
+ mask_k_bias=config.mask_k_bias,
1640
+ untie_cls_and_patch_norms=config.untie_cls_and_patch_norms,
1641
+ untie_global_and_local_cls_norm=config.untie_global_and_local_cls_norm,
1642
+ num_frames=config.num_frames,
1643
+ tubelet_size=config.tubelet_size,
1644
+ pos_embed_rope_base=config.pos_embed_rope_base,
1645
+ pos_embed_rope_rescale_coords=config.pos_embed_rope_rescale_coords,
1646
+ gated_attention=config.gated_attention,
1647
+ qk_norm=config.qk_norm,
1648
+ )
1649
+ self.post_init()
1650
+
1651
+ @torch.no_grad()
1652
+ def forward(self, pixel_values: Tensor, return_dict: bool = True, **kwargs):
1653
+ """pixel_values: image (B,3,H,W) or video (B,T,3,H,W). H,W multiples of patch_size."""
1654
+ out = self.backbone.forward_features(pixel_values)
1655
+ cls = out["x_norm_clstoken"]
1656
+ reg = out["x_storage_tokens"]
1657
+ patch = out["x_norm_patchtokens"]
1658
+ last_hidden = torch.cat([cls.unsqueeze(1), reg, patch], dim=1)
1659
+ if not return_dict:
1660
+ return (last_hidden, cls)
1661
+ return BaseModelOutputWithPooling(last_hidden_state=last_hidden, pooler_output=cls)
1662
+
1663
+
1664
+ AutoConfig_registered = False
1665
+ try:
1666
+ from transformers import AutoConfig, AutoModel
1667
+ AutoConfig.register("motif_vision", MotifVisionConfig)
1668
+ AutoModel.register(MotifVisionConfig, MotifVisionModel)
1669
+ AutoConfig_registered = True
1670
+ except Exception:
1671
+ pass
preprocessor_config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "image_processor_type": "BitImageProcessor",
3
+ "do_resize": true,
4
+ "size": {
5
+ "shortest_edge": 512
6
+ },
7
+ "do_center_crop": true,
8
+ "crop_size": {
9
+ "height": 512,
10
+ "width": 512
11
+ },
12
+ "do_rescale": true,
13
+ "rescale_factor": 0.00392156862745098,
14
+ "do_normalize": true,
15
+ "image_mean": [
16
+ 0.485,
17
+ 0.456,
18
+ 0.406
19
+ ],
20
+ "image_std": [
21
+ 0.229,
22
+ 0.224,
23
+ 0.225
24
+ ],
25
+ "resample": 3,
26
+ "do_convert_rgb": true
27
+ }