data-archetype commited on
Commit
75ff4df
·
0 Parent(s):

dinac3_96 v1.0

Browse files
.gitattributes ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
ATTRIBUTION.md ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Attribution
2
+
3
+ Licensing (code and weights) is stated once, in the License section of the model
4
+ card (README.md).
5
+
6
+ - **DINOv3 pretrained encoder:** Meta,
7
+ [official repository](https://github.com/facebookresearch/dinov3),
8
+ [paper](https://arxiv.org/abs/2508.10104),
9
+ [timm model](https://huggingface.co/timm/vit_base_patch16_dinov3.lvd_1689m).
10
+ The transformer blocks are used frozen; the input patch embedding is adapted
11
+ for this autoencoder. The weights are therefore distributed under the included
12
+ [DINOv3 License](LICENSE-DINOV3.md); see the model card's License section.
13
+ - **Encoder implementation:** provided by the separately installed
14
+ [timm](https://github.com/huggingface/pytorch-image-models) dependency, pinned
15
+ to `timm==1.0.26`.
16
+ - **semantic_vae** ([well9472/semantic_vae](https://huggingface.co/well9472/semantic_vae)):
17
+ the design follows its key ideas. These are a pretrained DINO encoder adapted
18
+ at its edges and a latent split into a reconstruction part and a DINO-derived
19
+ part. Its encoder has two DINOv2-B branches: one with a trainable patch
20
+ embedding, whose last six blocks are each LayerNormed and projected to 32
21
+ channels, and one fully frozen, whose final block is projected to 32 channels.
22
+ We also follow it in training the decoder mainly with a DINO feature loss and
23
+ in using VISReg as the latent regularizer.¹ No code or weights from it are included.
24
+ - **Decoder up-path:** the VQGAN f16 decoder layout of Esser, Rombach and Ommer,
25
+ [Taming Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2012.09841),
26
+ slightly modernized (per-position RMSNorm, folded upsampling) and behind a ViT
27
+ trunk.
28
+ - **Earlier decoder head (not in the release):** a pixel-level head after
29
+ [PixelDiT](https://arxiv.org/abs/2511.20645)
30
+ ([NVlabs/PixelDiT](https://github.com/NVlabs/PixelDiT)), with the FCDM block of
31
+ Kwon et al.,
32
+ [Reviving ConvNeXt for Efficient Convolutional Diffusion Models](https://arxiv.org/abs/2603.09408),
33
+ as its token mixer. The released encoder and decoder trunk were trained with it
34
+ before the convolutional head replaced it.
35
+ - **Latent regularizer:** [VISReg](https://arxiv.org/abs/2606.02572)
36
+ ([project page](https://haiyuwu.github.io/visreg/)), used during training only.
37
+ - **Optimizer (training only, not a runtime dependency):**
38
+ [dionw](https://github.com/JTriggerFish/dionw), a single-GPU implementation of
39
+ [Dion](https://arxiv.org/abs/2504.05295) /
40
+ [Dion3](https://arxiv.org/abs/2608.11612) with
41
+ [NorMuon](https://arxiv.org/abs/2510.05491) normalization, using kernels from
42
+ [microsoft/dion](https://github.com/microsoft/dion) and
43
+ [Dao-AILab/gram-newton-schulz](https://github.com/Dao-AILab/gram-newton-schulz).
44
+ - **Training data:** about 14 million images from a mix of public and licensed
45
+ or curated datasets: mostly photographs, plus book covers, a few text-heavy
46
+ datasets and 1% synthetic rendered text. No training images are redistributed.
47
+
48
+ ¹ semantic_vae's public model card names SIGReg (from LeJEPA) at 1e-3 as its
49
+ regularizer. We verified in its released training code that the regularizer is
50
+ VISReg, at 1e-3.
LICENSE ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction,
10
+ and distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by
13
+ the copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all
16
+ other entities that control, are controlled by, or are under common
17
+ control with that entity. For the purposes of this definition,
18
+ "control" means (i) the power, direct or indirect, to cause the
19
+ direction or management of such entity, whether by contract or
20
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
21
+ outstanding shares, or (iii) beneficial ownership of such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity
24
+ exercising permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation
28
+ source, and configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical
31
+ transformation or translation of a Source form, including but
32
+ not limited to compiled object code, generated documentation,
33
+ and conversions to other media types.
34
+
35
+ "Work" shall mean the work of authorship, whether in Source or
36
+ Object form, made available under the License, as indicated by a
37
+ copyright notice that is included in or attached to the work
38
+ (an example is provided in the Appendix below).
39
+
40
+ "Derivative Works" shall mean any work, whether in Source or Object
41
+ form, that is based on (or derived from) the Work and for which the
42
+ editorial revisions, annotations, elaborations, or other modifications
43
+ represent, as a whole, an original work of authorship. For the purposes
44
+ of this License, Derivative Works shall not include works that remain
45
+ separable from, or merely link (or bind by name) to the interfaces of,
46
+ the Work and Derivative Works thereof.
47
+
48
+ "Contribution" shall mean any work of authorship, including
49
+ the original version of the Work and any modifications or additions
50
+ to that Work or Derivative Works thereof, that is intentionally
51
+ submitted to Licensor for inclusion in the Work by the copyright owner
52
+ or by an individual or Legal Entity authorized to submit on behalf of
53
+ the copyright owner. For the purposes of this definition, "submitted"
54
+ means any form of electronic, verbal, or written communication sent
55
+ to the Licensor or its representatives, including but not limited to
56
+ communication on electronic mailing lists, source code control systems,
57
+ and issue tracking systems that are managed by, or on behalf of, the
58
+ Licensor for the purpose of discussing and improving the Work, but
59
+ excluding communication that is conspicuously marked or otherwise
60
+ designated in writing by the copyright owner as "Not a Contribution."
61
+
62
+ "Contributor" shall mean Licensor and any individual or Legal Entity
63
+ on behalf of whom a Contribution has been received by Licensor and
64
+ subsequently incorporated within the Work.
65
+
66
+ 2. Grant of Copyright License. Subject to the terms and conditions of
67
+ this License, each Contributor hereby grants to You a perpetual,
68
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
69
+ copyright license to reproduce, prepare Derivative Works of,
70
+ publicly display, publicly perform, sublicense, and distribute the
71
+ Work and such Derivative Works in Source or Object form.
72
+
73
+ 3. Grant of Patent License. Subject to the terms and conditions of
74
+ this License, each Contributor hereby grants to You a perpetual,
75
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
76
+ (except as stated in this section) patent license to make, have made,
77
+ use, offer to sell, sell, import, and otherwise transfer the Work,
78
+ where such license applies only to those patent claims licensable
79
+ by such Contributor that are necessarily infringed by their
80
+ Contribution(s) alone or by combination of their Contribution(s)
81
+ with the Work to which such Contribution(s) was submitted. If You
82
+ institute patent litigation against any entity (including a
83
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
84
+ or a Contribution incorporated within the Work constitutes direct
85
+ or contributory patent infringement, then any patent licenses
86
+ granted to You under this License for that Work shall terminate
87
+ as of the date such litigation is filed.
88
+
89
+ 4. Redistribution. You may reproduce and distribute copies of the
90
+ Work or Derivative Works thereof in any medium, with or without
91
+ modifications, and in Source or Object form, provided that You
92
+ meet the following conditions:
93
+
94
+ (a) You must give any other recipients of the Work or
95
+ Derivative Works a copy of this License; and
96
+
97
+ (b) You must cause any modified files to carry prominent notices
98
+ stating that You changed the files; and
99
+
100
+ (c) You must retain, in the Source form of any Derivative Works
101
+ that You distribute, all copyright, patent, trademark, and
102
+ attribution notices from the Source form of the Work,
103
+ excluding those notices that do not pertain to any part of
104
+ the Derivative Works; and
105
+
106
+ (d) If the Work includes a "NOTICE" text file as part of its
107
+ distribution, then any Derivative Works that You distribute must
108
+ include a readable copy of the attribution notices contained
109
+ within such NOTICE file, excluding those notices that do not
110
+ pertain to any part of the Derivative Works, in at least one
111
+ of the following places: within a NOTICE text file distributed
112
+ as part of the Derivative Works; within the Source form or
113
+ documentation, if provided along with the Derivative Works; or,
114
+ within a display generated by the Derivative Works, if and
115
+ wherever such third-party notices normally appear. The contents
116
+ of the NOTICE file are for informational purposes only and
117
+ do not modify the License. You may add Your own attribution
118
+ notices within Derivative Works that You distribute, alongside
119
+ or as an addendum to the NOTICE text from the Work, provided
120
+ that such additional attribution notices cannot be construed
121
+ as modifying the License.
122
+
123
+ You may add Your own copyright statement to Your modifications and
124
+ may provide additional or different license terms and conditions
125
+ for use, reproduction, or distribution of Your modifications, or
126
+ for any such Derivative Works as a whole, provided Your use,
127
+ reproduction, and distribution of the Work otherwise complies with
128
+ the conditions stated in this License.
129
+
130
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
131
+ any Contribution intentionally submitted for inclusion in the Work
132
+ by You to the Licensor shall be under the terms and conditions of
133
+ this License, without any additional terms or conditions.
134
+ Notwithstanding the above, nothing herein shall supersede or modify
135
+ the terms of any separate license agreement you may have executed
136
+ with Licensor regarding such Contributions.
137
+
138
+ 6. Trademarks. This License does not grant permission to use the trade
139
+ names, trademarks, service marks, or product names of the Licensor,
140
+ except as required for reasonable and customary use in describing the
141
+ origin of the Work and reproducing the content of the NOTICE file.
142
+
143
+ 7. Disclaimer of Warranty. Unless required by applicable law or
144
+ agreed to in writing, Licensor provides the Work (and each
145
+ Contributor provides its Contributions) on an "AS IS" BASIS,
146
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
147
+ implied, including, without limitation, any warranties or conditions
148
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
149
+ PARTICULAR PURPOSE. You are solely responsible for determining the
150
+ appropriateness of using or redistributing the Work and assume any
151
+ risks associated with Your exercise of permissions under this License.
152
+
153
+ 8. Limitation of Liability. In no event and under no legal theory,
154
+ whether in tort (including negligence), contract, or otherwise,
155
+ unless required by applicable law (such as deliberate and grossly
156
+ negligent acts) or agreed to in writing, shall any Contributor be
157
+ liable to You for damages, including any direct, indirect, special,
158
+ incidental, or consequential damages of any character arising as a
159
+ result of this License or out of the use or inability to use the
160
+ Work (including but not limited to damages for loss of goodwill,
161
+ work stoppage, computer failure or malfunction, or any and all
162
+ other commercial damages or losses), even if such Contributor
163
+ has been advised of the possibility of such damages.
164
+
165
+ 9. Accepting Warranty or Additional Liability. While redistributing
166
+ the Work or Derivative Works thereof, You may choose to offer,
167
+ and charge a fee for, acceptance of support, warranty, indemnity,
168
+ or other liability obligations and/or rights consistent with this
169
+ License. However, in accepting such obligations, You may act only
170
+ on Your own behalf and on Your sole responsibility, not on behalf
171
+ of any other Contributor, and only if You agree to indemnify,
172
+ defend, and hold each Contributor harmless for any liability
173
+ incurred by, or claims asserted against, such Contributor by reason
174
+ of your accepting any such warranty or additional liability.
175
+
176
+ END OF TERMS AND CONDITIONS
177
+
178
+ APPENDIX: How to apply the Apache License to your work.
179
+
180
+ To apply the Apache License to your work, attach the following
181
+ boilerplate notice, with the fields enclosed by brackets "[]"
182
+ replaced with your own identifying information. (Don't include
183
+ the brackets!) The text should be enclosed in the appropriate
184
+ comment syntax for the file format. We also recommend that a
185
+ file or class name and description of purpose be included on the
186
+ same "printed page" as the copyright notice for easier
187
+ identification within third-party archives.
188
+
189
+ Copyright [yyyy] [name of copyright owner]
190
+
191
+ Licensed under the Apache License, Version 2.0 (the "License");
192
+ you may not use this file except in compliance with the License.
193
+ You may obtain a copy of the License at
194
+
195
+ http://www.apache.org/licenses/LICENSE-2.0
196
+
197
+ Unless required by applicable law or agreed to in writing, software
198
+ distributed under the License is distributed on an "AS IS" BASIS,
199
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
200
+ See the License for the specific language governing permissions and
201
+ limitations under the License.
LICENSE-DINOV3.md ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # DINOv3 License
2
+
3
+ *Last Updated: August 19, 2025*
4
+
5
+ **“Agreement”** means the terms and conditions for use, reproduction, distribution and modification of the DINO Materials set forth herein.
6
+
7
+ **“DINO Materials”** means, collectively, Documentation and the models, software and algorithms, including machine-learning model code, trained model weights, inference-enabling code, training-enabling code, fine-tuning enabling code, and other elements of the foregoing distributed by Meta and made available under this Agreement.
8
+
9
+ **“Documentation”** means the specifications, manuals and documentation accompanying
10
+ DINO Materials distributed by Meta.
11
+
12
+ **“Licensee”** or **“you”** means you, or your employer or any other person or entity (if you are entering into this Agreement on such person or entity’s behalf), of the age required under applicable laws, rules or regulations to provide legal consent and that has legal authority to bind your employer or such other person or entity if you are entering in this Agreement on their behalf.
13
+
14
+ **“Meta”** or **“we”** means Meta Platforms Ireland Limited (if you are located in or, if you are an entity, your principal place of business is in the EEA or Switzerland) or Meta Platforms, Inc. (if you are located outside of the EEA or Switzerland).
15
+
16
+ **“Sanctions”** means any economic or trade sanctions or restrictions administered or enforced by the United States (including the Office of Foreign Assets Control of the U.S. Department of the Treasury (“OFAC”), the U.S. Department of State and the U.S. Department of Commerce), the United Nations, the European Union, or the United Kingdom.
17
+
18
+ **“Trade Controls”** means any of the following: Sanctions and applicable export and import controls.
19
+
20
+ By clicking “I Accept” below or by using or distributing any portion or element of the DINO Materials, you agree to be bound by this Agreement.
21
+
22
+ ## 1. License Rights and Redistribution.
23
+
24
+ a. <ins>Grant of Rights</ins>. You are granted a non-exclusive, worldwide, non-transferable and royalty-free limited license under Meta’s intellectual property or other rights owned by Meta embodied in the DINO Materials to use, reproduce, distribute, copy, create derivative works of, and make modifications to the DINO Materials.
25
+
26
+ b. <ins>Redistribution and Use</ins>.
27
+
28
+ i. Distribution of DINO Materials, and any derivative works thereof, are subject to the terms of this Agreement. If you distribute or make the DINO Materials, or any derivative works thereof, available to a third party, you may only do so under the terms of this Agreement and you shall provide a copy of this Agreement with any such DINO Materials.
29
+
30
+ ii. If you submit for publication the results of research you perform on, using, or otherwise in connection with DINO Materials, you must acknowledge the use of DINO Materials in your publication.
31
+
32
+ iii. Your use of the DINO Materials must comply with applicable laws and regulations, including Trade Control Laws and applicable privacy and data protection laws.
33
+
34
+ iv. Your use of the DINO Materials will not involve or encourage others to reverse engineer, decompile or discover the underlying components of the DINO Materials.
35
+
36
+ v. You are not the target of Trade Controls and your use of DINO Materials must comply with Trade Controls. You agree not to use, or permit others to use, DINO Materials for any activities subject to the International Traffic in Arms Regulations (ITAR) or end uses prohibited by Trade Controls, including those related to military or warfare purposes, nuclear industries or applications, espionage, or the development or use of guns or illegal weapons.
37
+
38
+ ## 2. User Support.
39
+
40
+ Your use of the DINO Materials is done at your own discretion; Meta does not process any information nor provide any service in relation to such use. Meta is under no obligation to provide any support services for the DINO Materials. Any support provided is “as is”, “with all faults”, and without warranty of any kind.
41
+
42
+ ## 3. Disclaimer of Warranty.
43
+
44
+ UNLESS REQUIRED BY APPLICABLE LAW, THE DINO MATERIALS AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED ON AN “AS IS” BASIS, WITHOUT WARRANTIES OF ANY KIND, AND META DISCLAIMS ALL WARRANTIES OF ANY KIND, BOTH EXPRESS AND IMPLIED, INCLUDING, WITHOUT LIMITATION, ANY WARRANTIES OF TITLE, NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. YOU ARE SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING OR REDISTRIBUTING THE DINO MATERIALS AND ASSUME ANY RISKS ASSOCIATED WITH YOUR USE OF THE DINO MATERIALS AND ANY OUTPUT AND RESULTS.
45
+
46
+ ## 4. Limitation of Liability.
47
+
48
+ IN NO EVENT WILL META OR ITS AFFILIATES BE LIABLE UNDER ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, TORT, NEGLIGENCE, PRODUCTS LIABILITY, OR OTHERWISE, ARISING OUT OF THIS AGREEMENT, FOR ANY LOST PROFITS OR ANY DIRECT OR INDIRECT, SPECIAL, CONSEQUENTIAL, INCIDENTAL, EXEMPLARY OR PUNITIVE DAMAGES, EVEN IF META OR ITS AFFILIATES HAVE BEEN ADVISED OF THE POSSIBILITY OF ANY OF THE FOREGOING.
49
+
50
+ ## 5. Intellectual Property.
51
+
52
+ a. Subject to Meta’s ownership of DINO Materials and derivatives made by or for Meta, with respect to any derivative works and modifications of the DINO Materials that are made by you, as between you and Meta, you are and will be the owner of such derivative works and modifications.
53
+
54
+ b. If you institute litigation or other proceedings against Meta or any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the DINO Materials, outputs or results, or any portion of any of the foregoing, constitutes infringement of intellectual property or other rights owned or licensable by you, then any licenses granted to you under this Agreement shall terminate as of the date such litigation or claim is filed or instituted. You will indemnify and hold harmless Meta from and against any claim by any third party arising out of or related to your use or distribution of the DINO Materials.
55
+
56
+ ## 6. Term and Termination.
57
+
58
+ The term of this Agreement will commence upon your acceptance of this Agreement or access to the DINO Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein. Meta may terminate this Agreement if you are in breach of any term or condition of this Agreement. Upon termination of this Agreement, you shall delete and cease use of the DINO Materials. Sections 3, 4 and 7 shall survive the termination of this Agreement.
59
+
60
+ ## 7. Governing Law and Jurisdiction.
61
+
62
+ This Agreement will be governed and construed under the laws of the State of California without regard to choice of law principles, and the UN Convention on Contracts for the International Sale of Goods does not apply to this Agreement. The courts of California shall have exclusive jurisdiction of any dispute arising out of this Agreement.
63
+
64
+ ## 8. Modifications and Amendments.
65
+
66
+ Meta may modify this Agreement from time to time; provided that they are similar in spirit to the current version of the Agreement, but may differ in detail to address new problems or concerns. All such changes will be effective immediately. Your continued use of the DINO Materials after any modification to this Agreement constitutes your agreement to such modification. Except as provided in this Agreement, no modification or addition to any provision of this Agreement will be binding unless it is in writing and signed by an authorized representative of both you and Meta.
README.md ADDED
@@ -0,0 +1,244 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: other
3
+ license_name: dinov3-license
4
+ license_link: https://huggingface.co/data-archetype/dinac3_96/blob/main/LICENSE-DINOV3.md
5
+ base_model: timm/vit_base_patch16_dinov3.lvd_1689m
6
+ tags:
7
+ - autoencoder
8
+ - image-reconstruction
9
+ - latent-space
10
+ - latent-diffusion
11
+ - dinov3
12
+ - pytorch
13
+ ---
14
+
15
+ # data-archetype/dinac3_96
16
+
17
+ **dinac3_96** is a deterministic semantic autoencoder for latent diffusion. It
18
+ maps an RGB image to a 96-channel latent grid at stride 16 and decodes it back in
19
+ a single forward pass.
20
+
21
+ It is strongly inspired by [semantic_vae][semantic-vae], the latent with the
22
+ fastest downstream DiT convergence we know of. dinac3_96 aims to keep that
23
+ convergence while improving reconstruction and simplifying the encoder to a
24
+ single DINOv3 pass instead of two.
25
+
26
+ - **Encoder:** a frozen [DINOv3][dinov3] ViT-B/16 behind a trainable patch
27
+ embedding. All twelve blocks are read, standardized per channel and mapped to
28
+ the latent by one trainable linear projection.
29
+ - **Latent:** 32 semantic channels, held during training to a fixed random
30
+ projection of DINOv3-B's blocks summed with fitted block weights, plus 64 reconstruction channels (reconstruction
31
+ losses and VISReg). No posterior noise.
32
+ - **Decoder:** a 4-block transformer trunk at width 1152 followed by a dense
33
+ convolutional up-path, a slightly modernized VQGAN-like decoder (stride 16 to
34
+ full resolution, three residual blocks per level). One pass, no diffusion
35
+ sampling.
36
+ - **Training losses:** mainly a DINOv3-B feature loss on the reconstruction, with
37
+ pixel and blurred-image MSE terms each at about a tenth of its gradient,
38
+ following semantic_vae's DINO-loss-trained decoder.
39
+
40
+ **[Technical report](TECHNICAL_REPORT.md)** ·
41
+ **[Reconstruction gallery](https://huggingface.co/spaces/data-archetype/dinac3-results)**
42
+ (39 images: original, reconstruction, RGB difference, latent PCA, per-image PSNR)
43
+
44
+ ## Generation benchmark
45
+
46
+ We evaluate a latent by training a class-conditional DiT on it (about 190M
47
+ parameters, 100k steps at batch 256 on ImageNet-1k, flow matching with SPRINT
48
+ routing). We score 10,000 samples against 810,000 real images on our
49
+ mixed-resolution ImageNet benchmark: four aspect ratios at 256- and 384-px areas.
50
+ Lower is better.[^fid] For each model: 50 NFEs, PDG 2.5–4.0, best settings kept.
51
+ Bold marks the better value in each column. The DiT was trained on dinac3_96's latent (the released latent
52
+ space, frozen since the joint phase), and its samples were decoded by the
53
+ released decoder.
54
+
55
+ | model | FID | MIND | Monge-DINO |
56
+ |---|---:|---:|---:|
57
+ | dinac3_96 | **`9.04`** | `8.29` | **`10.92`** |
58
+ | semantic_vae | `10.02` | **`7.42`** | `15.94` |
59
+
60
+ Monge-DINO uses DINOv3-B features, the same network whose features the decoder is
61
+ trained to match, so it is not independent of the training loss; FID and MIND, on
62
+ Inception features, are. The same holds for rMonge-DINOv3-B below.
63
+
64
+ Our blind pairwise comparison of the two models' generations put them on equal
65
+ footing.
66
+
67
+ Protocol, DiT settings and caveats:
68
+ [technical report, section 9.1](TECHNICAL_REPORT.md#91-generation-benchmark).
69
+
70
+ ## Reconstruction
71
+
72
+ **10,000 ImageNet images:** 10 real images per class with the generation
73
+ benchmark's aspect-ratio quotas (9 at 256-px area and 1 at 384 px per class).
74
+ Metrics compare the reconstructions with the same 10,000 originals.[^rfid] semantic_vae runs at its native bf16.
75
+
76
+ | model | rFID | rMIND | rMonge-DINOv3-B | PSNR mean | SSIM | LPIPS-VGG | LPIPS-Alex |
77
+ |---|---:|---:|---:|---:|---:|---:|---:|
78
+ | dinac3_96 | **`0.809`** | **`0.166`** | **`1.205`** | **`28.16`** | **`0.8195`** | **`0.0970`** | **`0.0362`** |
79
+ | semantic_vae | `2.149` | `0.820` | `5.994` | `25.19` | `0.7103` | `0.1715` | `0.0669` |
80
+
81
+ PSNR here: reconstructions clamped and rounded to uint8, peak 255, mean of
82
+ per-image values.
83
+
84
+ **2k PSNR benchmark** (the image set of our earlier releases):
85
+
86
+ | Model | Mean PSNR (dB) | Std (dB) | Median (dB) | P5 (dB) | P95 (dB) |
87
+ |---|---:|---:|---:|---:|---:|
88
+ | dinac3_96 | `32.15` | `5.11` | `31.64` | `24.27` | `40.70` |
89
+ | semantic_vae | `28.03` | `4.58` | `27.79` | `20.98` | `35.44` |
90
+ | dinac_ae_d2 | `35.59` | `4.87` | `35.40` | `27.89` | `43.51` |
91
+ | FLUX.2 VAE | `36.28` | `4.53` | `36.07` | `28.89` | `43.63` |
92
+
93
+ ## Latent interface
94
+
95
+ - 96 channels at stride 16; image height and width must be multiples of 16.
96
+ - Channels 0–63 are reconstruction channels, 64–95 semantic channels
97
+ (`free_channels(z)`, `semantic_channels(z)`).
98
+ - `encode(images)`: RGB in [-1, 1] → deterministic FP32 latents, whitened per
99
+ channel with the shipped statistics (what our DiTs were trained on).
100
+ - `decode(latents, height, width)`: whitened latents → FP32 RGB in about
101
+ [-1, 1], unclamped. Any multiple of 16 works, including sizes above 1024 px.
102
+ - `encode_raw` / `decode_raw` use the unwhitened latent; `whiten` / `dewhiten`
103
+ convert.
104
+
105
+ ## Precision
106
+
107
+ - Set the dtype only through `from_pretrained(dtype=...)`; dtype casts after
108
+ loading raise a `TypeError` (`.to(device)` works).
109
+ - `torch.bfloat16` (default) runs under bf16 autocast, bit-identical to the
110
+ training checkpoint under bf16 autocast. `torch.float32` runs without autocast.
111
+ - Encoder and decoder are compiled by default (about 20 s per module on the
112
+ first call). Pass `compile_encoder=False, compile_decoder=False` for eager.
113
+
114
+ ## Speed
115
+
116
+ RTX 5090, bf16, compiled (the default); mean of 20 timed batches after 3
117
+ warm-up batches. Peak VRAM includes the weights.
118
+
119
+ | Resolution | Batch | Encode ms/image | Decode ms/image | Peak VRAM encode / decode (GiB) |
120
+ |---:|---:|---:|---:|---:|
121
+ | `256x256` | `128` | `0.41` | `1.51` | `3.3` / `8.5` |
122
+ | `512x512` | `32` | `1.79` | `6.10` | `3.3` / `8.5` |
123
+ | `1024x1024` | `8` | `9.03` | `25.4` | `3.3` / `8.5` |
124
+ | `2048x2048` | `1` | `69.4` | `157` | `1.8` / `6.4` |
125
+
126
+ Eager mode decodes about 2.4× slower (1.8× at 2048²) and peaks at 12.7 GiB on
127
+ the batched rows.
128
+
129
+ **Size limit.** No resolution limit is coded; GPU memory is the practical one. On
130
+ a 32 GB RTX 5090, compiled decoding handles a single 4096 × 4096 image (15.8 s,
131
+ 24.5 GiB peak); eager decoding runs out of memory at that size.
132
+
133
+ ## Usage
134
+
135
+ Download the repository and install the package, in an environment with a
136
+ CUDA-enabled PyTorch (dependencies: `torch>=2.13`, `timm==1.0.26`, `safetensors`,
137
+ `huggingface-hub`):
138
+
139
+ ```bash
140
+ hf download data-archetype/dinac3_96 --local-dir dinac3_96
141
+ pip install ./dinac3_96
142
+ ```
143
+
144
+ ```python
145
+ import torch
146
+ from dinac3 import Dinac3
147
+
148
+ model = Dinac3.from_pretrained(
149
+ "data-archetype/dinac3_96", # or a local directory
150
+ device="cuda",
151
+ dtype=torch.bfloat16,
152
+ )
153
+
154
+ images = ... # [B, 3, H, W] on CUDA, RGB in [-1, 1], H and W multiples of 16
155
+
156
+ latents = model.encode(images) # [B, 96, H/16, W/16], whitened
157
+ semantic = model.semantic_channels(latents) # [B, 32, H/16, W/16]
158
+ recon = model.decode(latents, images.shape[-2], images.shape[-1])
159
+ recon = recon.clamp(-1, 1) # only for display or saving
160
+ ```
161
+
162
+ `from_pretrained(path_or_repo, *, device, dtype=torch.bfloat16,
163
+ compile_encoder=True, compile_decoder=True, revision=None)`:
164
+ - A `Path` is always a local directory.
165
+ - A `str` is a local directory if one exists.
166
+ - A `str` that looks like a path (it starts with `.`, `/` or `~`, or has more
167
+ than one `/`) but names no directory raises an error.
168
+ - Any other `str` is a Hub repository id; only `config.json` and
169
+ `model.safetensors` are downloaded.
170
+ - The DINOv3-B weights ship inside `model.safetensors`, so a local copy loads
171
+ offline.
172
+ - Inference is CUDA only.
173
+
174
+ ## Details
175
+
176
+ - **Network:** 164M parameters, 79M of them trained.
177
+ - Encoder (87M): frozen DINOv3 ViT-B/16 behind a trainable patch embedding;
178
+ all 12 blocks standardized, concatenated and projected linearly to 96
179
+ channels at stride 16.
180
+ - Decoder (78M): 1×1 projection to width 1152, 4 transformer blocks, then a
181
+ convolutional up-path from stride 16 to full resolution (512-channel
182
+ handoff, levels of 256, 256, 128 and 128 channels, three residual 3 × 3
183
+ blocks per level).
184
+ - **Training data:** about 14 million images from a mix of public and licensed
185
+ or curated datasets: mostly photographs, plus book covers, a few text-heavy
186
+ datasets and 1% synthetic rendered text. Aspect-ratio buckets, downsampled
187
+ only. No training images are redistributed with the model.
188
+ - **Training:**
189
+ - Losses: mainly a DINOv3-B all-block feature MSE on the reconstruction, plus
190
+ pixel and blurred-image MSE at about a tenth of its gradient; while the
191
+ encoder trained, also a semantic alignment loss and [VISReg][visreg].
192
+ - Steps: about 150k at 256/384 px (batch 128) plus about 25k fine-tuning at five
193
+ resolutions from 256 to 1024 px (batch 32). The free encoder weights (patch
194
+ embedding and output projection) were trained jointly with the decoder for
195
+ the first 52k steps and then frozen.
196
+ - Matrix weights trained with Dion ([dionw][dionw]), the rest with AdamW. The
197
+ released weights are EMA weights.
198
+ - **More:** the [technical report](TECHNICAL_REPORT.md) covers the semantic
199
+ target, the decoder heads tried, loss weights, training phases, ablations and
200
+ the benchmark protocol.
201
+ - **Related:** [semantic_vae][semantic-vae],
202
+ [DINAC-AE-D2](https://huggingface.co/data-archetype/dinac_ae_d2).
203
+
204
+ ## License
205
+
206
+ - **Code** (the `dinac3` package and scripts): Copyright 2026 data-archetype,
207
+ licensed under the Apache License 2.0, see [LICENSE](LICENSE).
208
+ - **Weights:** the DINOv3 License, see [LICENSE-DINOV3.md](LICENSE-DINOV3.md).
209
+ The weights contain DINOv3-B's weights and a trained derivative of its patch
210
+ embedding, so under the DINOv3 License (section 1.b.i) they may only be
211
+ distributed under that licence, which must accompany them. Its terms include:
212
+ - acknowledging the use of DINO Materials in publications (section 1.b.ii);
213
+ - no reverse engineering (section 1.b.iv);
214
+ - compliance with trade controls, which excludes ITAR-regulated uses and
215
+ military, nuclear, espionage and weapons applications (sections 1.b.iii and
216
+ 1.b.v).
217
+
218
+ The `license: other` tag above refers to the weights' DINOv3 License. See
219
+ [ATTRIBUTION.md](ATTRIBUTION.md) for sources and design credits.
220
+
221
+ ## Citation
222
+
223
+ ```bibtex
224
+ @misc{dinac3_96,
225
+ title = {dinac3_96: a DINOv3-based semantic autoencoder with a one-pass DINO-loss decoder},
226
+ author = {data-archetype},
227
+ email = {data-archetype@proton.me},
228
+ year = {2026},
229
+ month = oct,
230
+ url = {https://huggingface.co/data-archetype/dinac3_96},
231
+ }
232
+ ```
233
+
234
+ [dinov3]: https://arxiv.org/abs/2508.10104
235
+ [^fid]: Papers usually report FID on 50,000 square samples at a single
236
+ resolution, so our numbers are not directly comparable with published ones.
237
+
238
+ [^rfid]: Papers usually report rFID on the 50,000 ImageNet validation images at
239
+ a single square resolution, so our numbers are not directly comparable with
240
+ published ones.
241
+
242
+ [semantic-vae]: https://huggingface.co/well9472/semantic_vae
243
+ [visreg]: https://arxiv.org/abs/2606.02572
244
+ [dionw]: https://github.com/JTriggerFish/dionw
TECHNICAL_REPORT.md ADDED
@@ -0,0 +1,608 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # dinac3_96 Technical Report
2
+
3
+ `dinac3_96` is a deterministic semantic autoencoder for latent diffusion. A
4
+ frozen DINOv3 ViT-B/16 behind a trainable patch embedding is read at all twelve
5
+ blocks through one trainable linear projection, giving a 96-channel latent at
6
+ stride 16: 32 semantic channels held to a fixed DINOv3 target and 64 free
7
+ channels for reconstruction. A one-pass decoder (ViT trunk and dense
8
+ convolutional up-path) is trained mainly by a DINOv3 feature loss.
9
+
10
+ **Results in brief:** see the [README](README.md); protocols and full numbers in
11
+ [section 9](#9-results).
12
+
13
+ ## Contents
14
+
15
+ 1. [Overview](#1-overview)
16
+ 2. [Encoder](#2-encoder)
17
+ 3. [The semantic target](#3-the-semantic-target)
18
+ 4. [Decoder](#4-decoder)
19
+ 5. [Why a pure DINO loss](#5-why-a-pure-dino-loss)
20
+ 6. [Losses](#6-losses)
21
+ 7. [Data](#7-data)
22
+ 8. [Training](#8-training)
23
+ 9. [Results](#9-results)
24
+ 10. [References](#references)
25
+
26
+ ## 1. Overview
27
+
28
+ ```text
29
+ RGB image [B, 3, H, W], H and W divisible by 16
30
+ → trainable copy of DINOv3-B's patch-16 embedding
31
+ → 12 frozen DINOv3-B blocks
32
+ → every block's patch tokens through DINOv3's final LayerNorm,
33
+ standardized per channel, concatenated [B, 12 × 768, H/16, W/16]
34
+ → one trainable linear projection [B, 96, H/16, W/16]
35
+ channels 0–63: free (trained by reconstruction and VISReg)
36
+ channels 64–95: semantic (held to a fixed projected sum of DINOv3 blocks)
37
+
38
+ latent [B, 96, H/16, W/16]
39
+ → 1×1 projection to width 1152
40
+ → 4 ViT blocks at width 1152
41
+ → convolutional up-path: /16 → /8 → /4 → /2 → /1, three residual
42
+ blocks per level
43
+ → RGB [B, 3, H, W]
44
+ ```
45
+
46
+ Inputs are RGB in [-1, 1] of any size with height and width divisible by 16;
47
+ DINOv3's ImageNet normalization is applied internally.
48
+
49
+ The design keeps three ideas from [semantic_vae][semantic-vae]:
50
+
51
+ - **A pretrained DINO encoder adapted at its edges.** Its transformer blocks are
52
+ frozen; the patch embedding and the projection to the latent train.
53
+ - **A latent split into a DINO-derived part and a reconstruction part,**
54
+ similarly to our earlier [iRDiffAE][irdiffae] (half of its channels aligned to
55
+ DINOv2-S, half free) and `semantic_vae`.
56
+ - **A decoder trained by a DINO feature loss.** Like `semantic_vae`, we use an
57
+ (almost) exclusively DINO-based loss. We found this critical to the best
58
+ compromise between fast convergence of downstream diffusion models and good
59
+ reconstruction (sections 5 and 6).
60
+
61
+ The decoder head is a slightly modernized VQGAN-like convolutional up-path
62
+ behind a ViT trunk (section 4.2). Differences from `semantic_vae`:
63
+
64
+ - one encoder network instead of two, halving the encoder's compute;
65
+ - DINOv3-B instead of DINOv2-B;
66
+ - all twelve blocks read through one projection;
67
+ - semantic channels defined by a fixed projected sum of DINOv3-B's blocks with
68
+ fitted block weights (`semantic_vae` projects final-block features).
69
+
70
+ ## 2. Encoder
71
+
72
+ ### 2.1 Frozen DINOv3-B with a trainable patch embedding
73
+
74
+ [DINOv3][dinov3] ViT-B/16 (LVD-1689M): 12 blocks of width 768. Blocks, class and
75
+ register tokens and final LayerNorm are frozen. A trainable copy of the
76
+ pretrained patch embedding replaces the original; gradients reach it through the
77
+ frozen blocks. A second, unmodified DINOv3-B computes the targets of the semantic
78
+ alignment and of the DINO reconstruction loss.
79
+
80
+ **Why the patch embedding trains.** Self-supervised encoders discard colour,
81
+ intensity and fine-detail information that a decoder needs to recover. Training
82
+ the patch embedding lets the model recover some of it. Because DINOv3 is a
83
+ pre-norm residual transformer, this amounts to mixing pixel channels into DINO's
84
+ residual stream while the frozen blocks' features mostly stay as they were (the
85
+ alignment loss also pushes towards this).
86
+
87
+ ### 2.2 All twelve blocks through one projection
88
+
89
+ Patch tokens \\(h_l \in \mathbb{R}^{768}\\) of each block \\(l = 0, \dots, 11\\),
90
+ taken after DINOv3's final LayerNorm, are standardized with fixed statistics and
91
+ projected:
92
+
93
+ $$
94
+ \tilde h_l = \frac{h_l - \mu_l}{\sigma_l}, \qquad
95
+ x = [\tilde h_0; \dots; \tilde h_{11}] \in \mathbb{R}^{9216}, \qquad
96
+ z = W x + b \in \mathbb{R}^{96}.
97
+ $$
98
+
99
+ \\(\mu_l, \sigma_l\\) are fixed, measured once on training images with the
100
+ pretrained embedding.
101
+ \\(W\\), \\(b\\) and the patch embedding are the only trainable encoder parameters.
102
+
103
+ - **Per-channel standardization.** DINOv3-B's variance sits in a few
104
+ massive-activation channels (block 2: the top 8 of 768 hold 88%, participation
105
+ ratio 5.6; block 11: 11%, ratio 286). Unstandardized, the projection would
106
+ mostly copy them.
107
+ - **All twelve blocks, trainable projection.** Deep blocks carry semantics and
108
+ are smooth; shallow blocks carry detail. Our previous latents, fixed fusions of
109
+ all twelve blocks, gave good DiTs (FID 12.80) but capped
110
+ reconstruction, since only the patch embedding could add pixel information. A
111
+ trainable projection (32 semantic + 64 free) lifted reconstruction to 30–32 dB
112
+ with a diffusion decoder, while the semantic channels kept their class content
113
+ (probe 0.24, participation rank 29 of 32).
114
+
115
+ **Participation rank.** We measure how many channels of a group are effectively
116
+ used by the participation ratio \\((\sum_i \lambda_i)^2 / \sum_i \lambda_i^2\\)
117
+ of the eigenvalues \\(\lambda_i\\) of the group's token covariance. It equals
118
+ the number of channels when all eigenvalues are equal and 1 when a single one
119
+ dominates.
120
+
121
+ ### 2.3 Semantic and free channels
122
+
123
+ - **Semantic (last 32 rows):** initialized to the target map of section 3, so
124
+ equal to their target at step 0; held to it by an alignment loss during joint
125
+ training.
126
+ - **Free (64 rows):** random initialization with unit output variance; trained
127
+ only by the reconstruction losses and VISReg.
128
+
129
+ We chose 64 free channels over 32 for reconstruction. In our experiments a
130
+ 32 + 32 latent generated only marginally better (FID 8.85 against 9.04) but
131
+ reconstructed worse (ImageNet PSNR 26.24 against 28.16 dB, rFID 1.87 against
132
+ 0.81).
133
+
134
+ ### 2.4 Latent statistics and whitening
135
+
136
+ The latent is deterministic. Per-channel running means and variances, tracked
137
+ during joint training and frozen with the encoder, ship with the model. DiTs train
138
+ on whitened latents \\((z - \mu) / \sqrt{\sigma^2 + 10^{-4}}\\); the decoder
139
+ consumes the unwhitened latent (package interface in the [README](README.md)).
140
+
141
+ ## 3. The semantic target
142
+
143
+ ### 3.1 Why a fixed target
144
+
145
+ In our initial experiments we aligned the semantic channels to DINO through a
146
+ learned objective. These latents scored FID 14.1 to 19.1, far from `semantic_vae`'s 10.02, for a
147
+ common reason: each objective is best satisfied by DINO's highest-variance
148
+ directions, which are also its smoothest and carry relatively little class
149
+ information.
150
+
151
+ - **Learned up-projection with MSE or cosine** (latent → 768 → DINO features, as
152
+ in [DINAC-AE-D2][dinac-d2]): at fixed rank, MSE is minimized by the target's
153
+ top principal subspace (Eckart–Young); cosine on the full features behaves the
154
+ same.
155
+ - **CCA** between a channel block and DINO features: scale-invariant, so the
156
+ encoder picks the directions easiest to produce, in practice the same smooth
157
+ modes (60% subspace overlap with the top 32 principal directions).
158
+ - **CKA** on token similarities: dominated by the leading principal components of
159
+ both representations.
160
+ - **MSE to the class token:** constrains image-level content only, with no
161
+ per-token structure.
162
+
163
+ On DINOv3-B's final block, the top 32 principal directions have spectral slope
164
+ 2.25 (section 3.2) and class probe 0.167, against 1.84 and 0.291 for 32
165
+ class-discriminant directions.
166
+
167
+ A fixed target of the same width removes this choice. Each of the 32 semantic
168
+ channels regresses by per-channel MSE onto one channel of a fixed 32-channel map
169
+ of frozen DINO features, with no learned projection in between, so the map fully
170
+ specifies which directions the latent encodes. Every target channel has unit
171
+ variance, so each weighs equally in the loss. Section 3.2 gives the map.
172
+
173
+ ### 3.2 From a random projection to fitted layer weights
174
+
175
+ The target is built in two steps; only the second has free parameters.
176
+
177
+ **Step 1: one channel projection, shared by all blocks.** A single Gaussian
178
+ random matrix \\(P \in \mathbb{R}^{32 \times 768}\\), drawn once (seed 0) and kept
179
+ fixed, maps each of the 12 standardized blocks to 32 channels, which are
180
+ standardized again:
181
+
182
+ $$
183
+ y_l = \frac{P \tilde h_l - m_l}{s_l} \in \mathbb{R}^{32}, \qquad l = 0, \dots, 11.
184
+ $$
185
+
186
+ **Step 2: scalar block weights.** The 12 projected blocks are mixed channel by
187
+ channel with one scalar weight \\(w_l\\) per block (summing to 1) and rescaled to
188
+ unit variance, with \\(C_c\\) the 12 × 12 cross-block correlation of channel \\(c\\):
189
+
190
+ $$
191
+ t_c = \frac{\sum_l w_l\, y_{l,c}}{\sqrt{w^\top C_c\, w}}, \qquad c = 1, \dots, 32.
192
+ $$
193
+
194
+ The spectral slope \\(\alpha\\) is the exponent of a channel's radially
195
+ averaged power spectrum, \\(S(f) \propto f^{-\alpha}\\) (larger is smoother).
196
+
197
+ **Why a Gaussian random map.** Write a standardized block's token covariance as
198
+ \\(\Sigma = \sum_k \lambda_k u_k u_k^\top\\). \\(P\\) has i.i.d. \\(\mathcal{N}(0, 1)\\)
199
+ entries, so its distribution is rotation-invariant (\\(PU\\) is distributed as \\(P\\)
200
+ for any orthogonal \\(U\\)) and it favours no direction. Each output channel
201
+ \\(p_i^\top \tilde h\\) has expected variance
202
+ \\(\mathbb{E}[p_i^\top \Sigma p_i] = \sum_k \lambda_k\\): every principal mode
203
+ contributes in proportion to its variance, whereas a top-32 PCA projection keeps
204
+ modes 1 to 32 whole and drops all others.
205
+
206
+ **Choosing the block weights.** Equal block weights (\\(w_l = 1/12\\)) give a
207
+ latent that is too rough, with little class information (probe 0.123). We fitted
208
+ the 12 block weights to maximize an ImageNet class probe (ridge
209
+ classifier on image-mean tokens) subject to median \\(\alpha \le 1.10\\) and 90th
210
+ percentile \\(\le 1.20\\) at 256 px and median \\(\le 1.40\\) at 512 px, keeping the
211
+ spectrum centred near 1.0 at the 256–384 px training resolutions (`semantic_vae`:
212
+ median 1.09 at 256 px, probe 0.207). Search: an evolution strategy over
213
+ log-weights, selected on one half of the held-out images and reported on the
214
+ other (0.249 / 0.245).
215
+
216
+ The fit puts **0.53 on block 11** and **0.26 and 0.15 on blocks 0 and 1**; every
217
+ other block gets 0.02 or less. Block 11 carries the class content but alone is
218
+ too smooth (\\(\alpha\\) 1.64); the shallow blocks roughen it, while middle blocks
219
+ would add smoothness without class. Seven of eight random restarts converged to
220
+ the same pattern (block 11 at 0.45–0.56 plus blocks 0–2).
221
+
222
+ | 32-channel candidate (DINOv3-B unless noted) | \\(\alpha\\) at 256 px, p10 / median / p90 | \\(\alpha\\) at 512 px, median | class probe | participation rank |
223
+ | --- | --- | ---: | ---: | ---: |
224
+ | `semantic_vae` semantic block (DINOv2-B) | 0.99 / 1.09 / 1.19 | 1.39 | 0.207 | 24.0 |
225
+ | equal block weights, all 12 blocks | 0.78 / 0.88 / 0.96 | 1.20 | 0.123 | 27.8 |
226
+ | equal block weights, last 6 blocks | 1.03 / 1.07 / 1.15 | 1.43 | 0.154 | 27.2 |
227
+ | fitted block weights (probe projection) | 1.00 / 1.09 / 1.16 | 1.40 | 0.247 | 29.2 |
228
+ | **fitted block weights, as shipped (seed-0 Gaussian \\(P\\))** | **1.00 / 1.04 / 1.14** | **1.39** | **0.241** | **29.0** |
229
+
230
+ ### 3.3 Training the target in
231
+
232
+ The target is computed without gradient by the unmodified DINOv3-B on the same
233
+ image, projected by \\(P\\) and summed with the fitted block weights (section 3.2).
234
+ The encoder's trainable projection \\(W\\) is initialized with the same weights: its
235
+ 32 semantic rows are this composed linear map of the standardized blocks. With
236
+ the trainable patch embedding starting from its pretrained values, the semantic
237
+ channels equal their target at step 0. Training then updates the patch embedding
238
+ and \\(W\\), and the alignment loss, a per-token, per-channel MSE between the 32
239
+ semantic channels and the target, holds them close to it.
240
+
241
+ We weighted the alignment loss so that, on the semantic rows of \\(W\\), its gradient
242
+ is at least ten times that of the reconstruction losses, making sure the semantic
243
+ target dominates the pull of reconstruction.
244
+
245
+ ## 4. Decoder
246
+
247
+ Deterministic: one forward pass, no noisy-image input.
248
+
249
+ ### 4.1 Trunk
250
+
251
+ A 1×1 projection maps the 96 channels to width 1152, one token per 16 × 16
252
+ patch, followed by four ViT blocks: RMSNorm before and after the attention and
253
+ MLP branches, per-head query/key RMSNorm, axial 2D RoPE on patch indices, 18
254
+ heads of width 64, GELU MLP of ratio 4, global attention. In our
255
+ experiments an 8-block ViT trunk reconstructed 0.5–0.6 dB better than our
256
+ previous [FCDM][fcdm] trunk at matched steps, and extrapolated better (27.9 dB at 256 px
257
+ to 31.2 dB at 1024 px on whole photos).
258
+
259
+ ### 4.2 Convolutional up-path head
260
+
261
+ **Motivation.** Decoding each fixed 16 × 16 patch from its token, as in our
262
+ previous VAEs, did not give sharp enough detail, nor the right prior for good
263
+ reconstruction from noisy DiT latents. We therefore use a slightly modernized
264
+ [VQGAN][vqgan]-like decoder: the f16 up-path (levels, widths, three residual
265
+ blocks per level) with per-position RMSNorm and folded upsampling (below), and
266
+ our ViT trunk in place of its /16 section (input convolution, mid block and three
267
+ residual-plus-attention blocks, about 5% of its compute at 256 px).
268
+
269
+ ```text
270
+ trunk tokens [1152, H/16, W/16]
271
+ → RMSNorm → 1×1 conv 1152 → 512 → upsample ×2 (512) to /8
272
+ /8 ResBlock 512 → 256, ResBlock 256, ResBlock 256, upsample ×2 to /4
273
+ /4 ResBlock 256 × 3, upsample ×2 to /2
274
+ /2 ResBlock 256 → 128, ResBlock 128, ResBlock 128, upsample ×2 to /1
275
+ /1 ResBlock 128 × 3
276
+ → RMSNorm (affine) → SiLU → 3×3 conv 128 → 3 (zero-initialized) → RGB
277
+ ```
278
+
279
+ **Residual block** (VQGAN ResnetBlock, normalization replaced):
280
+
281
+ $$
282
+ h = \operatorname{conv}_{3\times3}\bigl(\operatorname{SiLU}(\operatorname{rms}(x))\bigr),\quad
283
+ h = \operatorname{conv}_{3\times3}\bigl(\operatorname{SiLU}(\operatorname{rms}(h))\bigr),\quad
284
+ \text{out} = \operatorname{shortcut}(x) + h,
285
+ $$
286
+
287
+ with a 1×1 convolution shortcut when the width changes and the identity
288
+ otherwise. Each block's second convolution and the output convolution are
289
+ zero-initialized, so a fresh block equals its shortcut and a fresh head outputs
290
+ zeros.
291
+
292
+ - **Per-position RMSNorm instead of GroupNorm.** Each `rms` normalizes one
293
+ position's channel vector (learned per-channel scale and shift) and never reads
294
+ other positions or images. GroupNorm's image-wide statistics make the output
295
+ depend on image size and distant regions, break tiled decoding, need a
296
+ reduction that does not fuse, and run in fp32 under autocast. Wan 2.1, KVAE and
297
+ DC-AE also use RMSNorm in some or all stages.
298
+ - **Upsampling.** Nearest ×2 then a 3 × 3 convolution, as in the FLUX and
299
+ `semantic_vae` decoders (no transposed-convolution checkerboards). Computed
300
+ exactly as a 4 × 4, stride-2 transposed convolution with a kernel folded from
301
+ the 3 × 3 weights: 16 instead of 36 multiply-adds per channel pair and no
302
+ 4×-size intermediate.
303
+ - **Full resolution kept.** In a diffusion-decoder comparison, ending at /2 with
304
+ a pixel-shuffle readout trained about twice as fast but lost fine detail
305
+ (lowest-noise validation loss 36% higher, visibly worse samples).
306
+
307
+ **Parameters.**
308
+
309
+ | part | parameters | trained |
310
+ | --- | ---: | --- |
311
+ | DINOv3-B backbone (12 blocks, class and 4 register tokens, final norm) | 85.05M | frozen |
312
+ | patch embedding (trainable copy) | 0.59M | yes |
313
+ | layer projection (9216 → 96, with bias) | 0.88M | yes |
314
+ | decoder trunk (1×1 latent projection and 4 ViT blocks at width 1152) | 63.85M | yes |
315
+ | convolutional up-path head | 14.05M | yes |
316
+ | **total** | **164.43M** (79.38M trained) | |
317
+
318
+ Encoder 86.53M (1.48M trained), decoder 77.90M.
319
+
320
+ **Resolution extrapolation test.** The decoder accepts any multiple of 16, including sizes
321
+ above its 1024-px training resolution. On 12 Pexels photos it scores 28.55 dB
322
+ PSNR at 512-px area, 29.72 dB at 1024-px area and 31.54 dB at stored size (3.9–4.2
323
+ megapixels). The 4-megapixel
324
+ reconstruction downscaled 2× scores 35.42 dB against the 1024-px input, above
325
+ the native 1024-px decode on all 12 images. 2× crops show no seams, swirls or
326
+ tiling. Rows and columns on 16-px patch boundaries carry 6–12% more error; a
327
+ 32-px border ring is 2–3 dB worse.
328
+
329
+ ## 5. Why a pure DINO loss
330
+
331
+ Compared with diffusion decoders and diffusion + DINO decoders:
332
+
333
+ 1. **Smoother latents.** Diffusion decoders push detail into the latent, which
334
+ the DiT then has to generate. From the same encoder at matched steps (10k),
335
+ the DINO-loss latent had median spectral slope 1.04 against
336
+ 0.97 (5th percentile 0.32 against 0.22) and 2.3 dB lower PSNR, but its DINO
337
+ loss at 10k (0.086) was already below the diffusion decoder's at 85k (0.104).
338
+ Every DINO-loss latent we benchmarked beat every diffusion-decoder latent on
339
+ FID, MIND and Monge-DINO.
340
+ 2. **Robust to the DiT's high-frequency errors.** Changing the DiT's sampler
341
+ moved Monge-DINO by 1.5–6.3 with our diffusion decoders and by 0.4, within
342
+ one standard error, with a DINO-loss decoder.
343
+ 3. **One step instead of about four** (our diffusion decoders used 4 Euler
344
+ steps), with more visible detail in our comparisons, including at 1024 px.
345
+
346
+ **Trade-off.** At the same latent width, diffusion decoders reach higher PSNR
347
+ (2.3 dB in the matched 10k-step comparison of point 1) and have fewer deterministic artifacts:
348
+ DINO-loss textures can swirl, and text is sometimes distorted.
349
+
350
+ ## 6. Losses
351
+
352
+ Weights are normalized to DINO = 1. Dion and AdamW are invariant to the overall
353
+ loss scale, so only the ratios matter.
354
+
355
+ | term | definition | weight | phases |
356
+ | --- | --- | ---: | --- |
357
+ | DINO reconstruction | MSE between frozen DINOv3-B features of the reconstruction and of the input, at all 12 blocks (each through DINOv3's final LayerNorm), averaged over blocks, tokens (patch and class) and channels | **1** | all |
358
+ | pixel MSE | MSE of RGB in [-1, 1] | **0.380** | all |
359
+ | blurred MSE | MSE between both images after a Gaussian blur with σ = 5 px | **0.505** | all |
360
+ | semantic alignment | per-token MSE of the 32 semantic channels to the fixed target (section 3) | **57.7** | joint |
361
+ | VISReg | [VISReg][visreg] on all 96 channels: centring, unit scale and a Gaussian-quantile shape term on 256 random projections, equal weights | **0.0202** | joint |
362
+
363
+ $$
364
+ \mathcal{L} = \mathcal{L}_{\text{DINO}}
365
+ + 0.380\,\mathcal{L}_{\text{pix}} + 0.505\,\mathcal{L}_{\text{blur}}
366
+ + 57.7\,\mathcal{L}_{\text{align}} + 0.0202\,\mathcal{L}_{\text{VISReg}}
367
+ $$
368
+
369
+ No diffusion, KL, adversarial or LPIPS loss. The
370
+ DINO loss is not standardized per block: in our experiments, blocks 2–6 carried 70% of its value and block 11 under 1%, following feature
371
+ scale.
372
+
373
+ ### 6.1 Why VISReg
374
+
375
+ [VISReg][visreg] has a centring term on channel means, a scale term on channel
376
+ standard deviations and a shape term matching the sorted values of 256 random
377
+ 1-D projections to standard-normal quantiles, all computed across the batch at
378
+ each spatial position, as in `semantic_vae` (its model card names SIGReg, but its
379
+ released training code uses VISReg).
380
+
381
+ - **Normalization.** Reconstruction is blind to an invertible per-channel affine
382
+ change of a deterministic latent, so nothing else fixes the free channels'
383
+ scale and offset. The centring term is their only mean anchor: in our experiments,
384
+ when it was outweighed by an alignment loss on a shared encoder, the detail
385
+ channels' means drifted by up to 3 standard deviations within a few
386
+ thousand steps.
387
+ - **Anti-collapse.** Projections are taken after per-channel standardization
388
+ (detached std), so the projection term acts on the correlation matrix \\(R\\):
389
+ unit-variance random projections ask for \\(u^\top R u \approx 1\\) for random
390
+ unit \\(u\\), pushing \\(R\\) weakly toward \\(I\\). This flattens the spectrum and raises
391
+ participation rank, penalizing a free block folded onto a few modes. Quantile
392
+ matching also constrains higher moments, but in our experiments 91% of the
393
+ quantile term was this projected-scale mismatch.
394
+
395
+ In our experiments, weakening VISReg or removing its projection term let the
396
+ free channels collapse: at 1/50 of the usual weight their participation rank fell
397
+ to 11 of 64 (33 at the usual weight), and a latent trained without the projection
398
+ term collapsed to 13 of 96.
399
+
400
+ We also tested a variant that keeps the centring and scale terms but replaces
401
+ quantile matching by matching only each projection's standard deviation to 1.
402
+ With \\(\hat z\\) the per-channel-standardized latent and \\(u_1, \dots, u_K\\) random
403
+ unit directions,
404
+
405
+ $$
406
+ \mathcal{L}_{\text{proj}} = \frac{1}{K} \sum_{k=1}^{K}
407
+ \bigl(\operatorname{std}(u_k^\top \hat z) - 1\bigr)^2
408
+ = \frac{1}{K} \sum_{k=1}^{K} \bigl(\sqrt{u_k^\top R\, u_k} - 1\bigr)^2,
409
+ $$
410
+
411
+ which acts on the correlation matrix \\(R\\) alone, without higher moments. It
412
+ performed similarly to VISReg, confirming the mechanism above. We nevertheless
413
+ kept VISReg, which is cheap and more standard.
414
+
415
+ ### 6.2 How the weights were set
416
+
417
+ Weights were set from per-term gradient norms, each term back-propagated
418
+ separately at a checkpoint over several image draws. Loss values are a poor
419
+ guide: in our experiments the blurred MSE had a weighted loss of only 2.4e-4 but
420
+ the largest decoder gradient, 3–4 times DINO's.
421
+
422
+ - **Pixel and blurred MSE:** each a tenth of DINO's gradient norm on the
423
+ decoder (both had cosine +0.15 to +0.33 with DINO's gradient); on this
424
+ model's convolutional head at 2k steps they measured 0.112 and 0.116. The pixel
425
+ MSE acts as a light regulariser on top of the DINO loss; the blurred term is
426
+ there because DINO losses are nearly blind to colour.
427
+ - **Alignment (57.7):** the ten-times rule of section 3.3.
428
+ - **VISReg (0.0202):** set mostly empirically from the latent's statistics and
429
+ not extensively ablated. We also tested a roughness penalty on the free
430
+ channels, a hinge keeping each channel's neighbour R² (how well a token is
431
+ predicted from its neighbours) above a threshold, and found it unnecessary
432
+ with the DINO loss.
433
+
434
+ ## 7. Data
435
+
436
+ - **Sources.** About 14 million images, mostly photographs,
437
+ plus book covers, a few text-heavy datasets and 1% synthetic rendered text
438
+ samples (single characters and simple words). A mix of public and licensed or
439
+ curated datasets; no training images are redistributed with the model.
440
+ - **Aspect-ratio buckets** from 1:4 to 4:1; one aspect ratio and resolution per
441
+ batch.
442
+ - **Downsampling only.** Antialiased bicubic resizing to the bucket's shape at the
443
+ target area (e.g. 896 × 1152 → 224 × 288 at 256 px); images smaller than a
444
+ target are never upsampled.
445
+ - **Resolution mixes.** Base stage: 256-px area for 90% of batches, 384 px for
446
+ 10%. Five-resolution stage: 256 / 384 / 512 / 768 / 1024 px at 20% each (256 as
447
+ square crops).
448
+
449
+ ## 8. Training
450
+
451
+ **Optimizer.** Weight matrices use [Dion][dion] / [Dion3][dion3] orthogonalized
452
+ updates with [NorMuon][normuon] normalization, via our single-GPU
453
+ [dionw][dionw], with the update RMS matched to AdamW's so one learning rate
454
+ serves both. Everything else uses AdamW (β = (0.9, 0.98)); no weight decay.
455
+ bf16 AMP, gradient clipping at 5, EMA with decay
456
+ 0.9995; the released weights are the EMA.
457
+
458
+ | stage | resolutions | batch | learning rate | steps |
459
+ | --- | --- | ---: | ---: | ---: |
460
+ | base | 256-px area (90%), 384 px (10%) | 128 | 1e-4 | about 150k |
461
+ | five resolutions | 256 to 1024 px, 20% each | 32 | 2.5e-5 | about 25k |
462
+
463
+ The free encoder weights (patch embedding and output projection) trained
464
+ jointly with the decoder for the first 52k steps and were then frozen: the
465
+ latent had stabilised, and freezing saves the encoder's backward pass.
466
+
467
+ ## 9. Results
468
+
469
+ ### 9.1 Generation benchmark
470
+
471
+ We judge a latent by a DiT trained on it: in our experiments, static latent
472
+ statistics (spectral slope, neighbour R², class probes, participation rank, PSNR)
473
+ did not predict DiT quality.
474
+
475
+ **Protocol.** Our class-conditional, mixed-resolution ImageNet benchmark.[^fid]
476
+
477
+ - **Generated set:** 10,000 images, 10 per class: 9 at 256-px area and 1 at
478
+ 384 px, in four aspect-ratio families (288 × 224, 288 × 192, 224 × 288 and
479
+ 192 × 288 at 256 px; 416 × 320, 448 × 288, 320 × 416 and 288 × 448 at 384 px)
480
+ with quotas from their training-data frequency. Every model uses the same
481
+ class, shape and seed manifest.
482
+ - **Reference set:** 810,000 real ImageNet training images (729,000 at 256 px and 81,000 at 384 px, 81 times each shape's generated
483
+ count), preprocessed as in training.
484
+ - **Metrics** (lower is better): [FID][fid] on Inception pool3 features; MIND
485
+ ([Monge Inception Distance][mind]), a sliced optimal-transport distance on the
486
+ same features (1,024 fixed random projections); Monge-DINO, the same
487
+ estimator on final-normalized DINOv3-B class tokens of 224-px centre crops. The
488
+ Monge ± values reflect only the randomness of the projections.
489
+ - **Guidance.** Path-drop guidance (PDG) contrasts the full model with the same model skipping its middle blocks,
490
+ both class-conditioned; no classifier-free guidance.
491
+
492
+ **The DiT.** A class-conditional flow-matching transformer with
493
+ [SPRINT][sprint] token routing, 190.43M parameters:
494
+
495
+ | setting | value |
496
+ | --- | --- |
497
+ | layers | 16 at width 896: 2 dense prefix, 12 routed middle, 2 dense suffix |
498
+ | attention | 14 heads of width 64; additive 2D sin-cos positions and axial 2D RoPE |
499
+ | MLP | GELU, ratio 4 |
500
+ | class conditioning | 4 learned tokens of width 256 per class, cross-attention (4 heads of width 64) in every layer; class dropped for 10% of samples |
501
+ | time conditioning | AdaLN per layer from a 256-wide time embedding |
502
+ | input / output | one token per latent cell (96 channels), latents whitened per channel with the VAE's statistics; velocity prediction |
503
+ | SPRINT mix, per sample | 80%: middle layers see one random token per 2 × 2 group (75% dropped); 10%: full grid; 10%: middle layers skipped. Middle outputs are scattered back, concatenated with the prefix output, scaled per channel by the timestep and projected back to width 896 |
504
+ | objective | flow matching, MSE on velocity, Beta(2, 2) timesteps with a resolution-dependent log-SNR shift relative to 256 px |
505
+ | data | ImageNet-1k (1.24M training images), aspect-ratio buckets, 90% at 256-px area, 10% at 384 px |
506
+ | optimizer | AdamW, learning rate 1e-4, β = (0.9, 0.98), weight decay 0, constant after 4,000 warmup steps, gradient clipping at 1 |
507
+ | batch, steps | 256, 100,000 |
508
+ | precision | bf16 AMP |
509
+ | EMA | 2,000-step half-life (decay ≈ 0.99965), from step 20,000; EMA weights used |
510
+
511
+ The latent has been fixed since the joint phase, so the released decoder
512
+ decodes exactly the latent the DiT learned.
513
+
514
+ `semantic_vae`'s DiT uses the same recipe, up to minor architectural details.
515
+
516
+ **Scores.** For each model: 50 NFEs, PDG 2.5–4.0, best settings kept. The DiT's
517
+ samples were decoded by the released convolutional decoder. Bold marks the
518
+ better value in each column.
519
+
520
+ | model | FID ↓ | MIND ↓ | Monge-DINOv3-B ↓ |
521
+ | --- | ---: | ---: | ---: |
522
+ | `dinac3_96` | **9.04** | 8.29 ± 0.19 | **10.92** ± 0.46 |
523
+ | `semantic_vae` | 10.02 | **7.42** ± 0.20 | 15.94 ± 0.66 |
524
+
525
+ Monge-DINO uses DINOv3-B features, as does the decoder's training loss, so it is
526
+ not independent of it; FID and MIND, on Inception features, are.
527
+
528
+ **Blind comparison.** Our blind pairwise comparison of the two models'
529
+ generations put them on equal footing.
530
+
531
+ **Limits.** One DiT seed per row; the benchmark's own noise has not been
532
+ measured. We treat FID differences of 0.1 to 0.2 as noise.
533
+
534
+ ### 9.2 Reconstruction
535
+
536
+ The main tables are in the [README](README.md#reconstruction); this section adds
537
+ detail.
538
+
539
+ **10,000 ImageNet images.** Originals preprocessed like the generation
540
+ benchmark's reference images. PSNR spread and the Monge uncertainties (from the
541
+ projections):
542
+
543
+ | model | rMIND | rMonge-DINOv3-B | PSNR median | P5 | P95 |
544
+ | --- | ---: | ---: | ---: | ---: | ---: |
545
+ | `dinac3_96` | **0.166** ± 0.002 | **1.205** ± 0.044 | **28.07** | **21.38** | **35.58** |
546
+ | `semantic_vae` | 0.820 ± 0.028 | 5.994 ± 0.216 | 25.09 | 18.83 | 31.82 |
547
+
548
+ **2k PSNR set.** 1,330 Pexels photos, 669 book covers and one printer calibration test image,
549
+ centre-cropped to a multiple of 16 at stored size (mostly about 1 megapixel).
550
+ PSNR between the float [-1, 1] input and the unclamped reconstruction, peak 2,
551
+ the DINAC-AE-D2 card's protocol. The DINAC-AE-D2 and FLUX.2 rows are quoted from
552
+ the [DINAC-AE-D2 card](https://huggingface.co/data-archetype/dinac_ae_d2); our
553
+ re-run of DINAC-AE-D2 reproduced it within 0.02 dB. The FLUX.2 row most likely
554
+ used a clamped reconstruction (inferred from its value); clamped, `dinac3_96`
555
+ scores 32.17 dB.
556
+
557
+ **By resolution.** Each 2k image resized down to an aspect-ratio bucket of the
558
+ given area; images too small are skipped, so subsets differ between rows
559
+ (compare within a row). Mean PSNR in dB:
560
+
561
+ | bucket | images | `dinac3_96` | `semantic_vae` | DINAC-AE-D2 |
562
+ | --- | ---: | ---: | ---: | ---: |
563
+ | 256 | 2,000 | 29.27 | 25.48 | 32.63 |
564
+ | 512 | 1,331 | 32.48 | 28.41 | 35.72 |
565
+ | 768 | 1,331 | 33.80 | 29.40 | 37.20 |
566
+ | 1024 | 1,124 | 34.06 | 29.81 | 37.45 |
567
+
568
+ **PSNR.** A decoder trained mainly by a DINO feature loss trades some pixel
569
+ accuracy for perceptual fidelity: on the 2k set `dinac3_96` is 3.4 dB below the
570
+ pixel-trained DINAC-AE-D2 and about 4 dB below FLUX.2 (the trade-off of section
571
+ 5). Text is the visible weakness: glyphs are sometimes distorted, plausibly
572
+ because the loss encodes text style more than glyph identity.
573
+
574
+ ## References
575
+
576
+ 1. well9472. [semantic_vae][semantic-vae]. Model and training notes.
577
+ 2. O. Siméoni et al. [DINOv3][dinov3]. arXiv:2508.10104, 2025.
578
+ 3. Y. Yu, W. Xiong, W. Nie et al. [PixelDiT: Pixel Diffusion Transformers for Image Generation][pixeldit]. arXiv:2511.20645, 2025.
579
+ 4. Kwon et al. [Reviving ConvNeXt for Efficient Convolutional Diffusion Models][fcdm] (FCDM). arXiv:2603.09408, 2026.
580
+ 5. K. Ahn, B. Xu, N. Abreu, Y. Fan et al. [Dion: Distributed Orthonormalized Updates][dion]. arXiv:2504.05295, 2025.
581
+ 6. N. Amsel, J. Zhang, K. Ahn, A. Naeimi et al. [Dion3: Full-Stack Orthogonal Updates][dion3]. arXiv:2608.11612, 2026.
582
+ 7. Z. Li, L. Liu, C. Liang, W. Chen et al. [NorMuon: Making Muon more efficient and scalable][normuon]. arXiv:2510.05491, 2025.
583
+ 8. [dionw][dionw]: single-GPU Dion with fused AdamW.
584
+ 9. Wu, Balestriero and Levine. [VISReg: Variance-Invariance-Sketching Regularization for JEPA training][visreg]. arXiv:2606.02572, 2026.
585
+ 10. Park et al. [Sprint: Sparse-Dense Residual Fusion for Efficient Diffusion Transformers][sprint]. arXiv:2510.21986, 2025.
586
+ 11. Berthet et al. [MIND: Monge Inception Distance for Generative Models Evaluation][mind]. arXiv:2605.06797, 2026.
587
+ 12. M. Heusel et al. [GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium][fid] (FID). arXiv:1706.08500, 2017.
588
+ 13. P. Esser, R. Rombach, B. Ommer. [Taming Transformers for High-Resolution Image Synthesis][vqgan] (VQGAN). arXiv:2012.09841, 2020.
589
+ 14. Data Archetype. [DINAC-AE-D2][dinac-d2].
590
+
591
+ [^fid]: Papers usually report FID on 50,000 square samples at a single
592
+ resolution, so our numbers are not directly comparable with published ones.
593
+
594
+ [semantic-vae]: https://huggingface.co/well9472/semantic_vae
595
+ [irdiffae]: https://huggingface.co/data-archetype/irdiffae-v1/blob/main/technical_report.md
596
+ [dinov3]: https://arxiv.org/abs/2508.10104
597
+ [pixeldit]: https://arxiv.org/abs/2511.20645
598
+ [fcdm]: https://arxiv.org/abs/2603.09408
599
+ [dion]: https://arxiv.org/abs/2504.05295
600
+ [dion3]: https://arxiv.org/abs/2608.11612
601
+ [normuon]: https://arxiv.org/abs/2510.05491
602
+ [dionw]: https://github.com/JTriggerFish/dionw
603
+ [visreg]: https://arxiv.org/abs/2606.02572
604
+ [sprint]: https://arxiv.org/abs/2510.21986
605
+ [mind]: https://arxiv.org/abs/2605.06797
606
+ [fid]: https://arxiv.org/abs/1706.08500
607
+ [vqgan]: https://arxiv.org/abs/2012.09841
608
+ [dinac-d2]: https://huggingface.co/data-archetype/dinac_ae_d2
config.json ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "format_version": 1,
3
+ "backbone": "vit_base_patch16_dinov3",
4
+ "latent_channels": 96,
5
+ "semantic_channels": 32,
6
+ "decoder_width": 1152,
7
+ "decoder_depth": 4,
8
+ "decoder_head_dim": 64,
9
+ "decoder_mlp_ratio": 4.0,
10
+ "conv_up_handoff_channels": 512,
11
+ "conv_up_channels": [
12
+ 256,
13
+ 256,
14
+ 128,
15
+ 128
16
+ ],
17
+ "conv_up_blocks_per_level": 3,
18
+ "latent_stats_eps": 0.0001
19
+ }
dinac3/__init__.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ """dinac3: deterministic image autoencoders with a DINOv3-aligned latent."""
2
+
3
+ from .config import CONFIG_FORMAT_VERSION, PATCH, Backbone, Dinac3Config
4
+ from .model import Dinac3
5
+
6
+ __all__ = ["CONFIG_FORMAT_VERSION", "PATCH", "Backbone", "Dinac3", "Dinac3Config"]
dinac3/config.py ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Explicit architecture configuration for dinac3 autoencoders."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import math
7
+ from dataclasses import asdict, dataclass, fields
8
+ from enum import Enum
9
+ from typing import TYPE_CHECKING, cast
10
+
11
+ if TYPE_CHECKING:
12
+ from pathlib import Path
13
+
14
+ # Pixels per latent token along each side (DINOv3-B patch size).
15
+ PATCH = 16
16
+ # Version of the config.json schema this package reads and writes.
17
+ CONFIG_FORMAT_VERSION = 1
18
+
19
+
20
+ class Backbone(str, Enum):
21
+ """Supported frozen encoder backbones (timm architecture names)."""
22
+
23
+ DINO_V3_B = "vit_base_patch16_dinov3"
24
+
25
+
26
+ def _require_int(data: dict[str, object], name: str) -> int:
27
+ """An integer JSON field (not a float such as ``64.0``, not a boolean).
28
+
29
+ Raises:
30
+ TypeError: If the value is not a JSON integer.
31
+ """
32
+
33
+ value = data[name]
34
+ if type(value) is not int:
35
+ raise TypeError(f"Config field {name} must be an integer, got {value!r}")
36
+ return value
37
+
38
+
39
+ def _require_int_tuple(data: dict[str, object], name: str) -> tuple[int, ...]:
40
+ """A JSON list of integers, as a tuple.
41
+
42
+ Raises:
43
+ TypeError: If the value is not a list of JSON integers.
44
+ """
45
+
46
+ value = data[name]
47
+ if not isinstance(value, list) or any(
48
+ type(item) is not int for item in cast("list[object]", value)
49
+ ):
50
+ raise TypeError(f"Config field {name} must be a list of integers: {value!r}")
51
+ return tuple(cast("list[int]", value))
52
+
53
+
54
+ def _require_float(data: dict[str, object], name: str) -> float:
55
+ """A real-number JSON field (an integer literal is accepted, a boolean not).
56
+
57
+ Raises:
58
+ TypeError: If the value is not a JSON number.
59
+ """
60
+
61
+ value = data[name]
62
+ if type(value) not in (int, float):
63
+ raise TypeError(f"Config field {name} must be a number, got {value!r}")
64
+ return float(cast("float", value))
65
+
66
+
67
+ def _require_str(data: dict[str, object], name: str) -> str:
68
+ """A string JSON field.
69
+
70
+ Raises:
71
+ TypeError: If the value is not a JSON string.
72
+ """
73
+
74
+ value = data[name]
75
+ if not isinstance(value, str):
76
+ raise TypeError(f"Config field {name} must be a string, got {value!r}")
77
+ return value
78
+
79
+
80
+ @dataclass(frozen=True)
81
+ class Dinac3Config:
82
+ """Architecture of one dinac3 model; every persisted field is required.
83
+
84
+ The latent has ``latent_channels`` channels at stride 16: the first
85
+ ``latent_channels - semantic_channels`` are free (detail) channels, the last
86
+ ``semantic_channels`` the DINOv3-aligned semantic channels.
87
+ """
88
+
89
+ format_version: int
90
+ backbone: Backbone
91
+ latent_channels: int
92
+ semantic_channels: int
93
+ decoder_width: int
94
+ decoder_depth: int
95
+ decoder_head_dim: int
96
+ decoder_mlp_ratio: float
97
+ conv_up_handoff_channels: int
98
+ conv_up_channels: tuple[int, ...]
99
+ conv_up_blocks_per_level: int
100
+ latent_stats_eps: float
101
+
102
+ def __post_init__(self) -> None:
103
+ """Reject unsupported topologies and invalid scalars.
104
+
105
+ Raises:
106
+ ValueError: On another format version, a non-positive width, depth
107
+ or count, a semantic block not smaller than the latent, a width
108
+ not divisible into heads, a head width not a multiple of 4, a
109
+ number of up-path levels other than log2(16) = 4, or a
110
+ non-finite or non-positive ratio or epsilon.
111
+ TypeError: If ``backbone`` is not a :class:`Backbone` or
112
+ ``conv_up_channels`` is not a tuple.
113
+ """
114
+
115
+ if self.format_version != CONFIG_FORMAT_VERSION:
116
+ raise ValueError(
117
+ f"config format_version {self.format_version} is not supported; "
118
+ f"this package reads version {CONFIG_FORMAT_VERSION}"
119
+ )
120
+ if not isinstance(self.backbone, Backbone):
121
+ raise TypeError("backbone must be a Backbone enum")
122
+ counts = (
123
+ self.latent_channels,
124
+ self.semantic_channels,
125
+ self.decoder_width,
126
+ self.decoder_depth,
127
+ self.decoder_head_dim,
128
+ self.conv_up_handoff_channels,
129
+ self.conv_up_blocks_per_level,
130
+ *self.conv_up_channels,
131
+ )
132
+ if any(value <= 0 for value in counts):
133
+ raise ValueError("Every width, depth and count must be positive")
134
+ if self.semantic_channels >= self.latent_channels:
135
+ raise ValueError("semantic_channels must be smaller than latent_channels")
136
+ if self.decoder_width % self.decoder_head_dim:
137
+ raise ValueError("decoder_width must be a multiple of decoder_head_dim")
138
+ if self.decoder_head_dim % 4:
139
+ raise ValueError("decoder_head_dim must be a multiple of 4 (2-D RoPE)")
140
+ if not isinstance(self.conv_up_channels, tuple):
141
+ raise TypeError("conv_up_channels must be a tuple of level widths")
142
+ if 2 ** len(self.conv_up_channels) != PATCH:
143
+ raise ValueError(
144
+ f"conv_up_channels needs one width per level from /{PATCH // 2} to "
145
+ f"/1 ({PATCH.bit_length() - 1} levels)"
146
+ )
147
+ scalars = (self.decoder_mlp_ratio, self.latent_stats_eps)
148
+ if not all(math.isfinite(v) and v > 0 for v in scalars):
149
+ raise ValueError(
150
+ "decoder_mlp_ratio and latent_stats_eps must be finite and positive"
151
+ )
152
+
153
+ @property
154
+ def free_channels(self) -> int:
155
+ """Number of leading free (detail) latent channels."""
156
+
157
+ return self.latent_channels - self.semantic_channels
158
+
159
+ def save(self, path: Path) -> None:
160
+ """Write every field as JSON."""
161
+
162
+ path.write_text(json.dumps(asdict(self), indent=2) + "\n")
163
+
164
+ @classmethod
165
+ def load(cls, path: Path) -> Dinac3Config:
166
+ """Load strictly: missing or unknown fields and wrong JSON types are errors.
167
+
168
+ Raises:
169
+ ValueError: If the JSON fields differ from the dataclass fields or
170
+ the format version is not this package's.
171
+ TypeError: If a field has the wrong JSON type (e.g. ``64.0`` or
172
+ ``true`` for an integer).
173
+ """
174
+
175
+ data = json.loads(path.read_text())
176
+ if not isinstance(data, dict):
177
+ raise TypeError(f"{path} must hold a JSON object")
178
+ values = cast("dict[str, object]", data)
179
+ expected = {field.name for field in fields(cls)}
180
+ if set(values) != expected:
181
+ raise ValueError(f"Config fields must be exactly {sorted(expected)}")
182
+ return cls(
183
+ format_version=_require_int(values, "format_version"),
184
+ backbone=Backbone(_require_str(values, "backbone")),
185
+ latent_channels=_require_int(values, "latent_channels"),
186
+ semantic_channels=_require_int(values, "semantic_channels"),
187
+ decoder_width=_require_int(values, "decoder_width"),
188
+ decoder_depth=_require_int(values, "decoder_depth"),
189
+ decoder_head_dim=_require_int(values, "decoder_head_dim"),
190
+ decoder_mlp_ratio=_require_float(values, "decoder_mlp_ratio"),
191
+ conv_up_handoff_channels=_require_int(values, "conv_up_handoff_channels"),
192
+ conv_up_channels=_require_int_tuple(values, "conv_up_channels"),
193
+ conv_up_blocks_per_level=_require_int(values, "conv_up_blocks_per_level"),
194
+ latent_stats_eps=_require_float(values, "latent_stats_eps"),
195
+ )
dinac3/conv_up_head.py ADDED
@@ -0,0 +1,180 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Deterministic convolutional up-path head: trunk tokens at /16 to RGB.
2
+
3
+ A handoff (per-position RMSNorm, 1x1 projection, upsampling to /8), then one
4
+ level per stride (/8, /4, /2, /1) of residual blocks, upsampling between levels,
5
+ and a readout (per-position RMSNorm, SiLU, 3x3 convolution to RGB). Every
6
+ upsampling is nearest x2 followed by a 3x3 convolution, computed exactly as a
7
+ stride-2 transposed convolution whose 4x4 kernel folds the 3x3 weights. No
8
+ attention, no time conditioning, no noise input: one deterministic pass.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from typing import TYPE_CHECKING
14
+
15
+ import torch
16
+ from torch import Tensor, nn
17
+ from torch.nn import functional as F
18
+
19
+ from .layers import ChannelRMSNorm, pointwise_conv, same_conv
20
+
21
+ if TYPE_CHECKING:
22
+ from .config import Dinac3Config
23
+
24
+ # Row k maps a 3x3 kernel's taps (offsets -1, 0, +1) onto tap k of a 4-tap,
25
+ # stride-2 transposed convolution: nearest x2 followed by the 3x3 convolution.
26
+ NEAREST_CONV_FOLD = (
27
+ (0.0, 0.0, 1.0),
28
+ (0.0, 1.0, 1.0),
29
+ (1.0, 1.0, 0.0),
30
+ (1.0, 0.0, 0.0),
31
+ )
32
+
33
+
34
+ class UpsampleConv(nn.Module):
35
+ """Nearest x2 then a 3x3 convolution, as one folded transposed convolution.
36
+
37
+ The 4x4 transposed kernel is folded once, at load time (:meth:`fold_kernel`),
38
+ from the FP32 3x3 weight in IEEE FP32 arithmetic, so a global TF32 matmul
39
+ setting cannot change it; the transposed convolution then runs in the
40
+ activation precision. Loading a state dict into a folded module folds again,
41
+ so the kernel never goes stale.
42
+ """
43
+
44
+ fold: Tensor
45
+ kernel: Tensor | None
46
+
47
+ def __init__(self, channels: int) -> None:
48
+ """Allocate the 3x3 convolution, the constant fold matrix and the
49
+ (not yet folded) kernel buffer."""
50
+
51
+ super().__init__()
52
+ self.conv = nn.Conv2d(channels, channels, 3, padding=1)
53
+ self.register_buffer("fold", torch.tensor(NEAREST_CONV_FOLD))
54
+ self.register_buffer("kernel", None, persistent=False)
55
+ self.register_load_state_dict_post_hook(_refold_after_load)
56
+
57
+ @torch.no_grad()
58
+ def fold_kernel(self) -> None:
59
+ """Fold the 3x3 weight into the ``[C_in, C_out, 4, 4]`` transposed kernel
60
+ (IEEE FP32 matmuls for the fold, restored afterwards).
61
+
62
+ Raises:
63
+ ValueError: If the stored fold matrix is not ``NEAREST_CONV_FOLD``.
64
+ """
65
+
66
+ expected = torch.tensor(NEAREST_CONV_FOLD, device=self.fold.device)
67
+ if not torch.equal(self.fold.float(), expected):
68
+ raise ValueError("UpsampleConv.fold is not the nearest x2 fold matrix")
69
+ matmul = torch.backends.cuda.matmul
70
+ previous = matmul.fp32_precision
71
+ matmul.fp32_precision = "ieee"
72
+ try:
73
+ with torch.autocast(device_type=self.fold.device.type, enabled=False):
74
+ fold = self.fold.to(dtype=self.conv.weight.dtype)
75
+ kernel = torch.einsum("ka,ocab,lb->ockl", fold, self.conv.weight, fold)
76
+ finally:
77
+ matmul.fp32_precision = previous
78
+ self.kernel = kernel.transpose(0, 1)
79
+
80
+ def forward(self, x: Tensor) -> Tensor:
81
+ """``[B, C, H, W]`` to ``[B, C, 2H, 2W]``.
82
+
83
+ Raises:
84
+ RuntimeError: Before :meth:`fold_kernel` (done by ``Dinac3.prepare``).
85
+ """
86
+
87
+ kernel = self.kernel
88
+ if kernel is None:
89
+ raise RuntimeError("UpsampleConv.fold_kernel() has not run")
90
+ return F.conv_transpose2d(x, kernel, self.conv.bias, stride=2, padding=1)
91
+
92
+
93
+ def _refold_after_load(module: nn.Module, _incompatible: object) -> None:
94
+ """After a state-dict load (children included), fold an already folded
95
+ upsampler again so its kernel never goes stale.
96
+
97
+ Raises:
98
+ TypeError: If registered on another module.
99
+ """
100
+
101
+ if not isinstance(module, UpsampleConv):
102
+ raise TypeError(f"Fold hook registered on {type(module)}")
103
+ if module.kernel is not None:
104
+ module.fold_kernel()
105
+
106
+
107
+ class ResBlock(nn.Module):
108
+ """``shortcut(x) + conv2(SiLU(norm2(conv1(SiLU(norm1(x))))))``, with
109
+ per-position RMSNorms carrying per-channel gains and biases, 3x3
110
+ convolutions and a 1x1 shortcut when the width changes."""
111
+
112
+ def __init__(self, in_channels: int, out_channels: int) -> None:
113
+ """Allocate norms, convolutions and the optional shortcut."""
114
+
115
+ super().__init__()
116
+ self.norm1 = ChannelRMSNorm(in_channels, affine=True)
117
+ self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
118
+ self.norm2 = ChannelRMSNorm(out_channels, affine=True)
119
+ self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
120
+ self.shortcut = (
121
+ nn.Conv2d(in_channels, out_channels, 1)
122
+ if in_channels != out_channels
123
+ else None
124
+ )
125
+
126
+ def forward(self, x: Tensor) -> Tensor:
127
+ """Residual update of a ``[B, C, H, W]`` map."""
128
+
129
+ h = same_conv(self.conv1, F.silu(self.norm1(x)))
130
+ h = same_conv(self.conv2, F.silu(self.norm2(h)))
131
+ skip = x if self.shortcut is None else pointwise_conv(self.shortcut, x)
132
+ return skip + h
133
+
134
+
135
+ class ConvUpHead(nn.Module):
136
+ """Trunk tokens ``[B, width, h, w]`` to RGB ``[B, 3, 16 h, 16 w]``."""
137
+
138
+ def __init__(self, config: Dinac3Config) -> None:
139
+ """Allocate handoff, residual blocks, upsamplers and readout."""
140
+
141
+ super().__init__()
142
+ handoff = config.conv_up_handoff_channels
143
+ self.blocks_per_level = config.conv_up_blocks_per_level
144
+ self.handoff_norm = ChannelRMSNorm(config.decoder_width, affine=False)
145
+ self.handoff_proj = nn.Conv2d(config.decoder_width, handoff, 1)
146
+ self.handoff_upsample = UpsampleConv(handoff)
147
+ blocks: list[ResBlock] = []
148
+ upsamples: list[UpsampleConv] = []
149
+ width = handoff
150
+ for index, channels in enumerate(config.conv_up_channels):
151
+ for _ in range(self.blocks_per_level):
152
+ blocks.append(ResBlock(width, channels))
153
+ width = channels
154
+ if index + 1 < len(config.conv_up_channels):
155
+ upsamples.append(UpsampleConv(width))
156
+ self.blocks = nn.ModuleList(blocks)
157
+ self.upsamples = nn.ModuleList(upsamples)
158
+ self.out_norm = ChannelRMSNorm(width, affine=True)
159
+ self.out_proj = nn.Conv2d(width, 3, 3, padding=1)
160
+
161
+ def fold_kernels(self) -> None:
162
+ """Fold every upsampler's kernel (after loading and any device move)."""
163
+
164
+ self.handoff_upsample.fold_kernel()
165
+ for upsampler in self.upsamples:
166
+ if not isinstance(upsampler, UpsampleConv):
167
+ raise TypeError(f"Unexpected upsampler {type(upsampler)}")
168
+ upsampler.fold_kernel()
169
+
170
+ def forward(self, features: Tensor) -> Tensor:
171
+ """Decode a contiguous NCHW token map."""
172
+
173
+ h = pointwise_conv(self.handoff_proj, self.handoff_norm(features))
174
+ h = self.handoff_upsample(h)
175
+ for index, block in enumerate(self.blocks):
176
+ h = block(h)
177
+ level, position = divmod(index, self.blocks_per_level)
178
+ if position == self.blocks_per_level - 1 and level < len(self.upsamples):
179
+ h = self.upsamples[level](h)
180
+ return same_conv(self.out_proj, F.silu(self.out_norm(h)))
dinac3/decoder.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Deterministic one-pass decoder: latent projection, ViT trunk and conv up-path."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING
6
+
7
+ from torch import Tensor, nn
8
+
9
+ from .conv_up_head import ConvUpHead
10
+ from .trunk import DitBlock, rope_tables
11
+
12
+ if TYPE_CHECKING:
13
+ from .config import Dinac3Config
14
+
15
+
16
+ class Decoder(nn.Module):
17
+ """Raw latents ``[B, C, h, w]`` to RGB ``[B, 3, 16 h, 16 w]``."""
18
+
19
+ def __init__(self, config: Dinac3Config) -> None:
20
+ """Allocate the latent projection, trunk blocks and conv up-path head."""
21
+
22
+ super().__init__()
23
+ width = config.decoder_width
24
+ self.head_dim = config.decoder_head_dim
25
+ self.latent_up = nn.Conv2d(config.latent_channels, width, 1)
26
+ self.trunk = nn.ModuleList(
27
+ [
28
+ DitBlock(width, config.decoder_head_dim, config.decoder_mlp_ratio)
29
+ for _ in range(config.decoder_depth)
30
+ ]
31
+ )
32
+ self.conv_up_head = ConvUpHead(config)
33
+
34
+ def forward(self, latents: Tensor) -> Tensor:
35
+ """Decode in one pass; latents are validated by the public API."""
36
+
37
+ x = self.latent_up(latents)
38
+ b, c, h, w = x.shape
39
+ sin, cos = rope_tables(h, w, head_dim=self.head_dim, device=x.device)
40
+ tokens = x.permute(0, 2, 3, 1).reshape(b, h * w, c)
41
+ for block in self.trunk:
42
+ tokens = block(tokens, sin, cos)
43
+ # The head reads the trunk's tokens as a contiguous NCHW map.
44
+ features = tokens.transpose(1, 2).reshape(b, c, h, w).contiguous()
45
+ return self.conv_up_head(features)
dinac3/encoder.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Frozen DINOv3-B behind a trained patch embedding, and the layer projection.
2
+
3
+ Every block's patch tokens (under the final LayerNorm) are standardized per
4
+ channel with fixed statistics, concatenated (12 x 768) and mapped by one linear
5
+ projection to the latent channels.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import TYPE_CHECKING, cast
11
+
12
+ import timm
13
+ import torch
14
+ from timm.layers.pos_embed_sincos import RotaryEmbeddingDinoV3
15
+ from torch import Tensor, nn
16
+
17
+ from .config import PATCH
18
+
19
+ if TYPE_CHECKING:
20
+ from timm.models.eva import Eva
21
+
22
+ from .config import Dinac3Config
23
+
24
+ IMAGENET_MEAN = (0.485, 0.456, 0.406)
25
+ IMAGENET_STD = (0.229, 0.224, 0.225)
26
+
27
+
28
+ def _require_supported_backbone(backbone: Eva) -> None:
29
+ """Fail fast if timm builds a backbone the explicit forward does not cover.
30
+
31
+ Raises:
32
+ ValueError: On absolute position embeddings, patch dropout, a pre-norm,
33
+ missing CLS/register tokens or a RoPE other than DINOv3's without
34
+ coordinate augmentation.
35
+ """
36
+
37
+ supported = (
38
+ backbone.pos_embed is None
39
+ and backbone.patch_drop is None
40
+ and isinstance(backbone.norm_pre, nn.Identity)
41
+ and backbone.cls_token is not None
42
+ and backbone.reg_token is not None
43
+ and isinstance(backbone.rope, RotaryEmbeddingDinoV3)
44
+ and not backbone.rope.aug_active
45
+ )
46
+ if not supported:
47
+ raise ValueError("Unsupported timm DINOv3 backbone structure for dinac3")
48
+
49
+
50
+ class Encoder(nn.Module):
51
+ """Images in [-1, 1] to raw latents ``[B, C, H / 16, W / 16]``."""
52
+
53
+ pixel_mean: Tensor
54
+ pixel_std: Tensor
55
+ layer_mean: Tensor
56
+ layer_std: Tensor
57
+
58
+ def __init__(self, config: Dinac3Config) -> None:
59
+ """Build the architecture only; every tensor comes from the artifact."""
60
+
61
+ super().__init__()
62
+ self.backbone = cast(
63
+ "Eva",
64
+ timm.create_model(
65
+ config.backbone.value,
66
+ pretrained=False,
67
+ num_classes=0,
68
+ dynamic_img_size=True,
69
+ dynamic_img_pad=True,
70
+ ),
71
+ )
72
+ _require_supported_backbone(self.backbone)
73
+ # timm keeps RoPE periods non-persistent; the artifact stores the exact
74
+ # values the training run used.
75
+ rope = cast("nn.Module", self.backbone.rope)
76
+ rope.register_buffer("periods", cast("Tensor", rope.periods), persistent=True)
77
+ self.prefix_tokens = self.backbone.num_prefix_tokens
78
+ width = self.backbone.embed_dim * len(self.backbone.blocks)
79
+ self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1))
80
+ self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1))
81
+ self.register_buffer("layer_mean", torch.zeros(width))
82
+ self.register_buffer("layer_std", torch.ones(width))
83
+ self.projection = nn.Linear(width, config.latent_channels)
84
+ self._rope_cache: dict[tuple[int, int, torch.device, torch.dtype], Tensor] = {}
85
+
86
+ def rope_embed(self, height: int, width: int) -> Tensor:
87
+ """The DINOv3 RoPE table of an image size, built once per token grid
88
+ (outside any compiled graph: timm builds coordinates on the CPU)."""
89
+
90
+ rope = cast("RotaryEmbeddingDinoV3", self.backbone.rope)
91
+ periods = cast("Tensor", rope.periods)
92
+ key = (height // PATCH, width // PATCH, periods.device, periods.dtype)
93
+ cached = self._rope_cache.get(key)
94
+ if cached is None:
95
+ with torch.inference_mode(False), torch.no_grad():
96
+ cached = rope.get_embed(shape=[key[0], key[1]])
97
+ self._rope_cache[key] = cached
98
+ return cached
99
+
100
+ def forward(self, images: Tensor, rope: Tensor) -> Tensor:
101
+ """Raw latents in the activation dtype; ``rope`` from :meth:`rope_embed`."""
102
+
103
+ dtype = self._prefix_tokens()[0].dtype
104
+ x = ((images.float().add(1.0).mul(0.5) - self.pixel_mean) / self.pixel_std).to(
105
+ dtype=dtype
106
+ )
107
+ features = self._block_features(x, rope)
108
+ tokens = torch.cat([feature.to(dtype=dtype) for feature in features], dim=-1)
109
+ standardized = (tokens.float() - self.layer_mean) / self.layer_std
110
+ latents = self.projection(standardized)
111
+ b, _, height, width = images.shape
112
+ return latents.transpose(1, 2).reshape(b, -1, height // PATCH, width // PATCH)
113
+
114
+ def _block_features(self, x: Tensor, rope: Tensor) -> list[Tensor]:
115
+ """Final-normed patch tokens of every block (timm's intermediates path)."""
116
+
117
+ backbone = self.backbone
118
+ tokens = backbone.patch_embed(x)
119
+ b, _, _, c = tokens.shape
120
+ tokens = tokens.view(b, -1, c)
121
+ cls_token, reg_token = self._prefix_tokens()
122
+ tokens = torch.cat(
123
+ [cls_token.expand(b, -1, -1), reg_token.expand(b, -1, -1), tokens], dim=1
124
+ )
125
+ features: list[Tensor] = []
126
+ for block in backbone.blocks:
127
+ tokens = block(tokens, rope=rope)
128
+ features.append(backbone.norm(tokens)[:, self.prefix_tokens :])
129
+ return features
130
+
131
+ def _prefix_tokens(self) -> tuple[Tensor, Tensor]:
132
+ """The CLS and register tokens (present: checked at construction)."""
133
+
134
+ cls_token, reg_token = self.backbone.cls_token, self.backbone.reg_token
135
+ if cls_token is None or reg_token is None:
136
+ raise RuntimeError("The DINOv3 backbone lost its CLS or register tokens")
137
+ return cls_token, reg_token
dinac3/layers.py ADDED
@@ -0,0 +1,97 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Normalization, MLP and convolution primitives with the training arithmetic.
2
+
3
+ Each op reproduces the training model's dtype handling under BF16 autocast
4
+ (reductions in FP32, results in the activation dtype), so BF16 inference matches
5
+ the training precision bit for bit.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import torch
11
+ from torch import Tensor, nn
12
+ from torch.nn import functional as F
13
+
14
+ NORM_EPS = 1e-6
15
+
16
+
17
+ class RMSNorm(nn.Module):
18
+ """RMS normalization over the last dimension, optional per-channel gain.
19
+
20
+ The gain is cast to the input dtype and the result returned in it.
21
+ """
22
+
23
+ def __init__(self, dim: int, *, affine: bool) -> None:
24
+ """Build the norm; ``affine`` adds a gain initialized to one."""
25
+
26
+ super().__init__()
27
+ self.dim = dim
28
+ self.weight = nn.Parameter(torch.ones(dim)) if affine else None
29
+
30
+ def forward(self, x: Tensor) -> Tensor:
31
+ """Normalize ``x`` ``[..., dim]``."""
32
+
33
+ weight = None if self.weight is None else self.weight.to(dtype=x.dtype)
34
+ return F.rms_norm(x, (self.dim,), weight, NORM_EPS).to(dtype=x.dtype)
35
+
36
+
37
+ class ChannelRMSNorm(nn.Module):
38
+ """Per-position RMS normalization over the channels of ``[B, C, H, W]``,
39
+ reduced in FP32 and applied in the input dtype, with an optional
40
+ per-channel gain and bias (cast to the input dtype)."""
41
+
42
+ def __init__(self, channels: int, *, affine: bool) -> None:
43
+ """Build the norm; ``affine`` adds a gain (one) and a bias (zero)."""
44
+
45
+ super().__init__()
46
+ self.weight = nn.Parameter(torch.ones(channels)) if affine else None
47
+ self.bias = nn.Parameter(torch.zeros(channels)) if affine else None
48
+
49
+ def forward(self, x: Tensor) -> Tensor:
50
+ """Normalize each position's channel vector."""
51
+
52
+ mean_square = torch.mean(
53
+ torch.square(x), dim=1, keepdim=True, dtype=torch.float32
54
+ )
55
+ y = x * torch.rsqrt(mean_square + NORM_EPS).to(dtype=x.dtype)
56
+ if self.weight is not None and self.bias is not None:
57
+ y = y * self.weight.view(1, -1, 1, 1).to(dtype=x.dtype)
58
+ y = y + self.bias.view(1, -1, 1, 1).to(dtype=x.dtype)
59
+ return y
60
+
61
+
62
+ class Mlp(nn.Module):
63
+ """Two linear layers with a GELU between them, both with biases."""
64
+
65
+ def __init__(self, dim: int, hidden: int) -> None:
66
+ """Allocate the up and down projections."""
67
+
68
+ super().__init__()
69
+ self.up = nn.Linear(dim, hidden)
70
+ self.down = nn.Linear(hidden, dim)
71
+
72
+ def forward(self, x: Tensor) -> Tensor:
73
+ """``down(gelu(up(x)))``."""
74
+
75
+ return self.down(F.gelu(self.up(x)))
76
+
77
+
78
+ def pointwise_conv(module: nn.Conv2d, x: Tensor) -> Tensor:
79
+ """A 1x1 convolution as a channel matmul (as the training graph runs it)."""
80
+
81
+ y = torch.einsum("bchw,oc->bohw", x, module.weight.flatten(1))
82
+ if module.bias is not None:
83
+ y = y + module.bias.view(1, -1, 1, 1).to(dtype=y.dtype)
84
+ return y
85
+
86
+
87
+ def same_conv(module: nn.Conv2d, x: Tensor) -> Tensor:
88
+ """A stride-1, same-padded convolution with ``module``'s parameters."""
89
+
90
+ return F.conv2d(
91
+ x,
92
+ module.weight,
93
+ module.bias,
94
+ stride=1,
95
+ padding=module.kernel_size[0] // 2,
96
+ groups=module.groups,
97
+ )
dinac3/model.py ADDED
@@ -0,0 +1,370 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Loading and inference API of dinac3 autoencoders."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ from collections.abc import Callable
7
+ from pathlib import Path
8
+ from typing import NoReturn
9
+
10
+ import torch
11
+ from safetensors.torch import load_file
12
+ from torch import Tensor, nn
13
+ from torch.fx.experimental import _config as fx_config
14
+
15
+ from .config import PATCH, Dinac3Config
16
+ from .decoder import Decoder
17
+ from .encoder import Encoder
18
+ from .precision import require_storage_policy
19
+
20
+ CONFIG_FILENAME = "config.json"
21
+ WEIGHTS_FILENAME = "model.safetensors"
22
+ SUPPORTED_DTYPES = (torch.bfloat16, torch.float32)
23
+
24
+
25
+ def _looks_like_local_path(name: str) -> bool:
26
+ """Whether a string names a filesystem path rather than a Hub ``org/name``
27
+ id: it starts with ``.``, ``/`` or ``~``, uses a path separator other than
28
+ the id's single ``/``, or has more than one ``/``."""
29
+
30
+ other_separator = os.sep != "/" and os.sep in name
31
+ return name.startswith((".", "/", "~")) or other_separator or name.count("/") > 1
32
+
33
+
34
+ def resolve_model_dir(path_or_repo: str | Path, *, revision: str | None) -> Path:
35
+ """A local artifact directory, or a Hugging Face Hub snapshot of a repo id.
36
+
37
+ A ``Path`` is always local. A ``str`` is local when it names an existing
38
+ directory; a string that looks like a path but names none is an error;
39
+ any other string is a Hub repository id, of which only ``config.json`` and
40
+ ``model.safetensors`` are downloaded.
41
+
42
+ Raises:
43
+ FileNotFoundError: If a local path does not name a directory.
44
+ ValueError: If ``revision`` is given for a local directory.
45
+ """
46
+
47
+ match path_or_repo:
48
+ case Path() as directory:
49
+ local = directory.expanduser()
50
+ case str() as name if Path(name).expanduser().is_dir():
51
+ local = Path(name).expanduser()
52
+ case str() as name if _looks_like_local_path(name):
53
+ raise FileNotFoundError(f"Local model path not found: {name}")
54
+ case str() as repo_id:
55
+ from huggingface_hub import snapshot_download
56
+
57
+ return Path(
58
+ snapshot_download(
59
+ repo_id,
60
+ revision=revision,
61
+ allow_patterns=[CONFIG_FILENAME, WEIGHTS_FILENAME],
62
+ )
63
+ )
64
+ case other:
65
+ raise TypeError(f"Expected a str or Path, got {type(other)}")
66
+ if not local.is_dir():
67
+ raise FileNotFoundError(f"Model directory does not exist: {local}")
68
+ if revision is not None:
69
+ raise ValueError(f"revision applies to Hub repositories, not to {local}")
70
+ return local
71
+
72
+
73
+ def compile_dynamic(call: Callable[..., Tensor]) -> Callable[..., Tensor]:
74
+ """``torch.compile`` with dynamic shapes and without duck sizing.
75
+
76
+ Dynamo otherwise assumes sizes that happen to be equal on the first call
77
+ (an image width and a token count, a batch and a grid side) stay equal, and
78
+ recompiles at the first shape that breaks the coincidence. Dimensions of
79
+ size 1 are still specialized: a batch of 1, or a 16-pixel image side (a
80
+ latent grid side of 1), compiles one more graph per such pattern.
81
+
82
+ Relies on PyTorch's private ``torch.fx.experimental._config.use_duck_shape``
83
+ (tested with PyTorch 2.13).
84
+ """
85
+
86
+ compiled = torch.compile(call, dynamic=True, fullgraph=True)
87
+
88
+ def traced(*args: Tensor) -> Tensor:
89
+ """Run the compiled call with duck sizing off (a per-thread setting)."""
90
+
91
+ with fx_config.patch(use_duck_shape=False): # ty: ignore[unresolved-attribute]
92
+ return compiled(*args)
93
+
94
+ return traced
95
+
96
+
97
+ class Dinac3(nn.Module):
98
+ """Deterministic image autoencoder with a DINOv3-aligned latent.
99
+
100
+ Images are RGB in [-1, 1] with height and width multiples of 16. The latent
101
+ has ``config.latent_channels`` channels at stride 16: free (detail)
102
+ channels first, the ``config.semantic_channels`` semantic channels last.
103
+ ``encode``/``decode`` use the whitened latent (zero mean, unit variance per
104
+ channel under the training distribution); ``encode_raw``/``decode_raw`` the
105
+ raw one.
106
+ """
107
+
108
+ latent_mean: Tensor
109
+ latent_var: Tensor
110
+
111
+ def __init__(self, config: Dinac3Config) -> None:
112
+ """Build the architecture without weights (no network access)."""
113
+
114
+ super().__init__()
115
+ self.config = config
116
+ self.encoder = Encoder(config)
117
+ self.decoder = Decoder(config)
118
+ self.register_buffer("latent_mean", torch.zeros(config.latent_channels))
119
+ self.register_buffer("latent_var", torch.ones(config.latent_channels))
120
+ self.compute_dtype = torch.bfloat16
121
+ self._encode_call: Callable[[Tensor, Tensor], Tensor] = self.encoder.forward
122
+ self._decode_call: Callable[[Tensor], Tensor] = self.decoder.forward
123
+
124
+ @classmethod
125
+ def from_pretrained(
126
+ cls,
127
+ path_or_repo: str | Path,
128
+ *,
129
+ device: torch.device | str,
130
+ dtype: torch.dtype = torch.bfloat16,
131
+ compile_encoder: bool = True,
132
+ compile_decoder: bool = True,
133
+ revision: str | None = None,
134
+ ) -> Dinac3:
135
+ """Load an artifact strictly and prepare it for CUDA inference.
136
+
137
+ Args:
138
+ path_or_repo: Local directory (``Path``, or ``str`` naming one) or
139
+ Hugging Face Hub repository id.
140
+ device: CUDA device.
141
+ dtype: ``torch.bfloat16`` (the stored mixed precision, BF16
142
+ autocast: the training precision) or ``torch.float32`` (every
143
+ tensor upcast, no autocast).
144
+ compile_encoder: ``torch.compile`` the encoder with dynamic shapes.
145
+ compile_decoder: ``torch.compile`` the decoder with dynamic shapes.
146
+ revision: Hub revision, for repository ids only.
147
+
148
+ Raises:
149
+ ValueError: On an unsupported dtype or device, or tensors stored in
150
+ another dtype than the storage policy.
151
+ """
152
+
153
+ directory = resolve_model_dir(path_or_repo, revision=revision)
154
+ config = Dinac3Config.load(directory / CONFIG_FILENAME)
155
+ state = load_file(str(directory / WEIGHTS_FILENAME))
156
+ require_storage_policy(state)
157
+ # Every tensor comes from the artifact: build on the meta device and
158
+ # let the loaded tensors become the parameters and buffers.
159
+ with torch.device("meta"):
160
+ model = cls(config)
161
+ model.load_state_dict(state, strict=True, assign=True)
162
+ model.prepare(
163
+ device=torch.device(device),
164
+ dtype=dtype,
165
+ compile_encoder=compile_encoder,
166
+ compile_decoder=compile_decoder,
167
+ )
168
+ return model
169
+
170
+ def prepare(
171
+ self,
172
+ *,
173
+ device: torch.device,
174
+ dtype: torch.dtype,
175
+ compile_encoder: bool,
176
+ compile_decoder: bool,
177
+ ) -> None:
178
+ """Move to ``device``, set the compute precision and compile once.
179
+
180
+ Raises:
181
+ ValueError: On a non-CUDA device or an unsupported dtype.
182
+ """
183
+
184
+ if device.type != "cuda":
185
+ raise ValueError("dinac3 inference requires a CUDA device")
186
+ if dtype not in SUPPORTED_DTYPES:
187
+ raise ValueError(f"dtype must be one of {SUPPORTED_DTYPES}, got {dtype}")
188
+ self.to(device=device)
189
+ if dtype == torch.float32:
190
+ for tensor in (*self.parameters(), *self.buffers()):
191
+ tensor.data = tensor.data.float()
192
+ self.compute_dtype = dtype
193
+ self.eval().requires_grad_(False)
194
+ self.decoder.conv_up_head.fold_kernels()
195
+ self._encode_call = (
196
+ compile_dynamic(self.encoder.forward)
197
+ if compile_encoder
198
+ else self.encoder.forward
199
+ )
200
+ self._decode_call = (
201
+ compile_dynamic(self.decoder.forward)
202
+ if compile_decoder
203
+ else self.decoder.forward
204
+ )
205
+
206
+ @torch.inference_mode()
207
+ def encode(self, images: Tensor) -> Tensor:
208
+ """Whitened FP32 latents ``[B, C, H / 16, W / 16]`` of [-1, 1] images."""
209
+
210
+ return self.whiten(self.encode_raw(images))
211
+
212
+ @torch.inference_mode()
213
+ def encode_raw(self, images: Tensor) -> Tensor:
214
+ """Raw (unwhitened) FP32 latents of [-1, 1] images."""
215
+
216
+ self._require_images(images)
217
+ rope = self.encoder.rope_embed(images.shape[-2], images.shape[-1])
218
+ with self._autocast():
219
+ return self._encode_call(images, rope).float()
220
+
221
+ @torch.inference_mode()
222
+ def decode(self, latents: Tensor, height: int, width: int) -> Tensor:
223
+ """FP32 RGB ``[B, 3, height, width]`` (unclamped, about [-1, 1]) from
224
+ whitened latents, in one deterministic pass."""
225
+
226
+ if (height, width) != (latents.shape[-2] * PATCH, latents.shape[-1] * PATCH):
227
+ raise ValueError(
228
+ f"{height}x{width} is not the 16x image of a "
229
+ f"{latents.shape[-2]}x{latents.shape[-1]} latent grid"
230
+ )
231
+ return self.decode_raw(self.dewhiten(latents))
232
+
233
+ @torch.inference_mode()
234
+ def decode_raw(self, latents: Tensor) -> Tensor:
235
+ """FP32 RGB from raw (unwhitened) latents, in one deterministic pass."""
236
+
237
+ self._require_latents(latents)
238
+ with self._autocast():
239
+ images = self._decode_call(latents.float())
240
+ return images.float().contiguous()
241
+
242
+ def whiten(self, latents: Tensor) -> Tensor:
243
+ """Raw to whitened latents, in FP32."""
244
+
245
+ mean, std = self._latent_stats()
246
+ return (latents.float() - mean) / std
247
+
248
+ def dewhiten(self, latents: Tensor) -> Tensor:
249
+ """Whitened to raw latents, in FP32."""
250
+
251
+ mean, std = self._latent_stats()
252
+ return latents.float() * std + mean
253
+
254
+ def semantic_channels(self, latents: Tensor) -> Tensor:
255
+ """The DINOv3-aligned channels: the last ``semantic_channels``."""
256
+
257
+ return latents[:, self.config.free_channels :]
258
+
259
+ def free_channels(self, latents: Tensor) -> Tensor:
260
+ """The free (detail) channels: all but the semantic ones."""
261
+
262
+ return latents[:, : self.config.free_channels]
263
+
264
+ def to(self, *args: object, **kwargs: object) -> Dinac3:
265
+ """Move to a device; dtype casts are rejected (see :meth:`_reject_cast`)."""
266
+
267
+ casts = (
268
+ "dtype" in kwargs
269
+ or "tensor" in kwargs
270
+ or any(isinstance(arg, torch.dtype | Tensor) for arg in args)
271
+ )
272
+ if casts:
273
+ self._reject_cast()
274
+ return super().to(*args, **kwargs) # ty: ignore[no-matching-overload]
275
+
276
+ def half(self) -> NoReturn:
277
+ """Rejected: see :meth:`_reject_cast`."""
278
+
279
+ self._reject_cast()
280
+
281
+ def bfloat16(self) -> NoReturn:
282
+ """Rejected: see :meth:`_reject_cast`."""
283
+
284
+ self._reject_cast()
285
+
286
+ def float(self) -> NoReturn:
287
+ """Rejected: see :meth:`_reject_cast`."""
288
+
289
+ self._reject_cast()
290
+
291
+ def double(self) -> NoReturn:
292
+ """Rejected: see :meth:`_reject_cast`."""
293
+
294
+ self._reject_cast()
295
+
296
+ def type(self, dst_type: object) -> NoReturn:
297
+ """Rejected: see :meth:`_reject_cast`."""
298
+
299
+ del dst_type
300
+ self._reject_cast()
301
+
302
+ def _reject_cast(self) -> NoReturn:
303
+ """Module-wide dtype casts would break the mixed storage policy.
304
+
305
+ Raises:
306
+ TypeError: Always, pointing to ``from_pretrained(dtype=...)``.
307
+ """
308
+
309
+ raise TypeError(
310
+ "dinac3 keeps a mixed storage policy (BF16 weights, FP32 residual-path, "
311
+ "statistics and RoPE tensors); a module-wide dtype cast would break it. "
312
+ "Choose the precision with Dinac3.from_pretrained(..., "
313
+ "dtype=torch.bfloat16 | torch.float32)."
314
+ )
315
+
316
+ def _latent_stats(self) -> tuple[Tensor, Tensor]:
317
+ """Per-channel ``(mean, std)`` ``[1, C, 1, 1]`` in FP32."""
318
+
319
+ mean = self.latent_mean.float().view(1, -1, 1, 1)
320
+ var = self.latent_var.float().view(1, -1, 1, 1)
321
+ return mean, torch.sqrt(var + self.config.latent_stats_eps)
322
+
323
+ def _autocast(self) -> torch.autocast:
324
+ """BF16 autocast for BF16 inference, disabled for FP32."""
325
+
326
+ return torch.autocast(
327
+ "cuda", dtype=torch.bfloat16, enabled=self.compute_dtype == torch.bfloat16
328
+ )
329
+
330
+ def _require_images(self, images: Tensor) -> None:
331
+ """Validate images at the API boundary, outside compiled graphs.
332
+
333
+ Raises:
334
+ ValueError: On a wrong shape, size, dtype or device.
335
+ """
336
+
337
+ if images.ndim != 4 or images.shape[1] != 3 or images.shape[0] < 1:
338
+ raise ValueError("images must have shape [B, 3, H, W] with B >= 1")
339
+ height, width = images.shape[-2:]
340
+ if min(height, width) < PATCH or height % PATCH or width % PATCH:
341
+ raise ValueError(f"Image sides must be positive multiples of {PATCH}")
342
+ self._require_tensor(images)
343
+
344
+ def _require_latents(self, latents: Tensor) -> None:
345
+ """Validate latents at the API boundary.
346
+
347
+ Raises:
348
+ ValueError: On a wrong shape, dtype or device.
349
+ """
350
+
351
+ if latents.ndim != 4 or latents.shape[1] != self.config.latent_channels:
352
+ raise ValueError(
353
+ f"latents must have shape [B, {self.config.latent_channels}, h, w]"
354
+ )
355
+ if min(latents.shape[0], *latents.shape[-2:]) < 1:
356
+ raise ValueError("Latent batch and grid sides must be positive")
357
+ self._require_tensor(latents)
358
+
359
+ def _require_tensor(self, tensor: Tensor) -> None:
360
+ """Require a floating tensor on the model's CUDA device.
361
+
362
+ Raises:
363
+ ValueError: If the tensor is elsewhere or not floating point.
364
+ """
365
+
366
+ device = self.latent_mean.device
367
+ if tensor.device != device:
368
+ raise ValueError(f"Inputs must be on the model device {device}")
369
+ if not tensor.is_floating_point():
370
+ raise ValueError("Inputs must be floating-point tensors")
dinac3/precision.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Storage dtype of every exported tensor.
2
+
3
+ BF16 for every weight that BF16 autocast consumes in BF16 anyway (matrices,
4
+ convolutions, biases, and norm gains and biases applied to BF16 activations):
5
+ storing them in BF16 gives bit-identical outputs. FP32 where the value takes
6
+ part in FP32 arithmetic: the frozen DINOv3 residual stream (CLS and register
7
+ tokens, layer scales, LayerNorms), its RoPE periods, the input and
8
+ layer-standardization statistics, the latent statistics used for whitening,
9
+ and the upsamplers' 3x3 weights and fold matrices (the 4x4 transposed kernel is
10
+ folded from them outside autocast, in their storage dtype).
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import re
16
+ from typing import TYPE_CHECKING
17
+
18
+ import torch
19
+
20
+ if TYPE_CHECKING:
21
+ from collections.abc import Mapping
22
+
23
+ from torch import Tensor
24
+
25
+ _FP32_NAMES = frozenset(
26
+ {
27
+ "latent_mean",
28
+ "latent_var",
29
+ "encoder.pixel_mean",
30
+ "encoder.pixel_std",
31
+ "encoder.layer_mean",
32
+ "encoder.layer_std",
33
+ "encoder.backbone.cls_token",
34
+ "encoder.backbone.reg_token",
35
+ "encoder.backbone.rope.periods",
36
+ "encoder.backbone.norm.weight",
37
+ "encoder.backbone.norm.bias",
38
+ }
39
+ )
40
+ _UPSAMPLER = re.compile(
41
+ r"decoder\.conv_up_head\.(handoff_upsample|upsamples\.\d+)\.(conv\.weight|fold)"
42
+ )
43
+ _BACKBONE_BLOCK_PREFIX = "encoder.backbone.blocks."
44
+ _FP32_BLOCK_SUFFIXES = (
45
+ ".gamma_1",
46
+ ".gamma_2",
47
+ ".norm1.weight",
48
+ ".norm1.bias",
49
+ ".norm2.weight",
50
+ ".norm2.bias",
51
+ )
52
+
53
+
54
+ def storage_dtype(name: str) -> torch.dtype:
55
+ """BF16 or FP32 storage of one exported tensor, by its state-dict name."""
56
+
57
+ fp32 = (
58
+ name in _FP32_NAMES
59
+ or (
60
+ name.startswith(_BACKBONE_BLOCK_PREFIX)
61
+ and name.endswith(_FP32_BLOCK_SUFFIXES)
62
+ )
63
+ or _UPSAMPLER.fullmatch(name) is not None
64
+ )
65
+ return torch.float32 if fp32 else torch.bfloat16
66
+
67
+
68
+ def require_storage_policy(state: Mapping[str, Tensor]) -> None:
69
+ """Require every tensor of an artifact to be stored in its policy dtype.
70
+
71
+ Raises:
72
+ ValueError: Naming every tensor stored in another dtype.
73
+ """
74
+
75
+ wrong = sorted(
76
+ name for name, tensor in state.items() if tensor.dtype != storage_dtype(name)
77
+ )
78
+ if wrong:
79
+ raise ValueError(f"Tensors not stored in their policy dtype: {wrong}")
dinac3/trunk.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The decoder's ViT trunk: RMS-sandwich blocks with 2-D RoPE."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch
6
+ from torch import Tensor, nn
7
+ from torch.nn import functional as F
8
+
9
+ from .layers import Mlp, RMSNorm
10
+
11
+ ROPE_BASE = 10_000.0
12
+
13
+
14
+ def rope_periods(head_dim: int, device: torch.device) -> Tensor:
15
+ """FP32 periods ``ROPE_BASE ** (2 i / (head_dim / 2))`` of the
16
+ ``head_dim / 4`` rotation frequencies per axis."""
17
+
18
+ exponents = (
19
+ 2.0
20
+ * torch.arange(head_dim // 4, device=device, dtype=torch.float32)
21
+ / (head_dim // 2)
22
+ )
23
+ return torch.tensor(ROPE_BASE, device=device, dtype=torch.float32) ** exponents
24
+
25
+
26
+ def rope_tables(
27
+ height: int, width: int, *, head_dim: int, device: torch.device
28
+ ) -> tuple[Tensor, Tensor]:
29
+ """FP32 ``(sin, cos)`` ``[H * W, head_dim]`` of axial 2-D RoPE on the token
30
+ grid: unnormalized patch-index coordinates (row, column), the
31
+ :func:`rope_periods`, adjacent-pair layout."""
32
+
33
+ periods = rope_periods(head_dim, device)
34
+ rows = torch.arange(0.0, float(height), device=device, dtype=torch.float32)
35
+ cols = torch.arange(0.0, float(width), device=device, dtype=torch.float32)
36
+ coords = torch.stack(torch.meshgrid(rows, cols, indexing="ij"), dim=-1).flatten(
37
+ 0, 1
38
+ )
39
+ angles = 1.0 * coords[:, :, None] / periods[None, None, :].expand(1, 2, -1)
40
+ angles = angles.repeat_interleave(2, dim=-1).flatten(1, 2)
41
+ return torch.sin(angles), torch.cos(angles)
42
+
43
+
44
+ def _rotate_pairs(x: Tensor) -> Tensor:
45
+ """Rotate consecutive channel pairs: ``(a, b) -> (-b, a)``."""
46
+
47
+ pairs = x.reshape(*x.shape[:-1], x.shape[-1] // 2, 2)
48
+ return torch.stack((-pairs[..., 1], pairs[..., 0]), dim=-1).reshape_as(x)
49
+
50
+
51
+ def apply_rope(x: Tensor, sin: Tensor, cos: Tensor) -> Tensor:
52
+ """Rotate ``[B, heads, N, head_dim]`` queries or keys in the tables' dtype."""
53
+
54
+ rotated = x.to(dtype=sin.dtype)
55
+ sin_b = sin.view(1, 1, *sin.shape)
56
+ cos_b = cos.view(1, 1, *cos.shape)
57
+ return ((rotated * cos_b) + (_rotate_pairs(rotated) * sin_b)).to(dtype=x.dtype)
58
+
59
+
60
+ class Attention(nn.Module):
61
+ """Bias-free multi-head self-attention with RMS-normalized queries and keys."""
62
+
63
+ def __init__(self, width: int, head_dim: int) -> None:
64
+ """Allocate the fused QKV and output projections and the q/k norms."""
65
+
66
+ super().__init__()
67
+ self.heads = width // head_dim
68
+ self.head_dim = head_dim
69
+ self.qkv = nn.Linear(width, 3 * width, bias=False)
70
+ self.proj_out = nn.Linear(width, width, bias=False)
71
+ self.q_norm = RMSNorm(head_dim, affine=True)
72
+ self.k_norm = RMSNorm(head_dim, affine=True)
73
+
74
+ def forward(self, x: Tensor, sin: Tensor, cos: Tensor) -> Tensor:
75
+ """Attend over all tokens of ``[B, N, width]``."""
76
+
77
+ b, n, width = x.shape
78
+ q, k, v = self.qkv(x).chunk(3, dim=-1)
79
+ shape = (b, n, self.heads, self.head_dim)
80
+ q = q.view(shape).transpose(1, 2).contiguous()
81
+ k = k.view(shape).transpose(1, 2).contiguous()
82
+ v = v.view(shape).transpose(1, 2).contiguous()
83
+ q = apply_rope(self.q_norm(q), sin, cos)
84
+ k = apply_rope(self.k_norm(k), sin, cos)
85
+ out = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=False)
86
+ return self.proj_out(out.transpose(1, 2).contiguous().view(b, n, width))
87
+
88
+
89
+ class DitBlock(nn.Module):
90
+ """ViT block: RMSNorm before and after attention and MLP."""
91
+
92
+ def __init__(self, width: int, head_dim: int, mlp_ratio: float) -> None:
93
+ """Allocate attention, MLP and the four norms."""
94
+
95
+ super().__init__()
96
+ self.attn_norm1 = RMSNorm(width, affine=True)
97
+ self.attn_norm2 = RMSNorm(width, affine=True)
98
+ self.mlp_norm1 = RMSNorm(width, affine=True)
99
+ self.mlp_norm2 = RMSNorm(width, affine=True)
100
+ self.attn = Attention(width, head_dim)
101
+ self.mlp = Mlp(width, int(mlp_ratio * width))
102
+
103
+ def forward(self, x: Tensor, sin: Tensor, cos: Tensor) -> Tensor:
104
+ """Residual attention then residual MLP on ``[B, N, width]`` tokens."""
105
+
106
+ x = x + self.attn_norm2(self.attn(self.attn_norm1(x), sin, cos))
107
+ return x + self.mlp_norm2(self.mlp(self.mlp_norm1(x)))
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2696c98a87d595afc386f068ab5f9fa99d6ff1f71e521145bc31371612b62856
3
+ size 336457454
pyproject.toml ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [build-system]
2
+ requires = ["setuptools>=68"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "dinac3"
7
+ version = "1.0"
8
+ description = "dinac3: deterministic image autoencoders with a DINOv3-aligned latent"
9
+ requires-python = ">=3.10"
10
+ dependencies = ["torch>=2.13", "timm==1.0.26", "safetensors>=0.5", "huggingface-hub>=0.28"]
11
+
12
+ [tool.setuptools.packages.find]
13
+ include = ["dinac3*"]