linoyts HF Staff commited on
Commit
8af79b6
·
verified ·
1 Parent(s): de67b13

Upload folder using huggingface_hub

Browse files
transformer/config.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_class_name": "MiniMaxH3PrunedTransformer3DModel",
3
+ "_diffusers_version": "0.40.0.dev0",
4
+ "_name_or_path": "multimodalart/MiniMax-H3-Pruned",
5
+ "adaln_rank": 8,
6
+ "attention_head_dim": 128,
7
+ "audio_in_channels": 32,
8
+ "auto_map": {
9
+ "AutoModel": "modeling_minimax_h3_pruned.MiniMaxH3PrunedTransformer3DModel"
10
+ },
11
+ "ffn_dim": 14336,
12
+ "final_norm_eps": 1e-05,
13
+ "freq_dim": 256,
14
+ "hidden_size": 5376,
15
+ "in_channels": 24,
16
+ "norm_eps": 1e-05,
17
+ "num_attention_heads": 56,
18
+ "num_layers": 50,
19
+ "num_refiner_layers": 2,
20
+ "patch_size": [
21
+ 1,
22
+ 2,
23
+ 2
24
+ ],
25
+ "qk_norm_eps": 1e-05,
26
+ "rope_freq_dim": 16,
27
+ "rope_theta": 10000.0,
28
+ "text_dim": 5120,
29
+ "time_embed_dim": 2688,
30
+ "time_embed_hidden_dim": 5376,
31
+ "time_table_size": 1025
32
+ }
transformer/diffusion_pytorch_model-00001-of-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e60a580e061f2ebadac22ffe9dffcadcc78155b89a65507f755b3b8957f588a3
3
+ size 9942657792
transformer/diffusion_pytorch_model-00002-of-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6e5c58a131303c65528dfa33717780e5fc676c46c275c0858ebc71b6c5d201c7
3
+ size 9736326144
transformer/diffusion_pytorch_model-00003-of-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87a0886428eb84a905d3b15e4e02bd72a08dbe0e9b4aa489a3390f4caf7d8a03
3
+ size 9967526304
transformer/diffusion_pytorch_model-00004-of-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:534f2baf9f9fc651f67f8d2f963238befe1924f423795c3a07e8f2caed008c6e
3
+ size 9967536424
transformer/diffusion_pytorch_model-00005-of-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e151b0ad2364b6123621a91c4e08f52ae2d608b081966f692be38d295250de7b
3
+ size 621489776
transformer/diffusion_pytorch_model.safetensors.index.json ADDED
@@ -0,0 +1,644 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "metadata": {
3
+ "total_size": 40235463200
4
+ },
5
+ "weight_map": {
6
+ "adaln_basis": "diffusion_pytorch_model-00001-of-00005.safetensors",
7
+ "adaln_mean": "diffusion_pytorch_model-00001-of-00005.safetensors",
8
+ "audio_proj_in.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
9
+ "audio_proj_in.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
10
+ "audio_proj_out.bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
11
+ "audio_proj_out.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
12
+ "context_embedder.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
13
+ "context_embedder.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
14
+ "norm_out.folded_bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
15
+ "norm_out.linear.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
16
+ "norm_out.norm.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
17
+ "proj_in.bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
18
+ "proj_in.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
19
+ "proj_out.bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
20
+ "proj_out.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
21
+ "time_embedder.table": "diffusion_pytorch_model-00001-of-00005.safetensors",
22
+ "token_refiner.final_norm.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
23
+ "token_refiner.refiner_blocks.0.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
24
+ "token_refiner.refiner_blocks.0.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
25
+ "token_refiner.refiner_blocks.0.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
26
+ "token_refiner.refiner_blocks.0.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
27
+ "token_refiner.refiner_blocks.0.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
28
+ "token_refiner.refiner_blocks.0.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
29
+ "token_refiner.refiner_blocks.0.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
30
+ "token_refiner.refiner_blocks.0.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
31
+ "token_refiner.refiner_blocks.0.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
32
+ "token_refiner.refiner_blocks.0.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
33
+ "token_refiner.refiner_blocks.1.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
34
+ "token_refiner.refiner_blocks.1.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
35
+ "token_refiner.refiner_blocks.1.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
36
+ "token_refiner.refiner_blocks.1.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
37
+ "token_refiner.refiner_blocks.1.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
38
+ "token_refiner.refiner_blocks.1.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
39
+ "token_refiner.refiner_blocks.1.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
40
+ "token_refiner.refiner_blocks.1.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
41
+ "token_refiner.refiner_blocks.1.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
42
+ "token_refiner.refiner_blocks.1.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
43
+ "transformer_blocks.0.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
44
+ "transformer_blocks.0.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
45
+ "transformer_blocks.0.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
46
+ "transformer_blocks.0.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
47
+ "transformer_blocks.0.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
48
+ "transformer_blocks.0.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
49
+ "transformer_blocks.0.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
50
+ "transformer_blocks.0.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
51
+ "transformer_blocks.0.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
52
+ "transformer_blocks.0.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
53
+ "transformer_blocks.0.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
54
+ "transformer_blocks.0.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
55
+ "transformer_blocks.1.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
56
+ "transformer_blocks.1.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
57
+ "transformer_blocks.1.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
58
+ "transformer_blocks.1.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
59
+ "transformer_blocks.1.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
60
+ "transformer_blocks.1.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
61
+ "transformer_blocks.1.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
62
+ "transformer_blocks.1.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
63
+ "transformer_blocks.1.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
64
+ "transformer_blocks.1.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
65
+ "transformer_blocks.1.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
66
+ "transformer_blocks.1.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
67
+ "transformer_blocks.10.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
68
+ "transformer_blocks.10.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
69
+ "transformer_blocks.10.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
70
+ "transformer_blocks.10.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
71
+ "transformer_blocks.10.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
72
+ "transformer_blocks.10.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
73
+ "transformer_blocks.10.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
74
+ "transformer_blocks.10.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
75
+ "transformer_blocks.10.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
76
+ "transformer_blocks.10.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
77
+ "transformer_blocks.10.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
78
+ "transformer_blocks.10.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
79
+ "transformer_blocks.11.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
80
+ "transformer_blocks.11.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
81
+ "transformer_blocks.11.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
82
+ "transformer_blocks.11.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
83
+ "transformer_blocks.11.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
84
+ "transformer_blocks.11.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
85
+ "transformer_blocks.11.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
86
+ "transformer_blocks.11.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
87
+ "transformer_blocks.11.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
88
+ "transformer_blocks.11.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
89
+ "transformer_blocks.11.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
90
+ "transformer_blocks.11.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
91
+ "transformer_blocks.12.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
92
+ "transformer_blocks.12.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
93
+ "transformer_blocks.12.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
94
+ "transformer_blocks.12.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
95
+ "transformer_blocks.12.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
96
+ "transformer_blocks.12.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
97
+ "transformer_blocks.12.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
98
+ "transformer_blocks.12.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
99
+ "transformer_blocks.12.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
100
+ "transformer_blocks.12.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
101
+ "transformer_blocks.12.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
102
+ "transformer_blocks.12.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
103
+ "transformer_blocks.13.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
104
+ "transformer_blocks.13.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
105
+ "transformer_blocks.13.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
106
+ "transformer_blocks.13.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
107
+ "transformer_blocks.13.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
108
+ "transformer_blocks.13.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
109
+ "transformer_blocks.13.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
110
+ "transformer_blocks.13.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
111
+ "transformer_blocks.13.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
112
+ "transformer_blocks.13.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
113
+ "transformer_blocks.13.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
114
+ "transformer_blocks.13.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
115
+ "transformer_blocks.14.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
116
+ "transformer_blocks.14.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
117
+ "transformer_blocks.14.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
118
+ "transformer_blocks.14.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
119
+ "transformer_blocks.14.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
120
+ "transformer_blocks.14.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
121
+ "transformer_blocks.14.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
122
+ "transformer_blocks.14.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
123
+ "transformer_blocks.14.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
124
+ "transformer_blocks.14.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
125
+ "transformer_blocks.14.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
126
+ "transformer_blocks.14.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
127
+ "transformer_blocks.15.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
128
+ "transformer_blocks.15.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
129
+ "transformer_blocks.15.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
130
+ "transformer_blocks.15.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
131
+ "transformer_blocks.15.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
132
+ "transformer_blocks.15.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
133
+ "transformer_blocks.15.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
134
+ "transformer_blocks.15.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
135
+ "transformer_blocks.15.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
136
+ "transformer_blocks.15.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
137
+ "transformer_blocks.15.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
138
+ "transformer_blocks.15.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
139
+ "transformer_blocks.16.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
140
+ "transformer_blocks.16.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
141
+ "transformer_blocks.16.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
142
+ "transformer_blocks.16.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
143
+ "transformer_blocks.16.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
144
+ "transformer_blocks.16.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
145
+ "transformer_blocks.16.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
146
+ "transformer_blocks.16.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
147
+ "transformer_blocks.16.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
148
+ "transformer_blocks.16.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
149
+ "transformer_blocks.16.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
150
+ "transformer_blocks.16.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
151
+ "transformer_blocks.17.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
152
+ "transformer_blocks.17.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
153
+ "transformer_blocks.17.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
154
+ "transformer_blocks.17.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
155
+ "transformer_blocks.17.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
156
+ "transformer_blocks.17.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
157
+ "transformer_blocks.17.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
158
+ "transformer_blocks.17.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
159
+ "transformer_blocks.17.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
160
+ "transformer_blocks.17.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
161
+ "transformer_blocks.17.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
162
+ "transformer_blocks.17.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
163
+ "transformer_blocks.18.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
164
+ "transformer_blocks.18.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
165
+ "transformer_blocks.18.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
166
+ "transformer_blocks.18.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
167
+ "transformer_blocks.18.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
168
+ "transformer_blocks.18.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
169
+ "transformer_blocks.18.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
170
+ "transformer_blocks.18.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
171
+ "transformer_blocks.18.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
172
+ "transformer_blocks.18.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
173
+ "transformer_blocks.18.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
174
+ "transformer_blocks.18.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
175
+ "transformer_blocks.19.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
176
+ "transformer_blocks.19.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
177
+ "transformer_blocks.19.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
178
+ "transformer_blocks.19.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
179
+ "transformer_blocks.19.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
180
+ "transformer_blocks.19.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
181
+ "transformer_blocks.19.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
182
+ "transformer_blocks.19.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
183
+ "transformer_blocks.19.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
184
+ "transformer_blocks.19.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
185
+ "transformer_blocks.19.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
186
+ "transformer_blocks.19.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
187
+ "transformer_blocks.2.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
188
+ "transformer_blocks.2.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
189
+ "transformer_blocks.2.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
190
+ "transformer_blocks.2.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
191
+ "transformer_blocks.2.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
192
+ "transformer_blocks.2.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
193
+ "transformer_blocks.2.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
194
+ "transformer_blocks.2.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
195
+ "transformer_blocks.2.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
196
+ "transformer_blocks.2.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
197
+ "transformer_blocks.2.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
198
+ "transformer_blocks.2.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
199
+ "transformer_blocks.20.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
200
+ "transformer_blocks.20.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
201
+ "transformer_blocks.20.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
202
+ "transformer_blocks.20.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
203
+ "transformer_blocks.20.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
204
+ "transformer_blocks.20.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
205
+ "transformer_blocks.20.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
206
+ "transformer_blocks.20.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
207
+ "transformer_blocks.20.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
208
+ "transformer_blocks.20.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
209
+ "transformer_blocks.20.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
210
+ "transformer_blocks.20.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
211
+ "transformer_blocks.21.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
212
+ "transformer_blocks.21.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
213
+ "transformer_blocks.21.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
214
+ "transformer_blocks.21.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
215
+ "transformer_blocks.21.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
216
+ "transformer_blocks.21.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
217
+ "transformer_blocks.21.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
218
+ "transformer_blocks.21.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
219
+ "transformer_blocks.21.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
220
+ "transformer_blocks.21.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
221
+ "transformer_blocks.21.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
222
+ "transformer_blocks.21.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
223
+ "transformer_blocks.22.adaln_proj.folded_bias": "diffusion_pytorch_model-00002-of-00005.safetensors",
224
+ "transformer_blocks.22.adaln_proj.linear.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
225
+ "transformer_blocks.22.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
226
+ "transformer_blocks.22.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
227
+ "transformer_blocks.22.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
228
+ "transformer_blocks.22.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
229
+ "transformer_blocks.22.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
230
+ "transformer_blocks.22.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
231
+ "transformer_blocks.22.ff.net.0.proj.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
232
+ "transformer_blocks.22.ff.net.2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
233
+ "transformer_blocks.22.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
234
+ "transformer_blocks.22.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
235
+ "transformer_blocks.23.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
236
+ "transformer_blocks.23.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
237
+ "transformer_blocks.23.attn.norm_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
238
+ "transformer_blocks.23.attn.norm_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
239
+ "transformer_blocks.23.attn.to_k.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
240
+ "transformer_blocks.23.attn.to_out.0.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
241
+ "transformer_blocks.23.attn.to_q.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
242
+ "transformer_blocks.23.attn.to_v.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
243
+ "transformer_blocks.23.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
244
+ "transformer_blocks.23.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
245
+ "transformer_blocks.23.norm1.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
246
+ "transformer_blocks.23.norm2.weight": "diffusion_pytorch_model-00002-of-00005.safetensors",
247
+ "transformer_blocks.24.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
248
+ "transformer_blocks.24.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
249
+ "transformer_blocks.24.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
250
+ "transformer_blocks.24.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
251
+ "transformer_blocks.24.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
252
+ "transformer_blocks.24.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
253
+ "transformer_blocks.24.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
254
+ "transformer_blocks.24.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
255
+ "transformer_blocks.24.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
256
+ "transformer_blocks.24.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
257
+ "transformer_blocks.24.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
258
+ "transformer_blocks.24.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
259
+ "transformer_blocks.25.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
260
+ "transformer_blocks.25.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
261
+ "transformer_blocks.25.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
262
+ "transformer_blocks.25.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
263
+ "transformer_blocks.25.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
264
+ "transformer_blocks.25.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
265
+ "transformer_blocks.25.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
266
+ "transformer_blocks.25.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
267
+ "transformer_blocks.25.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
268
+ "transformer_blocks.25.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
269
+ "transformer_blocks.25.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
270
+ "transformer_blocks.25.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
271
+ "transformer_blocks.26.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
272
+ "transformer_blocks.26.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
273
+ "transformer_blocks.26.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
274
+ "transformer_blocks.26.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
275
+ "transformer_blocks.26.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
276
+ "transformer_blocks.26.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
277
+ "transformer_blocks.26.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
278
+ "transformer_blocks.26.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
279
+ "transformer_blocks.26.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
280
+ "transformer_blocks.26.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
281
+ "transformer_blocks.26.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
282
+ "transformer_blocks.26.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
283
+ "transformer_blocks.27.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
284
+ "transformer_blocks.27.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
285
+ "transformer_blocks.27.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
286
+ "transformer_blocks.27.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
287
+ "transformer_blocks.27.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
288
+ "transformer_blocks.27.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
289
+ "transformer_blocks.27.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
290
+ "transformer_blocks.27.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
291
+ "transformer_blocks.27.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
292
+ "transformer_blocks.27.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
293
+ "transformer_blocks.27.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
294
+ "transformer_blocks.27.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
295
+ "transformer_blocks.28.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
296
+ "transformer_blocks.28.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
297
+ "transformer_blocks.28.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
298
+ "transformer_blocks.28.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
299
+ "transformer_blocks.28.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
300
+ "transformer_blocks.28.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
301
+ "transformer_blocks.28.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
302
+ "transformer_blocks.28.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
303
+ "transformer_blocks.28.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
304
+ "transformer_blocks.28.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
305
+ "transformer_blocks.28.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
306
+ "transformer_blocks.28.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
307
+ "transformer_blocks.29.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
308
+ "transformer_blocks.29.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
309
+ "transformer_blocks.29.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
310
+ "transformer_blocks.29.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
311
+ "transformer_blocks.29.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
312
+ "transformer_blocks.29.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
313
+ "transformer_blocks.29.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
314
+ "transformer_blocks.29.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
315
+ "transformer_blocks.29.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
316
+ "transformer_blocks.29.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
317
+ "transformer_blocks.29.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
318
+ "transformer_blocks.29.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
319
+ "transformer_blocks.3.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
320
+ "transformer_blocks.3.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
321
+ "transformer_blocks.3.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
322
+ "transformer_blocks.3.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
323
+ "transformer_blocks.3.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
324
+ "transformer_blocks.3.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
325
+ "transformer_blocks.3.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
326
+ "transformer_blocks.3.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
327
+ "transformer_blocks.3.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
328
+ "transformer_blocks.3.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
329
+ "transformer_blocks.3.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
330
+ "transformer_blocks.3.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
331
+ "transformer_blocks.30.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
332
+ "transformer_blocks.30.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
333
+ "transformer_blocks.30.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
334
+ "transformer_blocks.30.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
335
+ "transformer_blocks.30.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
336
+ "transformer_blocks.30.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
337
+ "transformer_blocks.30.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
338
+ "transformer_blocks.30.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
339
+ "transformer_blocks.30.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
340
+ "transformer_blocks.30.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
341
+ "transformer_blocks.30.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
342
+ "transformer_blocks.30.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
343
+ "transformer_blocks.31.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
344
+ "transformer_blocks.31.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
345
+ "transformer_blocks.31.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
346
+ "transformer_blocks.31.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
347
+ "transformer_blocks.31.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
348
+ "transformer_blocks.31.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
349
+ "transformer_blocks.31.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
350
+ "transformer_blocks.31.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
351
+ "transformer_blocks.31.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
352
+ "transformer_blocks.31.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
353
+ "transformer_blocks.31.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
354
+ "transformer_blocks.31.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
355
+ "transformer_blocks.32.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
356
+ "transformer_blocks.32.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
357
+ "transformer_blocks.32.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
358
+ "transformer_blocks.32.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
359
+ "transformer_blocks.32.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
360
+ "transformer_blocks.32.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
361
+ "transformer_blocks.32.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
362
+ "transformer_blocks.32.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
363
+ "transformer_blocks.32.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
364
+ "transformer_blocks.32.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
365
+ "transformer_blocks.32.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
366
+ "transformer_blocks.32.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
367
+ "transformer_blocks.33.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
368
+ "transformer_blocks.33.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
369
+ "transformer_blocks.33.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
370
+ "transformer_blocks.33.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
371
+ "transformer_blocks.33.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
372
+ "transformer_blocks.33.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
373
+ "transformer_blocks.33.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
374
+ "transformer_blocks.33.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
375
+ "transformer_blocks.33.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
376
+ "transformer_blocks.33.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
377
+ "transformer_blocks.33.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
378
+ "transformer_blocks.33.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
379
+ "transformer_blocks.34.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
380
+ "transformer_blocks.34.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
381
+ "transformer_blocks.34.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
382
+ "transformer_blocks.34.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
383
+ "transformer_blocks.34.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
384
+ "transformer_blocks.34.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
385
+ "transformer_blocks.34.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
386
+ "transformer_blocks.34.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
387
+ "transformer_blocks.34.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
388
+ "transformer_blocks.34.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
389
+ "transformer_blocks.34.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
390
+ "transformer_blocks.34.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
391
+ "transformer_blocks.35.adaln_proj.folded_bias": "diffusion_pytorch_model-00003-of-00005.safetensors",
392
+ "transformer_blocks.35.adaln_proj.linear.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
393
+ "transformer_blocks.35.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
394
+ "transformer_blocks.35.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
395
+ "transformer_blocks.35.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
396
+ "transformer_blocks.35.attn.to_out.0.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
397
+ "transformer_blocks.35.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
398
+ "transformer_blocks.35.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
399
+ "transformer_blocks.35.ff.net.0.proj.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
400
+ "transformer_blocks.35.ff.net.2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
401
+ "transformer_blocks.35.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
402
+ "transformer_blocks.35.norm2.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
403
+ "transformer_blocks.36.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
404
+ "transformer_blocks.36.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
405
+ "transformer_blocks.36.attn.norm_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
406
+ "transformer_blocks.36.attn.norm_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
407
+ "transformer_blocks.36.attn.to_k.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
408
+ "transformer_blocks.36.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
409
+ "transformer_blocks.36.attn.to_q.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
410
+ "transformer_blocks.36.attn.to_v.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
411
+ "transformer_blocks.36.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
412
+ "transformer_blocks.36.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
413
+ "transformer_blocks.36.norm1.weight": "diffusion_pytorch_model-00003-of-00005.safetensors",
414
+ "transformer_blocks.36.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
415
+ "transformer_blocks.37.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
416
+ "transformer_blocks.37.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
417
+ "transformer_blocks.37.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
418
+ "transformer_blocks.37.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
419
+ "transformer_blocks.37.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
420
+ "transformer_blocks.37.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
421
+ "transformer_blocks.37.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
422
+ "transformer_blocks.37.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
423
+ "transformer_blocks.37.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
424
+ "transformer_blocks.37.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
425
+ "transformer_blocks.37.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
426
+ "transformer_blocks.37.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
427
+ "transformer_blocks.38.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
428
+ "transformer_blocks.38.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
429
+ "transformer_blocks.38.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
430
+ "transformer_blocks.38.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
431
+ "transformer_blocks.38.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
432
+ "transformer_blocks.38.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
433
+ "transformer_blocks.38.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
434
+ "transformer_blocks.38.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
435
+ "transformer_blocks.38.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
436
+ "transformer_blocks.38.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
437
+ "transformer_blocks.38.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
438
+ "transformer_blocks.38.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
439
+ "transformer_blocks.39.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
440
+ "transformer_blocks.39.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
441
+ "transformer_blocks.39.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
442
+ "transformer_blocks.39.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
443
+ "transformer_blocks.39.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
444
+ "transformer_blocks.39.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
445
+ "transformer_blocks.39.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
446
+ "transformer_blocks.39.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
447
+ "transformer_blocks.39.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
448
+ "transformer_blocks.39.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
449
+ "transformer_blocks.39.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
450
+ "transformer_blocks.39.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
451
+ "transformer_blocks.4.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
452
+ "transformer_blocks.4.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
453
+ "transformer_blocks.4.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
454
+ "transformer_blocks.4.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
455
+ "transformer_blocks.4.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
456
+ "transformer_blocks.4.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
457
+ "transformer_blocks.4.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
458
+ "transformer_blocks.4.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
459
+ "transformer_blocks.4.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
460
+ "transformer_blocks.4.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
461
+ "transformer_blocks.4.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
462
+ "transformer_blocks.4.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
463
+ "transformer_blocks.40.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
464
+ "transformer_blocks.40.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
465
+ "transformer_blocks.40.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
466
+ "transformer_blocks.40.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
467
+ "transformer_blocks.40.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
468
+ "transformer_blocks.40.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
469
+ "transformer_blocks.40.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
470
+ "transformer_blocks.40.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
471
+ "transformer_blocks.40.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
472
+ "transformer_blocks.40.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
473
+ "transformer_blocks.40.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
474
+ "transformer_blocks.40.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
475
+ "transformer_blocks.41.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
476
+ "transformer_blocks.41.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
477
+ "transformer_blocks.41.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
478
+ "transformer_blocks.41.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
479
+ "transformer_blocks.41.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
480
+ "transformer_blocks.41.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
481
+ "transformer_blocks.41.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
482
+ "transformer_blocks.41.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
483
+ "transformer_blocks.41.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
484
+ "transformer_blocks.41.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
485
+ "transformer_blocks.41.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
486
+ "transformer_blocks.41.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
487
+ "transformer_blocks.42.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
488
+ "transformer_blocks.42.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
489
+ "transformer_blocks.42.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
490
+ "transformer_blocks.42.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
491
+ "transformer_blocks.42.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
492
+ "transformer_blocks.42.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
493
+ "transformer_blocks.42.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
494
+ "transformer_blocks.42.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
495
+ "transformer_blocks.42.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
496
+ "transformer_blocks.42.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
497
+ "transformer_blocks.42.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
498
+ "transformer_blocks.42.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
499
+ "transformer_blocks.43.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
500
+ "transformer_blocks.43.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
501
+ "transformer_blocks.43.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
502
+ "transformer_blocks.43.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
503
+ "transformer_blocks.43.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
504
+ "transformer_blocks.43.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
505
+ "transformer_blocks.43.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
506
+ "transformer_blocks.43.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
507
+ "transformer_blocks.43.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
508
+ "transformer_blocks.43.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
509
+ "transformer_blocks.43.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
510
+ "transformer_blocks.43.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
511
+ "transformer_blocks.44.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
512
+ "transformer_blocks.44.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
513
+ "transformer_blocks.44.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
514
+ "transformer_blocks.44.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
515
+ "transformer_blocks.44.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
516
+ "transformer_blocks.44.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
517
+ "transformer_blocks.44.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
518
+ "transformer_blocks.44.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
519
+ "transformer_blocks.44.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
520
+ "transformer_blocks.44.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
521
+ "transformer_blocks.44.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
522
+ "transformer_blocks.44.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
523
+ "transformer_blocks.45.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
524
+ "transformer_blocks.45.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
525
+ "transformer_blocks.45.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
526
+ "transformer_blocks.45.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
527
+ "transformer_blocks.45.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
528
+ "transformer_blocks.45.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
529
+ "transformer_blocks.45.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
530
+ "transformer_blocks.45.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
531
+ "transformer_blocks.45.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
532
+ "transformer_blocks.45.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
533
+ "transformer_blocks.45.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
534
+ "transformer_blocks.45.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
535
+ "transformer_blocks.46.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
536
+ "transformer_blocks.46.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
537
+ "transformer_blocks.46.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
538
+ "transformer_blocks.46.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
539
+ "transformer_blocks.46.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
540
+ "transformer_blocks.46.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
541
+ "transformer_blocks.46.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
542
+ "transformer_blocks.46.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
543
+ "transformer_blocks.46.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
544
+ "transformer_blocks.46.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
545
+ "transformer_blocks.46.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
546
+ "transformer_blocks.46.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
547
+ "transformer_blocks.47.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
548
+ "transformer_blocks.47.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
549
+ "transformer_blocks.47.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
550
+ "transformer_blocks.47.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
551
+ "transformer_blocks.47.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
552
+ "transformer_blocks.47.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
553
+ "transformer_blocks.47.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
554
+ "transformer_blocks.47.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
555
+ "transformer_blocks.47.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
556
+ "transformer_blocks.47.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
557
+ "transformer_blocks.47.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
558
+ "transformer_blocks.47.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
559
+ "transformer_blocks.48.adaln_proj.folded_bias": "diffusion_pytorch_model-00004-of-00005.safetensors",
560
+ "transformer_blocks.48.adaln_proj.linear.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
561
+ "transformer_blocks.48.attn.norm_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
562
+ "transformer_blocks.48.attn.norm_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
563
+ "transformer_blocks.48.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
564
+ "transformer_blocks.48.attn.to_out.0.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
565
+ "transformer_blocks.48.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
566
+ "transformer_blocks.48.attn.to_v.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
567
+ "transformer_blocks.48.ff.net.0.proj.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
568
+ "transformer_blocks.48.ff.net.2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
569
+ "transformer_blocks.48.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
570
+ "transformer_blocks.48.norm2.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
571
+ "transformer_blocks.49.adaln_proj.folded_bias": "diffusion_pytorch_model-00005-of-00005.safetensors",
572
+ "transformer_blocks.49.adaln_proj.linear.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
573
+ "transformer_blocks.49.attn.norm_k.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
574
+ "transformer_blocks.49.attn.norm_q.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
575
+ "transformer_blocks.49.attn.to_k.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
576
+ "transformer_blocks.49.attn.to_out.0.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
577
+ "transformer_blocks.49.attn.to_q.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
578
+ "transformer_blocks.49.attn.to_v.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
579
+ "transformer_blocks.49.ff.net.0.proj.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
580
+ "transformer_blocks.49.ff.net.2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
581
+ "transformer_blocks.49.norm1.weight": "diffusion_pytorch_model-00004-of-00005.safetensors",
582
+ "transformer_blocks.49.norm2.weight": "diffusion_pytorch_model-00005-of-00005.safetensors",
583
+ "transformer_blocks.5.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
584
+ "transformer_blocks.5.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
585
+ "transformer_blocks.5.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
586
+ "transformer_blocks.5.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
587
+ "transformer_blocks.5.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
588
+ "transformer_blocks.5.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
589
+ "transformer_blocks.5.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
590
+ "transformer_blocks.5.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
591
+ "transformer_blocks.5.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
592
+ "transformer_blocks.5.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
593
+ "transformer_blocks.5.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
594
+ "transformer_blocks.5.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
595
+ "transformer_blocks.6.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
596
+ "transformer_blocks.6.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
597
+ "transformer_blocks.6.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
598
+ "transformer_blocks.6.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
599
+ "transformer_blocks.6.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
600
+ "transformer_blocks.6.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
601
+ "transformer_blocks.6.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
602
+ "transformer_blocks.6.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
603
+ "transformer_blocks.6.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
604
+ "transformer_blocks.6.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
605
+ "transformer_blocks.6.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
606
+ "transformer_blocks.6.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
607
+ "transformer_blocks.7.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
608
+ "transformer_blocks.7.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
609
+ "transformer_blocks.7.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
610
+ "transformer_blocks.7.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
611
+ "transformer_blocks.7.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
612
+ "transformer_blocks.7.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
613
+ "transformer_blocks.7.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
614
+ "transformer_blocks.7.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
615
+ "transformer_blocks.7.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
616
+ "transformer_blocks.7.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
617
+ "transformer_blocks.7.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
618
+ "transformer_blocks.7.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
619
+ "transformer_blocks.8.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
620
+ "transformer_blocks.8.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
621
+ "transformer_blocks.8.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
622
+ "transformer_blocks.8.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
623
+ "transformer_blocks.8.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
624
+ "transformer_blocks.8.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
625
+ "transformer_blocks.8.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
626
+ "transformer_blocks.8.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
627
+ "transformer_blocks.8.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
628
+ "transformer_blocks.8.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
629
+ "transformer_blocks.8.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
630
+ "transformer_blocks.8.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
631
+ "transformer_blocks.9.adaln_proj.folded_bias": "diffusion_pytorch_model-00001-of-00005.safetensors",
632
+ "transformer_blocks.9.adaln_proj.linear.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
633
+ "transformer_blocks.9.attn.norm_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
634
+ "transformer_blocks.9.attn.norm_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
635
+ "transformer_blocks.9.attn.to_k.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
636
+ "transformer_blocks.9.attn.to_out.0.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
637
+ "transformer_blocks.9.attn.to_q.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
638
+ "transformer_blocks.9.attn.to_v.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
639
+ "transformer_blocks.9.ff.net.0.proj.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
640
+ "transformer_blocks.9.ff.net.2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
641
+ "transformer_blocks.9.norm1.weight": "diffusion_pytorch_model-00001-of-00005.safetensors",
642
+ "transformer_blocks.9.norm2.weight": "diffusion_pytorch_model-00001-of-00005.safetensors"
643
+ }
644
+ }
transformer/modeling_minimax_h3_pruned.py ADDED
@@ -0,0 +1,602 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2025 The MiniMax Team and The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """AdaLN-pruned MiniMax-H3 transformer.
15
+
16
+ Everything outside the timestep path is inherited from `MiniMaxH3Transformer3DModel`: the attention, the blocks, the
17
+ token refiner, the output heads and `forward` itself are the released implementation, unmodified. Only what feeds the
18
+ AdaLN projections changes.
19
+
20
+ Two other things this file adds. `enable_convrot` / `quantize_8bit`: an opt-in Hadamard conditioning of the
21
+ attention and feed-forward linears that makes 8-bit *compute* (int8 or fp8 dynamic activations, via torchao) land
22
+ within a rounding step of bfloat16. And `load_lora_adapter`, overridden to project a LoRA trained against the
23
+ *released* 2688-wide AdaLN projections onto these 8-wide ones. Both are inert unless used, so the plain pruned path
24
+ is byte for byte what it was.
25
+ """
26
+
27
+ import math
28
+ import re
29
+ from types import SimpleNamespace
30
+
31
+ import torch
32
+ import torch.nn as nn
33
+ import torch.nn.functional as F
34
+ from diffusers.configuration_utils import register_to_config
35
+ from diffusers.models.modeling_utils import get_parameter_dtype
36
+ from diffusers.models.transformers.transformer_minimax_h3 import (
37
+ MINIMAX_H3_MODALITY_NUM,
38
+ MiniMaxH3RotaryPosEmbed,
39
+ MiniMaxH3TokenRefiner,
40
+ MiniMaxH3Transformer3DModel,
41
+ MiniMaxH3TransformerBlock,
42
+ )
43
+ from diffusers.utils import logging
44
+
45
+
46
+ logger = logging.get_logger(__name__)
47
+
48
+ _HADAMARD_CACHE: dict = {}
49
+
50
+ # The 51 AdaLN projections, as a LoRA state dict names them, with or without a `transformer.` / `transformer_ref.`
51
+ # component prefix. Group 1 is the module path the model itself knows the projection by.
52
+ ADALN_LORA_A_KEY = re.compile(r"(?:^|\.)((?:transformer_blocks\.\d+\.adaln_proj|norm_out)\.linear)\.lora_A\.weight$")
53
+
54
+ # The linears ConvRot conditions: the block stack's attention and feed-forward projections, and nothing else.
55
+ # This is the set ComfyUI's `*_int8_convrot` checkpoints quantize (they carry one fused `qkv_proj`; the three
56
+ # split projections here share an input, so rotating each is the same transform). The AdaLN path, the patch
57
+ # projections, the output heads, the token refiner and every norm stay high precision.
58
+ CONVROT_SUFFIXES = ("attn.to_q", "attn.to_k", "attn.to_v", "attn.to_out.0", "ff.net.0.proj", "ff.net.2")
59
+
60
+
61
+ def _hadamard(size: int, device, dtype) -> torch.Tensor:
62
+ r"""Normalized *regular* Hadamard matrix of `size` - symmetric, orthogonal, therefore self-inverse.
63
+
64
+ Kronecker powers of the regular order-4 seed, which is why `size` must be a power of 4. Identical, entry for
65
+ entry, to the matrix `comfy_kitchen` builds for its `convrot` kernels.
66
+ """
67
+ key = (size, str(device), dtype)
68
+ if key not in _HADAMARD_CACHE:
69
+ if size < 4 or (size & (size - 1)) != 0 or math.log(size, 4) % 1 != 0:
70
+ raise ValueError(f"ConvRot group size must be a power of 4, got {size}")
71
+ seed = torch.tensor([[1, 1, 1, -1], [1, 1, -1, 1], [1, -1, 1, 1], [-1, 1, 1, 1]], dtype=dtype, device=device)
72
+ matrix, width = seed, 4
73
+ while width < size:
74
+ matrix, width = torch.kron(matrix, seed), width * 4
75
+ _HADAMARD_CACHE[key] = matrix / size**0.5
76
+ return _HADAMARD_CACHE[key]
77
+
78
+
79
+ def _rotate(tensor: torch.Tensor, group_size: int) -> torch.Tensor:
80
+ r"""`tensor @ blockdiag(H, ..., H)` over the last dimension, in the tensor's own dtype."""
81
+ shape = tensor.shape
82
+ if shape[-1] % group_size != 0:
83
+ raise ValueError(f"{shape[-1]} features is not a multiple of the ConvRot group size {group_size}")
84
+ matrix = _hadamard(group_size, tensor.device, tensor.dtype)
85
+ return torch.matmul(tensor.reshape(-1, shape[-1] // group_size, group_size), matrix).reshape(shape)
86
+
87
+
88
+ class MiniMaxH3ConvRotLinear(nn.Linear):
89
+ r"""`nn.Linear` that Hadamard-rotates its input; its weight already carries the same rotation.
90
+
91
+ `enable_convrot` bakes `W <- W @ H` into the weight, so computing `(x @ H) @ (W @ H)^T` returns `x @ W^T`
92
+ exactly - `H` is symmetric *and* orthogonal, so it is its own inverse. Nothing about the model's function
93
+ changes. What changes is the distribution a quantizer downstream of this module sees: every coordinate of
94
+ `x @ H` is a +-1 combination of `group_size` input channels, so a single outlier channel no longer sets the
95
+ scale for its whole row. That is all ConvRot is, and because both sides of one matmul are rotated back to
96
+ back, nothing has to commute with the AdaLN modulation.
97
+
98
+ Deliberately a bare `nn.Linear` subclass with no parameters of its own: `torchao.quantize_` still converts
99
+ it, the `state_dict` keys are unchanged, and PEFT wraps it as a `base_layer` - which leaves a LoRA's branch
100
+ reading the *unrotated* input, the basis LoRAs are trained in.
101
+ """
102
+
103
+ convrot_groupsize: int = 0
104
+
105
+ def forward(self, input: torch.Tensor) -> torch.Tensor: # noqa: A002
106
+ return F.linear(_rotate(input, self.convrot_groupsize), self.weight, self.bias)
107
+
108
+
109
+ class MiniMaxH3PrunedTimeEmbedder(nn.Module):
110
+ r"""The released timestep MLP, replaced by an interpolated table of AdaLN coordinates.
111
+
112
+ Every AdaLN projection in the released model consumes `silu(time_embedder(time_proj(t)))`, which depends on the
113
+ scalar timestep alone: over `t` in `[0, 1]` it traces a one-dimensional curve in `R^{time_embed_dim}`. A rank-8
114
+ affine subspace reproduces that curve to about 1.5e-5 relative RMS, so only the curve's coordinates in that
115
+ subspace are stored - sampled on a uniform grid of `table_size` timesteps and linearly interpolated in between.
116
+ The subspace offset is folded into the AdaLN biases and its basis into the AdaLN weights, which is why the
117
+ projections take an `adaln_rank`-wide input here instead of `time_embed_dim`.
118
+
119
+ The module stands in for `time_proj` and `time_embedder` together: it consumes the raw timestep, so the released
120
+ `forward` needs no change once `time_proj` is an identity.
121
+ """
122
+
123
+ def __init__(self, table_size: int = 1025, adaln_rank: int = 8) -> None:
124
+ super().__init__()
125
+ self.register_buffer("table", torch.zeros(table_size, adaln_rank), persistent=True)
126
+
127
+ @property
128
+ def linear_1(self):
129
+ r"""Answer the dtype question a diffusers before #14398 asks of the released timestep MLP.
130
+
131
+ `MiniMaxH3Transformer3DModel.forward` aligns the timestep it passes in with the timestep path's own dtype.
132
+ Since [#14398](https://github.com/huggingface/diffusers/pull/14398) it asks
133
+ `get_parameter_dtype(self.time_embedder)`, which on this module returns the table's float32 - there is no
134
+ parameter here, only the buffer. Before it, it read `self.time_embedder.linear_1.weight.dtype`, the first
135
+ `Linear` of the released timestep MLP this table replaces. The two questions have one answer on a pruned
136
+ checkpoint, so the older one is answered rather than raised: the table *is* the timestep path here.
137
+
138
+ A property returning a plain namespace, not a registered module: nothing about it reaches `state_dict`,
139
+ `named_modules`, `.to()`, a quantizer's scan or PEFT's target resolution, and it is read for a dtype and
140
+ never called. Delete it once every consumer runs a diffusers that carries #14398.
141
+ """
142
+ return SimpleNamespace(weight=self.table)
143
+
144
+ def forward(self, timestep: torch.Tensor) -> torch.Tensor:
145
+ table = self.table
146
+ steps = table.shape[0] - 1
147
+ position = timestep.to(table.dtype).flatten().clamp(0.0, 1.0) * steps
148
+ lower = position.floor().clamp(max=steps - 1).long()
149
+ weight = (position - lower).unsqueeze(-1)
150
+ return torch.lerp(table.index_select(0, lower), table.index_select(0, lower + 1), weight)
151
+
152
+
153
+ class MiniMaxH3PrunedTimeProj(nn.Module):
154
+ r"""Identity stand-in for `Timesteps`: the pruned time embedder indexes the raw timestep."""
155
+
156
+ def forward(self, timestep: torch.Tensor) -> torch.Tensor:
157
+ return timestep
158
+
159
+
160
+ class MiniMaxH3PrunedAdaLN(nn.Module):
161
+ r"""Shared by both pruned AdaLN modules: the folded float32 bias, plus the LoRA offsets that ride alongside it.
162
+
163
+ A LoRA trained against the released 2688-wide projection contributes `lora_B @ (lora_A @ x)` to the modulation.
164
+ With `x = mean + c @ basis`, that splits into a coordinate term the projected factors reproduce and a *constant*
165
+ term, `lora_B @ (lora_A @ mean)`, which no `Linear(8 -> out)` can express. That constant is held here, as a
166
+ per-adapter float32 buffer added to `folded_bias` - the same place, and the same precision, the fold's own
167
+ constant term lives in. It is deliberately not the projection's `bias`: rounding it into bfloat16 would spend a
168
+ full rounding step of the modulation on a term that is most of what the adapter does to the AdaLN path.
169
+ """
170
+
171
+ def __init__(self, out_features: int) -> None:
172
+ super().__init__()
173
+ self.register_buffer("folded_bias", torch.zeros(out_features), persistent=True)
174
+ # `{adapter name: buffer attribute}`. Plain state, not a submodule: the buffers themselves are what move
175
+ # with the module, and being non-persistent they stay out of the checkpoint, as an adapter should.
176
+ self._lora_adaln_offsets: dict[str, str] = {}
177
+
178
+ def register_lora_adaln_offset(self, adapter_name: str, offset: torch.Tensor) -> None:
179
+ r"""Attach one adapter's constant term, on the device and in the precision `folded_bias` is kept in."""
180
+ attribute = self._lora_adaln_offsets.get(adapter_name)
181
+ if attribute is None:
182
+ attribute = f"lora_adaln_offset_{len(self._lora_adaln_offsets)}"
183
+ self._lora_adaln_offsets[adapter_name] = attribute
184
+ value = offset.to(device=self.folded_bias.device, dtype=torch.float32)
185
+ self.register_buffer(attribute, value, persistent=False)
186
+
187
+ def lora_adaln_bias(self) -> torch.Tensor:
188
+ r"""`folded_bias` plus every active adapter's constant term at its current scaling.
189
+
190
+ PEFT owns everything this reads - `active_adapters`, `scaling`, `disable_adapters` - so the offsets follow
191
+ `set_adapters`, `disable_lora` and `delete_adapters` with no bookkeeping of their own. They apply whether or
192
+ not an adapter is merged: `fuse_lora` folds `lora_B @ lora_A` into the projection's weight, and there is
193
+ nowhere in a bias-free `Linear` for this term to be folded to.
194
+ """
195
+ bias = self.folded_bias
196
+ offsets = self._lora_adaln_offsets
197
+ if not offsets:
198
+ return bias
199
+ layer = self.linear
200
+ scaling = getattr(layer, "scaling", None)
201
+ if not isinstance(scaling, dict) or getattr(layer, "disable_adapters", False):
202
+ return bias
203
+ for adapter_name in getattr(layer, "active_adapters", ()):
204
+ attribute = offsets.get(adapter_name)
205
+ if attribute is None or adapter_name not in scaling:
206
+ continue
207
+ bias = bias + getattr(self, attribute) * float(scaling[adapter_name])
208
+ return bias
209
+
210
+
211
+ class MiniMaxH3PrunedAdaLayerNormModulation(MiniMaxH3PrunedAdaLN):
212
+ r"""`MiniMaxH3AdaLayerNormModulation` over the pruned timestep coordinates.
213
+
214
+ Two differences from the released module. It applies no `silu` - the table already holds the coordinates of the
215
+ activated curve. And the folded bias is a float32 buffer applied outside the projection rather than the
216
+ projection's own bias: it carries almost the entire modulation (the coordinate term contributes a few tenths of
217
+ it), so storing it in bfloat16 would put a full output-scale rounding step into every evaluation. Kept in
218
+ float32 it costs 0.4 MB per block and leaves the pruned AdaLN function closer to an exact float64 evaluation
219
+ than the released bfloat16 checkpoint's own arithmetic is.
220
+
221
+ `linear` stays a bias-free `nn.Linear` so PEFT wraps it exactly as it wraps the released projection.
222
+ """
223
+
224
+ def __init__(self, adaln_rank: int, hidden_size: int) -> None:
225
+ out_features = 6 * hidden_size * MINIMAX_H3_MODALITY_NUM
226
+ super().__init__(out_features)
227
+ self.hidden_size = hidden_size
228
+ self.linear = nn.Linear(adaln_rank, out_features, bias=False)
229
+
230
+ def forward(self, temb: torch.Tensor) -> tuple[torch.Tensor, ...]:
231
+ dtype = get_parameter_dtype(self.linear)
232
+ temb = self.linear(temb.to(dtype))
233
+ temb = (temb.float() + self.lora_adaln_bias()).to(dtype)
234
+ temb = temb.view(-1, 6 * self.hidden_size)
235
+ return temb.chunk(6, dim=-1)
236
+
237
+
238
+ class MiniMaxH3PrunedAdaLayerNormOut(MiniMaxH3PrunedAdaLN):
239
+ r"""`MiniMaxH3AdaLayerNormOut` over the pruned timestep coordinates; see the modulation module above."""
240
+
241
+ def __init__(self, hidden_size: int, adaln_rank: int, eps: float) -> None:
242
+ super().__init__(2 * hidden_size)
243
+ self.norm = nn.RMSNorm(hidden_size, eps=eps)
244
+ self.linear = nn.Linear(adaln_rank, 2 * hidden_size, bias=False)
245
+
246
+ def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor, timestep_indices: torch.Tensor) -> torch.Tensor:
247
+ dtype = get_parameter_dtype(self.linear)
248
+ temb = self.linear(temb.to(dtype))
249
+ shift, scale = (temb.float() + self.lora_adaln_bias()).to(dtype).chunk(2, dim=-1)
250
+ hidden_states = self.norm(hidden_states)
251
+ return hidden_states * (1.0 + scale.index_select(0, timestep_indices)) + shift.index_select(
252
+ 0, timestep_indices
253
+ )
254
+
255
+
256
+ class MiniMaxH3PrunedTransformer3DModel(MiniMaxH3Transformer3DModel):
257
+ r"""MiniMax-H3's DiT with the AdaLN input projections reduced to their reachable rank.
258
+
259
+ The released checkpoint spends 13.03B of its 33.14B parameters on the 50 per-block `adaln_proj.linear` matrices
260
+ plus `norm_out.linear`, all of which read the same 2688-wide timestep embedding. Because that embedding is a
261
+ function of the scalar timestep, its reachable set is a curve an 8-dimensional affine subspace covers to ~1.5e-5
262
+ relative RMS - far below one bfloat16 rounding step of the weights themselves. Folding the subspace into the
263
+ projections leaves an 8-wide input and removes 26 GB per partition.
264
+
265
+ Only what builds the timestep path differs from [`MiniMaxH3Transformer3DModel`]: `time_proj` becomes an identity,
266
+ `time_embedder` becomes [`MiniMaxH3PrunedTimeEmbedder`], and the AdaLN projections take `adaln_rank` inputs.
267
+ `forward` is inherited unchanged. The module names are the released ones, so LoRAs trained against a pruned
268
+ checkpoint - what the common trainers use by default - load natively, and `load_lora_adapter` projects LoRAs
269
+ trained against the released checkpoint's `time_embed_dim`-wide projections onto the same coordinates.
270
+
271
+ Args:
272
+ adaln_rank (`int`, defaults to `8`):
273
+ The width of the timestep coordinates every AdaLN projection consumes.
274
+ time_table_size (`int`, defaults to `1025`):
275
+ The number of uniformly spaced timesteps the coordinate table holds; values in between are interpolated
276
+ linearly.
277
+
278
+ Every other argument is [`MiniMaxH3Transformer3DModel`]'s and carries the same meaning. `freq_dim` and
279
+ `time_embed_hidden_dim` are kept in the config, unused, so a pruned config still records the shape of the
280
+ released timestep MLP it was folded from.
281
+ """
282
+
283
+ _supports_gradient_checkpointing = True
284
+ _no_split_modules = ["MiniMaxH3TransformerBlock", "MiniMaxH3TokenRefinerBlock", "MiniMaxH3PrunedAdaLayerNormOut"]
285
+ _repeated_blocks = ["MiniMaxH3TransformerBlock", "MiniMaxH3TokenRefinerBlock"]
286
+ _skip_layerwise_casting_patterns = ["norm"]
287
+ # The released checkpoint's mixed-precision split - patch projections, output heads and the timestep path in
288
+ # float32, the block stack in bfloat16 - plus the folded AdaLN biases, for the reason given on the modulation
289
+ # module. Entries are matched against the dot-separated segments of each parameter name.
290
+ _keep_in_fp32_modules = [
291
+ "proj_in",
292
+ "audio_proj_in",
293
+ "time_embedder",
294
+ "proj_out",
295
+ "audio_proj_out",
296
+ "rope",
297
+ "folded_bias",
298
+ "adaln_basis",
299
+ "adaln_mean",
300
+ ]
301
+
302
+ @register_to_config
303
+ def __init__(
304
+ self,
305
+ num_attention_heads: int = 56,
306
+ attention_head_dim: int = 128,
307
+ hidden_size: int = 5376,
308
+ num_layers: int = 50,
309
+ num_refiner_layers: int = 2,
310
+ ffn_dim: int = 14336,
311
+ in_channels: int = 24,
312
+ audio_in_channels: int = 32,
313
+ patch_size: tuple[int, int, int] = (1, 2, 2),
314
+ text_dim: int = 5120,
315
+ freq_dim: int = 256,
316
+ time_embed_hidden_dim: int = 5376,
317
+ time_embed_dim: int = 2688,
318
+ rope_freq_dim: int = 16,
319
+ rope_theta: float = 10000.0,
320
+ norm_eps: float = 1e-5,
321
+ qk_norm_eps: float = 1e-5,
322
+ final_norm_eps: float = 1e-5,
323
+ adaln_rank: int = 8,
324
+ time_table_size: int = 1025,
325
+ ) -> None:
326
+ # `MiniMaxH3Transformer3DModel.__init__` is itself wrapped by `register_to_config`, so calling it would
327
+ # register the released config over this one - and would allocate the 26 GB of AdaLN projections this class
328
+ # exists to avoid. The module tree is built here instead; everything but the timestep path is verbatim.
329
+ nn.Module.__init__(self)
330
+
331
+ video_patch_dim = in_channels * patch_size[0] * patch_size[1] * patch_size[2]
332
+
333
+ # 1. Per-modality input projections
334
+ self.proj_in = nn.Linear(video_patch_dim, hidden_size, bias=True)
335
+ self.audio_proj_in = nn.Linear(audio_in_channels, hidden_size, bias=True)
336
+ self.context_embedder = nn.Linear(text_dim, hidden_size, bias=True)
337
+
338
+ # 2. Timestep coordinates, shared by every AdaLN projection
339
+ self.time_proj = MiniMaxH3PrunedTimeProj()
340
+ self.time_embedder = MiniMaxH3PrunedTimeEmbedder(table_size=time_table_size, adaln_rank=adaln_rank)
341
+
342
+ # 2b. The affine map the fold was performed with: `silu(time_embedder(t)) ~= adaln_mean + c(t) @ adaln_basis`.
343
+ # Nothing in `forward` reads these - the folded projections already carry them. They are stored so that a
344
+ # LoRA trained on the released 2688-wide projections can be mapped onto these coordinates at load time;
345
+ # see `load_lora_adapter`. 97 KB per partition.
346
+ self.register_buffer("adaln_basis", torch.zeros(adaln_rank, time_embed_dim), persistent=True)
347
+ self.register_buffer("adaln_mean", torch.zeros(time_embed_dim), persistent=True)
348
+
349
+ # 3. Rotary embedding over the packed (t, h, w) grid
350
+ self.rope = MiniMaxH3RotaryPosEmbed(rope_freq_dim=rope_freq_dim, rope_theta=rope_theta)
351
+
352
+ # 4. Text stream refiner
353
+ self.token_refiner = MiniMaxH3TokenRefiner(
354
+ hidden_size=hidden_size,
355
+ num_attention_heads=num_attention_heads,
356
+ attention_head_dim=attention_head_dim,
357
+ ffn_dim=ffn_dim,
358
+ num_layers=num_refiner_layers,
359
+ norm_eps=norm_eps,
360
+ qk_norm_eps=qk_norm_eps,
361
+ final_norm_eps=final_norm_eps,
362
+ )
363
+
364
+ # 5. The block stack, with each block's AdaLN projection narrowed to the timestep coordinates. The block is
365
+ # built with `time_embed_dim=adaln_rank` so its own projection is already the right shape, then swapped
366
+ # for the pruned module, which drops the `silu` and moves the bias to float32.
367
+ self.transformer_blocks = nn.ModuleList(
368
+ [
369
+ MiniMaxH3TransformerBlock(
370
+ hidden_size=hidden_size,
371
+ num_attention_heads=num_attention_heads,
372
+ attention_head_dim=attention_head_dim,
373
+ ffn_dim=ffn_dim,
374
+ time_embed_dim=adaln_rank,
375
+ norm_eps=norm_eps,
376
+ qk_norm_eps=qk_norm_eps,
377
+ )
378
+ for _ in range(num_layers)
379
+ ]
380
+ )
381
+ for block in self.transformer_blocks:
382
+ block.adaln_proj = MiniMaxH3PrunedAdaLayerNormModulation(adaln_rank=adaln_rank, hidden_size=hidden_size)
383
+
384
+ # 6. Shared output norm and the two per-modality output heads
385
+ self.norm_out = MiniMaxH3PrunedAdaLayerNormOut(
386
+ hidden_size=hidden_size, adaln_rank=adaln_rank, eps=final_norm_eps
387
+ )
388
+ self.proj_out = nn.Linear(hidden_size, video_patch_dim, bias=True)
389
+ self.audio_proj_out = nn.Linear(hidden_size, audio_in_channels, bias=True)
390
+
391
+ self.gradient_checkpointing = False
392
+
393
+ # -- LoRA ---------------------------------------------------------------------------------------------------
394
+
395
+ def project_adaln_lora(self, state_dict: dict, prefix: str | None = None) -> tuple[dict, dict]:
396
+ r"""Map a LoRA's AdaLN factors from the released timestep embedding onto the pruned coordinates.
397
+
398
+ The released projection reads `x = silu(time_embedder(t))`, so a LoRA on it contributes
399
+
400
+ lora_B @ (lora_A @ x) = lora_B @ (lora_A @ (mean + basis.T @ c))
401
+ = (lora_B @ (lora_A @ basis.T)) @ c + lora_B @ (lora_A @ mean)
402
+
403
+ which is a rank-preserving `[rank, adaln_rank]` `lora_A` over the pruned coordinates plus a constant output
404
+ offset. Both are computed in float64 from the file's own factors; the only error is the rank-8 subspace's
405
+ own residual on the timestep curve, 1.5e-5 relative, which is ~250x below one bfloat16 step of the weights
406
+ being adapted.
407
+
408
+ Returns `(state_dict, {module path: offset})`, the state dict unchanged and the offsets empty when the AdaLN
409
+ factors are already `adaln_rank`-wide (a LoRA trained against a pruned checkpoint - most of them).
410
+
411
+ Every AdaLN module in a file has to be one or the other. A file that mixes widths is not something this can
412
+ half-apply, so it raises.
413
+
414
+ `prefix` scopes this to one component's keys, the same way `load_lora_adapter` scopes the load - a file that
415
+ names both partitions holds two different adapters under one roof, and only one of them is going into this
416
+ module.
417
+ """
418
+ adaln_rank = self.config.adaln_rank
419
+ time_embed_dim = self.config.time_embed_dim
420
+
421
+ modules = {}
422
+ for key in state_dict:
423
+ if prefix is not None and not key.startswith(f"{prefix}."):
424
+ continue
425
+ match = ADALN_LORA_A_KEY.search(key)
426
+ if match is not None:
427
+ modules[key] = match.group(1)
428
+ if not modules:
429
+ return state_dict, {}
430
+
431
+ widths = sorted({int(state_dict[key].shape[1]) for key in modules})
432
+ if widths == [adaln_rank]:
433
+ return state_dict, {}
434
+ if widths != [time_embed_dim]:
435
+ odd = sorted(
436
+ {modules[key] for key in modules if int(state_dict[key].shape[1]) != max(widths)},
437
+ key=lambda name: (name != "norm_out.linear", name),
438
+ )
439
+ raise ValueError(
440
+ f"This LoRA's {len(modules)} AdaLN projections read inputs of width {widths}. On this checkpoint they "
441
+ f"have to be uniformly {adaln_rank} wide (trained against a pruned checkpoint, loaded as they are) or "
442
+ f"uniformly {time_embed_dim} wide (trained against the released checkpoint, projected onto the pruned "
443
+ "coordinates at load). An adapter cannot be applied to some of its AdaLN modules and not others, so "
444
+ f"nothing was loaded. The minority width is on: {odd[:4]}{' ...' if len(odd) > 4 else ''}."
445
+ )
446
+
447
+ basis = self.adaln_basis
448
+ mean = self.adaln_mean
449
+ if basis.abs().sum() == 0:
450
+ raise ValueError(
451
+ "Projecting a released-checkpoint LoRA needs `adaln_basis` and `adaln_mean`, the affine map this "
452
+ "checkpoint's AdaLN projections were folded with, and this model was loaded without them. Re-download "
453
+ "the repository: they ship as `adaln_affine.safetensors` next to the weights."
454
+ )
455
+ basis = basis.double()
456
+ mean = mean.double()
457
+
458
+ projected = dict(state_dict)
459
+ offsets = {}
460
+ for key, module in modules.items():
461
+ lora_b_key = key[: -len("lora_A.weight")] + "lora_B.weight"
462
+ if lora_b_key not in state_dict:
463
+ raise ValueError(
464
+ f"{key} has no matching {lora_b_key}. The constant term of the projection is "
465
+ "`lora_B @ (lora_A @ mean)`, so both factors have to be present; nothing was loaded."
466
+ )
467
+ lora_a = state_dict[key].to(device=basis.device, dtype=torch.float64)
468
+ lora_b = state_dict[lora_b_key].to(device=basis.device, dtype=torch.float64)
469
+ # float32 rather than the file's bfloat16: the projection is exact arithmetic on the file's factors and
470
+ # there is no reason to round it twice. PEFT casts to the adapter's dtype when it loads them.
471
+ projected[key] = (lora_a @ basis.T).to(torch.float32).cpu().contiguous()
472
+ offsets[module] = (lora_b @ (lora_a @ mean)).to(torch.float32).cpu().contiguous()
473
+ return projected, offsets
474
+
475
+ def load_lora_adapter(self, pretrained_model_name_or_path_or_dict, prefix="transformer", hotswap=False, **kwargs):
476
+ r"""`PeftAdapterMixin.load_lora_adapter`, with the AdaLN projection of [`project_adaln_lora`] in front of it.
477
+
478
+ LoRAs trained against a pruned checkpoint pass through untouched. LoRAs trained against the released
479
+ checkpoint's 2688-wide AdaLN projections - the official turbo LoRA and its conversions - are mapped onto the
480
+ pruned coordinates here, which is the only thing that ever stopped them loading. Everything outside the AdaLN
481
+ path is identical between the two checkpoints and is neither inspected nor changed.
482
+ """
483
+ state_dict = pretrained_model_name_or_path_or_dict
484
+ if not isinstance(state_dict, dict):
485
+ from diffusers.loaders.lora_base import _fetch_state_dict
486
+
487
+ state_dict, _ = _fetch_state_dict(
488
+ pretrained_model_name_or_path_or_dict=state_dict,
489
+ weight_name=kwargs.get("weight_name"),
490
+ use_safetensors=kwargs.get("use_safetensors", True),
491
+ local_files_only=kwargs.get("local_files_only"),
492
+ cache_dir=kwargs.get("cache_dir"),
493
+ force_download=kwargs.get("force_download", False),
494
+ proxies=kwargs.get("proxies"),
495
+ token=kwargs.get("token"),
496
+ revision=kwargs.get("revision"),
497
+ subfolder=kwargs.get("subfolder"),
498
+ user_agent={"file_type": "attn_procs_weights", "framework": "pytorch"},
499
+ allow_pickle=False,
500
+ metadata=kwargs.get("metadata"),
501
+ )
502
+
503
+ state_dict, offsets = self.project_adaln_lora(state_dict, prefix=prefix)
504
+ if offsets:
505
+ logger.info(
506
+ f"Projecting {len(offsets)} AdaLN LoRA modules from the released {self.config.time_embed_dim}-wide "
507
+ f"timestep embedding onto this checkpoint's {self.config.adaln_rank} pruned coordinates; each one's "
508
+ "constant term is carried as a float32 offset on the modulation."
509
+ )
510
+
511
+ before = set(getattr(self, "peft_config", None) or ())
512
+ super().load_lora_adapter(state_dict, prefix=prefix, hotswap=hotswap, **kwargs)
513
+ if not offsets:
514
+ return
515
+
516
+ added = set(getattr(self, "peft_config", None) or ()) - before
517
+ adapter_name = added.pop() if len(added) == 1 else kwargs.get("adapter_name")
518
+ if adapter_name is None:
519
+ raise RuntimeError(
520
+ "The AdaLN projection could not tell which adapter was just loaded, so its constant terms were not "
521
+ f"attached and the adapter is incomplete. Adapters before: {sorted(before)}, after: "
522
+ f"{sorted(getattr(self, 'peft_config', None) or ())}. Pass `adapter_name` explicitly."
523
+ )
524
+ for module, offset in offsets.items():
525
+ self.get_submodule(module.rsplit(".", 1)[0]).register_lora_adaln_offset(adapter_name, offset)
526
+
527
+ # -- 8-bit compute -----------------------------------------------------------------------------------------
528
+ #
529
+ # `enable_convrot` and `quantize_8bit` are opt-in and change nothing until called.
530
+
531
+ def convrot_layers(self) -> list[str]:
532
+ r"""The attention and feed-forward linears ConvRot applies to (300 of them: 50 blocks x 6).
533
+
534
+ A PEFT-wrapped target appears as `....ff.net.2.base_layer` and is matched as such, so a model that
535
+ already has a LoRA attached rotates its *base* layers and leaves the adapters alone.
536
+ """
537
+ names = []
538
+ for name, module in self.named_modules():
539
+ if not isinstance(module, nn.Linear) or "transformer_blocks" not in name:
540
+ continue
541
+ stem = name[: -len(".base_layer")] if name.endswith(".base_layer") else name
542
+ if any(stem.endswith(suffix) for suffix in CONVROT_SUFFIXES):
543
+ names.append(name)
544
+ return names
545
+
546
+ def enable_convrot(self, group_size: int = 256) -> list[str]:
547
+ r"""Fold `W <- W @ H` into every ConvRot target and switch it to [`MiniMaxH3ConvRotLinear`].
548
+
549
+ The fold runs in float64 and rounds back to the weight's dtype once. Idempotence is not claimed: `H` is
550
+ an involution, so calling this twice restores the original weights while leaving the online rotation in
551
+ place, which is wrong. Call it once, on a freshly loaded model, before quantizing.
552
+ """
553
+ names = self.convrot_layers()
554
+ lookup = dict(self.named_modules())
555
+ for name in names:
556
+ module = lookup[name]
557
+ weight = module.weight
558
+ # In place: reassigning `weight.data` 300 times leaves the old storages to the caching allocator
559
+ # and fragments it badly on a card holding the whole model.
560
+ weight.data.copy_(_rotate(weight.data.double(), group_size))
561
+ module.__class__ = MiniMaxH3ConvRotLinear
562
+ module.convrot_groupsize = group_size
563
+ self._convrot_layers = names
564
+ return names
565
+
566
+ def convrot_filter(self, module: nn.Module, fqn: str) -> bool:
567
+ r"""`filter_fn` for `torchao.quantize_`: the ConvRot targets and nothing else."""
568
+ return fqn in getattr(self, "_convrot_layers", ()) or fqn in self.convrot_layers()
569
+
570
+ def quantize_8bit(self, config=None, group_size: int = 256, device=None):
571
+ r"""Rotate and quantize one linear at a time, so the working set is one weight rather than the model.
572
+
573
+ `config` is any torchao config; the default is int8 dynamic activations with int8 weights, which is the
574
+ configuration measured closest to bfloat16 on this model. `device` is where the fold and the
575
+ quantization run - point it at the GPU when the model itself is on the CPU and the transient cost is one
576
+ `[28672, 5376]` weight rather than 40 GB.
577
+
578
+ Returns `self`.
579
+ """
580
+ from torchao.quantization import Int8DynamicActivationInt8WeightConfig, quantize_
581
+
582
+ if config is None:
583
+ config = Int8DynamicActivationInt8WeightConfig()
584
+ names = self.convrot_layers()
585
+ lookup = dict(self.named_modules())
586
+ for name in names:
587
+ module = lookup[name]
588
+ home = module.weight.device
589
+ module.to(device or home)
590
+ weight = module.weight
591
+ # In place: reassigning `weight.data` 300 times leaves the old storages to the caching allocator
592
+ # and fragments it badly on a card holding the whole model.
593
+ weight.data.copy_(_rotate(weight.data.double(), group_size))
594
+ module.__class__ = MiniMaxH3ConvRotLinear
595
+ module.convrot_groupsize = group_size
596
+ quantize_(module, config)
597
+ module.to(home)
598
+ self._convrot_layers = names
599
+ return self
600
+
601
+
602
+ MiniMaxH3PrunedTransformer3DModel.register_for_auto_class("AutoModel")