Commit ·
75ff4df
0
Parent(s):
dinac3_96 v1.0
Browse files- .gitattributes +35 -0
- ATTRIBUTION.md +50 -0
- LICENSE +201 -0
- LICENSE-DINOV3.md +66 -0
- README.md +244 -0
- TECHNICAL_REPORT.md +608 -0
- config.json +19 -0
- dinac3/__init__.py +6 -0
- dinac3/config.py +195 -0
- dinac3/conv_up_head.py +180 -0
- dinac3/decoder.py +45 -0
- dinac3/encoder.py +137 -0
- dinac3/layers.py +97 -0
- dinac3/model.py +370 -0
- dinac3/precision.py +79 -0
- dinac3/trunk.py +107 -0
- model.safetensors +3 -0
- pyproject.toml +13 -0
.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*"]
|