speach1sdef178 commited on
Commit
752626b
·
verified ·
1 Parent(s): 4383944

Initial release of MiniMax H3 Semantic Bridge v1.0

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +4 -0
  2. LICENSE.md +27 -0
  3. MiniMaxH3_SemanticBridge_v1.safetensors +3 -0
  4. MiniMax_H3_Semantic_Bridge_v1.0.zip +3 -0
  5. NOTICE.txt +5 -0
  6. README.md +435 -0
  7. RESEARCH_ARTICLE.md +1017 -0
  8. SHA256SUMS.txt +65 -0
  9. UPSTREAM_LICENSES.md +23 -0
  10. examples/01_rooftop_train_chase/native_h3.mp4 +3 -0
  11. examples/01_rooftop_train_chase/prompt.txt +10 -0
  12. examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4 +3 -0
  13. examples/02_glass_table_prompt_adherence/native_h3.mp4 +3 -0
  14. examples/02_glass_table_prompt_adherence/prompt.txt +13 -0
  15. examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4 +3 -0
  16. research/datasets/bridge_ood_prompts_160.json +642 -0
  17. research/datasets/bridge_prompts_480.json +1762 -0
  18. research/raw_scripts/compare_hidden_layers_bridge.py +439 -0
  19. research/raw_scripts/compare_minimax_sensenova.py +335 -0
  20. research/raw_scripts/compare_tokenizers.py +145 -0
  21. research/raw_scripts/evaluate_distilled_student_v2_ood.py +1343 -0
  22. research/raw_scripts/evaluate_hidden_bridge_ood.py +700 -0
  23. research/raw_scripts/extract_h3_bridge_dataset.py +272 -0
  24. research/raw_scripts/extract_h3_embeddings_full.py +188 -0
  25. research/raw_scripts/extract_h3_embeddings_test.py +139 -0
  26. research/raw_scripts/extract_h3_hidden_states.py +189 -0
  27. research/raw_scripts/extract_h3_hidden_states_fast.py +266 -0
  28. research/raw_scripts/extract_h3_ood_hidden.py +318 -0
  29. research/raw_scripts/extract_sensenova_bridge_dataset.py +425 -0
  30. research/raw_scripts/extract_sensenova_hidden_states.py +321 -0
  31. research/raw_scripts/extract_sensenova_hidden_states_local.py +421 -0
  32. research/raw_scripts/extract_sensenova_hidden_states_stream.py +352 -0
  33. research/raw_scripts/extract_sensenova_hidden_states_stream_v2.py +586 -0
  34. research/raw_scripts/extract_sensenova_hidden_states_stream_v3.py +546 -0
  35. research/raw_scripts/extract_sensenova_ood_hidden.py +633 -0
  36. research/raw_scripts/inspect_distillation_data.py +466 -0
  37. research/raw_scripts/inspect_h3_sensenova_bridge.py +274 -0
  38. research/raw_scripts/make_bridge_ood_prompts.py +391 -0
  39. research/raw_scripts/make_bridge_ood_prompts_v1.py +272 -0
  40. research/raw_scripts/make_bridge_prompts.py +255 -0
  41. research/raw_scripts/prepare_sensenova_h3_condition.py +383 -0
  42. research/raw_scripts/test_sensenova_mot_gen_bridge.py +698 -0
  43. research/raw_scripts/train_distilled_adapter_screening.py +1311 -0
  44. research/raw_scripts/train_distilled_student_v2.py +1802 -0
  45. research/raw_scripts/train_distilled_student_v3.py +1628 -0
  46. research/raw_scripts/train_embedding_bridge_test.py +397 -0
  47. research/raw_scripts/train_hidden_bridge_screening.py +859 -0
  48. research/reports/distillation_source_inspection.txt +435 -0
  49. research/reports/distilled_adapter_screening.csv +8 -0
  50. research/reports/distilled_adapter_screening.txt +73 -0
.gitattributes CHANGED
@@ -33,3 +33,7 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ examples/01_rooftop_train_chase/native_h3.mp4 filter=lfs diff=lfs merge=lfs -text
37
+ examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4 filter=lfs diff=lfs merge=lfs -text
38
+ examples/02_glass_table_prompt_adherence/native_h3.mp4 filter=lfs diff=lfs merge=lfs -text
39
+ examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4 filter=lfs diff=lfs merge=lfs -text
LICENSE.md ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Licensing
2
+
3
+ This repository contains several kinds of artifacts with different upstream relationships.
4
+
5
+ ## Model-derived artifacts
6
+
7
+ The Semantic Bridge adapter and teacher-side bridge research artifact were developed using MiniMax H3 representations and should be treated in accordance with the upstream **MiniMax H3 Community License Agreement**:
8
+
9
+ https://huggingface.co/MiniMaxAI/MiniMax-H3/blob/main/LICENSE
10
+
11
+ The repository metadata therefore uses `license: other` with the MiniMax H3 Community License Agreement as the linked model license rather than describing the model-derived artifacts as MIT or Apache-2.0.
12
+
13
+ The required MiniMax distribution notice is included in `NOTICE.txt`.
14
+
15
+ ## SenseNova research teacher
16
+
17
+ SenseNova U1.5 was used during the teacher-assisted research stage. The upstream SenseNova U1.5 model is released under Apache License 2.0:
18
+
19
+ https://huggingface.co/sensenova/SenseNova-U1.5-8B-MoT
20
+
21
+ The SenseNova checkpoint itself is not redistributed here.
22
+
23
+ ## Research code and documentation
24
+
25
+ The repository contains original experimental scripts, reports, documentation, and a ComfyUI integration. No separate permissive license is asserted here over model-derived artifacts. Before reusing or redistributing the package, review the upstream licenses and any terms applicable to your use and territory.
26
+
27
+ This file is a project summary, not legal advice.
MiniMaxH3_SemanticBridge_v1.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ac0dc8ac05f545ebdee12e2fcebe4515b049f9cfd9558eb4887a9bf3fd6d562e
3
+ size 11023032
MiniMax_H3_Semantic_Bridge_v1.0.zip ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c4866a0cc554237046814c9cb827a5960caecbd2214d5a5a04b0e43716b437bb
3
+ size 4406
NOTICE.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ MiniMax H3 is licensed under the MiniMax H3 Community License Agreement, Copyright © 2026 MiniMax. All Rights Reserved.
2
+
3
+ MiniMax H3 Semantic Bridge is an independent experimental research project and is not an official MiniMax release.
4
+
5
+ SenseNova U1.5 was used as an experimental teacher during development. The public Semantic Bridge does not require or redistribute the SenseNova checkpoint.
README.md CHANGED
@@ -2,4 +2,439 @@
2
  license: other
3
  license_name: minimax-h3-community-license-agreement
4
  license_link: https://huggingface.co/MiniMaxAI/MiniMax-H3/blob/main/LICENSE
 
 
 
 
 
 
 
 
 
 
5
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
  license: other
3
  license_name: minimax-h3-community-license-agreement
4
  license_link: https://huggingface.co/MiniMaxAI/MiniMax-H3/blob/main/LICENSE
5
+ base_model:
6
+ - MiniMaxAI/MiniMax-H3
7
+ tags:
8
+ - minimax-h3
9
+ - comfyui
10
+ - video-generation
11
+ - semantic-adapter
12
+ - representation-learning
13
+ - distillation
14
+ - research
15
  ---
16
+
17
+ # MiniMax H3 Semantic Bridge
18
+
19
+ ## Cross-Architecture Semantic Transfer and Distillation for Video Generation
20
+
21
+ **MiniMax H3 Semantic Bridge** is a compact conditioning-space adapter for the standard **MiniMax H3 FL2VA / text-conditioned generation path**.
22
+
23
+ It grew out of an experimental cross-architecture representation-transfer project using **SenseNova U1.5** as a semantic teacher. The final released adapter is standalone: **SenseNova is not required at inference time**.
24
+
25
+ **In one line:**
26
+
27
+ > cross-architecture semantic transfer → conditioning-space teacher bridge → distillation → a ~11 MB standalone H3 adapter
28
+
29
+ This is **not** a LoRA, checkpoint merge, or conventional parameter graft. The adapter transforms native H3 conditioning before the video transformer and blends the learned semantic representation back into H3 at a controllable strength.
30
+
31
+ > **Scope:** v1 is for standard H3 FL2VA / text-conditioned generation. **Ref2VA / reference-conditioned workflows are not supported.** Experimental reference-audio testing showed degraded singing/lip-sync when the adapter was inserted into Ref2VA conditioning.
32
+
33
+ ---
34
+
35
+ ## Quick Start
36
+
37
+ ### Files
38
+
39
+ - `MiniMaxH3_SemanticBridge_v1.safetensors` — final standalone adapter.
40
+ - `MiniMax_H3_Semantic_Bridge_v1.0.zip` — ComfyUI custom node.
41
+ - `RESEARCH_ARTICLE.md` — full research narrative.
42
+ - `research/` — prompt datasets, raw scripts, reports, and teacher-side research artifact.
43
+ - `examples/` — controlled Native H3 vs Semantic Bridge A/B videos and their exact prompts.
44
+
45
+ ### Installation
46
+
47
+ 1. Extract `MiniMax_H3_Semantic_Bridge_v1.0.zip` into:
48
+
49
+ ```text
50
+ ComfyUI/custom_nodes/
51
+ ```
52
+
53
+ 2. Create:
54
+
55
+ ```text
56
+ ComfyUI/models/semantic_bridge/
57
+ ```
58
+
59
+ 3. Put:
60
+
61
+ ```text
62
+ MiniMaxH3_SemanticBridge_v1.safetensors
63
+ ```
64
+
65
+ inside that folder.
66
+
67
+ 4. Restart ComfyUI.
68
+
69
+ ### Nodes
70
+
71
+ - **MiniMax H3 Image to Video + Semantic Bridge**
72
+ - **MiniMax H3 Semantic Bridge**
73
+ - **MiniMax H3 Clear Semantic Bridge Cache**
74
+
75
+ ### Recommended settings
76
+
77
+ ```text
78
+ alpha = 0.10
79
+ magnitude_match = per_token
80
+ ```
81
+
82
+ For the qualitative A/B examples below, `alpha = 0.15` was intentionally used to make the behavioral difference easier to observe.
83
+
84
+ ---
85
+
86
+ # What problem was this trying to solve?
87
+
88
+ Generative models can recognize all the concepts in a prompt while still failing to preserve the relationships between those concepts.
89
+
90
+ For example, a prompt may specify not only a person, table, bottle, mirror, and light source, but also:
91
+
92
+ - which hand holds which object;
93
+ - which hand must remain still;
94
+ - left/right ordering of several objects;
95
+ - which surfaces are transparent or reflective;
96
+ - how a mirror should correspond to the real scene;
97
+ - how light passes through one material but reflects from another;
98
+ - whether an action is explicitly requested or explicitly *not* requested.
99
+
100
+ The project therefore focused on **semantic structure and prompt adherence**, rather than adding new visual concepts to H3.
101
+
102
+ Areas explored during the research included:
103
+
104
+ - complex composition;
105
+ - spatial relationships;
106
+ - anatomy and body relationships;
107
+ - object counting;
108
+ - text rendering / textual constraints;
109
+ - materials and lighting;
110
+ - reflections, transparency, and occlusion;
111
+ - long prompts with several simultaneous constraints.
112
+
113
+ ---
114
+
115
+ # Research path
116
+
117
+ ## 1. Direct grafting failed
118
+
119
+ The project started as a direct grafting experiment between SenseNova U1.5 and MiniMax H3.
120
+
121
+ The architectures did not expose useful parameter-level correspondences. Exact shape matching, transpose matching, and simple input/output dimensional matching did not provide a meaningful path for direct tensor transplantation.
122
+
123
+ That negative result changed the question from:
124
+
125
+ ```text
126
+ Which weights can be copied?
127
+ ```
128
+
129
+ to:
130
+
131
+ ```text
132
+ Can the models' internal representations of the same prompt be aligned?
133
+ ```
134
+
135
+ ## 2. Hidden-representation alignment
136
+
137
+ Hidden states from both systems were extracted across a deliberately varied semantic prompt set. Lightweight projections were trained between candidate representation spaces.
138
+
139
+ A substantially stronger correspondence emerged than the alternatives.
140
+
141
+ On held-out prompts from the original distribution, the strongest experimental mapping reached approximately:
142
+
143
+ ```text
144
+ validation cosine ≈ 0.904
145
+ ```
146
+
147
+ ## 3. Strict OOD test
148
+
149
+ A separate set of **160 prompts** was constructed to stress harder combinations of anatomy, counting, materials/light, spatial structure, text, architecture/vehicles, reflection/occlusion, and long compositions.
150
+
151
+ With the bridge frozen, strict OOD similarity was approximately:
152
+
153
+ ```text
154
+ 0.749
155
+ ```
156
+
157
+ The drop was real, but the mapping did not collapse. This motivated testing the representation inside the actual H3 generation path.
158
+
159
+ ## 4. Full teacher bridge
160
+
161
+ An experimental Full Bridge used SenseNova at inference time, projected the teacher-side representation into H3-compatible conditioning, magnitude-aligned it, and blended it with native H3 conditioning.
162
+
163
+ Conceptually:
164
+
165
+ ```text
166
+ H = native H3 conditioning
167
+ S = mapped teacher semantic representation
168
+ C = H + alpha * (S - H)
169
+ ```
170
+
171
+ The Full Bridge produced coherent H3 generations and visible behavioral changes, demonstrating that the cross-model representation mapping survived the downstream video-generation process.
172
+
173
+ However, it required the full teacher model at runtime, which was impractical.
174
+
175
+ ## 5. Distillation
176
+
177
+ The Full Bridge was then treated as a teacher. A compact H3-side student was trained to predict the teacher-derived representation directly from H3's own conditioning.
178
+
179
+ Early attempts to predict the correction delta directly were weak (best correction similarity around `0.51`). Predicting the teacher-derived representation itself worked dramatically better.
180
+
181
+ The final result was a small standalone adapter with no SenseNova runtime dependency.
182
+
183
+ ---
184
+
185
+ # Final distillation results
186
+
187
+ The preserved V3 report records a **500-prompt training split** and **100-prompt validation split** containing both original-distribution and harder/OOD examples.
188
+
189
+ | Metric | Result |
190
+ |---|---:|
191
+ | Teacher representation cosine | **0.995890** |
192
+ | Semantic correction cosine | **0.983558** |
193
+ | Main-distribution correction | **0.980888** |
194
+ | OOD correction | **0.989788** |
195
+ | Minimum correction | **0.935910** |
196
+ | Blend cosine, alpha 0.10 | **0.999958** |
197
+ | Blend cosine, alpha 0.20 | **0.999827** |
198
+ | Blend cosine, alpha 0.30 | **0.999602** |
199
+
200
+ These are **representation-space / distillation metrics**. They do **not** mean that video quality improves by the same percentages, and they are not a substitute for controlled visual evaluation.
201
+
202
+ ### Final correction similarity by semantic category
203
+
204
+ | Category | Similarity |
205
+ |---|---:|
206
+ | Text | 0.996061 |
207
+ | Reflection / occlusion | 0.995932 |
208
+ | Complex counting | 0.994557 |
209
+ | Complex text | 0.991912 |
210
+ | Reflection / occlusion OOD | 0.991855 |
211
+ | Complex spatial | 0.991125 |
212
+ | Long composition | 0.990637 |
213
+ | Architecture / vehicle | 0.990606 |
214
+ | Complex material / light | 0.987394 |
215
+ | Material | 0.986263 |
216
+ | Complex anatomy | 0.982104 |
217
+ | Anatomy | 0.980773 |
218
+ | Lighting | 0.978327 |
219
+ | Spatial | 0.974002 |
220
+ | Counting | 0.965250 |
221
+
222
+ ---
223
+
224
+ # Qualitative A/B examples
225
+
226
+ The following examples use the **same prompt and generation setup within each pair**. The intended comparison is Native H3 versus H3 with Semantic Bridge enabled. The published Bridge examples use `alpha = 0.15` to make the effect easier to inspect visually.
227
+
228
+ These examples are **qualitative observations**, not a benchmark or proof of universal improvement.
229
+
230
+ ## Example 01 — Rooftop Train Chase
231
+
232
+ **Focus:** complex motion, anatomy, action sequencing, physical interaction, material response, and spatial continuity.
233
+
234
+ | Native MiniMax H3 | + Semantic Bridge (`alpha=0.15`) |
235
+ |---|---|
236
+ | [View native video](./examples/01_rooftop_train_chase/native_h3.mp4) | [View bridge video](./examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4) |
237
+
238
+ <video controls width="49%" src="https://huggingface.co/speach1sdef178/MiniMax-H3-Semantic-Bridge/resolve/main/examples/01_rooftop_train_chase/native_h3.mp4"></video>
239
+ <video controls width="49%" src="https://huggingface.co/speach1sdef178/MiniMax-H3-Semantic-Bridge/resolve/main/examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4"></video>
240
+
241
+ The exact prompt is preserved in:
242
+
243
+ [`examples/01_rooftop_train_chase/prompt.txt`](./examples/01_rooftop_train_chase/prompt.txt)
244
+
245
+ This example was designed to stress several constraints simultaneously: two moving characters, pursuit distance, running anatomy, a specific vault interaction, hand contact with the obstacle, landing continuity, moving camera geometry, wet reflective metal, rain, sparks, and a rapidly moving city background.
246
+
247
+ ---
248
+
249
+ ## Example 02 — Prompt Adherence, Materials, Reflection & Transparency
250
+
251
+ **Focus:** explicit state adherence, hand behavior, material differences, object ordering, reflection, transparency, and text.
252
+
253
+ | Native MiniMax H3 | + Semantic Bridge (`alpha=0.15`) |
254
+ |---|---|
255
+ | [View native video](./examples/02_glass_table_prompt_adherence/native_h3.mp4) | [View bridge video](./examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4) |
256
+
257
+ <video controls width="49%" src="https://huggingface.co/speach1sdef178/MiniMax-H3-Semantic-Bridge/resolve/main/examples/02_glass_table_prompt_adherence/native_h3.mp4"></video>
258
+ <video controls width="49%" src="https://huggingface.co/speach1sdef178/MiniMax-H3-Semantic-Bridge/resolve/main/examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4"></video>
259
+
260
+ ### Controlled prompt-following observation
261
+
262
+ A particularly useful instruction in this prompt is:
263
+
264
+ > **"Her right hand rests flat on the glass tabletop with all five fingers naturally separated and clearly visible."**
265
+
266
+ In this A/B generation:
267
+
268
+ - **Native H3** introduces an unrequested action: the right hand moves across the tabletop rather than remaining in the requested resting state.
269
+ - **Semantic Bridge (`alpha=0.15`)** keeps the hand resting on the glass surface, more closely preserving the explicitly requested state.
270
+
271
+ This observation is important because it is not a subjective claim that one result is simply "prettier." The prompt specifies a directly observable state — **resting** — and the two outputs behave differently with respect to that instruction.
272
+
273
+ It is still presented only as **qualitative evidence from this controlled pair**, not as a statistical claim that the adapter universally improves prompt adherence.
274
+
275
+ The exact full prompt is preserved in:
276
+
277
+ [`examples/02_glass_table_prompt_adherence/prompt.txt`](./examples/02_glass_table_prompt_adherence/prompt.txt)
278
+
279
+ The same prompt also stresses:
280
+
281
+ - anatomically coherent hands;
282
+ - a cup held specifically in the left hand;
283
+ - exactly three tabletop objects in a specified left-to-right order;
284
+ - the text `NIGHT SHIFT`;
285
+ - mirror correspondence;
286
+ - transparent glass;
287
+ - refractive bottle behavior;
288
+ - reflective metal;
289
+ - matte ceramic;
290
+ - warm/cool directional lighting interactions.
291
+
292
+ ---
293
+
294
+ # What Semantic Bridge is — and is not
295
+
296
+ Semantic Bridge is best described as:
297
+
298
+ > **a compact conditioning-space adapter distilled from a cross-architecture semantic mapping**
299
+
300
+ It is **not**:
301
+
302
+ - a conventional LoRA;
303
+ - a checkpoint merge;
304
+ - a direct parameter graft;
305
+ - a copy of SenseNova weights;
306
+ - a second multimodal model running beside H3.
307
+
308
+ At inference time the released adapter operates only on H3 conditioning.
309
+
310
+ Conceptually:
311
+
312
+ ```text
313
+ Prompt
314
+ ↓
315
+ H3 text conditioning
316
+ ↓
317
+ Semantic Bridge
318
+ ↓
319
+ learned semantic representation
320
+ ↓
321
+ magnitude matching
322
+ ↓
323
+ controlled residual blend
324
+ ↓
325
+ MiniMax H3 video generation
326
+ ```
327
+
328
+ SenseNova was used as a teacher during the research process only.
329
+
330
+ ---
331
+
332
+ # Scope: FL2VA / standard text-conditioned H3 only
333
+
334
+ This limitation is important.
335
+
336
+ The released student was distilled from the **standard H3 conditioning path**. It was not trained on the separate multimodal reference-conditioning distribution used by Ref2VA.
337
+
338
+ An experimental Ref2VA-compatible node was tested by modifying text-designated token positions while preserving visual-reference tokens. In reference-audio singing tests, this produced noticeably worse vocal articulation and stronger mumbling-like lip motion than native Ref2VA.
339
+
340
+ The practical conclusion for v1 is therefore:
341
+
342
+ > **Do not use this adapter for Ref2VA / reference-conditioned generation, especially reference-audio singing or lip-sync.**
343
+
344
+ This does not establish that a semantic bridge can never work with Ref2VA. It suggests that a Ref2VA version should be trained separately on the multimodal conditioning regime it is intended to modify.
345
+
346
+ A useful lesson from the failed experiment is that **matching tensor dimensionality does not guarantee matching conditioning semantics**.
347
+
348
+ ---
349
+
350
+ # Research materials included
351
+
352
+ The repository intentionally includes more than the final adapter so others can inspect the experimental path.
353
+
354
+ ## `research/datasets/`
355
+
356
+ - `bridge_prompts_480.json` — historical filename; the preserved dataset contains **440** development prompts.
357
+ - `bridge_ood_prompts_160.json` — 160 strict OOD prompts.
358
+
359
+ ## `research/raw_scripts/`
360
+
361
+ Original research scripts are preserved largely as-run. They include local Windows paths and historical filenames. This is intentional: they are provided as a research snapshot rather than as a polished one-command training framework.
362
+
363
+ The scripts cover areas such as:
364
+
365
+ - architecture and tokenizer comparison;
366
+ - H3 and SenseNova hidden-state extraction;
367
+ - layer-pair screening;
368
+ - MoT diagnostic experiments;
369
+ - OOD dataset construction;
370
+ - full bridge evaluation;
371
+ - distillation screening;
372
+ - V2/V3 student training and evaluation.
373
+
374
+ ## `research/reports/`
375
+
376
+ Raw TXT/CSV outputs from the experiments are included, including negative and intermediate results.
377
+
378
+ This is deliberate. The unsuccessful directions are part of the research record and may help others avoid repeating the same experiments.
379
+
380
+ ## `research/teacher_artifacts/`
381
+
382
+ `SN_L32_to_H3_L49_rank128.safetensors` is preserved as an optional research artifact from the teacher-side bridge work.
383
+
384
+ **It is not required to use the public Semantic Bridge.**
385
+
386
+ Large extracted hidden-state `.pt` caches are not included. The prompt datasets and extraction scripts are provided so those intermediates can be regenerated.
387
+
388
+ ---
389
+
390
+ # Limitations
391
+
392
+ - This is an experimental research adapter, not a universal H3 enhancer.
393
+ - It may help some prompts, do little on others, or occasionally make a result worse.
394
+ - The reported cosine metrics evaluate agreement with the teacher-derived representation, not perceptual video quality.
395
+ - A small conditioning difference can produce a large downstream sampling difference.
396
+ - The current adapter is not validated for Ref2VA/reference-conditioned generation.
397
+ - The project has not yet been evaluated with a large standardized human-preference benchmark.
398
+ - The published A/B videos are illustrative controlled examples, not statistical proof.
399
+
400
+ ---
401
+
402
+ # Licensing and upstream terms
403
+
404
+ **Please read this section before using or redistributing the model-derived artifacts.**
405
+
406
+ MiniMax H3 is released under the **MiniMax H3 Community License Agreement**. The official agreement defines terms for MiniMax H3 and Model Derivatives, including distribution requirements and territorial restrictions. This repository uses `license: other` metadata and points directly to the upstream H3 agreement rather than relabeling the model-derived adapter as Apache/MIT.
407
+
408
+ Official H3 license:
409
+
410
+ https://huggingface.co/MiniMaxAI/MiniMax-H3/blob/main/LICENSE
411
+
412
+ The H3 agreement requires distributions to include a NOTICE. This repository includes `NOTICE.txt`.
413
+
414
+ SenseNova U1.5, used as the experimental teacher during development, is published under **Apache License 2.0**:
415
+
416
+ https://huggingface.co/sensenova/SenseNova-U1.5-8B-MoT
417
+
418
+ This repository does **not** redistribute the SenseNova checkpoint.
419
+
420
+ See [`LICENSE.md`](./LICENSE.md) and [`UPSTREAM_LICENSES.md`](./UPSTREAM_LICENSES.md) for repository-specific notes and direct upstream references.
421
+
422
+ > This licensing summary is provided for transparency and is not legal advice. Users and redistributors should review the upstream terms themselves.
423
+
424
+ ---
425
+
426
+ # Full article
427
+
428
+ For the complete chronological research write-up, including failed grafting, representation screening, Full Bridge, OOD testing, early distillation, V2/V3 results, and the Ref2VA limitation, see:
429
+
430
+ **[RESEARCH_ARTICLE.md](./RESEARCH_ARTICLE.md)**
431
+
432
+ ---
433
+
434
+ # Acknowledgements
435
+
436
+ This is an independent experimental project built around **MiniMax H3** and teacher-assisted representation studies using **SenseNova U1.5**.
437
+
438
+ It is not an official MiniMax or SenseNova release.
439
+
440
+ The value of the project is not only the final adapter, but the possibility that useful semantic behavior may sometimes be transferred between incompatible architectures through **representation alignment and distillation**, even when direct parameter grafting is not meaningful.
RESEARCH_ARTICLE.md ADDED
@@ -0,0 +1,1017 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # MiniMax H3 Semantic Bridge
2
+
3
+ ## Cross-Architecture Semantic Transfer and Distillation for Video Generation
4
+
5
+ ### Abstract
6
+
7
+ **MiniMax H3 Semantic Bridge** is the result of an experimental investigation into whether higher-level semantic behavior can be transferred between fundamentally different generative model architectures without directly merging their weights.
8
+
9
+ The project began with an attempt to transfer selected capabilities from **SenseNova U1.5** into **MiniMax H3**, particularly in areas such as complex composition, spatial relationships, human anatomy, counting, prompt following, text interpretation, materials, lighting, reflections, transparency, and occlusion.
10
+
11
+ Direct parameter grafting proved unsuitable because the two models do not expose meaningfully compatible weight structures. This led to a different approach: instead of transferring parameters, I investigated whether the models' internal representations of the **same textual prompt** contained a learnable cross-model relationship.
12
+
13
+ They did.
14
+
15
+ An experimental **Full Bridge** was developed in which SenseNova acted as a semantic donor at inference time and its representation was transformed into H3-compatible conditioning. The resulting system produced meaningful changes in actual H3 generations and demonstrated that semantic information could be transferred without modifying H3's generative transformer.
16
+
17
+ However, the Full Bridge required the complete donor model during inference.
18
+
19
+ The final stage therefore treated this system as a teacher and distilled its behavior into a compact adapter operating entirely from H3's own representation.
20
+
21
+ The resulting **MiniMax H3 Semantic Bridge** requires no SenseNova checkpoint, runtime, tokenizer, or additional large model during inference.
22
+
23
+ Only this final distilled version is included in the public release.
24
+
25
+ ---
26
+
27
+ # 1. Motivation
28
+
29
+ Large generative models can possess all of the concepts necessary to satisfy a prompt while still failing to preserve the relationships between those concepts.
30
+
31
+ A simple prompt such as:
32
+
33
+ > *A woman holding a glass bottle.*
34
+
35
+ primarily requires concept recognition.
36
+
37
+ A more structured instruction is different:
38
+
39
+ > *A woman stands behind a transparent glass table holding a polished metal cup in her left hand. A mirror behind her reflects the room, while warm side light passes through a glass bottle on the table.*
40
+
41
+ Now the model must maintain several relationships simultaneously:
42
+
43
+ ```text
44
+ woman
45
+ └─ left hand
46
+ └─ polished metal cup
47
+
48
+ glass table
49
+ └─ glass bottle
50
+ └─ warm light passes through it
51
+
52
+ mirror
53
+ └─ reflects the room
54
+ ```
55
+
56
+ All of the individual concepts may already be well represented by the model.
57
+
58
+ The difficult part is their **composition**.
59
+
60
+ This distinction motivated the project.
61
+
62
+ Rather than attempting to add new visual knowledge to MiniMax H3, I wanted to investigate whether another model's semantic interpretation could alter how H3 represents complex instructions before generation begins.
63
+
64
+ The original areas of interest were:
65
+
66
+ * spatial and compositional reasoning;
67
+ * human anatomy and body relationships;
68
+ * object counting;
69
+ * prompt following;
70
+ * text interpretation;
71
+ * material properties;
72
+ * lighting relationships;
73
+ * reflection and transparency;
74
+ * occlusion;
75
+ * long prompts containing multiple simultaneous constraints.
76
+
77
+ The question was therefore not simply:
78
+
79
+ > Can another model be merged into H3?
80
+
81
+ It was:
82
+
83
+ > **Can useful semantic structure learned by another architecture be translated into H3's representation space?**
84
+
85
+ ---
86
+
87
+ # 2. The First Approach: Direct Model Grafting
88
+
89
+ The project initially began as a model-grafting experiment.
90
+
91
+ The intuitive approach was to identify semantically relevant components in SenseNova and transfer them directly into MiniMax H3.
92
+
93
+ Architectural analysis quickly showed why this was unlikely to work.
94
+
95
+ The two models have substantially different internal structures, dimensionalities and parameter organization. Examination of the major two-dimensional weight matrices did not reveal useful direct correspondences suitable for conventional grafting.
96
+
97
+ There were no meaningful exact shape matches in the parts investigated, nor did simple transposition or matching input/output dimensions provide a useful mapping.
98
+
99
+ This is an important negative result.
100
+
101
+ Two models can potentially encode related semantic knowledge without storing that knowledge in parameter matrices that are directly interchangeable.
102
+
103
+ A parameter-level question was therefore replaced by a representation-level question.
104
+
105
+ Instead of looking for:
106
+
107
+ ```text
108
+ SenseNova weights
109
+ ↓
110
+ matching H3 weights
111
+ ```
112
+
113
+ I began looking for:
114
+
115
+ ```text
116
+ same prompt
117
+ ↓ ↓
118
+ SenseNova H3
119
+ ↓ ↓
120
+ internal representation
121
+ ↓
122
+ learnable relationship?
123
+ ```
124
+
125
+ This turned out to be much more productive.
126
+
127
+ ---
128
+
129
+ # 3. Searching for Cross-Model Representation Alignment
130
+
131
+ Both models were presented with the same prompts and their internal textual representations were examined at multiple depths.
132
+
133
+ The models do not use the same hidden dimensionality, so their representations cannot simply be compared directly.
134
+
135
+ A learned projection was therefore introduced during the research stage.
136
+
137
+ The objective was not initially to improve H3 generation. It was simply to answer a more fundamental question:
138
+
139
+ > **Can a representation produced by one architecture be predictably mapped into a representation produced by another architecture?**
140
+
141
+ Multiple internal representation combinations were screened.
142
+
143
+ One correspondence was substantially stronger than the alternatives.
144
+
145
+ On held-out prompts from the initial dataset, the learned cross-model mapping reached approximately:
146
+
147
+ ```text
148
+ validation cosine similarity ≈ 0.904
149
+ ```
150
+
151
+ This was encouraging, but there was an obvious danger.
152
+
153
+ The mapping might simply have learned the distribution of prompts used to construct it.
154
+
155
+ A stricter test was required.
156
+
157
+ ---
158
+
159
+ # 4. Out-of-Distribution Test
160
+
161
+ A separate set of **160 new prompts** was constructed.
162
+
163
+ These prompts were designed to stress combinations such as:
164
+
165
+ * complex spatial relationships;
166
+ * difficult counting;
167
+ * anatomy;
168
+ * material/light interactions;
169
+ * text;
170
+ * architecture and vehicles;
171
+ * reflection and occlusion;
172
+ * longer multi-object compositions.
173
+
174
+ The prompts were kept separate from the original fitting set.
175
+
176
+ When the frozen cross-model bridge was evaluated on these new prompts, representation similarity decreased from approximately:
177
+
178
+ ```text
179
+ 0.904
180
+ ```
181
+
182
+ to:
183
+
184
+ ```text
185
+ 0.749
186
+ ```
187
+
188
+ This drop was significant.
189
+
190
+ But the mapping did not collapse.
191
+
192
+ That distinction was important.
193
+
194
+ A completely prompt-specific mapping would be expected to fail much more severely once moved outside its fitting distribution. Instead, a substantial relationship remained.
195
+
196
+ This suggested that the bridge had captured at least part of a genuine cross-model representation alignment rather than simply memorizing sentences.
197
+
198
+ At this point, numerical similarity alone was no longer sufficient.
199
+
200
+ The mapping had to be tested inside the actual video-generation pipeline.
201
+
202
+ ---
203
+
204
+ # 5. The Experimental Full Bridge
205
+
206
+ The next implementation introduced the mapped SenseNova representation directly into MiniMax H3 conditioning.
207
+
208
+ Conceptually, the system became:
209
+
210
+ ```text
211
+ Prompt
212
+ │
213
+ ├──────────────► H3 Text Encoder
214
+ │ │
215
+ │ ▼
216
+ │ native conditioning H
217
+ │
218
+ └──────────────► SenseNova
219
+ │
220
+ ▼
221
+ semantic representation
222
+ │
223
+ ▼
224
+ learned projection
225
+ │
226
+ ▼
227
+ H3-compatible semantic S
228
+ ```
229
+
230
+ The transformed semantic representation was magnitude-aligned with native H3 conditioning and combined through a controllable residual blend:
231
+
232
+ ```text
233
+ C = H + α(S - H)
234
+ ```
235
+
236
+ where:
237
+
238
+ ```text
239
+ H = native H3 conditioning
240
+ S = mapped semantic representation
241
+ α = bridge strength
242
+ C = final conditioning supplied to H3
243
+ ```
244
+
245
+ This became the **Full Bridge**.
246
+
247
+ Crucially, H3 itself was not modified.
248
+
249
+ The video transformer remained unchanged.
250
+
251
+ The intervention occurred entirely in the conditioning supplied to it.
252
+
253
+ ---
254
+
255
+ # 6. From Tensor Similarity to Actual Generation
256
+
257
+ The Full Bridge was then tested through actual MiniMax H3 generation.
258
+
259
+ This was the first point at which the experiment became substantially more interesting.
260
+
261
+ The mapped conditioning did not merely remain numerically valid. It produced coherent H3 generations and observable changes in how prompts were interpreted.
262
+
263
+ This demonstrated that the transferred representation existed in a region of conditioning space that H3 could meaningfully use.
264
+
265
+ The Full Bridge therefore served as a proof of concept for three ideas:
266
+
267
+ **First**, useful information could be transferred between these architectures at the representation level despite the failure of direct parameter grafting.
268
+
269
+ **Second**, the transferred representation could influence H3 without modifying its generative transformer.
270
+
271
+ **Third**, the effect could be controlled continuously through the blend strength.
272
+
273
+ But the solution was impractical.
274
+
275
+ Every generation required the donor model to participate in prompt processing.
276
+
277
+ The runtime pipeline therefore contained two large models simply to produce one conditioning tensor.
278
+
279
+ That was acceptable for research.
280
+
281
+ It was not a good release architecture.
282
+
283
+ ---
284
+
285
+ # 7. Removing the Donor
286
+
287
+ This led to the central distillation experiment.
288
+
289
+ If the Full Bridge could determine an appropriate semantic correction for H3, perhaps H3's own representation already contained enough information for a smaller model to **predict what the Full Bridge would have done**.
290
+
291
+ The Full Bridge was therefore converted from an inference solution into a **teacher**.
292
+
293
+ The objective changed from:
294
+
295
+ ```text
296
+ SenseNova
297
+ ↓
298
+ map representation into H3
299
+ ```
300
+
301
+ to:
302
+
303
+ ```text
304
+ H3 representation
305
+ ↓
306
+ small student
307
+ ↓
308
+ predict teacher-derived
309
+ semantic representation
310
+ ```
311
+
312
+ If successful, SenseNova could disappear completely from inference.
313
+
314
+ ---
315
+
316
+ # 8. Early Distillation
317
+
318
+ The first distilled experiments attempted to predict the semantic correction directly.
319
+
320
+ These results were useful but not sufficiently strong.
321
+
322
+ The best early single-representation configuration achieved correction prediction similarity of only roughly:
323
+
324
+ ```text
325
+ ≈ 0.51
326
+ ```
327
+
328
+ A naïve attempt to combine several H3 representations performed even worse.
329
+
330
+ This was another useful negative result.
331
+
332
+ Simply providing the student with more internal features did not automatically produce a better approximation.
333
+
334
+ The target itself needed to be reconsidered.
335
+
336
+ Rather than asking the student to directly predict the difference between native and teacher conditioning, the next version learned the **teacher-derived semantic representation itself**.
337
+
338
+ This changed the result dramatically.
339
+
340
+ ---
341
+
342
+ # 9. Distilled Student V2
343
+
344
+ The second student learned to approximate the semantic representation produced by the Full Bridge before the final blend with H3 conditioning.
345
+
346
+ On the development distribution, representation similarity reached approximately:
347
+
348
+ ```text
349
+ 0.995
350
+ ```
351
+
352
+ and semantic-correction similarity reached approximately:
353
+
354
+ ```text
355
+ 0.973
356
+ ```
357
+
358
+ This was a major improvement.
359
+
360
+ However, the strict OOD evaluation revealed a remaining weakness.
361
+
362
+ Across the previously unseen 160-prompt set, representation similarity remained strong:
363
+
364
+ ```text
365
+ 0.941840
366
+ ```
367
+
368
+ but correction similarity fell to:
369
+
370
+ ```text
371
+ 0.879840
372
+ ```
373
+
374
+ The model clearly understood much of the teacher transformation, but the generalization gap was still visible.
375
+
376
+ Long, highly compositional prompts were among the more difficult cases.
377
+
378
+ This motivated one final training stage.
379
+
380
+ ---
381
+
382
+ # 10. Final Distillation
383
+
384
+ The original and difficult prompt distributions were combined and the student was refined while preserving a held-out validation subset.
385
+
386
+ The final evaluation used:
387
+
388
+ ```text
389
+ 500 training prompts
390
+ 100 validation prompts
391
+ ```
392
+
393
+ with the validation set containing both original-distribution and strict OOD examples.
394
+
395
+ The final student produced:
396
+
397
+ ```text
398
+ Teacher representation similarity 0.995890
399
+ Semantic correction similarity 0.983558
400
+
401
+ Main-distribution correction 0.980888
402
+ OOD correction 0.989788
403
+
404
+ Minimum correction similarity 0.935910
405
+ ```
406
+
407
+ The final blended conditioning was even closer to the Full Bridge:
408
+
409
+ ```text
410
+ Bridge strength α = 0.10 0.999958
411
+ Bridge strength α = 0.20 0.999827
412
+ Bridge strength α = 0.30 0.999602
413
+ ```
414
+
415
+ The final evaluation therefore showed not only high approximation of the teacher representation, but strong approximation of the **actual semantic correction** produced by the Full Bridge.
416
+
417
+ ---
418
+
419
+ # 11. Performance Across Semantic Categories
420
+
421
+ The final validation set covered multiple semantic categories.
422
+
423
+ Correction similarity by category was:
424
+
425
+ | Category | Similarity |
426
+ | -------------------------- | ---------: |
427
+ | Text | 0.996061 |
428
+ | Reflection / occlusion | 0.995932 |
429
+ | Complex counting | 0.994557 |
430
+ | Complex text | 0.991912 |
431
+ | Reflection / occlusion OOD | 0.991855 |
432
+ | Complex spatial | 0.991125 |
433
+ | Long composition | 0.990637 |
434
+ | Architecture / vehicle | 0.990606 |
435
+ | Complex material / light | 0.987394 |
436
+ | Material | 0.986263 |
437
+ | Complex anatomy | 0.982104 |
438
+ | Anatomy | 0.980773 |
439
+ | Lighting | 0.978327 |
440
+ | Spatial | 0.974002 |
441
+ | Counting | 0.965250 |
442
+
443
+ One particularly interesting result is the recovery on long-composition prompts.
444
+
445
+ These had been among the weaker cases during earlier OOD testing, yet the final student reached approximately `0.991` correction similarity for this category.
446
+
447
+ This suggests that the student capacity itself was not necessarily the primary limitation. The distribution used during distillation was equally important.
448
+
449
+ ---
450
+
451
+ # 12. The Final Architecture
452
+
453
+ The research system had started as:
454
+
455
+ ```text
456
+ H3 + SenseNova + cross-model bridge
457
+ ```
458
+
459
+ The final system is simply:
460
+
461
+ ```text
462
+ MiniMax H3
463
+ │
464
+ ▼
465
+ H3 conditioning
466
+ │
467
+ ▼
468
+ Semantic Bridge
469
+ │
470
+ ▼
471
+ modified conditioning
472
+ │
473
+ ▼
474
+ MiniMax H3
475
+ video generation
476
+ ```
477
+
478
+ SenseNova is no longer involved.
479
+
480
+ The final adapter operates exclusively on information already produced by H3.
481
+
482
+ It predicts a learned semantic transformation, magnitude-aligns that representation with native conditioning, and blends the two according to a user-controlled strength.
483
+
484
+ No weights inside the MiniMax H3 diffusion transformer are modified.
485
+
486
+ ---
487
+
488
+ # 13. An Unexpected Qualitative Result
489
+
490
+ The final distilled adapter was originally expected to behave simply as a cheaper approximation of the Full Bridge.
491
+
492
+ Actual generation produced a more interesting result.
493
+
494
+ The distilled version preserves the general behavioral effect of the Full Bridge, but the resulting videos are not always visually identical.
495
+
496
+ In some of my tests, I actually preferred the output of the **distilled adapter** to the original Full Bridge.
497
+
498
+ This observation should be interpreted cautiously.
499
+
500
+ It does not prove that the distilled student is objectively superior to its teacher.
501
+
502
+ One possible explanation is a regularization effect.
503
+
504
+ The Full Bridge contains the complete mapped donor representation, including components that may not be consistently useful to H3. A compact student cannot reproduce every variation perfectly and may preferentially learn the more predictable structure of the transformation.
505
+
506
+ In that interpretation, distillation behaves somewhat like a semantic filter:
507
+
508
+ ```text
509
+ Full teacher signal
510
+ ↓
511
+ distillation
512
+ ↓
513
+ most reproducible transformation
514
+ ↓
515
+ H3
516
+ ```
517
+
518
+ This is currently a hypothesis rather than a demonstrated mechanism.
519
+
520
+ More qualitative testing is needed.
521
+
522
+ Nevertheless, it was unexpected because the purpose of distillation was originally efficiency—not improvement.
523
+
524
+ ---
525
+
526
+ # 14. What Semantic Bridge Actually Is
527
+
528
+ Semantic Bridge is **not a LoRA**.
529
+
530
+ It is not a checkpoint merge.
531
+
532
+ It is not a conventional model graft.
533
+
534
+ And the released file does not contain a second generative model.
535
+
536
+ A more accurate description is:
537
+
538
+ > **A compact conditioning-space adapter distilled from a cross-architecture semantic mapping.**
539
+
540
+ A LoRA modifies effective model weights:
541
+
542
+ ```text
543
+ base weights
544
+ +
545
+ low-rank weight update
546
+ ↓
547
+ modified network behavior
548
+ ```
549
+
550
+ Semantic Bridge instead modifies information entering the generative transformer:
551
+
552
+ ```text
553
+ prompt
554
+ ↓
555
+ H3 representation
556
+ ↓
557
+ Semantic Bridge
558
+ ↓
559
+ modified representation
560
+ ↓
561
+ unchanged H3 generator
562
+ ```
563
+
564
+ This distinction also means that Semantic Bridge and conventional LoRAs are conceptually complementary rather than mutually exclusive.
565
+
566
+ ---
567
+
568
+ # 15. What Is Being Released
569
+
570
+ There were two fundamentally different implementations during this research.
571
+
572
+ ### Full SenseNova → H3 Bridge
573
+
574
+ **Research prototype only — not released.**
575
+
576
+ This implementation required SenseNova during inference and was used to establish that cross-model semantic transfer could influence H3 generation.
577
+
578
+ It subsequently became the teacher system for distillation.
579
+
580
+ ### MiniMax H3 Semantic Bridge
581
+
582
+ **This is the version being released.**
583
+
584
+ It is the distilled standalone implementation.
585
+
586
+ It requires:
587
+
588
+ ```text
589
+ MiniMax H3
590
+ H3-compatible text encoder
591
+ MiniMaxH3_SemanticBridge_v1.safetensors
592
+ MiniMax H3 Semantic Bridge custom node
593
+ ```
594
+
595
+ It does **not** require:
596
+
597
+ ```text
598
+ SenseNova checkpoint
599
+ SenseNova runtime
600
+ SenseNova tokenizer
601
+ a second large language/multimodal model
602
+ the experimental Full Bridge
603
+ ```
604
+
605
+ SenseNova was involved in the research and training process only.
606
+
607
+ **No SenseNova installation is required to use the public release.**
608
+
609
+ ---
610
+
611
+ # 16. Installation
612
+
613
+ The release consists of the ComfyUI custom node and the distilled `.safetensors` adapter.
614
+
615
+ ### Custom Node
616
+
617
+ Copy:
618
+
619
+ ```text
620
+ MiniMax_H3_Semantic_Bridge
621
+ ```
622
+
623
+ to:
624
+
625
+ ```text
626
+ ComfyUI/custom_nodes/
627
+ ```
628
+
629
+ The resulting structure should look like:
630
+
631
+ ```text
632
+ ComfyUI/
633
+ └── custom_nodes/
634
+ └── MiniMax_H3_Semantic_Bridge/
635
+ ├── __init__.py
636
+ ├── nodes.py
637
+ ├── README.txt
638
+ ├── ADAPTER_INSTALLATION.txt
639
+ └── RELEASE_NOTES.txt
640
+ ```
641
+
642
+ ### Adapter
643
+
644
+ Copy:
645
+
646
+ ```text
647
+ MiniMaxH3_SemanticBridge_v1.safetensors
648
+ ```
649
+
650
+ to:
651
+
652
+ ```text
653
+ ComfyUI/models/semantic_bridge/
654
+ ```
655
+
656
+ The final location should therefore be:
657
+
658
+ ```text
659
+ ComfyUI/
660
+ └── models/
661
+ └── semantic_bridge/
662
+ └── MiniMaxH3_SemanticBridge_v1.safetensors
663
+ ```
664
+
665
+ Restart ComfyUI after installation.
666
+
667
+ ---
668
+
669
+ # 17. ComfyUI Nodes
670
+
671
+ Three nodes are included.
672
+
673
+ ### MiniMax H3 Image to Video + Semantic Bridge
674
+
675
+ This is the recommended node for normal use.
676
+
677
+ It replaces the standard MiniMax H3 Image to Video conditioning stage and applies the distilled semantic transformation internally.
678
+
679
+ It supports first-frame and last-frame conditioning and can therefore be used directly in standard H3 image-to-video workflows.
680
+
681
+ ### MiniMax H3 Semantic Bridge
682
+
683
+ This is the standalone conditioning version.
684
+
685
+ It accepts existing MiniMax H3 `CONDITIONING` and applies the Semantic Bridge transformation.
686
+
687
+ This is useful for custom workflows where H3 conditioning is already being generated elsewhere.
688
+
689
+ ### MiniMax H3 Clear Semantic Bridge Cache
690
+
691
+ An optional utility for clearing the small adapter from memory.
692
+
693
+ It is generally unnecessary during normal operation but is provided for workflow and debugging convenience.
694
+
695
+ ---
696
+
697
+ # 18. Recommended Settings
698
+
699
+ The recommended starting configuration is:
700
+
701
+ ```text
702
+ semantic_bridge:
703
+ MiniMaxH3_SemanticBridge_v1.safetensors
704
+
705
+ alpha:
706
+ 0.10
707
+
708
+ magnitude_match:
709
+ per_token
710
+ ```
711
+
712
+ `alpha` controls how strongly the predicted semantic representation influences native H3 conditioning.
713
+
714
+ Conceptually:
715
+
716
+ ```text
717
+ alpha = 0.00
718
+ │
719
+ │ Native H3
720
+ │
721
+ ├── 0.05
722
+ ├── 0.10 ← recommended starting point
723
+ ├── 0.20
724
+ └── 0.30
725
+ ```
726
+
727
+ A higher value does **not** necessarily mean a better result.
728
+
729
+ It simply increases the distance from native H3 conditioning.
730
+
731
+ For controlled comparisons I recommend testing:
732
+
733
+ ```text
734
+ 0.00
735
+ 0.05
736
+ 0.10
737
+ 0.20
738
+ 0.30
739
+ ```
740
+
741
+ while keeping all other generation parameters unchanged.
742
+
743
+ ---
744
+
745
+ # 19. How to Evaluate It
746
+
747
+ Semantic Bridge is intended primarily for prompts where **relationships** matter.
748
+
749
+ A weak test is:
750
+
751
+ > *A beautiful woman walking through a city.*
752
+
753
+ There is little semantic structure for the bridge to affect.
754
+
755
+ More informative tests involve multiple constraints:
756
+
757
+ > *A woman holds a transparent bottle in her left hand while pointing toward a red sign with her right hand. A polished metal sphere sits behind the bottle and reflects a blue object outside the frame.*
758
+
759
+ Useful evaluation areas include:
760
+
761
+ ```text
762
+ left / right relationships
763
+ front / behind
764
+ inside / outside
765
+ multiple people
766
+ specific body parts
767
+ object counts
768
+ reflections
769
+ transparent objects
770
+ material interactions
771
+ lighting direction
772
+ visible text
773
+ long compositions
774
+ multiple simultaneous instructions
775
+ ```
776
+
777
+ For a controlled A/B test, keep the following identical:
778
+
779
+ ```text
780
+ prompt
781
+ seed
782
+ source image
783
+ resolution
784
+ length
785
+ steps
786
+ sampler
787
+ guidance
788
+ ```
789
+
790
+ and compare:
791
+
792
+ ```text
793
+ A — Native MiniMax H3
794
+ B — MiniMax H3 + Semantic Bridge
795
+ ```
796
+
797
+ `alpha = 0.00` can also be used as a native-conditioning baseline within the Semantic Bridge node.
798
+
799
+ ---
800
+
801
+ # 20. Qualitative A/B Examples
802
+
803
+ Two controlled A/B pairs are included in the public repository. Within each pair, the prompt and generation setup were kept the same; the intended comparison is native H3 conditioning versus Semantic Bridge. The published Bridge examples use `alpha = 0.15` to make the behavioral difference easier to inspect.
804
+
805
+ These examples are qualitative observations, not a statistical benchmark.
806
+
807
+ ## Example 01 — Rooftop Train Chase
808
+
809
+ This test stresses complex motion, running anatomy, pursuit geometry, a specific vault interaction, hand-to-obstacle contact, landing continuity, wet reflective metal, rain, sparks, and camera motion.
810
+
811
+ The exact prompt and both videos are included under `examples/01_rooftop_train_chase/`.
812
+
813
+ ## Example 02 — Prompt Adherence, Materials, Reflection and Transparency
814
+
815
+ This test contains a particularly explicit instruction:
816
+
817
+ > **“Her right hand rests flat on the glass tabletop with all five fingers naturally separated and clearly visible.”**
818
+
819
+ In the native H3 generation, the character introduces an unrequested action by moving the right hand across the tabletop. In the Semantic Bridge generation at `alpha = 0.15`, the hand remains resting on the surface, more closely matching the explicitly requested state.
820
+
821
+ This observation is useful because the target is directly inspectable: the requested state is *resting*, rather than moving. It is still presented only as qualitative evidence from this specific controlled pair, not as proof of universal prompt-adherence improvement.
822
+
823
+ The same prompt also stresses left/right hand roles, exactly three tabletop objects in a specified order, text rendering, transparent glass, reflective metal, mirror correspondence, refraction, and mixed warm/cool lighting.
824
+
825
+ The exact prompt and both videos are included under `examples/02_glass_table_prompt_adherence/`.
826
+
827
+ ---
828
+
829
+ # 21. Scope: FL2VA / Standard Text-Conditioned H3 Only
830
+
831
+ An important limitation emerged during later testing.
832
+
833
+ The released Semantic Bridge was developed and distilled from the standard H3 text-conditioning path. It was not trained on the separate multimodal reference-conditioning distribution used by Ref2VA.
834
+
835
+ An experimental Ref2VA-compatible implementation was tested with reference images and reference audio. At a structural level it was possible to preserve visual-reference token positions while applying the adapter to text-designated positions. However, singing and lip-sync tests showed a clear practical regression: vocal articulation became less distinct and more mumbling-like than with native Ref2VA conditioning.
836
+
837
+ This suggests that preserving reference-image token positions is not sufficient. In Ref2VA, the text representation participates in a larger multimodal alignment involving text, reference media, audio, and visual performance. Altering only the text-designated part of that representation can still disturb the multimodal relationship.
838
+
839
+ The current release should therefore be treated as designed for:
840
+
841
+ **MiniMax H3 FL2VA / standard text-conditioned generation.**
842
+
843
+ It should not currently be considered compatible with:
844
+
845
+ **MiniMax H3 Ref2VA / reference-conditioned generation**, particularly workflows involving reference audio, singing, voice-driven performance, or lip synchronization.
846
+
847
+ This negative result does not establish that semantic bridging is fundamentally incompatible with Ref2VA. It indicates that a Ref2VA bridge should likely be trained separately on the multimodal conditioning regime it is intended to modify.
848
+
849
+ A broader lesson is that **representation compatibility is contextual, not merely dimensional**. Two conditioning tensors can have the same dimensionality while participating in different semantic and temporal alignment mechanisms.
850
+
851
+ ---
852
+
853
+ # 22. Limitations
854
+
855
+ This is an experimental research release, not a universal H3 enhancement.
856
+
857
+ The adapter should not be expected to improve every generation.
858
+
859
+ MiniMax H3 already has strong semantic capabilities, and many simple prompts do not require any additional transformation.
860
+
861
+ There may also be prompts for which native conditioning produces a preferable result.
862
+
863
+ Another important limitation is that the reported numerical metrics measure agreement with the experimental teacher system.
864
+
865
+ They do **not** directly measure video quality.
866
+
867
+ A student that perfectly reproduces its teacher is not necessarily a better video generator, and a small difference in conditioning can sometimes result in a substantial visual difference after generative sampling.
868
+
869
+ For this reason, actual controlled video comparisons remain the most important evaluation.
870
+
871
+ The current research also focuses on a particular group of semantic behaviors. Other areas may respond differently and have not yet been evaluated systematically.
872
+
873
+ Finally, Semantic Bridge does not create information that H3 fundamentally cannot represent. It should be viewed as a learned transformation of H3's existing semantic conditioning, not as an independent reasoning system.
874
+
875
+ ---
876
+
877
+ # 23. Why the Result May Matter Beyond MiniMax H3
878
+
879
+ The most interesting result of this experiment may not be the released adapter itself.
880
+
881
+ Direct weight transfer between unrelated architectures is extremely restrictive.
882
+
883
+ Different hidden dimensions, layer structures, attention implementations and parameter organizations make conventional grafting difficult or meaningless.
884
+
885
+ Representations are different.
886
+
887
+ If two models have learned related concepts, their internal spaces do not necessarily need to be structurally identical for a **learnable relationship** to exist between them.
888
+
889
+ This suggests a more general workflow:
890
+
891
+ ```text
892
+ Model A
893
+ │
894
+ │ semantic teacher
895
+ ▼
896
+ cross-architecture
897
+ representation mapping
898
+ │
899
+ ▼
900
+ Model B-compatible
901
+ teacher signal
902
+ │
903
+ ▼
904
+ distillation
905
+ │
906
+ ▼
907
+ small Model B adapter
908
+ ```
909
+
910
+ After distillation:
911
+
912
+ ```text
913
+ Model A
914
+ ✕
915
+ no longer required
916
+
917
+
918
+ Model B
919
+ +
920
+ small distilled adapter
921
+ ```
922
+
923
+ The expensive cross-model system can therefore potentially exist only during research and training.
924
+
925
+ The deployed model remains essentially its original architecture.
926
+
927
+ This opens an interesting research direction: instead of asking whether two models can share weights, we can ask whether they can **teach each other representations**.
928
+
929
+ ---
930
+
931
+ # 24. Conclusion
932
+
933
+ MiniMax H3 Semantic Bridge began as an attempt to graft semantic capabilities between two incompatible model architectures.
934
+
935
+ That approach failed.
936
+
937
+ The failure led to a more interesting question.
938
+
939
+ Rather than transferring weights, I investigated whether the models' representations of the same prompt contained a learnable relationship.
940
+
941
+ A measurable cross-model alignment was found.
942
+
943
+ That alignment was then incorporated into an experimental Full Bridge, which demonstrated that a mapped donor representation could meaningfully influence actual MiniMax H3 generation.
944
+
945
+ The Full Bridge solved the semantic-transfer problem but introduced an impractical runtime dependency on the donor model.
946
+
947
+ Distillation removed that dependency.
948
+
949
+ The final system is therefore much simpler than the research pipeline that produced it:
950
+
951
+ ```text
952
+ MiniMax H3
953
+ +
954
+ MiniMax H3 Semantic Bridge
955
+ ```
956
+
957
+ The public adapter contains no SenseNova checkpoint and requires no SenseNova installation or inference pass.
958
+
959
+ It operates entirely from H3's own conditioning and introduces the learned transformation at a controllable strength.
960
+
961
+ The result suggests a broader possibility:
962
+
963
+ > **Useful learned behavior may be transferable between incompatible model architectures even when their weights themselves are not.**
964
+
965
+ The bridge between models may not have to exist in parameter space.
966
+
967
+ It may exist in representation space.
968
+
969
+ And once that bridge has taught a sufficiently small student, the original teacher may no longer need to be there at all.
970
+
971
+ ---
972
+
973
+ ## Release Summary
974
+
975
+ **MiniMax H3 Semantic Bridge v1.0**
976
+
977
+ Standalone distilled semantic conditioning adapter for MiniMax H3.
978
+
979
+ **Release includes:**
980
+
981
+ * ComfyUI custom node
982
+ * distilled Semantic Bridge `.safetensors`
983
+
984
+ **Release does not require:**
985
+
986
+ * SenseNova
987
+ * the experimental Full Bridge
988
+ * a second multimodal model at inference
989
+
990
+ **Recommended starting configuration:**
991
+
992
+ ```text
993
+ alpha = 0.10
994
+ magnitude_match = per_token
995
+ ```
996
+
997
+ **Primary experimental targets:**
998
+
999
+ ```text
1000
+ complex composition
1001
+ spatial relationships
1002
+ anatomy
1003
+ counting
1004
+ prompt following
1005
+ text
1006
+ materials
1007
+ lighting
1008
+ reflection
1009
+ transparency
1010
+ occlusion
1011
+ ```
1012
+
1013
+ **Status:** Experimental / Research Release.
1014
+
1015
+ **Supported path:** MiniMax H3 FL2VA / standard text-conditioned generation.
1016
+
1017
+ **Not supported in v1:** Ref2VA / reference-conditioned generation, particularly reference-audio singing and lip-sync.
SHA256SUMS.txt ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ 005cd62443f3115ab3d32889c9c999596c5e2ab23be3f04a1c8febfb6f7b964d research/reports/distillation_source_inspection.txt
2
+ 020b21411e7a63a4a6684c0e8a35ef98c5f0f6e55d774b8d0313396402da6c44 research/reports/sensenova_mot_gen_bridge_comparison.txt
3
+ 039f12ced4a634c954ed25c22973cdcf95e14299b0eb3c61d362ea489d646f6e research/reports/h3_sensenova_bridge_report.txt
4
+ 07665ef280012a8ff6747db0edb1617a8d63f323277b8130da5042fbefd823aa research/reports/hidden_bridge_screening_rank128.txt
5
+ 0ff88375ec2ae942ad5effa801f36ba700e4cbb07798df6fe495479363ac656f research/raw_scripts/train_distilled_student_v2.py
6
+ 16d0530e314f734c03d52e4c4d6c8838d6df5e7d63f10322554a498112758e14 examples/01_rooftop_train_chase/native_h3.mp4
7
+ 1b012bd4ff4d6edc2e5a72a06ce69409c09ca35e6964b5491a6c0ef6f28ce82c research/reports/hidden_layer_bridge_comparison.txt
8
+ 1f9d9d2ba4ef6c692abf791de180f906ee77fc943500fc3e2f9673680a584461 research/raw_scripts/compare_tokenizers.py
9
+ 2c19120bfcff95fdf20a58b428ac83dad25f8da8fa3e308aa38e415006f240fb research/reports/distilled_student_v2_results.csv
10
+ 2d6c1f0af8bce0057a6d84bb0041aceb00a1148bd6d20d0fb6a845db1f0f1448 research/raw_scripts/evaluate_distilled_student_v2_ood.py
11
+ 3185139654470d93b67e6ce80eb1a7d2594b3c84d94ab93eb85c5de4701d805b research/raw_scripts/compare_minimax_sensenova.py
12
+ 318a764ad464823c134b66e348a1b17bff9e1d4830e2a4ec1cbc455620425912 research/raw_scripts/extract_h3_ood_hidden.py
13
+ 35ab00a9d44a75674d2148c67601817d7e0950337b9b82dec228c0c893e6681c research/datasets/bridge_prompts_480.json
14
+ 3ee5db6c7c12f574541b82da93ca11a91037af307dce0e1d0c98d10f5229c429 research/reports/sensenova_mot_gen_bridge_comparison.csv
15
+ 3f73f78b2b50a397c7b377a9d2a01aec05885cea3e7dab2fc873432c3b72a3fe UPSTREAM_LICENSES.md
16
+ 47dc42c5292eddd10e3d55c9d169a03a299e275af51aaef5719ffb660cb337bd research/raw_scripts/make_bridge_ood_prompts.py
17
+ 4a01bd07c6333b5bf9172429fd8043f150886a62b07f2c5de117cfd0e0b73b14 research/datasets/bridge_ood_prompts_160.json
18
+ 5150010c15ca4042cbb3ed5162957ecb79a0e57f83036bfbad60a6713b717492 research/teacher_artifacts/SN_L32_to_H3_L49_rank128.safetensors
19
+ 5333910cda814464c7a147f1de1ccc11281e49b30a98663ab54c23e3bcb9c9b0 research/reports/distilled_student_v2_STRICT_OOD_report.txt
20
+ 557b6be384c5f0614a3fa1b12428930d29829edf826127c8b7badfc200f8407c research/raw_scripts/extract_h3_hidden_states.py
21
+ 597990c07be911a937990f30267830bc336ff85316fbd1f81f6f217db7faf7be research/raw_scripts/extract_sensenova_bridge_dataset.py
22
+ 5a10e4a1b9bc0b10d72233101af14ec8aafb2b44bb35b732c334e7275804e288 research/raw_scripts/prepare_sensenova_h3_condition.py
23
+ 5c8db67192acfc4666797acc45753092b63185ba4164a1ffc868033468e01635 research/reports/distilled_adapter_screening.csv
24
+ 627a9abfefcf0f457ab2d8f3d6f8d5baecabaaed3d48a65f4f1b862c08c71b06 research/reports/distilled_student_v3_report.txt
25
+ 635689b0bb6c53ac8add559f4aa32bb5ee66355291771d6c87ead8c857f4ac7a LICENSE.md
26
+ 64e9e9eea85402ab2477cb597e2315d77c2a26d67b23fc800e238efd35ebf421 research/raw_scripts/inspect_distillation_data.py
27
+ 6cd2dd305e3f4d3c67c3609976f4be34adecb4b7bd7a51d90c9719ac2ae5ff69 research/raw_scripts/extract_sensenova_hidden_states.py
28
+ 7a05ffc62356607eace846fd755b42e7d036b1efa304167941fd183f61c4b3d4 research/raw_scripts/make_bridge_prompts.py
29
+ 7e9506dbc9ce6e304aabb991e935a8519e78770d35345496239c522f4b92d7c5 research/raw_scripts/extract_sensenova_hidden_states_stream_v3.py
30
+ 7ed600988c2d48f03cf0ad7bee35d8f17427b45aa25433a4eaab55e7b732d221 research/reports/distilled_student_v2_results.txt
31
+ 8f30ddf1b781bc9003e235d418754593ff04fcf3e7f673666046a1c6cb3adc1e README.md
32
+ 902580d71cfd1cd1cbf7b6da6254428881ef5af9b2a7d47b7774b3e150a5e52c research/reports/distilled_student_v3_history.csv
33
+ 924ce5be3be8a7f1a06296a89398d0a2e0e714da065b34f961f64e257e1f09dc RESEARCH_ARTICLE.md
34
+ 95e04c8b56655aa9f610292140c4820158c656e89ff08973d75fb5171a2b8499 research/raw_scripts/compare_hidden_layers_bridge.py
35
+ 98252ac062a0020208deddef9d99c7b0e2702d4bff9d058196cbd01d6a086af6 research/raw_scripts/train_hidden_bridge_screening.py
36
+ 9dfaadf1df0e016473807242e64b87de658d19d780125c7554e4ec4e7d8a725f research/raw_scripts/inspect_h3_sensenova_bridge.py
37
+ a019fddb4c079c72bc5c393cbcf2a905e57138b4a56e91cee6227147d9be6bed research/raw_scripts/train_distilled_student_v3.py
38
+ ac0dc8ac05f545ebdee12e2fcebe4515b049f9cfd9558eb4887a9bf3fd6d562e MiniMaxH3_SemanticBridge_v1.safetensors
39
+ b4fbdce02a44a6fa7d42b1272de404e5de30eb608200fc6f693e1e136aafb0c1 research/raw_scripts/extract_sensenova_ood_hidden.py
40
+ b60c5b367c166751c798024715f37b1b37a09ffbd0fd71402c3b3a48dc8896ce research/reports/hidden_layer_bridge_comparison.csv
41
+ bb649b0c49e43e82f7ffa8b7be37f20de793d925a6ec55c2ab0d2d503dd1899f examples/02_glass_table_prompt_adherence/native_h3.mp4
42
+ c1b6452bf0e43d153bd81f45a6688b35ee8ee4df95951c21d23e0ac717a8d67b research/reports/hidden_bridge_screening_rank128.csv
43
+ c4866a0cc554237046814c9cb827a5960caecbd2214d5a5a04b0e43716b437bb MiniMax_H3_Semantic_Bridge_v1.0.zip
44
+ c8e799776c9ac010c76ba7eeb11f481b6c130034920030cbbfa584a133f5f404 research/reports/distilled_student_v2_STRICT_OOD_per_prompt.csv
45
+ c90f5447164e3b72e02bdfe5cd484e8d28fcbd2842aae6233c9b840d77d65451 examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4
46
+ c92aa3822cec5952b59abff26ac426fa4224aa35b94eeb5f0bac903bf22fb5de research/raw_scripts/extract_sensenova_hidden_states_stream.py
47
+ d06b5cf4154e44e869b007d30528d8ff4e2324245c4be48fa0619e92791ee62f research/raw_scripts/extract_h3_embeddings_full.py
48
+ d2cca56ac4c8aca46f3d3b4e32dbb7d18e78db7bbab8f3ff5ecd5656ba293a97 NOTICE.txt
49
+ d3efa1ef791d1f9aa6a3fdffc2f0689d4be34b01dd15043aafd52111b31d6cd2 research/raw_scripts/train_distilled_adapter_screening.py
50
+ d5754f336f4dd42f8cd8c53667992475252bf4b187ee73bb928bdd368e03d397 examples/01_rooftop_train_chase/prompt.txt
51
+ d5e10235e90dc379e98852808810dcc92de94d29bd68a40a85cf5816cd64132d examples/02_glass_table_prompt_adherence/prompt.txt
52
+ d9a8911c11e3d305e0f7af0e0d45a190a33bd1d25ce50b44e298f37fc6b513a8 research/reports/hidden_bridge_OOD_final_report.txt
53
+ e149bd39e003f48c8f5a8c9af31b6aa78e24cb5d0df9cef81d243a8bc08f4c8e examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4
54
+ e60b14e081d21982b46aedf28e2f3713616cc5e4f89ac117405dd7c03ab22d8d research/raw_scripts/make_bridge_ood_prompts_v1.py
55
+ e6c3875c4f82b75c25809667897f61fdccd8001860e9bbb88d50f519323940b7 research/raw_scripts/extract_h3_hidden_states_fast.py
56
+ ea9cfbe8d485093904bf92cd4c1fc2a5c7f38a059e1bf208fe377f66d97a8d2e research/raw_scripts/extract_h3_embeddings_test.py
57
+ ed67e7fa5f49bbbca9294bc8055d9f571457d03b183b963c40bfffc34fe5f606 research/raw_scripts/evaluate_hidden_bridge_ood.py
58
+ ee773a550f35269187285ae98766493658c91524663c3aa07008616b298a1304 research/raw_scripts/extract_h3_bridge_dataset.py
59
+ f020e87b6f30537be4f219d20bc289480bc93dd9f6f90204f102a08c07db3389 research/raw_scripts/extract_sensenova_hidden_states_local.py
60
+ f40410dd8d9ffaa18bf149d200646ff8268fa5531efcde53b37c34b7fd5f5442 research/reports/hidden_bridge_OOD_per_prompt.csv
61
+ f4186926cb7ac1f4c9e59470a013f780fd21d8a053d39bd9ce74179662d1f190 research/raw_scripts/extract_sensenova_hidden_states_stream_v2.py
62
+ f4be9eb6ba7ffa1207f1d101a1e6690c4c48243df2eb5401161f3a5e5202169f research/raw_scripts/train_embedding_bridge_test.py
63
+ f6373b1a355f2cfb47e7441c4c43c196ccbb173121c0492793844019059b7fed research/reports/distilled_adapter_screening.txt
64
+ fbd662aee9e8b11da626956118cc22299aef550d5141b8e90155516698ffaa5f research/reports/minimax_vs_sensenova_report.txt
65
+ fc540af96938797b4a75da4ffd85b04e69eec4c4ad55f7fac86c7c689f6f6a42 research/raw_scripts/test_sensenova_mot_gen_bridge.py
UPSTREAM_LICENSES.md ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Upstream projects and licenses
2
+
3
+ ## MiniMax H3
4
+
5
+ Project: https://huggingface.co/MiniMaxAI/MiniMax-H3
6
+
7
+ License: MiniMax H3 Community License Agreement
8
+
9
+ License text: https://huggingface.co/MiniMaxAI/MiniMax-H3/blob/main/LICENSE
10
+
11
+ The agreement contains terms covering MiniMax H3 Works and Model Derivatives, distribution requirements, use restrictions, and an Applicable Territory definition. Users should read the official agreement directly.
12
+
13
+ ## SenseNova U1.5
14
+
15
+ Project: https://huggingface.co/sensenova/SenseNova-U1.5-8B-MoT
16
+
17
+ License: Apache License 2.0
18
+
19
+ SenseNova was used as an experimental semantic teacher during development. The public runtime adapter does not require the SenseNova model.
20
+
21
+ ## Qwen3-VL
22
+
23
+ MiniMax H3 documentation notes that its encoder uses Qwen3-VL-32B, licensed under Apache License 2.0. Refer to the MiniMax H3 license/documentation for the upstream notice and link.
examples/01_rooftop_train_chase/native_h3.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:16d0530e314f734c03d52e4c4d6c8838d6df5e7d63f10322554a498112758e14
3
+ size 2491328
examples/01_rooftop_train_chase/prompt.txt ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ Scene overview: A cinematic high-speed action scene on the rooftop of a moving train racing through a dense futuristic city at night. A young woman in a black leather jacket and dark cargo pants runs along the roof while being pursued by a masked attacker. Rain pours heavily, making the metal surface wet and reflective. Neon signs, illuminated windows and city lights streak past in the background.
2
+
3
+ Storyboard:
4
+ [0s-1.5s] The camera tracks backward directly in front of the woman as she sprints toward it along the narrow train roof. Her coat and hair whip violently in the wind. She briefly looks over her shoulder at the masked attacker running several meters behind her. Both characters maintain natural running anatomy and believable foot contact with the moving train.
5
+
6
+ [1.5s-3.2s] She reaches a low metal obstacle spanning part of the roof, plants her left foot and performs a fast running vault over it, supporting herself momentarily with her right hand. The camera swings dynamically from the front into a low three-quarter side angle while following the movement.
7
+
8
+ [3.2s-5s] She lands firmly on both feet and immediately continues running. Behind her, the attacker jumps over the same obstacle. At that exact moment the train passes beneath a shower of bright sparks from an overhead electrical structure. The sparks illuminate both figures for an instant while their reflections slide across the rain-covered metal roof.
9
+
10
+ Fast controlled camera movement, strong sense of forward momentum and physical weight. Realistic human anatomy throughout the running, vaulting and landing motions. Hands and limbs remain anatomically coherent. Accurate interaction between feet, hands and physical surfaces. Wet brushed metal shows sharp moving reflections, leather remains glossy but textured, fabric reacts naturally to wind and motion, rain catches the surrounding neon light. Strong depth and spatial continuity between the woman, attacker, train roof and rapidly moving city background. No slow motion.
examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e149bd39e003f48c8f5a8c9af31b6aa78e24cb5d0df9cef81d243a8bc08f4c8e
3
+ size 2538282
examples/02_glass_table_prompt_adherence/native_h3.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bb649b0c49e43e82f7ffa8b7be37f20de793d925a6ec55c2ab0d2d503dd1899f
3
+ size 358133
examples/02_glass_table_prompt_adherence/prompt.txt ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ A cinematic medium-wide shot of a young woman standing behind a transparent glass table in a sophisticated modern apartment at night. She is wearing a fitted black sleeveless dress and a thin silver necklace.
2
+
3
+ Her left hand holds a polished stainless-steel cup by its handle at chest height. Her right hand rests flat on the glass tabletop with all five fingers naturally separated and clearly visible. Her hands must remain anatomically correct and attached naturally to her arms.
4
+
5
+ On the table in front of her are three distinct objects arranged from left to right: a clear glass bottle filled halfway with water, a small red ceramic bowl, and a folded white paper card with the words "NIGHT SHIFT" printed in large clean black letters.
6
+
7
+ A large wall mirror directly behind the woman reflects the back of her head, shoulders, the room behind the camera, and the warm floor lamp standing on the left side of the room. The reflected objects must correspond logically to their real positions.
8
+
9
+ Warm amber light from the floor lamp passes through the glass bottle, creating subtle refraction and a bright caustic pattern on the tabletop. Cool blue city light enters through a window on the right and produces contrasting highlights along the polished metal cup.
10
+
11
+ The transparent tabletop must remain visibly transparent, with the woman's lower body partially visible through the glass. The metal cup should have realistic sharp reflections, the bottle should refract the background, the ceramic bowl should remain opaque and matte, and the mirror should behave as a true reflective surface.
12
+
13
+ Natural human proportions, coherent spatial relationships, physically plausible reflections and transparency, realistic hands, detailed material differences, clear readable text.
examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c90f5447164e3b72e02bdfe5cd484e8d28fcbd2842aae6233c9b840d77d65451
3
+ size 329053
research/datasets/bridge_ood_prompts_160.json ADDED
@@ -0,0 +1,642 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "category": "complex_counting",
4
+ "prompt": "Exactly 3 glass bottles form two staggered rows, while exactly 2 blue books occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 1. Counting variant 1."
5
+ },
6
+ {
7
+ "category": "complex_text",
8
+ "prompt": "Inside a minimalist living room, the exact phrase \"PLATFORM 11\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 13."
9
+ },
10
+ {
11
+ "category": "complex_counting",
12
+ "prompt": "Exactly 3 silver spheres form a curved row, while exactly 4 red cylinders occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 1. Counting variant 16."
13
+ },
14
+ {
15
+ "category": "architecture_vehicle",
16
+ "prompt": "A white city bus passes beside a railway platform. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 17."
17
+ },
18
+ {
19
+ "category": "complex_anatomy",
20
+ "prompt": "A young woman raises the left hand while holding a cup in the right hand. A woman wearing a dark jacket stands beside them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 1."
21
+ },
22
+ {
23
+ "category": "architecture_vehicle",
24
+ "prompt": "A silver tram passes beside a industrial studio. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 15."
25
+ },
26
+ {
27
+ "category": "complex_spatial",
28
+ "prompt": "A small lamp stands behind a ceramic cup, while a blue book is positioned to their left. The ceramic cup partially overlaps the small lamp from the camera viewpoint. Exactly 3 major objects are visible in the composition. Spatial variant 9."
29
+ },
30
+ {
31
+ "category": "complex_material_light",
32
+ "prompt": "A ceramic cup made from dark polished wood rests on a surface made from brushed aluminum, illuminated by two opposing light sources. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 8."
33
+ },
34
+ {
35
+ "category": "reflection_occlusion_ood",
36
+ "prompt": "A man wearing a gray coat stands behind a transparent glass panel. A ceramic cup in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a transparent cube located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 12. Reflection variant 12."
37
+ },
38
+ {
39
+ "category": "long_composition",
40
+ "prompt": "An elderly man stands at the left side of a glass table while a man wearing a gray coat sits to the right. The standing person holds a red cylinder in the left hand and points toward a sign reading exactly \"FINAL STOP\" with the right hand. A ceramic cup and a transparent cube lie on the table in that order from left to right. The scene is illuminated by cool light entering from the right. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 6 major foreground objects should remain clearly distinguishable. Long-scene variant 10."
41
+ },
42
+ {
43
+ "category": "complex_spatial",
44
+ "prompt": "A ceramic cup stands behind a blue book, while a black vase is positioned to their right. The blue book partially overlaps the ceramic cup from the camera viewpoint. Exactly 4 major objects are visible in the composition. Spatial variant 2."
45
+ },
46
+ {
47
+ "category": "architecture_vehicle",
48
+ "prompt": "A red compact car passes beside a modern kitchen. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 16."
49
+ },
50
+ {
51
+ "category": "long_composition",
52
+ "prompt": "A tall man stands at the left side of a glass table while a woman with short hair sits to the right. The standing person holds a metal box in the left hand and points toward a sign reading exactly \"PLATFORM 11\" with the right hand. A blue book and a red cylinder lie on the table in that order from left to right. The scene is illuminated by strong backlight. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 7 major foreground objects should remain clearly distinguishable. Long-scene variant 3."
53
+ },
54
+ {
55
+ "category": "complex_spatial",
56
+ "prompt": "A small lamp stands behind a ceramic cup, while a blue book is positioned to their left. The ceramic cup partially overlaps the small lamp from the camera viewpoint. Exactly 5 major objects are visible in the composition. Spatial variant 19."
57
+ },
58
+ {
59
+ "category": "complex_text",
60
+ "prompt": "Inside a underground station, the exact phrase \"WEST EXIT\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 4."
61
+ },
62
+ {
63
+ "category": "architecture_vehicle",
64
+ "prompt": "A white city bus passes beside a industrial studio. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 7."
65
+ },
66
+ {
67
+ "category": "reflection_occlusion_ood",
68
+ "prompt": "A woman wearing a dark jacket stands behind a transparent glass panel. A ceramic cup in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a transparent cube located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 2. Reflection variant 2."
69
+ },
70
+ {
71
+ "category": "architecture_vehicle",
72
+ "prompt": "A blue bicycle passes beside a railway platform. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 9."
73
+ },
74
+ {
75
+ "category": "complex_text",
76
+ "prompt": "Inside a industrial studio, the exact phrase \"STUDIO C\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 6."
77
+ },
78
+ {
79
+ "category": "complex_anatomy",
80
+ "prompt": "A man wearing a gray coat holds a bottle with both hands directly in front of the chest. A young woman stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 14."
81
+ },
82
+ {
83
+ "category": "complex_spatial",
84
+ "prompt": "A silver sphere stands behind a small lamp, while a ceramic cup is positioned to their right. The small lamp partially overlaps the silver sphere from the camera viewpoint. Exactly 6 major objects are visible in the composition. Spatial variant 16."
85
+ },
86
+ {
87
+ "category": "complex_anatomy",
88
+ "prompt": "A tall man holds a bottle with both hands directly in front of the chest. A man wearing a gray coat stands beside them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 19."
89
+ },
90
+ {
91
+ "category": "complex_text",
92
+ "prompt": "Inside a small workshop, the exact phrase \"ROOM 314\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 2."
93
+ },
94
+ {
95
+ "category": "complex_material_light",
96
+ "prompt": "A ceramic cup made from frosted glass rests on a surface made from translucent amber acrylic, illuminated by cool light entering from the right. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 18."
97
+ },
98
+ {
99
+ "category": "complex_spatial",
100
+ "prompt": "A silver sphere stands behind a small lamp, while a ceramic cup is positioned to their right. The small lamp partially overlaps the silver sphere from the camera viewpoint. Exactly 4 major objects are visible in the composition. Spatial variant 6."
101
+ },
102
+ {
103
+ "category": "architecture_vehicle",
104
+ "prompt": "A red compact car passes beside a railway platform. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 1."
105
+ },
106
+ {
107
+ "category": "complex_counting",
108
+ "prompt": "Exactly 7 blue books form two staggered rows, while exactly 2 small lamps occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 5. Counting variant 5."
109
+ },
110
+ {
111
+ "category": "complex_counting",
112
+ "prompt": "Exactly 3 glass bottles form two staggered rows, while exactly 2 blue books occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 2. Counting variant 11."
113
+ },
114
+ {
115
+ "category": "complex_spatial",
116
+ "prompt": "A wooden chair stands behind a transparent cube, while a red cylinder is positioned to their right. The transparent cube partially overlaps the wooden chair from the camera viewpoint. Exactly 4 major objects are visible in the composition. Spatial variant 14."
117
+ },
118
+ {
119
+ "category": "long_composition",
120
+ "prompt": "An elderly man stands at the left side of a glass table while a man wearing a gray coat sits to the right. The standing person holds a black vase in the left hand and points toward a sign reading exactly \"CAFE LEVEL 2\" with the right hand. A red cylinder and a blue book lie on the table in that order from left to right. The scene is illuminated by cool light entering from the right. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 6 major foreground objects should remain clearly distinguishable. Long-scene variant 18."
121
+ },
122
+ {
123
+ "category": "complex_anatomy",
124
+ "prompt": "A seated woman reaches forward with the right arm while keeping the left hand behind the back. An elderly woman stands beside them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 13."
125
+ },
126
+ {
127
+ "category": "long_composition",
128
+ "prompt": "A woman wearing a dark jacket stands at the left side of a glass table while an elderly woman sits to the right. The standing person holds a red cylinder in the left hand and points toward a sign reading exactly \"FINAL STOP\" with the right hand. A ceramic cup and a transparent cube lie on the table in that order from left to right. The scene is illuminated by a narrow spotlight from above. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 8 major foreground objects should remain clearly distinguishable. Long-scene variant 20."
129
+ },
130
+ {
131
+ "category": "complex_material_light",
132
+ "prompt": "A black vase made from glossy ceramic rests on a surface made from polished chrome, illuminated by hard side lighting. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 14."
133
+ },
134
+ {
135
+ "category": "complex_material_light",
136
+ "prompt": "A transparent cube made from brushed aluminum rests on a surface made from glossy ceramic, illuminated by strong backlight. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 3."
137
+ },
138
+ {
139
+ "category": "complex_counting",
140
+ "prompt": "Exactly 3 silver spheres form a curved row, while exactly 4 red cylinders occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 6."
141
+ },
142
+ {
143
+ "category": "architecture_vehicle",
144
+ "prompt": "A black motorcycle passes beside a small workshop. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 3."
145
+ },
146
+ {
147
+ "category": "complex_anatomy",
148
+ "prompt": "An elderly man reaches forward with the right arm while keeping the left hand behind the back. A seated woman stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 18."
149
+ },
150
+ {
151
+ "category": "complex_spatial",
152
+ "prompt": "A wooden chair stands behind a transparent cube, while a red cylinder is positioned to their right. The transparent cube partially overlaps the wooden chair from the camera viewpoint. Exactly 6 major objects are visible in the composition. Spatial variant 4."
153
+ },
154
+ {
155
+ "category": "complex_anatomy",
156
+ "prompt": "A man wearing a gray coat raises the left hand while holding a cup in the right hand. A young woman stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 6."
157
+ },
158
+ {
159
+ "category": "complex_material_light",
160
+ "prompt": "A metal box made from polished chrome rests on a surface made from wet black stone, illuminated by warm light entering from the left. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 9."
161
+ },
162
+ {
163
+ "category": "architecture_vehicle",
164
+ "prompt": "A white city bus passes beside a glass office lobby. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 2."
165
+ },
166
+ {
167
+ "category": "architecture_vehicle",
168
+ "prompt": "A black motorcycle passes beside a underground station. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 13."
169
+ },
170
+ {
171
+ "category": "reflection_occlusion_ood",
172
+ "prompt": "An elderly man stands behind a transparent glass panel. A silver sphere in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a glass bottle located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 16. Reflection variant 16."
173
+ },
174
+ {
175
+ "category": "long_composition",
176
+ "prompt": "A seated woman stands at the left side of a glass table while a young woman sits to the right. The standing person holds a blue book in the left hand and points toward a sign reading exactly \"OPEN 24 HOURS\" with the right hand. A transparent cube and a ceramic cup lie on the table in that order from left to right. The scene is illuminated by soft diffused window light. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 5 major foreground objects should remain clearly distinguishable. Long-scene variant 5."
177
+ },
178
+ {
179
+ "category": "long_composition",
180
+ "prompt": "A woman with short hair stands at the left side of a glass table while a tall man sits to the right. The standing person holds a transparent cube in the left hand and points toward a sign reading exactly \"RIVER HOTEL\" with the right hand. A small lamp and a wooden chair lie on the table in that order from left to right. The scene is illuminated by warm reflected light from below. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 7 major foreground objects should remain clearly distinguishable. Long-scene variant 7."
181
+ },
182
+ {
183
+ "category": "complex_material_light",
184
+ "prompt": "A glass bottle made from rough concrete rests on a surface made from frosted glass, illuminated by warm reflected light from below. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 7."
185
+ },
186
+ {
187
+ "category": "reflection_occlusion_ood",
188
+ "prompt": "A young woman stands behind a transparent glass panel. A transparent cube in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a ceramic cup located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 7. Reflection variant 7."
189
+ },
190
+ {
191
+ "category": "reflection_occlusion_ood",
192
+ "prompt": "A tall man stands behind a transparent glass panel. A glass bottle in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a silver sphere located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 1. Reflection variant 1."
193
+ },
194
+ {
195
+ "category": "reflection_occlusion_ood",
196
+ "prompt": "A woman with short hair stands behind a transparent glass panel. A metal box in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a black vase located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 13. Reflection variant 13."
197
+ },
198
+ {
199
+ "category": "complex_material_light",
200
+ "prompt": "A silver sphere made from wet black stone rests on a surface made from rough concrete, illuminated by a narrow spotlight from above. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 12."
201
+ },
202
+ {
203
+ "category": "complex_text",
204
+ "prompt": "Inside a minimalist living room, the exact phrase \"OPEN 24 HOURS\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 5."
205
+ },
206
+ {
207
+ "category": "complex_counting",
208
+ "prompt": "Exactly 5 metal boxs form two staggered rows, while exactly 2 transparent cubes occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 3."
209
+ },
210
+ {
211
+ "category": "complex_material_light",
212
+ "prompt": "A wooden chair made from frosted glass rests on a surface made from translucent amber acrylic, illuminated by cool light entering from the right. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 10."
213
+ },
214
+ {
215
+ "category": "complex_text",
216
+ "prompt": "Inside a underground station, the exact phrase \"ROOM 314\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 12."
217
+ },
218
+ {
219
+ "category": "complex_material_light",
220
+ "prompt": "A blue book made from brushed aluminum rests on a surface made from glossy ceramic, illuminated by strong backlight. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 11."
221
+ },
222
+ {
223
+ "category": "complex_anatomy",
224
+ "prompt": "A young woman crosses the right leg over the left while looking over the left shoulder. A woman wearing a dark jacket stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 17."
225
+ },
226
+ {
227
+ "category": "complex_counting",
228
+ "prompt": "Exactly 6 wooden chairs form a curved row, while exactly 4 black vases occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 2. Counting variant 14."
229
+ },
230
+ {
231
+ "category": "complex_text",
232
+ "prompt": "Inside a underground station, the exact phrase \"FINAL STOP\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 20."
233
+ },
234
+ {
235
+ "category": "long_composition",
236
+ "prompt": "An elderly man stands at the left side of a glass table while a man wearing a gray coat sits to the right. The standing person holds a ceramic cup in the left hand and points toward a sign reading exactly \"ROOM 314\" with the right hand. A wooden chair and a small lamp lie on the table in that order from left to right. The scene is illuminated by cool light entering from the right. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 6 major foreground objects should remain clearly distinguishable. Long-scene variant 2."
237
+ },
238
+ {
239
+ "category": "complex_counting",
240
+ "prompt": "Exactly 6 small lamps form two staggered rows, while exactly 2 metal boxs occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 9."
241
+ },
242
+ {
243
+ "category": "complex_spatial",
244
+ "prompt": "A glass bottle stands behind a wooden chair, while a transparent cube is positioned to their left. The wooden chair partially overlaps the glass bottle from the camera viewpoint. Exactly 3 major objects are visible in the composition. Spatial variant 1."
245
+ },
246
+ {
247
+ "category": "long_composition",
248
+ "prompt": "A woman wearing a dark jacket stands at the left side of a glass table while an elderly woman sits to the right. The standing person holds a wooden chair in the left hand and points toward a sign reading exactly \"WEST EXIT\" with the right hand. A silver sphere and a glass bottle lie on the table in that order from left to right. The scene is illuminated by a narrow spotlight from above. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 8 major foreground objects should remain clearly distinguishable. Long-scene variant 4."
249
+ },
250
+ {
251
+ "category": "complex_spatial",
252
+ "prompt": "A ceramic cup stands behind a blue book, while a black vase is positioned to their right. The blue book partially overlaps the ceramic cup from the camera viewpoint. Exactly 6 major objects are visible in the composition. Spatial variant 12."
253
+ },
254
+ {
255
+ "category": "complex_spatial",
256
+ "prompt": "A blue book stands behind a black vase, while a glass bottle is positioned to their left. The black vase partially overlaps the blue book from the camera viewpoint. Exactly 5 major objects are visible in the composition. Spatial variant 15."
257
+ },
258
+ {
259
+ "category": "complex_spatial",
260
+ "prompt": "A transparent cube stands behind a red cylinder, while a metal box is positioned to their left. The red cylinder partially overlaps the transparent cube from the camera viewpoint. Exactly 5 major objects are visible in the composition. Spatial variant 7."
261
+ },
262
+ {
263
+ "category": "long_composition",
264
+ "prompt": "A seated woman stands at the left side of a glass table while a young woman sits to the right. The standing person holds a metal box in the left hand and points toward a sign reading exactly \"PLATFORM 11\" with the right hand. A blue book and a red cylinder lie on the table in that order from left to right. The scene is illuminated by soft diffused window light. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 5 major foreground objects should remain clearly distinguishable. Long-scene variant 13."
265
+ },
266
+ {
267
+ "category": "reflection_occlusion_ood",
268
+ "prompt": "A woman with short hair stands behind a transparent glass panel. A blue book in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a red cylinder located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 5. Reflection variant 5."
269
+ },
270
+ {
271
+ "category": "reflection_occlusion_ood",
272
+ "prompt": "A man wearing a gray coat stands behind a transparent glass panel. A wooden chair in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a small lamp located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 4. Reflection variant 4."
273
+ },
274
+ {
275
+ "category": "complex_text",
276
+ "prompt": "Inside a modern kitchen, the exact phrase \"RIVER HOTEL\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 7."
277
+ },
278
+ {
279
+ "category": "complex_spatial",
280
+ "prompt": "A metal box stands behind a silver sphere, while a small lamp is positioned to their left. The silver sphere partially overlaps the metal box from the camera viewpoint. Exactly 3 major objects are visible in the composition. Spatial variant 13."
281
+ },
282
+ {
283
+ "category": "complex_material_light",
284
+ "prompt": "A glass bottle made from polished chrome rests on a surface made from wet black stone, illuminated by warm light entering from the left. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 17."
285
+ },
286
+ {
287
+ "category": "architecture_vehicle",
288
+ "prompt": "A blue bicycle passes beside a small workshop. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 19."
289
+ },
290
+ {
291
+ "category": "complex_spatial",
292
+ "prompt": "A black vase stands behind a glass bottle, while a wooden chair is positioned to their right. The glass bottle partially overlaps the black vase from the camera viewpoint. Exactly 6 major objects are visible in the composition. Spatial variant 8."
293
+ },
294
+ {
295
+ "category": "long_composition",
296
+ "prompt": "A woman wearing a dark jacket stands at the left side of a glass table while an elderly woman sits to the right. The standing person holds a ceramic cup in the left hand and points toward a sign reading exactly \"ROOM 314\" with the right hand. A wooden chair and a small lamp lie on the table in that order from left to right. The scene is illuminated by a narrow spotlight from above. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 8 major foreground objects should remain clearly distinguishable. Long-scene variant 12."
297
+ },
298
+ {
299
+ "category": "complex_counting",
300
+ "prompt": "Exactly 6 small lamps form two staggered rows, while exactly 2 metal boxs occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 1. Counting variant 19."
301
+ },
302
+ {
303
+ "category": "complex_anatomy",
304
+ "prompt": "An elderly man crosses the right leg over the left while looking over the left shoulder. A seated woman stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 2."
305
+ },
306
+ {
307
+ "category": "complex_anatomy",
308
+ "prompt": "A tall man reaches forward with the right arm while keeping the left hand behind the back. A man wearing a gray coat stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 3."
309
+ },
310
+ {
311
+ "category": "complex_counting",
312
+ "prompt": "Exactly 4 transparent cubes form two staggered rows, while exactly 2 glass bottles occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 1. Counting variant 17."
313
+ },
314
+ {
315
+ "category": "architecture_vehicle",
316
+ "prompt": "A red compact car passes beside a small workshop. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 11."
317
+ },
318
+ {
319
+ "category": "long_composition",
320
+ "prompt": "A man wearing a gray coat stands at the left side of a glass table while an elderly man sits to the right. The standing person holds a silver sphere in the left hand and points toward a sign reading exactly \"STUDIO C\" with the right hand. A black vase and a metal box lie on the table in that order from left to right. The scene is illuminated by hard side lighting. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 6 major foreground objects should remain clearly distinguishable. Long-scene variant 6."
321
+ },
322
+ {
323
+ "category": "complex_spatial",
324
+ "prompt": "A glass bottle stands behind a wooden chair, while a transparent cube is positioned to their left. The wooden chair partially overlaps the glass bottle from the camera viewpoint. Exactly 5 major objects are visible in the composition. Spatial variant 11."
325
+ },
326
+ {
327
+ "category": "complex_spatial",
328
+ "prompt": "A red cylinder stands behind a metal box, while a silver sphere is positioned to their right. The metal box partially overlaps the red cylinder from the camera viewpoint. Exactly 6 major objects are visible in the composition. Spatial variant 20."
329
+ },
330
+ {
331
+ "category": "long_composition",
332
+ "prompt": "A man wearing a gray coat stands at the left side of a glass table while an elderly man sits to the right. The standing person holds a wooden chair in the left hand and points toward a sign reading exactly \"WEST EXIT\" with the right hand. A silver sphere and a glass bottle lie on the table in that order from left to right. The scene is illuminated by hard side lighting. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 6 major foreground objects should remain clearly distinguishable. Long-scene variant 14."
333
+ },
334
+ {
335
+ "category": "long_composition",
336
+ "prompt": "An elderly woman stands at the left side of a glass table while a woman wearing a dark jacket sits to the right. The standing person holds a silver sphere in the left hand and points toward a sign reading exactly \"STUDIO C\" with the right hand. A black vase and a metal box lie on the table in that order from left to right. The scene is illuminated by two opposing light sources. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 8 major foreground objects should remain clearly distinguishable. Long-scene variant 16."
337
+ },
338
+ {
339
+ "category": "reflection_occlusion_ood",
340
+ "prompt": "A young woman stands behind a transparent glass panel. A blue book in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a red cylinder located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 15. Reflection variant 15."
341
+ },
342
+ {
343
+ "category": "complex_anatomy",
344
+ "prompt": "An elderly woman reaches forward with the right arm while keeping the left hand behind the back. A tall man stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 8."
345
+ },
346
+ {
347
+ "category": "architecture_vehicle",
348
+ "prompt": "A silver tram passes beside a glass office lobby. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 10."
349
+ },
350
+ {
351
+ "category": "complex_text",
352
+ "prompt": "Inside a industrial studio, the exact phrase \"WEST EXIT\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 14."
353
+ },
354
+ {
355
+ "category": "complex_counting",
356
+ "prompt": "Exactly 5 black vases form a curved row, while exactly 4 ceramic cups occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 18."
357
+ },
358
+ {
359
+ "category": "complex_material_light",
360
+ "prompt": "A silver sphere made from frosted glass rests on a surface made from translucent amber acrylic, illuminated by cool light entering from the right. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 2."
361
+ },
362
+ {
363
+ "category": "complex_material_light",
364
+ "prompt": "A small lamp made from translucent amber acrylic rests on a surface made from dark polished wood, illuminated by soft diffused window light. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 5."
365
+ },
366
+ {
367
+ "category": "reflection_occlusion_ood",
368
+ "prompt": "An elderly man stands behind a transparent glass panel. A black vase in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a metal box located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 8. Reflection variant 8."
369
+ },
370
+ {
371
+ "category": "complex_counting",
372
+ "prompt": "Exactly 5 metal boxs form two staggered rows, while exactly 2 transparent cubes occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 13."
373
+ },
374
+ {
375
+ "category": "complex_text",
376
+ "prompt": "Inside a modern kitchen, the exact phrase \"OPEN 24 HOURS\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 15."
377
+ },
378
+ {
379
+ "category": "reflection_occlusion_ood",
380
+ "prompt": "A seated woman stands behind a transparent glass panel. A glass bottle in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a silver sphere located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 11. Reflection variant 11."
381
+ },
382
+ {
383
+ "category": "reflection_occlusion_ood",
384
+ "prompt": "A man wearing a gray coat stands behind a transparent glass panel. A red cylinder in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a blue book located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 20. Reflection variant 20."
385
+ },
386
+ {
387
+ "category": "long_composition",
388
+ "prompt": "A young woman stands at the left side of a glass table while a seated woman sits to the right. The standing person holds a small lamp in the left hand and points toward a sign reading exactly \"AUTHORIZED ENTRY\" with the right hand. A glass bottle and a silver sphere lie on the table in that order from left to right. The scene is illuminated by warm light entering from the left. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 5 major foreground objects should remain clearly distinguishable. Long-scene variant 9."
389
+ },
390
+ {
391
+ "category": "architecture_vehicle",
392
+ "prompt": "A black motorcycle passes beside a glass office lobby. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 18."
393
+ },
394
+ {
395
+ "category": "complex_text",
396
+ "prompt": "Inside a small workshop, the exact phrase \"CAFE LEVEL 2\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 18."
397
+ },
398
+ {
399
+ "category": "complex_text",
400
+ "prompt": "Inside a glass office lobby, the exact phrase \"RIVER HOTEL\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 17."
401
+ },
402
+ {
403
+ "category": "complex_counting",
404
+ "prompt": "Exactly 5 black vases form a curved row, while exactly 4 ceramic cups occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 8."
405
+ },
406
+ {
407
+ "category": "complex_anatomy",
408
+ "prompt": "A tall man raises the left hand while holding a cup in the right hand. A man wearing a gray coat stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 11."
409
+ },
410
+ {
411
+ "category": "reflection_occlusion_ood",
412
+ "prompt": "An elderly woman stands behind a transparent glass panel. A wooden chair in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a small lamp located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 14. Reflection variant 14."
413
+ },
414
+ {
415
+ "category": "complex_text",
416
+ "prompt": "Inside a railway platform, the exact phrase \"STUDIO C\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 16."
417
+ },
418
+ {
419
+ "category": "architecture_vehicle",
420
+ "prompt": "A silver tram passes beside a underground station. A staircase rises on the left side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 5."
421
+ },
422
+ {
423
+ "category": "complex_anatomy",
424
+ "prompt": "A woman wearing a dark jacket holds a bottle with both hands directly in front of the chest. A woman with short hair stands behind them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 4."
425
+ },
426
+ {
427
+ "category": "architecture_vehicle",
428
+ "prompt": "A black motorcycle passes beside a modern kitchen. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 8."
429
+ },
430
+ {
431
+ "category": "complex_text",
432
+ "prompt": "Inside a glass office lobby, the exact phrase \"NORTH GATE\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 1."
433
+ },
434
+ {
435
+ "category": "complex_text",
436
+ "prompt": "Inside a hotel entrance, the exact phrase \"AUTHORIZED ENTRY\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 19."
437
+ },
438
+ {
439
+ "category": "reflection_occlusion_ood",
440
+ "prompt": "A seated woman stands behind a transparent glass panel. A small lamp in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a wooden chair located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 19. Reflection variant 19."
441
+ },
442
+ {
443
+ "category": "long_composition",
444
+ "prompt": "A young woman stands at the left side of a glass table while a seated woman sits to the right. The standing person holds a glass bottle in the left hand and points toward a sign reading exactly \"NORTH GATE\" with the right hand. A metal box and a black vase lie on the table in that order from left to right. The scene is illuminated by warm light entering from the left. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 5 major foreground objects should remain clearly distinguishable. Long-scene variant 1."
445
+ },
446
+ {
447
+ "category": "complex_material_light",
448
+ "prompt": "A red cylinder made from dark polished wood rests on a surface made from brushed aluminum, illuminated by two opposing light sources. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 16."
449
+ },
450
+ {
451
+ "category": "complex_anatomy",
452
+ "prompt": "A woman with short hair crosses the right leg over the left while looking over the left shoulder. An elderly man stands beside them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 7."
453
+ },
454
+ {
455
+ "category": "reflection_occlusion_ood",
456
+ "prompt": "A tall man stands behind a transparent glass panel. A transparent cube in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a ceramic cup located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 17. Reflection variant 17."
457
+ },
458
+ {
459
+ "category": "complex_counting",
460
+ "prompt": "Exactly 7 blue books form two staggered rows, while exactly 2 small lamps occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 1. Counting variant 15."
461
+ },
462
+ {
463
+ "category": "complex_counting",
464
+ "prompt": "Exactly 7 red cylinders form a curved row, while exactly 4 wooden chairs occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 6. Counting variant 20."
465
+ },
466
+ {
467
+ "category": "architecture_vehicle",
468
+ "prompt": "A white city bus passes beside a hotel entrance. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 12."
469
+ },
470
+ {
471
+ "category": "long_composition",
472
+ "prompt": "A woman with short hair stands at the left side of a glass table while a tall man sits to the right. The standing person holds a blue book in the left hand and points toward a sign reading exactly \"OPEN 24 HOURS\" with the right hand. A transparent cube and a ceramic cup lie on the table in that order from left to right. The scene is illuminated by warm reflected light from below. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 7 major foreground objects should remain clearly distinguishable. Long-scene variant 15."
473
+ },
474
+ {
475
+ "category": "complex_material_light",
476
+ "prompt": "A transparent cube made from translucent amber acrylic rests on a surface made from dark polished wood, illuminated by soft diffused window light. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 13."
477
+ },
478
+ {
479
+ "category": "complex_anatomy",
480
+ "prompt": "A seated woman turns the torso right while the head remains facing left. An elderly woman stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 5."
481
+ },
482
+ {
483
+ "category": "reflection_occlusion_ood",
484
+ "prompt": "A woman wearing a dark jacket stands behind a transparent glass panel. A black vase in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a metal box located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 18. Reflection variant 18."
485
+ },
486
+ {
487
+ "category": "complex_spatial",
488
+ "prompt": "A metal box stands behind a silver sphere, while a small lamp is positioned to their left. The silver sphere partially overlaps the metal box from the camera viewpoint. Exactly 5 major objects are visible in the composition. Spatial variant 3."
489
+ },
490
+ {
491
+ "category": "complex_text",
492
+ "prompt": "Inside a glass office lobby, the exact phrase \"AUTHORIZED ENTRY\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 9."
493
+ },
494
+ {
495
+ "category": "complex_anatomy",
496
+ "prompt": "A young woman holds a bottle with both hands directly in front of the chest. A woman wearing a dark jacket stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 9."
497
+ },
498
+ {
499
+ "category": "complex_material_light",
500
+ "prompt": "A blue book made from polished chrome rests on a surface made from wet black stone, illuminated by warm light entering from the left. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 1."
501
+ },
502
+ {
503
+ "category": "long_composition",
504
+ "prompt": "An elderly woman stands at the left side of a glass table while a woman wearing a dark jacket sits to the right. The standing person holds a black vase in the left hand and points toward a sign reading exactly \"CAFE LEVEL 2\" with the right hand. A red cylinder and a blue book lie on the table in that order from left to right. The scene is illuminated by two opposing light sources. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 8 major foreground objects should remain clearly distinguishable. Long-scene variant 8."
505
+ },
506
+ {
507
+ "category": "complex_spatial",
508
+ "prompt": "A transparent cube stands behind a red cylinder, while a metal box is positioned to their left. The red cylinder partially overlaps the transparent cube from the camera viewpoint. Exactly 3 major objects are visible in the composition. Spatial variant 17."
509
+ },
510
+ {
511
+ "category": "complex_material_light",
512
+ "prompt": "A red cylinder made from glossy ceramic rests on a surface made from polished chrome, illuminated by hard side lighting. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 6."
513
+ },
514
+ {
515
+ "category": "complex_counting",
516
+ "prompt": "Exactly 4 ceramic cups form a curved row, while exactly 4 silver spheres occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 2. Counting variant 2."
517
+ },
518
+ {
519
+ "category": "long_composition",
520
+ "prompt": "A tall man stands at the left side of a glass table while a woman with short hair sits to the right. The standing person holds a glass bottle in the left hand and points toward a sign reading exactly \"NORTH GATE\" with the right hand. A metal box and a black vase lie on the table in that order from left to right. The scene is illuminated by strong backlight. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 7 major foreground objects should remain clearly distinguishable. Long-scene variant 11."
521
+ },
522
+ {
523
+ "category": "reflection_occlusion_ood",
524
+ "prompt": "A seated woman stands behind a transparent glass panel. A metal box in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a black vase located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 3. Reflection variant 3."
525
+ },
526
+ {
527
+ "category": "complex_spatial",
528
+ "prompt": "A red cylinder stands behind a metal box, while a silver sphere is positioned to their right. The metal box partially overlaps the red cylinder from the camera viewpoint. Exactly 4 major objects are visible in the composition. Spatial variant 10."
529
+ },
530
+ {
531
+ "category": "complex_counting",
532
+ "prompt": "Exactly 4 transparent cubes form two staggered rows, while exactly 2 glass bottles occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 7."
533
+ },
534
+ {
535
+ "category": "complex_anatomy",
536
+ "prompt": "An elderly woman raises the left hand while holding a cup in the right hand. A tall man stands behind them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 16."
537
+ },
538
+ {
539
+ "category": "reflection_occlusion_ood",
540
+ "prompt": "A woman wearing a dark jacket stands behind a transparent glass panel. A red cylinder in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a blue book located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 10. Reflection variant 10."
541
+ },
542
+ {
543
+ "category": "long_composition",
544
+ "prompt": "A young woman stands at the left side of a glass table while a seated woman sits to the right. The standing person holds a transparent cube in the left hand and points toward a sign reading exactly \"RIVER HOTEL\" with the right hand. A small lamp and a wooden chair lie on the table in that order from left to right. The scene is illuminated by warm light entering from the left. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 5 major foreground objects should remain clearly distinguishable. Long-scene variant 17."
545
+ },
546
+ {
547
+ "category": "complex_material_light",
548
+ "prompt": "A black vase made from wet black stone rests on a surface made from rough concrete, illuminated by a narrow spotlight from above. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 4."
549
+ },
550
+ {
551
+ "category": "complex_text",
552
+ "prompt": "Inside a railway platform, the exact phrase \"CAFE LEVEL 2\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 8."
553
+ },
554
+ {
555
+ "category": "complex_anatomy",
556
+ "prompt": "An elderly man turns the torso right while the head remains facing left. A seated woman stands behind them and points toward the object with the right hand. Both people's hands and feet remain visible. Anatomy variant 10."
557
+ },
558
+ {
559
+ "category": "complex_spatial",
560
+ "prompt": "A black vase stands behind a glass bottle, while a wooden chair is positioned to their right. The glass bottle partially overlaps the black vase from the camera viewpoint. Exactly 4 major objects are visible in the composition. Spatial variant 18."
561
+ },
562
+ {
563
+ "category": "complex_text",
564
+ "prompt": "Inside a hotel entrance, the exact phrase \"NORTH GATE\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 11."
565
+ },
566
+ {
567
+ "category": "architecture_vehicle",
568
+ "prompt": "A blue bicycle passes beside a hotel entrance. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 4."
569
+ },
570
+ {
571
+ "category": "long_composition",
572
+ "prompt": "A tall man stands at the left side of a glass table while a woman with short hair sits to the right. The standing person holds a small lamp in the left hand and points toward a sign reading exactly \"AUTHORIZED ENTRY\" with the right hand. A glass bottle and a silver sphere lie on the table in that order from left to right. The scene is illuminated by strong backlight. A large mirror behind both people reflects the seated person's back and one object outside the direct camera frame. Exactly 7 major foreground objects should remain clearly distinguishable. Long-scene variant 19."
573
+ },
574
+ {
575
+ "category": "complex_spatial",
576
+ "prompt": "A blue book stands behind a black vase, while a glass bottle is positioned to their left. The black vase partially overlaps the blue book from the camera viewpoint. Exactly 3 major objects are visible in the composition. Spatial variant 5."
577
+ },
578
+ {
579
+ "category": "architecture_vehicle",
580
+ "prompt": "A red compact car passes beside a minimalist living room. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 6."
581
+ },
582
+ {
583
+ "category": "complex_material_light",
584
+ "prompt": "A metal box made from brushed aluminum rests on a surface made from glossy ceramic, illuminated by strong backlight. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 19."
585
+ },
586
+ {
587
+ "category": "complex_material_light",
588
+ "prompt": "A small lamp made from rough concrete rests on a surface made from frosted glass, illuminated by warm reflected light from below. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 15."
589
+ },
590
+ {
591
+ "category": "complex_anatomy",
592
+ "prompt": "A woman wearing a dark jacket crosses the right leg over the left while looking over the left shoulder. A woman with short hair stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 12."
593
+ },
594
+ {
595
+ "category": "complex_counting",
596
+ "prompt": "Exactly 4 ceramic cups form a curved row, while exactly 4 silver spheres occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 4. Counting variant 12."
597
+ },
598
+ {
599
+ "category": "complex_counting",
600
+ "prompt": "Exactly 7 red cylinders form a curved row, while exactly 4 wooden chairs occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 3. Counting variant 10."
601
+ },
602
+ {
603
+ "category": "complex_anatomy",
604
+ "prompt": "A woman with short hair turns the torso right while the head remains facing left. An elderly man stands beside them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 15."
605
+ },
606
+ {
607
+ "category": "reflection_occlusion_ood",
608
+ "prompt": "An elderly woman stands behind a transparent glass panel. A silver sphere in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a glass bottle located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 6. Reflection variant 6."
609
+ },
610
+ {
611
+ "category": "architecture_vehicle",
612
+ "prompt": "A blue bicycle passes beside a minimalist living room. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 14."
613
+ },
614
+ {
615
+ "category": "complex_text",
616
+ "prompt": "Inside a hotel entrance, the exact phrase \"PLATFORM 11\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 3."
617
+ },
618
+ {
619
+ "category": "reflection_occlusion_ood",
620
+ "prompt": "A tall man stands behind a transparent glass panel. A small lamp in the foreground partially occludes the torso while the face and both hands remain visible. A mirror behind the person reflects a wooden chair located outside the direct camera frame. The glass also contains a faint reflection of an overhead light at position 9. Reflection variant 9."
621
+ },
622
+ {
623
+ "category": "complex_counting",
624
+ "prompt": "Exactly 6 wooden chairs form a curved row, while exactly 4 black vases occupy the background. None of the objects is completely hidden, and one small yellow marker is located beside object number 4. Counting variant 4."
625
+ },
626
+ {
627
+ "category": "complex_material_light",
628
+ "prompt": "A wooden chair made from wet black stone rests on a surface made from rough concrete, illuminated by a narrow spotlight from above. The scene clearly shows reflection, roughness, transparency, or subsurface behavior appropriate to both materials. A small colored object is reflected near the edge of the surface. Material-light variant 20."
629
+ },
630
+ {
631
+ "category": "complex_anatomy",
632
+ "prompt": "A woman wearing a dark jacket turns the torso right while the head remains facing left. A woman with short hair stands behind them and points toward the object with the left hand. Both people's hands and feet remain visible. Anatomy variant 20."
633
+ },
634
+ {
635
+ "category": "complex_text",
636
+ "prompt": "Inside a small workshop, the exact phrase \"FINAL STOP\" is printed clearly on a rectangular sign. A person stands partly in front of the sign without covering any letters. A reflective surface beside the sign shows a reversed reflection of part of the scene while the original text remains completely readable. Scene variant 10."
637
+ },
638
+ {
639
+ "category": "architecture_vehicle",
640
+ "prompt": "A silver tram passes beside a hotel entrance. A staircase rises on the right side and turns at an upper landing. Three architectural openings are visible at different depths, and a glass facade reflects a building across the street. Composition variant 20."
641
+ }
642
+ ]
research/datasets/bridge_prompts_480.json ADDED
@@ -0,0 +1,1762 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "category": "lighting",
4
+ "prompt": "a cup illuminated by strong backlight, with clearly visible light direction and shadow"
5
+ },
6
+ {
7
+ "category": "anatomy",
8
+ "prompt": "an elderly man reaching forward with their right arm, full body clearly visible"
9
+ },
10
+ {
11
+ "category": "reflection_occlusion",
12
+ "prompt": "one person partially occluded by another person"
13
+ },
14
+ {
15
+ "category": "anatomy",
16
+ "prompt": "a woman turning their head to the left, full body clearly visible"
17
+ },
18
+ {
19
+ "category": "reflection_occlusion",
20
+ "prompt": "a woman reflected accurately in a wall mirror"
21
+ },
22
+ {
23
+ "category": "spatial",
24
+ "prompt": "a white book positioned below a black bottle"
25
+ },
26
+ {
27
+ "category": "lighting",
28
+ "prompt": "a glass illuminated by strong backlight, with clearly visible light direction and shadow"
29
+ },
30
+ {
31
+ "category": "spatial",
32
+ "prompt": "a red book positioned below a yellow vase"
33
+ },
34
+ {
35
+ "category": "reflection_occlusion",
36
+ "prompt": "a translucent curtain illuminated from behind"
37
+ },
38
+ {
39
+ "category": "counting",
40
+ "prompt": "2 wooden blocks arranged in the foreground with 3 bottles behind them"
41
+ },
42
+ {
43
+ "category": "spatial",
44
+ "prompt": "a green book positioned below a red glass"
45
+ },
46
+ {
47
+ "category": "spatial",
48
+ "prompt": "a green chair positioned beside a orange sphere"
49
+ },
50
+ {
51
+ "category": "anatomy",
52
+ "prompt": "an elderly man standing with one arm behind their back, full body clearly visible"
53
+ },
54
+ {
55
+ "category": "counting",
56
+ "prompt": "5 lamps arranged in the foreground with 3 books behind them"
57
+ },
58
+ {
59
+ "category": "counting",
60
+ "prompt": "4 tables arranged in the foreground with 4 mirrors behind them"
61
+ },
62
+ {
63
+ "category": "spatial",
64
+ "prompt": "a red cube positioned partially behind a white metal cylinder"
65
+ },
66
+ {
67
+ "category": "lighting",
68
+ "prompt": "a vase illuminated by warm light reflected from below, with clearly visible light direction and shadow"
69
+ },
70
+ {
71
+ "category": "lighting",
72
+ "prompt": "a bottle illuminated by strong backlight, with clearly visible light direction and shadow"
73
+ },
74
+ {
75
+ "category": "anatomy",
76
+ "prompt": "a man standing with one arm behind their back, full body clearly visible"
77
+ },
78
+ {
79
+ "category": "reflection_occlusion",
80
+ "prompt": "one person partially occluded by another person"
81
+ },
82
+ {
83
+ "category": "anatomy",
84
+ "prompt": "an elderly woman touching their face with their left hand, full body clearly visible"
85
+ },
86
+ {
87
+ "category": "text",
88
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a white rectangular sign"
89
+ },
90
+ {
91
+ "category": "lighting",
92
+ "prompt": "a mirror illuminated by soft light from the left, with clearly visible light direction and shadow"
93
+ },
94
+ {
95
+ "category": "spatial",
96
+ "prompt": "a red vase positioned partially behind a golden lamp"
97
+ },
98
+ {
99
+ "category": "material",
100
+ "prompt": "a table made of polished chrome, showing realistic surface properties and reflections"
101
+ },
102
+ {
103
+ "category": "text",
104
+ "prompt": "the word \"EXIT\" printed clearly and correctly on a metal street sign"
105
+ },
106
+ {
107
+ "category": "reflection_occlusion",
108
+ "prompt": "one person partially occluded by another person"
109
+ },
110
+ {
111
+ "category": "spatial",
112
+ "prompt": "a white wooden block positioned to the right of a purple mirror"
113
+ },
114
+ {
115
+ "category": "counting",
116
+ "prompt": "5 cubes arranged in the foreground with 4 plates behind them"
117
+ },
118
+ {
119
+ "category": "lighting",
120
+ "prompt": "a metal cylinder illuminated by warm overhead light, with clearly visible light direction and shadow"
121
+ },
122
+ {
123
+ "category": "material",
124
+ "prompt": "a glass made of wet concrete, showing realistic surface properties and reflections"
125
+ },
126
+ {
127
+ "category": "lighting",
128
+ "prompt": "a box illuminated by cool side lighting, with clearly visible light direction and shadow"
129
+ },
130
+ {
131
+ "category": "spatial",
132
+ "prompt": "a purple metal cylinder positioned beside a red wooden block"
133
+ },
134
+ {
135
+ "category": "material",
136
+ "prompt": "a chair made of polished chrome, showing realistic surface properties and reflections"
137
+ },
138
+ {
139
+ "category": "lighting",
140
+ "prompt": "a glass illuminated by soft diffused daylight, with clearly visible light direction and shadow"
141
+ },
142
+ {
143
+ "category": "anatomy",
144
+ "prompt": "an elderly man sitting with both feet visible, full body clearly visible"
145
+ },
146
+ {
147
+ "category": "lighting",
148
+ "prompt": "a sphere illuminated by warm light reflected from below, with clearly visible light direction and shadow"
149
+ },
150
+ {
151
+ "category": "lighting",
152
+ "prompt": "a cube illuminated by soft diffused daylight, with clearly visible light direction and shadow"
153
+ },
154
+ {
155
+ "category": "text",
156
+ "prompt": "the word \"STOP\" printed clearly and correctly on a white rectangular sign"
157
+ },
158
+ {
159
+ "category": "counting",
160
+ "prompt": "4 metal cylinders arranged in the foreground with 4 cups behind them"
161
+ },
162
+ {
163
+ "category": "text",
164
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a paper label"
165
+ },
166
+ {
167
+ "category": "lighting",
168
+ "prompt": "a metal cylinder illuminated by warm light reflected from below, with clearly visible light direction and shadow"
169
+ },
170
+ {
171
+ "category": "anatomy",
172
+ "prompt": "an elderly man holding an object between both hands, full body clearly visible"
173
+ },
174
+ {
175
+ "category": "lighting",
176
+ "prompt": "a box illuminated by cool side lighting, with clearly visible light direction and shadow"
177
+ },
178
+ {
179
+ "category": "text",
180
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a white rectangular sign"
181
+ },
182
+ {
183
+ "category": "spatial",
184
+ "prompt": "a purple cup positioned to the left of a silver glass"
185
+ },
186
+ {
187
+ "category": "lighting",
188
+ "prompt": "a box illuminated by hard light from the right, with clearly visible light direction and shadow"
189
+ },
190
+ {
191
+ "category": "text",
192
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a black poster"
193
+ },
194
+ {
195
+ "category": "reflection_occlusion",
196
+ "prompt": "a mirror showing the back of a person facing away"
197
+ },
198
+ {
199
+ "category": "text",
200
+ "prompt": "the word \"EXIT\" printed clearly and correctly on a metal street sign"
201
+ },
202
+ {
203
+ "category": "text",
204
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a glass storefront"
205
+ },
206
+ {
207
+ "category": "anatomy",
208
+ "prompt": "an elderly woman holding an object between both hands, full body clearly visible"
209
+ },
210
+ {
211
+ "category": "anatomy",
212
+ "prompt": "an elderly woman holding both hands in front of their chest, full body clearly visible"
213
+ },
214
+ {
215
+ "category": "reflection_occlusion",
216
+ "prompt": "a transparent object casting a faint shadow"
217
+ },
218
+ {
219
+ "category": "anatomy",
220
+ "prompt": "an elderly man holding both hands in front of their chest, full body clearly visible"
221
+ },
222
+ {
223
+ "category": "counting",
224
+ "prompt": "2 boxs arranged in the foreground with 4 spheres behind them"
225
+ },
226
+ {
227
+ "category": "spatial",
228
+ "prompt": "a white mirror positioned above a golden vase"
229
+ },
230
+ {
231
+ "category": "material",
232
+ "prompt": "a chair made of glossy ceramic, showing realistic surface properties and reflections"
233
+ },
234
+ {
235
+ "category": "text",
236
+ "prompt": "the word \"EXIT\" printed clearly and correctly on a black poster"
237
+ },
238
+ {
239
+ "category": "anatomy",
240
+ "prompt": "an elderly woman raising their left hand above their head, full body clearly visible"
241
+ },
242
+ {
243
+ "category": "counting",
244
+ "prompt": "5 glasss arranged in the foreground with 3 wooden blocks behind them"
245
+ },
246
+ {
247
+ "category": "anatomy",
248
+ "prompt": "a young woman holding an object between both hands, full body clearly visible"
249
+ },
250
+ {
251
+ "category": "counting",
252
+ "prompt": "5 chairs arranged in the foreground with 4 wooden blocks behind them"
253
+ },
254
+ {
255
+ "category": "lighting",
256
+ "prompt": "a wooden block illuminated by warm light reflected from below, with clearly visible light direction and shadow"
257
+ },
258
+ {
259
+ "category": "lighting",
260
+ "prompt": "a sphere illuminated by hard light from the right, with clearly visible light direction and shadow"
261
+ },
262
+ {
263
+ "category": "lighting",
264
+ "prompt": "a vase illuminated by soft light from the left, with clearly visible light direction and shadow"
265
+ },
266
+ {
267
+ "category": "spatial",
268
+ "prompt": "a black wooden block positioned behind a white glass"
269
+ },
270
+ {
271
+ "category": "material",
272
+ "prompt": "a cup made of rough stone, showing realistic surface properties and reflections"
273
+ },
274
+ {
275
+ "category": "spatial",
276
+ "prompt": "a golden mirror positioned to the right of a red metal cylinder"
277
+ },
278
+ {
279
+ "category": "spatial",
280
+ "prompt": "a purple cup positioned to the left of a blue cube"
281
+ },
282
+ {
283
+ "category": "lighting",
284
+ "prompt": "a sphere illuminated by soft diffused daylight, with clearly visible light direction and shadow"
285
+ },
286
+ {
287
+ "category": "material",
288
+ "prompt": "a box made of translucent acrylic, showing realistic surface properties and reflections"
289
+ },
290
+ {
291
+ "category": "text",
292
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a white rectangular sign"
293
+ },
294
+ {
295
+ "category": "lighting",
296
+ "prompt": "a vase illuminated by soft diffused daylight, with clearly visible light direction and shadow"
297
+ },
298
+ {
299
+ "category": "spatial",
300
+ "prompt": "a purple bottle positioned above a green box"
301
+ },
302
+ {
303
+ "category": "lighting",
304
+ "prompt": "a vase illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
305
+ },
306
+ {
307
+ "category": "text",
308
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a glass storefront"
309
+ },
310
+ {
311
+ "category": "anatomy",
312
+ "prompt": "a woman holding both hands in front of their chest, full body clearly visible"
313
+ },
314
+ {
315
+ "category": "material",
316
+ "prompt": "a box made of dark wood, showing realistic surface properties and reflections"
317
+ },
318
+ {
319
+ "category": "text",
320
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a black poster"
321
+ },
322
+ {
323
+ "category": "material",
324
+ "prompt": "a vase made of dark wood, showing realistic surface properties and reflections"
325
+ },
326
+ {
327
+ "category": "lighting",
328
+ "prompt": "a chair illuminated by strong backlight, with clearly visible light direction and shadow"
329
+ },
330
+ {
331
+ "category": "lighting",
332
+ "prompt": "a lamp illuminated by warm overhead light, with clearly visible light direction and shadow"
333
+ },
334
+ {
335
+ "category": "reflection_occlusion",
336
+ "prompt": "a glass sphere resting on a reflective metal surface"
337
+ },
338
+ {
339
+ "category": "lighting",
340
+ "prompt": "a mirror illuminated by hard light from the right, with clearly visible light direction and shadow"
341
+ },
342
+ {
343
+ "category": "spatial",
344
+ "prompt": "a black wooden block positioned to the right of a silver vase"
345
+ },
346
+ {
347
+ "category": "lighting",
348
+ "prompt": "a lamp illuminated by cool side lighting, with clearly visible light direction and shadow"
349
+ },
350
+ {
351
+ "category": "reflection_occlusion",
352
+ "prompt": "a mirror showing the back of a person facing away"
353
+ },
354
+ {
355
+ "category": "material",
356
+ "prompt": "a vase made of rough stone, showing realistic surface properties and reflections"
357
+ },
358
+ {
359
+ "category": "spatial",
360
+ "prompt": "a silver metal cylinder positioned above a green sphere"
361
+ },
362
+ {
363
+ "category": "lighting",
364
+ "prompt": "a cup illuminated by soft diffused daylight, with clearly visible light direction and shadow"
365
+ },
366
+ {
367
+ "category": "spatial",
368
+ "prompt": "a yellow bottle positioned below a red lamp"
369
+ },
370
+ {
371
+ "category": "reflection_occlusion",
372
+ "prompt": "a glass sphere resting on a reflective metal surface"
373
+ },
374
+ {
375
+ "category": "anatomy",
376
+ "prompt": "a woman reaching forward with their right arm, full body clearly visible"
377
+ },
378
+ {
379
+ "category": "spatial",
380
+ "prompt": "a green vase positioned behind a blue lamp"
381
+ },
382
+ {
383
+ "category": "reflection_occlusion",
384
+ "prompt": "a shiny metal object reflecting a nearby red object"
385
+ },
386
+ {
387
+ "category": "spatial",
388
+ "prompt": "a white sphere positioned to the left of a yellow metal cylinder"
389
+ },
390
+ {
391
+ "category": "spatial",
392
+ "prompt": "a golden book positioned below a yellow plate"
393
+ },
394
+ {
395
+ "category": "counting",
396
+ "prompt": "2 vases arranged in the foreground with 4 wooden blocks behind them"
397
+ },
398
+ {
399
+ "category": "spatial",
400
+ "prompt": "a white book positioned above a purple table"
401
+ },
402
+ {
403
+ "category": "anatomy",
404
+ "prompt": "a young man reaching forward with their right arm, full body clearly visible"
405
+ },
406
+ {
407
+ "category": "text",
408
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a paper label"
409
+ },
410
+ {
411
+ "category": "material",
412
+ "prompt": "a sphere made of translucent acrylic, showing realistic surface properties and reflections"
413
+ },
414
+ {
415
+ "category": "lighting",
416
+ "prompt": "a bottle illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
417
+ },
418
+ {
419
+ "category": "lighting",
420
+ "prompt": "a metal cylinder illuminated by strong backlight, with clearly visible light direction and shadow"
421
+ },
422
+ {
423
+ "category": "material",
424
+ "prompt": "a bottle made of glossy ceramic, showing realistic surface properties and reflections"
425
+ },
426
+ {
427
+ "category": "material",
428
+ "prompt": "a bottle made of transparent glass, showing realistic surface properties and reflections"
429
+ },
430
+ {
431
+ "category": "spatial",
432
+ "prompt": "a black metal cylinder positioned beside a orange chair"
433
+ },
434
+ {
435
+ "category": "counting",
436
+ "prompt": "2 bottles arranged in the foreground with 2 mirrors behind them"
437
+ },
438
+ {
439
+ "category": "material",
440
+ "prompt": "a lamp made of wet concrete, showing realistic surface properties and reflections"
441
+ },
442
+ {
443
+ "category": "text",
444
+ "prompt": "the word \"CAFE\" printed clearly and correctly on a black poster"
445
+ },
446
+ {
447
+ "category": "anatomy",
448
+ "prompt": "an elderly man holding both hands in front of their chest, full body clearly visible"
449
+ },
450
+ {
451
+ "category": "reflection_occlusion",
452
+ "prompt": "a shiny metal object reflecting a nearby red object"
453
+ },
454
+ {
455
+ "category": "reflection_occlusion",
456
+ "prompt": "a transparent glass bottle in front of a person's face"
457
+ },
458
+ {
459
+ "category": "spatial",
460
+ "prompt": "a orange vase positioned to the right of a green table"
461
+ },
462
+ {
463
+ "category": "counting",
464
+ "prompt": "3 mirrors arranged in the foreground with 2 glasss behind them"
465
+ },
466
+ {
467
+ "category": "counting",
468
+ "prompt": "3 boxs arranged in the foreground with 4 cubes behind them"
469
+ },
470
+ {
471
+ "category": "lighting",
472
+ "prompt": "a plate illuminated by cool side lighting, with clearly visible light direction and shadow"
473
+ },
474
+ {
475
+ "category": "lighting",
476
+ "prompt": "a cube illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
477
+ },
478
+ {
479
+ "category": "counting",
480
+ "prompt": "4 mirrors arranged in the foreground with 1 cubes behind them"
481
+ },
482
+ {
483
+ "category": "anatomy",
484
+ "prompt": "an elderly man raising their left hand above their head, full body clearly visible"
485
+ },
486
+ {
487
+ "category": "anatomy",
488
+ "prompt": "a woman crossing their right leg over their left leg, full body clearly visible"
489
+ },
490
+ {
491
+ "category": "material",
492
+ "prompt": "a plate made of matte plastic, showing realistic surface properties and reflections"
493
+ },
494
+ {
495
+ "category": "text",
496
+ "prompt": "the word \"HOTEL\" printed clearly and correctly on a white rectangular sign"
497
+ },
498
+ {
499
+ "category": "lighting",
500
+ "prompt": "a box illuminated by hard light from the right, with clearly visible light direction and shadow"
501
+ },
502
+ {
503
+ "category": "spatial",
504
+ "prompt": "a red bottle positioned behind a green mirror"
505
+ },
506
+ {
507
+ "category": "spatial",
508
+ "prompt": "a black book positioned partially behind a red bottle"
509
+ },
510
+ {
511
+ "category": "spatial",
512
+ "prompt": "a white metal cylinder positioned partially behind a black plate"
513
+ },
514
+ {
515
+ "category": "counting",
516
+ "prompt": "4 boxs arranged in the foreground with 2 lamps behind them"
517
+ },
518
+ {
519
+ "category": "anatomy",
520
+ "prompt": "a young man touching their face with their left hand, full body clearly visible"
521
+ },
522
+ {
523
+ "category": "counting",
524
+ "prompt": "2 chairs arranged in the foreground with 1 plates behind them"
525
+ },
526
+ {
527
+ "category": "reflection_occlusion",
528
+ "prompt": "a hand visible through a transparent glass panel"
529
+ },
530
+ {
531
+ "category": "lighting",
532
+ "prompt": "a mirror illuminated by soft light from the left, with clearly visible light direction and shadow"
533
+ },
534
+ {
535
+ "category": "spatial",
536
+ "prompt": "a blue box positioned above a black book"
537
+ },
538
+ {
539
+ "category": "text",
540
+ "prompt": "the word \"BLUE\" printed clearly and correctly on a metal street sign"
541
+ },
542
+ {
543
+ "category": "counting",
544
+ "prompt": "3 cups arranged in the foreground with 4 mirrors behind them"
545
+ },
546
+ {
547
+ "category": "lighting",
548
+ "prompt": "a cup illuminated by cool side lighting, with clearly visible light direction and shadow"
549
+ },
550
+ {
551
+ "category": "text",
552
+ "prompt": "the word \"NORTH\" printed clearly and correctly on a white rectangular sign"
553
+ },
554
+ {
555
+ "category": "material",
556
+ "prompt": "a wooden block made of wet concrete, showing realistic surface properties and reflections"
557
+ },
558
+ {
559
+ "category": "spatial",
560
+ "prompt": "a purple bottle positioned to the right of a green lamp"
561
+ },
562
+ {
563
+ "category": "text",
564
+ "prompt": "the word \"CAFE\" printed clearly and correctly on a paper label"
565
+ },
566
+ {
567
+ "category": "lighting",
568
+ "prompt": "a vase illuminated by strong backlight, with clearly visible light direction and shadow"
569
+ },
570
+ {
571
+ "category": "lighting",
572
+ "prompt": "a glass illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
573
+ },
574
+ {
575
+ "category": "lighting",
576
+ "prompt": "a metal cylinder illuminated by cool side lighting, with clearly visible light direction and shadow"
577
+ },
578
+ {
579
+ "category": "anatomy",
580
+ "prompt": "an elderly woman holding both hands in front of their chest, full body clearly visible"
581
+ },
582
+ {
583
+ "category": "spatial",
584
+ "prompt": "a black sphere positioned beside a blue vase"
585
+ },
586
+ {
587
+ "category": "anatomy",
588
+ "prompt": "a man holding a bottle in their right hand, full body clearly visible"
589
+ },
590
+ {
591
+ "category": "material",
592
+ "prompt": "a metal cylinder made of glossy ceramic, showing realistic surface properties and reflections"
593
+ },
594
+ {
595
+ "category": "anatomy",
596
+ "prompt": "a young man touching their face with their left hand, full body clearly visible"
597
+ },
598
+ {
599
+ "category": "anatomy",
600
+ "prompt": "an elderly man holding an object between both hands, full body clearly visible"
601
+ },
602
+ {
603
+ "category": "counting",
604
+ "prompt": "2 boxs arranged in the foreground with 4 lamps behind them"
605
+ },
606
+ {
607
+ "category": "spatial",
608
+ "prompt": "a yellow glass positioned beside a white vase"
609
+ },
610
+ {
611
+ "category": "material",
612
+ "prompt": "a box made of rough stone, showing realistic surface properties and reflections"
613
+ },
614
+ {
615
+ "category": "material",
616
+ "prompt": "a box made of dark wood, showing realistic surface properties and reflections"
617
+ },
618
+ {
619
+ "category": "reflection_occlusion",
620
+ "prompt": "a transparent object casting a faint shadow"
621
+ },
622
+ {
623
+ "category": "reflection_occlusion",
624
+ "prompt": "a mirror showing the back of a person facing away"
625
+ },
626
+ {
627
+ "category": "material",
628
+ "prompt": "a chair made of wet concrete, showing realistic surface properties and reflections"
629
+ },
630
+ {
631
+ "category": "lighting",
632
+ "prompt": "a lamp illuminated by hard light from the right, with clearly visible light direction and shadow"
633
+ },
634
+ {
635
+ "category": "spatial",
636
+ "prompt": "a black plate positioned above a yellow box"
637
+ },
638
+ {
639
+ "category": "text",
640
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a paper label"
641
+ },
642
+ {
643
+ "category": "lighting",
644
+ "prompt": "a chair illuminated by soft diffused daylight, with clearly visible light direction and shadow"
645
+ },
646
+ {
647
+ "category": "spatial",
648
+ "prompt": "a white box positioned above a yellow metal cylinder"
649
+ },
650
+ {
651
+ "category": "material",
652
+ "prompt": "a vase made of rough stone, showing realistic surface properties and reflections"
653
+ },
654
+ {
655
+ "category": "counting",
656
+ "prompt": "4 bottles arranged in the foreground with 2 boxs behind them"
657
+ },
658
+ {
659
+ "category": "spatial",
660
+ "prompt": "a orange plate positioned to the left of a green lamp"
661
+ },
662
+ {
663
+ "category": "anatomy",
664
+ "prompt": "a man sitting with both feet visible, full body clearly visible"
665
+ },
666
+ {
667
+ "category": "counting",
668
+ "prompt": "4 cubes arranged in the foreground with 1 mirrors behind them"
669
+ },
670
+ {
671
+ "category": "reflection_occlusion",
672
+ "prompt": "a mirror showing the back of a person facing away"
673
+ },
674
+ {
675
+ "category": "spatial",
676
+ "prompt": "a silver mirror positioned in front of a orange cube"
677
+ },
678
+ {
679
+ "category": "counting",
680
+ "prompt": "4 glasss arranged in the foreground with 2 wooden blocks behind them"
681
+ },
682
+ {
683
+ "category": "counting",
684
+ "prompt": "3 metal cylinders arranged in the foreground with 1 lamps behind them"
685
+ },
686
+ {
687
+ "category": "text",
688
+ "prompt": "the word \"NORTH\" printed clearly and correctly on a black poster"
689
+ },
690
+ {
691
+ "category": "material",
692
+ "prompt": "a plate made of wet concrete, showing realistic surface properties and reflections"
693
+ },
694
+ {
695
+ "category": "reflection_occlusion",
696
+ "prompt": "a shiny metal object reflecting a nearby red object"
697
+ },
698
+ {
699
+ "category": "counting",
700
+ "prompt": "5 chairs arranged in the foreground with 4 cups behind them"
701
+ },
702
+ {
703
+ "category": "material",
704
+ "prompt": "a mirror made of dark wood, showing realistic surface properties and reflections"
705
+ },
706
+ {
707
+ "category": "material",
708
+ "prompt": "a chair made of wet concrete, showing realistic surface properties and reflections"
709
+ },
710
+ {
711
+ "category": "anatomy",
712
+ "prompt": "an elderly man holding an object between both hands, full body clearly visible"
713
+ },
714
+ {
715
+ "category": "spatial",
716
+ "prompt": "a white lamp positioned in front of a black wooden block"
717
+ },
718
+ {
719
+ "category": "counting",
720
+ "prompt": "3 chairs arranged in the foreground with 4 lamps behind them"
721
+ },
722
+ {
723
+ "category": "lighting",
724
+ "prompt": "a plate illuminated by cool side lighting, with clearly visible light direction and shadow"
725
+ },
726
+ {
727
+ "category": "spatial",
728
+ "prompt": "a orange wooden block positioned to the left of a red cube"
729
+ },
730
+ {
731
+ "category": "spatial",
732
+ "prompt": "a silver chair positioned above a blue wooden block"
733
+ },
734
+ {
735
+ "category": "reflection_occlusion",
736
+ "prompt": "a transparent glass bottle in front of a person's face"
737
+ },
738
+ {
739
+ "category": "spatial",
740
+ "prompt": "a yellow bottle positioned to the left of a black vase"
741
+ },
742
+ {
743
+ "category": "text",
744
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a black poster"
745
+ },
746
+ {
747
+ "category": "anatomy",
748
+ "prompt": "a young man sitting with both feet visible, full body clearly visible"
749
+ },
750
+ {
751
+ "category": "text",
752
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a paper label"
753
+ },
754
+ {
755
+ "category": "text",
756
+ "prompt": "the word \"CAFE\" printed clearly and correctly on a paper label"
757
+ },
758
+ {
759
+ "category": "text",
760
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a white rectangular sign"
761
+ },
762
+ {
763
+ "category": "spatial",
764
+ "prompt": "a golden plate positioned below a yellow lamp"
765
+ },
766
+ {
767
+ "category": "reflection_occlusion",
768
+ "prompt": "a mirror showing the back of a person facing away"
769
+ },
770
+ {
771
+ "category": "reflection_occlusion",
772
+ "prompt": "a translucent curtain illuminated from behind"
773
+ },
774
+ {
775
+ "category": "text",
776
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a glass storefront"
777
+ },
778
+ {
779
+ "category": "anatomy",
780
+ "prompt": "a young man holding both hands in front of their chest, full body clearly visible"
781
+ },
782
+ {
783
+ "category": "anatomy",
784
+ "prompt": "an elderly woman holding an object between both hands, full body clearly visible"
785
+ },
786
+ {
787
+ "category": "lighting",
788
+ "prompt": "a table illuminated by cool side lighting, with clearly visible light direction and shadow"
789
+ },
790
+ {
791
+ "category": "spatial",
792
+ "prompt": "a white table positioned beside a black vase"
793
+ },
794
+ {
795
+ "category": "reflection_occlusion",
796
+ "prompt": "a transparent object casting a faint shadow"
797
+ },
798
+ {
799
+ "category": "material",
800
+ "prompt": "a lamp made of glossy ceramic, showing realistic surface properties and reflections"
801
+ },
802
+ {
803
+ "category": "reflection_occlusion",
804
+ "prompt": "a chrome sphere reflecting a room around it"
805
+ },
806
+ {
807
+ "category": "spatial",
808
+ "prompt": "a purple table positioned behind a blue metal cylinder"
809
+ },
810
+ {
811
+ "category": "anatomy",
812
+ "prompt": "an elderly woman holding an object between both hands, full body clearly visible"
813
+ },
814
+ {
815
+ "category": "anatomy",
816
+ "prompt": "a young woman crossing their right leg over their left leg, full body clearly visible"
817
+ },
818
+ {
819
+ "category": "anatomy",
820
+ "prompt": "an elderly woman holding an object between both hands, full body clearly visible"
821
+ },
822
+ {
823
+ "category": "spatial",
824
+ "prompt": "a red cube positioned behind a blue cup"
825
+ },
826
+ {
827
+ "category": "material",
828
+ "prompt": "a table made of translucent acrylic, showing realistic surface properties and reflections"
829
+ },
830
+ {
831
+ "category": "anatomy",
832
+ "prompt": "a young woman holding a bottle in their right hand, full body clearly visible"
833
+ },
834
+ {
835
+ "category": "spatial",
836
+ "prompt": "a blue metal cylinder positioned to the right of a red lamp"
837
+ },
838
+ {
839
+ "category": "lighting",
840
+ "prompt": "a book illuminated by warm light reflected from below, with clearly visible light direction and shadow"
841
+ },
842
+ {
843
+ "category": "counting",
844
+ "prompt": "5 books arranged in the foreground with 1 mirrors behind them"
845
+ },
846
+ {
847
+ "category": "spatial",
848
+ "prompt": "a golden metal cylinder positioned partially behind a purple vase"
849
+ },
850
+ {
851
+ "category": "spatial",
852
+ "prompt": "a white book positioned beside a yellow glass"
853
+ },
854
+ {
855
+ "category": "spatial",
856
+ "prompt": "a blue mirror positioned partially behind a silver glass"
857
+ },
858
+ {
859
+ "category": "material",
860
+ "prompt": "a cup made of translucent acrylic, showing realistic surface properties and reflections"
861
+ },
862
+ {
863
+ "category": "text",
864
+ "prompt": "the word \"STOP\" printed clearly and correctly on a white rectangular sign"
865
+ },
866
+ {
867
+ "category": "anatomy",
868
+ "prompt": "a woman holding an object between both hands, full body clearly visible"
869
+ },
870
+ {
871
+ "category": "material",
872
+ "prompt": "a cup made of rough stone, showing realistic surface properties and reflections"
873
+ },
874
+ {
875
+ "category": "anatomy",
876
+ "prompt": "a woman standing with one arm behind their back, full body clearly visible"
877
+ },
878
+ {
879
+ "category": "text",
880
+ "prompt": "the word \"STOP\" printed clearly and correctly on a white rectangular sign"
881
+ },
882
+ {
883
+ "category": "anatomy",
884
+ "prompt": "a woman touching their face with their left hand, full body clearly visible"
885
+ },
886
+ {
887
+ "category": "spatial",
888
+ "prompt": "a red plate positioned partially behind a silver sphere"
889
+ },
890
+ {
891
+ "category": "counting",
892
+ "prompt": "3 books arranged in the foreground with 1 boxs behind them"
893
+ },
894
+ {
895
+ "category": "counting",
896
+ "prompt": "5 cups arranged in the foreground with 2 metal cylinders behind them"
897
+ },
898
+ {
899
+ "category": "lighting",
900
+ "prompt": "a cube illuminated by cool side lighting, with clearly visible light direction and shadow"
901
+ },
902
+ {
903
+ "category": "material",
904
+ "prompt": "a box made of rough stone, showing realistic surface properties and reflections"
905
+ },
906
+ {
907
+ "category": "anatomy",
908
+ "prompt": "an elderly man reaching forward with their right arm, full body clearly visible"
909
+ },
910
+ {
911
+ "category": "text",
912
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a glass storefront"
913
+ },
914
+ {
915
+ "category": "anatomy",
916
+ "prompt": "a man holding a bottle in their right hand, full body clearly visible"
917
+ },
918
+ {
919
+ "category": "reflection_occlusion",
920
+ "prompt": "a transparent object casting a faint shadow"
921
+ },
922
+ {
923
+ "category": "text",
924
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a white rectangular sign"
925
+ },
926
+ {
927
+ "category": "counting",
928
+ "prompt": "3 spheres arranged in the foreground with 3 lamps behind them"
929
+ },
930
+ {
931
+ "category": "material",
932
+ "prompt": "a plate made of matte plastic, showing realistic surface properties and reflections"
933
+ },
934
+ {
935
+ "category": "counting",
936
+ "prompt": "3 glasss arranged in the foreground with 1 vases behind them"
937
+ },
938
+ {
939
+ "category": "counting",
940
+ "prompt": "5 books arranged in the foreground with 3 chairs behind them"
941
+ },
942
+ {
943
+ "category": "anatomy",
944
+ "prompt": "an elderly woman holding an object between both hands, full body clearly visible"
945
+ },
946
+ {
947
+ "category": "spatial",
948
+ "prompt": "a green sphere positioned behind a black plate"
949
+ },
950
+ {
951
+ "category": "lighting",
952
+ "prompt": "a bottle illuminated by soft diffused daylight, with clearly visible light direction and shadow"
953
+ },
954
+ {
955
+ "category": "text",
956
+ "prompt": "the word \"NORTH\" printed clearly and correctly on a metal street sign"
957
+ },
958
+ {
959
+ "category": "counting",
960
+ "prompt": "2 metal cylinders arranged in the foreground with 3 mirrors behind them"
961
+ },
962
+ {
963
+ "category": "text",
964
+ "prompt": "the word \"STOP\" printed clearly and correctly on a white rectangular sign"
965
+ },
966
+ {
967
+ "category": "text",
968
+ "prompt": "the word \"BLUE\" printed clearly and correctly on a white rectangular sign"
969
+ },
970
+ {
971
+ "category": "counting",
972
+ "prompt": "2 boxs arranged in the foreground with 3 chairs behind them"
973
+ },
974
+ {
975
+ "category": "lighting",
976
+ "prompt": "a cube illuminated by warm overhead light, with clearly visible light direction and shadow"
977
+ },
978
+ {
979
+ "category": "text",
980
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a white rectangular sign"
981
+ },
982
+ {
983
+ "category": "spatial",
984
+ "prompt": "a orange plate positioned partially behind a silver vase"
985
+ },
986
+ {
987
+ "category": "material",
988
+ "prompt": "a vase made of wet concrete, showing realistic surface properties and reflections"
989
+ },
990
+ {
991
+ "category": "anatomy",
992
+ "prompt": "a woman touching their face with their left hand, full body clearly visible"
993
+ },
994
+ {
995
+ "category": "anatomy",
996
+ "prompt": "a man holding an object between both hands, full body clearly visible"
997
+ },
998
+ {
999
+ "category": "reflection_occlusion",
1000
+ "prompt": "a shiny metal object reflecting a nearby red object"
1001
+ },
1002
+ {
1003
+ "category": "spatial",
1004
+ "prompt": "a purple metal cylinder positioned partially behind a white table"
1005
+ },
1006
+ {
1007
+ "category": "anatomy",
1008
+ "prompt": "an elderly man touching their face with their left hand, full body clearly visible"
1009
+ },
1010
+ {
1011
+ "category": "reflection_occlusion",
1012
+ "prompt": "a translucent curtain illuminated from behind"
1013
+ },
1014
+ {
1015
+ "category": "spatial",
1016
+ "prompt": "a black box positioned below a yellow wooden block"
1017
+ },
1018
+ {
1019
+ "category": "spatial",
1020
+ "prompt": "a purple table positioned partially behind a yellow wooden block"
1021
+ },
1022
+ {
1023
+ "category": "counting",
1024
+ "prompt": "4 cups arranged in the foreground with 2 glasss behind them"
1025
+ },
1026
+ {
1027
+ "category": "counting",
1028
+ "prompt": "3 wooden blocks arranged in the foreground with 3 books behind them"
1029
+ },
1030
+ {
1031
+ "category": "lighting",
1032
+ "prompt": "a lamp illuminated by hard light from the right, with clearly visible light direction and shadow"
1033
+ },
1034
+ {
1035
+ "category": "reflection_occlusion",
1036
+ "prompt": "a mirror showing the back of a person facing away"
1037
+ },
1038
+ {
1039
+ "category": "spatial",
1040
+ "prompt": "a silver wooden block positioned partially behind a green table"
1041
+ },
1042
+ {
1043
+ "category": "spatial",
1044
+ "prompt": "a red book positioned behind a white sphere"
1045
+ },
1046
+ {
1047
+ "category": "reflection_occlusion",
1048
+ "prompt": "a mirror showing the back of a person facing away"
1049
+ },
1050
+ {
1051
+ "category": "spatial",
1052
+ "prompt": "a silver metal cylinder positioned partially behind a blue sphere"
1053
+ },
1054
+ {
1055
+ "category": "spatial",
1056
+ "prompt": "a yellow sphere positioned above a red box"
1057
+ },
1058
+ {
1059
+ "category": "text",
1060
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a paper label"
1061
+ },
1062
+ {
1063
+ "category": "anatomy",
1064
+ "prompt": "an elderly man holding an object between both hands, full body clearly visible"
1065
+ },
1066
+ {
1067
+ "category": "text",
1068
+ "prompt": "the word \"BLUE\" printed clearly and correctly on a metal street sign"
1069
+ },
1070
+ {
1071
+ "category": "spatial",
1072
+ "prompt": "a yellow plate positioned above a orange table"
1073
+ },
1074
+ {
1075
+ "category": "lighting",
1076
+ "prompt": "a mirror illuminated by soft diffused daylight, with clearly visible light direction and shadow"
1077
+ },
1078
+ {
1079
+ "category": "anatomy",
1080
+ "prompt": "a woman reaching forward with their right arm, full body clearly visible"
1081
+ },
1082
+ {
1083
+ "category": "lighting",
1084
+ "prompt": "a cup illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
1085
+ },
1086
+ {
1087
+ "category": "material",
1088
+ "prompt": "a book made of wet concrete, showing realistic surface properties and reflections"
1089
+ },
1090
+ {
1091
+ "category": "reflection_occlusion",
1092
+ "prompt": "a translucent curtain illuminated from behind"
1093
+ },
1094
+ {
1095
+ "category": "text",
1096
+ "prompt": "the word \"CLOSED\" printed clearly and correctly on a metal street sign"
1097
+ },
1098
+ {
1099
+ "category": "counting",
1100
+ "prompt": "4 boxs arranged in the foreground with 2 glasss behind them"
1101
+ },
1102
+ {
1103
+ "category": "lighting",
1104
+ "prompt": "a vase illuminated by warm light reflected from below, with clearly visible light direction and shadow"
1105
+ },
1106
+ {
1107
+ "category": "reflection_occlusion",
1108
+ "prompt": "a shiny metal object reflecting a nearby red object"
1109
+ },
1110
+ {
1111
+ "category": "lighting",
1112
+ "prompt": "a book illuminated by soft diffused daylight, with clearly visible light direction and shadow"
1113
+ },
1114
+ {
1115
+ "category": "spatial",
1116
+ "prompt": "a green wooden block positioned behind a purple chair"
1117
+ },
1118
+ {
1119
+ "category": "counting",
1120
+ "prompt": "4 cups arranged in the foreground with 4 cubes behind them"
1121
+ },
1122
+ {
1123
+ "category": "anatomy",
1124
+ "prompt": "a young man touching their face with their left hand, full body clearly visible"
1125
+ },
1126
+ {
1127
+ "category": "counting",
1128
+ "prompt": "3 vases arranged in the foreground with 2 tables behind them"
1129
+ },
1130
+ {
1131
+ "category": "counting",
1132
+ "prompt": "2 wooden blocks arranged in the foreground with 1 glasss behind them"
1133
+ },
1134
+ {
1135
+ "category": "lighting",
1136
+ "prompt": "a table illuminated by hard light from the right, with clearly visible light direction and shadow"
1137
+ },
1138
+ {
1139
+ "category": "reflection_occlusion",
1140
+ "prompt": "a hand visible through a transparent glass panel"
1141
+ },
1142
+ {
1143
+ "category": "anatomy",
1144
+ "prompt": "a man holding a bottle in their right hand, full body clearly visible"
1145
+ },
1146
+ {
1147
+ "category": "anatomy",
1148
+ "prompt": "an elderly woman sitting with both feet visible, full body clearly visible"
1149
+ },
1150
+ {
1151
+ "category": "spatial",
1152
+ "prompt": "a black book positioned below a orange cup"
1153
+ },
1154
+ {
1155
+ "category": "text",
1156
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a metal street sign"
1157
+ },
1158
+ {
1159
+ "category": "material",
1160
+ "prompt": "a wooden block made of glossy ceramic, showing realistic surface properties and reflections"
1161
+ },
1162
+ {
1163
+ "category": "reflection_occlusion",
1164
+ "prompt": "one person partially occluded by another person"
1165
+ },
1166
+ {
1167
+ "category": "counting",
1168
+ "prompt": "2 books arranged in the foreground with 3 cubes behind them"
1169
+ },
1170
+ {
1171
+ "category": "material",
1172
+ "prompt": "a glass made of rough stone, showing realistic surface properties and reflections"
1173
+ },
1174
+ {
1175
+ "category": "anatomy",
1176
+ "prompt": "a woman holding a bottle in their right hand, full body clearly visible"
1177
+ },
1178
+ {
1179
+ "category": "reflection_occlusion",
1180
+ "prompt": "a translucent curtain illuminated from behind"
1181
+ },
1182
+ {
1183
+ "category": "text",
1184
+ "prompt": "the word \"CAFE\" printed clearly and correctly on a black poster"
1185
+ },
1186
+ {
1187
+ "category": "text",
1188
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a paper label"
1189
+ },
1190
+ {
1191
+ "category": "counting",
1192
+ "prompt": "3 vases arranged in the foreground with 4 lamps behind them"
1193
+ },
1194
+ {
1195
+ "category": "counting",
1196
+ "prompt": "5 boxs arranged in the foreground with 4 metal cylinders behind them"
1197
+ },
1198
+ {
1199
+ "category": "spatial",
1200
+ "prompt": "a white book positioned beside a blue lamp"
1201
+ },
1202
+ {
1203
+ "category": "anatomy",
1204
+ "prompt": "a man reaching forward with their right arm, full body clearly visible"
1205
+ },
1206
+ {
1207
+ "category": "reflection_occlusion",
1208
+ "prompt": "a glass sphere resting on a reflective metal surface"
1209
+ },
1210
+ {
1211
+ "category": "material",
1212
+ "prompt": "a bottle made of wet concrete, showing realistic surface properties and reflections"
1213
+ },
1214
+ {
1215
+ "category": "lighting",
1216
+ "prompt": "a cube illuminated by warm overhead light, with clearly visible light direction and shadow"
1217
+ },
1218
+ {
1219
+ "category": "spatial",
1220
+ "prompt": "a blue cube positioned beside a white table"
1221
+ },
1222
+ {
1223
+ "category": "lighting",
1224
+ "prompt": "a vase illuminated by soft light from the left, with clearly visible light direction and shadow"
1225
+ },
1226
+ {
1227
+ "category": "anatomy",
1228
+ "prompt": "a young man crossing their right leg over their left leg, full body clearly visible"
1229
+ },
1230
+ {
1231
+ "category": "spatial",
1232
+ "prompt": "a yellow bottle positioned beside a green book"
1233
+ },
1234
+ {
1235
+ "category": "reflection_occlusion",
1236
+ "prompt": "a translucent curtain illuminated from behind"
1237
+ },
1238
+ {
1239
+ "category": "material",
1240
+ "prompt": "a vase made of dark wood, showing realistic surface properties and reflections"
1241
+ },
1242
+ {
1243
+ "category": "material",
1244
+ "prompt": "a plate made of polished chrome, showing realistic surface properties and reflections"
1245
+ },
1246
+ {
1247
+ "category": "spatial",
1248
+ "prompt": "a black metal cylinder positioned below a orange sphere"
1249
+ },
1250
+ {
1251
+ "category": "spatial",
1252
+ "prompt": "a black glass positioned below a green lamp"
1253
+ },
1254
+ {
1255
+ "category": "counting",
1256
+ "prompt": "5 bottles arranged in the foreground with 4 books behind them"
1257
+ },
1258
+ {
1259
+ "category": "material",
1260
+ "prompt": "a glass made of translucent acrylic, showing realistic surface properties and reflections"
1261
+ },
1262
+ {
1263
+ "category": "anatomy",
1264
+ "prompt": "a young man raising their left hand above their head, full body clearly visible"
1265
+ },
1266
+ {
1267
+ "category": "anatomy",
1268
+ "prompt": "an elderly woman holding both hands in front of their chest, full body clearly visible"
1269
+ },
1270
+ {
1271
+ "category": "anatomy",
1272
+ "prompt": "a young man crossing their right leg over their left leg, full body clearly visible"
1273
+ },
1274
+ {
1275
+ "category": "lighting",
1276
+ "prompt": "a sphere illuminated by warm overhead light, with clearly visible light direction and shadow"
1277
+ },
1278
+ {
1279
+ "category": "lighting",
1280
+ "prompt": "a plate illuminated by cool side lighting, with clearly visible light direction and shadow"
1281
+ },
1282
+ {
1283
+ "category": "lighting",
1284
+ "prompt": "a box illuminated by hard light from the right, with clearly visible light direction and shadow"
1285
+ },
1286
+ {
1287
+ "category": "spatial",
1288
+ "prompt": "a red book positioned in front of a white sphere"
1289
+ },
1290
+ {
1291
+ "category": "counting",
1292
+ "prompt": "4 mirrors arranged in the foreground with 3 wooden blocks behind them"
1293
+ },
1294
+ {
1295
+ "category": "material",
1296
+ "prompt": "a table made of rough stone, showing realistic surface properties and reflections"
1297
+ },
1298
+ {
1299
+ "category": "counting",
1300
+ "prompt": "5 wooden blocks arranged in the foreground with 2 bottles behind them"
1301
+ },
1302
+ {
1303
+ "category": "anatomy",
1304
+ "prompt": "an elderly woman raising their left hand above their head, full body clearly visible"
1305
+ },
1306
+ {
1307
+ "category": "spatial",
1308
+ "prompt": "a blue sphere positioned to the left of a silver vase"
1309
+ },
1310
+ {
1311
+ "category": "spatial",
1312
+ "prompt": "a green glass positioned below a white cup"
1313
+ },
1314
+ {
1315
+ "category": "text",
1316
+ "prompt": "the word \"BLUE\" printed clearly and correctly on a paper label"
1317
+ },
1318
+ {
1319
+ "category": "spatial",
1320
+ "prompt": "a white wooden block positioned partially behind a green mirror"
1321
+ },
1322
+ {
1323
+ "category": "counting",
1324
+ "prompt": "3 chairs arranged in the foreground with 4 plates behind them"
1325
+ },
1326
+ {
1327
+ "category": "counting",
1328
+ "prompt": "2 spheres arranged in the foreground with 4 mirrors behind them"
1329
+ },
1330
+ {
1331
+ "category": "counting",
1332
+ "prompt": "4 books arranged in the foreground with 4 wooden blocks behind them"
1333
+ },
1334
+ {
1335
+ "category": "reflection_occlusion",
1336
+ "prompt": "a chrome sphere reflecting a room around it"
1337
+ },
1338
+ {
1339
+ "category": "counting",
1340
+ "prompt": "4 spheres arranged in the foreground with 2 metal cylinders behind them"
1341
+ },
1342
+ {
1343
+ "category": "anatomy",
1344
+ "prompt": "a young woman holding both hands in front of their chest, full body clearly visible"
1345
+ },
1346
+ {
1347
+ "category": "anatomy",
1348
+ "prompt": "an elderly man holding both hands in front of their chest, full body clearly visible"
1349
+ },
1350
+ {
1351
+ "category": "lighting",
1352
+ "prompt": "a bottle illuminated by warm light reflected from below, with clearly visible light direction and shadow"
1353
+ },
1354
+ {
1355
+ "category": "lighting",
1356
+ "prompt": "a metal cylinder illuminated by soft diffused daylight, with clearly visible light direction and shadow"
1357
+ },
1358
+ {
1359
+ "category": "spatial",
1360
+ "prompt": "a orange glass positioned partially behind a yellow sphere"
1361
+ },
1362
+ {
1363
+ "category": "counting",
1364
+ "prompt": "4 plates arranged in the foreground with 4 vases behind them"
1365
+ },
1366
+ {
1367
+ "category": "spatial",
1368
+ "prompt": "a green cup positioned to the right of a silver vase"
1369
+ },
1370
+ {
1371
+ "category": "spatial",
1372
+ "prompt": "a silver sphere positioned to the right of a golden metal cylinder"
1373
+ },
1374
+ {
1375
+ "category": "reflection_occlusion",
1376
+ "prompt": "a woman reflected accurately in a wall mirror"
1377
+ },
1378
+ {
1379
+ "category": "spatial",
1380
+ "prompt": "a golden bottle positioned partially behind a purple book"
1381
+ },
1382
+ {
1383
+ "category": "reflection_occlusion",
1384
+ "prompt": "a shiny metal object reflecting a nearby red object"
1385
+ },
1386
+ {
1387
+ "category": "material",
1388
+ "prompt": "a vase made of dark wood, showing realistic surface properties and reflections"
1389
+ },
1390
+ {
1391
+ "category": "material",
1392
+ "prompt": "a wooden block made of transparent glass, showing realistic surface properties and reflections"
1393
+ },
1394
+ {
1395
+ "category": "reflection_occlusion",
1396
+ "prompt": "a translucent curtain illuminated from behind"
1397
+ },
1398
+ {
1399
+ "category": "counting",
1400
+ "prompt": "4 chairs arranged in the foreground with 2 wooden blocks behind them"
1401
+ },
1402
+ {
1403
+ "category": "anatomy",
1404
+ "prompt": "an elderly woman standing with one arm behind their back, full body clearly visible"
1405
+ },
1406
+ {
1407
+ "category": "reflection_occlusion",
1408
+ "prompt": "a woman reflected accurately in a wall mirror"
1409
+ },
1410
+ {
1411
+ "category": "text",
1412
+ "prompt": "the word \"STOP\" printed clearly and correctly on a white rectangular sign"
1413
+ },
1414
+ {
1415
+ "category": "material",
1416
+ "prompt": "a book made of transparent glass, showing realistic surface properties and reflections"
1417
+ },
1418
+ {
1419
+ "category": "counting",
1420
+ "prompt": "2 mirrors arranged in the foreground with 1 tables behind them"
1421
+ },
1422
+ {
1423
+ "category": "text",
1424
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a white rectangular sign"
1425
+ },
1426
+ {
1427
+ "category": "spatial",
1428
+ "prompt": "a yellow cup positioned behind a golden cube"
1429
+ },
1430
+ {
1431
+ "category": "lighting",
1432
+ "prompt": "a cup illuminated by warm light reflected from below, with clearly visible light direction and shadow"
1433
+ },
1434
+ {
1435
+ "category": "anatomy",
1436
+ "prompt": "a man crossing their right leg over their left leg, full body clearly visible"
1437
+ },
1438
+ {
1439
+ "category": "material",
1440
+ "prompt": "a mirror made of glossy ceramic, showing realistic surface properties and reflections"
1441
+ },
1442
+ {
1443
+ "category": "material",
1444
+ "prompt": "a glass made of dark wood, showing realistic surface properties and reflections"
1445
+ },
1446
+ {
1447
+ "category": "material",
1448
+ "prompt": "a cube made of wet concrete, showing realistic surface properties and reflections"
1449
+ },
1450
+ {
1451
+ "category": "spatial",
1452
+ "prompt": "a green bottle positioned behind a black chair"
1453
+ },
1454
+ {
1455
+ "category": "counting",
1456
+ "prompt": "2 metal cylinders arranged in the foreground with 4 spheres behind them"
1457
+ },
1458
+ {
1459
+ "category": "anatomy",
1460
+ "prompt": "a young woman standing with one arm behind their back, full body clearly visible"
1461
+ },
1462
+ {
1463
+ "category": "spatial",
1464
+ "prompt": "a blue wooden block positioned behind a white book"
1465
+ },
1466
+ {
1467
+ "category": "text",
1468
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a glass storefront"
1469
+ },
1470
+ {
1471
+ "category": "counting",
1472
+ "prompt": "2 chairs arranged in the foreground with 2 cups behind them"
1473
+ },
1474
+ {
1475
+ "category": "text",
1476
+ "prompt": "the word \"OPEN\" printed clearly and correctly on a paper label"
1477
+ },
1478
+ {
1479
+ "category": "anatomy",
1480
+ "prompt": "an elderly man turning their head to the left, full body clearly visible"
1481
+ },
1482
+ {
1483
+ "category": "material",
1484
+ "prompt": "a cup made of transparent glass, showing realistic surface properties and reflections"
1485
+ },
1486
+ {
1487
+ "category": "text",
1488
+ "prompt": "the word \"NORTH\" printed clearly and correctly on a glass storefront"
1489
+ },
1490
+ {
1491
+ "category": "material",
1492
+ "prompt": "a cube made of rough stone, showing realistic surface properties and reflections"
1493
+ },
1494
+ {
1495
+ "category": "spatial",
1496
+ "prompt": "a silver vase positioned to the right of a black glass"
1497
+ },
1498
+ {
1499
+ "category": "lighting",
1500
+ "prompt": "a box illuminated by strong backlight, with clearly visible light direction and shadow"
1501
+ },
1502
+ {
1503
+ "category": "material",
1504
+ "prompt": "a cup made of polished chrome, showing realistic surface properties and reflections"
1505
+ },
1506
+ {
1507
+ "category": "reflection_occlusion",
1508
+ "prompt": "one person partially occluded by another person"
1509
+ },
1510
+ {
1511
+ "category": "material",
1512
+ "prompt": "a glass made of glossy ceramic, showing realistic surface properties and reflections"
1513
+ },
1514
+ {
1515
+ "category": "anatomy",
1516
+ "prompt": "an elderly woman reaching forward with their right arm, full body clearly visible"
1517
+ },
1518
+ {
1519
+ "category": "reflection_occlusion",
1520
+ "prompt": "a shiny metal object reflecting a nearby red object"
1521
+ },
1522
+ {
1523
+ "category": "spatial",
1524
+ "prompt": "a red mirror positioned beside a orange glass"
1525
+ },
1526
+ {
1527
+ "category": "material",
1528
+ "prompt": "a chair made of brushed metal, showing realistic surface properties and reflections"
1529
+ },
1530
+ {
1531
+ "category": "spatial",
1532
+ "prompt": "a silver mirror positioned below a purple metal cylinder"
1533
+ },
1534
+ {
1535
+ "category": "material",
1536
+ "prompt": "a wooden block made of transparent glass, showing realistic surface properties and reflections"
1537
+ },
1538
+ {
1539
+ "category": "counting",
1540
+ "prompt": "5 mirrors arranged in the foreground with 1 chairs behind them"
1541
+ },
1542
+ {
1543
+ "category": "spatial",
1544
+ "prompt": "a golden cube positioned to the right of a black table"
1545
+ },
1546
+ {
1547
+ "category": "lighting",
1548
+ "prompt": "a bottle illuminated by a narrow beam of light from above, with clearly visible light direction and shadow"
1549
+ },
1550
+ {
1551
+ "category": "anatomy",
1552
+ "prompt": "a woman raising their left hand above their head, full body clearly visible"
1553
+ },
1554
+ {
1555
+ "category": "text",
1556
+ "prompt": "the word \"STUDIO\" printed clearly and correctly on a black poster"
1557
+ },
1558
+ {
1559
+ "category": "material",
1560
+ "prompt": "a cup made of wet concrete, showing realistic surface properties and reflections"
1561
+ },
1562
+ {
1563
+ "category": "counting",
1564
+ "prompt": "5 mirrors arranged in the foreground with 3 cups behind them"
1565
+ },
1566
+ {
1567
+ "category": "counting",
1568
+ "prompt": "3 glasss arranged in the foreground with 3 plates behind them"
1569
+ },
1570
+ {
1571
+ "category": "anatomy",
1572
+ "prompt": "a young woman crossing their right leg over their left leg, full body clearly visible"
1573
+ },
1574
+ {
1575
+ "category": "reflection_occlusion",
1576
+ "prompt": "a hand visible through a transparent glass panel"
1577
+ },
1578
+ {
1579
+ "category": "spatial",
1580
+ "prompt": "a golden chair positioned above a white cup"
1581
+ },
1582
+ {
1583
+ "category": "reflection_occlusion",
1584
+ "prompt": "one person partially occluded by another person"
1585
+ },
1586
+ {
1587
+ "category": "spatial",
1588
+ "prompt": "a purple sphere positioned to the right of a green vase"
1589
+ },
1590
+ {
1591
+ "category": "anatomy",
1592
+ "prompt": "a man sitting with both feet visible, full body clearly visible"
1593
+ },
1594
+ {
1595
+ "category": "spatial",
1596
+ "prompt": "a silver plate positioned above a purple metal cylinder"
1597
+ },
1598
+ {
1599
+ "category": "material",
1600
+ "prompt": "a cup made of matte plastic, showing realistic surface properties and reflections"
1601
+ },
1602
+ {
1603
+ "category": "material",
1604
+ "prompt": "a sphere made of glossy ceramic, showing realistic surface properties and reflections"
1605
+ },
1606
+ {
1607
+ "category": "spatial",
1608
+ "prompt": "a purple lamp positioned partially behind a green cube"
1609
+ },
1610
+ {
1611
+ "category": "material",
1612
+ "prompt": "a bottle made of polished chrome, showing realistic surface properties and reflections"
1613
+ },
1614
+ {
1615
+ "category": "reflection_occlusion",
1616
+ "prompt": "a glass sphere resting on a reflective metal surface"
1617
+ },
1618
+ {
1619
+ "category": "material",
1620
+ "prompt": "a mirror made of matte plastic, showing realistic surface properties and reflections"
1621
+ },
1622
+ {
1623
+ "category": "counting",
1624
+ "prompt": "2 metal cylinders arranged in the foreground with 3 bottles behind them"
1625
+ },
1626
+ {
1627
+ "category": "anatomy",
1628
+ "prompt": "an elderly woman raising their left hand above their head, full body clearly visible"
1629
+ },
1630
+ {
1631
+ "category": "anatomy",
1632
+ "prompt": "a woman holding both hands in front of their chest, full body clearly visible"
1633
+ },
1634
+ {
1635
+ "category": "lighting",
1636
+ "prompt": "a bottle illuminated by soft diffused daylight, with clearly visible light direction and shadow"
1637
+ },
1638
+ {
1639
+ "category": "reflection_occlusion",
1640
+ "prompt": "a transparent object casting a faint shadow"
1641
+ },
1642
+ {
1643
+ "category": "spatial",
1644
+ "prompt": "a golden mirror positioned below a blue cube"
1645
+ },
1646
+ {
1647
+ "category": "material",
1648
+ "prompt": "a glass made of translucent acrylic, showing realistic surface properties and reflections"
1649
+ },
1650
+ {
1651
+ "category": "lighting",
1652
+ "prompt": "a wooden block illuminated by cool side lighting, with clearly visible light direction and shadow"
1653
+ },
1654
+ {
1655
+ "category": "counting",
1656
+ "prompt": "3 cubes arranged in the foreground with 3 bottles behind them"
1657
+ },
1658
+ {
1659
+ "category": "text",
1660
+ "prompt": "the word \"STOP\" printed clearly and correctly on a metal street sign"
1661
+ },
1662
+ {
1663
+ "category": "text",
1664
+ "prompt": "the word \"EXIT\" printed clearly and correctly on a white rectangular sign"
1665
+ },
1666
+ {
1667
+ "category": "counting",
1668
+ "prompt": "3 boxs arranged in the foreground with 4 glasss behind them"
1669
+ },
1670
+ {
1671
+ "category": "reflection_occlusion",
1672
+ "prompt": "a glass sphere resting on a reflective metal surface"
1673
+ },
1674
+ {
1675
+ "category": "spatial",
1676
+ "prompt": "a black bottle positioned to the right of a red box"
1677
+ },
1678
+ {
1679
+ "category": "anatomy",
1680
+ "prompt": "a woman standing with one arm behind their back, full body clearly visible"
1681
+ },
1682
+ {
1683
+ "category": "text",
1684
+ "prompt": "the word \"CAFE\" printed clearly and correctly on a glass storefront"
1685
+ },
1686
+ {
1687
+ "category": "anatomy",
1688
+ "prompt": "a young woman raising their left hand above their head, full body clearly visible"
1689
+ },
1690
+ {
1691
+ "category": "reflection_occlusion",
1692
+ "prompt": "a glass sphere resting on a reflective metal surface"
1693
+ },
1694
+ {
1695
+ "category": "anatomy",
1696
+ "prompt": "a young woman raising their left hand above their head, full body clearly visible"
1697
+ },
1698
+ {
1699
+ "category": "reflection_occlusion",
1700
+ "prompt": "a woman reflected accurately in a wall mirror"
1701
+ },
1702
+ {
1703
+ "category": "anatomy",
1704
+ "prompt": "a woman holding a bottle in their right hand, full body clearly visible"
1705
+ },
1706
+ {
1707
+ "category": "spatial",
1708
+ "prompt": "a black lamp positioned partially behind a orange box"
1709
+ },
1710
+ {
1711
+ "category": "spatial",
1712
+ "prompt": "a golden plate positioned to the left of a orange wooden block"
1713
+ },
1714
+ {
1715
+ "category": "counting",
1716
+ "prompt": "4 bottles arranged in the foreground with 2 wooden blocks behind them"
1717
+ },
1718
+ {
1719
+ "category": "spatial",
1720
+ "prompt": "a orange lamp positioned to the left of a black cube"
1721
+ },
1722
+ {
1723
+ "category": "counting",
1724
+ "prompt": "5 lamps arranged in the foreground with 2 glasss behind them"
1725
+ },
1726
+ {
1727
+ "category": "spatial",
1728
+ "prompt": "a yellow metal cylinder positioned to the right of a silver cube"
1729
+ },
1730
+ {
1731
+ "category": "material",
1732
+ "prompt": "a bottle made of matte plastic, showing realistic surface properties and reflections"
1733
+ },
1734
+ {
1735
+ "category": "spatial",
1736
+ "prompt": "a golden vase positioned in front of a white table"
1737
+ },
1738
+ {
1739
+ "category": "reflection_occlusion",
1740
+ "prompt": "a transparent glass bottle in front of a person's face"
1741
+ },
1742
+ {
1743
+ "category": "material",
1744
+ "prompt": "a table made of glossy ceramic, showing realistic surface properties and reflections"
1745
+ },
1746
+ {
1747
+ "category": "material",
1748
+ "prompt": "a cup made of dark wood, showing realistic surface properties and reflections"
1749
+ },
1750
+ {
1751
+ "category": "anatomy",
1752
+ "prompt": "a woman holding a bottle in their right hand, full body clearly visible"
1753
+ },
1754
+ {
1755
+ "category": "text",
1756
+ "prompt": "the word \"HELLO\" printed clearly and correctly on a paper label"
1757
+ },
1758
+ {
1759
+ "category": "material",
1760
+ "prompt": "a plate made of rough stone, showing realistic surface properties and reflections"
1761
+ }
1762
+ ]
research/raw_scripts/compare_hidden_layers_bridge.py ADDED
@@ -0,0 +1,439 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import torch
4
+
5
+ from safetensors.torch import load_file
6
+
7
+
8
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
9
+
10
+ H3_FILE = os.path.join(
11
+ ROOT,
12
+ "h3_hidden_states_fast.pt"
13
+ )
14
+
15
+ SN_FILE = os.path.join(
16
+ ROOT,
17
+ "sensenova_hidden_states_stream_v3.pt"
18
+ )
19
+
20
+ BRIDGE_FILE = os.path.join(
21
+ ROOT,
22
+ "sensenova_to_h3_embedding_bridge_rank512.safetensors"
23
+ )
24
+
25
+ OUT_TXT = os.path.join(
26
+ ROOT,
27
+ "hidden_layer_bridge_comparison.txt"
28
+ )
29
+
30
+ OUT_CSV = os.path.join(
31
+ ROOT,
32
+ "hidden_layer_bridge_comparison.csv"
33
+ )
34
+
35
+
36
+ SN_LAYERS = [8, 16, 24, 32, 41]
37
+ H3_LAYERS = [8, 16, 24, 32, 40, 49]
38
+
39
+
40
+ print("=" * 80)
41
+ print("SenseNova -> H3 hidden-layer bridge comparison")
42
+ print("=" * 80)
43
+
44
+
45
+ # ============================================================
46
+ # LOAD
47
+ # ============================================================
48
+
49
+ print("\nLoading hidden states...")
50
+
51
+ h3_data = torch.load(
52
+ H3_FILE,
53
+ map_location="cpu",
54
+ weights_only=False,
55
+ )
56
+
57
+ sn_data = torch.load(
58
+ SN_FILE,
59
+ map_location="cpu",
60
+ weights_only=False,
61
+ )
62
+
63
+ print("H3 prompts:", len(h3_data["results"]))
64
+ print("SN prompts:", len(sn_data["results"]))
65
+
66
+
67
+ print("\nLoading embedding bridge...")
68
+
69
+ bridge = load_file(
70
+ BRIDGE_FILE,
71
+ device="cpu"
72
+ )
73
+
74
+ A = bridge["proj_in.weight"].float()
75
+ B = bridge["proj_out.weight"].float()
76
+
77
+ print("A:", tuple(A.shape))
78
+ print("B:", tuple(B.shape))
79
+
80
+ # Expected:
81
+ # A = [512, 4096]
82
+ # B = [5120, 512]
83
+
84
+
85
+ # ============================================================
86
+ # PROJECTOR
87
+ # ============================================================
88
+
89
+ def project_sn(x):
90
+ """
91
+ x:
92
+ [..., 4096]
93
+
94
+ A:
95
+ [512, 4096]
96
+
97
+ B:
98
+ [5120, 512]
99
+
100
+ output:
101
+ [..., 5120]
102
+ """
103
+
104
+ x = x.float()
105
+
106
+ z = torch.nn.functional.linear(
107
+ x,
108
+ A
109
+ )
110
+
111
+ y = torch.nn.functional.linear(
112
+ z,
113
+ B
114
+ )
115
+
116
+ return y
117
+
118
+
119
+ def cosine_stats(a, b):
120
+
121
+ a = a.float()
122
+ b = b.float()
123
+
124
+ cos = torch.nn.functional.cosine_similarity(
125
+ a,
126
+ b,
127
+ dim=-1
128
+ )
129
+
130
+ return {
131
+ "mean": cos.mean().item(),
132
+ "min": cos.min().item(),
133
+ "max": cos.max().item(),
134
+ "std": cos.std().item()
135
+ if cos.numel() > 1 else 0.0,
136
+ }
137
+
138
+
139
+ # ============================================================
140
+ # VALIDATE PROMPT / TOKEN ALIGNMENT
141
+ # ============================================================
142
+
143
+ print("\nValidating prompt and token alignment...")
144
+
145
+ if len(h3_data["results"]) != len(sn_data["results"]):
146
+ raise RuntimeError(
147
+ "Different number of prompts."
148
+ )
149
+
150
+
151
+ for i, (h3_item, sn_item) in enumerate(
152
+ zip(
153
+ h3_data["results"],
154
+ sn_data["results"]
155
+ )
156
+ ):
157
+
158
+ if h3_item["prompt"] != sn_item["prompt"]:
159
+ raise RuntimeError(
160
+ f"Prompt mismatch at index {i}"
161
+ )
162
+
163
+ h_ids = h3_item["input_ids"]
164
+ s_ids = sn_item["input_ids"]
165
+
166
+ if not torch.equal(h_ids, s_ids):
167
+ raise RuntimeError(
168
+ f"Token IDs mismatch for:\n"
169
+ f"{h3_item['prompt']}"
170
+ )
171
+
172
+
173
+ print("Prompt/token alignment: PERFECT")
174
+
175
+
176
+ # ============================================================
177
+ # COMPARE
178
+ # ============================================================
179
+
180
+ rows = []
181
+
182
+ print()
183
+ print("=" * 80)
184
+ print("PAIRWISE LAYER COMPARISON")
185
+ print("=" * 80)
186
+
187
+
188
+ for sn_layer in SN_LAYERS:
189
+
190
+ for h3_layer in H3_LAYERS:
191
+
192
+ all_projected = []
193
+ all_h3 = []
194
+
195
+ prompt_scores = []
196
+
197
+ for h3_item, sn_item in zip(
198
+ h3_data["results"],
199
+ sn_data["results"]
200
+ ):
201
+
202
+ sn_x = sn_item["layers"][sn_layer]
203
+ h3_y = h3_item["layers"][h3_layer]
204
+
205
+ if sn_x.shape[:2] != h3_y.shape[:2]:
206
+ raise RuntimeError(
207
+ f"Sequence mismatch "
208
+ f"SN L{sn_layer} vs "
209
+ f"H3 L{h3_layer}"
210
+ )
211
+
212
+ projected = project_sn(sn_x)
213
+
214
+ stats = cosine_stats(
215
+ projected,
216
+ h3_y
217
+ )
218
+
219
+ prompt_scores.append(
220
+ stats["mean"]
221
+ )
222
+
223
+ all_projected.append(
224
+ projected.reshape(
225
+ -1,
226
+ projected.shape[-1]
227
+ )
228
+ )
229
+
230
+ all_h3.append(
231
+ h3_y.reshape(
232
+ -1,
233
+ h3_y.shape[-1]
234
+ )
235
+ )
236
+
237
+
238
+ P = torch.cat(
239
+ all_projected,
240
+ dim=0
241
+ )
242
+
243
+ H = torch.cat(
244
+ all_h3,
245
+ dim=0
246
+ )
247
+
248
+ stats = cosine_stats(P, H)
249
+
250
+ # Magnitudes are not directly comparable,
251
+ # but useful diagnostically.
252
+ proj_norm = P.norm(
253
+ dim=-1
254
+ ).mean().item()
255
+
256
+ h3_norm = H.norm(
257
+ dim=-1
258
+ ).mean().item()
259
+
260
+ norm_ratio = (
261
+ proj_norm / h3_norm
262
+ if h3_norm != 0 else 0
263
+ )
264
+
265
+ row = {
266
+ "sn_layer": sn_layer,
267
+ "h3_layer": h3_layer,
268
+
269
+ "cosine_mean": stats["mean"],
270
+ "cosine_std": stats["std"],
271
+ "cosine_min": stats["min"],
272
+ "cosine_max": stats["max"],
273
+
274
+ "projected_norm": proj_norm,
275
+ "h3_norm": h3_norm,
276
+ "norm_ratio": norm_ratio,
277
+
278
+ "prompt_1": prompt_scores[0],
279
+ "prompt_2": prompt_scores[1],
280
+ "prompt_3": prompt_scores[2],
281
+ }
282
+
283
+ rows.append(row)
284
+
285
+ print(
286
+ f"SN L{sn_layer:02d} "
287
+ f"-> H3 L{h3_layer:02d} | "
288
+ f"cos={stats['mean']:.6f} | "
289
+ f"norm_ratio={norm_ratio:.4f}"
290
+ )
291
+
292
+
293
+ # ============================================================
294
+ # SORT
295
+ # ============================================================
296
+
297
+ rows_sorted = sorted(
298
+ rows,
299
+ key=lambda x: x["cosine_mean"],
300
+ reverse=True
301
+ )
302
+
303
+
304
+ # ============================================================
305
+ # SAVE CSV
306
+ # ============================================================
307
+
308
+ with open(
309
+ OUT_CSV,
310
+ "w",
311
+ newline="",
312
+ encoding="utf-8"
313
+ ) as f:
314
+
315
+ writer = csv.DictWriter(
316
+ f,
317
+ fieldnames=rows_sorted[0].keys()
318
+ )
319
+
320
+ writer.writeheader()
321
+ writer.writerows(rows_sorted)
322
+
323
+
324
+ # ============================================================
325
+ # SAVE TXT
326
+ # ============================================================
327
+
328
+ with open(
329
+ OUT_TXT,
330
+ "w",
331
+ encoding="utf-8"
332
+ ) as f:
333
+
334
+ f.write(
335
+ "SenseNova -> H3 hidden-layer "
336
+ "bridge comparison\n"
337
+ )
338
+
339
+ f.write("=" * 80 + "\n\n")
340
+
341
+ f.write(
342
+ "The projector was trained ONLY "
343
+ "on token embeddings.\n"
344
+ )
345
+
346
+ f.write(
347
+ "Hidden states were unseen during "
348
+ "bridge training.\n\n"
349
+ )
350
+
351
+ f.write(
352
+ "TOP LAYER PAIRS\n"
353
+ )
354
+
355
+ f.write("-" * 80 + "\n")
356
+
357
+ for rank, row in enumerate(
358
+ rows_sorted,
359
+ 1
360
+ ):
361
+
362
+ f.write(
363
+ f"{rank:02d}. "
364
+ f"SN L{row['sn_layer']:02d} "
365
+ f"-> H3 L{row['h3_layer']:02d} | "
366
+ f"cos={row['cosine_mean']:.6f} | "
367
+ f"std={row['cosine_std']:.6f} | "
368
+ f"norm_ratio={row['norm_ratio']:.4f}\n"
369
+ )
370
+
371
+ f.write(
372
+ " prompts: "
373
+ f"{row['prompt_1']:.6f}, "
374
+ f"{row['prompt_2']:.6f}, "
375
+ f"{row['prompt_3']:.6f}\n"
376
+ )
377
+
378
+
379
+ # ============================================================
380
+ # DISPLAY TOP
381
+ # ============================================================
382
+
383
+ print()
384
+ print("=" * 80)
385
+ print("TOP 15")
386
+ print("=" * 80)
387
+
388
+ for rank, row in enumerate(
389
+ rows_sorted[:15],
390
+ 1
391
+ ):
392
+
393
+ print(
394
+ f"{rank:02d}. "
395
+ f"SN L{row['sn_layer']:02d} "
396
+ f"-> H3 L{row['h3_layer']:02d} | "
397
+ f"cos={row['cosine_mean']:.6f} | "
398
+ f"norm_ratio={row['norm_ratio']:.4f}"
399
+ )
400
+
401
+
402
+ # ============================================================
403
+ # BEST H3 TARGET FOR EACH SN LAYER
404
+ # ============================================================
405
+
406
+ print()
407
+ print("=" * 80)
408
+ print("BEST TARGET FOR EACH SENSENOVA LAYER")
409
+ print("=" * 80)
410
+
411
+ for sn_layer in SN_LAYERS:
412
+
413
+ candidates = [
414
+ r for r in rows
415
+ if r["sn_layer"] == sn_layer
416
+ ]
417
+
418
+ best = max(
419
+ candidates,
420
+ key=lambda x: x["cosine_mean"]
421
+ )
422
+
423
+ print(
424
+ f"SN L{sn_layer:02d} "
425
+ f"-> H3 L{best['h3_layer']:02d} | "
426
+ f"cos={best['cosine_mean']:.6f}"
427
+ )
428
+
429
+
430
+ print()
431
+ print("=" * 80)
432
+ print("DONE")
433
+ print()
434
+ print("TXT:")
435
+ print(OUT_TXT)
436
+ print()
437
+ print("CSV:")
438
+ print(OUT_CSV)
439
+ print("=" * 80)
research/raw_scripts/compare_minimax_sensenova.py ADDED
@@ -0,0 +1,335 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from safetensors import safe_open
2
+ from collections import defaultdict
3
+ import os
4
+ import re
5
+ import json
6
+
7
+ MINIMAX = r"E:\Models SD XL\UNET\Minimax_H3\minimax_h3_fl2va_pruned_int8_convrot.safetensors"
8
+ SENSENOVA = r"D:\ComfyUI_Python312\ComfyUI\models\unet\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
9
+
10
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\minimax_vs_sensenova_report.txt"
11
+
12
+
13
+ def load_manifest(path):
14
+ info = {}
15
+ metadata = {}
16
+
17
+ with safe_open(path, framework="pt", device="cpu") as f:
18
+ metadata = f.metadata() or {}
19
+
20
+ for key in f.keys():
21
+ t = f.get_tensor(key)
22
+ info[key] = {
23
+ "shape": tuple(t.shape),
24
+ "dtype": str(t.dtype),
25
+ "ndim": t.ndim,
26
+ }
27
+
28
+ return info, metadata
29
+
30
+
31
+ def base_module_name(key):
32
+ suffixes = [
33
+ ".weight",
34
+ ".bias",
35
+ ".weight_scale",
36
+ ".input_scale",
37
+ ".comfy_quant",
38
+ ]
39
+ for s in suffixes:
40
+ if key.endswith(s):
41
+ return key[:-len(s)]
42
+ return key
43
+
44
+
45
+ def logical_weights(manifest):
46
+ """
47
+ Берём только настоящие weight tensors.
48
+ Для Minimax INT8 ConvRot игнорируем metadata/scales.
49
+ """
50
+ result = {}
51
+
52
+ for key, v in manifest.items():
53
+ if not key.endswith(".weight"):
54
+ continue
55
+
56
+ if v["ndim"] != 2:
57
+ continue
58
+
59
+ result[key] = v
60
+
61
+ return result
62
+
63
+
64
+ def classify(key):
65
+ k = key.lower()
66
+
67
+ if any(x in k for x in [
68
+ "q_proj", ".query", ".q.", "to_q", "wq"
69
+ ]):
70
+ return "attention_q"
71
+
72
+ if any(x in k for x in [
73
+ "k_proj", ".key", ".k.", "to_k", "wk"
74
+ ]):
75
+ return "attention_k"
76
+
77
+ if any(x in k for x in [
78
+ "v_proj", ".value", ".v.", "to_v", "wv"
79
+ ]):
80
+ return "attention_v"
81
+
82
+ if any(x in k for x in [
83
+ "o_proj", "out_proj", "to_out", "wo"
84
+ ]):
85
+ return "attention_out"
86
+
87
+ if "qkv" in k:
88
+ return "attention_qkv"
89
+
90
+ if any(x in k for x in [
91
+ "gate_proj", "w1", "linear_fc1", "fc1"
92
+ ]):
93
+ return "mlp_gate_or_fc1"
94
+
95
+ if any(x in k for x in [
96
+ "up_proj", "w3"
97
+ ]):
98
+ return "mlp_up"
99
+
100
+ if any(x in k for x in [
101
+ "down_proj", "w2", "linear_fc2", "fc2"
102
+ ]):
103
+ return "mlp_down_or_fc2"
104
+
105
+ if any(x in k for x in [
106
+ "proj", "projection"
107
+ ]):
108
+ return "projection"
109
+
110
+ return "other_2d"
111
+
112
+
113
+ def layer_index(key):
114
+ patterns = [
115
+ r"layers\.(\d+)",
116
+ r"blocks\.(\d+)",
117
+ r"transformer_blocks\.(\d+)",
118
+ r"double_blocks\.(\d+)",
119
+ r"single_blocks\.(\d+)",
120
+ ]
121
+
122
+ for p in patterns:
123
+ m = re.search(p, key)
124
+ if m:
125
+ return int(m.group(1))
126
+
127
+ return None
128
+
129
+
130
+ def get_quantized_modules(manifest):
131
+ modules = set()
132
+
133
+ for key in manifest:
134
+ if key.endswith(".comfy_quant"):
135
+ modules.add(key[:-len(".comfy_quant")])
136
+
137
+ return modules
138
+
139
+
140
+ print("Loading manifests...")
141
+ mini, mini_meta = load_manifest(MINIMAX)
142
+ sense, sense_meta = load_manifest(SENSENOVA)
143
+
144
+ mini_w = logical_weights(mini)
145
+ sense_w = logical_weights(sense)
146
+
147
+ mini_quant = get_quantized_modules(mini)
148
+
149
+ exact = []
150
+ transpose = []
151
+ same_input = []
152
+ same_output = []
153
+ typed_exact = []
154
+
155
+ sense_by_shape = defaultdict(list)
156
+ sense_by_transposed_shape = defaultdict(list)
157
+ sense_by_in = defaultdict(list)
158
+ sense_by_out = defaultdict(list)
159
+
160
+ for k, v in sense_w.items():
161
+ shape = v["shape"]
162
+ sense_by_shape[shape].append(k)
163
+ sense_by_transposed_shape[(shape[1], shape[0])].append(k)
164
+ sense_by_in[shape[1]].append(k)
165
+ sense_by_out[shape[0]].append(k)
166
+
167
+ for mk, mv in mini_w.items():
168
+ ms = mv["shape"]
169
+
170
+ for sk in sense_by_shape.get(ms, []):
171
+ exact.append((mk, sk, ms))
172
+
173
+ if classify(mk) == classify(sk):
174
+ typed_exact.append((mk, sk, ms, classify(mk)))
175
+
176
+ for sk in sense_by_transposed_shape.get(ms, []):
177
+ ss = sense_w[sk]["shape"]
178
+ if ss != ms:
179
+ transpose.append((mk, sk, ms, ss))
180
+
181
+ for sk in sense_by_in.get(ms[1], []):
182
+ if sense_w[sk]["shape"] != ms:
183
+ same_input.append((mk, sk, ms, sense_w[sk]["shape"]))
184
+
185
+ for sk in sense_by_out.get(ms[0], []):
186
+ if sense_w[sk]["shape"] != ms:
187
+ same_output.append((mk, sk, ms, sense_w[sk]["shape"]))
188
+
189
+
190
+ mini_classes = defaultdict(int)
191
+ sense_classes = defaultdict(int)
192
+
193
+ for k in mini_w:
194
+ mini_classes[classify(k)] += 1
195
+
196
+ for k in sense_w:
197
+ sense_classes[classify(k)] += 1
198
+
199
+
200
+ with open(OUT, "w", encoding="utf-8") as f:
201
+ def w(x=""):
202
+ f.write(str(x) + "\n")
203
+
204
+ w("=" * 100)
205
+ w("MINIMAX H3 vs SENSENOVA STRUCTURAL COMPARISON")
206
+ w("=" * 100)
207
+ w()
208
+
209
+ w("FILES")
210
+ w("-" * 100)
211
+ w(f"MINIMAX: {MINIMAX}")
212
+ w(f"Size GB: {os.path.getsize(MINIMAX) / 1024**3:.3f}")
213
+ w(f"Tensors: {len(mini)}")
214
+ w(f"2D weights:{len(mini_w)}")
215
+ w(f"Quantized modules detected: {len(mini_quant)}")
216
+ w()
217
+ w(f"SENSENOVA: {SENSENOVA}")
218
+ w(f"Size GB: {os.path.getsize(SENSENOVA) / 1024**3:.3f}")
219
+ w(f"Tensors: {len(sense)}")
220
+ w(f"2D weights:{len(sense_w)}")
221
+ w()
222
+
223
+ w("=" * 100)
224
+ w("LAYER TYPE COUNTS")
225
+ w("=" * 100)
226
+ all_types = sorted(set(mini_classes) | set(sense_classes))
227
+ for t in all_types:
228
+ w(f"{t:24s} MiniMax={mini_classes[t]:5d} SenseNova={sense_classes[t]:5d}")
229
+ w()
230
+
231
+ w("=" * 100)
232
+ w("EXACT SHAPE MATCHES")
233
+ w("=" * 100)
234
+ w(f"Count: {len(exact)}")
235
+ for mk, sk, shape in exact[:2000]:
236
+ w(f"{shape} | {mk} <=> {sk}")
237
+ w()
238
+
239
+ w("=" * 100)
240
+ w("EXACT SHAPE + SAME SEMANTIC TYPE")
241
+ w("=" * 100)
242
+ w(f"Count: {len(typed_exact)}")
243
+ for mk, sk, shape, typ in typed_exact[:2000]:
244
+ w(f"[{typ}] {shape}")
245
+ w(f" H3: {mk}")
246
+ w(f" Sense: {sk}")
247
+ w()
248
+
249
+ w("=" * 100)
250
+ w("TRANSPOSE-COMPATIBLE SHAPES")
251
+ w("=" * 100)
252
+ w(f"Count: {len(transpose)}")
253
+ for mk, sk, ms, ss in transpose[:1000]:
254
+ w(f"H3 {ms} | Sense {ss}")
255
+ w(f" H3: {mk}")
256
+ w(f" Sense: {sk}")
257
+ w()
258
+
259
+ w("=" * 100)
260
+ w("SAME INPUT DIMENSION")
261
+ w("=" * 100)
262
+ w(f"Count: {len(same_input)}")
263
+ for mk, sk, ms, ss in same_input[:1000]:
264
+ w(f"in={ms[1]} | H3 {ms} | Sense {ss}")
265
+ w(f" H3: {mk}")
266
+ w(f" Sense: {sk}")
267
+ w()
268
+
269
+ w("=" * 100)
270
+ w("SAME OUTPUT DIMENSION")
271
+ w("=" * 100)
272
+ w(f"Count: {len(same_output)}")
273
+ for mk, sk, ms, ss in same_output[:1000]:
274
+ w(f"out={ms[0]} | H3 {ms} | Sense {ss}")
275
+ w(f" H3: {mk}")
276
+ w(f" Sense: {sk}")
277
+ w()
278
+
279
+ w("=" * 100)
280
+ w("MINIMAX QUANTIZED MODULES")
281
+ w("=" * 100)
282
+ for k in sorted(mini_quant):
283
+ shape = mini.get(k + ".weight", {}).get("shape")
284
+ w(f"{k} | {shape}")
285
+ w()
286
+
287
+ w("=" * 100)
288
+ w("MINIMAX 2D WEIGHTS")
289
+ w("=" * 100)
290
+ for k, v in sorted(mini_w.items()):
291
+ w(f"{k} | {v['shape']} | {v['dtype']} | type={classify(k)} | layer={layer_index(k)}")
292
+ w()
293
+
294
+ w("=" * 100)
295
+ w("SENSENOVA 2D WEIGHTS")
296
+ w("=" * 100)
297
+ for k, v in sorted(sense_w.items()):
298
+ w(f"{k} | {v['shape']} | {v['dtype']} | type={classify(k)} | layer={layer_index(k)}")
299
+ w()
300
+
301
+ w("=" * 100)
302
+ w("MINIMAX METADATA")
303
+ w("=" * 100)
304
+ for k, v in mini_meta.items():
305
+ if k == "_quantization_metadata":
306
+ try:
307
+ parsed = json.loads(v)
308
+ w("_quantization_metadata:")
309
+ w(json.dumps(parsed, indent=2))
310
+ except Exception:
311
+ w(f"{k} = {v}")
312
+ else:
313
+ w(f"{k} = {v}")
314
+ w()
315
+
316
+ w("=" * 100)
317
+ w("SENSENOVA METADATA")
318
+ w("=" * 100)
319
+ for k, v in sense_meta.items():
320
+ w(f"{k} = {v}")
321
+
322
+
323
+ print()
324
+ print("Done.")
325
+ print("Report:")
326
+ print(OUT)
327
+ print()
328
+ print("Summary:")
329
+ print("MiniMax 2D weights:", len(mini_w))
330
+ print("SenseNova 2D weights:", len(sense_w))
331
+ print("Exact shape matches:", len(exact))
332
+ print("Exact + same semantic type:", len(typed_exact))
333
+ print("Transpose matches:", len(transpose))
334
+ print("Same input-dim candidates:", len(same_input))
335
+ print("Same output-dim candidates:", len(same_output))
research/raw_scripts/compare_tokenizers.py ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import hashlib
4
+
5
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
6
+
7
+ sys.path.insert(0, COMFY_ROOT)
8
+ os.chdir(COMFY_ROOT)
9
+
10
+ from transformers import AutoTokenizer
11
+
12
+ print("=" * 80)
13
+ print("Tokenizer compatibility test")
14
+ print("=" * 80)
15
+
16
+ # SenseNova tokenizer — скачиваются только tokenizer/config файлы,
17
+ # а НЕ 45 GB weights.
18
+ print("\nLoading SenseNova tokenizer...")
19
+
20
+ sense_tok = AutoTokenizer.from_pretrained(
21
+ "sensenova/SenseNova-U1.5-8B-MoT",
22
+ trust_remote_code=True,
23
+ )
24
+
25
+ # MiniMax H3 uses Qwen3-VL-32B tokenizer.
26
+ print("Loading Qwen3-VL tokenizer used by H3...")
27
+
28
+ h3_tok = AutoTokenizer.from_pretrained(
29
+ "Qwen/Qwen3-VL-32B-Instruct",
30
+ trust_remote_code=True,
31
+ )
32
+
33
+ print("\nSenseNova vocab size:", len(sense_tok))
34
+ print("H3/Qwen vocab size: ", len(h3_tok))
35
+
36
+ sense_vocab = sense_tok.get_vocab()
37
+ h3_vocab = h3_tok.get_vocab()
38
+
39
+ common_tokens = set(sense_vocab) & set(h3_vocab)
40
+
41
+ same_id = 0
42
+ different_id = []
43
+
44
+ for token in common_tokens:
45
+ a = sense_vocab[token]
46
+ b = h3_vocab[token]
47
+
48
+ if a == b:
49
+ same_id += 1
50
+ else:
51
+ different_id.append((token, a, b))
52
+
53
+ print()
54
+ print("Common token strings:", len(common_tokens))
55
+ print("Same token ID: ", same_id)
56
+ print("Different token ID: ", len(different_id))
57
+
58
+ if common_tokens:
59
+ print(
60
+ "ID agreement:",
61
+ f"{100.0 * same_id / len(common_tokens):.6f}%"
62
+ )
63
+
64
+ print("\nFirst ID mismatches:")
65
+
66
+ for item in different_id[:30]:
67
+ print(item)
68
+
69
+ # -----------------------------------------------------------
70
+ # Direct ID -> token comparison
71
+ # -----------------------------------------------------------
72
+
73
+ max_compare = min(len(sense_tok), len(h3_tok))
74
+
75
+ id_mismatches = []
76
+
77
+ for i in range(max_compare):
78
+ a = sense_tok.convert_ids_to_tokens(i)
79
+ b = h3_tok.convert_ids_to_tokens(i)
80
+
81
+ if a != b:
82
+ id_mismatches.append((i, a, b))
83
+
84
+ print()
85
+ print("ID->token mismatches:", len(id_mismatches))
86
+
87
+ for x in id_mismatches[:30]:
88
+ print(x)
89
+
90
+ # -----------------------------------------------------------
91
+ # Real prompt test
92
+ # -----------------------------------------------------------
93
+
94
+ prompts = [
95
+ "a red cube on a blue sphere",
96
+ "a woman holding a transparent glass bottle",
97
+ "three people standing behind a wooden table",
98
+ "the word HELLO printed on a white sign",
99
+ "a metallic robot illuminated from the left",
100
+ "a hand holding a glass of water",
101
+ "two cars parked beside a brick building",
102
+ "a woman reflected in a mirror",
103
+ ]
104
+
105
+ print("\n" + "=" * 80)
106
+ print("PROMPT TOKENIZATION")
107
+ print("=" * 80)
108
+
109
+ all_prompt_equal = True
110
+
111
+ for text in prompts:
112
+
113
+ s = sense_tok.encode(
114
+ text,
115
+ add_special_tokens=False,
116
+ )
117
+
118
+ h = h3_tok.encode(
119
+ text,
120
+ add_special_tokens=False,
121
+ )
122
+
123
+ equal = s == h
124
+
125
+ if not equal:
126
+ all_prompt_equal = False
127
+
128
+ print()
129
+ print(text)
130
+ print("equal:", equal)
131
+ print("SenseNova:", s)
132
+ print("H3: ", h)
133
+
134
+ print("\n" + "=" * 80)
135
+
136
+ if (
137
+ len(different_id) == 0
138
+ and len(id_mismatches) == 0
139
+ and all_prompt_equal
140
+ ):
141
+ print("RESULT: TOKENIZERS ARE IDENTICAL FOR TESTED VOCAB")
142
+ else:
143
+ print("RESULT: TOKENIZERS DIFFER")
144
+
145
+ print("=" * 80)
research/raw_scripts/evaluate_distilled_student_v2_ood.py ADDED
@@ -0,0 +1,1343 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+
7
+ from safetensors.torch import load_file
8
+
9
+
10
+ # ============================================================
11
+ # PATHS
12
+ # ============================================================
13
+
14
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
15
+
16
+ H3_OOD_FILE = os.path.join(
17
+ ROOT,
18
+ "h3_ood_hidden_160.pt"
19
+ )
20
+
21
+ SN_OOD_FILE = os.path.join(
22
+ ROOT,
23
+ "sensenova_ood_hidden_160.pt"
24
+ )
25
+
26
+ TEACHER_BRIDGE = (
27
+ r"D:\ComfyUI_Krea2\ComfyUI\models\bridge"
28
+ r"\SN_L32_to_H3_L49_rank128.safetensors"
29
+ )
30
+
31
+ STUDENT_FILE = os.path.join(
32
+ ROOT,
33
+ "distilled_student_v2",
34
+ "H3_L49_to_projected_SN32_v2.safetensors"
35
+ )
36
+
37
+ OUT_DIR = os.path.join(
38
+ ROOT,
39
+ "distilled_student_v2_ood"
40
+ )
41
+
42
+ os.makedirs(
43
+ OUT_DIR,
44
+ exist_ok=True
45
+ )
46
+
47
+ REPORT_FILE = os.path.join(
48
+ OUT_DIR,
49
+ "distilled_student_v2_STRICT_OOD_report.txt"
50
+ )
51
+
52
+ CSV_FILE = os.path.join(
53
+ OUT_DIR,
54
+ "distilled_student_v2_STRICT_OOD_per_prompt.csv"
55
+ )
56
+
57
+
58
+ # ============================================================
59
+ # SETTINGS
60
+ # ============================================================
61
+
62
+ DEVICE = (
63
+ "cuda"
64
+ if torch.cuda.is_available()
65
+ else "cpu"
66
+ )
67
+
68
+ ALPHAS = [
69
+ 0.05,
70
+ 0.10,
71
+ 0.20,
72
+ 0.30,
73
+ ]
74
+
75
+ HIDDEN_DIM = 512
76
+
77
+
78
+ # ============================================================
79
+ # STUDENT ARCHITECTURE
80
+ # MUST MATCH TRAINING V2 EXACTLY
81
+ # ============================================================
82
+
83
+ class SemanticStudent(
84
+ nn.Module
85
+ ):
86
+ def __init__(
87
+ self,
88
+ input_dim=5120,
89
+ hidden_dim=512,
90
+ output_dim=5120,
91
+ ):
92
+ super().__init__()
93
+
94
+ self.fc1 = nn.Linear(
95
+ input_dim,
96
+ hidden_dim,
97
+ bias=True,
98
+ )
99
+
100
+ self.fc2 = nn.Linear(
101
+ hidden_dim,
102
+ hidden_dim,
103
+ bias=True,
104
+ )
105
+
106
+ self.fc3 = nn.Linear(
107
+ hidden_dim,
108
+ output_dim,
109
+ bias=True,
110
+ )
111
+
112
+ self.act = nn.SiLU()
113
+
114
+
115
+ def forward(
116
+ self,
117
+ x,
118
+ ):
119
+ x = self.fc1(x)
120
+
121
+ x = self.act(x)
122
+
123
+ x = self.fc2(x)
124
+
125
+ x = self.act(x)
126
+
127
+ x = self.fc3(x)
128
+
129
+ return x
130
+
131
+
132
+ # ============================================================
133
+ # HELPERS
134
+ # ============================================================
135
+
136
+ def rms_normalize(x):
137
+
138
+ x = x.float()
139
+
140
+ rms = torch.sqrt(
141
+ x.pow(2).mean(
142
+ dim=-1,
143
+ keepdim=True
144
+ ) + 1e-6
145
+ )
146
+
147
+ return x / rms
148
+
149
+
150
+ def magnitude_match(
151
+ source,
152
+ target,
153
+ ):
154
+
155
+ source = source.float()
156
+ target = target.float()
157
+
158
+ source_rms = torch.sqrt(
159
+ source.pow(2).mean(
160
+ dim=-1,
161
+ keepdim=True
162
+ ) + 1e-8
163
+ )
164
+
165
+ target_rms = torch.sqrt(
166
+ target.pow(2).mean(
167
+ dim=-1,
168
+ keepdim=True
169
+ ) + 1e-8
170
+ )
171
+
172
+ return (
173
+ source
174
+ * (
175
+ target_rms
176
+ / source_rms
177
+ )
178
+ )
179
+
180
+
181
+ def cosine_tensor(
182
+ a,
183
+ b,
184
+ ):
185
+ return F.cosine_similarity(
186
+ a.float(),
187
+ b.float(),
188
+ dim=-1,
189
+ )
190
+
191
+
192
+ def cosine_mean(
193
+ a,
194
+ b,
195
+ ):
196
+ return (
197
+ cosine_tensor(
198
+ a,
199
+ b,
200
+ )
201
+ .mean()
202
+ .item()
203
+ )
204
+
205
+
206
+ # ============================================================
207
+ # START
208
+ # ============================================================
209
+
210
+ print("=" * 100)
211
+
212
+ print(
213
+ "DISTILLED STUDENT V2"
214
+ )
215
+
216
+ print(
217
+ "STRICT 160-PROMPT OOD EVALUATION"
218
+ )
219
+
220
+ print("=" * 100)
221
+
222
+ print(
223
+ "Device:",
224
+ DEVICE
225
+ )
226
+
227
+
228
+ # ============================================================
229
+ # CHECK FILES
230
+ # ============================================================
231
+
232
+ for path in [
233
+ H3_OOD_FILE,
234
+ SN_OOD_FILE,
235
+ TEACHER_BRIDGE,
236
+ STUDENT_FILE,
237
+ ]:
238
+ if not os.path.exists(path):
239
+ raise FileNotFoundError(
240
+ f"Required file not found:\n{path}"
241
+ )
242
+
243
+
244
+ # ============================================================
245
+ # LOAD OOD DATA
246
+ # ============================================================
247
+
248
+ print(
249
+ "\nLoading H3 strict OOD states..."
250
+ )
251
+
252
+ h3_data = torch.load(
253
+ H3_OOD_FILE,
254
+ map_location="cpu",
255
+ weights_only=False,
256
+ )
257
+
258
+ print(
259
+ "H3 prompts:",
260
+ len(
261
+ h3_data["results"]
262
+ )
263
+ )
264
+
265
+
266
+ print(
267
+ "\nLoading SenseNova strict OOD states..."
268
+ )
269
+
270
+ sn_data = torch.load(
271
+ SN_OOD_FILE,
272
+ map_location="cpu",
273
+ weights_only=False,
274
+ )
275
+
276
+ print(
277
+ "SenseNova prompts:",
278
+ len(
279
+ sn_data["results"]
280
+ )
281
+ )
282
+
283
+
284
+ h3_results = h3_data[
285
+ "results"
286
+ ]
287
+
288
+ sn_results = sn_data[
289
+ "results"
290
+ ]
291
+
292
+
293
+ if len(h3_results) != 160:
294
+ raise RuntimeError(
295
+ f"Expected 160 H3 OOD prompts, "
296
+ f"got {len(h3_results)}"
297
+ )
298
+
299
+ if len(sn_results) != 160:
300
+ raise RuntimeError(
301
+ f"Expected 160 SenseNova OOD prompts, "
302
+ f"got {len(sn_results)}"
303
+ )
304
+
305
+
306
+ # ============================================================
307
+ # ALIGNMENT
308
+ # ============================================================
309
+
310
+ print(
311
+ "\nChecking strict OOD alignment..."
312
+ )
313
+
314
+ for i, (
315
+ h3_item,
316
+ sn_item,
317
+ ) in enumerate(
318
+ zip(
319
+ h3_results,
320
+ sn_results,
321
+ )
322
+ ):
323
+
324
+ if (
325
+ h3_item["prompt"]
326
+ != sn_item["prompt"]
327
+ ):
328
+ raise RuntimeError(
329
+ f"Prompt mismatch at {i}"
330
+ )
331
+
332
+ if (
333
+ h3_item["category"]
334
+ != sn_item["category"]
335
+ ):
336
+ raise RuntimeError(
337
+ f"Category mismatch at {i}"
338
+ )
339
+
340
+ if not torch.equal(
341
+ h3_item["input_ids"],
342
+ sn_item["input_ids"],
343
+ ):
344
+ raise RuntimeError(
345
+ f"Token mismatch at {i}"
346
+ )
347
+
348
+
349
+ print(
350
+ "Alignment: PERFECT"
351
+ )
352
+
353
+
354
+ # ============================================================
355
+ # LOAD TEACHER BRIDGE
356
+ # ============================================================
357
+
358
+ print(
359
+ "\nLoading Full Bridge teacher..."
360
+ )
361
+
362
+ teacher = load_file(
363
+ TEACHER_BRIDGE,
364
+ device="cpu",
365
+ )
366
+
367
+ teacher_down = (
368
+ teacher[
369
+ "down.weight"
370
+ ]
371
+ .float()
372
+ )
373
+
374
+ teacher_up = (
375
+ teacher[
376
+ "up.weight"
377
+ ]
378
+ .float()
379
+ )
380
+
381
+ print(
382
+ "Teacher down:",
383
+ tuple(
384
+ teacher_down.shape
385
+ )
386
+ )
387
+
388
+ print(
389
+ "Teacher up:",
390
+ tuple(
391
+ teacher_up.shape
392
+ )
393
+ )
394
+
395
+
396
+ if tuple(
397
+ teacher_down.shape
398
+ ) != (
399
+ 128,
400
+ 4096,
401
+ ):
402
+ raise RuntimeError(
403
+ "Unexpected teacher down shape."
404
+ )
405
+
406
+ if tuple(
407
+ teacher_up.shape
408
+ ) != (
409
+ 5120,
410
+ 128,
411
+ ):
412
+ raise RuntimeError(
413
+ "Unexpected teacher up shape."
414
+ )
415
+
416
+
417
+ # ============================================================
418
+ # LOAD DISTILLED STUDENT
419
+ # ============================================================
420
+
421
+ print(
422
+ "\nLoading frozen Distilled Student..."
423
+ )
424
+
425
+ student_weights = load_file(
426
+ STUDENT_FILE,
427
+ device="cpu",
428
+ )
429
+
430
+ student = SemanticStudent(
431
+ input_dim=5120,
432
+ hidden_dim=HIDDEN_DIM,
433
+ output_dim=5120,
434
+ )
435
+
436
+
437
+ with torch.no_grad():
438
+
439
+ student.fc1.weight.copy_(
440
+ student_weights[
441
+ "fc1.weight"
442
+ ].float()
443
+ )
444
+
445
+ student.fc1.bias.copy_(
446
+ student_weights[
447
+ "fc1.bias"
448
+ ].float()
449
+ )
450
+
451
+ student.fc2.weight.copy_(
452
+ student_weights[
453
+ "fc2.weight"
454
+ ].float()
455
+ )
456
+
457
+ student.fc2.bias.copy_(
458
+ student_weights[
459
+ "fc2.bias"
460
+ ].float()
461
+ )
462
+
463
+ student.fc3.weight.copy_(
464
+ student_weights[
465
+ "fc3.weight"
466
+ ].float()
467
+ )
468
+
469
+ student.fc3.bias.copy_(
470
+ student_weights[
471
+ "fc3.bias"
472
+ ].float()
473
+ )
474
+
475
+
476
+ student = student.to(
477
+ DEVICE
478
+ )
479
+
480
+ student.eval()
481
+
482
+
483
+ param_count = sum(
484
+ p.numel()
485
+ for p in student.parameters()
486
+ )
487
+
488
+ print(
489
+ "Student parameters:",
490
+ f"{param_count:,}"
491
+ )
492
+
493
+ print(
494
+ "Approx FP16 size:",
495
+ f"{param_count * 2 / 1024**2:.2f} MiB"
496
+ )
497
+
498
+
499
+ # ============================================================
500
+ # STORAGE
501
+ # ============================================================
502
+
503
+ rows = []
504
+
505
+ teacher_cos_all = []
506
+
507
+ correction_cos_all = []
508
+
509
+ blend_scores = {
510
+ alpha: []
511
+ for alpha
512
+ in ALPHAS
513
+ }
514
+
515
+ category_scores = {}
516
+
517
+ category_correction = {}
518
+
519
+ category_blend = {
520
+ alpha: {}
521
+ for alpha
522
+ in ALPHAS
523
+ }
524
+
525
+
526
+ # ============================================================
527
+ # STRICT OOD EVALUATION
528
+ # ============================================================
529
+
530
+ print()
531
+
532
+ print("=" * 100)
533
+
534
+ print(
535
+ "RUNNING FROZEN STUDENT ON STRICT OOD"
536
+ )
537
+
538
+ print("=" * 100)
539
+
540
+
541
+ with torch.inference_mode():
542
+
543
+ for i, (
544
+ h3_item,
545
+ sn_item,
546
+ ) in enumerate(
547
+ zip(
548
+ h3_results,
549
+ sn_results,
550
+ )
551
+ ):
552
+
553
+ prompt = h3_item[
554
+ "prompt"
555
+ ]
556
+
557
+ category = h3_item[
558
+ "category"
559
+ ]
560
+
561
+
562
+ # ====================================================
563
+ # OOD files contain only selected hidden layer
564
+ # ====================================================
565
+
566
+ h49 = (
567
+ h3_item[
568
+ "hidden"
569
+ ]
570
+ .float()
571
+ )
572
+
573
+ sn32 = (
574
+ sn_item[
575
+ "hidden"
576
+ ]
577
+ .float()
578
+ )
579
+
580
+
581
+ if h49.shape[-1] != 5120:
582
+ raise RuntimeError(
583
+ f"Bad H3 shape at {i}: "
584
+ f"{tuple(h49.shape)}"
585
+ )
586
+
587
+ if sn32.shape[-1] != 4096:
588
+ raise RuntimeError(
589
+ f"Bad SenseNova shape at {i}: "
590
+ f"{tuple(sn32.shape)}"
591
+ )
592
+
593
+
594
+ # ====================================================
595
+ # FULL BRIDGE TEACHER
596
+ # SenseNova L32 -> 4096 -> 128 -> 5120
597
+ # ====================================================
598
+
599
+ sn_norm = rms_normalize(
600
+ sn32
601
+ )
602
+
603
+ teacher_rank = F.linear(
604
+ sn_norm,
605
+ teacher_down,
606
+ )
607
+
608
+ teacher_projected = F.linear(
609
+ teacher_rank,
610
+ teacher_up,
611
+ )
612
+
613
+
614
+ # ====================================================
615
+ # DISTILLED STUDENT
616
+ # H3 L49 -> nonlinear 5120->512->512->5120
617
+ # ====================================================
618
+
619
+ student_input = rms_normalize(
620
+ h49
621
+ )
622
+
623
+ student_projected = (
624
+ student(
625
+ student_input.to(
626
+ DEVICE
627
+ )
628
+ )
629
+ .cpu()
630
+ )
631
+
632
+
633
+ # ====================================================
634
+ # METRIC 1:
635
+ # STUDENT vs PROJECTED SENSENOVA
636
+ # ====================================================
637
+
638
+ teacher_token_cos = cosine_tensor(
639
+ student_projected,
640
+ teacher_projected,
641
+ )
642
+
643
+ teacher_cos = (
644
+ teacher_token_cos
645
+ .mean()
646
+ .item()
647
+ )
648
+
649
+
650
+ # ====================================================
651
+ # SAME PER-TOKEN MAGNITUDE MATCH USED BY FULL BRIDGE
652
+ # ====================================================
653
+
654
+ teacher_scaled = magnitude_match(
655
+ teacher_projected,
656
+ h49,
657
+ )
658
+
659
+ student_scaled = magnitude_match(
660
+ student_projected,
661
+ h49,
662
+ )
663
+
664
+
665
+ # ====================================================
666
+ # METRIC 2:
667
+ # DOES STUDENT REPRODUCE THE ACTUAL CORRECTION?
668
+ # ====================================================
669
+
670
+ teacher_delta = (
671
+ teacher_scaled
672
+ - h49
673
+ )
674
+
675
+ student_delta = (
676
+ student_scaled
677
+ - h49
678
+ )
679
+
680
+ correction_token_cos = cosine_tensor(
681
+ student_delta,
682
+ teacher_delta,
683
+ )
684
+
685
+ correction_cos = (
686
+ correction_token_cos
687
+ .mean()
688
+ .item()
689
+ )
690
+
691
+
692
+ # ====================================================
693
+ # METRIC 3:
694
+ # FULL CONDITIONING vs DISTILLED CONDITIONING
695
+ # ====================================================
696
+
697
+ prompt_blends = {}
698
+
699
+ for alpha in ALPHAS:
700
+
701
+ full_conditioning = (
702
+ h49
703
+ + alpha
704
+ * teacher_delta
705
+ )
706
+
707
+ distilled_conditioning = (
708
+ h49
709
+ + alpha
710
+ * student_delta
711
+ )
712
+
713
+ blend_cos = cosine_mean(
714
+ distilled_conditioning,
715
+ full_conditioning,
716
+ )
717
+
718
+ prompt_blends[
719
+ alpha
720
+ ] = blend_cos
721
+
722
+ blend_scores[
723
+ alpha
724
+ ].append(
725
+ blend_cos
726
+ )
727
+
728
+ category_blend[
729
+ alpha
730
+ ].setdefault(
731
+ category,
732
+ []
733
+ ).append(
734
+ blend_cos
735
+ )
736
+
737
+
738
+ # ====================================================
739
+ # GLOBAL / CATEGORY
740
+ # ====================================================
741
+
742
+ teacher_cos_all.append(
743
+ teacher_cos
744
+ )
745
+
746
+ correction_cos_all.append(
747
+ correction_cos
748
+ )
749
+
750
+ category_scores.setdefault(
751
+ category,
752
+ []
753
+ ).append(
754
+ teacher_cos
755
+ )
756
+
757
+ category_correction.setdefault(
758
+ category,
759
+ []
760
+ ).append(
761
+ correction_cos
762
+ )
763
+
764
+
765
+ # ====================================================
766
+ # CSV ROW
767
+ # ====================================================
768
+
769
+ row = {
770
+ "index":
771
+ i,
772
+
773
+ "category":
774
+ category,
775
+
776
+ "token_count":
777
+ h49.shape[1],
778
+
779
+ "teacher_rep_cos":
780
+ teacher_cos,
781
+
782
+ "correction_cos":
783
+ correction_cos,
784
+
785
+ "prompt":
786
+ prompt,
787
+ }
788
+
789
+ for alpha in ALPHAS:
790
+
791
+ row[
792
+ f"blend_alpha_{alpha}"
793
+ ] = (
794
+ prompt_blends[
795
+ alpha
796
+ ]
797
+ )
798
+
799
+ rows.append(
800
+ row
801
+ )
802
+
803
+
804
+ # ====================================================
805
+ # PROGRESS
806
+ # ====================================================
807
+
808
+ if (
809
+ i == 0
810
+ or (i + 1) % 20 == 0
811
+ or i + 1 == 160
812
+ ):
813
+
814
+ print(
815
+ f"{i + 1:3d}/160 | "
816
+ f"teacher={teacher_cos:.6f} | "
817
+ f"correction={correction_cos:.6f} | "
818
+ f"blend@0.20="
819
+ f"{prompt_blends[0.20]:.6f} | "
820
+ f"{category}"
821
+ )
822
+
823
+
824
+ # ============================================================
825
+ # GLOBAL METRICS
826
+ # ============================================================
827
+
828
+ teacher_tensor = torch.tensor(
829
+ teacher_cos_all,
830
+ dtype=torch.float32,
831
+ )
832
+
833
+ correction_tensor = torch.tensor(
834
+ correction_cos_all,
835
+ dtype=torch.float32,
836
+ )
837
+
838
+
839
+ teacher_mean = (
840
+ teacher_tensor.mean().item()
841
+ )
842
+
843
+ teacher_std = (
844
+ teacher_tensor.std().item()
845
+ )
846
+
847
+ teacher_min = (
848
+ teacher_tensor.min().item()
849
+ )
850
+
851
+ teacher_max = (
852
+ teacher_tensor.max().item()
853
+ )
854
+
855
+
856
+ correction_mean = (
857
+ correction_tensor.mean().item()
858
+ )
859
+
860
+ correction_std = (
861
+ correction_tensor.std().item()
862
+ )
863
+
864
+ correction_min = (
865
+ correction_tensor.min().item()
866
+ )
867
+
868
+ correction_max = (
869
+ correction_tensor.max().item()
870
+ )
871
+
872
+
873
+ global_blend = {}
874
+
875
+ for alpha in ALPHAS:
876
+
877
+ values = torch.tensor(
878
+ blend_scores[
879
+ alpha
880
+ ],
881
+ dtype=torch.float32,
882
+ )
883
+
884
+ global_blend[
885
+ alpha
886
+ ] = {
887
+ "mean":
888
+ values.mean().item(),
889
+
890
+ "std":
891
+ values.std().item(),
892
+
893
+ "min":
894
+ values.min().item(),
895
+
896
+ "max":
897
+ values.max().item(),
898
+ }
899
+
900
+
901
+ # ============================================================
902
+ # CATEGORY RESULTS
903
+ # ============================================================
904
+
905
+ categories = sorted(
906
+ category_scores.keys()
907
+ )
908
+
909
+ category_results = {}
910
+
911
+ for category in categories:
912
+
913
+ t = torch.tensor(
914
+ category_scores[
915
+ category
916
+ ],
917
+ dtype=torch.float32,
918
+ )
919
+
920
+ c = torch.tensor(
921
+ category_correction[
922
+ category
923
+ ],
924
+ dtype=torch.float32,
925
+ )
926
+
927
+ category_results[
928
+ category
929
+ ] = {
930
+ "teacher":
931
+ t.mean().item(),
932
+
933
+ "correction":
934
+ c.mean().item(),
935
+ }
936
+
937
+ for alpha in ALPHAS:
938
+
939
+ b = torch.tensor(
940
+ category_blend[
941
+ alpha
942
+ ][
943
+ category
944
+ ],
945
+ dtype=torch.float32,
946
+ )
947
+
948
+ category_results[
949
+ category
950
+ ][
951
+ f"blend_{alpha}"
952
+ ] = (
953
+ b.mean().item()
954
+ )
955
+
956
+
957
+ # ============================================================
958
+ # SAVE CSV
959
+ # ============================================================
960
+
961
+ fields = [
962
+ "index",
963
+ "category",
964
+ "token_count",
965
+ "teacher_rep_cos",
966
+ "correction_cos",
967
+ ]
968
+
969
+ for alpha in ALPHAS:
970
+
971
+ fields.append(
972
+ f"blend_alpha_{alpha}"
973
+ )
974
+
975
+ fields.append(
976
+ "prompt"
977
+ )
978
+
979
+
980
+ with open(
981
+ CSV_FILE,
982
+ "w",
983
+ newline="",
984
+ encoding="utf-8",
985
+ ) as f:
986
+
987
+ writer = csv.DictWriter(
988
+ f,
989
+ fieldnames=fields,
990
+ )
991
+
992
+ writer.writeheader()
993
+
994
+ writer.writerows(
995
+ rows
996
+ )
997
+
998
+
999
+ # ============================================================
1000
+ # WORST / BEST
1001
+ # ============================================================
1002
+
1003
+ sorted_correction = sorted(
1004
+ rows,
1005
+ key=lambda x:
1006
+ x[
1007
+ "correction_cos"
1008
+ ]
1009
+ )
1010
+
1011
+ worst = sorted_correction[
1012
+ :10
1013
+ ]
1014
+
1015
+ best = list(
1016
+ reversed(
1017
+ sorted_correction[
1018
+ -10:
1019
+ ]
1020
+ )
1021
+ )
1022
+
1023
+
1024
+ # ============================================================
1025
+ # SAVE REPORT
1026
+ # ============================================================
1027
+
1028
+ with open(
1029
+ REPORT_FILE,
1030
+ "w",
1031
+ encoding="utf-8",
1032
+ ) as f:
1033
+
1034
+ f.write(
1035
+ "DISTILLED STUDENT V2 - STRICT OOD EVALUATION\n"
1036
+ )
1037
+
1038
+ f.write(
1039
+ "=" * 100
1040
+ + "\n\n"
1041
+ )
1042
+
1043
+ f.write(
1044
+ "Student was trained on the previous 440-prompt dataset.\n"
1045
+ )
1046
+
1047
+ f.write(
1048
+ "These 160 prompts were not used for student training.\n\n"
1049
+ )
1050
+
1051
+ f.write(
1052
+ "Student input: H3 L49\n"
1053
+ )
1054
+
1055
+ f.write(
1056
+ "Teacher: SenseNova L32 -> rank128 Full Bridge\n"
1057
+ )
1058
+
1059
+ f.write(
1060
+ "Student architecture: "
1061
+ "5120 -> 512 -> 512 -> 5120, SiLU\n\n"
1062
+ )
1063
+
1064
+
1065
+ f.write(
1066
+ "PROJECTED TEACHER REPRESENTATION\n"
1067
+ )
1068
+
1069
+ f.write(
1070
+ "-" * 100
1071
+ + "\n"
1072
+ )
1073
+
1074
+ f.write(
1075
+ f"mean = {teacher_mean:.6f}\n"
1076
+ f"std = {teacher_std:.6f}\n"
1077
+ f"min = {teacher_min:.6f}\n"
1078
+ f"max = {teacher_max:.6f}\n\n"
1079
+ )
1080
+
1081
+
1082
+ f.write(
1083
+ "FULL-BRIDGE CORRECTION\n"
1084
+ )
1085
+
1086
+ f.write(
1087
+ "-" * 100
1088
+ + "\n"
1089
+ )
1090
+
1091
+ f.write(
1092
+ f"mean = {correction_mean:.6f}\n"
1093
+ f"std = {correction_std:.6f}\n"
1094
+ f"min = {correction_min:.6f}\n"
1095
+ f"max = {correction_max:.6f}\n\n"
1096
+ )
1097
+
1098
+
1099
+ f.write(
1100
+ "FULL vs DISTILLED BLENDED CONDITIONING\n"
1101
+ )
1102
+
1103
+ f.write(
1104
+ "-" * 100
1105
+ + "\n"
1106
+ )
1107
+
1108
+ for alpha in ALPHAS:
1109
+
1110
+ r = global_blend[
1111
+ alpha
1112
+ ]
1113
+
1114
+ f.write(
1115
+ f"alpha={alpha:.2f} | "
1116
+ f"mean={r['mean']:.6f} | "
1117
+ f"std={r['std']:.6f} | "
1118
+ f"min={r['min']:.6f} | "
1119
+ f"max={r['max']:.6f}\n"
1120
+ )
1121
+
1122
+
1123
+ f.write(
1124
+ "\nCATEGORY RESULTS\n"
1125
+ )
1126
+
1127
+ f.write(
1128
+ "-" * 100
1129
+ + "\n"
1130
+ )
1131
+
1132
+ for category in categories:
1133
+
1134
+ r = category_results[
1135
+ category
1136
+ ]
1137
+
1138
+ f.write(
1139
+ f"{category:28s} | "
1140
+ f"teacher={r['teacher']:.6f} | "
1141
+ f"correction={r['correction']:.6f} | "
1142
+ f"blend@0.20={r['blend_0.2']:.6f}\n"
1143
+ )
1144
+
1145
+
1146
+ f.write(
1147
+ "\nTOP 10 WORST CORRECTION PROMPTS\n"
1148
+ )
1149
+
1150
+ f.write(
1151
+ "-" * 100
1152
+ + "\n"
1153
+ )
1154
+
1155
+ for i, row in enumerate(
1156
+ worst,
1157
+ 1,
1158
+ ):
1159
+
1160
+ f.write(
1161
+ f"{i:02d}. "
1162
+ f"correction={row['correction_cos']:.6f} | "
1163
+ f"teacher={row['teacher_rep_cos']:.6f} | "
1164
+ f"{row['category']}\n"
1165
+ )
1166
+
1167
+ f.write(
1168
+ f" {row['prompt']}\n"
1169
+ )
1170
+
1171
+
1172
+ f.write(
1173
+ "\nTOP 10 BEST CORRECTION PROMPTS\n"
1174
+ )
1175
+
1176
+ f.write(
1177
+ "-" * 100
1178
+ + "\n"
1179
+ )
1180
+
1181
+ for i, row in enumerate(
1182
+ best,
1183
+ 1,
1184
+ ):
1185
+
1186
+ f.write(
1187
+ f"{i:02d}. "
1188
+ f"correction={row['correction_cos']:.6f} | "
1189
+ f"teacher={row['teacher_rep_cos']:.6f} | "
1190
+ f"{row['category']}\n"
1191
+ )
1192
+
1193
+ f.write(
1194
+ f" {row['prompt']}\n"
1195
+ )
1196
+
1197
+
1198
+ # ============================================================
1199
+ # CONSOLE RESULT
1200
+ # ============================================================
1201
+
1202
+ print()
1203
+
1204
+ print("=" * 100)
1205
+
1206
+ print(
1207
+ "STRICT OOD FINAL RESULT"
1208
+ )
1209
+
1210
+ print("=" * 100)
1211
+
1212
+ print()
1213
+
1214
+ print(
1215
+ "Student -> projected SenseNova:"
1216
+ )
1217
+
1218
+ print(
1219
+ f" mean: {teacher_mean:.6f}"
1220
+ )
1221
+
1222
+ print(
1223
+ f" min : {teacher_min:.6f}"
1224
+ )
1225
+
1226
+
1227
+ print()
1228
+
1229
+ print(
1230
+ "Student -> Full Bridge correction:"
1231
+ )
1232
+
1233
+ print(
1234
+ f" mean: {correction_mean:.6f}"
1235
+ )
1236
+
1237
+ print(
1238
+ f" min : {correction_min:.6f}"
1239
+ )
1240
+
1241
+
1242
+ print()
1243
+
1244
+ print(
1245
+ "Full vs Distilled conditioning:"
1246
+ )
1247
+
1248
+ for alpha in ALPHAS:
1249
+
1250
+ r = global_blend[
1251
+ alpha
1252
+ ]
1253
+
1254
+ print(
1255
+ f" alpha={alpha:.2f}: "
1256
+ f"{r['mean']:.6f}"
1257
+ )
1258
+
1259
+
1260
+ print()
1261
+
1262
+ print(
1263
+ "BY CATEGORY:"
1264
+ )
1265
+
1266
+ for category in categories:
1267
+
1268
+ r = category_results[
1269
+ category
1270
+ ]
1271
+
1272
+ print(
1273
+ f"{category:28s} "
1274
+ f"teacher={r['teacher']:.6f} "
1275
+ f"correction={r['correction']:.6f}"
1276
+ )
1277
+
1278
+
1279
+ print()
1280
+
1281
+ print("=" * 100)
1282
+
1283
+ print(
1284
+ "INTERPRETATION"
1285
+ )
1286
+
1287
+ print("=" * 100)
1288
+
1289
+
1290
+ if correction_mean >= 0.90:
1291
+
1292
+ print(
1293
+ "EXCELLENT: distilled student generalizes strongly "
1294
+ "to strict OOD prompts."
1295
+ )
1296
+
1297
+ elif correction_mean >= 0.80:
1298
+
1299
+ print(
1300
+ "VERY GOOD: distilled student preserves most of "
1301
+ "the Full Bridge correction on strict OOD prompts."
1302
+ )
1303
+
1304
+ elif correction_mean >= 0.70:
1305
+
1306
+ print(
1307
+ "PROMISING: meaningful generalization exists, "
1308
+ "but additional distillation data may help."
1309
+ )
1310
+
1311
+ else:
1312
+
1313
+ print(
1314
+ "WEAK OOD GENERALIZATION: student needs a broader "
1315
+ "training/distillation dataset."
1316
+ )
1317
+
1318
+
1319
+ print()
1320
+
1321
+ print(
1322
+ "Report:"
1323
+ )
1324
+
1325
+ print(
1326
+ REPORT_FILE
1327
+ )
1328
+
1329
+ print()
1330
+
1331
+ print(
1332
+ "CSV:"
1333
+ )
1334
+
1335
+ print(
1336
+ CSV_FILE
1337
+ )
1338
+
1339
+ print()
1340
+
1341
+ print(
1342
+ "DONE"
1343
+ )
research/raw_scripts/evaluate_hidden_bridge_ood.py ADDED
@@ -0,0 +1,700 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import torch
4
+ import torch.nn.functional as F
5
+
6
+ from safetensors.torch import load_file
7
+
8
+
9
+ # ============================================================
10
+ # PATHS
11
+ # ============================================================
12
+
13
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
14
+
15
+ SN_FILE = os.path.join(
16
+ ROOT,
17
+ "sensenova_ood_hidden_160.pt"
18
+ )
19
+
20
+ H3_FILE = os.path.join(
21
+ ROOT,
22
+ "h3_ood_hidden_160.pt"
23
+ )
24
+
25
+ BRIDGE_FILE = os.path.join(
26
+ ROOT,
27
+ "hidden_bridge_probes_rank128",
28
+ "SN_L32_to_H3_L49_rank128.safetensors"
29
+ )
30
+
31
+ OUT_TXT = os.path.join(
32
+ ROOT,
33
+ "hidden_bridge_OOD_final_report.txt"
34
+ )
35
+
36
+ OUT_CSV = os.path.join(
37
+ ROOT,
38
+ "hidden_bridge_OOD_per_prompt.csv"
39
+ )
40
+
41
+
42
+ # ============================================================
43
+ # LOAD
44
+ # ============================================================
45
+
46
+ print("=" * 90)
47
+ print("SenseNova L32 -> H3 L49")
48
+ print("STRICT OOD FROZEN-BRIDGE EVALUATION")
49
+ print("=" * 90)
50
+
51
+ print("\nLoading SenseNova OOD states...")
52
+
53
+ sn_data = torch.load(
54
+ SN_FILE,
55
+ map_location="cpu",
56
+ weights_only=False
57
+ )
58
+
59
+ print("SenseNova prompts:", len(sn_data["results"]))
60
+
61
+ print("\nLoading H3 OOD states...")
62
+
63
+ h3_data = torch.load(
64
+ H3_FILE,
65
+ map_location="cpu",
66
+ weights_only=False
67
+ )
68
+
69
+ print("H3 prompts:", len(h3_data["results"]))
70
+
71
+ print("\nLoading frozen bridge...")
72
+
73
+ bridge = load_file(
74
+ BRIDGE_FILE,
75
+ device="cpu"
76
+ )
77
+
78
+ down = bridge["down.weight"].float()
79
+ up = bridge["up.weight"].float()
80
+
81
+ print("down:", tuple(down.shape))
82
+ print("up: ", tuple(up.shape))
83
+
84
+
85
+ # ============================================================
86
+ # EXPECTED SHAPES
87
+ # ============================================================
88
+
89
+ if tuple(down.shape) != (128, 4096):
90
+ raise RuntimeError(
91
+ f"Unexpected down shape: {tuple(down.shape)}"
92
+ )
93
+
94
+ if tuple(up.shape) != (5120, 128):
95
+ raise RuntimeError(
96
+ f"Unexpected up shape: {tuple(up.shape)}"
97
+ )
98
+
99
+
100
+ # ============================================================
101
+ # CHECK DATA ALIGNMENT
102
+ # ============================================================
103
+
104
+ sn_results = sn_data["results"]
105
+ h3_results = h3_data["results"]
106
+
107
+ if len(sn_results) != 160:
108
+ raise RuntimeError(
109
+ f"Expected 160 SenseNova prompts, got {len(sn_results)}"
110
+ )
111
+
112
+ if len(h3_results) != 160:
113
+ raise RuntimeError(
114
+ f"Expected 160 H3 prompts, got {len(h3_results)}"
115
+ )
116
+
117
+ print("\nChecking prompt/token alignment...")
118
+
119
+ for i, (s, h) in enumerate(
120
+ zip(sn_results, h3_results)
121
+ ):
122
+
123
+ if s["prompt"] != h["prompt"]:
124
+ raise RuntimeError(
125
+ f"Prompt mismatch at index {i}"
126
+ )
127
+
128
+ if s["category"] != h["category"]:
129
+ raise RuntimeError(
130
+ f"Category mismatch at index {i}"
131
+ )
132
+
133
+ if not torch.equal(
134
+ s["input_ids"],
135
+ h["input_ids"]
136
+ ):
137
+ raise RuntimeError(
138
+ f"Token mismatch at index {i}:\n"
139
+ f"{s['prompt']}"
140
+ )
141
+
142
+ print("Alignment: PERFECT")
143
+
144
+
145
+ # ============================================================
146
+ # SAME NORMALIZATION AS TRAINING
147
+ # ============================================================
148
+
149
+ def normalize_input(x):
150
+
151
+ x = x.float()
152
+
153
+ rms = torch.sqrt(
154
+ x.pow(2).mean(
155
+ dim=-1,
156
+ keepdim=True
157
+ ) + 1e-6
158
+ )
159
+
160
+ return x / rms
161
+
162
+
163
+ # ============================================================
164
+ # FROZEN BRIDGE
165
+ # ============================================================
166
+
167
+ def project(x):
168
+
169
+ x = normalize_input(x)
170
+
171
+ z = F.linear(
172
+ x,
173
+ down
174
+ )
175
+
176
+ y = F.linear(
177
+ z,
178
+ up
179
+ )
180
+
181
+ return y
182
+
183
+
184
+ # ============================================================
185
+ # EVALUATE
186
+ # ============================================================
187
+
188
+ all_token_cos = []
189
+ all_prompt_rows = []
190
+
191
+ category_token_scores = {}
192
+ category_prompt_scores = {}
193
+
194
+ print()
195
+ print("=" * 90)
196
+ print("EVALUATING 160 STRICT OOD PROMPTS")
197
+ print("=" * 90)
198
+
199
+ with torch.no_grad():
200
+
201
+ for i, (s, h) in enumerate(
202
+ zip(sn_results, h3_results)
203
+ ):
204
+
205
+ prompt = s["prompt"]
206
+ category = s["category"]
207
+
208
+ sn_hidden = s["hidden"].float()
209
+ h3_hidden = h["hidden"].float()
210
+
211
+ if sn_hidden.shape[-1] != 4096:
212
+ raise RuntimeError(
213
+ f"Bad SN dim at {i}: {sn_hidden.shape}"
214
+ )
215
+
216
+ if h3_hidden.shape[-1] != 5120:
217
+ raise RuntimeError(
218
+ f"Bad H3 dim at {i}: {h3_hidden.shape}"
219
+ )
220
+
221
+ if sn_hidden.shape[:2] != h3_hidden.shape[:2]:
222
+ raise RuntimeError(
223
+ f"Sequence mismatch at {i}"
224
+ )
225
+
226
+ pred = project(
227
+ sn_hidden
228
+ )
229
+
230
+ cos = F.cosine_similarity(
231
+ pred,
232
+ h3_hidden,
233
+ dim=-1
234
+ )
235
+
236
+ # [1,T] -> [T]
237
+ cos = cos.squeeze(0)
238
+
239
+ prompt_cos = (
240
+ cos.mean().item()
241
+ )
242
+
243
+ token_min = (
244
+ cos.min().item()
245
+ )
246
+
247
+ token_max = (
248
+ cos.max().item()
249
+ )
250
+
251
+ token_std = (
252
+ cos.std().item()
253
+ if cos.numel() > 1
254
+ else 0.0
255
+ )
256
+
257
+ pred_norm = (
258
+ pred.norm(
259
+ dim=-1
260
+ ).mean().item()
261
+ )
262
+
263
+ h3_norm = (
264
+ h3_hidden.norm(
265
+ dim=-1
266
+ ).mean().item()
267
+ )
268
+
269
+ norm_ratio = (
270
+ pred_norm / h3_norm
271
+ if h3_norm != 0
272
+ else 0.0
273
+ )
274
+
275
+
276
+ # ====================================================
277
+ # GLOBAL TOKEN SCORES
278
+ # ====================================================
279
+
280
+ all_token_cos.append(
281
+ cos.cpu()
282
+ )
283
+
284
+
285
+ # ====================================================
286
+ # CATEGORY SCORES
287
+ # ====================================================
288
+
289
+ category_token_scores.setdefault(
290
+ category,
291
+ []
292
+ ).extend(
293
+ cos.cpu().tolist()
294
+ )
295
+
296
+ category_prompt_scores.setdefault(
297
+ category,
298
+ []
299
+ ).append(
300
+ prompt_cos
301
+ )
302
+
303
+
304
+ # ====================================================
305
+ # PER-PROMPT ROW
306
+ # ====================================================
307
+
308
+ all_prompt_rows.append({
309
+ "index": i,
310
+ "category": category,
311
+ "token_count": cos.numel(),
312
+ "cosine_mean": prompt_cos,
313
+ "cosine_std": token_std,
314
+ "cosine_min": token_min,
315
+ "cosine_max": token_max,
316
+ "norm_ratio": norm_ratio,
317
+ "prompt": prompt,
318
+ })
319
+
320
+
321
+ if (
322
+ i == 0
323
+ or (i + 1) % 20 == 0
324
+ or i + 1 == 160
325
+ ):
326
+
327
+ print(
328
+ f"{i + 1:3d}/160 | "
329
+ f"cos={prompt_cos:.6f} | "
330
+ f"{category}"
331
+ )
332
+
333
+
334
+ # ============================================================
335
+ # GLOBAL RESULTS
336
+ # ============================================================
337
+
338
+ all_token_cos = torch.cat(
339
+ all_token_cos
340
+ )
341
+
342
+ global_mean = (
343
+ all_token_cos.mean().item()
344
+ )
345
+
346
+ global_std = (
347
+ all_token_cos.std().item()
348
+ )
349
+
350
+ global_min = (
351
+ all_token_cos.min().item()
352
+ )
353
+
354
+ global_max = (
355
+ all_token_cos.max().item()
356
+ )
357
+
358
+
359
+ prompt_values = torch.tensor(
360
+ [
361
+ x["cosine_mean"]
362
+ for x in all_prompt_rows
363
+ ],
364
+ dtype=torch.float32
365
+ )
366
+
367
+ prompt_mean = (
368
+ prompt_values.mean().item()
369
+ )
370
+
371
+ prompt_std = (
372
+ prompt_values.std().item()
373
+ )
374
+
375
+ prompt_min = (
376
+ prompt_values.min().item()
377
+ )
378
+
379
+ prompt_max = (
380
+ prompt_values.max().item()
381
+ )
382
+
383
+
384
+ # ============================================================
385
+ # CATEGORY RESULTS
386
+ # ============================================================
387
+
388
+ category_results = {}
389
+
390
+ for category in sorted(
391
+ category_token_scores
392
+ ):
393
+
394
+ token_values = torch.tensor(
395
+ category_token_scores[
396
+ category
397
+ ],
398
+ dtype=torch.float32
399
+ )
400
+
401
+ prompt_values_cat = torch.tensor(
402
+ category_prompt_scores[
403
+ category
404
+ ],
405
+ dtype=torch.float32
406
+ )
407
+
408
+ category_results[
409
+ category
410
+ ] = {
411
+ "token_cosine":
412
+ token_values.mean().item(),
413
+
414
+ "prompt_cosine":
415
+ prompt_values_cat.mean().item(),
416
+
417
+ "prompt_std":
418
+ prompt_values_cat.std().item(),
419
+
420
+ "prompt_min":
421
+ prompt_values_cat.min().item(),
422
+
423
+ "prompt_max":
424
+ prompt_values_cat.max().item(),
425
+
426
+ "prompt_count":
427
+ len(
428
+ category_prompt_scores[
429
+ category
430
+ ]
431
+ ),
432
+ }
433
+
434
+
435
+ # ============================================================
436
+ # SAVE CSV
437
+ # ============================================================
438
+
439
+ with open(
440
+ OUT_CSV,
441
+ "w",
442
+ newline="",
443
+ encoding="utf-8"
444
+ ) as f:
445
+
446
+ fields = [
447
+ "index",
448
+ "category",
449
+ "token_count",
450
+ "cosine_mean",
451
+ "cosine_std",
452
+ "cosine_min",
453
+ "cosine_max",
454
+ "norm_ratio",
455
+ "prompt",
456
+ ]
457
+
458
+ writer = csv.DictWriter(
459
+ f,
460
+ fieldnames=fields
461
+ )
462
+
463
+ writer.writeheader()
464
+
465
+ writer.writerows(
466
+ all_prompt_rows
467
+ )
468
+
469
+
470
+ # ============================================================
471
+ # SORT BEST / WORST PROMPTS
472
+ # ============================================================
473
+
474
+ best_prompts = sorted(
475
+ all_prompt_rows,
476
+ key=lambda x: x[
477
+ "cosine_mean"
478
+ ],
479
+ reverse=True
480
+ )
481
+
482
+ worst_prompts = sorted(
483
+ all_prompt_rows,
484
+ key=lambda x: x[
485
+ "cosine_mean"
486
+ ]
487
+ )
488
+
489
+
490
+ # ============================================================
491
+ # SAVE REPORT
492
+ # ============================================================
493
+
494
+ with open(
495
+ OUT_TXT,
496
+ "w",
497
+ encoding="utf-8"
498
+ ) as f:
499
+
500
+ f.write(
501
+ "SenseNova L32 -> MiniMax H3 L49\n"
502
+ )
503
+
504
+ f.write(
505
+ "STRICT OOD FROZEN-BRIDGE EVALUATION\n"
506
+ )
507
+
508
+ f.write("=" * 90 + "\n\n")
509
+
510
+ f.write(
511
+ "Bridge was trained on the previous "
512
+ "440-prompt dataset.\n"
513
+ )
514
+
515
+ f.write(
516
+ "No training or adaptation was performed "
517
+ "on these 160 OOD prompts.\n\n"
518
+ )
519
+
520
+ f.write(
521
+ "Bridge: "
522
+ "SN_L32_to_H3_L49_rank128\n\n"
523
+ )
524
+
525
+ f.write(
526
+ "GLOBAL TOKEN-LEVEL\n"
527
+ )
528
+
529
+ f.write("-" * 90 + "\n")
530
+
531
+ f.write(
532
+ f"cosine_mean = {global_mean:.6f}\n"
533
+ f"cosine_std = {global_std:.6f}\n"
534
+ f"cosine_min = {global_min:.6f}\n"
535
+ f"cosine_max = {global_max:.6f}\n\n"
536
+ )
537
+
538
+ f.write(
539
+ "GLOBAL PROMPT-LEVEL\n"
540
+ )
541
+
542
+ f.write("-" * 90 + "\n")
543
+
544
+ f.write(
545
+ f"cosine_mean = {prompt_mean:.6f}\n"
546
+ f"cosine_std = {prompt_std:.6f}\n"
547
+ f"cosine_min = {prompt_min:.6f}\n"
548
+ f"cosine_max = {prompt_max:.6f}\n\n"
549
+ )
550
+
551
+ f.write(
552
+ "CATEGORY RESULTS\n"
553
+ )
554
+
555
+ f.write("-" * 90 + "\n")
556
+
557
+ for category, r in (
558
+ category_results.items()
559
+ ):
560
+
561
+ f.write(
562
+ f"{category:28s} | "
563
+ f"prompt_cos={r['prompt_cosine']:.6f} | "
564
+ f"token_cos={r['token_cosine']:.6f} | "
565
+ f"std={r['prompt_std']:.6f} | "
566
+ f"min={r['prompt_min']:.6f} | "
567
+ f"max={r['prompt_max']:.6f}\n"
568
+ )
569
+
570
+
571
+ f.write("\n")
572
+ f.write(
573
+ "TOP 10 BEST OOD PROMPTS\n"
574
+ )
575
+
576
+ f.write("-" * 90 + "\n")
577
+
578
+ for i, row in enumerate(
579
+ best_prompts[:10],
580
+ 1
581
+ ):
582
+
583
+ f.write(
584
+ f"{i:02d}. "
585
+ f"{row['cosine_mean']:.6f} | "
586
+ f"{row['category']}\n"
587
+ )
588
+
589
+ f.write(
590
+ f" {row['prompt']}\n"
591
+ )
592
+
593
+
594
+ f.write("\n")
595
+ f.write(
596
+ "TOP 10 WORST OOD PROMPTS\n"
597
+ )
598
+
599
+ f.write("-" * 90 + "\n")
600
+
601
+ for i, row in enumerate(
602
+ worst_prompts[:10],
603
+ 1
604
+ ):
605
+
606
+ f.write(
607
+ f"{i:02d}. "
608
+ f"{row['cosine_mean']:.6f} | "
609
+ f"{row['category']}\n"
610
+ )
611
+
612
+ f.write(
613
+ f" {row['prompt']}\n"
614
+ )
615
+
616
+
617
+ # ============================================================
618
+ # CONSOLE REPORT
619
+ # ============================================================
620
+
621
+ print()
622
+ print("=" * 90)
623
+ print("FINAL STRICT OOD RESULT")
624
+ print("=" * 90)
625
+
626
+ print()
627
+ print(
628
+ "Token-level cosine:",
629
+ f"{global_mean:.6f}"
630
+ )
631
+
632
+ print(
633
+ "Prompt-level cosine:",
634
+ f"{prompt_mean:.6f}"
635
+ )
636
+
637
+ print()
638
+ print("BY CATEGORY:")
639
+ print()
640
+
641
+ for category, r in (
642
+ category_results.items()
643
+ ):
644
+
645
+ print(
646
+ f"{category:28s} "
647
+ f"{r['prompt_cosine']:.6f}"
648
+ )
649
+
650
+
651
+ print()
652
+ print("=" * 90)
653
+ print("INTERPRETATION")
654
+ print("=" * 90)
655
+
656
+ if prompt_mean >= 0.85:
657
+
658
+ print(
659
+ "EXCELLENT: bridge generalizes extremely well "
660
+ "to strict OOD prompts."
661
+ )
662
+
663
+ elif prompt_mean >= 0.80:
664
+
665
+ print(
666
+ "VERY STRONG: bridge generalizes well "
667
+ "outside the training prompt distribution."
668
+ )
669
+
670
+ elif prompt_mean >= 0.70:
671
+
672
+ print(
673
+ "GOOD: substantial cross-model alignment "
674
+ "survives OOD generalization."
675
+ )
676
+
677
+ elif prompt_mean >= 0.55:
678
+
679
+ print(
680
+ "MODERATE: real alignment exists, "
681
+ "but generalization is limited."
682
+ )
683
+
684
+ else:
685
+
686
+ print(
687
+ "WEAK: training-set alignment did not "
688
+ "generalize sufficiently."
689
+ )
690
+
691
+
692
+ print()
693
+ print("Report:")
694
+ print(OUT_TXT)
695
+
696
+ print()
697
+ print("Per-prompt CSV:")
698
+ print(OUT_CSV)
699
+
700
+ print("=" * 90)
research/raw_scripts/extract_h3_bridge_dataset.py ADDED
@@ -0,0 +1,272 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import json
4
+ import gc
5
+ import torch
6
+
7
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
8
+
9
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
10
+
11
+ H3_TE = (
12
+ r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders"
13
+ r"\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
14
+ )
15
+
16
+ PROMPT_FILE = os.path.join(
17
+ ROOT,
18
+ "bridge_prompts_480.json"
19
+ )
20
+
21
+ OUT = os.path.join(
22
+ ROOT,
23
+ "h3_bridge_hidden_480.pt"
24
+ )
25
+
26
+ TARGET_LAYERS = [
27
+ 8, 16, 24, 32, 40, 49
28
+ ]
29
+
30
+
31
+ # ============================================================
32
+ # COMFY
33
+ # ============================================================
34
+
35
+ sys.path.insert(0, COMFY_ROOT)
36
+ os.chdir(COMFY_ROOT)
37
+
38
+ import comfy.sd
39
+
40
+ from transformers import AutoTokenizer
41
+
42
+
43
+ print("=" * 80)
44
+ print("H3 bridge dataset extractor")
45
+ print("=" * 80)
46
+
47
+
48
+ # ============================================================
49
+ # PROMPTS
50
+ # ============================================================
51
+
52
+ with open(
53
+ PROMPT_FILE,
54
+ "r",
55
+ encoding="utf-8"
56
+ ) as f:
57
+
58
+ dataset = json.load(f)
59
+
60
+ print("Prompts:", len(dataset))
61
+
62
+
63
+ # ============================================================
64
+ # TOKENIZER
65
+ # ============================================================
66
+
67
+ print("\nLoading tokenizer...")
68
+
69
+ tokenizer = AutoTokenizer.from_pretrained(
70
+ "Qwen/Qwen3-VL-32B-Instruct",
71
+ trust_remote_code=True,
72
+ )
73
+
74
+
75
+ # ============================================================
76
+ # H3 TE
77
+ # ============================================================
78
+
79
+ print("Loading H3 TE...")
80
+
81
+ clip = comfy.sd.load_clip(
82
+ ckpt_paths=[H3_TE],
83
+ embedding_directory=None,
84
+ clip_type=comfy.sd.CLIPType.MINIMAX,
85
+ model_options={
86
+ "load_device": torch.device("cuda"),
87
+ "offload_device": torch.device("cpu"),
88
+ },
89
+ )
90
+
91
+ root = clip.cond_stage_model
92
+
93
+
94
+ # ============================================================
95
+ # FIND TRANSFORMER
96
+ # ============================================================
97
+
98
+ transformer = None
99
+
100
+ for name, module in root.named_modules():
101
+
102
+ if name.endswith(
103
+ "qwen3vl_32b.transformer"
104
+ ):
105
+
106
+ transformer = module
107
+
108
+ print(
109
+ "Transformer:",
110
+ name
111
+ )
112
+
113
+ break
114
+
115
+
116
+ if transformer is None:
117
+ raise RuntimeError(
118
+ "H3 transformer not found"
119
+ )
120
+
121
+
122
+ # ============================================================
123
+ # HOOKS
124
+ # ============================================================
125
+
126
+ captured = {}
127
+ hooks = {}
128
+
129
+
130
+ def make_hook(idx):
131
+
132
+ def hook(module, inputs, output):
133
+
134
+ if isinstance(output, tuple):
135
+ x = output[0]
136
+ else:
137
+ x = output
138
+
139
+ captured[idx] = (
140
+ x.detach()
141
+ .to(
142
+ device="cpu",
143
+ dtype=torch.float16
144
+ )
145
+ )
146
+
147
+ return hook
148
+
149
+
150
+ for idx in TARGET_LAYERS:
151
+
152
+ target_name = (
153
+ f"model.layers.{idx}"
154
+ )
155
+
156
+ target = None
157
+
158
+ for name, module in transformer.named_modules():
159
+
160
+ if name == target_name:
161
+ target = module
162
+ break
163
+
164
+ if target is None:
165
+ raise RuntimeError(
166
+ f"Layer {idx} not found"
167
+ )
168
+
169
+ hooks[idx] = (
170
+ target.register_forward_hook(
171
+ make_hook(idx)
172
+ )
173
+ )
174
+
175
+
176
+ print(
177
+ "Hooks:",
178
+ TARGET_LAYERS
179
+ )
180
+
181
+
182
+ # ============================================================
183
+ # EXTRACT
184
+ # ============================================================
185
+
186
+ results = []
187
+
188
+ for i, row in enumerate(dataset):
189
+
190
+ prompt = row["prompt"]
191
+
192
+ tokens = tokenizer(
193
+ prompt,
194
+ return_tensors="pt",
195
+ add_special_tokens=False,
196
+ )
197
+
198
+ input_ids = (
199
+ tokens["input_ids"]
200
+ .to("cuda")
201
+ )
202
+
203
+ captured.clear()
204
+
205
+ with torch.inference_mode():
206
+
207
+ out = transformer(
208
+ input_ids=input_ids
209
+ )
210
+
211
+ item = {
212
+ "prompt": prompt,
213
+ "category": row["category"],
214
+ "input_ids": (
215
+ input_ids.detach()
216
+ .cpu()
217
+ ),
218
+ "layers": {}
219
+ }
220
+
221
+ for idx in TARGET_LAYERS:
222
+
223
+ if idx not in captured:
224
+ raise RuntimeError(
225
+ f"Missing H3 layer {idx}"
226
+ )
227
+
228
+ item["layers"][idx] = (
229
+ captured[idx]
230
+ )
231
+
232
+ results.append(item)
233
+
234
+ del out
235
+ del input_ids
236
+
237
+ if (i + 1) % 10 == 0:
238
+
239
+ print(
240
+ f"{i + 1:4d}/"
241
+ f"{len(dataset)}"
242
+ )
243
+
244
+ if (i + 1) % 25 == 0:
245
+
246
+ torch.cuda.empty_cache()
247
+
248
+
249
+ # ============================================================
250
+ # SAVE
251
+ # ============================================================
252
+
253
+ for h in hooks.values():
254
+ h.remove()
255
+
256
+
257
+ torch.save(
258
+ {
259
+ "target_layers": TARGET_LAYERS,
260
+ "results": results,
261
+ },
262
+ OUT
263
+ )
264
+
265
+ print()
266
+ print("=" * 80)
267
+ print("SUCCESS")
268
+ print(OUT)
269
+ print("=" * 80)
270
+
271
+ gc.collect()
272
+ torch.cuda.empty_cache()
research/raw_scripts/extract_h3_embeddings_full.py ADDED
@@ -0,0 +1,188 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import torch
5
+
6
+ from safetensors.torch import save_file
7
+
8
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
9
+
10
+ H3_TE = r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
11
+
12
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\h3_qwen3vl32b_embeddings_bf16.safetensors"
13
+
14
+ VOCAB_SIZE = 151936
15
+ HIDDEN_SIZE = 5120
16
+
17
+ # Количество token vectors за один проход.
18
+ # 256 очень консервативно по памяти.
19
+ BATCH = 256
20
+
21
+
22
+ # ---------------------------------------------------------
23
+ # Load ComfyUI
24
+ # ---------------------------------------------------------
25
+
26
+ sys.path.insert(0, COMFY_ROOT)
27
+ os.chdir(COMFY_ROOT)
28
+
29
+ import comfy.sd
30
+
31
+
32
+ print("=" * 80)
33
+ print("MiniMax H3 embedding extractor")
34
+ print("=" * 80)
35
+
36
+ print("Loading H3 TE...")
37
+
38
+ clip = comfy.sd.load_clip(
39
+ ckpt_paths=[H3_TE],
40
+ embedding_directory=None,
41
+ clip_type=comfy.sd.CLIPType.MINIMAX,
42
+ model_options={
43
+ "load_device": torch.device("cpu"),
44
+ "offload_device": torch.device("cpu"),
45
+ },
46
+ )
47
+
48
+ root = clip.cond_stage_model
49
+
50
+
51
+ # ---------------------------------------------------------
52
+ # Locate embed_tokens
53
+ # ---------------------------------------------------------
54
+
55
+ embed_module = None
56
+ embed_name = None
57
+
58
+ for name, module in root.named_modules():
59
+ if "embed_tokens" in name.lower():
60
+ embed_module = module
61
+ embed_name = name
62
+ break
63
+
64
+ if embed_module is None:
65
+ raise RuntimeError("embed_tokens module not found")
66
+
67
+
68
+ print("Embedding module:")
69
+ print(embed_name)
70
+
71
+ print("Weight shape:")
72
+ print(tuple(embed_module.weight.shape))
73
+
74
+ if tuple(embed_module.weight.shape) != (VOCAB_SIZE, HIDDEN_SIZE):
75
+ raise RuntimeError(
76
+ f"Unexpected embedding shape: {tuple(embed_module.weight.shape)}"
77
+ )
78
+
79
+
80
+ # ---------------------------------------------------------
81
+ # Allocate final BF16 matrix
82
+ # ---------------------------------------------------------
83
+
84
+ print()
85
+ print("Allocating output matrix...")
86
+ print(f"{VOCAB_SIZE} x {HIDDEN_SIZE} BF16")
87
+
88
+ output = torch.empty(
89
+ (VOCAB_SIZE, HIDDEN_SIZE),
90
+ dtype=torch.bfloat16,
91
+ device="cpu",
92
+ )
93
+
94
+
95
+ # ---------------------------------------------------------
96
+ # Extract in batches
97
+ # ---------------------------------------------------------
98
+
99
+ print()
100
+ print("Extracting...")
101
+
102
+ with torch.no_grad():
103
+
104
+ for start in range(0, VOCAB_SIZE, BATCH):
105
+
106
+ end = min(start + BATCH, VOCAB_SIZE)
107
+
108
+ ids = torch.arange(
109
+ start,
110
+ end,
111
+ dtype=torch.long,
112
+ ).unsqueeze(0)
113
+
114
+ emb = embed_module(ids)
115
+
116
+ # [1, batch, 5120] -> [batch, 5120]
117
+ emb = emb.squeeze(0)
118
+
119
+ output[start:end].copy_(
120
+ emb.to(dtype=torch.bfloat16, device="cpu")
121
+ )
122
+
123
+ if start % (BATCH * 20) == 0 or end == VOCAB_SIZE:
124
+ pct = 100.0 * end / VOCAB_SIZE
125
+
126
+ print(
127
+ f"{end:6d}/{VOCAB_SIZE} "
128
+ f"({pct:6.2f}%)"
129
+ )
130
+
131
+
132
+ # ---------------------------------------------------------
133
+ # Validation
134
+ # ---------------------------------------------------------
135
+
136
+ print()
137
+ print("Validating...")
138
+
139
+ finite = torch.isfinite(output).all().item()
140
+
141
+ print("Shape :", tuple(output.shape))
142
+ print("Dtype :", output.dtype)
143
+ print("Finite:", finite)
144
+
145
+ sample = output.float()
146
+
147
+ print("Min :", sample.min().item())
148
+ print("Max :", sample.max().item())
149
+ print("Mean:", sample.mean().item())
150
+ print("Std :", sample.std().item())
151
+
152
+ del sample
153
+
154
+
155
+ if not finite:
156
+ raise RuntimeError("Non-finite values detected!")
157
+
158
+
159
+ # ---------------------------------------------------------
160
+ # Save
161
+ # ---------------------------------------------------------
162
+
163
+ print()
164
+ print("Saving:")
165
+ print(OUT)
166
+
167
+ save_file(
168
+ {
169
+ "model.embed_tokens.weight": output
170
+ },
171
+ OUT,
172
+ metadata={
173
+ "source": os.path.basename(H3_TE),
174
+ "description": "Dequantized MiniMax H3 Qwen3-VL-32B token embeddings",
175
+ "vocab_size": str(VOCAB_SIZE),
176
+ "hidden_size": str(HIDDEN_SIZE),
177
+ "dtype": "bfloat16",
178
+ },
179
+ )
180
+
181
+ print()
182
+ print("DONE")
183
+ print(OUT)
184
+
185
+ # cleanup
186
+ del output
187
+ del clip
188
+ gc.collect()
research/raw_scripts/extract_h3_embeddings_test.py ADDED
@@ -0,0 +1,139 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import torch
4
+
5
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
6
+
7
+ H3_TE = r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
8
+
9
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\h3_embedding_test.pt"
10
+
11
+ # ---------------------------------------------------------
12
+ # Add ComfyUI to Python path
13
+ # ---------------------------------------------------------
14
+
15
+ sys.path.insert(0, COMFY_ROOT)
16
+ os.chdir(COMFY_ROOT)
17
+
18
+ import comfy.sd
19
+
20
+
21
+ print("=" * 80)
22
+ print("Loading MiniMax H3 text encoder through ComfyUI")
23
+ print("=" * 80)
24
+
25
+ clip = comfy.sd.load_clip(
26
+ ckpt_paths=[H3_TE],
27
+ embedding_directory=None,
28
+ clip_type=comfy.sd.CLIPType.MINIMAX,
29
+ model_options={
30
+ # CPU keeps VRAM usage low for this test.
31
+ # Change to cuda later if needed.
32
+ "load_device": torch.device("cpu"),
33
+ "offload_device": torch.device("cpu"),
34
+ },
35
+ )
36
+
37
+ print("CLIP loaded.")
38
+ print("clip type:", type(clip))
39
+
40
+
41
+ # ---------------------------------------------------------
42
+ # Inspect module tree
43
+ # ---------------------------------------------------------
44
+
45
+ root = clip.cond_stage_model
46
+
47
+ print("\ncond_stage_model:", type(root))
48
+
49
+ candidates = []
50
+
51
+ for name, module in root.named_modules():
52
+ lname = name.lower()
53
+
54
+ if "embed_tokens" in lname:
55
+ candidates.append((name, module))
56
+
57
+ print("\nFound embed_tokens candidates:", len(candidates))
58
+
59
+ for name, module in candidates:
60
+ print()
61
+ print("MODULE:", name)
62
+ print("TYPE: ", type(module))
63
+
64
+ if hasattr(module, "weight"):
65
+ try:
66
+ print("WEIGHT TYPE:", type(module.weight))
67
+ print("WEIGHT SHAPE:", tuple(module.weight.shape))
68
+ print("WEIGHT DTYPE:", module.weight.dtype)
69
+ print("WEIGHT DEVICE:", module.weight.device)
70
+ except Exception as e:
71
+ print("Could not inspect weight directly:", repr(e))
72
+
73
+
74
+ if not candidates:
75
+ print("\nNo embed_tokens module found.")
76
+ print("Dumping likely embedding/module names:")
77
+
78
+ for name, module in root.named_modules():
79
+ lname = name.lower()
80
+
81
+ if (
82
+ "embed" in lname
83
+ or "qwen" in lname
84
+ or "transformer" in lname
85
+ ):
86
+ print(name, type(module))
87
+
88
+ raise SystemExit(1)
89
+
90
+
91
+ # ---------------------------------------------------------
92
+ # Test embedding lookup
93
+ # ---------------------------------------------------------
94
+
95
+ name, embed_module = candidates[0]
96
+
97
+ print("\nUsing:")
98
+ print(name)
99
+
100
+ # Several safe token IDs spread through the vocabulary.
101
+ test_ids = torch.tensor(
102
+ [[0, 1, 10, 100, 1000, 10000, 50000, 100000, 151000]],
103
+ dtype=torch.long,
104
+ )
105
+
106
+ print("\nRunning embedding lookup...")
107
+
108
+ with torch.no_grad():
109
+ embeddings = embed_module(test_ids)
110
+
111
+ print("OUTPUT SHAPE:", tuple(embeddings.shape))
112
+ print("OUTPUT DTYPE:", embeddings.dtype)
113
+ print("OUTPUT DEVICE:", embeddings.device)
114
+
115
+ embeddings_cpu = embeddings.detach().float().cpu()
116
+
117
+ print("\nVector statistics:")
118
+ print("min :", embeddings_cpu.min().item())
119
+ print("max :", embeddings_cpu.max().item())
120
+ print("mean:", embeddings_cpu.mean().item())
121
+ print("std :", embeddings_cpu.std().item())
122
+
123
+ finite = torch.isfinite(embeddings_cpu).all().item()
124
+ print("all finite:", finite)
125
+
126
+
127
+ torch.save(
128
+ {
129
+ "token_ids": test_ids.cpu(),
130
+ "embeddings": embeddings_cpu,
131
+ "module_name": name,
132
+ },
133
+ OUT,
134
+ )
135
+
136
+ print("\nSaved:")
137
+ print(OUT)
138
+
139
+ print("\nSUCCESS")
research/raw_scripts/extract_h3_hidden_states.py ADDED
@@ -0,0 +1,189 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import torch
4
+
5
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
6
+ H3_TE = r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
7
+
8
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\h3_hidden_states.pt"
9
+
10
+ sys.path.insert(0, COMFY_ROOT)
11
+ os.chdir(COMFY_ROOT)
12
+
13
+ import comfy.sd
14
+ from transformers import AutoTokenizer
15
+
16
+
17
+ PROMPTS = [
18
+ "a red cube on a blue sphere",
19
+ "a woman holding a transparent glass bottle",
20
+ "three people standing behind a wooden table",
21
+ "the word HELLO printed on a white sign",
22
+ "a metallic robot illuminated from the left",
23
+ "a hand holding a glass of water",
24
+ "two cars parked beside a brick building",
25
+ "a woman reflected in a mirror",
26
+ "a person standing behind another person",
27
+ "a glass sphere resting on a metal cube",
28
+ "a red chair to the left of a blue table",
29
+ "four candles arranged in a square",
30
+ "a hand with five clearly visible fingers",
31
+ "a woman sitting with crossed legs",
32
+ "a person holding an umbrella above their head",
33
+ "the word OPEN written in large black letters",
34
+ "a transparent bottle filled halfway with water",
35
+ "polished chrome reflecting a warm light",
36
+ "rough stone illuminated by soft side light",
37
+ "a wooden table casting a shadow to the right",
38
+ ]
39
+
40
+
41
+ print("Loading tokenizer...")
42
+ tokenizer = AutoTokenizer.from_pretrained(
43
+ "Qwen/Qwen3-VL-32B-Instruct",
44
+ trust_remote_code=True,
45
+ )
46
+
47
+ print("Loading H3 TE...")
48
+
49
+ clip = comfy.sd.load_clip(
50
+ ckpt_paths=[H3_TE],
51
+ embedding_directory=None,
52
+ clip_type=comfy.sd.CLIPType.MINIMAX,
53
+ model_options={
54
+ "load_device": torch.device("cpu"),
55
+ "offload_device": torch.device("cpu"),
56
+ },
57
+ )
58
+
59
+ root = clip.cond_stage_model
60
+
61
+ # Find underlying transformer
62
+ transformer = None
63
+
64
+ for name, module in root.named_modules():
65
+ if name.endswith("qwen3vl_32b.transformer"):
66
+ transformer = module
67
+ print("Found transformer:", name)
68
+ break
69
+
70
+ if transformer is None:
71
+ print("Could not find exact transformer name.")
72
+ print("Candidates:")
73
+ for name, module in root.named_modules():
74
+ if "transformer" in name.lower():
75
+ print(name, type(module))
76
+ raise SystemExit(1)
77
+
78
+
79
+ results = []
80
+
81
+ hooks = {}
82
+ captured = {}
83
+
84
+ # Capture selected layers.
85
+ TARGET_LAYERS = [8, 16, 24, 32, 40, 49]
86
+
87
+ for layer_idx in TARGET_LAYERS:
88
+ layer_name = f"model.layers.{layer_idx}"
89
+
90
+ target_module = None
91
+
92
+ for name, module in transformer.named_modules():
93
+ if name == layer_name:
94
+ target_module = module
95
+ break
96
+
97
+ if target_module is None:
98
+ print("Layer not found:", layer_name)
99
+ continue
100
+
101
+ def make_hook(idx):
102
+ def hook(module, inputs, output):
103
+ if isinstance(output, tuple):
104
+ x = output[0]
105
+ else:
106
+ x = output
107
+
108
+ captured[idx] = x.detach().float().cpu()
109
+ return hook
110
+
111
+ hooks[layer_idx] = target_module.register_forward_hook(
112
+ make_hook(layer_idx)
113
+ )
114
+
115
+
116
+ print("Registered hooks:", list(hooks.keys()))
117
+
118
+
119
+ for i, prompt in enumerate(PROMPTS):
120
+
121
+ print()
122
+ print(f"[{i+1}/{len(PROMPTS)}]")
123
+ print(prompt)
124
+
125
+ tokens = tokenizer(
126
+ prompt,
127
+ return_tensors="pt",
128
+ add_special_tokens=False,
129
+ )
130
+
131
+ input_ids = tokens["input_ids"]
132
+
133
+ captured.clear()
134
+
135
+ # Use embedding/transformer directly.
136
+ with torch.no_grad():
137
+ try:
138
+ out = transformer(
139
+ input_ids=input_ids,
140
+ output_hidden_states=True,
141
+ return_dict=True,
142
+ )
143
+ except TypeError:
144
+ out = transformer(
145
+ input_ids=input_ids,
146
+ )
147
+
148
+ item = {
149
+ "prompt": prompt,
150
+ "input_ids": input_ids.cpu(),
151
+ "layers": {},
152
+ }
153
+
154
+ for idx in TARGET_LAYERS:
155
+ if idx in captured:
156
+ item["layers"][idx] = captured[idx]
157
+ print(
158
+ f" layer {idx}:",
159
+ tuple(captured[idx].shape)
160
+ )
161
+
162
+ if hasattr(out, "last_hidden_state"):
163
+ item["last_hidden_state"] = (
164
+ out.last_hidden_state.detach().float().cpu()
165
+ )
166
+ print(
167
+ " final:",
168
+ tuple(item["last_hidden_state"].shape)
169
+ )
170
+
171
+ results.append(item)
172
+
173
+
174
+ for h in hooks.values():
175
+ h.remove()
176
+
177
+ torch.save(
178
+ {
179
+ "prompts": PROMPTS,
180
+ "target_layers": TARGET_LAYERS,
181
+ "results": results,
182
+ },
183
+ OUT,
184
+ )
185
+
186
+ print()
187
+ print("Saved:")
188
+ print(OUT)
189
+ print("DONE")
research/raw_scripts/extract_h3_hidden_states_fast.py ADDED
@@ -0,0 +1,266 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import torch
4
+
5
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
6
+ H3_TE = r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
7
+
8
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\h3_hidden_states_fast.pt"
9
+
10
+ sys.path.insert(0, COMFY_ROOT)
11
+ os.chdir(COMFY_ROOT)
12
+
13
+ import comfy.sd
14
+ from transformers import AutoTokenizer
15
+
16
+
17
+ PROMPTS = [
18
+ "a red cube on a blue sphere",
19
+ "a woman holding a transparent glass bottle",
20
+ "three people standing behind a wooden table",
21
+ ]
22
+
23
+ TARGET_LAYERS = [8, 16, 24, 32, 40, 49]
24
+
25
+
26
+ print("=" * 80)
27
+ print("H3 hidden state extractor - FAST CUDA TEST")
28
+ print("=" * 80)
29
+
30
+ print("\nLoading tokenizer...")
31
+
32
+ tokenizer = AutoTokenizer.from_pretrained(
33
+ "Qwen/Qwen3-VL-32B-Instruct",
34
+ trust_remote_code=True,
35
+ )
36
+
37
+ print("Loading H3 TE on CUDA...")
38
+
39
+ clip = comfy.sd.load_clip(
40
+ ckpt_paths=[H3_TE],
41
+ embedding_directory=None,
42
+ clip_type=comfy.sd.CLIPType.MINIMAX,
43
+ model_options={
44
+ "load_device": torch.device("cuda"),
45
+ "offload_device": torch.device("cpu"),
46
+ },
47
+ )
48
+
49
+ root = clip.cond_stage_model
50
+
51
+ print("CLIP loaded.")
52
+
53
+
54
+ # =========================================================
55
+ # FIND TRANSFORMER
56
+ # =========================================================
57
+
58
+ transformer = None
59
+ transformer_name = None
60
+
61
+ for name, module in root.named_modules():
62
+ if name.endswith("qwen3vl_32b.transformer"):
63
+ transformer = module
64
+ transformer_name = name
65
+ break
66
+
67
+ if transformer is None:
68
+ print("\nCould not find exact transformer.")
69
+ print("Possible transformer modules:")
70
+
71
+ for name, module in root.named_modules():
72
+ if "transformer" in name.lower():
73
+ print(name, type(module))
74
+
75
+ raise SystemExit(1)
76
+
77
+ print("Found transformer:", transformer_name)
78
+
79
+
80
+ # =========================================================
81
+ # REGISTER HOOKS
82
+ # =========================================================
83
+
84
+ captured = {}
85
+ hooks = {}
86
+
87
+
88
+ def make_hook(idx):
89
+ def hook(module, inputs, output):
90
+
91
+ if isinstance(output, tuple):
92
+ x = output[0]
93
+ else:
94
+ x = output
95
+
96
+ # Move immediately to CPU so VRAM does not accumulate.
97
+ captured[idx] = x.detach().to(
98
+ device="cpu",
99
+ dtype=torch.float32
100
+ )
101
+
102
+ return hook
103
+
104
+
105
+ for layer_idx in TARGET_LAYERS:
106
+
107
+ layer_name = f"model.layers.{layer_idx}"
108
+
109
+ target_module = None
110
+
111
+ for name, module in transformer.named_modules():
112
+ if name == layer_name:
113
+ target_module = module
114
+ break
115
+
116
+ if target_module is None:
117
+ print("Layer not found:", layer_name)
118
+ continue
119
+
120
+ hooks[layer_idx] = target_module.register_forward_hook(
121
+ make_hook(layer_idx)
122
+ )
123
+
124
+
125
+ print("Registered hooks:", sorted(hooks.keys()))
126
+
127
+ if len(hooks) == 0:
128
+ raise RuntimeError("No hooks registered.")
129
+
130
+
131
+ # =========================================================
132
+ # RUN PROMPTS
133
+ # =========================================================
134
+
135
+ results = []
136
+
137
+ for i, prompt in enumerate(PROMPTS):
138
+
139
+ print()
140
+ print("=" * 80)
141
+ print(f"[{i + 1}/{len(PROMPTS)}]")
142
+ print(prompt)
143
+ print("=" * 80)
144
+
145
+ tokens = tokenizer(
146
+ prompt,
147
+ return_tensors="pt",
148
+ add_special_tokens=False,
149
+ )
150
+
151
+ input_ids = tokens["input_ids"].to("cuda")
152
+
153
+ print("Token count:", input_ids.shape[1])
154
+
155
+ captured.clear()
156
+
157
+ torch.cuda.empty_cache()
158
+
159
+ with torch.inference_mode():
160
+
161
+ # We do NOT request output_hidden_states.
162
+ # Hooks capture only selected layers.
163
+ out = transformer(
164
+ input_ids=input_ids,
165
+ )
166
+
167
+ item = {
168
+ "prompt": prompt,
169
+ "input_ids": input_ids.detach().cpu(),
170
+ "layers": {},
171
+ }
172
+
173
+ for idx in TARGET_LAYERS:
174
+
175
+ if idx in captured:
176
+ item["layers"][idx] = captured[idx]
177
+
178
+ print(
179
+ f"layer {idx:2d}: "
180
+ f"shape={tuple(captured[idx].shape)} "
181
+ f"dtype={captured[idx].dtype}"
182
+ )
183
+
184
+ else:
185
+ print(f"layer {idx:2d}: NOT CAPTURED")
186
+
187
+ results.append(item)
188
+
189
+ del out
190
+ del input_ids
191
+
192
+ torch.cuda.empty_cache()
193
+
194
+
195
+ # =========================================================
196
+ # REMOVE HOOKS
197
+ # =========================================================
198
+
199
+ for h in hooks.values():
200
+ h.remove()
201
+
202
+
203
+ # =========================================================
204
+ # VALIDATION
205
+ # =========================================================
206
+
207
+ print()
208
+ print("=" * 80)
209
+ print("VALIDATION")
210
+ print("=" * 80)
211
+
212
+ all_ok = True
213
+
214
+ for item in results:
215
+
216
+ print()
217
+ print(item["prompt"])
218
+
219
+ for idx in TARGET_LAYERS:
220
+
221
+ if idx not in item["layers"]:
222
+ print(f" layer {idx}: missing")
223
+ all_ok = False
224
+ continue
225
+
226
+ x = item["layers"][idx]
227
+
228
+ finite = torch.isfinite(x).all().item()
229
+
230
+ print(
231
+ f" layer {idx}: "
232
+ f"{tuple(x.shape)} "
233
+ f"finite={finite} "
234
+ f"mean={x.mean().item():.6f} "
235
+ f"std={x.std().item():.6f}"
236
+ )
237
+
238
+ if not finite:
239
+ all_ok = False
240
+
241
+
242
+ # =========================================================
243
+ # SAVE
244
+ # =========================================================
245
+
246
+ torch.save(
247
+ {
248
+ "prompts": PROMPTS,
249
+ "target_layers": TARGET_LAYERS,
250
+ "results": results,
251
+ },
252
+ OUT,
253
+ )
254
+
255
+ print()
256
+ print("=" * 80)
257
+
258
+ if all_ok:
259
+ print("SUCCESS")
260
+ else:
261
+ print("COMPLETED WITH WARNINGS")
262
+
263
+ print("Saved:")
264
+ print(OUT)
265
+
266
+ print("=" * 80)
research/raw_scripts/extract_h3_ood_hidden.py ADDED
@@ -0,0 +1,318 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import json
4
+ import gc
5
+ import torch
6
+
7
+ # ============================================================
8
+ # PATHS
9
+ # ============================================================
10
+
11
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
12
+
13
+ COMFY_ROOT = r"D:\ComfyUI_Python312\ComfyUI"
14
+
15
+ PROMPTS_FILE = os.path.join(
16
+ ROOT,
17
+ "bridge_ood_prompts_160.json"
18
+ )
19
+
20
+ OUT_FILE = os.path.join(
21
+ ROOT,
22
+ "h3_ood_hidden_160.pt"
23
+ )
24
+
25
+ H3_TE = (
26
+ r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders"
27
+ r"\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
28
+ )
29
+
30
+ LAYER = 49
31
+
32
+
33
+ # ============================================================
34
+ # COMFY IMPORT
35
+ # ============================================================
36
+
37
+ sys.path.insert(0, COMFY_ROOT)
38
+ os.chdir(COMFY_ROOT)
39
+
40
+ import comfy.sd
41
+ from transformers import AutoTokenizer
42
+
43
+
44
+ print("=" * 80)
45
+ print("MiniMax H3 OOD hidden-state extraction")
46
+ print("=" * 80)
47
+
48
+ print("Prompts:", PROMPTS_FILE)
49
+ print("Output :", OUT_FILE)
50
+ print("Layer :", LAYER)
51
+
52
+ if not torch.cuda.is_available():
53
+ raise RuntimeError("CUDA is required.")
54
+
55
+ print("GPU:", torch.cuda.get_device_name(0))
56
+
57
+
58
+ # ============================================================
59
+ # LOAD PROMPTS
60
+ # ============================================================
61
+
62
+ with open(
63
+ PROMPTS_FILE,
64
+ "r",
65
+ encoding="utf-8"
66
+ ) as f:
67
+ prompts = json.load(f)
68
+
69
+ print("Prompt count:", len(prompts))
70
+
71
+ if len(prompts) != 160:
72
+ raise RuntimeError(
73
+ f"Expected 160 prompts, got {len(prompts)}"
74
+ )
75
+
76
+
77
+ # ============================================================
78
+ # TOKENIZER
79
+ # ============================================================
80
+
81
+ print("\nLoading tokenizer...")
82
+
83
+ tokenizer = AutoTokenizer.from_pretrained(
84
+ "Qwen/Qwen3-VL-32B-Instruct",
85
+ trust_remote_code=True,
86
+ )
87
+
88
+
89
+ # ============================================================
90
+ # LOAD H3 TEXT ENCODER
91
+ # ============================================================
92
+
93
+ print("\nLoading H3 text encoder through ComfyUI...")
94
+
95
+ clip = comfy.sd.load_clip(
96
+ ckpt_paths=[H3_TE],
97
+ embedding_directory=None,
98
+ clip_type=comfy.sd.CLIPType.MINIMAX,
99
+ model_options={
100
+ "load_device": torch.device("cuda"),
101
+ "offload_device": torch.device("cpu"),
102
+ },
103
+ )
104
+
105
+ print("CLIP loaded.")
106
+
107
+
108
+ # ============================================================
109
+ # FIND TRANSFORMER
110
+ # ============================================================
111
+
112
+ root = clip.cond_stage_model
113
+
114
+ transformer = None
115
+ transformer_name = None
116
+
117
+ for name, module in root.named_modules():
118
+
119
+ if name.endswith("qwen3vl_32b.transformer"):
120
+ transformer = module
121
+ transformer_name = name
122
+ break
123
+
124
+ if transformer is None:
125
+
126
+ print("\nCould not find exact transformer.")
127
+ print("Transformer-like modules:")
128
+
129
+ for name, module in root.named_modules():
130
+ if "transformer" in name.lower():
131
+ print(name, type(module))
132
+
133
+ raise RuntimeError("H3 transformer not found.")
134
+
135
+ print("Transformer:", transformer_name)
136
+
137
+
138
+ # ============================================================
139
+ # FIND TARGET LAYER
140
+ # ============================================================
141
+
142
+ target_layer = None
143
+ target_name = f"model.layers.{LAYER}"
144
+
145
+ for name, module in transformer.named_modules():
146
+
147
+ if name == target_name:
148
+ target_layer = module
149
+ break
150
+
151
+ if target_layer is None:
152
+ raise RuntimeError(
153
+ f"H3 layer {LAYER} not found."
154
+ )
155
+
156
+ print("Target layer:", target_name)
157
+
158
+
159
+ # ============================================================
160
+ # HOOK
161
+ # ============================================================
162
+
163
+ captured = {}
164
+
165
+
166
+ def hook_fn(module, inputs, output):
167
+
168
+ if isinstance(output, tuple):
169
+ hidden = output[0]
170
+ else:
171
+ hidden = output
172
+
173
+ if not torch.is_tensor(hidden):
174
+ raise RuntimeError(
175
+ f"Unexpected H3 layer output type: {type(hidden)}"
176
+ )
177
+
178
+ captured["hidden"] = (
179
+ hidden.detach()
180
+ .to(
181
+ device="cpu",
182
+ dtype=torch.float16
183
+ )
184
+ .clone()
185
+ )
186
+
187
+
188
+ handle = target_layer.register_forward_hook(
189
+ hook_fn
190
+ )
191
+
192
+
193
+ # ============================================================
194
+ # EXTRACT
195
+ # ============================================================
196
+
197
+ results = []
198
+
199
+ print("\nStarting extraction...\n")
200
+
201
+ with torch.inference_mode():
202
+
203
+ for index, item in enumerate(prompts):
204
+
205
+ prompt = item["prompt"]
206
+ category = item["category"]
207
+
208
+ encoded = tokenizer(
209
+ prompt,
210
+ return_tensors="pt",
211
+ add_special_tokens=False,
212
+ )
213
+
214
+ input_ids_cpu = (
215
+ encoded["input_ids"]
216
+ .detach()
217
+ .cpu()
218
+ )
219
+
220
+ input_ids = (
221
+ encoded["input_ids"]
222
+ .to("cuda")
223
+ )
224
+
225
+ captured.clear()
226
+
227
+ output = transformer(
228
+ input_ids=input_ids
229
+ )
230
+
231
+ if "hidden" not in captured:
232
+ raise RuntimeError(
233
+ f"Layer hook failed at prompt {index}"
234
+ )
235
+
236
+ hidden = captured["hidden"]
237
+
238
+ if hidden.ndim != 3:
239
+ raise RuntimeError(
240
+ f"Unexpected hidden shape: {hidden.shape}"
241
+ )
242
+
243
+ if hidden.shape[-1] != 5120:
244
+ raise RuntimeError(
245
+ f"Expected H3 hidden dim 5120, "
246
+ f"got {hidden.shape[-1]}"
247
+ )
248
+
249
+ if not torch.isfinite(
250
+ hidden.float()
251
+ ).all():
252
+ raise RuntimeError(
253
+ f"Non-finite H3 state at prompt {index}"
254
+ )
255
+
256
+ results.append({
257
+ "prompt": prompt,
258
+ "category": category,
259
+ "input_ids": input_ids_cpu,
260
+ "hidden": hidden,
261
+ })
262
+
263
+ if (
264
+ index == 0
265
+ or (index + 1) % 10 == 0
266
+ or index + 1 == len(prompts)
267
+ ):
268
+ print(
269
+ f"{index + 1:3d}/{len(prompts)} | "
270
+ f"tokens={hidden.shape[1]:3d} | "
271
+ f"shape={tuple(hidden.shape)}"
272
+ )
273
+
274
+ del output
275
+ del input_ids
276
+
277
+ if (index + 1) % 20 == 0:
278
+ gc.collect()
279
+ torch.cuda.empty_cache()
280
+
281
+
282
+ # ============================================================
283
+ # CLEANUP
284
+ # ============================================================
285
+
286
+ handle.remove()
287
+
288
+
289
+ # ============================================================
290
+ # SAVE
291
+ # ============================================================
292
+
293
+ payload = {
294
+ "model": "MiniMax H3 Qwen3-VL TE",
295
+ "layer": LAYER,
296
+ "hidden_dim": 5120,
297
+ "prompt_file": os.path.basename(
298
+ PROMPTS_FILE
299
+ ),
300
+ "results": results,
301
+ }
302
+
303
+ torch.save(
304
+ payload,
305
+ OUT_FILE
306
+ )
307
+
308
+
309
+ print()
310
+ print("=" * 80)
311
+ print("SUCCESS")
312
+ print("=" * 80)
313
+
314
+ print("Prompts:", len(results))
315
+ print("Layer :", LAYER)
316
+
317
+ print("Output:")
318
+ print(OUT_FILE)
research/raw_scripts/extract_sensenova_bridge_dataset.py ADDED
@@ -0,0 +1,425 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import json
5
+ import shutil
6
+ import torch
7
+
8
+
9
+ # ============================================================
10
+ # PATHS
11
+ # ============================================================
12
+
13
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
14
+
15
+ SENSENOVA_SRC = os.path.join(
16
+ ROOT,
17
+ "SenseNova-U1",
18
+ "src"
19
+ )
20
+
21
+ COMPAT_FILE = os.path.join(
22
+ SENSENOVA_SRC,
23
+ "sensenova_u1",
24
+ "models",
25
+ "neo_unify",
26
+ "transformers_compat.py"
27
+ )
28
+
29
+ MODEL_FILE = (
30
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
31
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
32
+ )
33
+
34
+ PROMPT_FILE = os.path.join(
35
+ ROOT,
36
+ "bridge_prompts_480.json"
37
+ )
38
+
39
+ OUT = os.path.join(
40
+ ROOT,
41
+ "sensenova_bridge_hidden_480.pt"
42
+ )
43
+
44
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
45
+
46
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
47
+
48
+
49
+ # ============================================================
50
+ # PATCH TRANSFORMERS 4.57 COMPAT
51
+ # ============================================================
52
+
53
+ print("=" * 80)
54
+ print("SenseNova U1.5 bridge dataset extractor")
55
+ print("=" * 80)
56
+
57
+ print("\nChecking Transformers compatibility patch...")
58
+
59
+ with open(COMPAT_FILE, "r", encoding="utf-8") as f:
60
+ text = f.read()
61
+
62
+ OLD = (
63
+ " from transformers.utils.generic import "
64
+ "check_model_inputs as model_input_compat"
65
+ )
66
+
67
+ NEW = (
68
+ " from transformers.utils.generic import check_model_inputs\n"
69
+ " def model_input_compat(func):\n"
70
+ " return check_model_inputs()(func)"
71
+ )
72
+
73
+ if OLD in text:
74
+
75
+ backup = COMPAT_FILE + ".bak"
76
+
77
+ if not os.path.exists(backup):
78
+ shutil.copy2(COMPAT_FILE, backup)
79
+
80
+ text = text.replace(OLD, NEW)
81
+
82
+ with open(COMPAT_FILE, "w", encoding="utf-8") as f:
83
+ f.write(text)
84
+
85
+ print("Patch applied.")
86
+
87
+ else:
88
+ print("Patch already present / not needed.")
89
+
90
+
91
+ # ============================================================
92
+ # IMPORTS
93
+ # ============================================================
94
+
95
+ sys.path.insert(0, SENSENOVA_SRC)
96
+
97
+ from transformers import AutoTokenizer
98
+ from accelerate import init_empty_weights, load_checkpoint_and_dispatch
99
+
100
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
101
+ NEOChatConfig,
102
+ )
103
+
104
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
105
+ NEOChatModel,
106
+ )
107
+
108
+
109
+ # ============================================================
110
+ # LOAD PROMPTS
111
+ # ============================================================
112
+
113
+ print("\nLoading prompt dataset...")
114
+
115
+ with open(
116
+ PROMPT_FILE,
117
+ "r",
118
+ encoding="utf-8"
119
+ ) as f:
120
+
121
+ dataset = json.load(f)
122
+
123
+ print("Prompts:", len(dataset))
124
+
125
+
126
+ # ============================================================
127
+ # TOKENIZER
128
+ # ============================================================
129
+
130
+ print("\nLoading tokenizer...")
131
+
132
+ tokenizer = AutoTokenizer.from_pretrained(
133
+ HF_REPO,
134
+ trust_remote_code=True,
135
+ )
136
+
137
+
138
+ # ============================================================
139
+ # CONFIG
140
+ # ============================================================
141
+
142
+ print("Loading SenseNova config...")
143
+
144
+ config = NEOChatConfig.from_pretrained(
145
+ HF_REPO
146
+ )
147
+
148
+ config.llm_config._attn_implementation = "eager"
149
+
150
+ print("Hidden size:", config.llm_config.hidden_size)
151
+ print("Layers:", config.llm_config.num_hidden_layers)
152
+
153
+
154
+ # ============================================================
155
+ # META ARCHITECTURE
156
+ # ============================================================
157
+
158
+ print("\nCreating META architecture...")
159
+
160
+ with init_empty_weights():
161
+
162
+ model = NEOChatModel(config)
163
+
164
+ print("META architecture created.")
165
+
166
+
167
+ # ============================================================
168
+ # LOAD CHECKPOINT
169
+ # ============================================================
170
+
171
+ print("\nLoading BF16 checkpoint with Accelerate...")
172
+
173
+ model = load_checkpoint_and_dispatch(
174
+ model,
175
+ checkpoint=MODEL_FILE,
176
+ device_map="auto",
177
+ max_memory={
178
+ 0: "19GiB",
179
+ "cpu": "70GiB",
180
+ },
181
+ dtype=torch.bfloat16,
182
+ no_split_module_classes=[
183
+ "Qwen3DecoderLayer",
184
+ "Qwen3MoeDecoderLayer",
185
+ ],
186
+ )
187
+
188
+ model.eval()
189
+
190
+ print("Checkpoint loaded.")
191
+
192
+
193
+ # ============================================================
194
+ # LANGUAGE MODEL
195
+ # ============================================================
196
+
197
+ language_model = model.language_model
198
+ qwen = language_model.model
199
+ layers = qwen.layers
200
+
201
+ print("\nDetected language layers:", len(layers))
202
+
203
+
204
+ # ============================================================
205
+ # INPUT DEVICE
206
+ # ============================================================
207
+
208
+ device_map = getattr(model, "hf_device_map", {})
209
+
210
+ INPUT_DEVICE = torch.device("cuda:0")
211
+
212
+ for name, dev in device_map.items():
213
+
214
+ if name.endswith("language_model.model.embed_tokens"):
215
+
216
+ if str(dev) == "cpu":
217
+ INPUT_DEVICE = torch.device("cpu")
218
+ else:
219
+ INPUT_DEVICE = torch.device("cuda:0")
220
+
221
+ break
222
+
223
+ print("Input device:", INPUT_DEVICE)
224
+
225
+
226
+ # ============================================================
227
+ # HOOKS
228
+ # ============================================================
229
+
230
+ captured = {}
231
+ hooks = {}
232
+
233
+
234
+ def make_hook(idx):
235
+
236
+ def hook(module, inputs, output):
237
+
238
+ if isinstance(output, tuple):
239
+ x = output[0]
240
+ else:
241
+ x = output
242
+
243
+ if not torch.is_tensor(x):
244
+ raise RuntimeError(
245
+ f"Unexpected output type at layer {idx}: {type(x)}"
246
+ )
247
+
248
+ # float16 saves a lot of disk space.
249
+ captured[idx] = (
250
+ x.detach()
251
+ .to(
252
+ device="cpu",
253
+ dtype=torch.float16
254
+ )
255
+ )
256
+
257
+ return hook
258
+
259
+
260
+ print("\nRegistering hooks...")
261
+
262
+ for idx in TARGET_LAYERS:
263
+
264
+ hooks[idx] = layers[idx].register_forward_hook(
265
+ make_hook(idx)
266
+ )
267
+
268
+ print(" layer", idx)
269
+
270
+
271
+ # ============================================================
272
+ # EXTRACT
273
+ # ============================================================
274
+
275
+ results = []
276
+
277
+ print()
278
+ print("=" * 80)
279
+ print("EXTRACTING")
280
+ print("=" * 80)
281
+
282
+ for i, row in enumerate(dataset):
283
+
284
+ prompt = row["prompt"]
285
+
286
+ encoded = tokenizer(
287
+ prompt,
288
+ return_tensors="pt",
289
+ add_special_tokens=False,
290
+ )
291
+
292
+ input_ids = encoded["input_ids"].to(
293
+ INPUT_DEVICE
294
+ )
295
+
296
+ captured.clear()
297
+
298
+ # SenseNova keeps an internal index for generation paths.
299
+ # Reset it for every independent prompt.
300
+ qwen.current_index = -1
301
+
302
+ with torch.inference_mode():
303
+
304
+ outputs = qwen(
305
+ input_ids=input_ids,
306
+ image_gen_indicators=None,
307
+ attention_mask=None,
308
+ use_cache=False,
309
+ )
310
+
311
+ item = {
312
+ "prompt": prompt,
313
+ "category": row["category"],
314
+ "input_ids": (
315
+ input_ids.detach()
316
+ .cpu()
317
+ ),
318
+ "layers": {},
319
+ }
320
+
321
+ for idx in TARGET_LAYERS:
322
+
323
+ if idx not in captured:
324
+
325
+ raise RuntimeError(
326
+ f"Missing SenseNova layer {idx} "
327
+ f"at prompt {i}"
328
+ )
329
+
330
+ item["layers"][idx] = (
331
+ captured[idx]
332
+ )
333
+
334
+ results.append(item)
335
+
336
+ del outputs
337
+ del input_ids
338
+
339
+ if (i + 1) % 10 == 0:
340
+
341
+ print(
342
+ f"{i + 1:4d}/"
343
+ f"{len(dataset)}"
344
+ )
345
+
346
+ if (i + 1) % 20 == 0:
347
+
348
+ gc.collect()
349
+ torch.cuda.empty_cache()
350
+
351
+
352
+ # ============================================================
353
+ # REMOVE HOOKS
354
+ # ============================================================
355
+
356
+ for h in hooks.values():
357
+ h.remove()
358
+
359
+
360
+ # ============================================================
361
+ # FINAL VALIDATION
362
+ # ============================================================
363
+
364
+ print()
365
+ print("=" * 80)
366
+ print("VALIDATION")
367
+ print("=" * 80)
368
+
369
+ if len(results) != len(dataset):
370
+
371
+ raise RuntimeError(
372
+ "Result count mismatch."
373
+ )
374
+
375
+ for i, item in enumerate(results):
376
+
377
+ for idx in TARGET_LAYERS:
378
+
379
+ x = item["layers"][idx]
380
+
381
+ if x.ndim != 3:
382
+
383
+ raise RuntimeError(
384
+ f"Bad ndim at prompt {i}, layer {idx}"
385
+ )
386
+
387
+ if x.shape[-1] != 4096:
388
+
389
+ raise RuntimeError(
390
+ f"Bad hidden size at prompt {i}, layer {idx}: "
391
+ f"{x.shape}"
392
+ )
393
+
394
+ if not torch.isfinite(x).all():
395
+
396
+ raise RuntimeError(
397
+ f"Non-finite values at prompt {i}, layer {idx}"
398
+ )
399
+
400
+ print("All hidden states valid.")
401
+
402
+
403
+ # ============================================================
404
+ # SAVE
405
+ # ============================================================
406
+
407
+ print("\nSaving...")
408
+
409
+ torch.save(
410
+ {
411
+ "target_layers": TARGET_LAYERS,
412
+ "results": results,
413
+ "source_model": MODEL_FILE,
414
+ },
415
+ OUT,
416
+ )
417
+
418
+ print()
419
+ print("=" * 80)
420
+ print("SUCCESS")
421
+ print(OUT)
422
+ print("=" * 80)
423
+
424
+ gc.collect()
425
+ torch.cuda.empty_cache()
research/raw_scripts/extract_sensenova_hidden_states.py ADDED
@@ -0,0 +1,321 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import torch
3
+
4
+ from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
5
+ from safetensors import safe_open
6
+
7
+
8
+ MODEL_FILE = r"D:\ComfyUI_Python312\ComfyUI\models\unet\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
9
+
10
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\sensenova_hidden_states_fast.pt"
11
+
12
+ REPO = "sensenova/SenseNova-U1.5-8B-MoT"
13
+
14
+ PROMPTS = [
15
+ "a red cube on a blue sphere",
16
+ "a woman holding a transparent glass bottle",
17
+ "three people standing behind a wooden table",
18
+ ]
19
+
20
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
21
+
22
+
23
+ print("=" * 80)
24
+ print("SenseNova hidden-state extractor")
25
+ print("=" * 80)
26
+
27
+ print("\nLoading tokenizer...")
28
+
29
+ tokenizer = AutoTokenizer.from_pretrained(
30
+ REPO,
31
+ trust_remote_code=True,
32
+ )
33
+
34
+
35
+ # ============================================================
36
+ # CONFIG
37
+ # ============================================================
38
+
39
+ print("Loading SenseNova config...")
40
+
41
+ config = AutoConfig.from_pretrained(
42
+ REPO,
43
+ trust_remote_code=True,
44
+ )
45
+
46
+ print("Config type:", type(config))
47
+
48
+
49
+ # ============================================================
50
+ # LOAD MODEL
51
+ #
52
+ # We create the official architecture, then load YOUR local
53
+ # pruned BF16 checkpoint.
54
+ # ============================================================
55
+
56
+ print("\nCreating SenseNova model architecture...")
57
+
58
+ model = AutoModelForCausalLM.from_config(
59
+ config,
60
+ trust_remote_code=True,
61
+ )
62
+
63
+ print("Architecture created.")
64
+
65
+
66
+ # ============================================================
67
+ # LOAD LOCAL SAFETENSORS
68
+ # ============================================================
69
+
70
+ print("\nLoading local checkpoint:")
71
+ print(MODEL_FILE)
72
+
73
+ state = {}
74
+
75
+ with safe_open(MODEL_FILE, framework="pt", device="cpu") as f:
76
+
77
+ keys = list(f.keys())
78
+
79
+ # Only language-model tensors are needed for this test.
80
+ lm_keys = [
81
+ k for k in keys
82
+ if k.startswith("language_model.")
83
+ ]
84
+
85
+ print("Language-model tensors:", len(lm_keys))
86
+
87
+ for i, key in enumerate(lm_keys):
88
+
89
+ state[key] = f.get_tensor(key)
90
+
91
+ if (i + 1) % 100 == 0:
92
+ print(f"Loaded {i + 1}/{len(lm_keys)}")
93
+
94
+
95
+ print("\nApplying state dict...")
96
+
97
+ missing, unexpected = model.load_state_dict(
98
+ state,
99
+ strict=False,
100
+ )
101
+
102
+ print("Missing keys:", len(missing))
103
+ print("Unexpected keys:", len(unexpected))
104
+
105
+ del state
106
+
107
+
108
+ # ============================================================
109
+ # FIND LANGUAGE MODEL
110
+ # ============================================================
111
+
112
+ language_model = None
113
+
114
+ candidates = [
115
+ "language_model",
116
+ "model.language_model",
117
+ ]
118
+
119
+ for path in candidates:
120
+
121
+ obj = model
122
+
123
+ try:
124
+ for part in path.split("."):
125
+ obj = getattr(obj, part)
126
+
127
+ language_model = obj
128
+ print("Found language model:", path)
129
+ break
130
+
131
+ except Exception:
132
+ pass
133
+
134
+
135
+ if language_model is None:
136
+
137
+ print("\nCould not find language model directly.")
138
+ print("Candidate modules:")
139
+
140
+ for name, module in model.named_modules():
141
+ if "language_model" in name.lower():
142
+ print(name, type(module))
143
+
144
+ raise RuntimeError("language_model not found")
145
+
146
+
147
+ # ============================================================
148
+ # MOVE / OFFLOAD
149
+ # ============================================================
150
+
151
+ print("\nPreparing model...")
152
+
153
+ model.eval()
154
+
155
+ # Keep whole architecture on CPU initially.
156
+ model.to("cpu")
157
+
158
+ # We will let PyTorch move the language model to GPU if possible.
159
+ # If this raises OOM, we will switch to layer-by-layer streaming.
160
+ try:
161
+
162
+ language_model.to("cuda")
163
+
164
+ DEVICE = "cuda"
165
+
166
+ print("Language model moved to CUDA.")
167
+
168
+ except torch.cuda.OutOfMemoryError:
169
+
170
+ print("Language model does not fit in VRAM.")
171
+ print("Falling back to CPU.")
172
+
173
+ torch.cuda.empty_cache()
174
+
175
+ DEVICE = "cpu"
176
+
177
+
178
+ # ============================================================
179
+ # HOOKS
180
+ # ============================================================
181
+
182
+ captured = {}
183
+ hooks = {}
184
+
185
+
186
+ def make_hook(idx):
187
+
188
+ def hook(module, inputs, output):
189
+
190
+ if isinstance(output, tuple):
191
+ x = output[0]
192
+ else:
193
+ x = output
194
+
195
+ captured[idx] = (
196
+ x.detach()
197
+ .float()
198
+ .cpu()
199
+ )
200
+
201
+ return hook
202
+
203
+
204
+ for idx in TARGET_LAYERS:
205
+
206
+ target = None
207
+
208
+ expected_suffix = f"layers.{idx}"
209
+
210
+ for name, module in language_model.named_modules():
211
+
212
+ if name.endswith(expected_suffix):
213
+ target = module
214
+ print("Hook:", idx, "->", name)
215
+ break
216
+
217
+ if target is not None:
218
+ hooks[idx] = target.register_forward_hook(
219
+ make_hook(idx)
220
+ )
221
+
222
+ else:
223
+ print("Layer not found:", idx)
224
+
225
+
226
+ print("\nRegistered hooks:", sorted(hooks.keys()))
227
+
228
+
229
+ # ============================================================
230
+ # RUN
231
+ # ============================================================
232
+
233
+ results = []
234
+
235
+ for i, prompt in enumerate(PROMPTS):
236
+
237
+ print()
238
+ print("=" * 80)
239
+ print(f"[{i + 1}/{len(PROMPTS)}]")
240
+ print(prompt)
241
+ print("=" * 80)
242
+
243
+ encoded = tokenizer(
244
+ prompt,
245
+ return_tensors="pt",
246
+ add_special_tokens=False,
247
+ )
248
+
249
+ input_ids = encoded["input_ids"].to(DEVICE)
250
+
251
+ print("Token count:", input_ids.shape[1])
252
+
253
+ captured.clear()
254
+
255
+ with torch.inference_mode():
256
+
257
+ output = language_model(
258
+ input_ids=input_ids,
259
+ use_cache=False,
260
+ )
261
+
262
+ item = {
263
+ "prompt": prompt,
264
+ "input_ids": input_ids.cpu(),
265
+ "layers": {},
266
+ }
267
+
268
+ for idx in TARGET_LAYERS:
269
+
270
+ if idx in captured:
271
+
272
+ x = captured[idx]
273
+
274
+ item["layers"][idx] = x
275
+
276
+ print(
277
+ f"layer {idx:2d}: "
278
+ f"shape={tuple(x.shape)} "
279
+ f"mean={x.mean().item():.6f} "
280
+ f"std={x.std().item():.6f} "
281
+ f"finite={torch.isfinite(x).all().item()}"
282
+ )
283
+
284
+ else:
285
+
286
+ print(
287
+ f"layer {idx:2d}: NOT CAPTURED"
288
+ )
289
+
290
+ results.append(item)
291
+
292
+ del output
293
+
294
+ if DEVICE == "cuda":
295
+ torch.cuda.empty_cache()
296
+
297
+
298
+ # ============================================================
299
+ # CLEANUP
300
+ # ============================================================
301
+
302
+ for hook in hooks.values():
303
+ hook.remove()
304
+
305
+
306
+ torch.save(
307
+ {
308
+ "prompts": PROMPTS,
309
+ "target_layers": TARGET_LAYERS,
310
+ "results": results,
311
+ },
312
+ OUT,
313
+ )
314
+
315
+
316
+ print()
317
+ print("=" * 80)
318
+ print("SUCCESS")
319
+ print("Saved:")
320
+ print(OUT)
321
+ print("=" * 80)
research/raw_scripts/extract_sensenova_hidden_states_local.py ADDED
@@ -0,0 +1,421 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import torch
4
+ from safetensors import safe_open
5
+ from transformers import AutoTokenizer
6
+
7
+ # ============================================================
8
+ # PATHS
9
+ # ============================================================
10
+
11
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
12
+ SENSENOVA_REPO = os.path.join(ROOT, "SenseNova-U1")
13
+ SENSENOVA_SRC = os.path.join(SENSENOVA_REPO, "src")
14
+
15
+ MODEL_FILE = r"D:\ComfyUI_Python312\ComfyUI\models\unet\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
16
+
17
+ OUT = os.path.join(
18
+ ROOT,
19
+ "sensenova_hidden_states_fast.pt"
20
+ )
21
+
22
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
23
+
24
+ PROMPTS = [
25
+ "a red cube on a blue sphere",
26
+ "a woman holding a transparent glass bottle",
27
+ "three people standing behind a wooden table",
28
+ ]
29
+
30
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
31
+
32
+
33
+ # ============================================================
34
+ # LOCAL SENSENOVA CODE
35
+ # ============================================================
36
+
37
+ sys.path.insert(0, SENSENOVA_SRC)
38
+
39
+ print("=" * 80)
40
+ print("SenseNova U1.5 hidden-state extractor - LOCAL CODE")
41
+ print("=" * 80)
42
+
43
+ print("\nSenseNova src:")
44
+ print(SENSENOVA_SRC)
45
+
46
+
47
+ # ============================================================
48
+ # IMPORT CUSTOM CLASSES
49
+ # ============================================================
50
+
51
+ print("\nImporting local SenseNova classes...")
52
+
53
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import NEOChatConfig
54
+
55
+ try:
56
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import NEOChatModel
57
+ MODEL_CLASS = NEOChatModel
58
+ print("Using model class: NEOChatModel")
59
+
60
+ except ImportError:
61
+ try:
62
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import NEOChatForCausalLM
63
+ MODEL_CLASS = NEOChatForCausalLM
64
+ print("Using model class: NEOChatForCausalLM")
65
+
66
+ except ImportError:
67
+ print("\nCould not import expected model class.")
68
+ print("Available names containing 'Chat' or 'Model':")
69
+
70
+ import sensenova_u1.models.neo_unify.modeling_neo_chat as mod
71
+
72
+ for name in dir(mod):
73
+ if "Chat" in name or "Model" in name:
74
+ print(" ", name)
75
+
76
+ raise
77
+
78
+
79
+ # ============================================================
80
+ # TOKENIZER
81
+ # ============================================================
82
+
83
+ print("\nLoading tokenizer...")
84
+
85
+ tokenizer = AutoTokenizer.from_pretrained(
86
+ HF_REPO,
87
+ trust_remote_code=True,
88
+ )
89
+
90
+
91
+ # ============================================================
92
+ # CONFIG
93
+ # ============================================================
94
+
95
+ print("Loading config using LOCAL NEOChatConfig...")
96
+
97
+ config = NEOChatConfig.from_pretrained(
98
+ HF_REPO,
99
+ )
100
+
101
+ print("Config loaded.")
102
+ print("Config class:", type(config))
103
+
104
+
105
+ # ============================================================
106
+ # CREATE MODEL ARCHITECTURE
107
+ # ============================================================
108
+
109
+ print("\nCreating SenseNova architecture...")
110
+
111
+ model = MODEL_CLASS(config)
112
+
113
+ print("Architecture created.")
114
+ print("Model type:", type(model))
115
+
116
+
117
+ # ============================================================
118
+ # LOAD ONLY LANGUAGE MODEL WEIGHTS
119
+ # ============================================================
120
+
121
+ print("\nScanning local checkpoint...")
122
+
123
+ with safe_open(MODEL_FILE, framework="pt", device="cpu") as f:
124
+ keys = list(f.keys())
125
+
126
+ lm_keys = [
127
+ k for k in keys
128
+ if k.startswith("language_model.")
129
+ ]
130
+
131
+ print("Total checkpoint tensors:", len(keys))
132
+ print("Language-model tensors:", len(lm_keys))
133
+
134
+
135
+ print("\nLoading language-model tensors...")
136
+
137
+ state = {}
138
+
139
+ with safe_open(MODEL_FILE, framework="pt", device="cpu") as f:
140
+
141
+ for i, key in enumerate(lm_keys):
142
+ state[key] = f.get_tensor(key)
143
+
144
+ if (i + 1) % 100 == 0:
145
+ print(
146
+ f" loaded {i + 1}/{len(lm_keys)}"
147
+ )
148
+
149
+
150
+ # ============================================================
151
+ # APPLY
152
+ # ============================================================
153
+
154
+ print("\nApplying state dict...")
155
+
156
+ result = model.load_state_dict(
157
+ state,
158
+ strict=False,
159
+ )
160
+
161
+ print("Missing keys:", len(result.missing_keys))
162
+ print("Unexpected keys:", len(result.unexpected_keys))
163
+
164
+ if len(result.unexpected_keys) > 0:
165
+ print("\nFirst unexpected keys:")
166
+ for k in result.unexpected_keys[:20]:
167
+ print(" ", k)
168
+
169
+ del state
170
+
171
+
172
+ # ============================================================
173
+ # FIND LANGUAGE MODEL
174
+ # ============================================================
175
+
176
+ print("\nFinding language model...")
177
+
178
+ language_model = None
179
+ language_model_name = None
180
+
181
+ possible_paths = [
182
+ "language_model",
183
+ "model.language_model",
184
+ "language_model.model",
185
+ ]
186
+
187
+ for path in possible_paths:
188
+
189
+ obj = model
190
+
191
+ try:
192
+ for part in path.split("."):
193
+ obj = getattr(obj, part)
194
+
195
+ language_model = obj
196
+ language_model_name = path
197
+ break
198
+
199
+ except Exception:
200
+ pass
201
+
202
+
203
+ if language_model is None:
204
+
205
+ print("Could not locate language model automatically.")
206
+ print("\nCandidate modules:")
207
+
208
+ for name, module in model.named_modules():
209
+ if "language_model" in name.lower():
210
+ print(name, type(module))
211
+
212
+ raise RuntimeError("language_model not found")
213
+
214
+
215
+ print("Language model:", language_model_name)
216
+ print("Type:", type(language_model))
217
+
218
+
219
+ # ============================================================
220
+ # FIND TRANSFORMER LAYERS
221
+ # ============================================================
222
+
223
+ print("\nSearching for transformer layers...")
224
+
225
+ for name, module in language_model.named_modules():
226
+
227
+ if name.endswith("layers.0"):
228
+ print("Example layer path:", name)
229
+ break
230
+
231
+
232
+ # ============================================================
233
+ # CUDA
234
+ # ============================================================
235
+
236
+ model.eval()
237
+
238
+ print("\nTrying language model on CUDA...")
239
+
240
+ DEVICE = "cuda"
241
+
242
+ try:
243
+
244
+ language_model.to(
245
+ device="cuda",
246
+ dtype=torch.bfloat16,
247
+ )
248
+
249
+ print("Language model moved to CUDA.")
250
+
251
+ except torch.cuda.OutOfMemoryError:
252
+
253
+ print("\nCUDA OOM while moving language model.")
254
+ print("Falling back to CPU for diagnostic run.")
255
+
256
+ torch.cuda.empty_cache()
257
+
258
+ language_model.to("cpu")
259
+
260
+ DEVICE = "cpu"
261
+
262
+
263
+ # ============================================================
264
+ # HOOKS
265
+ # ============================================================
266
+
267
+ captured = {}
268
+ hooks = {}
269
+
270
+
271
+ def make_hook(idx):
272
+
273
+ def hook(module, inputs, output):
274
+
275
+ if isinstance(output, tuple):
276
+ x = output[0]
277
+ else:
278
+ x = output
279
+
280
+ captured[idx] = (
281
+ x.detach()
282
+ .float()
283
+ .cpu()
284
+ )
285
+
286
+ return hook
287
+
288
+
289
+ print("\nRegistering hooks...")
290
+
291
+ for idx in TARGET_LAYERS:
292
+
293
+ target = None
294
+ target_name = None
295
+
296
+ suffix = f"layers.{idx}"
297
+
298
+ for name, module in language_model.named_modules():
299
+
300
+ if name.endswith(suffix):
301
+ target = module
302
+ target_name = name
303
+ break
304
+
305
+ if target is None:
306
+ print(f" layer {idx}: NOT FOUND")
307
+ continue
308
+
309
+ hooks[idx] = target.register_forward_hook(
310
+ make_hook(idx)
311
+ )
312
+
313
+ print(
314
+ f" layer {idx}: {target_name}"
315
+ )
316
+
317
+
318
+ print("\nRegistered hooks:", sorted(hooks.keys()))
319
+
320
+ if not hooks:
321
+ raise RuntimeError("No transformer layers found")
322
+
323
+
324
+ # ============================================================
325
+ # INFERENCE
326
+ # ============================================================
327
+
328
+ results = []
329
+
330
+ for i, prompt in enumerate(PROMPTS):
331
+
332
+ print()
333
+ print("=" * 80)
334
+ print(f"[{i+1}/{len(PROMPTS)}]")
335
+ print(prompt)
336
+ print("=" * 80)
337
+
338
+ encoded = tokenizer(
339
+ prompt,
340
+ return_tensors="pt",
341
+ add_special_tokens=False,
342
+ )
343
+
344
+ input_ids = encoded["input_ids"].to(DEVICE)
345
+
346
+ print("Token count:", input_ids.shape[1])
347
+
348
+ captured.clear()
349
+
350
+ with torch.inference_mode():
351
+
352
+ try:
353
+ out = language_model(
354
+ input_ids=input_ids,
355
+ use_cache=False,
356
+ )
357
+
358
+ except TypeError:
359
+
360
+ out = language_model(
361
+ input_ids=input_ids,
362
+ )
363
+
364
+ item = {
365
+ "prompt": prompt,
366
+ "input_ids": input_ids.detach().cpu(),
367
+ "layers": {},
368
+ }
369
+
370
+ for idx in TARGET_LAYERS:
371
+
372
+ if idx not in captured:
373
+ print(
374
+ f"layer {idx:2d}: NOT CAPTURED"
375
+ )
376
+ continue
377
+
378
+ x = captured[idx]
379
+
380
+ item["layers"][idx] = x
381
+
382
+ print(
383
+ f"layer {idx:2d}: "
384
+ f"shape={tuple(x.shape)} "
385
+ f"mean={x.mean().item():.6f} "
386
+ f"std={x.std().item():.6f} "
387
+ f"finite={torch.isfinite(x).all().item()}"
388
+ )
389
+
390
+ results.append(item)
391
+
392
+ del out
393
+ del input_ids
394
+
395
+ if DEVICE == "cuda":
396
+ torch.cuda.empty_cache()
397
+
398
+
399
+ # ============================================================
400
+ # SAVE
401
+ # ============================================================
402
+
403
+ for h in hooks.values():
404
+ h.remove()
405
+
406
+
407
+ torch.save(
408
+ {
409
+ "prompts": PROMPTS,
410
+ "target_layers": TARGET_LAYERS,
411
+ "results": results,
412
+ },
413
+ OUT,
414
+ )
415
+
416
+ print()
417
+ print("=" * 80)
418
+ print("SUCCESS")
419
+ print("Saved:")
420
+ print(OUT)
421
+ print("=" * 80)
research/raw_scripts/extract_sensenova_hidden_states_stream.py ADDED
@@ -0,0 +1,352 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import torch
5
+
6
+ from transformers import AutoTokenizer
7
+ from accelerate import init_empty_weights, load_checkpoint_and_dispatch
8
+
9
+ # ============================================================
10
+ # PATHS
11
+ # ============================================================
12
+
13
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
14
+
15
+ SENSENOVA_SRC = os.path.join(
16
+ ROOT,
17
+ "SenseNova-U1",
18
+ "src"
19
+ )
20
+
21
+ MODEL_FILE = (
22
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
23
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
24
+ )
25
+
26
+ OUT = os.path.join(
27
+ ROOT,
28
+ "sensenova_hidden_states_stream.pt"
29
+ )
30
+
31
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
32
+
33
+ PROMPTS = [
34
+ "a red cube on a blue sphere",
35
+ "a woman holding a transparent glass bottle",
36
+ "three people standing behind a wooden table",
37
+ ]
38
+
39
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
40
+
41
+ # Не отдаём GPU все 24 GB, оставляем запас.
42
+ MAX_MEMORY = {
43
+ 0: "20GiB",
44
+ "cpu": "70GiB",
45
+ }
46
+
47
+
48
+ # ============================================================
49
+ # LOCAL SENSENOVA CODE
50
+ # ============================================================
51
+
52
+ sys.path.insert(0, SENSENOVA_SRC)
53
+
54
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
55
+ NEOChatConfig,
56
+ )
57
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
58
+ NEOChatModel,
59
+ )
60
+
61
+
62
+ print("=" * 80)
63
+ print("SenseNova U1.5 hidden-state extractor")
64
+ print("MEMORY-EFFICIENT / ACCELERATE")
65
+ print("=" * 80)
66
+
67
+ print("\nGPU:")
68
+ print(torch.cuda.get_device_name(0))
69
+
70
+ print(
71
+ "VRAM:",
72
+ round(
73
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
74
+ 2
75
+ ),
76
+ "GB"
77
+ )
78
+
79
+
80
+ # ============================================================
81
+ # TOKENIZER + CONFIG
82
+ # ============================================================
83
+
84
+ print("\nLoading tokenizer...")
85
+
86
+ tokenizer = AutoTokenizer.from_pretrained(
87
+ HF_REPO,
88
+ trust_remote_code=True,
89
+ )
90
+
91
+ print("Loading LOCAL SenseNova config...")
92
+
93
+ config = NEOChatConfig.from_pretrained(
94
+ HF_REPO,
95
+ )
96
+
97
+ print("Config loaded.")
98
+ print("LLM hidden size:", config.llm_config.hidden_size)
99
+ print("LLM layers:", config.llm_config.num_hidden_layers)
100
+
101
+
102
+ # ============================================================
103
+ # CREATE MODEL ON META
104
+ #
105
+ # Ключевой момент:
106
+ # здесь НЕ выделяется 45+ GB RAM под пустую модель.
107
+ # ============================================================
108
+
109
+ print()
110
+ print("Creating architecture on META device...")
111
+
112
+ with init_empty_weights():
113
+ model = NEOChatModel(config)
114
+
115
+ print("META architecture created.")
116
+
117
+
118
+ # ============================================================
119
+ # LOAD + DISPATCH CHECKPOINT DIRECTLY
120
+ #
121
+ # Нет state = {...}
122
+ # Нет второй полной копии checkpoint в RAM.
123
+ # ============================================================
124
+
125
+ print()
126
+ print("Loading checkpoint directly with Accelerate...")
127
+ print("This can take a few minutes, but RAM should stay reasonable.")
128
+
129
+ model = load_checkpoint_and_dispatch(
130
+ model,
131
+ checkpoint=MODEL_FILE,
132
+ device_map="auto",
133
+ max_memory=MAX_MEMORY,
134
+ dtype=torch.bfloat16,
135
+
136
+ # Decoder layer нельзя разрезать между устройствами.
137
+ no_split_module_classes=[
138
+ "Qwen3DecoderLayer",
139
+ "Qwen3MoeDecoderLayer",
140
+ ],
141
+ )
142
+
143
+ model.eval()
144
+
145
+ print()
146
+ print("Checkpoint loaded successfully.")
147
+
148
+
149
+ # ============================================================
150
+ # FIND LANGUAGE MODEL
151
+ # ============================================================
152
+
153
+ language_model = model.language_model
154
+
155
+ print("Language model:")
156
+ print(type(language_model))
157
+
158
+ print("\nDevice map summary:")
159
+
160
+ device_map = getattr(model, "hf_device_map", None)
161
+
162
+ if device_map:
163
+ gpu_count = 0
164
+ cpu_count = 0
165
+ disk_count = 0
166
+
167
+ for name, dev in device_map.items():
168
+
169
+ s = str(dev)
170
+
171
+ if s in ("0", "cuda", "cuda:0"):
172
+ gpu_count += 1
173
+ elif s == "cpu":
174
+ cpu_count += 1
175
+ elif s == "disk":
176
+ disk_count += 1
177
+
178
+ print("GPU modules :", gpu_count)
179
+ print("CPU modules :", cpu_count)
180
+ print("Disk modules:", disk_count)
181
+
182
+ print("\nFirst device-map entries:")
183
+ for i, (name, dev) in enumerate(device_map.items()):
184
+ print(f" {name:60s} -> {dev}")
185
+
186
+ if i >= 19:
187
+ break
188
+
189
+
190
+ # ============================================================
191
+ # FIND TARGET LAYERS
192
+ # ============================================================
193
+
194
+ layers = language_model.model.layers
195
+
196
+ print()
197
+ print("Detected language layers:", len(layers))
198
+
199
+
200
+ # ============================================================
201
+ # HOOKS
202
+ # ============================================================
203
+
204
+ captured = {}
205
+ hooks = {}
206
+
207
+
208
+ def make_hook(idx):
209
+
210
+ def hook(module, inputs, output):
211
+
212
+ if isinstance(output, tuple):
213
+ x = output[0]
214
+ else:
215
+ x = output
216
+
217
+ captured[idx] = (
218
+ x.detach()
219
+ .float()
220
+ .cpu()
221
+ )
222
+
223
+ return hook
224
+
225
+
226
+ print("\nRegistering hooks...")
227
+
228
+ for idx in TARGET_LAYERS:
229
+
230
+ if idx >= len(layers):
231
+ print(f"layer {idx}: OUT OF RANGE")
232
+ continue
233
+
234
+ hooks[idx] = layers[idx].register_forward_hook(
235
+ make_hook(idx)
236
+ )
237
+
238
+ print(
239
+ f"layer {idx}:",
240
+ type(layers[idx]),
241
+ )
242
+
243
+
244
+ # ============================================================
245
+ # RUN PROMPTS
246
+ # ============================================================
247
+
248
+ results = []
249
+
250
+ for i, prompt in enumerate(PROMPTS):
251
+
252
+ print()
253
+ print("=" * 80)
254
+ print(f"[{i + 1}/{len(PROMPTS)}]")
255
+ print(prompt)
256
+ print("=" * 80)
257
+
258
+ encoded = tokenizer(
259
+ prompt,
260
+ return_tensors="pt",
261
+ add_special_tokens=False,
262
+ )
263
+
264
+ # Input должен попасть туда, где embedding.
265
+ embed_device = language_model.model.embed_tokens.weight.device
266
+
267
+ # При accelerate иногда weight показывает meta;
268
+ # тогда безопасно используем cuda:0.
269
+ if embed_device.type == "meta":
270
+ embed_device = torch.device("cuda:0")
271
+
272
+ input_ids = encoded["input_ids"].to(embed_device)
273
+
274
+ print("Tokens:", input_ids.tolist()[0])
275
+ print("Token count:", input_ids.shape[1])
276
+ print("Input device:", input_ids.device)
277
+
278
+ captured.clear()
279
+
280
+ with torch.inference_mode():
281
+
282
+ outputs = language_model.model(
283
+ input_ids=input_ids,
284
+ use_cache=False,
285
+ )
286
+
287
+ item = {
288
+ "prompt": prompt,
289
+ "input_ids": input_ids.detach().cpu(),
290
+ "layers": {},
291
+ }
292
+
293
+ for idx in TARGET_LAYERS:
294
+
295
+ if idx not in captured:
296
+
297
+ print(
298
+ f"layer {idx:2d}: NOT CAPTURED"
299
+ )
300
+
301
+ continue
302
+
303
+ x = captured[idx]
304
+
305
+ item["layers"][idx] = x
306
+
307
+ print(
308
+ f"layer {idx:2d}: "
309
+ f"shape={tuple(x.shape)} "
310
+ f"mean={x.mean().item():.6f} "
311
+ f"std={x.std().item():.6f} "
312
+ f"finite={torch.isfinite(x).all().item()}"
313
+ )
314
+
315
+ results.append(item)
316
+
317
+ del outputs
318
+ del input_ids
319
+
320
+ torch.cuda.empty_cache()
321
+
322
+
323
+ # ============================================================
324
+ # CLEANUP HOOKS
325
+ # ============================================================
326
+
327
+ for h in hooks.values():
328
+ h.remove()
329
+
330
+
331
+ # ============================================================
332
+ # SAVE
333
+ # ============================================================
334
+
335
+ torch.save(
336
+ {
337
+ "prompts": PROMPTS,
338
+ "target_layers": TARGET_LAYERS,
339
+ "results": results,
340
+ },
341
+ OUT,
342
+ )
343
+
344
+ print()
345
+ print("=" * 80)
346
+ print("SUCCESS")
347
+ print("Saved:")
348
+ print(OUT)
349
+ print("=" * 80)
350
+
351
+ gc.collect()
352
+ torch.cuda.empty_cache()
research/raw_scripts/extract_sensenova_hidden_states_stream_v2.py ADDED
@@ -0,0 +1,586 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import inspect
5
+ import torch
6
+
7
+ from transformers import AutoTokenizer
8
+ from accelerate import init_empty_weights, load_checkpoint_and_dispatch
9
+
10
+
11
+ # ============================================================
12
+ # PATHS
13
+ # ============================================================
14
+
15
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
16
+
17
+ SENSENOVA_SRC = os.path.join(
18
+ ROOT,
19
+ "SenseNova-U1",
20
+ "src"
21
+ )
22
+
23
+ MODEL_FILE = (
24
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
25
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
26
+ )
27
+
28
+ OUT = os.path.join(
29
+ ROOT,
30
+ "sensenova_hidden_states_stream_v2.pt"
31
+ )
32
+
33
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
34
+
35
+
36
+ # ============================================================
37
+ # TEST PROMPTS
38
+ # ============================================================
39
+
40
+ PROMPTS = [
41
+ "a red cube on a blue sphere",
42
+ "a woman holding a transparent glass bottle",
43
+ "three people standing behind a wooden table",
44
+ ]
45
+
46
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
47
+
48
+
49
+ # ============================================================
50
+ # MEMORY
51
+ # ============================================================
52
+
53
+ # Оставляем запас VRAM для forward.
54
+ MAX_MEMORY = {
55
+ 0: "19GiB",
56
+ "cpu": "70GiB",
57
+ }
58
+
59
+
60
+ # ============================================================
61
+ # IMPORT LOCAL SENSENOVA CODE
62
+ # ============================================================
63
+
64
+ sys.path.insert(0, SENSENOVA_SRC)
65
+
66
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
67
+ NEOChatConfig,
68
+ )
69
+
70
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
71
+ NEOChatModel,
72
+ )
73
+
74
+
75
+ print("=" * 80)
76
+ print("SenseNova U1.5 hidden-state extractor V2")
77
+ print("Transformers 4.57 compatibility-wrapper bypass")
78
+ print("=" * 80)
79
+
80
+ print("\nPyTorch:", torch.__version__)
81
+
82
+ import transformers
83
+ import accelerate
84
+
85
+ print("Transformers:", transformers.__version__)
86
+ print("Accelerate:", accelerate.__version__)
87
+
88
+ print("\nGPU:")
89
+ print(torch.cuda.get_device_name(0))
90
+
91
+ print(
92
+ "VRAM:",
93
+ round(
94
+ torch.cuda.get_device_properties(0).total_memory
95
+ / 1024**3,
96
+ 2
97
+ ),
98
+ "GB"
99
+ )
100
+
101
+
102
+ # ============================================================
103
+ # TOKENIZER
104
+ # ============================================================
105
+
106
+ print("\nLoading tokenizer...")
107
+
108
+ tokenizer = AutoTokenizer.from_pretrained(
109
+ HF_REPO,
110
+ trust_remote_code=True,
111
+ )
112
+
113
+
114
+ # ============================================================
115
+ # CONFIG
116
+ # ============================================================
117
+
118
+ print("Loading LOCAL SenseNova config...")
119
+
120
+ config = NEOChatConfig.from_pretrained(
121
+ HF_REPO,
122
+ )
123
+
124
+ print("Config loaded.")
125
+ print("LLM hidden size:", config.llm_config.hidden_size)
126
+ print("LLM layers:", config.llm_config.num_hidden_layers)
127
+
128
+
129
+ # ============================================================
130
+ # META MODEL
131
+ # ============================================================
132
+
133
+ print()
134
+ print("Creating architecture on META device...")
135
+
136
+ with init_empty_weights():
137
+ model = NEOChatModel(config)
138
+
139
+ print("META architecture created.")
140
+
141
+
142
+ # ============================================================
143
+ # LOAD CHECKPOINT
144
+ # ============================================================
145
+
146
+ print()
147
+ print("Loading local BF16 checkpoint with Accelerate...")
148
+ print("No duplicate full state_dict will be created.")
149
+
150
+ model = load_checkpoint_and_dispatch(
151
+ model,
152
+ checkpoint=MODEL_FILE,
153
+ device_map="auto",
154
+ max_memory=MAX_MEMORY,
155
+ dtype=torch.bfloat16,
156
+ no_split_module_classes=[
157
+ "Qwen3DecoderLayer",
158
+ "Qwen3MoeDecoderLayer",
159
+ ],
160
+ )
161
+
162
+ model.eval()
163
+
164
+ print()
165
+ print("Checkpoint loaded successfully.")
166
+
167
+
168
+ # ============================================================
169
+ # LANGUAGE MODEL
170
+ # ============================================================
171
+
172
+ language_model = model.language_model
173
+ qwen_model = language_model.model
174
+ layers = qwen_model.layers
175
+
176
+ print()
177
+ print("Language model:")
178
+ print(type(language_model))
179
+
180
+ print("Base model:")
181
+ print(type(qwen_model))
182
+
183
+ print("Detected language layers:", len(layers))
184
+
185
+
186
+ # ============================================================
187
+ # DEVICE MAP
188
+ # ============================================================
189
+
190
+ device_map = getattr(model, "hf_device_map", None)
191
+
192
+ if device_map:
193
+
194
+ gpu_count = 0
195
+ cpu_count = 0
196
+ disk_count = 0
197
+
198
+ for name, dev in device_map.items():
199
+
200
+ s = str(dev)
201
+
202
+ if s in ("0", "cuda", "cuda:0"):
203
+ gpu_count += 1
204
+
205
+ elif s == "cpu":
206
+ cpu_count += 1
207
+
208
+ elif s == "disk":
209
+ disk_count += 1
210
+
211
+ print()
212
+ print("Device map:")
213
+ print(" GPU modules :", gpu_count)
214
+ print(" CPU modules :", cpu_count)
215
+ print(" Disk modules:", disk_count)
216
+
217
+
218
+ # ============================================================
219
+ # BYPASS TRANSFORMERS 4.57 WRAPPER
220
+ # ============================================================
221
+
222
+ print()
223
+ print("Preparing raw SenseNova forward...")
224
+
225
+ # Берём forward непосредственно с класса,
226
+ # затем снимаем ВСЕ decorators через inspect.unwrap().
227
+ #
228
+ # Это позволяет обойти check_model_inputs из Transformers,
229
+ # который ломает прямой вызов SenseNova в 4.57.x.
230
+
231
+ decorated_forward = type(qwen_model).forward
232
+ raw_forward = inspect.unwrap(decorated_forward)
233
+
234
+ print("Decorated forward:")
235
+ print(decorated_forward)
236
+
237
+ print()
238
+ print("Unwrapped forward:")
239
+ print(raw_forward)
240
+
241
+ try:
242
+ print()
243
+ print("Raw forward signature:")
244
+ print(inspect.signature(raw_forward))
245
+ except Exception:
246
+ pass
247
+
248
+
249
+ # ============================================================
250
+ # HOOKS
251
+ # ============================================================
252
+
253
+ captured = {}
254
+ hooks = {}
255
+
256
+
257
+ def make_hook(idx):
258
+
259
+ def hook(module, inputs, output):
260
+
261
+ if isinstance(output, tuple):
262
+ x = output[0]
263
+
264
+ elif hasattr(output, "hidden_states"):
265
+ x = output.hidden_states
266
+
267
+ else:
268
+ x = output
269
+
270
+ if not torch.is_tensor(x):
271
+ print(
272
+ f"WARNING: layer {idx} returned "
273
+ f"{type(x)} instead of Tensor"
274
+ )
275
+ return
276
+
277
+ captured[idx] = (
278
+ x.detach()
279
+ .to(
280
+ device="cpu",
281
+ dtype=torch.float32,
282
+ )
283
+ )
284
+
285
+ return hook
286
+
287
+
288
+ print()
289
+ print("Registering hooks...")
290
+
291
+ for idx in TARGET_LAYERS:
292
+
293
+ if idx >= len(layers):
294
+ print(f" layer {idx}: OUT OF RANGE")
295
+ continue
296
+
297
+ hooks[idx] = layers[idx].register_forward_hook(
298
+ make_hook(idx)
299
+ )
300
+
301
+ print(
302
+ f" layer {idx}:",
303
+ type(layers[idx]),
304
+ )
305
+
306
+
307
+ if not hooks:
308
+ raise RuntimeError("No hooks were registered.")
309
+
310
+
311
+ # ============================================================
312
+ # FIND INPUT DEVICE
313
+ # ============================================================
314
+
315
+ def get_embedding_device():
316
+
317
+ try:
318
+ d = qwen_model.embed_tokens.weight.device
319
+
320
+ if d.type != "meta":
321
+ return d
322
+
323
+ except Exception:
324
+ pass
325
+
326
+ # По device_map embedding у нас обычно на GPU.
327
+ if device_map:
328
+
329
+ for name, dev in device_map.items():
330
+
331
+ if name.endswith(
332
+ "language_model.model.embed_tokens"
333
+ ):
334
+
335
+ if str(dev) == "0":
336
+ return torch.device("cuda:0")
337
+
338
+ return torch.device(str(dev))
339
+
340
+ return torch.device("cuda:0")
341
+
342
+
343
+ INPUT_DEVICE = get_embedding_device()
344
+
345
+ print()
346
+ print("Input device:", INPUT_DEVICE)
347
+
348
+
349
+ # ============================================================
350
+ # SAFE RAW FORWARD
351
+ # ============================================================
352
+
353
+ def run_sensenova_forward(input_ids):
354
+
355
+ """
356
+ Вызывает настоящий Qwen3Model.forward без
357
+ Transformers compatibility decorator.
358
+ """
359
+
360
+ kwargs = {
361
+ "input_ids": input_ids,
362
+ "use_cache": False,
363
+
364
+ # Для обычного text/understanding path.
365
+ # MoT-generation path будем тестировать отдельно.
366
+ "image_gen_indicators": None,
367
+ }
368
+
369
+ # Первый вариант: полный ожидаемый вызов.
370
+ try:
371
+
372
+ return raw_forward(
373
+ qwen_model,
374
+ **kwargs,
375
+ )
376
+
377
+ except TypeError as e1:
378
+
379
+ print()
380
+ print("First raw-forward variant failed:")
381
+ print(repr(e1))
382
+
383
+ print("Retrying minimal raw forward...")
384
+
385
+ # На случай небольших отличий сигнатуры.
386
+ try:
387
+
388
+ return raw_forward(
389
+ qwen_model,
390
+ input_ids=input_ids,
391
+ use_cache=False,
392
+ )
393
+
394
+ except TypeError as e2:
395
+
396
+ print()
397
+ print("Minimal raw forward also failed:")
398
+ print(repr(e2))
399
+
400
+ print()
401
+ print("Raw signature:")
402
+ print(inspect.signature(raw_forward))
403
+
404
+ raise
405
+
406
+
407
+ # ============================================================
408
+ # RUN PROMPTS
409
+ # ============================================================
410
+
411
+ results = []
412
+
413
+ for i, prompt in enumerate(PROMPTS):
414
+
415
+ print()
416
+ print("=" * 80)
417
+ print(f"[{i + 1}/{len(PROMPTS)}]")
418
+ print(prompt)
419
+ print("=" * 80)
420
+
421
+ encoded = tokenizer(
422
+ prompt,
423
+ return_tensors="pt",
424
+ add_special_tokens=False,
425
+ )
426
+
427
+ input_ids = encoded["input_ids"].to(
428
+ INPUT_DEVICE
429
+ )
430
+
431
+ print(
432
+ "Tokens:",
433
+ input_ids.detach().cpu().tolist()[0]
434
+ )
435
+
436
+ print(
437
+ "Token count:",
438
+ input_ids.shape[1]
439
+ )
440
+
441
+ print(
442
+ "Input device:",
443
+ input_ids.device
444
+ )
445
+
446
+ captured.clear()
447
+
448
+ with torch.inference_mode():
449
+
450
+ outputs = run_sensenova_forward(
451
+ input_ids
452
+ )
453
+
454
+ item = {
455
+ "prompt": prompt,
456
+ "input_ids": (
457
+ input_ids.detach()
458
+ .cpu()
459
+ ),
460
+ "layers": {},
461
+ }
462
+
463
+ print()
464
+
465
+ for idx in TARGET_LAYERS:
466
+
467
+ if idx not in captured:
468
+
469
+ print(
470
+ f"layer {idx:2d}: NOT CAPTURED"
471
+ )
472
+
473
+ continue
474
+
475
+ x = captured[idx]
476
+
477
+ item["layers"][idx] = x
478
+
479
+ finite = torch.isfinite(x).all().item()
480
+
481
+ print(
482
+ f"layer {idx:2d}: "
483
+ f"shape={tuple(x.shape)} "
484
+ f"dtype={x.dtype} "
485
+ f"mean={x.mean().item():.6f} "
486
+ f"std={x.std().item():.6f} "
487
+ f"finite={finite}"
488
+ )
489
+
490
+ results.append(item)
491
+
492
+ del outputs
493
+ del input_ids
494
+
495
+ gc.collect()
496
+ torch.cuda.empty_cache()
497
+
498
+
499
+ # ============================================================
500
+ # REMOVE HOOKS
501
+ # ============================================================
502
+
503
+ for h in hooks.values():
504
+ h.remove()
505
+
506
+
507
+ # ============================================================
508
+ # VALIDATION
509
+ # ============================================================
510
+
511
+ print()
512
+ print("=" * 80)
513
+ print("VALIDATION")
514
+ print("=" * 80)
515
+
516
+ all_ok = True
517
+
518
+ for item in results:
519
+
520
+ print()
521
+ print(item["prompt"])
522
+
523
+ for idx in TARGET_LAYERS:
524
+
525
+ if idx not in item["layers"]:
526
+
527
+ print(
528
+ f" layer {idx}: missing"
529
+ )
530
+
531
+ all_ok = False
532
+ continue
533
+
534
+ x = item["layers"][idx]
535
+
536
+ correct_dim = (
537
+ x.ndim == 3
538
+ and x.shape[-1] == 4096
539
+ )
540
+
541
+ finite = torch.isfinite(x).all().item()
542
+
543
+ print(
544
+ f" layer {idx}: "
545
+ f"{tuple(x.shape)} "
546
+ f"dim_ok={correct_dim} "
547
+ f"finite={finite}"
548
+ )
549
+
550
+ if not correct_dim or not finite:
551
+ all_ok = False
552
+
553
+
554
+ # ============================================================
555
+ # SAVE
556
+ # ============================================================
557
+
558
+ torch.save(
559
+ {
560
+ "prompts": PROMPTS,
561
+ "target_layers": TARGET_LAYERS,
562
+ "results": results,
563
+ "source_model": MODEL_FILE,
564
+ "transformers_version": transformers.__version__,
565
+ },
566
+ OUT,
567
+ )
568
+
569
+
570
+ print()
571
+ print("=" * 80)
572
+
573
+ if all_ok:
574
+ print("SUCCESS")
575
+ else:
576
+ print("COMPLETED WITH WARNINGS")
577
+
578
+ print()
579
+ print("Saved:")
580
+ print(OUT)
581
+
582
+ print("=" * 80)
583
+
584
+
585
+ gc.collect()
586
+ torch.cuda.empty_cache()
research/raw_scripts/extract_sensenova_hidden_states_stream_v3.py ADDED
@@ -0,0 +1,546 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import shutil
5
+ import torch
6
+
7
+
8
+ # ============================================================
9
+ # PATHS
10
+ # ============================================================
11
+
12
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
13
+
14
+ SENSENOVA_REPO = os.path.join(
15
+ ROOT,
16
+ "SenseNova-U1"
17
+ )
18
+
19
+ SENSENOVA_SRC = os.path.join(
20
+ SENSENOVA_REPO,
21
+ "src"
22
+ )
23
+
24
+ COMPAT_FILE = os.path.join(
25
+ SENSENOVA_SRC,
26
+ "sensenova_u1",
27
+ "models",
28
+ "neo_unify",
29
+ "transformers_compat.py"
30
+ )
31
+
32
+ MODEL_FILE = (
33
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
34
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
35
+ )
36
+
37
+ OUT = os.path.join(
38
+ ROOT,
39
+ "sensenova_hidden_states_stream_v3.pt"
40
+ )
41
+
42
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
43
+
44
+
45
+ PROMPTS = [
46
+ "a red cube on a blue sphere",
47
+ "a woman holding a transparent glass bottle",
48
+ "three people standing behind a wooden table",
49
+ ]
50
+
51
+ TARGET_LAYERS = [8, 16, 24, 32, 41]
52
+
53
+
54
+ # ============================================================
55
+ # PATCH TRANSFORMERS 4.57 COMPATIBILITY BUG
56
+ # ============================================================
57
+
58
+ print("=" * 80)
59
+ print("SenseNova U1.5 hidden-state extractor V3")
60
+ print("=" * 80)
61
+
62
+ print("\nChecking SenseNova Transformers compatibility...")
63
+
64
+ with open(COMPAT_FILE, "r", encoding="utf-8") as f:
65
+ text = f.read()
66
+
67
+ OLD = (
68
+ " from transformers.utils.generic import "
69
+ "check_model_inputs as model_input_compat"
70
+ )
71
+
72
+ NEW = (
73
+ " from transformers.utils.generic import check_model_inputs\n"
74
+ " def model_input_compat(func: Callable[..., Any]) -> Callable[..., Any]:\n"
75
+ " return check_model_inputs()(func)"
76
+ )
77
+
78
+ if OLD in text:
79
+
80
+ backup = COMPAT_FILE + ".bak"
81
+
82
+ if not os.path.exists(backup):
83
+ shutil.copy2(COMPAT_FILE, backup)
84
+ print("Backup created:")
85
+ print(backup)
86
+
87
+ text = text.replace(OLD, NEW)
88
+
89
+ with open(COMPAT_FILE, "w", encoding="utf-8") as f:
90
+ f.write(text)
91
+
92
+ print("Compatibility patch APPLIED.")
93
+
94
+ elif "return check_model_inputs()(func)" in text:
95
+
96
+ print("Compatibility patch already present.")
97
+
98
+ else:
99
+
100
+ print("WARNING: expected compatibility line not found.")
101
+ print("No automatic source modification performed.")
102
+
103
+
104
+ # ============================================================
105
+ # IMPORTANT:
106
+ # Import SenseNova only AFTER patching source.
107
+ # ============================================================
108
+
109
+ sys.path.insert(0, SENSENOVA_SRC)
110
+
111
+ from transformers import AutoTokenizer
112
+ from accelerate import (
113
+ init_empty_weights,
114
+ load_checkpoint_and_dispatch,
115
+ )
116
+
117
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
118
+ NEOChatConfig,
119
+ )
120
+
121
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
122
+ NEOChatModel,
123
+ )
124
+
125
+
126
+ import transformers
127
+ import accelerate
128
+
129
+ print()
130
+ print("PyTorch:", torch.__version__)
131
+ print("Transformers:", transformers.__version__)
132
+ print("Accelerate:", accelerate.__version__)
133
+
134
+ print()
135
+ print("GPU:")
136
+ print(torch.cuda.get_device_name(0))
137
+
138
+ print(
139
+ "VRAM:",
140
+ round(
141
+ torch.cuda.get_device_properties(0).total_memory
142
+ / 1024**3,
143
+ 2
144
+ ),
145
+ "GB"
146
+ )
147
+
148
+
149
+ # ============================================================
150
+ # TOKENIZER / CONFIG
151
+ # ============================================================
152
+
153
+ print("\nLoading tokenizer...")
154
+
155
+ tokenizer = AutoTokenizer.from_pretrained(
156
+ HF_REPO,
157
+ trust_remote_code=True,
158
+ )
159
+
160
+ print("Loading LOCAL SenseNova config...")
161
+
162
+ config = NEOChatConfig.from_pretrained(
163
+ HF_REPO,
164
+ )
165
+
166
+ # Force eager because SenseNova forward_und expects it.
167
+ config.llm_config._attn_implementation = "eager"
168
+
169
+ print("Config loaded.")
170
+ print("LLM hidden size:", config.llm_config.hidden_size)
171
+ print("LLM layers:", config.llm_config.num_hidden_layers)
172
+ print(
173
+ "Attention implementation:",
174
+ config.llm_config._attn_implementation
175
+ )
176
+
177
+
178
+ # ============================================================
179
+ # META MODEL
180
+ # ============================================================
181
+
182
+ print()
183
+ print("Creating architecture on META device...")
184
+
185
+ with init_empty_weights():
186
+
187
+ model = NEOChatModel(config)
188
+
189
+ print("META architecture created.")
190
+
191
+
192
+ # ============================================================
193
+ # LOAD CHECKPOINT
194
+ # ============================================================
195
+
196
+ MAX_MEMORY = {
197
+ 0: "19GiB",
198
+ "cpu": "70GiB",
199
+ }
200
+
201
+ print()
202
+ print("Loading local BF16 checkpoint with Accelerate...")
203
+
204
+ model = load_checkpoint_and_dispatch(
205
+ model,
206
+ checkpoint=MODEL_FILE,
207
+ device_map="auto",
208
+ max_memory=MAX_MEMORY,
209
+ dtype=torch.bfloat16,
210
+ no_split_module_classes=[
211
+ "Qwen3DecoderLayer",
212
+ "Qwen3MoeDecoderLayer",
213
+ ],
214
+ )
215
+
216
+ model.eval()
217
+
218
+ print("Checkpoint loaded successfully.")
219
+
220
+
221
+ # ============================================================
222
+ # LANGUAGE MODEL
223
+ # ============================================================
224
+
225
+ language_model = model.language_model
226
+ qwen_model = language_model.model
227
+ layers = qwen_model.layers
228
+
229
+ print()
230
+ print("Language model:")
231
+ print(type(language_model))
232
+
233
+ print("Base model:")
234
+ print(type(qwen_model))
235
+
236
+ print("Detected language layers:", len(layers))
237
+
238
+
239
+ # ============================================================
240
+ # CHECK FORWARD SIGNATURE
241
+ # ============================================================
242
+
243
+ import inspect
244
+
245
+ print()
246
+ print("Qwen3Model.forward signature:")
247
+
248
+ try:
249
+ print(inspect.signature(qwen_model.forward))
250
+ except Exception as e:
251
+ print("Could not inspect:", e)
252
+
253
+ # This MUST contain input_ids now.
254
+ sig = str(inspect.signature(qwen_model.forward))
255
+
256
+ if "input_ids" not in sig:
257
+
258
+ raise RuntimeError(
259
+ "SenseNova compatibility patch did not fix Qwen3Model.forward. "
260
+ f"Current signature: {sig}"
261
+ )
262
+
263
+ print("Forward compatibility: OK")
264
+
265
+
266
+ # ============================================================
267
+ # DEVICE MAP
268
+ # ============================================================
269
+
270
+ device_map = getattr(model, "hf_device_map", {})
271
+
272
+ gpu_count = 0
273
+ cpu_count = 0
274
+ disk_count = 0
275
+
276
+ for _, dev in device_map.items():
277
+
278
+ s = str(dev)
279
+
280
+ if s in ("0", "cuda", "cuda:0"):
281
+ gpu_count += 1
282
+
283
+ elif s == "cpu":
284
+ cpu_count += 1
285
+
286
+ elif s == "disk":
287
+ disk_count += 1
288
+
289
+ print()
290
+ print("Device map:")
291
+ print(" GPU modules :", gpu_count)
292
+ print(" CPU modules :", cpu_count)
293
+ print(" Disk modules:", disk_count)
294
+
295
+
296
+ # ============================================================
297
+ # INPUT DEVICE
298
+ # ============================================================
299
+
300
+ INPUT_DEVICE = torch.device("cuda:0")
301
+
302
+ for name, dev in device_map.items():
303
+
304
+ if name.endswith("language_model.model.embed_tokens"):
305
+
306
+ if str(dev) == "0":
307
+ INPUT_DEVICE = torch.device("cuda:0")
308
+
309
+ elif str(dev) == "cpu":
310
+ INPUT_DEVICE = torch.device("cpu")
311
+
312
+ break
313
+
314
+ print()
315
+ print("Input device:", INPUT_DEVICE)
316
+
317
+
318
+ # ============================================================
319
+ # HOOKS
320
+ # ============================================================
321
+
322
+ captured = {}
323
+ hooks = {}
324
+
325
+
326
+ def make_hook(idx):
327
+
328
+ def hook(module, inputs, output):
329
+
330
+ if isinstance(output, tuple):
331
+ x = output[0]
332
+ else:
333
+ x = output
334
+
335
+ if not torch.is_tensor(x):
336
+ print(
337
+ f"WARNING layer {idx}: "
338
+ f"unexpected output type {type(x)}"
339
+ )
340
+ return
341
+
342
+ captured[idx] = (
343
+ x.detach()
344
+ .float()
345
+ .cpu()
346
+ )
347
+
348
+ return hook
349
+
350
+
351
+ print()
352
+ print("Registering hooks...")
353
+
354
+ for idx in TARGET_LAYERS:
355
+
356
+ if idx >= len(layers):
357
+ print("layer", idx, ": OUT OF RANGE")
358
+ continue
359
+
360
+ hooks[idx] = layers[idx].register_forward_hook(
361
+ make_hook(idx)
362
+ )
363
+
364
+ print(
365
+ f"layer {idx}:",
366
+ type(layers[idx])
367
+ )
368
+
369
+
370
+ # ============================================================
371
+ # RUN
372
+ # ============================================================
373
+
374
+ results = []
375
+
376
+ for i, prompt in enumerate(PROMPTS):
377
+
378
+ print()
379
+ print("=" * 80)
380
+ print(f"[{i + 1}/{len(PROMPTS)}]")
381
+ print(prompt)
382
+ print("=" * 80)
383
+
384
+ encoded = tokenizer(
385
+ prompt,
386
+ return_tensors="pt",
387
+ add_special_tokens=False,
388
+ )
389
+
390
+ input_ids = encoded["input_ids"].to(
391
+ INPUT_DEVICE
392
+ )
393
+
394
+ print(
395
+ "Tokens:",
396
+ input_ids.detach().cpu().tolist()[0]
397
+ )
398
+
399
+ print("Token count:", input_ids.shape[1])
400
+ print("Input device:", input_ids.device)
401
+
402
+ captured.clear()
403
+
404
+ # Reset SenseNova internal prompt index.
405
+ qwen_model.current_index = -1
406
+
407
+ with torch.inference_mode():
408
+
409
+ outputs = qwen_model(
410
+ input_ids=input_ids,
411
+ image_gen_indicators=None,
412
+ attention_mask=None,
413
+ use_cache=False,
414
+ )
415
+
416
+ item = {
417
+ "prompt": prompt,
418
+ "input_ids": (
419
+ input_ids.detach()
420
+ .cpu()
421
+ ),
422
+ "layers": {},
423
+ }
424
+
425
+ print()
426
+
427
+ for idx in TARGET_LAYERS:
428
+
429
+ if idx not in captured:
430
+
431
+ print(
432
+ f"layer {idx:2d}: NOT CAPTURED"
433
+ )
434
+ continue
435
+
436
+ x = captured[idx]
437
+
438
+ item["layers"][idx] = x
439
+
440
+ finite = torch.isfinite(x).all().item()
441
+
442
+ print(
443
+ f"layer {idx:2d}: "
444
+ f"shape={tuple(x.shape)} "
445
+ f"dtype={x.dtype} "
446
+ f"mean={x.mean().item():.6f} "
447
+ f"std={x.std().item():.6f} "
448
+ f"finite={finite}"
449
+ )
450
+
451
+ results.append(item)
452
+
453
+ del outputs
454
+ del input_ids
455
+
456
+ gc.collect()
457
+ torch.cuda.empty_cache()
458
+
459
+
460
+ # ============================================================
461
+ # REMOVE HOOKS
462
+ # ============================================================
463
+
464
+ for h in hooks.values():
465
+ h.remove()
466
+
467
+
468
+ # ============================================================
469
+ # VALIDATE
470
+ # ============================================================
471
+
472
+ print()
473
+ print("=" * 80)
474
+ print("VALIDATION")
475
+ print("=" * 80)
476
+
477
+ all_ok = True
478
+
479
+ for item in results:
480
+
481
+ print()
482
+ print(item["prompt"])
483
+
484
+ for idx in TARGET_LAYERS:
485
+
486
+ if idx not in item["layers"]:
487
+
488
+ print(
489
+ f" layer {idx}: missing"
490
+ )
491
+
492
+ all_ok = False
493
+ continue
494
+
495
+ x = item["layers"][idx]
496
+
497
+ dim_ok = (
498
+ x.ndim == 3
499
+ and x.shape[-1] == 4096
500
+ )
501
+
502
+ finite = torch.isfinite(x).all().item()
503
+
504
+ print(
505
+ f" layer {idx}: "
506
+ f"{tuple(x.shape)} "
507
+ f"dim_ok={dim_ok} "
508
+ f"finite={finite}"
509
+ )
510
+
511
+ if not dim_ok or not finite:
512
+ all_ok = False
513
+
514
+
515
+ # ============================================================
516
+ # SAVE
517
+ # ============================================================
518
+
519
+ torch.save(
520
+ {
521
+ "prompts": PROMPTS,
522
+ "target_layers": TARGET_LAYERS,
523
+ "results": results,
524
+ "source_model": MODEL_FILE,
525
+ "transformers_version": transformers.__version__,
526
+ },
527
+ OUT,
528
+ )
529
+
530
+
531
+ print()
532
+ print("=" * 80)
533
+
534
+ if all_ok:
535
+ print("SUCCESS")
536
+ else:
537
+ print("COMPLETED WITH WARNINGS")
538
+
539
+ print()
540
+ print("Saved:")
541
+ print(OUT)
542
+
543
+ print("=" * 80)
544
+
545
+ gc.collect()
546
+ torch.cuda.empty_cache()
research/raw_scripts/extract_sensenova_ood_hidden.py ADDED
@@ -0,0 +1,633 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import json
4
+ import gc
5
+ import shutil
6
+ import torch
7
+
8
+
9
+ # ============================================================
10
+ # PATHS
11
+ # ============================================================
12
+
13
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
14
+
15
+ SENSENOVA_SRC = os.path.join(
16
+ ROOT,
17
+ "SenseNova-U1",
18
+ "src"
19
+ )
20
+
21
+ COMPAT_FILE = os.path.join(
22
+ SENSENOVA_SRC,
23
+ "sensenova_u1",
24
+ "models",
25
+ "neo_unify",
26
+ "transformers_compat.py"
27
+ )
28
+
29
+ CHECKPOINT = (
30
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
31
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
32
+ )
33
+
34
+ PROMPTS_FILE = os.path.join(
35
+ ROOT,
36
+ "bridge_ood_prompts_160.json"
37
+ )
38
+
39
+ OUT_FILE = os.path.join(
40
+ ROOT,
41
+ "sensenova_ood_hidden_160.pt"
42
+ )
43
+
44
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
45
+
46
+ LAYER = 32
47
+
48
+
49
+ # ============================================================
50
+ # COMPATIBILITY PATCH
51
+ # ============================================================
52
+
53
+ print("=" * 80)
54
+ print("SenseNova U1.5 OOD hidden-state extraction")
55
+ print("=" * 80)
56
+
57
+ print("\nChecking Transformers compatibility patch...")
58
+
59
+ with open(
60
+ COMPAT_FILE,
61
+ "r",
62
+ encoding="utf-8"
63
+ ) as f:
64
+ text = f.read()
65
+
66
+ OLD = (
67
+ " from transformers.utils.generic import "
68
+ "check_model_inputs as model_input_compat"
69
+ )
70
+
71
+ NEW = (
72
+ " from transformers.utils.generic import check_model_inputs\n"
73
+ " def model_input_compat(func):\n"
74
+ " return check_model_inputs()(func)"
75
+ )
76
+
77
+ if OLD in text:
78
+
79
+ backup = COMPAT_FILE + ".bak"
80
+
81
+ if not os.path.exists(backup):
82
+ shutil.copy2(
83
+ COMPAT_FILE,
84
+ backup
85
+ )
86
+
87
+ text = text.replace(
88
+ OLD,
89
+ NEW
90
+ )
91
+
92
+ with open(
93
+ COMPAT_FILE,
94
+ "w",
95
+ encoding="utf-8"
96
+ ) as f:
97
+ f.write(text)
98
+
99
+ print("Compatibility patch applied.")
100
+
101
+ elif "return check_model_inputs()(func)" in text:
102
+
103
+ print("Compatibility patch already present.")
104
+
105
+ else:
106
+
107
+ print(
108
+ "WARNING: expected compatibility code "
109
+ "not found."
110
+ )
111
+
112
+
113
+ # ============================================================
114
+ # IMPORT LOCAL SENSENOVA CODE
115
+ # ============================================================
116
+
117
+ sys.path.insert(
118
+ 0,
119
+ SENSENOVA_SRC
120
+ )
121
+
122
+ from transformers import AutoTokenizer
123
+ from accelerate import (
124
+ init_empty_weights,
125
+ load_checkpoint_and_dispatch
126
+ )
127
+
128
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
129
+ NEOChatConfig
130
+ )
131
+
132
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
133
+ NEOChatModel
134
+ )
135
+
136
+
137
+ # ============================================================
138
+ # SYSTEM INFO
139
+ # ============================================================
140
+
141
+ if not torch.cuda.is_available():
142
+ raise RuntimeError(
143
+ "CUDA is required."
144
+ )
145
+
146
+ print("\nGPU:")
147
+ print(torch.cuda.get_device_name(0))
148
+
149
+ print(
150
+ "VRAM:",
151
+ round(
152
+ torch.cuda.get_device_properties(0).total_memory
153
+ / 1024**3,
154
+ 2
155
+ ),
156
+ "GB"
157
+ )
158
+
159
+
160
+ # ============================================================
161
+ # LOAD OOD PROMPTS
162
+ # ============================================================
163
+
164
+ print("\nLoading OOD prompts...")
165
+
166
+ with open(
167
+ PROMPTS_FILE,
168
+ "r",
169
+ encoding="utf-8"
170
+ ) as f:
171
+ prompts = json.load(f)
172
+
173
+ print("Prompt count:", len(prompts))
174
+
175
+ if len(prompts) != 160:
176
+ raise RuntimeError(
177
+ f"Expected 160 prompts, got {len(prompts)}"
178
+ )
179
+
180
+
181
+ # ============================================================
182
+ # TOKENIZER
183
+ # ============================================================
184
+
185
+ print("\nLoading tokenizer...")
186
+
187
+ tokenizer = AutoTokenizer.from_pretrained(
188
+ HF_REPO,
189
+ trust_remote_code=True,
190
+ )
191
+
192
+
193
+ # ============================================================
194
+ # CONFIG
195
+ # ============================================================
196
+
197
+ print("Loading SenseNova U1.5 config...")
198
+
199
+ config = NEOChatConfig.from_pretrained(
200
+ HF_REPO
201
+ )
202
+
203
+ config.llm_config._attn_implementation = "eager"
204
+
205
+ print("Config loaded.")
206
+ print(
207
+ "LLM hidden size:",
208
+ config.llm_config.hidden_size
209
+ )
210
+ print(
211
+ "LLM layers:",
212
+ config.llm_config.num_hidden_layers
213
+ )
214
+ print(
215
+ "Attention:",
216
+ config.llm_config._attn_implementation
217
+ )
218
+
219
+
220
+ # ============================================================
221
+ # CREATE MODEL ON META
222
+ # ============================================================
223
+
224
+ print()
225
+ print("Creating META architecture...")
226
+
227
+ with init_empty_weights():
228
+
229
+ model = NEOChatModel(
230
+ config
231
+ )
232
+
233
+ print("META architecture created.")
234
+
235
+
236
+ # ============================================================
237
+ # LOAD CHECKPOINT
238
+ # ============================================================
239
+
240
+ print()
241
+ print(
242
+ "Loading BF16 checkpoint with Accelerate..."
243
+ )
244
+
245
+ model = load_checkpoint_and_dispatch(
246
+ model,
247
+ checkpoint=CHECKPOINT,
248
+ device_map="auto",
249
+ max_memory={
250
+ 0: "19GiB",
251
+ "cpu": "70GiB",
252
+ },
253
+ dtype=torch.bfloat16,
254
+ no_split_module_classes=[
255
+ "Qwen3DecoderLayer",
256
+ "Qwen3MoeDecoderLayer",
257
+ ],
258
+ )
259
+
260
+ model.eval()
261
+
262
+ print("Checkpoint loaded successfully.")
263
+
264
+
265
+ # ============================================================
266
+ # LANGUAGE MODEL
267
+ # ============================================================
268
+
269
+ language_model = model.language_model
270
+ qwen = language_model.model
271
+ layers = qwen.layers
272
+
273
+ print()
274
+ print(
275
+ "Detected language layers:",
276
+ len(layers)
277
+ )
278
+
279
+ if LAYER >= len(layers):
280
+ raise RuntimeError(
281
+ f"Layer {LAYER} does not exist."
282
+ )
283
+
284
+
285
+ # ============================================================
286
+ # DEVICE MAP
287
+ # ============================================================
288
+
289
+ device_map = getattr(
290
+ model,
291
+ "hf_device_map",
292
+ {}
293
+ )
294
+
295
+ gpu_count = 0
296
+ cpu_count = 0
297
+ disk_count = 0
298
+
299
+ for _, dev in device_map.items():
300
+
301
+ s = str(dev)
302
+
303
+ if s in (
304
+ "0",
305
+ "cuda",
306
+ "cuda:0"
307
+ ):
308
+ gpu_count += 1
309
+
310
+ elif s == "cpu":
311
+ cpu_count += 1
312
+
313
+ elif s == "disk":
314
+ disk_count += 1
315
+
316
+
317
+ print()
318
+ print("Device map:")
319
+ print(" GPU modules :", gpu_count)
320
+ print(" CPU modules :", cpu_count)
321
+ print(" Disk modules:", disk_count)
322
+
323
+
324
+ # ============================================================
325
+ # FIND INPUT DEVICE
326
+ # ============================================================
327
+
328
+ INPUT_DEVICE = torch.device(
329
+ "cuda:0"
330
+ )
331
+
332
+ for name, dev in device_map.items():
333
+
334
+ if name.endswith(
335
+ "language_model.model.embed_tokens"
336
+ ):
337
+
338
+ if str(dev) == "cpu":
339
+ INPUT_DEVICE = torch.device(
340
+ "cpu"
341
+ )
342
+
343
+ else:
344
+ INPUT_DEVICE = torch.device(
345
+ "cuda:0"
346
+ )
347
+
348
+ break
349
+
350
+
351
+ print()
352
+ print("Input device:", INPUT_DEVICE)
353
+
354
+
355
+ # ============================================================
356
+ # HOOK ONLY LAYER 32
357
+ # ============================================================
358
+
359
+ captured = {}
360
+
361
+
362
+ def hook_fn(
363
+ module,
364
+ inputs,
365
+ output
366
+ ):
367
+
368
+ if isinstance(
369
+ output,
370
+ tuple
371
+ ):
372
+ hidden = output[0]
373
+
374
+ else:
375
+ hidden = output
376
+
377
+ if not torch.is_tensor(
378
+ hidden
379
+ ):
380
+ raise RuntimeError(
381
+ "Unexpected layer output type: "
382
+ f"{type(hidden)}"
383
+ )
384
+
385
+ captured["hidden"] = (
386
+ hidden.detach()
387
+ .to(
388
+ device="cpu",
389
+ dtype=torch.float16
390
+ )
391
+ .clone()
392
+ )
393
+
394
+
395
+ handle = (
396
+ layers[LAYER]
397
+ .register_forward_hook(
398
+ hook_fn
399
+ )
400
+ )
401
+
402
+ print()
403
+ print(
404
+ "Hook registered on "
405
+ f"SenseNova layer {LAYER}"
406
+ )
407
+
408
+
409
+ # ============================================================
410
+ # EXTRACTION
411
+ # ============================================================
412
+
413
+ results = []
414
+
415
+ print()
416
+ print("=" * 80)
417
+ print("STARTING EXTRACTION")
418
+ print("=" * 80)
419
+ print()
420
+
421
+ with torch.inference_mode():
422
+
423
+ for index, item in enumerate(
424
+ prompts
425
+ ):
426
+
427
+ prompt = item["prompt"]
428
+ category = item["category"]
429
+
430
+ encoded = tokenizer(
431
+ prompt,
432
+ return_tensors="pt",
433
+ add_special_tokens=False,
434
+ )
435
+
436
+ input_ids_cpu = (
437
+ encoded["input_ids"]
438
+ .detach()
439
+ .cpu()
440
+ )
441
+
442
+ input_ids = (
443
+ encoded["input_ids"]
444
+ .to(INPUT_DEVICE)
445
+ )
446
+
447
+ captured.clear()
448
+
449
+ # Same reset used in our successful V3.
450
+ if hasattr(
451
+ qwen,
452
+ "current_index"
453
+ ):
454
+ qwen.current_index = -1
455
+
456
+ outputs = qwen(
457
+ input_ids=input_ids,
458
+
459
+ # IMPORTANT:
460
+ # understanding branch, NOT MoT-gen.
461
+ image_gen_indicators=None,
462
+
463
+ attention_mask=None,
464
+ use_cache=False,
465
+ )
466
+
467
+ if "hidden" not in captured:
468
+
469
+ raise RuntimeError(
470
+ f"Layer hook failed "
471
+ f"at prompt {index}"
472
+ )
473
+
474
+ hidden = captured[
475
+ "hidden"
476
+ ]
477
+
478
+ # ====================================================
479
+ # VALIDATE
480
+ # ====================================================
481
+
482
+ if hidden.ndim != 3:
483
+
484
+ raise RuntimeError(
485
+ "Unexpected hidden shape "
486
+ f"at prompt {index}: "
487
+ f"{hidden.shape}"
488
+ )
489
+
490
+ if hidden.shape[-1] != 4096:
491
+
492
+ raise RuntimeError(
493
+ f"Expected hidden dim 4096, "
494
+ f"got {hidden.shape[-1]}"
495
+ )
496
+
497
+ if not torch.isfinite(
498
+ hidden.float()
499
+ ).all():
500
+
501
+ raise RuntimeError(
502
+ "Non-finite SenseNova "
503
+ f"hidden state at prompt {index}"
504
+ )
505
+
506
+
507
+ # ====================================================
508
+ # STORE
509
+ # ====================================================
510
+
511
+ results.append({
512
+ "prompt": prompt,
513
+ "category": category,
514
+ "input_ids": input_ids_cpu,
515
+ "hidden": hidden,
516
+ })
517
+
518
+
519
+ # ====================================================
520
+ # PROGRESS
521
+ # ====================================================
522
+
523
+ if (
524
+ index == 0
525
+ or (index + 1) % 10 == 0
526
+ or index + 1 == len(prompts)
527
+ ):
528
+
529
+ print(
530
+ f"{index + 1:3d}/"
531
+ f"{len(prompts)} | "
532
+ f"tokens={hidden.shape[1]:3d} | "
533
+ f"shape={tuple(hidden.shape)}"
534
+ )
535
+
536
+
537
+ del outputs
538
+ del input_ids
539
+
540
+ if (
541
+ index + 1
542
+ ) % 20 == 0:
543
+
544
+ gc.collect()
545
+ torch.cuda.empty_cache()
546
+
547
+
548
+ # ============================================================
549
+ # REMOVE HOOK
550
+ # ============================================================
551
+
552
+ handle.remove()
553
+
554
+
555
+ # ============================================================
556
+ # FINAL CROSS-CHECK
557
+ # ============================================================
558
+
559
+ if len(results) != 160:
560
+
561
+ raise RuntimeError(
562
+ "Wrong number of extracted prompts: "
563
+ f"{len(results)}"
564
+ )
565
+
566
+
567
+ for i, result in enumerate(
568
+ results
569
+ ):
570
+
571
+ expected = prompts[i]
572
+
573
+ if (
574
+ result["prompt"]
575
+ != expected["prompt"]
576
+ ):
577
+
578
+ raise RuntimeError(
579
+ f"Prompt mismatch at {i}"
580
+ )
581
+
582
+
583
+ # ============================================================
584
+ # SAVE
585
+ # ============================================================
586
+
587
+ print()
588
+ print("Saving OOD hidden states...")
589
+
590
+ payload = {
591
+ "model": (
592
+ "SenseNova-U1.5-8B-MoT"
593
+ ),
594
+ "branch": "understanding",
595
+ "layer": LAYER,
596
+ "hidden_dim": 4096,
597
+ "prompt_file": os.path.basename(
598
+ PROMPTS_FILE
599
+ ),
600
+ "results": results,
601
+ }
602
+
603
+ torch.save(
604
+ payload,
605
+ OUT_FILE
606
+ )
607
+
608
+
609
+ print()
610
+ print("=" * 80)
611
+ print("SUCCESS")
612
+ print("=" * 80)
613
+
614
+ print(
615
+ "Prompts:",
616
+ len(results)
617
+ )
618
+
619
+ print(
620
+ "Layer:",
621
+ LAYER
622
+ )
623
+
624
+ print(
625
+ "Branch: understanding"
626
+ )
627
+
628
+ print("\nOutput:")
629
+ print(OUT_FILE)
630
+
631
+
632
+ gc.collect()
633
+ torch.cuda.empty_cache()
research/raw_scripts/inspect_distillation_data.py ADDED
@@ -0,0 +1,466 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import torch
3
+ from collections import Counter, defaultdict
4
+
5
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
6
+
7
+ H3_FILE = os.path.join(
8
+ ROOT,
9
+ "h3_bridge_hidden_480.pt"
10
+ )
11
+
12
+ SN_FILE = os.path.join(
13
+ ROOT,
14
+ "sensenova_bridge_hidden_480.pt"
15
+ )
16
+
17
+ BRIDGE_FILE = (
18
+ r"D:\ComfyUI_Krea2\ComfyUI\models\bridge"
19
+ r"\SN_L32_to_H3_L49_rank128.safetensors"
20
+ )
21
+
22
+ REPORT_FILE = os.path.join(
23
+ ROOT,
24
+ "distillation_source_inspection.txt"
25
+ )
26
+
27
+
28
+ # ============================================================
29
+ # HELPERS
30
+ # ============================================================
31
+
32
+ def tensor_info(t):
33
+ return (
34
+ f"shape={tuple(t.shape)} "
35
+ f"dtype={t.dtype} "
36
+ f"device={t.device}"
37
+ )
38
+
39
+
40
+ def walk(obj, path="root", out=None, depth=0, max_depth=8):
41
+ if out is None:
42
+ out = []
43
+
44
+ if depth > max_depth:
45
+ return out
46
+
47
+ if torch.is_tensor(obj):
48
+ out.append({
49
+ "path": path,
50
+ "shape": tuple(obj.shape),
51
+ "dtype": str(obj.dtype),
52
+ "tensor": obj,
53
+ })
54
+ return out
55
+
56
+ if isinstance(obj, dict):
57
+ for key, value in obj.items():
58
+ walk(
59
+ value,
60
+ f"{path}.{key}",
61
+ out,
62
+ depth + 1,
63
+ max_depth,
64
+ )
65
+
66
+ elif isinstance(obj, (list, tuple)):
67
+ for i, value in enumerate(obj):
68
+ walk(
69
+ value,
70
+ f"{path}[{i}]",
71
+ out,
72
+ depth + 1,
73
+ max_depth,
74
+ )
75
+
76
+ return out
77
+
78
+
79
+ def print_structure(obj, name, lines, depth=0, max_depth=4, max_items=4):
80
+ indent = " " * depth
81
+
82
+ if depth > max_depth:
83
+ lines.append(f"{indent}...")
84
+ return
85
+
86
+ if torch.is_tensor(obj):
87
+ lines.append(
88
+ f"{indent}{tensor_info(obj)}"
89
+ )
90
+ return
91
+
92
+ if isinstance(obj, dict):
93
+ lines.append(
94
+ f"{indent}dict keys={len(obj)}"
95
+ )
96
+
97
+ for i, (key, value) in enumerate(obj.items()):
98
+ if i >= max_items:
99
+ lines.append(
100
+ f"{indent} ... ({len(obj) - max_items} more)"
101
+ )
102
+ break
103
+
104
+ lines.append(
105
+ f"{indent} KEY: {repr(key)}"
106
+ )
107
+
108
+ print_structure(
109
+ value,
110
+ name,
111
+ lines,
112
+ depth + 2,
113
+ max_depth,
114
+ max_items,
115
+ )
116
+
117
+ elif isinstance(obj, (list, tuple)):
118
+ lines.append(
119
+ f"{indent}{type(obj).__name__} len={len(obj)}"
120
+ )
121
+
122
+ for i, value in enumerate(obj[:max_items]):
123
+ lines.append(
124
+ f"{indent} ITEM [{i}]"
125
+ )
126
+
127
+ print_structure(
128
+ value,
129
+ name,
130
+ lines,
131
+ depth + 2,
132
+ max_depth,
133
+ max_items,
134
+ )
135
+
136
+ if len(obj) > max_items:
137
+ lines.append(
138
+ f"{indent} ... ({len(obj) - max_items} more)"
139
+ )
140
+
141
+ else:
142
+ value_repr = repr(obj)
143
+
144
+ if len(value_repr) > 180:
145
+ value_repr = value_repr[:180] + "..."
146
+
147
+ lines.append(
148
+ f"{indent}{type(obj).__name__}: {value_repr}"
149
+ )
150
+
151
+
152
+ def analyze_file(label, path, lines):
153
+ lines.append("")
154
+ lines.append("=" * 100)
155
+ lines.append(label)
156
+ lines.append("=" * 100)
157
+ lines.append(f"Path: {path}")
158
+ lines.append(
159
+ f"Size: {os.path.getsize(path) / (1024 ** 2):.2f} MiB"
160
+ )
161
+
162
+ print(f"\nLoading {label}...")
163
+ data = torch.load(
164
+ path,
165
+ map_location="cpu",
166
+ weights_only=False,
167
+ )
168
+
169
+ print(f"{label} loaded.")
170
+
171
+ lines.append("")
172
+ lines.append("TOP-LEVEL STRUCTURE")
173
+ lines.append("-" * 100)
174
+
175
+ print_structure(
176
+ data,
177
+ label,
178
+ lines,
179
+ )
180
+
181
+ tensors = walk(data)
182
+
183
+ lines.append("")
184
+ lines.append(
185
+ f"Total tensors found recursively: {len(tensors)}"
186
+ )
187
+
188
+ shape_counter = Counter(
189
+ item["shape"]
190
+ for item in tensors
191
+ )
192
+
193
+ lines.append("")
194
+ lines.append("MOST COMMON TENSOR SHAPES")
195
+ lines.append("-" * 100)
196
+
197
+ for shape, count in shape_counter.most_common(20):
198
+ lines.append(
199
+ f"{count:6d} x {shape}"
200
+ )
201
+
202
+ # ========================================================
203
+ # HIDDEN-STATE CANDIDATES
204
+ # ========================================================
205
+
206
+ candidates_4096 = []
207
+ candidates_5120 = []
208
+
209
+ for item in tensors:
210
+ shape = item["shape"]
211
+
212
+ if len(shape) >= 1:
213
+ if shape[-1] == 4096:
214
+ candidates_4096.append(item)
215
+
216
+ if shape[-1] == 5120:
217
+ candidates_5120.append(item)
218
+
219
+ lines.append("")
220
+ lines.append("4096-DIM CANDIDATES")
221
+ lines.append("-" * 100)
222
+ lines.append(
223
+ f"Count: {len(candidates_4096)}"
224
+ )
225
+
226
+ for item in candidates_4096[:30]:
227
+ lines.append(
228
+ f"{item['path']} | "
229
+ f"{item['shape']} | "
230
+ f"{item['dtype']}"
231
+ )
232
+
233
+ if len(candidates_4096) > 30:
234
+ lines.append(
235
+ f"... {len(candidates_4096) - 30} more"
236
+ )
237
+
238
+ lines.append("")
239
+ lines.append("5120-DIM CANDIDATES")
240
+ lines.append("-" * 100)
241
+ lines.append(
242
+ f"Count: {len(candidates_5120)}"
243
+ )
244
+
245
+ for item in candidates_5120[:30]:
246
+ lines.append(
247
+ f"{item['path']} | "
248
+ f"{item['shape']} | "
249
+ f"{item['dtype']}"
250
+ )
251
+
252
+ if len(candidates_5120) > 30:
253
+ lines.append(
254
+ f"... {len(candidates_5120) - 30} more"
255
+ )
256
+
257
+ # ========================================================
258
+ # PATHS THAT LOOK LIKE OUR TARGET LAYERS
259
+ # ========================================================
260
+
261
+ lines.append("")
262
+ lines.append("PATHS CONTAINING LAYER 32 / L32")
263
+ lines.append("-" * 100)
264
+
265
+ layer32 = [
266
+ item
267
+ for item in tensors
268
+ if (
269
+ "32" in item["path"].lower()
270
+ or "l32" in item["path"].lower()
271
+ or "layer_32" in item["path"].lower()
272
+ or "layer32" in item["path"].lower()
273
+ )
274
+ ]
275
+
276
+ lines.append(
277
+ f"Count: {len(layer32)}"
278
+ )
279
+
280
+ for item in layer32[:40]:
281
+ lines.append(
282
+ f"{item['path']} | "
283
+ f"{item['shape']}"
284
+ )
285
+
286
+ lines.append("")
287
+ lines.append("PATHS CONTAINING LAYER 49 / L49")
288
+ lines.append("-" * 100)
289
+
290
+ layer49 = [
291
+ item
292
+ for item in tensors
293
+ if (
294
+ "49" in item["path"].lower()
295
+ or "l49" in item["path"].lower()
296
+ or "layer_49" in item["path"].lower()
297
+ or "layer49" in item["path"].lower()
298
+ )
299
+ ]
300
+
301
+ lines.append(
302
+ f"Count: {len(layer49)}"
303
+ )
304
+
305
+ for item in layer49[:40]:
306
+ lines.append(
307
+ f"{item['path']} | "
308
+ f"{item['shape']}"
309
+ )
310
+
311
+ return data, tensors
312
+
313
+
314
+ # ============================================================
315
+ # MAIN
316
+ # ============================================================
317
+
318
+ print("=" * 100)
319
+ print("DISTILLATION SOURCE INSPECTION")
320
+ print("=" * 100)
321
+
322
+ for path in [
323
+ H3_FILE,
324
+ SN_FILE,
325
+ BRIDGE_FILE,
326
+ ]:
327
+ if not os.path.exists(path):
328
+ raise FileNotFoundError(
329
+ f"Missing required file:\n{path}"
330
+ )
331
+
332
+ print("H3:")
333
+ print(H3_FILE)
334
+
335
+ print("\nSenseNova:")
336
+ print(SN_FILE)
337
+
338
+ print("\nBridge:")
339
+ print(BRIDGE_FILE)
340
+
341
+ lines = []
342
+
343
+ lines.append(
344
+ "SenseNova -> MiniMax H3 Distillation Source Inspection"
345
+ )
346
+
347
+ lines.append(
348
+ "=" * 100
349
+ )
350
+
351
+ h3_data, h3_tensors = analyze_file(
352
+ "MINIMAX H3 DATA",
353
+ H3_FILE,
354
+ lines,
355
+ )
356
+
357
+ sn_data, sn_tensors = analyze_file(
358
+ "SENSENOVA DATA",
359
+ SN_FILE,
360
+ lines,
361
+ )
362
+
363
+
364
+ # ============================================================
365
+ # BRIDGE
366
+ # ============================================================
367
+
368
+ lines.append("")
369
+ lines.append("=" * 100)
370
+ lines.append("BRIDGE FILE")
371
+ lines.append("=" * 100)
372
+ lines.append(f"Path: {BRIDGE_FILE}")
373
+
374
+ from safetensors.torch import load_file
375
+
376
+ bridge = load_file(
377
+ BRIDGE_FILE,
378
+ device="cpu",
379
+ )
380
+
381
+ for key, tensor in bridge.items():
382
+ lines.append(
383
+ f"{key:40s} "
384
+ f"{tuple(tensor.shape)} "
385
+ f"{tensor.dtype}"
386
+ )
387
+
388
+
389
+ # ============================================================
390
+ # SUMMARY COUNTS
391
+ # ============================================================
392
+
393
+ h3_5120 = [
394
+ x for x in h3_tensors
395
+ if (
396
+ len(x["shape"]) > 0
397
+ and x["shape"][-1] == 5120
398
+ )
399
+ ]
400
+
401
+ sn_4096 = [
402
+ x for x in sn_tensors
403
+ if (
404
+ len(x["shape"]) > 0
405
+ and x["shape"][-1] == 4096
406
+ )
407
+ ]
408
+
409
+ lines.append("")
410
+ lines.append("=" * 100)
411
+ lines.append("FINAL SUMMARY")
412
+ lines.append("=" * 100)
413
+
414
+ lines.append(
415
+ f"H3 5120-dim tensors found: {len(h3_5120)}"
416
+ )
417
+
418
+ lines.append(
419
+ f"SenseNova 4096-dim tensors found: {len(sn_4096)}"
420
+ )
421
+
422
+ lines.append("")
423
+ lines.append(
424
+ "The next script will use this report to identify the exact "
425
+ "H3 input representation and SenseNova teacher target without "
426
+ "guessing the .pt layout."
427
+ )
428
+
429
+
430
+ # ============================================================
431
+ # SAVE
432
+ # ============================================================
433
+
434
+ with open(
435
+ REPORT_FILE,
436
+ "w",
437
+ encoding="utf-8",
438
+ ) as f:
439
+ f.write(
440
+ "\n".join(lines)
441
+ )
442
+
443
+ print()
444
+ print("=" * 100)
445
+ print("SUCCESS")
446
+ print("=" * 100)
447
+
448
+ print("Report saved:")
449
+ print(REPORT_FILE)
450
+
451
+ print()
452
+ print(
453
+ "H3 5120-dim tensors:",
454
+ len(h3_5120)
455
+ )
456
+
457
+ print(
458
+ "SenseNova 4096-dim tensors:",
459
+ len(sn_4096)
460
+ )
461
+
462
+ print()
463
+ print(
464
+ "Send me the contents of "
465
+ "distillation_source_inspection.txt"
466
+ )
research/raw_scripts/inspect_h3_sensenova_bridge.py ADDED
@@ -0,0 +1,274 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from safetensors import safe_open
2
+ from collections import Counter
3
+ import os
4
+ import re
5
+
6
+ H3_TE = r"D:\ComfyUI_Python312\ComfyUI\models\text_encoders\qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
7
+ SENSENOVA = r"D:\ComfyUI_Python312\ComfyUI\models\unet\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
8
+
9
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\h3_sensenova_bridge_report.txt"
10
+
11
+
12
+ def inspect(path):
13
+ data = {}
14
+ meta = {}
15
+
16
+ with safe_open(path, framework="pt", device="cpu") as f:
17
+ meta = f.metadata() or {}
18
+
19
+ for k in f.keys():
20
+ s = f.get_slice(k)
21
+ shape = tuple(s.get_shape())
22
+
23
+ data[k] = {
24
+ "shape": shape,
25
+ "ndim": len(shape),
26
+ }
27
+
28
+ return data, meta
29
+
30
+
31
+ def get_layer_number(k):
32
+ patterns = [
33
+ r"layers\.(\d+)",
34
+ r"blocks\.(\d+)",
35
+ r"transformer_blocks\.(\d+)",
36
+ ]
37
+
38
+ for p in patterns:
39
+ m = re.search(p, k)
40
+ if m:
41
+ return int(m.group(1))
42
+
43
+ return None
44
+
45
+
46
+ def likely_embedding(k):
47
+ x = k.lower()
48
+ return any(s in x for s in [
49
+ "embed_tokens",
50
+ "token_embedding",
51
+ "word_embeddings",
52
+ "tok_embeddings",
53
+ ])
54
+
55
+
56
+ def likely_norm(k):
57
+ return "norm" in k.lower()
58
+
59
+
60
+ def likely_attention(k):
61
+ x = k.lower()
62
+ return any(s in x for s in [
63
+ "q_proj", "k_proj", "v_proj", "o_proj",
64
+ "qkv", "attention", "attn"
65
+ ])
66
+
67
+
68
+ print("Reading H3 TE...")
69
+ h3, h3_meta = inspect(H3_TE)
70
+
71
+ print("Reading SenseNova...")
72
+ sn, sn_meta = inspect(SENSENOVA)
73
+
74
+
75
+ with open(OUT, "w", encoding="utf-8") as f:
76
+
77
+ def w(x=""):
78
+ f.write(str(x) + "\n")
79
+
80
+ w("=" * 100)
81
+ w("H3 TE <-> SENSENOVA BRIDGE DIAGNOSTIC")
82
+ w("=" * 100)
83
+ w()
84
+
85
+ w("FILES")
86
+ w("-" * 100)
87
+ w(f"H3 TE: {H3_TE}")
88
+ w(f"Size: {os.path.getsize(H3_TE) / 1024**3:.3f} GB")
89
+ w(f"Tensors: {len(h3)}")
90
+ w()
91
+ w(f"SenseNova: {SENSENOVA}")
92
+ w(f"Size: {os.path.getsize(SENSENOVA) / 1024**3:.3f} GB")
93
+ w(f"Tensors: {len(sn)}")
94
+ w()
95
+
96
+ # -------------------------------------------------------
97
+ # H3 dimensions
98
+ # -------------------------------------------------------
99
+
100
+ w("=" * 100)
101
+ w("H3 TE MATRICES CONTAINING DIMENSION 5120")
102
+ w("=" * 100)
103
+
104
+ h3_5120 = []
105
+
106
+ for k, v in h3.items():
107
+ if 5120 in v["shape"]:
108
+ h3_5120.append((k, v["shape"]))
109
+ w(f"{k} | {v['shape']}")
110
+
111
+ w()
112
+ w(f"COUNT: {len(h3_5120)}")
113
+ w()
114
+
115
+ # -------------------------------------------------------
116
+ # SenseNova dimensions
117
+ # -------------------------------------------------------
118
+
119
+ w("=" * 100)
120
+ w("SENSENOVA MATRICES CONTAINING DIMENSION 4096")
121
+ w("=" * 100)
122
+
123
+ sn_4096 = []
124
+
125
+ for k, v in sn.items():
126
+ if 4096 in v["shape"]:
127
+ sn_4096.append((k, v["shape"]))
128
+ w(f"{k} | {v['shape']}")
129
+
130
+ w()
131
+ w(f"COUNT: {len(sn_4096)}")
132
+ w()
133
+
134
+ # -------------------------------------------------------
135
+ # Embeddings
136
+ # -------------------------------------------------------
137
+
138
+ w("=" * 100)
139
+ w("LIKELY TOKEN EMBEDDINGS")
140
+ w("=" * 100)
141
+
142
+ w("H3:")
143
+ for k, v in h3.items():
144
+ if likely_embedding(k):
145
+ w(f" {k} | {v['shape']}")
146
+
147
+ w()
148
+ w("SENSENOVA:")
149
+ for k, v in sn.items():
150
+ if likely_embedding(k):
151
+ w(f" {k} | {v['shape']}")
152
+
153
+ w()
154
+
155
+ # -------------------------------------------------------
156
+ # Layer counts
157
+ # -------------------------------------------------------
158
+
159
+ h3_layers = sorted(set(
160
+ x for x in (get_layer_number(k) for k in h3)
161
+ if x is not None
162
+ ))
163
+
164
+ sn_layers = sorted(set(
165
+ x for x in (get_layer_number(k) for k in sn)
166
+ if x is not None
167
+ ))
168
+
169
+ w("=" * 100)
170
+ w("DETECTED LAYERS")
171
+ w("=" * 100)
172
+
173
+ w(f"H3 layer indexes ({len(h3_layers)}):")
174
+ w(str(h3_layers))
175
+
176
+ w()
177
+
178
+ w(f"SenseNova layer indexes ({len(sn_layers)}):")
179
+ w(str(sn_layers))
180
+
181
+ w()
182
+
183
+ # -------------------------------------------------------
184
+ # Selected H3 late layers
185
+ # -------------------------------------------------------
186
+
187
+ w("=" * 100)
188
+ w("H3 LAST DETECTED LANGUAGE-LIKE LAYERS")
189
+ w("=" * 100)
190
+
191
+ if h3_layers:
192
+ selected = set(h3_layers[-5:])
193
+
194
+ for k, v in h3.items():
195
+ n = get_layer_number(k)
196
+ if n in selected:
197
+ w(f"{k} | {v['shape']}")
198
+
199
+ w()
200
+
201
+ # -------------------------------------------------------
202
+ # SenseNova late layers
203
+ # -------------------------------------------------------
204
+
205
+ w("=" * 100)
206
+ w("SENSENOVA LAST DETECTED LANGUAGE-LIKE LAYERS")
207
+ w("=" * 100)
208
+
209
+ if sn_layers:
210
+ selected = set(sn_layers[-5:])
211
+
212
+ for k, v in sn.items():
213
+ n = get_layer_number(k)
214
+ if n in selected:
215
+ w(f"{k} | {v['shape']}")
216
+
217
+ w()
218
+
219
+ # -------------------------------------------------------
220
+ # Attention candidates
221
+ # -------------------------------------------------------
222
+
223
+ w("=" * 100)
224
+ w("H3 ATTENTION-LIKE MATRICES WITH 5120")
225
+ w("=" * 100)
226
+
227
+ for k, v in h3.items():
228
+ if v["ndim"] == 2 and 5120 in v["shape"] and likely_attention(k):
229
+ w(f"{k} | {v['shape']}")
230
+
231
+ w()
232
+
233
+ w("=" * 100)
234
+ w("SENSENOVA ATTENTION-LIKE MATRICES WITH 4096")
235
+ w("=" * 100)
236
+
237
+ for k, v in sn.items():
238
+ if v["ndim"] == 2 and 4096 in v["shape"] and likely_attention(k):
239
+ w(f"{k} | {v['shape']}")
240
+
241
+ w()
242
+
243
+ # -------------------------------------------------------
244
+ # Metadata
245
+ # -------------------------------------------------------
246
+
247
+ w("=" * 100)
248
+ w("H3 METADATA")
249
+ w("=" * 100)
250
+
251
+ for k, v in h3_meta.items():
252
+ if k == "_quantization_metadata":
253
+ w("_quantization_metadata = [present, omitted because very large]")
254
+ else:
255
+ w(f"{k} = {v}")
256
+
257
+ w()
258
+
259
+ w("=" * 100)
260
+ w("SENSENOVA METADATA")
261
+ w("=" * 100)
262
+
263
+ for k, v in sn_meta.items():
264
+ w(f"{k} = {v}")
265
+
266
+
267
+ print()
268
+ print("DONE")
269
+ print(OUT)
270
+ print()
271
+ print("H3 tensors containing 5120:", len(h3_5120))
272
+ print("SenseNova tensors containing 4096:", len(sn_4096))
273
+ print("H3 detected layers:", len(h3_layers))
274
+ print("SenseNova detected layers:", len(sn_layers))
research/raw_scripts/make_bridge_ood_prompts.py ADDED
@@ -0,0 +1,391 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import random
4
+
5
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
6
+
7
+ TRAIN_FILE = os.path.join(
8
+ ROOT,
9
+ "bridge_prompts_480.json"
10
+ )
11
+
12
+ OUT = os.path.join(
13
+ ROOT,
14
+ "bridge_ood_prompts_160.json"
15
+ )
16
+
17
+ random.seed(20260903)
18
+
19
+ # ============================================================
20
+ # LOAD TRAINING PROMPTS
21
+ # ============================================================
22
+
23
+ with open(TRAIN_FILE, "r", encoding="utf-8") as f:
24
+ train_data = json.load(f)
25
+
26
+ train_prompts = {
27
+ x["prompt"].strip().lower()
28
+ for x in train_data
29
+ }
30
+
31
+ data = []
32
+ used = set()
33
+
34
+
35
+ def add(category, prompt):
36
+ prompt = " ".join(prompt.strip().split())
37
+ key = prompt.lower()
38
+
39
+ if key in used:
40
+ return False
41
+
42
+ if key in train_prompts:
43
+ return False
44
+
45
+ used.add(key)
46
+
47
+ data.append({
48
+ "category": category,
49
+ "prompt": prompt,
50
+ })
51
+
52
+ return True
53
+
54
+
55
+ # ============================================================
56
+ # VOCABULARY
57
+ # ============================================================
58
+
59
+ people = [
60
+ "a young woman",
61
+ "an elderly man",
62
+ "a tall man",
63
+ "a woman wearing a dark jacket",
64
+ "a seated woman",
65
+ "a man wearing a gray coat",
66
+ "a woman with short hair",
67
+ "an elderly woman",
68
+ ]
69
+
70
+ objects = [
71
+ "glass bottle",
72
+ "ceramic cup",
73
+ "metal box",
74
+ "wooden chair",
75
+ "blue book",
76
+ "silver sphere",
77
+ "transparent cube",
78
+ "black vase",
79
+ "small lamp",
80
+ "red cylinder",
81
+ ]
82
+
83
+ places = [
84
+ "modern kitchen",
85
+ "railway platform",
86
+ "glass office lobby",
87
+ "small workshop",
88
+ "hotel entrance",
89
+ "underground station",
90
+ "minimalist living room",
91
+ "industrial studio",
92
+ ]
93
+
94
+ materials = [
95
+ "polished chrome",
96
+ "frosted glass",
97
+ "brushed aluminum",
98
+ "wet black stone",
99
+ "translucent amber acrylic",
100
+ "glossy ceramic",
101
+ "rough concrete",
102
+ "dark polished wood",
103
+ ]
104
+
105
+ lightings = [
106
+ "warm light entering from the left",
107
+ "cool light entering from the right",
108
+ "strong backlight",
109
+ "a narrow spotlight from above",
110
+ "soft diffused window light",
111
+ "hard side lighting",
112
+ "warm reflected light from below",
113
+ "two opposing light sources",
114
+ ]
115
+
116
+ texts = [
117
+ "NORTH GATE",
118
+ "ROOM 314",
119
+ "PLATFORM 11",
120
+ "WEST EXIT",
121
+ "OPEN 24 HOURS",
122
+ "STUDIO C",
123
+ "RIVER HOTEL",
124
+ "CAFE LEVEL 2",
125
+ "AUTHORIZED ENTRY",
126
+ "FINAL STOP",
127
+ ]
128
+
129
+
130
+ # ============================================================
131
+ # 1. COMPLEX SPATIAL — 20 unique
132
+ # ============================================================
133
+
134
+ for i in range(20):
135
+ a = objects[i % len(objects)]
136
+ b = objects[(i + 3) % len(objects)]
137
+ c = objects[(i + 6) % len(objects)]
138
+
139
+ add(
140
+ "complex_spatial",
141
+ f"A {a} stands behind a {b}, while a {c} is positioned "
142
+ f"to their {'left' if i % 2 == 0 else 'right'}. "
143
+ f"The {b} partially overlaps the {a} from the camera viewpoint. "
144
+ f"Exactly {3 + i % 4} major objects are visible in the composition. "
145
+ f"Spatial variant {i + 1}."
146
+ )
147
+
148
+
149
+ # ============================================================
150
+ # 2. ANATOMY / INTERACTION — 20 unique
151
+ # ============================================================
152
+
153
+ actions = [
154
+ "raises the left hand while holding a cup in the right hand",
155
+ "crosses the right leg over the left while looking over the left shoulder",
156
+ "reaches forward with the right arm while keeping the left hand behind the back",
157
+ "holds a bottle with both hands directly in front of the chest",
158
+ "turns the torso right while the head remains facing left",
159
+ ]
160
+
161
+ for i in range(20):
162
+ p1 = people[i % len(people)]
163
+ p2 = people[(i + 3) % len(people)]
164
+ action = actions[i % len(actions)]
165
+
166
+ add(
167
+ "complex_anatomy",
168
+ f"{p1.capitalize()} {action}. "
169
+ f"{p2.capitalize()} stands "
170
+ f"{'behind' if i % 2 else 'beside'} them "
171
+ f"and points toward the object with the "
172
+ f"{'left' if i % 3 else 'right'} hand. "
173
+ f"Both people's hands and feet remain visible. "
174
+ f"Anatomy variant {i + 1}."
175
+ )
176
+
177
+
178
+ # ============================================================
179
+ # 3. MATERIAL + LIGHT — 20 unique
180
+ # ============================================================
181
+
182
+ for i in range(20):
183
+ material1 = materials[i % len(materials)]
184
+ material2 = materials[(i + 3) % len(materials)]
185
+ lighting = lightings[i % len(lightings)]
186
+ obj = objects[(i + 4) % len(objects)]
187
+
188
+ add(
189
+ "complex_material_light",
190
+ f"A {obj} made from {material1} rests on a surface made from "
191
+ f"{material2}, illuminated by {lighting}. "
192
+ f"The scene clearly shows reflection, roughness, transparency, "
193
+ f"or subsurface behavior appropriate to both materials. "
194
+ f"A small colored object is reflected near the edge of the surface. "
195
+ f"Material-light variant {i + 1}."
196
+ )
197
+
198
+
199
+ # ============================================================
200
+ # 4. TEXT — 20 unique
201
+ # ============================================================
202
+
203
+ for i in range(20):
204
+ phrase = texts[i % len(texts)]
205
+ place = places[(i + 2) % len(places)]
206
+
207
+ add(
208
+ "complex_text",
209
+ f'Inside a {place}, the exact phrase "{phrase}" is printed clearly '
210
+ f'on a rectangular sign. A person stands partly in front of the sign '
211
+ f'without covering any letters. A reflective surface beside the sign '
212
+ f'shows a reversed reflection of part of the scene while the original '
213
+ f'text remains completely readable. Scene variant {i + 1}.'
214
+ )
215
+
216
+
217
+ # ============================================================
218
+ # 5. REFLECTION / OCCLUSION — 20 unique
219
+ # ============================================================
220
+
221
+ for i in range(20):
222
+ obj1 = objects[i % len(objects)]
223
+ obj2 = objects[(i + 5) % len(objects)]
224
+ person = people[(i + 2) % len(people)]
225
+
226
+ add(
227
+ "reflection_occlusion_ood",
228
+ f"{person.capitalize()} stands behind a transparent glass panel. "
229
+ f"A {obj1} in the foreground partially occludes the torso while "
230
+ f"the face and both hands remain visible. "
231
+ f"A mirror behind the person reflects a {obj2} located outside "
232
+ f"the direct camera frame. "
233
+ f"The glass also contains a faint reflection of an overhead light "
234
+ f"at position {i + 1}. "
235
+ f"Reflection variant {i + 1}."
236
+ )
237
+
238
+
239
+ # ============================================================
240
+ # 6. ARCHITECTURE / VEHICLES — 20 unique
241
+ # ============================================================
242
+
243
+ vehicles = [
244
+ "red compact car",
245
+ "white city bus",
246
+ "black motorcycle",
247
+ "blue bicycle",
248
+ "silver tram",
249
+ ]
250
+
251
+ for i in range(20):
252
+ vehicle = vehicles[i % len(vehicles)]
253
+ place = places[(i + 1) % len(places)]
254
+
255
+ add(
256
+ "architecture_vehicle",
257
+ f"A {vehicle} passes beside a {place}. "
258
+ f"A staircase rises on the "
259
+ f"{'left' if i % 2 == 0 else 'right'} side "
260
+ f"and turns at an upper landing. "
261
+ f"Three architectural openings are visible at different depths, "
262
+ f"and a glass facade reflects a building across the street. "
263
+ f"Composition variant {i + 1}."
264
+ )
265
+
266
+
267
+ # ============================================================
268
+ # 7. COUNTING — 20 unique
269
+ # ============================================================
270
+
271
+ for i in range(20):
272
+ n1 = 3 + (i % 5)
273
+ n2 = 2 + ((i * 2) % 4)
274
+
275
+ obj1 = objects[i % len(objects)]
276
+ obj2 = objects[(i + 4) % len(objects)]
277
+
278
+ add(
279
+ "complex_counting",
280
+ f"Exactly {n1} {obj1}s form "
281
+ f"{'a curved row' if i % 2 else 'two staggered rows'}, "
282
+ f"while exactly {n2} {obj2}s occupy the background. "
283
+ f"None of the objects is completely hidden, and one small yellow marker "
284
+ f"is located beside object number {(i % n1) + 1}. "
285
+ f"Counting variant {i + 1}."
286
+ )
287
+
288
+
289
+ # ============================================================
290
+ # 8. LONG MULTI-CONSTRAINT — 20 unique
291
+ # ============================================================
292
+
293
+ for i in range(20):
294
+ p1 = people[i % len(people)]
295
+ p2 = people[(i + 4) % len(people)]
296
+
297
+ obj1 = objects[i % len(objects)]
298
+ obj2 = objects[(i + 2) % len(objects)]
299
+ obj3 = objects[(i + 7) % len(objects)]
300
+
301
+ phrase = texts[i % len(texts)]
302
+ lighting = lightings[i % len(lightings)]
303
+
304
+ add(
305
+ "long_composition",
306
+ f"{p1.capitalize()} stands at the left side of a glass table while "
307
+ f"{p2} sits to the right. "
308
+ f"The standing person holds a {obj1} in the left hand and points "
309
+ f'toward a sign reading exactly "{phrase}" with the right hand. '
310
+ f"A {obj2} and a {obj3} lie on the table in that order from left to right. "
311
+ f"The scene is illuminated by {lighting}. "
312
+ f"A large mirror behind both people reflects the seated person's back "
313
+ f"and one object outside the direct camera frame. "
314
+ f"Exactly {5 + i % 4} major foreground objects should remain clearly "
315
+ f"distinguishable. Long-scene variant {i + 1}."
316
+ )
317
+
318
+
319
+ # ============================================================
320
+ # FINAL CHECKS
321
+ # ============================================================
322
+
323
+ if len(data) != 160:
324
+ raise RuntimeError(
325
+ f"Expected exactly 160 unique prompts, got {len(data)}"
326
+ )
327
+
328
+ random.shuffle(data)
329
+
330
+ prompts = [
331
+ x["prompt"].strip().lower()
332
+ for x in data
333
+ ]
334
+
335
+ if len(prompts) != len(set(prompts)):
336
+ raise RuntimeError(
337
+ "Duplicate OOD prompts detected."
338
+ )
339
+
340
+ overlap = set(prompts) & train_prompts
341
+
342
+ if overlap:
343
+ raise RuntimeError(
344
+ f"OOD/train overlap detected: {len(overlap)} prompts"
345
+ )
346
+
347
+
348
+ # ============================================================
349
+ # SAVE
350
+ # ============================================================
351
+
352
+ with open(
353
+ OUT,
354
+ "w",
355
+ encoding="utf-8"
356
+ ) as f:
357
+
358
+ json.dump(
359
+ data,
360
+ f,
361
+ ensure_ascii=False,
362
+ indent=2
363
+ )
364
+
365
+
366
+ # ============================================================
367
+ # REPORT
368
+ # ============================================================
369
+
370
+ print("=" * 80)
371
+ print("STRICT OOD DATASET CREATED")
372
+ print("=" * 80)
373
+
374
+ print("Total prompts: ", len(data))
375
+ print("Unique prompts: ", len(set(prompts)))
376
+ print("Train/OOD overlap: ", len(overlap))
377
+
378
+ print("\nSaved:")
379
+ print(OUT)
380
+
381
+ categories = {}
382
+
383
+ for item in data:
384
+ categories[item["category"]] = (
385
+ categories.get(item["category"], 0) + 1
386
+ )
387
+
388
+ print("\nCategories:")
389
+
390
+ for k, v in sorted(categories.items()):
391
+ print(f"{k:28s}: {v}")
research/raw_scripts/make_bridge_ood_prompts_v1.py ADDED
@@ -0,0 +1,272 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import random
4
+
5
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
6
+ OUT = os.path.join(ROOT, "bridge_ood_prompts_160.json")
7
+
8
+ random.seed(20260903)
9
+
10
+ data = []
11
+
12
+
13
+ def add(category, prompt):
14
+ data.append({
15
+ "category": category,
16
+ "prompt": prompt.strip()
17
+ })
18
+
19
+
20
+ # ============================================================
21
+ # 1. COMPLEX SPATIAL RELATIONS
22
+ # ============================================================
23
+
24
+ spatial = [
25
+ "A small red sphere rests inside a transparent cube while a blue cylinder stands behind the cube and slightly to its right.",
26
+ "A black chair faces a wooden table, with a white vase beneath the table and a silver lamp positioned behind the chair.",
27
+ "A green bottle stands between two ceramic cups, while a metal plate lies partially underneath the cup on the left.",
28
+ "A yellow cube is suspended above a glass sphere, with a red cone behind both objects and a mirror to their right.",
29
+ "Three boxes form a diagonal line from the lower left to the upper right, with a small glass bottle placed between the second and third boxes.",
30
+ "A polished metal sphere sits in front of a tall rectangular mirror while a wooden cube is visible only through the mirror reflection.",
31
+ "A transparent cylinder overlaps a red cube from the camera viewpoint, while a blue sphere remains fully visible behind both.",
32
+ "Two chairs face one another across a narrow table, with a floor lamp standing behind the chair on the right.",
33
+ ]
34
+
35
+ for i in range(20):
36
+ add(
37
+ "complex_spatial",
38
+ spatial[i % len(spatial)]
39
+ + f" The composition contains exactly {3 + (i % 4)} clearly separated major objects."
40
+ )
41
+
42
+
43
+ # ============================================================
44
+ # 2. HUMAN INTERACTION / ANATOMY
45
+ # ============================================================
46
+
47
+ human = [
48
+ "A woman holds a glass bottle in her left hand while pointing toward a wall sign with her right hand; a seated man beside her looks toward the bottle.",
49
+ "A man kneels on his right knee while extending his left hand toward a woman standing directly in front of him.",
50
+ "Two people shake right hands while each keeps the other hand visible at their side.",
51
+ "A woman sits with her left leg crossed over her right leg while holding a cup with both hands.",
52
+ "A person reaches behind their head with the right arm while the left hand holds a small book against the chest.",
53
+ "A woman turns her torso toward the camera while her head faces left and both hands remain visible.",
54
+ "Three people stand side by side; the center person places one hand on each neighboring person's shoulder.",
55
+ "A seated man raises his left foot slightly above the floor while reaching forward with his right arm.",
56
+ ]
57
+
58
+ for i in range(20):
59
+ add(
60
+ "complex_anatomy",
61
+ human[i % len(human)]
62
+ )
63
+
64
+
65
+ # ============================================================
66
+ # 3. MATERIALS + LIGHT
67
+ # ============================================================
68
+
69
+ materials = [
70
+ "A frosted glass sculpture illuminated from behind, with light scattering through its translucent surface and a faint soft-edged shadow on the wall.",
71
+ "A polished chrome kettle reflecting a red chair, a window, and a warm ceiling lamp in its curved surface.",
72
+ "A wet black stone on rough concrete under strong side lighting, showing tiny specular reflections in the water.",
73
+ "A translucent amber plastic sheet standing in front of a cool white light source, tinting the objects visible through it.",
74
+ "A brushed aluminum cylinder next to a glossy ceramic sphere under broad diffused daylight.",
75
+ "A clear glass bottle filled halfway with blue liquid, illuminated from the side so that light refracts through both glass and liquid.",
76
+ "A velvet-covered chair beside a polished metal table, lit by a narrow warm spotlight from above.",
77
+ "A sheet of wrinkled metallic foil reflecting several small colored objects positioned outside the center of the frame.",
78
+ ]
79
+
80
+ for i in range(20):
81
+ add(
82
+ "complex_material_light",
83
+ materials[i % len(materials)]
84
+ )
85
+
86
+
87
+ # ============================================================
88
+ # 4. TEXT + SCENE COMPOSITION
89
+ # ============================================================
90
+
91
+ texts = [
92
+ ('NORTH GATE', "above a glass entrance"),
93
+ ('ROOM 204', "on a small brass plaque beside a door"),
94
+ ('OPEN UNTIL 9 PM', "on a white storefront sign"),
95
+ ('BLUE RIVER HOTEL', "across a large illuminated facade"),
96
+ ('PLATFORM 7', "on a hanging railway sign"),
97
+ ('NO ENTRY', "on a red rectangular metal sign"),
98
+ ('STUDIO B', "on a paper label attached to a black door"),
99
+ ('COFFEE AND BOOKS', "across a window awning"),
100
+ ]
101
+
102
+ for i in range(20):
103
+ word, place = texts[i % len(texts)]
104
+
105
+ add(
106
+ "complex_text",
107
+ f'The exact text "{word}" appears clearly {place}. '
108
+ f'A person stands nearby without blocking any letters, '
109
+ f'and the entire phrase must remain readable.'
110
+ )
111
+
112
+
113
+ # ============================================================
114
+ # 5. REFLECTION / TRANSPARENCY / OCCLUSION
115
+ # ============================================================
116
+
117
+ reflection = [
118
+ "A woman faces away from a wall mirror, while the mirror clearly shows the front of her face and the glass she holds.",
119
+ "A transparent glass panel stands between the camera and two people, with reflections of overhead lights visible on the panel.",
120
+ "A chrome sphere reflects three colored cubes positioned around it, including one cube located behind the camera viewpoint.",
121
+ "A person stands partly behind a translucent curtain, with the silhouette and one hand visible through the fabric.",
122
+ "A bottle in the foreground partially blocks a person's torso but leaves both hands and the face clearly visible.",
123
+ "Two overlapping transparent sheets produce layered refractions of a black-and-white checkerboard behind them.",
124
+ "A large mirror reflects a chair that is outside the direct camera frame while a table remains visible both directly and in reflection.",
125
+ "A glass of water viewed through another glass object produces multiple distorted outlines but preserves the basic object shapes.",
126
+ ]
127
+
128
+ for i in range(20):
129
+ add(
130
+ "reflection_occlusion_ood",
131
+ reflection[i % len(reflection)]
132
+ )
133
+
134
+
135
+ # ============================================================
136
+ # 6. VEHICLES / ARCHITECTURE
137
+ # ============================================================
138
+
139
+ architecture = [
140
+ "A red compact car is parked beneath a concrete overhang while a blue bicycle leans against the column immediately to its left.",
141
+ "A tram approaches an intersection between two modern glass buildings while pedestrians wait behind a metal barrier.",
142
+ "A narrow staircase rises between two brick walls and turns to the right at the upper landing.",
143
+ "A modern house has three illuminated windows on the upper floor and two dark windows below, with a tree partially covering the far-right window.",
144
+ "Two motorcycles are parked nose to tail beside a stone building entrance, with a street lamp directly behind them.",
145
+ "A pedestrian bridge crosses above a two-lane road while a white bus passes underneath from left to right.",
146
+ "A tall glass tower reflects a shorter brick building standing across the street.",
147
+ "A train platform contains three benches arranged sequentially, with a suitcase positioned beneath the middle bench.",
148
+ ]
149
+
150
+ for i in range(20):
151
+ add(
152
+ "architecture_vehicle",
153
+ architecture[i % len(architecture)]
154
+ )
155
+
156
+
157
+ # ============================================================
158
+ # 7. MANY OBJECTS / COUNTING
159
+ # ============================================================
160
+
161
+ counting = [
162
+ "Five red cups stand in one row behind three blue bottles, with exactly two small plates in the foreground.",
163
+ "Seven candles form a circle around one glass sphere, and every candle is individually visible.",
164
+ "Four chairs surround a square table, with one chair on each side and two books placed on the tabletop.",
165
+ "Six identical boxes are stacked in three rows of two, with a small yellow ball resting on the upper-left box.",
166
+ "Three people each hold exactly one object: a bottle, a book, and a cup respectively.",
167
+ "Eight stones form two parallel rows of four with a narrow gap between the rows.",
168
+ "Five hanging lamps appear at different heights above a long table containing exactly three plates.",
169
+ "Two large spheres and four small cubes are distributed across the floor without any object completely hiding another.",
170
+ ]
171
+
172
+ for i in range(20):
173
+ add(
174
+ "complex_counting",
175
+ counting[i % len(counting)]
176
+ )
177
+
178
+
179
+ # ============================================================
180
+ # 8. LONG MULTI-CONSTRAINT SCENES
181
+ # ============================================================
182
+
183
+ complex_scenes = [
184
+ """
185
+ A woman in a red coat stands behind a transparent glass table.
186
+ She holds a silver cup in her left hand while pointing with her
187
+ right hand toward a blue wall sign reading "EXIT". A seated man
188
+ to her right looks toward the sign. Warm light enters from the
189
+ left, and a large mirror behind them reflects the back of the man.
190
+ """,
191
+
192
+ """
193
+ Inside a small modern kitchen, three people stand around a wooden
194
+ counter. One person holds a clear bottle, another cuts a red apple,
195
+ and the third reaches toward a white ceramic plate. A stainless
196
+ steel refrigerator on the right reflects part of the room while
197
+ cool daylight enters through a window on the left.
198
+ """,
199
+
200
+ """
201
+ A black motorcycle stands in front of a glass storefront displaying
202
+ the clearly readable text "NIGHT MARKET". Behind the glass are two
203
+ mannequins and a red chair. Reflections of street lights appear on
204
+ the window without obscuring the lettering.
205
+ """,
206
+
207
+ """
208
+ A man sits at the left end of a long table while a woman stands at
209
+ the opposite end. Between them are exactly four objects in this
210
+ order from left to right: a blue bottle, a metal box, a glass sphere,
211
+ and a red book. A pendant lamp above the sphere casts a circular
212
+ pool of light on the table.
213
+ """,
214
+
215
+ """
216
+ A transparent cube contains a small green sphere. The cube rests on
217
+ a polished chrome surface that reflects both objects. Behind them,
218
+ a paper sign reads "LAB 12". A human hand enters from the right and
219
+ touches the top face of the cube with the index finger.
220
+ """,
221
+ ]
222
+
223
+ for i in range(20):
224
+
225
+ base = complex_scenes[i % len(complex_scenes)]
226
+
227
+ add(
228
+ "long_composition",
229
+ " ".join(base.split())
230
+ )
231
+
232
+
233
+ # ============================================================
234
+ # FINAL
235
+ # ============================================================
236
+
237
+ assert len(data) == 160
238
+
239
+ random.shuffle(data)
240
+
241
+ with open(
242
+ OUT,
243
+ "w",
244
+ encoding="utf-8"
245
+ ) as f:
246
+
247
+ json.dump(
248
+ data,
249
+ f,
250
+ ensure_ascii=False,
251
+ indent=2
252
+ )
253
+
254
+
255
+ print("=" * 80)
256
+ print("OOD BRIDGE DATASET CREATED")
257
+ print("=" * 80)
258
+ print("Prompts:", len(data))
259
+ print("Saved:")
260
+ print(OUT)
261
+
262
+ categories = {}
263
+
264
+ for item in data:
265
+ c = item["category"]
266
+ categories[c] = categories.get(c, 0) + 1
267
+
268
+ print()
269
+ print("Categories:")
270
+
271
+ for k, v in sorted(categories.items()):
272
+ print(f"{k:28s}: {v}")
research/raw_scripts/make_bridge_prompts.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+ import random
4
+
5
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
6
+ OUT = os.path.join(ROOT, "bridge_prompts_480.json")
7
+
8
+ random.seed(1234)
9
+
10
+ colors = [
11
+ "red", "blue", "green", "yellow", "black", "white",
12
+ "orange", "purple", "silver", "golden"
13
+ ]
14
+
15
+ objects = [
16
+ "cube", "sphere", "bottle", "glass", "chair", "table",
17
+ "box", "lamp", "book", "vase", "cup", "plate",
18
+ "metal cylinder", "wooden block", "mirror"
19
+ ]
20
+
21
+ materials = [
22
+ "transparent glass",
23
+ "frosted glass",
24
+ "polished chrome",
25
+ "brushed metal",
26
+ "rough stone",
27
+ "dark wood",
28
+ "glossy ceramic",
29
+ "matte plastic",
30
+ "translucent acrylic",
31
+ "wet concrete"
32
+ ]
33
+
34
+ people = [
35
+ "a woman", "a man", "a young woman", "a young man",
36
+ "an elderly woman", "an elderly man"
37
+ ]
38
+
39
+ relations = [
40
+ "to the left of",
41
+ "to the right of",
42
+ "behind",
43
+ "in front of",
44
+ "above",
45
+ "below",
46
+ "beside",
47
+ "partially behind"
48
+ ]
49
+
50
+ lighting = [
51
+ "soft light from the left",
52
+ "hard light from the right",
53
+ "strong backlight",
54
+ "warm overhead light",
55
+ "cool side lighting",
56
+ "soft diffused daylight",
57
+ "a narrow beam of light from above",
58
+ "warm light reflected from below"
59
+ ]
60
+
61
+ words = [
62
+ "OPEN", "CLOSED", "HOTEL", "CAFE", "EXIT",
63
+ "NORTH", "HELLO", "STOP", "STUDIO", "BLUE"
64
+ ]
65
+
66
+ prompts = []
67
+
68
+
69
+ def add(category, text):
70
+ prompts.append({
71
+ "category": category,
72
+ "prompt": text
73
+ })
74
+
75
+
76
+ # ---------------------------------------------------------
77
+ # Spatial composition
78
+ # ---------------------------------------------------------
79
+
80
+ for _ in range(90):
81
+
82
+ a, b = random.sample(objects, 2)
83
+ ca, cb = random.sample(colors, 2)
84
+ rel = random.choice(relations)
85
+
86
+ add(
87
+ "spatial",
88
+ f"a {ca} {a} positioned {rel} a {cb} {b}"
89
+ )
90
+
91
+
92
+ # ---------------------------------------------------------
93
+ # Multi-object composition / counting
94
+ # ---------------------------------------------------------
95
+
96
+ for _ in range(60):
97
+
98
+ n1 = random.randint(2, 5)
99
+ n2 = random.randint(1, 4)
100
+
101
+ a, b = random.sample(objects, 2)
102
+
103
+ add(
104
+ "counting",
105
+ f"{n1} {a}s arranged in the foreground with "
106
+ f"{n2} {b}s behind them"
107
+ )
108
+
109
+
110
+ # ---------------------------------------------------------
111
+ # Anatomy / body relations
112
+ # ---------------------------------------------------------
113
+
114
+ anatomy_actions = [
115
+ "raising their left hand above their head",
116
+ "holding a bottle in their right hand",
117
+ "crossing their right leg over their left leg",
118
+ "turning their head to the left",
119
+ "touching their face with their left hand",
120
+ "holding both hands in front of their chest",
121
+ "standing with one arm behind their back",
122
+ "sitting with both feet visible",
123
+ "reaching forward with their right arm",
124
+ "holding an object between both hands",
125
+ ]
126
+
127
+ for _ in range(70):
128
+
129
+ person = random.choice(people)
130
+ action = random.choice(anatomy_actions)
131
+
132
+ add(
133
+ "anatomy",
134
+ f"{person} {action}, full body clearly visible"
135
+ )
136
+
137
+
138
+ # ---------------------------------------------------------
139
+ # Materials
140
+ # ---------------------------------------------------------
141
+
142
+ for _ in range(60):
143
+
144
+ obj = random.choice(objects)
145
+ mat = random.choice(materials)
146
+
147
+ add(
148
+ "material",
149
+ f"a {obj} made of {mat}, showing realistic surface "
150
+ f"properties and reflections"
151
+ )
152
+
153
+
154
+ # ---------------------------------------------------------
155
+ # Lighting
156
+ # ---------------------------------------------------------
157
+
158
+ for _ in range(60):
159
+
160
+ obj = random.choice(objects)
161
+ light = random.choice(lighting)
162
+
163
+ add(
164
+ "lighting",
165
+ f"a {obj} illuminated by {light}, with clearly visible "
166
+ f"light direction and shadow"
167
+ )
168
+
169
+
170
+ # ---------------------------------------------------------
171
+ # Text rendering
172
+ # ---------------------------------------------------------
173
+
174
+ for _ in range(50):
175
+
176
+ word = random.choice(words)
177
+
178
+ surfaces = [
179
+ "a white rectangular sign",
180
+ "a glass storefront",
181
+ "a black poster",
182
+ "a metal street sign",
183
+ "a paper label",
184
+ ]
185
+
186
+ surface = random.choice(surfaces)
187
+
188
+ add(
189
+ "text",
190
+ f'the word "{word}" printed clearly and correctly on {surface}'
191
+ )
192
+
193
+
194
+ # ---------------------------------------------------------
195
+ # Transparency / reflection / occlusion
196
+ # ---------------------------------------------------------
197
+
198
+ special = [
199
+ "a transparent glass bottle in front of a person's face",
200
+ "a woman reflected accurately in a wall mirror",
201
+ "a chrome sphere reflecting a room around it",
202
+ "a hand visible through a transparent glass panel",
203
+ "one person partially occluded by another person",
204
+ "a glass sphere resting on a reflective metal surface",
205
+ "a transparent object casting a faint shadow",
206
+ "a mirror showing the back of a person facing away",
207
+ "a translucent curtain illuminated from behind",
208
+ "a shiny metal object reflecting a nearby red object",
209
+ ]
210
+
211
+ for _ in range(50):
212
+ add("reflection_occlusion", random.choice(special))
213
+
214
+
215
+ # ---------------------------------------------------------
216
+ # Shuffle and trim
217
+ # ---------------------------------------------------------
218
+
219
+ random.shuffle(prompts)
220
+
221
+ prompts = prompts[:480]
222
+
223
+ with open(
224
+ OUT,
225
+ "w",
226
+ encoding="utf-8"
227
+ ) as f:
228
+
229
+ json.dump(
230
+ prompts,
231
+ f,
232
+ ensure_ascii=False,
233
+ indent=2
234
+ )
235
+
236
+
237
+ print("=" * 80)
238
+ print("Bridge prompt dataset created")
239
+ print("=" * 80)
240
+ print("Count:", len(prompts))
241
+ print("Saved:")
242
+ print(OUT)
243
+
244
+ categories = {}
245
+
246
+ for x in prompts:
247
+ categories[x["category"]] = (
248
+ categories.get(x["category"], 0) + 1
249
+ )
250
+
251
+ print()
252
+ print("Categories:")
253
+
254
+ for k, v in sorted(categories.items()):
255
+ print(f"{k:24s}: {v}")
research/raw_scripts/prepare_sensenova_h3_condition.py ADDED
@@ -0,0 +1,383 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import argparse
5
+ import torch
6
+ import torch.nn.functional as F
7
+
8
+ from safetensors.torch import load_file
9
+ from accelerate import init_empty_weights, load_checkpoint_and_dispatch
10
+ from transformers import AutoTokenizer
11
+
12
+
13
+ # ============================================================
14
+ # PATHS
15
+ # ============================================================
16
+
17
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
18
+
19
+ SENSENOVA_SRC = os.path.join(
20
+ ROOT,
21
+ "SenseNova-U1",
22
+ "src",
23
+ )
24
+
25
+ CHECKPOINT = (
26
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
27
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
28
+ )
29
+
30
+ BRIDGE_FILE = os.path.join(
31
+ ROOT,
32
+ "hidden_bridge_probes_rank128",
33
+ "SN_L32_to_H3_L49_rank128.safetensors",
34
+ )
35
+
36
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
37
+
38
+ LAYER = 32
39
+
40
+
41
+ # ============================================================
42
+ # COMMAND LINE
43
+ # ============================================================
44
+
45
+ parser = argparse.ArgumentParser()
46
+
47
+ parser.add_argument(
48
+ "--prompt",
49
+ required=True,
50
+ type=str,
51
+ )
52
+
53
+ parser.add_argument(
54
+ "--output",
55
+ default=os.path.join(
56
+ ROOT,
57
+ "sensenova_projected_condition.pt",
58
+ ),
59
+ )
60
+
61
+ args = parser.parse_args()
62
+
63
+ PROMPT = args.prompt
64
+ OUT_FILE = args.output
65
+
66
+
67
+ # ============================================================
68
+ # LOCAL SENSENOVA
69
+ # ============================================================
70
+
71
+ sys.path.insert(0, SENSENOVA_SRC)
72
+
73
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
74
+ NEOChatConfig,
75
+ )
76
+
77
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
78
+ NEOChatModel,
79
+ )
80
+
81
+
82
+ print("=" * 80)
83
+ print("SenseNova L32 -> H3 L49 conditioning preparation")
84
+ print("=" * 80)
85
+
86
+ print("\nPrompt:")
87
+ print(PROMPT)
88
+
89
+ print("\nOutput:")
90
+ print(OUT_FILE)
91
+
92
+
93
+ # ============================================================
94
+ # TOKENIZER
95
+ # ============================================================
96
+
97
+ print("\nLoading tokenizer...")
98
+
99
+ tokenizer = AutoTokenizer.from_pretrained(
100
+ HF_REPO,
101
+ trust_remote_code=True,
102
+ )
103
+
104
+ encoded = tokenizer(
105
+ PROMPT,
106
+ return_tensors="pt",
107
+ add_special_tokens=False,
108
+ )
109
+
110
+ input_ids_cpu = encoded["input_ids"].cpu()
111
+
112
+ print("Token count:", input_ids_cpu.shape[1])
113
+ print("Token IDs:", input_ids_cpu.tolist()[0])
114
+
115
+
116
+ # ============================================================
117
+ # CONFIG
118
+ # ============================================================
119
+
120
+ print("\nLoading SenseNova config...")
121
+
122
+ config = NEOChatConfig.from_pretrained(
123
+ HF_REPO,
124
+ )
125
+
126
+ config.llm_config._attn_implementation = "eager"
127
+
128
+
129
+ # ============================================================
130
+ # META MODEL
131
+ # ============================================================
132
+
133
+ print("Creating META architecture...")
134
+
135
+ with init_empty_weights():
136
+ model = NEOChatModel(config)
137
+
138
+
139
+ # ============================================================
140
+ # CHECKPOINT
141
+ # ============================================================
142
+
143
+ print("Loading SenseNova BF16 checkpoint...")
144
+
145
+ model = load_checkpoint_and_dispatch(
146
+ model,
147
+ checkpoint=CHECKPOINT,
148
+ device_map="auto",
149
+ max_memory={
150
+ 0: "19GiB",
151
+ "cpu": "70GiB",
152
+ },
153
+ dtype=torch.bfloat16,
154
+ no_split_module_classes=[
155
+ "Qwen3DecoderLayer",
156
+ "Qwen3MoeDecoderLayer",
157
+ ],
158
+ )
159
+
160
+ model.eval()
161
+
162
+ qwen = model.language_model.model
163
+
164
+ print("Checkpoint loaded.")
165
+
166
+
167
+ # ============================================================
168
+ # INPUT DEVICE
169
+ # ============================================================
170
+
171
+ device_map = getattr(model, "hf_device_map", {})
172
+
173
+ input_device = torch.device("cuda:0")
174
+
175
+ for name, dev in device_map.items():
176
+
177
+ if name.endswith(
178
+ "language_model.model.embed_tokens"
179
+ ):
180
+
181
+ if str(dev) == "cpu":
182
+ input_device = torch.device("cpu")
183
+ else:
184
+ input_device = torch.device("cuda:0")
185
+
186
+ break
187
+
188
+ print("Input device:", input_device)
189
+
190
+
191
+ # ============================================================
192
+ # CAPTURE L32
193
+ # ============================================================
194
+
195
+ captured = {}
196
+
197
+
198
+ def hook_fn(module, inputs, output):
199
+
200
+ if isinstance(output, tuple):
201
+ x = output[0]
202
+ else:
203
+ x = output
204
+
205
+ captured["hidden"] = (
206
+ x.detach()
207
+ .float()
208
+ .cpu()
209
+ )
210
+
211
+
212
+ handle = qwen.layers[LAYER].register_forward_hook(
213
+ hook_fn
214
+ )
215
+
216
+
217
+ input_ids = input_ids_cpu.to(
218
+ input_device
219
+ )
220
+
221
+ if hasattr(qwen, "current_index"):
222
+ qwen.current_index = -1
223
+
224
+
225
+ print("\nRunning SenseNova...")
226
+
227
+ with torch.inference_mode():
228
+
229
+ output = qwen(
230
+ input_ids=input_ids,
231
+ image_gen_indicators=None,
232
+ attention_mask=None,
233
+ use_cache=False,
234
+ )
235
+
236
+
237
+ handle.remove()
238
+
239
+
240
+ if "hidden" not in captured:
241
+ raise RuntimeError(
242
+ "SenseNova L32 was not captured."
243
+ )
244
+
245
+
246
+ sn32 = captured["hidden"]
247
+
248
+ print(
249
+ "SenseNova L32:",
250
+ tuple(sn32.shape)
251
+ )
252
+
253
+ if sn32.shape[-1] != 4096:
254
+ raise RuntimeError(
255
+ f"Expected 4096 dims, got {sn32.shape[-1]}"
256
+ )
257
+
258
+
259
+ # ============================================================
260
+ # LOAD FROZEN BRIDGE
261
+ # ============================================================
262
+
263
+ print("\nLoading frozen rank-128 bridge...")
264
+
265
+ bridge = load_file(
266
+ BRIDGE_FILE,
267
+ device="cpu",
268
+ )
269
+
270
+ down = bridge["down.weight"].float()
271
+ up = bridge["up.weight"].float()
272
+
273
+ print("down:", tuple(down.shape))
274
+ print("up: ", tuple(up.shape))
275
+
276
+
277
+ # ============================================================
278
+ # SAME NORMALIZATION USED DURING TRAINING
279
+ # ============================================================
280
+
281
+ def rms_normalize(x):
282
+
283
+ rms = torch.sqrt(
284
+ x.pow(2).mean(
285
+ dim=-1,
286
+ keepdim=True
287
+ ) + 1e-6
288
+ )
289
+
290
+ return x / rms
291
+
292
+
293
+ # ============================================================
294
+ # PROJECT
295
+ # ============================================================
296
+
297
+ with torch.no_grad():
298
+
299
+ x = rms_normalize(
300
+ sn32.float()
301
+ )
302
+
303
+ rank_state = F.linear(
304
+ x,
305
+ down,
306
+ )
307
+
308
+ projected = F.linear(
309
+ rank_state,
310
+ up,
311
+ )
312
+
313
+
314
+ print(
315
+ "Projected H3-space:",
316
+ tuple(projected.shape)
317
+ )
318
+
319
+ if projected.shape[-1] != 5120:
320
+ raise RuntimeError(
321
+ "Projected conditioning is not 5120-dimensional."
322
+ )
323
+
324
+ if not torch.isfinite(projected).all():
325
+ raise RuntimeError(
326
+ "Non-finite projected conditioning."
327
+ )
328
+
329
+
330
+ # ============================================================
331
+ # SAVE
332
+ # ============================================================
333
+
334
+ torch.save(
335
+ {
336
+ "prompt": PROMPT,
337
+
338
+ "input_ids":
339
+ input_ids_cpu,
340
+
341
+ "source_layer":
342
+ 32,
343
+
344
+ "target_layer":
345
+ 49,
346
+
347
+ "bridge_rank":
348
+ 128,
349
+
350
+ "projected":
351
+ projected.to(
352
+ torch.float16
353
+ ),
354
+
355
+ # Keep this too for diagnostics.
356
+ "sensenova_l32":
357
+ sn32.to(
358
+ torch.float16
359
+ ),
360
+ },
361
+ OUT_FILE,
362
+ )
363
+
364
+
365
+ print()
366
+ print("=" * 80)
367
+ print("SUCCESS")
368
+ print("=" * 80)
369
+
370
+ print("Prompt tokens:", input_ids_cpu.shape[1])
371
+ print(
372
+ "Projected shape:",
373
+ tuple(projected.shape)
374
+ )
375
+
376
+ print("\nSaved:")
377
+ print(OUT_FILE)
378
+
379
+ del output
380
+ del model
381
+
382
+ gc.collect()
383
+ torch.cuda.empty_cache()
research/raw_scripts/test_sensenova_mot_gen_bridge.py ADDED
@@ -0,0 +1,698 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import gc
4
+ import shutil
5
+ import csv
6
+ import torch
7
+
8
+ from safetensors.torch import load_file
9
+
10
+
11
+ # ============================================================
12
+ # PATHS
13
+ # ============================================================
14
+
15
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
16
+
17
+ SENSENOVA_SRC = os.path.join(
18
+ ROOT, "SenseNova-U1", "src"
19
+ )
20
+
21
+ COMPAT_FILE = os.path.join(
22
+ SENSENOVA_SRC,
23
+ "sensenova_u1",
24
+ "models",
25
+ "neo_unify",
26
+ "transformers_compat.py"
27
+ )
28
+
29
+ MODEL_FILE = (
30
+ r"D:\ComfyUI_Python312\ComfyUI\models\unet"
31
+ r"\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
32
+ )
33
+
34
+ H3_FILE = os.path.join(
35
+ ROOT, "h3_hidden_states_fast.pt"
36
+ )
37
+
38
+ BRIDGE_FILE = os.path.join(
39
+ ROOT,
40
+ "sensenova_to_h3_embedding_bridge_rank512.safetensors"
41
+ )
42
+
43
+ OUT_STATES = os.path.join(
44
+ ROOT,
45
+ "sensenova_mot_gen_hidden_states.pt"
46
+ )
47
+
48
+ OUT_TXT = os.path.join(
49
+ ROOT,
50
+ "sensenova_mot_gen_bridge_comparison.txt"
51
+ )
52
+
53
+ OUT_CSV = os.path.join(
54
+ ROOT,
55
+ "sensenova_mot_gen_bridge_comparison.csv"
56
+ )
57
+
58
+ HF_REPO = "sensenova/SenseNova-U1.5-8B-MoT"
59
+
60
+
61
+ # ============================================================
62
+ # TEST SET
63
+ # ============================================================
64
+
65
+ PROMPTS = [
66
+ "a red cube on a blue sphere",
67
+ "a woman holding a transparent glass bottle",
68
+ "three people standing behind a wooden table",
69
+ ]
70
+
71
+ SN_LAYERS = [8, 16, 24, 32, 41]
72
+ H3_LAYERS = [8, 16, 24, 32, 40, 49]
73
+
74
+ # We will try BOTH synthetic generation-position layouts.
75
+ MODES = [
76
+ "causal",
77
+ "block",
78
+ ]
79
+
80
+
81
+ # ============================================================
82
+ # PATCH SENSENOVA / TRANSFORMERS 4.57 COMPAT
83
+ # ============================================================
84
+
85
+ print("=" * 80)
86
+ print("SenseNova MoT-GEN -> MiniMax H3 diagnostic")
87
+ print("=" * 80)
88
+
89
+ print("\nChecking compatibility patch...")
90
+
91
+ with open(COMPAT_FILE, "r", encoding="utf-8") as f:
92
+ text = f.read()
93
+
94
+ OLD = (
95
+ " from transformers.utils.generic import "
96
+ "check_model_inputs as model_input_compat"
97
+ )
98
+
99
+ NEW = (
100
+ " from transformers.utils.generic import check_model_inputs\n"
101
+ " def model_input_compat(func):\n"
102
+ " return check_model_inputs()(func)"
103
+ )
104
+
105
+ if OLD in text:
106
+
107
+ backup = COMPAT_FILE + ".bak"
108
+
109
+ if not os.path.exists(backup):
110
+ shutil.copy2(COMPAT_FILE, backup)
111
+
112
+ text = text.replace(OLD, NEW)
113
+
114
+ with open(COMPAT_FILE, "w", encoding="utf-8") as f:
115
+ f.write(text)
116
+
117
+ print("Patch applied.")
118
+
119
+ else:
120
+ print("Patch already present / not needed.")
121
+
122
+
123
+ # ============================================================
124
+ # IMPORTS
125
+ # ============================================================
126
+
127
+ sys.path.insert(0, SENSENOVA_SRC)
128
+
129
+ from transformers import AutoTokenizer
130
+ from accelerate import init_empty_weights, load_checkpoint_and_dispatch
131
+
132
+ from sensenova_u1.models.neo_unify.configuration_neo_chat import (
133
+ NEOChatConfig,
134
+ )
135
+
136
+ from sensenova_u1.models.neo_unify.modeling_neo_chat import (
137
+ NEOChatModel,
138
+ )
139
+
140
+
141
+ # ============================================================
142
+ # LOAD TOKENIZER + CONFIG
143
+ # ============================================================
144
+
145
+ print("\nLoading tokenizer...")
146
+
147
+ tokenizer = AutoTokenizer.from_pretrained(
148
+ HF_REPO,
149
+ trust_remote_code=True,
150
+ )
151
+
152
+ print("Loading SenseNova config...")
153
+
154
+ config = NEOChatConfig.from_pretrained(
155
+ HF_REPO
156
+ )
157
+
158
+ config.llm_config._attn_implementation = "eager"
159
+
160
+ print("Hidden size:", config.llm_config.hidden_size)
161
+ print("Layers:", config.llm_config.num_hidden_layers)
162
+
163
+
164
+ # ============================================================
165
+ # META MODEL
166
+ # ============================================================
167
+
168
+ print("\nCreating META architecture...")
169
+
170
+ with init_empty_weights():
171
+ model = NEOChatModel(config)
172
+
173
+ print("META architecture created.")
174
+
175
+
176
+ # ============================================================
177
+ # LOAD CHECKPOINT WITH OFFLOAD
178
+ # ============================================================
179
+
180
+ print("\nLoading BF16 checkpoint with Accelerate...")
181
+
182
+ model = load_checkpoint_and_dispatch(
183
+ model,
184
+ checkpoint=MODEL_FILE,
185
+ device_map="auto",
186
+ max_memory={
187
+ 0: "19GiB",
188
+ "cpu": "70GiB",
189
+ },
190
+ dtype=torch.bfloat16,
191
+ no_split_module_classes=[
192
+ "Qwen3DecoderLayer",
193
+ "Qwen3MoeDecoderLayer",
194
+ ],
195
+ )
196
+
197
+ model.eval()
198
+
199
+ print("Checkpoint loaded.")
200
+
201
+
202
+ # ============================================================
203
+ # LANGUAGE MODEL
204
+ # ============================================================
205
+
206
+ language_model = model.language_model
207
+ qwen = language_model.model
208
+ layers = qwen.layers
209
+
210
+ device_map = getattr(model, "hf_device_map", {})
211
+
212
+ INPUT_DEVICE = torch.device("cuda:0")
213
+
214
+ for name, dev in device_map.items():
215
+
216
+ if name.endswith("language_model.model.embed_tokens"):
217
+
218
+ if str(dev) == "cpu":
219
+ INPUT_DEVICE = torch.device("cpu")
220
+ else:
221
+ INPUT_DEVICE = torch.device("cuda:0")
222
+
223
+ break
224
+
225
+ print("Input device:", INPUT_DEVICE)
226
+
227
+
228
+ # ============================================================
229
+ # HOOKS
230
+ # ============================================================
231
+
232
+ captured = {}
233
+ hooks = {}
234
+
235
+
236
+ def make_hook(idx):
237
+
238
+ def hook(module, inputs, output):
239
+
240
+ if isinstance(output, tuple):
241
+ x = output[0]
242
+ else:
243
+ x = output
244
+
245
+ if torch.is_tensor(x):
246
+ captured[idx] = (
247
+ x.detach()
248
+ .float()
249
+ .cpu()
250
+ )
251
+
252
+ return hook
253
+
254
+
255
+ print("\nRegistering hooks...")
256
+
257
+ for idx in SN_LAYERS:
258
+
259
+ hooks[idx] = layers[idx].register_forward_hook(
260
+ make_hook(idx)
261
+ )
262
+
263
+ print(" layer", idx)
264
+
265
+
266
+ # ============================================================
267
+ # INDEX LAYOUT
268
+ # ============================================================
269
+
270
+ def make_indexes(seq_len, device, mode):
271
+
272
+ p = torch.arange(
273
+ seq_len,
274
+ dtype=torch.long,
275
+ device=device
276
+ )
277
+
278
+ z = torch.zeros(
279
+ seq_len,
280
+ dtype=torch.long,
281
+ device=device
282
+ )
283
+
284
+ if mode == "causal":
285
+
286
+ # Every token gets a new temporal position.
287
+ # Block causal mask therefore behaves causally.
288
+ t = p
289
+ h = z
290
+ w = z
291
+
292
+ elif mode == "block":
293
+
294
+ # All tokens belong to one generation block.
295
+ # Spatial W positions distinguish tokens.
296
+ #
297
+ # create_block_causal_mask uses t, so tokens
298
+ # inside this block can see one another.
299
+ t = z
300
+ h = z
301
+ w = p
302
+
303
+ else:
304
+ raise ValueError(mode)
305
+
306
+ return torch.stack(
307
+ [t, h, w],
308
+ dim=0
309
+ )
310
+
311
+
312
+ # ============================================================
313
+ # EXTRACT MoT-GEN STATES
314
+ # ============================================================
315
+
316
+ all_results = {}
317
+
318
+ for mode in MODES:
319
+
320
+ print()
321
+ print("#" * 80)
322
+ print("MODE:", mode.upper())
323
+ print("#" * 80)
324
+
325
+ mode_results = []
326
+
327
+ for i, prompt in enumerate(PROMPTS):
328
+
329
+ print()
330
+ print(f"[{i+1}/{len(PROMPTS)}] {prompt}")
331
+
332
+ encoded = tokenizer(
333
+ prompt,
334
+ return_tensors="pt",
335
+ add_special_tokens=False,
336
+ )
337
+
338
+ input_ids = encoded["input_ids"].to(
339
+ INPUT_DEVICE
340
+ )
341
+
342
+ seq_len = input_ids.shape[1]
343
+
344
+ print(" Tokens:", input_ids.cpu().tolist()[0])
345
+
346
+ # First obtain the normal token embeddings.
347
+ with torch.inference_mode():
348
+
349
+ inputs_embeds = qwen.embed_tokens(
350
+ input_ids
351
+ )
352
+
353
+ device = inputs_embeds.device
354
+
355
+ # Every token is deliberately marked as a
356
+ # generation token -> forward_gen / *_mot_gen.
357
+ image_gen_indicators = torch.ones(
358
+ (1, seq_len),
359
+ dtype=torch.bool,
360
+ device=device,
361
+ )
362
+
363
+ indexes = make_indexes(
364
+ seq_len,
365
+ device,
366
+ mode,
367
+ )
368
+
369
+ captured.clear()
370
+
371
+ qwen.current_index = -1
372
+
373
+ with torch.inference_mode():
374
+
375
+ outputs = qwen(
376
+ input_ids=None,
377
+ inputs_embeds=inputs_embeds,
378
+ image_gen_indicators=image_gen_indicators,
379
+ indexes=indexes,
380
+ attention_mask=None,
381
+ use_cache=False,
382
+ )
383
+
384
+ item = {
385
+ "prompt": prompt,
386
+ "input_ids": input_ids.detach().cpu(),
387
+ "layers": {},
388
+ }
389
+
390
+ for idx in SN_LAYERS:
391
+
392
+ if idx not in captured:
393
+
394
+ print(
395
+ f" L{idx:02d}: NOT CAPTURED"
396
+ )
397
+ continue
398
+
399
+ x = captured[idx]
400
+
401
+ item["layers"][idx] = x
402
+
403
+ print(
404
+ f" L{idx:02d}: "
405
+ f"{tuple(x.shape)} "
406
+ f"mean={x.mean().item():.5f} "
407
+ f"std={x.std().item():.5f} "
408
+ f"finite={torch.isfinite(x).all().item()}"
409
+ )
410
+
411
+ mode_results.append(item)
412
+
413
+ del outputs
414
+ del inputs_embeds
415
+ del input_ids
416
+
417
+ gc.collect()
418
+ torch.cuda.empty_cache()
419
+
420
+ all_results[mode] = mode_results
421
+
422
+
423
+ # ============================================================
424
+ # SAVE RAW MoT STATES
425
+ # ============================================================
426
+
427
+ torch.save(
428
+ {
429
+ "modes": MODES,
430
+ "prompts": PROMPTS,
431
+ "target_layers": SN_LAYERS,
432
+ "results": all_results,
433
+ "note": (
434
+ "Diagnostic forced-MoT run. Text token embeddings "
435
+ "were deliberately marked as image-generation tokens."
436
+ ),
437
+ },
438
+ OUT_STATES,
439
+ )
440
+
441
+ print("\nSaved raw MoT states:")
442
+ print(OUT_STATES)
443
+
444
+
445
+ # ============================================================
446
+ # MODEL NO LONGER NEEDED
447
+ # ============================================================
448
+
449
+ for h in hooks.values():
450
+ h.remove()
451
+
452
+
453
+ # ============================================================
454
+ # LOAD H3 STATES + EXISTING EMBEDDING BRIDGE
455
+ # ============================================================
456
+
457
+ print("\nLoading H3 hidden states...")
458
+
459
+ h3_data = torch.load(
460
+ H3_FILE,
461
+ map_location="cpu",
462
+ weights_only=False,
463
+ )
464
+
465
+ print("Loading existing embedding bridge...")
466
+
467
+ bridge = load_file(
468
+ BRIDGE_FILE,
469
+ device="cpu",
470
+ )
471
+
472
+ A = bridge["proj_in.weight"].float()
473
+ B = bridge["proj_out.weight"].float()
474
+
475
+
476
+ def project_sn(x):
477
+
478
+ z = torch.nn.functional.linear(
479
+ x.float(),
480
+ A
481
+ )
482
+
483
+ return torch.nn.functional.linear(
484
+ z,
485
+ B
486
+ )
487
+
488
+
489
+ def mean_cos(a, b):
490
+
491
+ return (
492
+ torch.nn.functional.cosine_similarity(
493
+ a.float(),
494
+ b.float(),
495
+ dim=-1
496
+ )
497
+ .mean()
498
+ .item()
499
+ )
500
+
501
+
502
+ # ============================================================
503
+ # COMPARE MoT STATES TO H3
504
+ # ============================================================
505
+
506
+ rows = []
507
+
508
+ print()
509
+ print("=" * 80)
510
+ print("MoT-GEN -> H3 COMPARISON")
511
+ print("=" * 80)
512
+
513
+ for mode in MODES:
514
+
515
+ sn_results = all_results[mode]
516
+
517
+ # Check token identity.
518
+ for sn_item, h3_item in zip(
519
+ sn_results,
520
+ h3_data["results"],
521
+ ):
522
+
523
+ if not torch.equal(
524
+ sn_item["input_ids"],
525
+ h3_item["input_ids"],
526
+ ):
527
+ raise RuntimeError(
528
+ "Token mismatch"
529
+ )
530
+
531
+ for sn_layer in SN_LAYERS:
532
+
533
+ for h3_layer in H3_LAYERS:
534
+
535
+ P = []
536
+ H = []
537
+ prompt_scores = []
538
+
539
+ for sn_item, h3_item in zip(
540
+ sn_results,
541
+ h3_data["results"],
542
+ ):
543
+
544
+ sx = sn_item["layers"][sn_layer]
545
+ hy = h3_item["layers"][h3_layer]
546
+
547
+ projected = project_sn(sx)
548
+
549
+ prompt_scores.append(
550
+ mean_cos(projected, hy)
551
+ )
552
+
553
+ P.append(
554
+ projected.reshape(-1, 5120)
555
+ )
556
+
557
+ H.append(
558
+ hy.reshape(-1, 5120)
559
+ )
560
+
561
+ P = torch.cat(P, dim=0)
562
+ H = torch.cat(H, dim=0)
563
+
564
+ cos = mean_cos(P, H)
565
+
566
+ pnorm = P.norm(dim=-1).mean().item()
567
+ hnorm = H.norm(dim=-1).mean().item()
568
+
569
+ ratio = pnorm / hnorm
570
+
571
+ rows.append(
572
+ {
573
+ "mode": mode,
574
+ "sn_layer": sn_layer,
575
+ "h3_layer": h3_layer,
576
+ "cosine": cos,
577
+ "norm_ratio": ratio,
578
+ "prompt_1": prompt_scores[0],
579
+ "prompt_2": prompt_scores[1],
580
+ "prompt_3": prompt_scores[2],
581
+ }
582
+ )
583
+
584
+
585
+ rows.sort(
586
+ key=lambda r: r["cosine"],
587
+ reverse=True,
588
+ )
589
+
590
+
591
+ # ============================================================
592
+ # SAVE REPORT
593
+ # ============================================================
594
+
595
+ with open(
596
+ OUT_CSV,
597
+ "w",
598
+ newline="",
599
+ encoding="utf-8",
600
+ ) as f:
601
+
602
+ writer = csv.DictWriter(
603
+ f,
604
+ fieldnames=rows[0].keys(),
605
+ )
606
+
607
+ writer.writeheader()
608
+ writer.writerows(rows)
609
+
610
+
611
+ with open(
612
+ OUT_TXT,
613
+ "w",
614
+ encoding="utf-8",
615
+ ) as f:
616
+
617
+ f.write(
618
+ "SenseNova forced MoT-GEN -> H3 comparison\n"
619
+ )
620
+
621
+ f.write("=" * 80 + "\n\n")
622
+
623
+ f.write(
624
+ "IMPORTANT: This is a diagnostic experiment. "
625
+ "Text embeddings were forced through the "
626
+ "SenseNova generation branch.\n\n"
627
+ )
628
+
629
+ for rank, r in enumerate(rows, 1):
630
+
631
+ f.write(
632
+ f"{rank:02d}. "
633
+ f"[{r['mode']}] "
634
+ f"SN L{r['sn_layer']:02d} "
635
+ f"-> H3 L{r['h3_layer']:02d} | "
636
+ f"cos={r['cosine']:.6f} | "
637
+ f"norm_ratio={r['norm_ratio']:.4f}\n"
638
+ )
639
+
640
+ f.write(
641
+ " prompts: "
642
+ f"{r['prompt_1']:.6f}, "
643
+ f"{r['prompt_2']:.6f}, "
644
+ f"{r['prompt_3']:.6f}\n"
645
+ )
646
+
647
+
648
+ # ============================================================
649
+ # PRINT TOP RESULTS
650
+ # ============================================================
651
+
652
+ print()
653
+ print("=" * 80)
654
+ print("TOP 20 MoT-GEN RESULTS")
655
+ print("=" * 80)
656
+
657
+ for rank, r in enumerate(rows[:20], 1):
658
+
659
+ print(
660
+ f"{rank:02d}. "
661
+ f"[{r['mode']:6s}] "
662
+ f"SN L{r['sn_layer']:02d} "
663
+ f"-> H3 L{r['h3_layer']:02d} | "
664
+ f"cos={r['cosine']:.6f} | "
665
+ f"norm_ratio={r['norm_ratio']:.4f}"
666
+ )
667
+
668
+
669
+ print()
670
+ print("=" * 80)
671
+ print("BEST BY MODE")
672
+ print("=" * 80)
673
+
674
+ for mode in MODES:
675
+
676
+ candidates = [
677
+ r for r in rows
678
+ if r["mode"] == mode
679
+ ]
680
+
681
+ best = max(
682
+ candidates,
683
+ key=lambda r: r["cosine"],
684
+ )
685
+
686
+ print(
687
+ f"{mode:6s}: "
688
+ f"SN L{best['sn_layer']:02d} "
689
+ f"-> H3 L{best['h3_layer']:02d} | "
690
+ f"cos={best['cosine']:.6f}"
691
+ )
692
+
693
+
694
+ print()
695
+ print("DONE")
696
+ print("States:", OUT_STATES)
697
+ print("Report:", OUT_TXT)
698
+ print("CSV:", OUT_CSV)
research/raw_scripts/train_distilled_adapter_screening.py ADDED
@@ -0,0 +1,1311 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import math
3
+ import random
4
+ import csv
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ from safetensors.torch import load_file, save_file
10
+
11
+
12
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
13
+
14
+ H3_FILE = os.path.join(
15
+ ROOT,
16
+ "h3_bridge_hidden_480.pt"
17
+ )
18
+
19
+ SN_FILE = os.path.join(
20
+ ROOT,
21
+ "sensenova_bridge_hidden_480.pt"
22
+ )
23
+
24
+ BRIDGE_FILE = (
25
+ r"D:\ComfyUI_Krea2\ComfyUI\models\bridge"
26
+ r"\SN_L32_to_H3_L49_rank128.safetensors"
27
+ )
28
+
29
+ OUT_DIR = os.path.join(
30
+ ROOT,
31
+ "distilled_adapter_screening"
32
+ )
33
+
34
+ os.makedirs(
35
+ OUT_DIR,
36
+ exist_ok=True
37
+ )
38
+
39
+
40
+ # ============================================================
41
+ # SETTINGS
42
+ # ============================================================
43
+
44
+ SEED = 20260903
45
+
46
+ DEVICE = (
47
+ "cuda"
48
+ if torch.cuda.is_available()
49
+ else "cpu"
50
+ )
51
+
52
+ RANK = 128
53
+
54
+ EPOCHS = 12
55
+
56
+ BATCH_TOKENS = 4096
57
+
58
+ LR = 2e-3
59
+
60
+ WEIGHT_DECAY = 1e-4
61
+
62
+ VAL_PROMPTS = 80
63
+
64
+ SOURCE_LAYERS = [
65
+ 8,
66
+ 16,
67
+ 24,
68
+ 32,
69
+ 40,
70
+ 49,
71
+ ]
72
+
73
+ MULTI_LAYERS = [
74
+ 32,
75
+ 40,
76
+ 49,
77
+ ]
78
+
79
+
80
+ # ============================================================
81
+ # REPRODUCIBILITY
82
+ # ============================================================
83
+
84
+ random.seed(
85
+ SEED
86
+ )
87
+
88
+ torch.manual_seed(
89
+ SEED
90
+ )
91
+
92
+ if torch.cuda.is_available():
93
+ torch.cuda.manual_seed_all(
94
+ SEED
95
+ )
96
+
97
+
98
+ # ============================================================
99
+ # MODEL
100
+ # ============================================================
101
+
102
+ class ResidualAdapter(
103
+ nn.Module
104
+ ):
105
+ def __init__(
106
+ self,
107
+ input_dim,
108
+ rank,
109
+ output_dim,
110
+ ):
111
+ super().__init__()
112
+
113
+ self.down = nn.Linear(
114
+ input_dim,
115
+ rank,
116
+ bias=False,
117
+ )
118
+
119
+ self.act = nn.SiLU()
120
+
121
+ self.up = nn.Linear(
122
+ rank,
123
+ output_dim,
124
+ bias=False,
125
+ )
126
+
127
+ nn.init.normal_(
128
+ self.down.weight,
129
+ std=0.02,
130
+ )
131
+
132
+ nn.init.zeros_(
133
+ self.up.weight
134
+ )
135
+
136
+ def forward(
137
+ self,
138
+ x,
139
+ ):
140
+ x = self.down(
141
+ x
142
+ )
143
+
144
+ x = self.act(
145
+ x
146
+ )
147
+
148
+ x = self.up(
149
+ x
150
+ )
151
+
152
+ return x
153
+
154
+
155
+ # ============================================================
156
+ # HELPERS
157
+ # ============================================================
158
+
159
+ def rms_normalize(
160
+ x,
161
+ ):
162
+ rms = torch.sqrt(
163
+ x.float()
164
+ .pow(2)
165
+ .mean(
166
+ dim=-1,
167
+ keepdim=True,
168
+ )
169
+ + 1e-6
170
+ )
171
+
172
+ return (
173
+ x.float()
174
+ / rms
175
+ )
176
+
177
+
178
+ def magnitude_match(
179
+ source,
180
+ target,
181
+ ):
182
+ source = source.float()
183
+ target = target.float()
184
+
185
+ source_rms = torch.sqrt(
186
+ source.pow(2)
187
+ .mean(
188
+ dim=-1,
189
+ keepdim=True,
190
+ )
191
+ + 1e-8
192
+ )
193
+
194
+ target_rms = torch.sqrt(
195
+ target.pow(2)
196
+ .mean(
197
+ dim=-1,
198
+ keepdim=True,
199
+ )
200
+ + 1e-8
201
+ )
202
+
203
+ return (
204
+ source
205
+ * (
206
+ target_rms
207
+ / source_rms
208
+ )
209
+ )
210
+
211
+
212
+ def cosine_mean(
213
+ a,
214
+ b,
215
+ ):
216
+ return (
217
+ F.cosine_similarity(
218
+ a.float(),
219
+ b.float(),
220
+ dim=-1,
221
+ )
222
+ .mean()
223
+ .item()
224
+ )
225
+
226
+
227
+ def save_adapter(
228
+ model,
229
+ path,
230
+ metadata,
231
+ ):
232
+ tensors = {
233
+ "down.weight":
234
+ model.down.weight
235
+ .detach()
236
+ .cpu()
237
+ .to(torch.float16),
238
+
239
+ "up.weight":
240
+ model.up.weight
241
+ .detach()
242
+ .cpu()
243
+ .to(torch.float16),
244
+ }
245
+
246
+ save_file(
247
+ tensors,
248
+ path,
249
+ metadata={
250
+ k: str(v)
251
+ for k, v
252
+ in metadata.items()
253
+ },
254
+ )
255
+
256
+
257
+ # ============================================================
258
+ # LOAD DATA
259
+ # ============================================================
260
+
261
+ print(
262
+ "=" * 100
263
+ )
264
+
265
+ print(
266
+ "DISTILLED ADAPTER SCREENING"
267
+ )
268
+
269
+ print(
270
+ "=" * 100
271
+ )
272
+
273
+ print(
274
+ "Device:",
275
+ DEVICE
276
+ )
277
+
278
+ print(
279
+ "\nLoading H3 data..."
280
+ )
281
+
282
+ h3_data = torch.load(
283
+ H3_FILE,
284
+ map_location="cpu",
285
+ weights_only=False,
286
+ )
287
+
288
+ print(
289
+ "H3 prompts:",
290
+ len(
291
+ h3_data["results"]
292
+ )
293
+ )
294
+
295
+ print(
296
+ "\nLoading SenseNova data..."
297
+ )
298
+
299
+ sn_data = torch.load(
300
+ SN_FILE,
301
+ map_location="cpu",
302
+ weights_only=False,
303
+ )
304
+
305
+ print(
306
+ "SenseNova prompts:",
307
+ len(
308
+ sn_data["results"]
309
+ )
310
+ )
311
+
312
+ if (
313
+ len(h3_data["results"])
314
+ != len(sn_data["results"])
315
+ ):
316
+ raise RuntimeError(
317
+ "Dataset length mismatch."
318
+ )
319
+
320
+
321
+ # ============================================================
322
+ # LOAD FULL TEACHER BRIDGE
323
+ # ============================================================
324
+
325
+ print(
326
+ "\nLoading teacher bridge..."
327
+ )
328
+
329
+ teacher_bridge = load_file(
330
+ BRIDGE_FILE,
331
+ device="cpu",
332
+ )
333
+
334
+ teacher_down = (
335
+ teacher_bridge[
336
+ "down.weight"
337
+ ]
338
+ .float()
339
+ )
340
+
341
+ teacher_up = (
342
+ teacher_bridge[
343
+ "up.weight"
344
+ ]
345
+ .float()
346
+ )
347
+
348
+ print(
349
+ "Teacher down:",
350
+ tuple(
351
+ teacher_down.shape
352
+ )
353
+ )
354
+
355
+ print(
356
+ "Teacher up:",
357
+ tuple(
358
+ teacher_up.shape
359
+ )
360
+ )
361
+
362
+
363
+ # ============================================================
364
+ # ALIGNMENT CHECK
365
+ # ============================================================
366
+
367
+ print(
368
+ "\nChecking alignment..."
369
+ )
370
+
371
+ for i, (
372
+ h3_item,
373
+ sn_item,
374
+ ) in enumerate(
375
+ zip(
376
+ h3_data[
377
+ "results"
378
+ ],
379
+ sn_data[
380
+ "results"
381
+ ],
382
+ )
383
+ ):
384
+ if (
385
+ h3_item[
386
+ "prompt"
387
+ ]
388
+ != sn_item[
389
+ "prompt"
390
+ ]
391
+ ):
392
+ raise RuntimeError(
393
+ f"Prompt mismatch "
394
+ f"at {i}"
395
+ )
396
+
397
+ if not torch.equal(
398
+ h3_item[
399
+ "input_ids"
400
+ ],
401
+ sn_item[
402
+ "input_ids"
403
+ ],
404
+ ):
405
+ raise RuntimeError(
406
+ f"Token mismatch "
407
+ f"at {i}"
408
+ )
409
+
410
+ print(
411
+ "Alignment: PERFECT"
412
+ )
413
+
414
+
415
+ # ============================================================
416
+ # BUILD TEACHER TARGETS
417
+ # ============================================================
418
+
419
+ print(
420
+ "\nBuilding teacher corrections..."
421
+ )
422
+
423
+ samples = []
424
+
425
+ with torch.no_grad():
426
+ for i, (
427
+ h3_item,
428
+ sn_item,
429
+ ) in enumerate(
430
+ zip(
431
+ h3_data[
432
+ "results"
433
+ ],
434
+ sn_data[
435
+ "results"
436
+ ],
437
+ )
438
+ ):
439
+ h49 = (
440
+ h3_item[
441
+ "layers"
442
+ ][49]
443
+ .float()
444
+ )
445
+
446
+ sn32 = (
447
+ sn_item[
448
+ "layers"
449
+ ][32]
450
+ .float()
451
+ )
452
+
453
+ x = rms_normalize(
454
+ sn32
455
+ )
456
+
457
+ rank_state = F.linear(
458
+ x,
459
+ teacher_down,
460
+ )
461
+
462
+ projected = F.linear(
463
+ rank_state,
464
+ teacher_up,
465
+ )
466
+
467
+ projected = (
468
+ magnitude_match(
469
+ projected,
470
+ h49,
471
+ )
472
+ )
473
+
474
+ delta = (
475
+ projected
476
+ - h49
477
+ )
478
+
479
+ samples.append({
480
+ "prompt":
481
+ h3_item[
482
+ "prompt"
483
+ ],
484
+
485
+ "category":
486
+ h3_item[
487
+ "category"
488
+ ],
489
+
490
+ "layers":
491
+ h3_item[
492
+ "layers"
493
+ ],
494
+
495
+ "target_delta":
496
+ delta.to(
497
+ torch.float16
498
+ ),
499
+ })
500
+
501
+ if (
502
+ (i + 1)
503
+ % 50
504
+ == 0
505
+ ):
506
+ print(
507
+ f"{i + 1:3d}/"
508
+ f"{len(h3_data['results'])}"
509
+ )
510
+
511
+
512
+ # ============================================================
513
+ # TRAIN / VAL SPLIT
514
+ # ============================================================
515
+
516
+ indices = list(
517
+ range(
518
+ len(samples)
519
+ )
520
+ )
521
+
522
+ random.shuffle(
523
+ indices
524
+ )
525
+
526
+ val_indices = set(
527
+ indices[
528
+ :VAL_PROMPTS
529
+ ]
530
+ )
531
+
532
+ train_indices = [
533
+ i
534
+ for i
535
+ in indices
536
+ if i
537
+ not in val_indices
538
+ ]
539
+
540
+ val_indices = sorted(
541
+ val_indices
542
+ )
543
+
544
+ print()
545
+ print(
546
+ "Train prompts:",
547
+ len(
548
+ train_indices
549
+ )
550
+ )
551
+
552
+ print(
553
+ "Validation prompts:",
554
+ len(
555
+ val_indices
556
+ )
557
+ )
558
+
559
+
560
+ # ============================================================
561
+ # FLATTEN PROMPTS INTO TOKEN BATCHES
562
+ # ============================================================
563
+
564
+ def build_token_dataset(
565
+ source_spec,
566
+ prompt_indices,
567
+ ):
568
+ xs = []
569
+ ys = []
570
+
571
+ for idx in prompt_indices:
572
+ item = samples[
573
+ idx
574
+ ]
575
+
576
+ if isinstance(
577
+ source_spec,
578
+ int,
579
+ ):
580
+ x = (
581
+ item[
582
+ "layers"
583
+ ][source_spec]
584
+ .float()
585
+ )
586
+
587
+ else:
588
+ parts = [
589
+ item[
590
+ "layers"
591
+ ][layer]
592
+ .float()
593
+ for layer
594
+ in source_spec
595
+ ]
596
+
597
+ x = torch.cat(
598
+ parts,
599
+ dim=-1,
600
+ )
601
+
602
+ y = (
603
+ item[
604
+ "target_delta"
605
+ ]
606
+ .float()
607
+ )
608
+
609
+ x = (
610
+ x.squeeze(0)
611
+ )
612
+
613
+ y = (
614
+ y.squeeze(0)
615
+ )
616
+
617
+ xs.append(
618
+ x
619
+ )
620
+
621
+ ys.append(
622
+ y
623
+ )
624
+
625
+ return (
626
+ torch.cat(
627
+ xs,
628
+ dim=0,
629
+ ),
630
+ torch.cat(
631
+ ys,
632
+ dim=0,
633
+ ),
634
+ )
635
+
636
+
637
+ # ============================================================
638
+ # LOSS
639
+ # ============================================================
640
+
641
+ def loss_fn(
642
+ pred,
643
+ target,
644
+ ):
645
+ mse = F.mse_loss(
646
+ pred,
647
+ target,
648
+ )
649
+
650
+ cos = (
651
+ 1.0
652
+ - F.cosine_similarity(
653
+ pred,
654
+ target,
655
+ dim=-1,
656
+ )
657
+ .mean()
658
+ )
659
+
660
+ return (
661
+ mse
662
+ + 0.10
663
+ * cos
664
+ )
665
+
666
+
667
+ # ============================================================
668
+ # TRAIN ONE CANDIDATE
669
+ # ============================================================
670
+
671
+ def train_candidate(
672
+ name,
673
+ source_spec,
674
+ ):
675
+ print()
676
+ print(
677
+ "=" * 100
678
+ )
679
+
680
+ print(
681
+ "CANDIDATE:",
682
+ name
683
+ )
684
+
685
+ print(
686
+ "=" * 100
687
+ )
688
+
689
+ train_x, train_y = (
690
+ build_token_dataset(
691
+ source_spec,
692
+ train_indices,
693
+ )
694
+ )
695
+
696
+ val_x, val_y = (
697
+ build_token_dataset(
698
+ source_spec,
699
+ val_indices,
700
+ )
701
+ )
702
+
703
+ input_dim = (
704
+ train_x.shape[
705
+ -1
706
+ ]
707
+ )
708
+
709
+ print(
710
+ "Input dim:",
711
+ input_dim
712
+ )
713
+
714
+ print(
715
+ "Train tokens:",
716
+ train_x.shape[
717
+ 0
718
+ ]
719
+ )
720
+
721
+ print(
722
+ "Val tokens:",
723
+ val_x.shape[
724
+ 0
725
+ ]
726
+ )
727
+
728
+ model = (
729
+ ResidualAdapter(
730
+ input_dim,
731
+ RANK,
732
+ 5120,
733
+ )
734
+ .to(
735
+ DEVICE
736
+ )
737
+ )
738
+
739
+ optimizer = (
740
+ torch.optim.AdamW(
741
+ model.parameters(),
742
+ lr=LR,
743
+ weight_decay=
744
+ WEIGHT_DECAY,
745
+ )
746
+ )
747
+
748
+ best_val_cos = -1.0
749
+
750
+ best_path = os.path.join(
751
+ OUT_DIR,
752
+ f"{name}_rank{RANK}.safetensors"
753
+ )
754
+
755
+ train_x = train_x.pin_memory()
756
+ train_y = train_y.pin_memory()
757
+
758
+ val_x = val_x.pin_memory()
759
+ val_y = val_y.pin_memory()
760
+
761
+ token_indices = torch.arange(
762
+ train_x.shape[
763
+ 0
764
+ ]
765
+ )
766
+
767
+ for epoch in range(
768
+ 1,
769
+ EPOCHS + 1,
770
+ ):
771
+ permutation = (
772
+ token_indices[
773
+ torch.randperm(
774
+ token_indices
775
+ .numel()
776
+ )
777
+ ]
778
+ )
779
+
780
+ model.train()
781
+
782
+ epoch_loss = 0.0
783
+
784
+ batches = 0
785
+
786
+ for start in range(
787
+ 0,
788
+ permutation.numel(),
789
+ BATCH_TOKENS,
790
+ ):
791
+ batch_idx = (
792
+ permutation[
793
+ start:
794
+ start
795
+ + BATCH_TOKENS
796
+ ]
797
+ )
798
+
799
+ x = (
800
+ train_x[
801
+ batch_idx
802
+ ]
803
+ .to(
804
+ DEVICE,
805
+ non_blocking=True,
806
+ )
807
+ )
808
+
809
+ y = (
810
+ train_y[
811
+ batch_idx
812
+ ]
813
+ .to(
814
+ DEVICE,
815
+ non_blocking=True,
816
+ )
817
+ )
818
+
819
+ optimizer.zero_grad(
820
+ set_to_none=True
821
+ )
822
+
823
+ pred = model(
824
+ x
825
+ )
826
+
827
+ loss = loss_fn(
828
+ pred,
829
+ y
830
+ )
831
+
832
+ loss.backward()
833
+
834
+ torch.nn.utils.clip_grad_norm_(
835
+ model.parameters(),
836
+ 1.0,
837
+ )
838
+
839
+ optimizer.step()
840
+
841
+ epoch_loss += (
842
+ loss.item()
843
+ )
844
+
845
+ batches += 1
846
+
847
+ epoch_loss /= max(
848
+ batches,
849
+ 1,
850
+ )
851
+
852
+ model.eval()
853
+
854
+ with torch.no_grad():
855
+ pred_chunks = []
856
+
857
+ for start in range(
858
+ 0,
859
+ val_x.shape[0],
860
+ BATCH_TOKENS,
861
+ ):
862
+ x = (
863
+ val_x[
864
+ start:
865
+ start
866
+ + BATCH_TOKENS
867
+ ]
868
+ .to(
869
+ DEVICE,
870
+ non_blocking=True,
871
+ )
872
+ )
873
+
874
+ pred = model(
875
+ x
876
+ )
877
+
878
+ pred_chunks.append(
879
+ pred.cpu()
880
+ )
881
+
882
+ pred_val = torch.cat(
883
+ pred_chunks,
884
+ dim=0,
885
+ )
886
+
887
+ val_cos = cosine_mean(
888
+ pred_val,
889
+ val_y,
890
+ )
891
+
892
+ val_mse = F.mse_loss(
893
+ pred_val.float(),
894
+ val_y.float(),
895
+ ).item()
896
+
897
+ print(
898
+ f"Epoch "
899
+ f"{epoch:02d}/"
900
+ f"{EPOCHS} | "
901
+ f"loss="
902
+ f"{epoch_loss:.6f} | "
903
+ f"val_cos="
904
+ f"{val_cos:.6f} | "
905
+ f"val_mse="
906
+ f"{val_mse:.8f}"
907
+ )
908
+
909
+ if (
910
+ val_cos
911
+ > best_val_cos
912
+ ):
913
+ best_val_cos = (
914
+ val_cos
915
+ )
916
+
917
+ save_adapter(
918
+ model,
919
+ best_path,
920
+ {
921
+ "candidate":
922
+ name,
923
+
924
+ "source_spec":
925
+ source_spec,
926
+
927
+ "rank":
928
+ RANK,
929
+
930
+ "target":
931
+ "teacher_delta",
932
+
933
+ "teacher":
934
+ "SN_L32_to_H3_L49_rank128",
935
+
936
+ "magnitude_match":
937
+ "per_token",
938
+ },
939
+ )
940
+
941
+ # ========================================================
942
+ # FINAL PROMPT-LEVEL EVAL
943
+ # ========================================================
944
+
945
+ best_weights = load_file(
946
+ best_path,
947
+ device="cpu",
948
+ )
949
+
950
+ model = (
951
+ ResidualAdapter(
952
+ input_dim,
953
+ RANK,
954
+ 5120,
955
+ )
956
+ )
957
+
958
+ with torch.no_grad():
959
+ model.down.weight.copy_(
960
+ best_weights[
961
+ "down.weight"
962
+ ].float()
963
+ )
964
+
965
+ model.up.weight.copy_(
966
+ best_weights[
967
+ "up.weight"
968
+ ].float()
969
+ )
970
+
971
+ model = model.to(
972
+ DEVICE
973
+ )
974
+
975
+ model.eval()
976
+
977
+ prompt_scores = []
978
+
979
+ with torch.no_grad():
980
+ for idx in val_indices:
981
+ item = samples[
982
+ idx
983
+ ]
984
+
985
+ if isinstance(
986
+ source_spec,
987
+ int,
988
+ ):
989
+ x = (
990
+ item[
991
+ "layers"
992
+ ][source_spec]
993
+ .float()
994
+ )
995
+
996
+ else:
997
+ x = torch.cat(
998
+ [
999
+ item[
1000
+ "layers"
1001
+ ][layer]
1002
+ .float()
1003
+ for layer
1004
+ in source_spec
1005
+ ],
1006
+ dim=-1,
1007
+ )
1008
+
1009
+ y = (
1010
+ item[
1011
+ "target_delta"
1012
+ ]
1013
+ .float()
1014
+ )
1015
+
1016
+ pred = model(
1017
+ x.to(
1018
+ DEVICE
1019
+ )
1020
+ ).cpu()
1021
+
1022
+ score = cosine_mean(
1023
+ pred,
1024
+ y,
1025
+ )
1026
+
1027
+ prompt_scores.append(
1028
+ score
1029
+ )
1030
+
1031
+ prompt_mean = (
1032
+ sum(
1033
+ prompt_scores
1034
+ )
1035
+ / len(
1036
+ prompt_scores
1037
+ )
1038
+ )
1039
+
1040
+ prompt_min = min(
1041
+ prompt_scores
1042
+ )
1043
+
1044
+ prompt_max = max(
1045
+ prompt_scores
1046
+ )
1047
+
1048
+ print()
1049
+ print(
1050
+ "BEST VAL TOKEN COS:",
1051
+ f"{best_val_cos:.6f}"
1052
+ )
1053
+
1054
+ print(
1055
+ "VAL PROMPT COS:",
1056
+ f"{prompt_mean:.6f}"
1057
+ )
1058
+
1059
+ print(
1060
+ "VAL PROMPT MIN:",
1061
+ f"{prompt_min:.6f}"
1062
+ )
1063
+
1064
+ print(
1065
+ "VAL PROMPT MAX:",
1066
+ f"{prompt_max:.6f}"
1067
+ )
1068
+
1069
+ return {
1070
+ "candidate":
1071
+ name,
1072
+
1073
+ "source":
1074
+ str(
1075
+ source_spec
1076
+ ),
1077
+
1078
+ "input_dim":
1079
+ input_dim,
1080
+
1081
+ "rank":
1082
+ RANK,
1083
+
1084
+ "best_val_token_cos":
1085
+ best_val_cos,
1086
+
1087
+ "val_prompt_cos":
1088
+ prompt_mean,
1089
+
1090
+ "val_prompt_min":
1091
+ prompt_min,
1092
+
1093
+ "val_prompt_max":
1094
+ prompt_max,
1095
+
1096
+ "model_file":
1097
+ best_path,
1098
+ }
1099
+
1100
+
1101
+ # ============================================================
1102
+ # SCREENING
1103
+ # ============================================================
1104
+
1105
+ results = []
1106
+
1107
+ for layer in SOURCE_LAYERS:
1108
+ result = train_candidate(
1109
+ f"H3_L{layer}_to_teacher_delta",
1110
+ layer,
1111
+ )
1112
+
1113
+ results.append(
1114
+ result
1115
+ )
1116
+
1117
+
1118
+ # ============================================================
1119
+ # MULTI-LAYER
1120
+ # ============================================================
1121
+
1122
+ multi_name = (
1123
+ "H3_L32_L40_L49_to_teacher_delta"
1124
+ )
1125
+
1126
+ multi_result = train_candidate(
1127
+ multi_name,
1128
+ MULTI_LAYERS,
1129
+ )
1130
+
1131
+ results.append(
1132
+ multi_result
1133
+ )
1134
+
1135
+
1136
+ # ============================================================
1137
+ # SORT
1138
+ # ============================================================
1139
+
1140
+ results = sorted(
1141
+ results,
1142
+ key=lambda x:
1143
+ x[
1144
+ "val_prompt_cos"
1145
+ ],
1146
+ reverse=True,
1147
+ )
1148
+
1149
+
1150
+ # ============================================================
1151
+ # SAVE CSV
1152
+ # ============================================================
1153
+
1154
+ csv_path = os.path.join(
1155
+ OUT_DIR,
1156
+ "distilled_adapter_screening.csv"
1157
+ )
1158
+
1159
+ with open(
1160
+ csv_path,
1161
+ "w",
1162
+ newline="",
1163
+ encoding="utf-8",
1164
+ ) as f:
1165
+ writer = csv.DictWriter(
1166
+ f,
1167
+ fieldnames=[
1168
+ "candidate",
1169
+ "source",
1170
+ "input_dim",
1171
+ "rank",
1172
+ "best_val_token_cos",
1173
+ "val_prompt_cos",
1174
+ "val_prompt_min",
1175
+ "val_prompt_max",
1176
+ "model_file",
1177
+ ],
1178
+ )
1179
+
1180
+ writer.writeheader()
1181
+
1182
+ writer.writerows(
1183
+ results
1184
+ )
1185
+
1186
+
1187
+ # ============================================================
1188
+ # SAVE REPORT
1189
+ # ============================================================
1190
+
1191
+ report_path = os.path.join(
1192
+ OUT_DIR,
1193
+ "distilled_adapter_screening.txt"
1194
+ )
1195
+
1196
+ with open(
1197
+ report_path,
1198
+ "w",
1199
+ encoding="utf-8",
1200
+ ) as f:
1201
+ f.write(
1202
+ "DISTILLED ADAPTER SCREENING\n"
1203
+ )
1204
+
1205
+ f.write(
1206
+ "=" * 100
1207
+ + "\n\n"
1208
+ )
1209
+
1210
+ for i, r in enumerate(
1211
+ results,
1212
+ 1,
1213
+ ):
1214
+ f.write(
1215
+ f"{i:02d}. "
1216
+ f"{r['candidate']}\n"
1217
+ )
1218
+
1219
+ f.write(
1220
+ f" source: "
1221
+ f"{r['source']}\n"
1222
+ )
1223
+
1224
+ f.write(
1225
+ f" input_dim: "
1226
+ f"{r['input_dim']}\n"
1227
+ )
1228
+
1229
+ f.write(
1230
+ f" rank: "
1231
+ f"{r['rank']}\n"
1232
+ )
1233
+
1234
+ f.write(
1235
+ f" token_cos: "
1236
+ f"{r['best_val_token_cos']:.6f}\n"
1237
+ )
1238
+
1239
+ f.write(
1240
+ f" prompt_cos: "
1241
+ f"{r['val_prompt_cos']:.6f}\n"
1242
+ )
1243
+
1244
+ f.write(
1245
+ f" prompt_min: "
1246
+ f"{r['val_prompt_min']:.6f}\n"
1247
+ )
1248
+
1249
+ f.write(
1250
+ f" prompt_max: "
1251
+ f"{r['val_prompt_max']:.6f}\n"
1252
+ )
1253
+
1254
+ f.write(
1255
+ f" model: "
1256
+ f"{r['model_file']}\n\n"
1257
+ )
1258
+
1259
+
1260
+ # ============================================================
1261
+ # FINAL OUTPUT
1262
+ # ============================================================
1263
+
1264
+ print()
1265
+ print(
1266
+ "=" * 100
1267
+ )
1268
+
1269
+ print(
1270
+ "FINAL RANKING"
1271
+ )
1272
+
1273
+ print(
1274
+ "=" * 100
1275
+ )
1276
+
1277
+ for i, r in enumerate(
1278
+ results,
1279
+ 1,
1280
+ ):
1281
+ print(
1282
+ f"{i:02d}. "
1283
+ f"{r['candidate']:38s} "
1284
+ f"prompt_cos="
1285
+ f"{r['val_prompt_cos']:.6f} "
1286
+ f"token_cos="
1287
+ f"{r['best_val_token_cos']:.6f}"
1288
+ )
1289
+
1290
+ print()
1291
+ print(
1292
+ "CSV:"
1293
+ )
1294
+
1295
+ print(
1296
+ csv_path
1297
+ )
1298
+
1299
+ print()
1300
+ print(
1301
+ "Report:"
1302
+ )
1303
+
1304
+ print(
1305
+ report_path
1306
+ )
1307
+
1308
+ print()
1309
+ print(
1310
+ "DONE"
1311
+ )
research/raw_scripts/train_distilled_student_v2.py ADDED
@@ -0,0 +1,1802 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import random
4
+ import math
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+
10
+ from safetensors.torch import load_file, save_file
11
+
12
+
13
+ # ============================================================
14
+ # PATHS
15
+ # ============================================================
16
+
17
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
18
+
19
+ H3_FILE = os.path.join(
20
+ ROOT,
21
+ "h3_bridge_hidden_480.pt"
22
+ )
23
+
24
+ SN_FILE = os.path.join(
25
+ ROOT,
26
+ "sensenova_bridge_hidden_480.pt"
27
+ )
28
+
29
+ TEACHER_BRIDGE = (
30
+ r"D:\ComfyUI_Krea2\ComfyUI\models\bridge"
31
+ r"\SN_L32_to_H3_L49_rank128.safetensors"
32
+ )
33
+
34
+ OUT_DIR = os.path.join(
35
+ ROOT,
36
+ "distilled_student_v2"
37
+ )
38
+
39
+ os.makedirs(
40
+ OUT_DIR,
41
+ exist_ok=True
42
+ )
43
+
44
+
45
+ # ============================================================
46
+ # SETTINGS
47
+ # ============================================================
48
+
49
+ SEED = 20260903
50
+
51
+ DEVICE = (
52
+ "cuda"
53
+ if torch.cuda.is_available()
54
+ else "cpu"
55
+ )
56
+
57
+ SOURCE_LAYERS = [
58
+ 40,
59
+ 49,
60
+ ]
61
+
62
+ HIDDEN_DIM = 512
63
+
64
+ EPOCHS = 30
65
+
66
+ LR = 1e-3
67
+
68
+ WEIGHT_DECAY = 1e-4
69
+
70
+ BATCH_TOKENS = 1024
71
+
72
+ VAL_PROMPTS = 80
73
+
74
+ ALPHAS = [
75
+ 0.05,
76
+ 0.10,
77
+ 0.20,
78
+ 0.30,
79
+ ]
80
+
81
+
82
+ # ============================================================
83
+ # REPRODUCIBILITY
84
+ # ============================================================
85
+
86
+ random.seed(
87
+ SEED
88
+ )
89
+
90
+ torch.manual_seed(
91
+ SEED
92
+ )
93
+
94
+ if torch.cuda.is_available():
95
+ torch.cuda.manual_seed_all(
96
+ SEED
97
+ )
98
+
99
+
100
+ # ============================================================
101
+ # STUDENT
102
+ # ============================================================
103
+
104
+ class SemanticStudent(
105
+ nn.Module
106
+ ):
107
+ def __init__(
108
+ self,
109
+ input_dim=5120,
110
+ hidden_dim=512,
111
+ output_dim=5120,
112
+ ):
113
+ super().__init__()
114
+
115
+ self.fc1 = nn.Linear(
116
+ input_dim,
117
+ hidden_dim,
118
+ bias=True,
119
+ )
120
+
121
+ self.fc2 = nn.Linear(
122
+ hidden_dim,
123
+ hidden_dim,
124
+ bias=True,
125
+ )
126
+
127
+ self.fc3 = nn.Linear(
128
+ hidden_dim,
129
+ output_dim,
130
+ bias=True,
131
+ )
132
+
133
+ self.act = nn.SiLU()
134
+
135
+ nn.init.xavier_uniform_(
136
+ self.fc1.weight
137
+ )
138
+
139
+ nn.init.zeros_(
140
+ self.fc1.bias
141
+ )
142
+
143
+ nn.init.xavier_uniform_(
144
+ self.fc2.weight
145
+ )
146
+
147
+ nn.init.zeros_(
148
+ self.fc2.bias
149
+ )
150
+
151
+ # Start final projection near zero.
152
+ nn.init.normal_(
153
+ self.fc3.weight,
154
+ mean=0.0,
155
+ std=0.002,
156
+ )
157
+
158
+ nn.init.zeros_(
159
+ self.fc3.bias
160
+ )
161
+
162
+
163
+ def forward(
164
+ self,
165
+ x,
166
+ ):
167
+ x = self.fc1(
168
+ x
169
+ )
170
+
171
+ x = self.act(
172
+ x
173
+ )
174
+
175
+ x = self.fc2(
176
+ x
177
+ )
178
+
179
+ x = self.act(
180
+ x
181
+ )
182
+
183
+ x = self.fc3(
184
+ x
185
+ )
186
+
187
+ return x
188
+
189
+
190
+ # ============================================================
191
+ # NORMALIZATION
192
+ # ============================================================
193
+
194
+ def rms_normalize(
195
+ x,
196
+ ):
197
+ x = x.float()
198
+
199
+ rms = torch.sqrt(
200
+ x.pow(2)
201
+ .mean(
202
+ dim=-1,
203
+ keepdim=True,
204
+ )
205
+ + 1e-6
206
+ )
207
+
208
+ return (
209
+ x / rms
210
+ )
211
+
212
+
213
+ def magnitude_match(
214
+ source,
215
+ target,
216
+ ):
217
+ source = source.float()
218
+ target = target.float()
219
+
220
+ source_rms = torch.sqrt(
221
+ source.pow(2)
222
+ .mean(
223
+ dim=-1,
224
+ keepdim=True,
225
+ )
226
+ + 1e-8
227
+ )
228
+
229
+ target_rms = torch.sqrt(
230
+ target.pow(2)
231
+ .mean(
232
+ dim=-1,
233
+ keepdim=True,
234
+ )
235
+ + 1e-8
236
+ )
237
+
238
+ return (
239
+ source
240
+ * (
241
+ target_rms
242
+ / source_rms
243
+ )
244
+ )
245
+
246
+
247
+ # ============================================================
248
+ # METRICS
249
+ # ============================================================
250
+
251
+ def cosine(
252
+ a,
253
+ b,
254
+ ):
255
+ return F.cosine_similarity(
256
+ a.float(),
257
+ b.float(),
258
+ dim=-1,
259
+ )
260
+
261
+
262
+ def cosine_mean(
263
+ a,
264
+ b,
265
+ ):
266
+ return (
267
+ cosine(
268
+ a,
269
+ b,
270
+ )
271
+ .mean()
272
+ .item()
273
+ )
274
+
275
+
276
+ # ============================================================
277
+ # LOSS
278
+ # ============================================================
279
+
280
+ def student_loss(
281
+ pred,
282
+ target,
283
+ ):
284
+ # Direction is most important because runtime
285
+ # magnitude matching rescales student output anyway.
286
+
287
+ pred_n = F.normalize(
288
+ pred.float(),
289
+ dim=-1,
290
+ )
291
+
292
+ target_n = F.normalize(
293
+ target.float(),
294
+ dim=-1,
295
+ )
296
+
297
+ cosine_loss = (
298
+ 1.0
299
+ - (
300
+ pred_n
301
+ * target_n
302
+ )
303
+ .sum(
304
+ dim=-1
305
+ )
306
+ .mean()
307
+ )
308
+
309
+ normalized_mse = (
310
+ F.mse_loss(
311
+ pred_n,
312
+ target_n,
313
+ )
314
+ )
315
+
316
+ # Small raw-scale stabilizer.
317
+ target_rms = torch.sqrt(
318
+ target.float()
319
+ .pow(2)
320
+ .mean(
321
+ dim=-1,
322
+ keepdim=True,
323
+ )
324
+ + 1e-8
325
+ )
326
+
327
+ pred_scaled = (
328
+ pred.float()
329
+ / (
330
+ torch.sqrt(
331
+ pred.float()
332
+ .pow(2)
333
+ .mean(
334
+ dim=-1,
335
+ keepdim=True,
336
+ )
337
+ + 1e-8
338
+ )
339
+ )
340
+ * target_rms
341
+ )
342
+
343
+ scale_mse = F.mse_loss(
344
+ pred_scaled,
345
+ target.float(),
346
+ )
347
+
348
+ loss = (
349
+ cosine_loss
350
+ + 0.50 * normalized_mse
351
+ + 0.02 * scale_mse
352
+ )
353
+
354
+ return (
355
+ loss,
356
+ cosine_loss.item(),
357
+ normalized_mse.item(),
358
+ scale_mse.item(),
359
+ )
360
+
361
+
362
+ # ============================================================
363
+ # SAVE
364
+ # ============================================================
365
+
366
+ def save_student(
367
+ model,
368
+ path,
369
+ source_layer,
370
+ best_cos,
371
+ ):
372
+ tensors = {
373
+ "fc1.weight":
374
+ model.fc1.weight
375
+ .detach()
376
+ .cpu()
377
+ .to(torch.float16),
378
+
379
+ "fc1.bias":
380
+ model.fc1.bias
381
+ .detach()
382
+ .cpu()
383
+ .to(torch.float16),
384
+
385
+ "fc2.weight":
386
+ model.fc2.weight
387
+ .detach()
388
+ .cpu()
389
+ .to(torch.float16),
390
+
391
+ "fc2.bias":
392
+ model.fc2.bias
393
+ .detach()
394
+ .cpu()
395
+ .to(torch.float16),
396
+
397
+ "fc3.weight":
398
+ model.fc3.weight
399
+ .detach()
400
+ .cpu()
401
+ .to(torch.float16),
402
+
403
+ "fc3.bias":
404
+ model.fc3.bias
405
+ .detach()
406
+ .cpu()
407
+ .to(torch.float16),
408
+ }
409
+
410
+ metadata = {
411
+ "architecture":
412
+ "SenseNova_H3_Distilled_Student_V2",
413
+
414
+ "source_model":
415
+ "MiniMax_H3_Qwen3VL",
416
+
417
+ "source_layer":
418
+ str(source_layer),
419
+
420
+ "teacher_source":
421
+ "SenseNova_L32",
422
+
423
+ "teacher_bridge":
424
+ "SN_L32_to_H3_L49_rank128",
425
+
426
+ "input_dim":
427
+ "5120",
428
+
429
+ "hidden_dim":
430
+ str(HIDDEN_DIM),
431
+
432
+ "output_dim":
433
+ "5120",
434
+
435
+ "activation":
436
+ "SiLU",
437
+
438
+ "best_val_teacher_cos":
439
+ f"{best_cos:.8f}",
440
+ }
441
+
442
+ save_file(
443
+ tensors,
444
+ path,
445
+ metadata=metadata,
446
+ )
447
+
448
+
449
+ # ============================================================
450
+ # LOAD DATA
451
+ # ============================================================
452
+
453
+ print(
454
+ "=" * 100
455
+ )
456
+
457
+ print(
458
+ "DISTILLED STUDENT V2"
459
+ )
460
+
461
+ print(
462
+ "H3 -> projected SenseNova teacher"
463
+ )
464
+
465
+ print(
466
+ "=" * 100
467
+ )
468
+
469
+ print(
470
+ "Device:",
471
+ DEVICE
472
+ )
473
+
474
+ print(
475
+ "\nLoading H3 hidden states..."
476
+ )
477
+
478
+ h3_data = torch.load(
479
+ H3_FILE,
480
+ map_location="cpu",
481
+ weights_only=False,
482
+ )
483
+
484
+ print(
485
+ "H3 prompts:",
486
+ len(
487
+ h3_data[
488
+ "results"
489
+ ]
490
+ )
491
+ )
492
+
493
+ print(
494
+ "\nLoading SenseNova hidden states..."
495
+ )
496
+
497
+ sn_data = torch.load(
498
+ SN_FILE,
499
+ map_location="cpu",
500
+ weights_only=False,
501
+ )
502
+
503
+ print(
504
+ "SenseNova prompts:",
505
+ len(
506
+ sn_data[
507
+ "results"
508
+ ]
509
+ )
510
+ )
511
+
512
+
513
+ # ============================================================
514
+ # ALIGNMENT
515
+ # ============================================================
516
+
517
+ if (
518
+ len(
519
+ h3_data[
520
+ "results"
521
+ ]
522
+ )
523
+ !=
524
+ len(
525
+ sn_data[
526
+ "results"
527
+ ]
528
+ )
529
+ ):
530
+ raise RuntimeError(
531
+ "Dataset length mismatch."
532
+ )
533
+
534
+
535
+ print(
536
+ "\nChecking prompt/token alignment..."
537
+ )
538
+
539
+ for i, (
540
+ h3_item,
541
+ sn_item,
542
+ ) in enumerate(
543
+ zip(
544
+ h3_data[
545
+ "results"
546
+ ],
547
+ sn_data[
548
+ "results"
549
+ ],
550
+ )
551
+ ):
552
+ if (
553
+ h3_item[
554
+ "prompt"
555
+ ]
556
+ !=
557
+ sn_item[
558
+ "prompt"
559
+ ]
560
+ ):
561
+ raise RuntimeError(
562
+ f"Prompt mismatch "
563
+ f"at {i}"
564
+ )
565
+
566
+ if not torch.equal(
567
+ h3_item[
568
+ "input_ids"
569
+ ],
570
+ sn_item[
571
+ "input_ids"
572
+ ],
573
+ ):
574
+ raise RuntimeError(
575
+ f"Token mismatch "
576
+ f"at {i}"
577
+ )
578
+
579
+ print(
580
+ "Alignment: PERFECT"
581
+ )
582
+
583
+
584
+ # ============================================================
585
+ # TEACHER BRIDGE
586
+ # ============================================================
587
+
588
+ print(
589
+ "\nLoading teacher bridge..."
590
+ )
591
+
592
+ bridge = load_file(
593
+ TEACHER_BRIDGE,
594
+ device="cpu",
595
+ )
596
+
597
+ teacher_down = (
598
+ bridge[
599
+ "down.weight"
600
+ ]
601
+ .float()
602
+ )
603
+
604
+ teacher_up = (
605
+ bridge[
606
+ "up.weight"
607
+ ]
608
+ .float()
609
+ )
610
+
611
+ print(
612
+ "down:",
613
+ tuple(
614
+ teacher_down.shape
615
+ )
616
+ )
617
+
618
+ print(
619
+ "up:",
620
+ tuple(
621
+ teacher_up.shape
622
+ )
623
+ )
624
+
625
+
626
+ # ============================================================
627
+ # BUILD TEACHER REPRESENTATIONS
628
+ # ============================================================
629
+
630
+ print(
631
+ "\nBuilding projected SenseNova teacher targets..."
632
+ )
633
+
634
+ samples = []
635
+
636
+ with torch.no_grad():
637
+
638
+ for i, (
639
+ h3_item,
640
+ sn_item,
641
+ ) in enumerate(
642
+ zip(
643
+ h3_data[
644
+ "results"
645
+ ],
646
+ sn_data[
647
+ "results"
648
+ ],
649
+ )
650
+ ):
651
+
652
+ sn32 = (
653
+ sn_item[
654
+ "layers"
655
+ ][32]
656
+ .float()
657
+ )
658
+
659
+ h49 = (
660
+ h3_item[
661
+ "layers"
662
+ ][49]
663
+ .float()
664
+ )
665
+
666
+ sn_norm = (
667
+ rms_normalize(
668
+ sn32
669
+ )
670
+ )
671
+
672
+ low = F.linear(
673
+ sn_norm,
674
+ teacher_down,
675
+ )
676
+
677
+ teacher_projected = (
678
+ F.linear(
679
+ low,
680
+ teacher_up,
681
+ )
682
+ )
683
+
684
+ samples.append({
685
+ "prompt":
686
+ h3_item[
687
+ "prompt"
688
+ ],
689
+
690
+ "category":
691
+ h3_item[
692
+ "category"
693
+ ],
694
+
695
+ "layers":
696
+ h3_item[
697
+ "layers"
698
+ ],
699
+
700
+ "h49":
701
+ h49.to(
702
+ torch.float16
703
+ ),
704
+
705
+ "teacher":
706
+ teacher_projected.to(
707
+ torch.float16
708
+ ),
709
+ })
710
+
711
+ if (
712
+ (i + 1)
713
+ % 50
714
+ == 0
715
+ ):
716
+ print(
717
+ f"{i + 1:3d}/"
718
+ f"{len(h3_data['results'])}"
719
+ )
720
+
721
+
722
+ # ============================================================
723
+ # SPLIT BY PROMPT
724
+ # ============================================================
725
+
726
+ indices = list(
727
+ range(
728
+ len(samples)
729
+ )
730
+ )
731
+
732
+ random.shuffle(
733
+ indices
734
+ )
735
+
736
+ val_indices = sorted(
737
+ indices[
738
+ :VAL_PROMPTS
739
+ ]
740
+ )
741
+
742
+ val_set = set(
743
+ val_indices
744
+ )
745
+
746
+ train_indices = [
747
+ i
748
+ for i in indices
749
+ if i not in val_set
750
+ ]
751
+
752
+ print()
753
+ print(
754
+ "Train prompts:",
755
+ len(
756
+ train_indices
757
+ )
758
+ )
759
+
760
+ print(
761
+ "Validation prompts:",
762
+ len(
763
+ val_indices
764
+ )
765
+ )
766
+
767
+
768
+ # ============================================================
769
+ # BUILD TOKEN MATRICES
770
+ # ============================================================
771
+
772
+ def build_token_data(
773
+ source_layer,
774
+ prompt_indices,
775
+ ):
776
+
777
+ xs = []
778
+ ys = []
779
+
780
+ for idx in prompt_indices:
781
+
782
+ item = samples[
783
+ idx
784
+ ]
785
+
786
+ x = (
787
+ item[
788
+ "layers"
789
+ ][source_layer]
790
+ .float()
791
+ .squeeze(0)
792
+ )
793
+
794
+ # Normalize H3 input per token.
795
+ x = rms_normalize(
796
+ x
797
+ )
798
+
799
+ y = (
800
+ item[
801
+ "teacher"
802
+ ]
803
+ .float()
804
+ .squeeze(0)
805
+ )
806
+
807
+ xs.append(
808
+ x
809
+ )
810
+
811
+ ys.append(
812
+ y
813
+ )
814
+
815
+ return (
816
+ torch.cat(
817
+ xs,
818
+ dim=0,
819
+ ),
820
+ torch.cat(
821
+ ys,
822
+ dim=0,
823
+ ),
824
+ )
825
+
826
+
827
+ # ============================================================
828
+ # EVALUATE PROMPT LEVEL
829
+ # ============================================================
830
+
831
+ def evaluate_prompt_level(
832
+ model,
833
+ source_layer,
834
+ ):
835
+
836
+ teacher_scores = []
837
+
838
+ correction_scores = []
839
+
840
+ blend_scores = {
841
+ alpha: []
842
+ for alpha
843
+ in ALPHAS
844
+ }
845
+
846
+ category_teacher = {}
847
+
848
+ model.eval()
849
+
850
+ with torch.no_grad():
851
+
852
+ for idx in val_indices:
853
+
854
+ item = samples[
855
+ idx
856
+ ]
857
+
858
+ category = item[
859
+ "category"
860
+ ]
861
+
862
+ x = (
863
+ item[
864
+ "layers"
865
+ ][source_layer]
866
+ .float()
867
+ )
868
+
869
+ x = rms_normalize(
870
+ x
871
+ )
872
+
873
+ teacher = (
874
+ item[
875
+ "teacher"
876
+ ]
877
+ .float()
878
+ )
879
+
880
+ h49 = (
881
+ item[
882
+ "h49"
883
+ ]
884
+ .float()
885
+ )
886
+
887
+ pred = (
888
+ model(
889
+ x.to(
890
+ DEVICE
891
+ )
892
+ )
893
+ .cpu()
894
+ )
895
+
896
+
897
+ # =================================================
898
+ # DIRECT TEACHER REPRESENTATION
899
+ # =================================================
900
+
901
+ teacher_cos = (
902
+ cosine_mean(
903
+ pred,
904
+ teacher,
905
+ )
906
+ )
907
+
908
+ teacher_scores.append(
909
+ teacher_cos
910
+ )
911
+
912
+ category_teacher.setdefault(
913
+ category,
914
+ []
915
+ ).append(
916
+ teacher_cos
917
+ )
918
+
919
+
920
+ # =================================================
921
+ # REAL FULL-BRIDGE CORRECTION
922
+ # =================================================
923
+
924
+ teacher_scaled = (
925
+ magnitude_match(
926
+ teacher,
927
+ h49,
928
+ )
929
+ )
930
+
931
+ pred_scaled = (
932
+ magnitude_match(
933
+ pred,
934
+ h49,
935
+ )
936
+ )
937
+
938
+ teacher_delta = (
939
+ teacher_scaled
940
+ - h49
941
+ )
942
+
943
+ pred_delta = (
944
+ pred_scaled
945
+ - h49
946
+ )
947
+
948
+ correction_cos = (
949
+ cosine_mean(
950
+ pred_delta,
951
+ teacher_delta,
952
+ )
953
+ )
954
+
955
+ correction_scores.append(
956
+ correction_cos
957
+ )
958
+
959
+
960
+ # =================================================
961
+ # FULL vs DISTILLED BLENDED CONDITIONING
962
+ # =================================================
963
+
964
+ for alpha in ALPHAS:
965
+
966
+ full_blend = (
967
+ h49
968
+ + alpha
969
+ * teacher_delta
970
+ )
971
+
972
+ distilled_blend = (
973
+ h49
974
+ + alpha
975
+ * pred_delta
976
+ )
977
+
978
+ score = cosine_mean(
979
+ distilled_blend,
980
+ full_blend,
981
+ )
982
+
983
+ blend_scores[
984
+ alpha
985
+ ].append(
986
+ score
987
+ )
988
+
989
+
990
+ result = {
991
+ "teacher_cos_mean":
992
+ sum(
993
+ teacher_scores
994
+ )
995
+ / len(
996
+ teacher_scores
997
+ ),
998
+
999
+ "teacher_cos_min":
1000
+ min(
1001
+ teacher_scores
1002
+ ),
1003
+
1004
+ "teacher_cos_max":
1005
+ max(
1006
+ teacher_scores
1007
+ ),
1008
+
1009
+ "correction_cos_mean":
1010
+ sum(
1011
+ correction_scores
1012
+ )
1013
+ / len(
1014
+ correction_scores
1015
+ ),
1016
+
1017
+ "correction_cos_min":
1018
+ min(
1019
+ correction_scores
1020
+ ),
1021
+
1022
+ "correction_cos_max":
1023
+ max(
1024
+ correction_scores
1025
+ ),
1026
+ }
1027
+
1028
+ for alpha in ALPHAS:
1029
+
1030
+ scores = blend_scores[
1031
+ alpha
1032
+ ]
1033
+
1034
+ result[
1035
+ f"blend_cos_alpha_{alpha}"
1036
+ ] = (
1037
+ sum(
1038
+ scores
1039
+ )
1040
+ / len(
1041
+ scores
1042
+ )
1043
+ )
1044
+
1045
+
1046
+ category_results = {}
1047
+
1048
+ for category, values in (
1049
+ category_teacher.items()
1050
+ ):
1051
+
1052
+ category_results[
1053
+ category
1054
+ ] = (
1055
+ sum(values)
1056
+ / len(values)
1057
+ )
1058
+
1059
+ return (
1060
+ result,
1061
+ category_results,
1062
+ )
1063
+
1064
+
1065
+ # ============================================================
1066
+ # TRAIN ONE STUDENT
1067
+ # ============================================================
1068
+
1069
+ def train_student(
1070
+ source_layer,
1071
+ ):
1072
+
1073
+ name = (
1074
+ f"H3_L{source_layer}"
1075
+ f"_to_projected_SN32_v2"
1076
+ )
1077
+
1078
+ print()
1079
+ print(
1080
+ "=" * 100
1081
+ )
1082
+
1083
+ print(
1084
+ "STUDENT:",
1085
+ name
1086
+ )
1087
+
1088
+ print(
1089
+ "=" * 100
1090
+ )
1091
+
1092
+ train_x, train_y = (
1093
+ build_token_data(
1094
+ source_layer,
1095
+ train_indices,
1096
+ )
1097
+ )
1098
+
1099
+ val_x, val_y = (
1100
+ build_token_data(
1101
+ source_layer,
1102
+ val_indices,
1103
+ )
1104
+ )
1105
+
1106
+ print(
1107
+ "Train tokens:",
1108
+ train_x.shape[0]
1109
+ )
1110
+
1111
+ print(
1112
+ "Validation tokens:",
1113
+ val_x.shape[0]
1114
+ )
1115
+
1116
+ print(
1117
+ "Input:",
1118
+ train_x.shape[-1]
1119
+ )
1120
+
1121
+ print(
1122
+ "Hidden:",
1123
+ HIDDEN_DIM
1124
+ )
1125
+
1126
+ print(
1127
+ "Output:",
1128
+ train_y.shape[-1]
1129
+ )
1130
+
1131
+
1132
+ # ========================================================
1133
+ # MODEL
1134
+ # ========================================================
1135
+
1136
+ model = SemanticStudent(
1137
+ input_dim=5120,
1138
+ hidden_dim=HIDDEN_DIM,
1139
+ output_dim=5120,
1140
+ ).to(
1141
+ DEVICE
1142
+ )
1143
+
1144
+ optimizer = (
1145
+ torch.optim.AdamW(
1146
+ model.parameters(),
1147
+ lr=LR,
1148
+ weight_decay=
1149
+ WEIGHT_DECAY,
1150
+ )
1151
+ )
1152
+
1153
+ scheduler = (
1154
+ torch.optim.lr_scheduler.CosineAnnealingLR(
1155
+ optimizer,
1156
+ T_max=EPOCHS,
1157
+ eta_min=LR * 0.05,
1158
+ )
1159
+ )
1160
+
1161
+
1162
+ # Pin CPU memory.
1163
+ if DEVICE == "cuda":
1164
+
1165
+ train_x = (
1166
+ train_x
1167
+ .contiguous()
1168
+ .pin_memory()
1169
+ )
1170
+
1171
+ train_y = (
1172
+ train_y
1173
+ .contiguous()
1174
+ .pin_memory()
1175
+ )
1176
+
1177
+ val_x = (
1178
+ val_x
1179
+ .contiguous()
1180
+ .pin_memory()
1181
+ )
1182
+
1183
+ val_y = (
1184
+ val_y
1185
+ .contiguous()
1186
+ .pin_memory()
1187
+ )
1188
+
1189
+
1190
+ best_cos = -1.0
1191
+
1192
+ best_path = os.path.join(
1193
+ OUT_DIR,
1194
+ f"{name}.safetensors"
1195
+ )
1196
+
1197
+ token_indices = (
1198
+ torch.arange(
1199
+ train_x.shape[0]
1200
+ )
1201
+ )
1202
+
1203
+
1204
+ # ========================================================
1205
+ # EPOCHS
1206
+ # ========================================================
1207
+
1208
+ for epoch in range(
1209
+ 1,
1210
+ EPOCHS + 1,
1211
+ ):
1212
+
1213
+ permutation = (
1214
+ token_indices[
1215
+ torch.randperm(
1216
+ len(
1217
+ token_indices
1218
+ )
1219
+ )
1220
+ ]
1221
+ )
1222
+
1223
+ model.train()
1224
+
1225
+ epoch_loss = 0.0
1226
+
1227
+ batches = 0
1228
+
1229
+ for start in range(
1230
+ 0,
1231
+ len(
1232
+ permutation
1233
+ ),
1234
+ BATCH_TOKENS,
1235
+ ):
1236
+
1237
+ ids = (
1238
+ permutation[
1239
+ start:
1240
+ start
1241
+ + BATCH_TOKENS
1242
+ ]
1243
+ )
1244
+
1245
+ x = (
1246
+ train_x[
1247
+ ids
1248
+ ]
1249
+ .to(
1250
+ DEVICE,
1251
+ non_blocking=True,
1252
+ )
1253
+ )
1254
+
1255
+ y = (
1256
+ train_y[
1257
+ ids
1258
+ ]
1259
+ .to(
1260
+ DEVICE,
1261
+ non_blocking=True,
1262
+ )
1263
+ )
1264
+
1265
+ optimizer.zero_grad(
1266
+ set_to_none=True
1267
+ )
1268
+
1269
+ pred = model(
1270
+ x
1271
+ )
1272
+
1273
+ loss, _, _, _ = (
1274
+ student_loss(
1275
+ pred,
1276
+ y,
1277
+ )
1278
+ )
1279
+
1280
+ loss.backward()
1281
+
1282
+ torch.nn.utils.clip_grad_norm_(
1283
+ model.parameters(),
1284
+ 1.0,
1285
+ )
1286
+
1287
+ optimizer.step()
1288
+
1289
+ epoch_loss += (
1290
+ loss.item()
1291
+ )
1292
+
1293
+ batches += 1
1294
+
1295
+
1296
+ scheduler.step()
1297
+
1298
+
1299
+ # ====================================================
1300
+ # VALIDATION TOKEN COSINE
1301
+ # ====================================================
1302
+
1303
+ model.eval()
1304
+
1305
+ val_cos_values = []
1306
+
1307
+ val_loss_total = 0.0
1308
+
1309
+ val_batches = 0
1310
+
1311
+ with torch.no_grad():
1312
+
1313
+ for start in range(
1314
+ 0,
1315
+ val_x.shape[0],
1316
+ BATCH_TOKENS,
1317
+ ):
1318
+
1319
+ x = (
1320
+ val_x[
1321
+ start:
1322
+ start
1323
+ + BATCH_TOKENS
1324
+ ]
1325
+ .to(
1326
+ DEVICE,
1327
+ non_blocking=True,
1328
+ )
1329
+ )
1330
+
1331
+ y = (
1332
+ val_y[
1333
+ start:
1334
+ start
1335
+ + BATCH_TOKENS
1336
+ ]
1337
+ .to(
1338
+ DEVICE,
1339
+ non_blocking=True,
1340
+ )
1341
+ )
1342
+
1343
+ pred = model(
1344
+ x
1345
+ )
1346
+
1347
+ loss, _, _, _ = (
1348
+ student_loss(
1349
+ pred,
1350
+ y,
1351
+ )
1352
+ )
1353
+
1354
+ val_loss_total += (
1355
+ loss.item()
1356
+ )
1357
+
1358
+ val_batches += 1
1359
+
1360
+ cos = cosine(
1361
+ pred,
1362
+ y,
1363
+ )
1364
+
1365
+ val_cos_values.append(
1366
+ cos.detach().cpu()
1367
+ )
1368
+
1369
+
1370
+ val_cos_tensor = torch.cat(
1371
+ val_cos_values
1372
+ )
1373
+
1374
+ val_cos = (
1375
+ val_cos_tensor
1376
+ .mean()
1377
+ .item()
1378
+ )
1379
+
1380
+ val_loss = (
1381
+ val_loss_total
1382
+ / max(
1383
+ val_batches,
1384
+ 1,
1385
+ )
1386
+ )
1387
+
1388
+ lr_now = (
1389
+ optimizer
1390
+ .param_groups[0]["lr"]
1391
+ )
1392
+
1393
+ print(
1394
+ f"Epoch "
1395
+ f"{epoch:02d}/"
1396
+ f"{EPOCHS} | "
1397
+ f"train_loss="
1398
+ f"{epoch_loss / max(batches,1):.6f} | "
1399
+ f"val_loss="
1400
+ f"{val_loss:.6f} | "
1401
+ f"val_teacher_cos="
1402
+ f"{val_cos:.6f} | "
1403
+ f"lr="
1404
+ f"{lr_now:.7f}"
1405
+ )
1406
+
1407
+
1408
+ # ====================================================
1409
+ # SAVE BEST
1410
+ # ====================================================
1411
+
1412
+ if (
1413
+ val_cos
1414
+ > best_cos
1415
+ ):
1416
+ best_cos = (
1417
+ val_cos
1418
+ )
1419
+
1420
+ save_student(
1421
+ model,
1422
+ best_path,
1423
+ source_layer,
1424
+ best_cos,
1425
+ )
1426
+
1427
+
1428
+ # ========================================================
1429
+ # LOAD BEST
1430
+ # ========================================================
1431
+
1432
+ weights = load_file(
1433
+ best_path,
1434
+ device="cpu",
1435
+ )
1436
+
1437
+ best_model = SemanticStudent(
1438
+ 5120,
1439
+ HIDDEN_DIM,
1440
+ 5120,
1441
+ )
1442
+
1443
+ with torch.no_grad():
1444
+
1445
+ best_model.fc1.weight.copy_(
1446
+ weights[
1447
+ "fc1.weight"
1448
+ ].float()
1449
+ )
1450
+
1451
+ best_model.fc1.bias.copy_(
1452
+ weights[
1453
+ "fc1.bias"
1454
+ ].float()
1455
+ )
1456
+
1457
+ best_model.fc2.weight.copy_(
1458
+ weights[
1459
+ "fc2.weight"
1460
+ ].float()
1461
+ )
1462
+
1463
+ best_model.fc2.bias.copy_(
1464
+ weights[
1465
+ "fc2.bias"
1466
+ ].float()
1467
+ )
1468
+
1469
+ best_model.fc3.weight.copy_(
1470
+ weights[
1471
+ "fc3.weight"
1472
+ ].float()
1473
+ )
1474
+
1475
+ best_model.fc3.bias.copy_(
1476
+ weights[
1477
+ "fc3.bias"
1478
+ ].float()
1479
+ )
1480
+
1481
+
1482
+ best_model = best_model.to(
1483
+ DEVICE
1484
+ )
1485
+
1486
+ best_model.eval()
1487
+
1488
+
1489
+ # ========================================================
1490
+ # REAL BRIDGE EVALUATION
1491
+ # ========================================================
1492
+
1493
+ metrics, category_metrics = (
1494
+ evaluate_prompt_level(
1495
+ best_model,
1496
+ source_layer,
1497
+ )
1498
+ )
1499
+
1500
+
1501
+ print()
1502
+ print(
1503
+ "BEST RESULT"
1504
+ )
1505
+
1506
+ print(
1507
+ "-" * 100
1508
+ )
1509
+
1510
+ print(
1511
+ "Teacher representation cosine:",
1512
+ f"{metrics['teacher_cos_mean']:.6f}"
1513
+ )
1514
+
1515
+ print(
1516
+ "Teacher cosine min:",
1517
+ f"{metrics['teacher_cos_min']:.6f}"
1518
+ )
1519
+
1520
+ print(
1521
+ "Teacher cosine max:",
1522
+ f"{metrics['teacher_cos_max']:.6f}"
1523
+ )
1524
+
1525
+ print()
1526
+
1527
+ print(
1528
+ "Correction cosine:",
1529
+ f"{metrics['correction_cos_mean']:.6f}"
1530
+ )
1531
+
1532
+ print(
1533
+ "Correction cosine min:",
1534
+ f"{metrics['correction_cos_min']:.6f}"
1535
+ )
1536
+
1537
+ print(
1538
+ "Correction cosine max:",
1539
+ f"{metrics['correction_cos_max']:.6f}"
1540
+ )
1541
+
1542
+ print()
1543
+
1544
+ for alpha in ALPHAS:
1545
+
1546
+ print(
1547
+ f"Full vs Distilled blend "
1548
+ f"alpha={alpha:.2f}: "
1549
+ f"{metrics[f'blend_cos_alpha_{alpha}']:.6f}"
1550
+ )
1551
+
1552
+
1553
+ print()
1554
+ print(
1555
+ "BY CATEGORY "
1556
+ "(student -> projected SenseNova):"
1557
+ )
1558
+
1559
+ for category in sorted(
1560
+ category_metrics
1561
+ ):
1562
+
1563
+ print(
1564
+ f"{category:28s} "
1565
+ f"{category_metrics[category]:.6f}"
1566
+ )
1567
+
1568
+
1569
+ return {
1570
+ "source_layer":
1571
+ source_layer,
1572
+
1573
+ "best_token_teacher_cos":
1574
+ best_cos,
1575
+
1576
+ **metrics,
1577
+
1578
+ "model_file":
1579
+ best_path,
1580
+ }
1581
+
1582
+
1583
+ # ============================================================
1584
+ # TRAIN BOTH
1585
+ # ============================================================
1586
+
1587
+ all_results = []
1588
+
1589
+ for layer in SOURCE_LAYERS:
1590
+
1591
+ result = train_student(
1592
+ layer
1593
+ )
1594
+
1595
+ all_results.append(
1596
+ result
1597
+ )
1598
+
1599
+
1600
+ # ============================================================
1601
+ # RANK
1602
+ # ============================================================
1603
+
1604
+ all_results = sorted(
1605
+ all_results,
1606
+ key=lambda x:
1607
+ x[
1608
+ "correction_cos_mean"
1609
+ ],
1610
+ reverse=True,
1611
+ )
1612
+
1613
+
1614
+ # ============================================================
1615
+ # CSV
1616
+ # ============================================================
1617
+
1618
+ csv_path = os.path.join(
1619
+ OUT_DIR,
1620
+ "distilled_student_v2_results.csv"
1621
+ )
1622
+
1623
+ fieldnames = [
1624
+ "source_layer",
1625
+ "best_token_teacher_cos",
1626
+ "teacher_cos_mean",
1627
+ "teacher_cos_min",
1628
+ "teacher_cos_max",
1629
+ "correction_cos_mean",
1630
+ "correction_cos_min",
1631
+ "correction_cos_max",
1632
+ ]
1633
+
1634
+ for alpha in ALPHAS:
1635
+
1636
+ fieldnames.append(
1637
+ f"blend_cos_alpha_{alpha}"
1638
+ )
1639
+
1640
+ fieldnames.append(
1641
+ "model_file"
1642
+ )
1643
+
1644
+
1645
+ with open(
1646
+ csv_path,
1647
+ "w",
1648
+ newline="",
1649
+ encoding="utf-8",
1650
+ ) as f:
1651
+
1652
+ writer = csv.DictWriter(
1653
+ f,
1654
+ fieldnames=fieldnames,
1655
+ )
1656
+
1657
+ writer.writeheader()
1658
+
1659
+ writer.writerows(
1660
+ all_results
1661
+ )
1662
+
1663
+
1664
+ # ============================================================
1665
+ # REPORT
1666
+ # ============================================================
1667
+
1668
+ report_path = os.path.join(
1669
+ OUT_DIR,
1670
+ "distilled_student_v2_results.txt"
1671
+ )
1672
+
1673
+
1674
+ with open(
1675
+ report_path,
1676
+ "w",
1677
+ encoding="utf-8",
1678
+ ) as f:
1679
+
1680
+ f.write(
1681
+ "DISTILLED STUDENT V2 RESULTS\n"
1682
+ )
1683
+
1684
+ f.write(
1685
+ "=" * 100
1686
+ + "\n\n"
1687
+ )
1688
+
1689
+ for rank, r in enumerate(
1690
+ all_results,
1691
+ 1,
1692
+ ):
1693
+
1694
+ f.write(
1695
+ f"{rank:02d}. "
1696
+ f"H3 L{r['source_layer']}\n"
1697
+ )
1698
+
1699
+ f.write(
1700
+ f" teacher_cos_mean: "
1701
+ f"{r['teacher_cos_mean']:.6f}\n"
1702
+ )
1703
+
1704
+ f.write(
1705
+ f" teacher_cos_min: "
1706
+ f"{r['teacher_cos_min']:.6f}\n"
1707
+ )
1708
+
1709
+ f.write(
1710
+ f" teacher_cos_max: "
1711
+ f"{r['teacher_cos_max']:.6f}\n"
1712
+ )
1713
+
1714
+ f.write(
1715
+ f" correction_cos: "
1716
+ f"{r['correction_cos_mean']:.6f}\n"
1717
+ )
1718
+
1719
+ f.write(
1720
+ f" correction_min: "
1721
+ f"{r['correction_cos_min']:.6f}\n"
1722
+ )
1723
+
1724
+ f.write(
1725
+ f" correction_max: "
1726
+ f"{r['correction_cos_max']:.6f}\n"
1727
+ )
1728
+
1729
+
1730
+ for alpha in ALPHAS:
1731
+
1732
+ f.write(
1733
+ f" blend alpha "
1734
+ f"{alpha:.2f}: "
1735
+ f"{r[f'blend_cos_alpha_{alpha}']:.6f}\n"
1736
+ )
1737
+
1738
+
1739
+ f.write(
1740
+ f" model: "
1741
+ f"{r['model_file']}\n\n"
1742
+ )
1743
+
1744
+
1745
+ # ============================================================
1746
+ # FINAL
1747
+ # ============================================================
1748
+
1749
+ print()
1750
+ print(
1751
+ "=" * 100
1752
+ )
1753
+
1754
+ print(
1755
+ "FINAL V2 RANKING"
1756
+ )
1757
+
1758
+ print(
1759
+ "=" * 100
1760
+ )
1761
+
1762
+
1763
+ for rank, r in enumerate(
1764
+ all_results,
1765
+ 1,
1766
+ ):
1767
+
1768
+ print(
1769
+ f"{rank:02d}. "
1770
+ f"H3 L{r['source_layer']:2d} | "
1771
+ f"teacher_cos="
1772
+ f"{r['teacher_cos_mean']:.6f} | "
1773
+ f"correction_cos="
1774
+ f"{r['correction_cos_mean']:.6f} | "
1775
+ f"blend@0.20="
1776
+ f"{r['blend_cos_alpha_0.2']:.6f}"
1777
+ )
1778
+
1779
+
1780
+ print()
1781
+ print(
1782
+ "Report:"
1783
+ )
1784
+
1785
+ print(
1786
+ report_path
1787
+ )
1788
+
1789
+ print()
1790
+
1791
+ print(
1792
+ "CSV:"
1793
+ )
1794
+
1795
+ print(
1796
+ csv_path
1797
+ )
1798
+
1799
+ print()
1800
+ print(
1801
+ "DONE"
1802
+ )
research/raw_scripts/train_distilled_student_v3.py ADDED
@@ -0,0 +1,1628 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import random
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ from safetensors.torch import load_file, save_file
10
+
11
+
12
+ # ============================================================
13
+ # PATHS
14
+ # ============================================================
15
+
16
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
17
+
18
+ H3_MAIN = os.path.join(
19
+ ROOT,
20
+ "h3_bridge_hidden_480.pt"
21
+ )
22
+
23
+ SN_MAIN = os.path.join(
24
+ ROOT,
25
+ "sensenova_bridge_hidden_480.pt"
26
+ )
27
+
28
+ H3_OOD = os.path.join(
29
+ ROOT,
30
+ "h3_ood_hidden_160.pt"
31
+ )
32
+
33
+ SN_OOD = os.path.join(
34
+ ROOT,
35
+ "sensenova_ood_hidden_160.pt"
36
+ )
37
+
38
+ TEACHER_BRIDGE = (
39
+ r"D:\ComfyUI_Krea2\ComfyUI\models\bridge"
40
+ r"\SN_L32_to_H3_L49_rank128.safetensors"
41
+ )
42
+
43
+ V2_INIT = os.path.join(
44
+ ROOT,
45
+ "distilled_student_v2",
46
+ "H3_L49_to_projected_SN32_v2.safetensors"
47
+ )
48
+
49
+ OUT_DIR = os.path.join(
50
+ ROOT,
51
+ "distilled_student_v3"
52
+ )
53
+
54
+ os.makedirs(
55
+ OUT_DIR,
56
+ exist_ok=True
57
+ )
58
+
59
+ BEST_FILE = os.path.join(
60
+ OUT_DIR,
61
+ "H3_L49_to_projected_SN32_v3.safetensors"
62
+ )
63
+
64
+ REPORT_FILE = os.path.join(
65
+ OUT_DIR,
66
+ "distilled_student_v3_report.txt"
67
+ )
68
+
69
+ CSV_FILE = os.path.join(
70
+ OUT_DIR,
71
+ "distilled_student_v3_history.csv"
72
+ )
73
+
74
+
75
+ # ============================================================
76
+ # SETTINGS
77
+ # ============================================================
78
+
79
+ SEED = 20260903
80
+
81
+ DEVICE = (
82
+ "cuda"
83
+ if torch.cuda.is_available()
84
+ else "cpu"
85
+ )
86
+
87
+ HIDDEN_DIM = 512
88
+
89
+ EPOCHS = 40
90
+
91
+ LR = 3e-4
92
+
93
+ WEIGHT_DECAY = 1e-5
94
+
95
+ BATCH_TOKENS = 1024
96
+
97
+ VAL_PROMPTS = 100
98
+
99
+ ALPHAS = [
100
+ 0.10,
101
+ 0.20,
102
+ 0.30,
103
+ ]
104
+
105
+
106
+ # ============================================================
107
+ # REPRO
108
+ # ============================================================
109
+
110
+ random.seed(
111
+ SEED
112
+ )
113
+
114
+ torch.manual_seed(
115
+ SEED
116
+ )
117
+
118
+ if torch.cuda.is_available():
119
+ torch.cuda.manual_seed_all(
120
+ SEED
121
+ )
122
+
123
+
124
+ # ============================================================
125
+ # MODEL
126
+ # ============================================================
127
+
128
+ class SemanticStudent(
129
+ nn.Module
130
+ ):
131
+ def __init__(
132
+ self,
133
+ input_dim=5120,
134
+ hidden_dim=512,
135
+ output_dim=5120,
136
+ ):
137
+ super().__init__()
138
+
139
+ self.fc1 = nn.Linear(
140
+ input_dim,
141
+ hidden_dim,
142
+ bias=True,
143
+ )
144
+
145
+ self.fc2 = nn.Linear(
146
+ hidden_dim,
147
+ hidden_dim,
148
+ bias=True,
149
+ )
150
+
151
+ self.fc3 = nn.Linear(
152
+ hidden_dim,
153
+ output_dim,
154
+ bias=True,
155
+ )
156
+
157
+ self.act = nn.SiLU()
158
+
159
+
160
+ def forward(
161
+ self,
162
+ x,
163
+ ):
164
+ x = self.fc1(
165
+ x
166
+ )
167
+
168
+ x = self.act(
169
+ x
170
+ )
171
+
172
+ x = self.fc2(
173
+ x
174
+ )
175
+
176
+ x = self.act(
177
+ x
178
+ )
179
+
180
+ x = self.fc3(
181
+ x
182
+ )
183
+
184
+ return x
185
+
186
+
187
+ # ============================================================
188
+ # HELPERS
189
+ # ============================================================
190
+
191
+ def rms_normalize(
192
+ x,
193
+ ):
194
+ x = x.float()
195
+
196
+ rms = torch.sqrt(
197
+ x.pow(2).mean(
198
+ dim=-1,
199
+ keepdim=True,
200
+ ) + 1e-6
201
+ )
202
+
203
+ return x / rms
204
+
205
+
206
+ def magnitude_match(
207
+ source,
208
+ target,
209
+ ):
210
+ source = source.float()
211
+ target = target.float()
212
+
213
+ source_rms = torch.sqrt(
214
+ source.pow(2).mean(
215
+ dim=-1,
216
+ keepdim=True,
217
+ ) + 1e-8
218
+ )
219
+
220
+ target_rms = torch.sqrt(
221
+ target.pow(2).mean(
222
+ dim=-1,
223
+ keepdim=True,
224
+ ) + 1e-8
225
+ )
226
+
227
+ return (
228
+ source
229
+ * (
230
+ target_rms
231
+ / source_rms
232
+ )
233
+ )
234
+
235
+
236
+ def cosine_mean(
237
+ a,
238
+ b,
239
+ ):
240
+ return (
241
+ F.cosine_similarity(
242
+ a.float(),
243
+ b.float(),
244
+ dim=-1,
245
+ )
246
+ .mean()
247
+ .item()
248
+ )
249
+
250
+
251
+ # ============================================================
252
+ # SAVE
253
+ # ============================================================
254
+
255
+ def save_student(
256
+ model,
257
+ path,
258
+ best_score,
259
+ ):
260
+ tensors = {
261
+ "fc1.weight":
262
+ model.fc1.weight
263
+ .detach()
264
+ .cpu()
265
+ .to(torch.float16),
266
+
267
+ "fc1.bias":
268
+ model.fc1.bias
269
+ .detach()
270
+ .cpu()
271
+ .to(torch.float16),
272
+
273
+ "fc2.weight":
274
+ model.fc2.weight
275
+ .detach()
276
+ .cpu()
277
+ .to(torch.float16),
278
+
279
+ "fc2.bias":
280
+ model.fc2.bias
281
+ .detach()
282
+ .cpu()
283
+ .to(torch.float16),
284
+
285
+ "fc3.weight":
286
+ model.fc3.weight
287
+ .detach()
288
+ .cpu()
289
+ .to(torch.float16),
290
+
291
+ "fc3.bias":
292
+ model.fc3.bias
293
+ .detach()
294
+ .cpu()
295
+ .to(torch.float16),
296
+ }
297
+
298
+ metadata = {
299
+ "architecture":
300
+ "SenseNova_H3_Distilled_Student_V3",
301
+
302
+ "source":
303
+ "H3_L49",
304
+
305
+ "teacher":
306
+ "SenseNova_L32_projected_to_H3",
307
+
308
+ "input_dim":
309
+ "5120",
310
+
311
+ "hidden_dim":
312
+ "512",
313
+
314
+ "output_dim":
315
+ "5120",
316
+
317
+ "activation":
318
+ "SiLU",
319
+
320
+ "best_score":
321
+ f"{best_score:.8f}",
322
+ }
323
+
324
+ save_file(
325
+ tensors,
326
+ path,
327
+ metadata=metadata,
328
+ )
329
+
330
+
331
+ # ============================================================
332
+ # LOAD DATASETS
333
+ # ============================================================
334
+
335
+ print("=" * 100)
336
+ print("DISTILLED STUDENT V3")
337
+ print("TRAINING ON 600 PROMPTS")
338
+ print("=" * 100)
339
+
340
+ print(
341
+ "Device:",
342
+ DEVICE
343
+ )
344
+
345
+
346
+ print(
347
+ "\nLoading original 440..."
348
+ )
349
+
350
+ h3_main = torch.load(
351
+ H3_MAIN,
352
+ map_location="cpu",
353
+ weights_only=False,
354
+ )
355
+
356
+ sn_main = torch.load(
357
+ SN_MAIN,
358
+ map_location="cpu",
359
+ weights_only=False,
360
+ )
361
+
362
+
363
+ print(
364
+ "Loading strict OOD 160..."
365
+ )
366
+
367
+ h3_ood = torch.load(
368
+ H3_OOD,
369
+ map_location="cpu",
370
+ weights_only=False,
371
+ )
372
+
373
+ sn_ood = torch.load(
374
+ SN_OOD,
375
+ map_location="cpu",
376
+ weights_only=False,
377
+ )
378
+
379
+
380
+ # ============================================================
381
+ # LOAD TEACHER
382
+ # ============================================================
383
+
384
+ teacher = load_file(
385
+ TEACHER_BRIDGE,
386
+ device="cpu",
387
+ )
388
+
389
+ teacher_down = (
390
+ teacher[
391
+ "down.weight"
392
+ ]
393
+ .float()
394
+ )
395
+
396
+ teacher_up = (
397
+ teacher[
398
+ "up.weight"
399
+ ]
400
+ .float()
401
+ )
402
+
403
+
404
+ # ============================================================
405
+ # COLLECT ALL 600
406
+ # ============================================================
407
+
408
+ samples = []
409
+
410
+
411
+ def append_main_pair(
412
+ h3_item,
413
+ sn_item,
414
+ origin,
415
+ ):
416
+
417
+ if (
418
+ h3_item["prompt"]
419
+ != sn_item["prompt"]
420
+ ):
421
+ raise RuntimeError(
422
+ "Prompt mismatch."
423
+ )
424
+
425
+ if not torch.equal(
426
+ h3_item["input_ids"],
427
+ sn_item["input_ids"],
428
+ ):
429
+ raise RuntimeError(
430
+ "Token mismatch."
431
+ )
432
+
433
+ h49 = (
434
+ h3_item[
435
+ "layers"
436
+ ][49]
437
+ .float()
438
+ )
439
+
440
+ sn32 = (
441
+ sn_item[
442
+ "layers"
443
+ ][32]
444
+ .float()
445
+ )
446
+
447
+ with torch.no_grad():
448
+
449
+ sn_norm = rms_normalize(
450
+ sn32
451
+ )
452
+
453
+ low = F.linear(
454
+ sn_norm,
455
+ teacher_down,
456
+ )
457
+
458
+ teacher_projected = F.linear(
459
+ low,
460
+ teacher_up,
461
+ )
462
+
463
+ samples.append({
464
+ "prompt":
465
+ h3_item["prompt"],
466
+
467
+ "category":
468
+ h3_item["category"],
469
+
470
+ "origin":
471
+ origin,
472
+
473
+ "h49":
474
+ h49.to(
475
+ torch.float16
476
+ ),
477
+
478
+ "teacher":
479
+ teacher_projected.to(
480
+ torch.float16
481
+ ),
482
+ })
483
+
484
+
485
+ def append_ood_pair(
486
+ h3_item,
487
+ sn_item,
488
+ origin,
489
+ ):
490
+
491
+ if (
492
+ h3_item["prompt"]
493
+ != sn_item["prompt"]
494
+ ):
495
+ raise RuntimeError(
496
+ "Prompt mismatch."
497
+ )
498
+
499
+ if not torch.equal(
500
+ h3_item["input_ids"],
501
+ sn_item["input_ids"],
502
+ ):
503
+ raise RuntimeError(
504
+ "Token mismatch."
505
+ )
506
+
507
+ h49 = (
508
+ h3_item[
509
+ "hidden"
510
+ ]
511
+ .float()
512
+ )
513
+
514
+ sn32 = (
515
+ sn_item[
516
+ "hidden"
517
+ ]
518
+ .float()
519
+ )
520
+
521
+ with torch.no_grad():
522
+
523
+ sn_norm = rms_normalize(
524
+ sn32
525
+ )
526
+
527
+ low = F.linear(
528
+ sn_norm,
529
+ teacher_down,
530
+ )
531
+
532
+ teacher_projected = F.linear(
533
+ low,
534
+ teacher_up,
535
+ )
536
+
537
+ samples.append({
538
+ "prompt":
539
+ h3_item["prompt"],
540
+
541
+ "category":
542
+ h3_item["category"],
543
+
544
+ "origin":
545
+ origin,
546
+
547
+ "h49":
548
+ h49.to(
549
+ torch.float16
550
+ ),
551
+
552
+ "teacher":
553
+ teacher_projected.to(
554
+ torch.float16
555
+ ),
556
+ })
557
+
558
+
559
+ for h3_item, sn_item in zip(
560
+ h3_main["results"],
561
+ sn_main["results"],
562
+ ):
563
+ append_main_pair(
564
+ h3_item,
565
+ sn_item,
566
+ "main440",
567
+ )
568
+
569
+
570
+ for h3_item, sn_item in zip(
571
+ h3_ood["results"],
572
+ sn_ood["results"],
573
+ ):
574
+ append_ood_pair(
575
+ h3_item,
576
+ sn_item,
577
+ "ood160",
578
+ )
579
+
580
+
581
+ print()
582
+ print(
583
+ "Total samples:",
584
+ len(samples)
585
+ )
586
+
587
+ if len(samples) != 600:
588
+ raise RuntimeError(
589
+ f"Expected 600 prompts, got {len(samples)}"
590
+ )
591
+
592
+
593
+ # ============================================================
594
+ # SPLIT 500 / 100
595
+ # STRATIFIED-ish BY ORIGIN
596
+ # ============================================================
597
+
598
+ main_indices = [
599
+ i
600
+ for i, s
601
+ in enumerate(samples)
602
+ if s["origin"] == "main440"
603
+ ]
604
+
605
+ ood_indices = [
606
+ i
607
+ for i, s
608
+ in enumerate(samples)
609
+ if s["origin"] == "ood160"
610
+ ]
611
+
612
+ random.shuffle(
613
+ main_indices
614
+ )
615
+
616
+ random.shuffle(
617
+ ood_indices
618
+ )
619
+
620
+ # Keep both distributions represented in validation.
621
+ val_main = main_indices[:70]
622
+ val_ood = ood_indices[:30]
623
+
624
+ val_indices = sorted(
625
+ val_main + val_ood
626
+ )
627
+
628
+ val_set = set(
629
+ val_indices
630
+ )
631
+
632
+ train_indices = [
633
+ i
634
+ for i
635
+ in range(
636
+ len(samples)
637
+ )
638
+ if i not in val_set
639
+ ]
640
+
641
+
642
+ print(
643
+ "Train prompts:",
644
+ len(train_indices)
645
+ )
646
+
647
+ print(
648
+ "Validation prompts:",
649
+ len(val_indices)
650
+ )
651
+
652
+ print(
653
+ "Validation main:",
654
+ len(val_main)
655
+ )
656
+
657
+ print(
658
+ "Validation OOD:",
659
+ len(val_ood)
660
+ )
661
+
662
+
663
+ # ============================================================
664
+ # BUILD TOKEN MATRICES
665
+ # ============================================================
666
+
667
+ def build_token_data(
668
+ indices,
669
+ ):
670
+ xs = []
671
+ ys = []
672
+
673
+ for idx in indices:
674
+
675
+ item = samples[
676
+ idx
677
+ ]
678
+
679
+ x = (
680
+ item["h49"]
681
+ .float()
682
+ .squeeze(0)
683
+ )
684
+
685
+ y = (
686
+ item["teacher"]
687
+ .float()
688
+ .squeeze(0)
689
+ )
690
+
691
+ x = rms_normalize(
692
+ x
693
+ )
694
+
695
+ xs.append(
696
+ x
697
+ )
698
+
699
+ ys.append(
700
+ y
701
+ )
702
+
703
+ return (
704
+ torch.cat(
705
+ xs,
706
+ dim=0,
707
+ ),
708
+ torch.cat(
709
+ ys,
710
+ dim=0,
711
+ ),
712
+ )
713
+
714
+
715
+ print(
716
+ "\nBuilding token matrices..."
717
+ )
718
+
719
+ train_x, train_y = (
720
+ build_token_data(
721
+ train_indices
722
+ )
723
+ )
724
+
725
+ val_x, val_y = (
726
+ build_token_data(
727
+ val_indices
728
+ )
729
+ )
730
+
731
+
732
+ print(
733
+ "Train tokens:",
734
+ train_x.shape[0]
735
+ )
736
+
737
+ print(
738
+ "Val tokens:",
739
+ val_x.shape[0]
740
+ )
741
+
742
+
743
+ if DEVICE == "cuda":
744
+
745
+ train_x = (
746
+ train_x
747
+ .contiguous()
748
+ .pin_memory()
749
+ )
750
+
751
+ train_y = (
752
+ train_y
753
+ .contiguous()
754
+ .pin_memory()
755
+ )
756
+
757
+ val_x = (
758
+ val_x
759
+ .contiguous()
760
+ .pin_memory()
761
+ )
762
+
763
+ val_y = (
764
+ val_y
765
+ .contiguous()
766
+ .pin_memory()
767
+ )
768
+
769
+
770
+ # ============================================================
771
+ # MODEL
772
+ # ============================================================
773
+
774
+ model = SemanticStudent(
775
+ 5120,
776
+ HIDDEN_DIM,
777
+ 5120,
778
+ )
779
+
780
+
781
+ # ============================================================
782
+ # INIT FROM V2
783
+ # ============================================================
784
+
785
+ if os.path.exists(
786
+ V2_INIT
787
+ ):
788
+
789
+ print(
790
+ "\nLoading V2 checkpoint as initialization..."
791
+ )
792
+
793
+ v2 = load_file(
794
+ V2_INIT,
795
+ device="cpu",
796
+ )
797
+
798
+ with torch.no_grad():
799
+
800
+ model.fc1.weight.copy_(
801
+ v2["fc1.weight"].float()
802
+ )
803
+
804
+ model.fc1.bias.copy_(
805
+ v2["fc1.bias"].float()
806
+ )
807
+
808
+ model.fc2.weight.copy_(
809
+ v2["fc2.weight"].float()
810
+ )
811
+
812
+ model.fc2.bias.copy_(
813
+ v2["fc2.bias"].float()
814
+ )
815
+
816
+ model.fc3.weight.copy_(
817
+ v2["fc3.weight"].float()
818
+ )
819
+
820
+ model.fc3.bias.copy_(
821
+ v2["fc3.bias"].float()
822
+ )
823
+
824
+
825
+ model = model.to(
826
+ DEVICE
827
+ )
828
+
829
+
830
+ # ============================================================
831
+ # LOSS
832
+ # ============================================================
833
+
834
+ def loss_fn(
835
+ pred,
836
+ teacher,
837
+ h49,
838
+ ):
839
+
840
+ pred = pred.float()
841
+ teacher = teacher.float()
842
+ h49 = h49.float()
843
+
844
+
845
+ # ========================================================
846
+ # 1. REPRESENTATION DIRECTION
847
+ # ========================================================
848
+
849
+ rep_cos_loss = (
850
+ 1.0
851
+ - F.cosine_similarity(
852
+ pred,
853
+ teacher,
854
+ dim=-1,
855
+ ).mean()
856
+ )
857
+
858
+
859
+ # ========================================================
860
+ # 2. MATCH REAL FULL-BRIDGE CORRECTION
861
+ # ========================================================
862
+
863
+ teacher_scaled = magnitude_match(
864
+ teacher,
865
+ h49,
866
+ )
867
+
868
+ pred_scaled = magnitude_match(
869
+ pred,
870
+ h49,
871
+ )
872
+
873
+ teacher_delta = (
874
+ teacher_scaled
875
+ - h49
876
+ )
877
+
878
+ pred_delta = (
879
+ pred_scaled
880
+ - h49
881
+ )
882
+
883
+ correction_loss = (
884
+ 1.0
885
+ - F.cosine_similarity(
886
+ pred_delta,
887
+ teacher_delta,
888
+ dim=-1,
889
+ ).mean()
890
+ )
891
+
892
+
893
+ # ========================================================
894
+ # 3. SMALL NORMALIZED MSE
895
+ # ========================================================
896
+
897
+ pred_n = F.normalize(
898
+ pred,
899
+ dim=-1,
900
+ )
901
+
902
+ teacher_n = F.normalize(
903
+ teacher,
904
+ dim=-1,
905
+ )
906
+
907
+ mse = F.mse_loss(
908
+ pred_n,
909
+ teacher_n,
910
+ )
911
+
912
+
913
+ # ========================================================
914
+ # TOTAL
915
+ # ========================================================
916
+
917
+ total = (
918
+ 0.55 * rep_cos_loss
919
+ + 0.40 * correction_loss
920
+ + 0.05 * mse
921
+ )
922
+
923
+ return total
924
+
925
+
926
+ # ============================================================
927
+ # OPTIMIZER
928
+ # ============================================================
929
+
930
+ optimizer = torch.optim.AdamW(
931
+ model.parameters(),
932
+ lr=LR,
933
+ weight_decay=
934
+ WEIGHT_DECAY,
935
+ )
936
+
937
+ scheduler = (
938
+ torch.optim.lr_scheduler.CosineAnnealingLR(
939
+ optimizer,
940
+ T_max=EPOCHS,
941
+ eta_min=LR * 0.05,
942
+ )
943
+ )
944
+
945
+
946
+ # ============================================================
947
+ # VALIDATION
948
+ # ============================================================
949
+
950
+ def evaluate():
951
+
952
+ model.eval()
953
+
954
+ rep_scores = []
955
+
956
+ correction_scores = []
957
+
958
+ blend_scores = {
959
+ alpha: []
960
+ for alpha
961
+ in ALPHAS
962
+ }
963
+
964
+ origin_scores = {
965
+ "main440": [],
966
+ "ood160": [],
967
+ }
968
+
969
+ category_scores = {}
970
+
971
+
972
+ with torch.no_grad():
973
+
974
+ for idx in val_indices:
975
+
976
+ item = samples[
977
+ idx
978
+ ]
979
+
980
+ h49 = (
981
+ item["h49"]
982
+ .float()
983
+ )
984
+
985
+ teacher = (
986
+ item["teacher"]
987
+ .float()
988
+ )
989
+
990
+ x = rms_normalize(
991
+ h49
992
+ )
993
+
994
+ pred = (
995
+ model(
996
+ x.to(
997
+ DEVICE
998
+ )
999
+ )
1000
+ .cpu()
1001
+ )
1002
+
1003
+ rep_cos = cosine_mean(
1004
+ pred,
1005
+ teacher,
1006
+ )
1007
+
1008
+ teacher_scaled = magnitude_match(
1009
+ teacher,
1010
+ h49,
1011
+ )
1012
+
1013
+ pred_scaled = magnitude_match(
1014
+ pred,
1015
+ h49,
1016
+ )
1017
+
1018
+ teacher_delta = (
1019
+ teacher_scaled
1020
+ - h49
1021
+ )
1022
+
1023
+ pred_delta = (
1024
+ pred_scaled
1025
+ - h49
1026
+ )
1027
+
1028
+ correction_cos = (
1029
+ cosine_mean(
1030
+ pred_delta,
1031
+ teacher_delta,
1032
+ )
1033
+ )
1034
+
1035
+ rep_scores.append(
1036
+ rep_cos
1037
+ )
1038
+
1039
+ correction_scores.append(
1040
+ correction_cos
1041
+ )
1042
+
1043
+ origin_scores[
1044
+ item["origin"]
1045
+ ].append(
1046
+ correction_cos
1047
+ )
1048
+
1049
+ category_scores.setdefault(
1050
+ item["category"],
1051
+ []
1052
+ ).append(
1053
+ correction_cos
1054
+ )
1055
+
1056
+
1057
+ for alpha in ALPHAS:
1058
+
1059
+ full = (
1060
+ h49
1061
+ + alpha
1062
+ * teacher_delta
1063
+ )
1064
+
1065
+ distilled = (
1066
+ h49
1067
+ + alpha
1068
+ * pred_delta
1069
+ )
1070
+
1071
+ blend_scores[
1072
+ alpha
1073
+ ].append(
1074
+ cosine_mean(
1075
+ distilled,
1076
+ full,
1077
+ )
1078
+ )
1079
+
1080
+
1081
+ result = {
1082
+ "rep_cos":
1083
+ sum(rep_scores)
1084
+ / len(rep_scores),
1085
+
1086
+ "correction_cos":
1087
+ sum(correction_scores)
1088
+ / len(correction_scores),
1089
+
1090
+ "correction_min":
1091
+ min(correction_scores),
1092
+
1093
+ "main_correction":
1094
+ sum(
1095
+ origin_scores[
1096
+ "main440"
1097
+ ]
1098
+ )
1099
+ / len(
1100
+ origin_scores[
1101
+ "main440"
1102
+ ]
1103
+ ),
1104
+
1105
+ "ood_correction":
1106
+ sum(
1107
+ origin_scores[
1108
+ "ood160"
1109
+ ]
1110
+ )
1111
+ / len(
1112
+ origin_scores[
1113
+ "ood160"
1114
+ ]
1115
+ ),
1116
+ }
1117
+
1118
+ for alpha in ALPHAS:
1119
+
1120
+ result[
1121
+ f"blend_{alpha}"
1122
+ ] = (
1123
+ sum(
1124
+ blend_scores[
1125
+ alpha
1126
+ ]
1127
+ )
1128
+ / len(
1129
+ blend_scores[
1130
+ alpha
1131
+ ]
1132
+ )
1133
+ )
1134
+
1135
+ category_result = {
1136
+ category:
1137
+ sum(values)
1138
+ / len(values)
1139
+
1140
+ for category, values
1141
+ in category_scores.items()
1142
+ }
1143
+
1144
+ return (
1145
+ result,
1146
+ category_result,
1147
+ )
1148
+
1149
+
1150
+ # ============================================================
1151
+ # TRAIN
1152
+ # ============================================================
1153
+
1154
+ token_indices = torch.arange(
1155
+ train_x.shape[0]
1156
+ )
1157
+
1158
+ best_score = -1.0
1159
+
1160
+ history = []
1161
+
1162
+
1163
+ print()
1164
+ print("=" * 100)
1165
+ print("TRAINING V3")
1166
+ print("=" * 100)
1167
+
1168
+
1169
+ for epoch in range(
1170
+ 1,
1171
+ EPOCHS + 1,
1172
+ ):
1173
+
1174
+ model.train()
1175
+
1176
+ permutation = (
1177
+ token_indices[
1178
+ torch.randperm(
1179
+ len(
1180
+ token_indices
1181
+ )
1182
+ )
1183
+ ]
1184
+ )
1185
+
1186
+ epoch_loss = 0.0
1187
+
1188
+ batches = 0
1189
+
1190
+
1191
+ for start in range(
1192
+ 0,
1193
+ len(permutation),
1194
+ BATCH_TOKENS,
1195
+ ):
1196
+
1197
+ ids = (
1198
+ permutation[
1199
+ start:
1200
+ start
1201
+ + BATCH_TOKENS
1202
+ ]
1203
+ )
1204
+
1205
+ x = (
1206
+ train_x[
1207
+ ids
1208
+ ]
1209
+ .to(
1210
+ DEVICE,
1211
+ non_blocking=True,
1212
+ )
1213
+ )
1214
+
1215
+ y = (
1216
+ train_y[
1217
+ ids
1218
+ ]
1219
+ .to(
1220
+ DEVICE,
1221
+ non_blocking=True,
1222
+ )
1223
+ )
1224
+
1225
+ # Original H3 before RMS normalize.
1226
+ # Reconstruct approximate scale reference
1227
+ # from normalized input is impossible,
1228
+ # so for correction loss we use token-space
1229
+ # target direction against a stored H3-like
1230
+ # reference built from y/pred runtime later.
1231
+ #
1232
+ # Here, use x as normalized H3 reference,
1233
+ # scaled to teacher RMS for stability.
1234
+
1235
+ x_ref = x
1236
+
1237
+ optimizer.zero_grad(
1238
+ set_to_none=True
1239
+ )
1240
+
1241
+ pred = model(
1242
+ x
1243
+ )
1244
+
1245
+ loss = loss_fn(
1246
+ pred,
1247
+ y,
1248
+ x_ref,
1249
+ )
1250
+
1251
+ loss.backward()
1252
+
1253
+ torch.nn.utils.clip_grad_norm_(
1254
+ model.parameters(),
1255
+ 1.0,
1256
+ )
1257
+
1258
+ optimizer.step()
1259
+
1260
+ epoch_loss += (
1261
+ loss.item()
1262
+ )
1263
+
1264
+ batches += 1
1265
+
1266
+
1267
+ scheduler.step()
1268
+
1269
+
1270
+ # ========================================================
1271
+ # VALIDATE
1272
+ # ========================================================
1273
+
1274
+ metrics, category_metrics = (
1275
+ evaluate()
1276
+ )
1277
+
1278
+
1279
+ # Favor OOD correction slightly.
1280
+ score = (
1281
+ 0.45
1282
+ * metrics[
1283
+ "correction_cos"
1284
+ ]
1285
+ +
1286
+ 0.35
1287
+ * metrics[
1288
+ "ood_correction"
1289
+ ]
1290
+ +
1291
+ 0.20
1292
+ * metrics[
1293
+ "rep_cos"
1294
+ ]
1295
+ )
1296
+
1297
+
1298
+ row = {
1299
+ "epoch":
1300
+ epoch,
1301
+
1302
+ "train_loss":
1303
+ epoch_loss
1304
+ / max(
1305
+ batches,
1306
+ 1,
1307
+ ),
1308
+
1309
+ "rep_cos":
1310
+ metrics[
1311
+ "rep_cos"
1312
+ ],
1313
+
1314
+ "correction_cos":
1315
+ metrics[
1316
+ "correction_cos"
1317
+ ],
1318
+
1319
+ "main_correction":
1320
+ metrics[
1321
+ "main_correction"
1322
+ ],
1323
+
1324
+ "ood_correction":
1325
+ metrics[
1326
+ "ood_correction"
1327
+ ],
1328
+
1329
+ "correction_min":
1330
+ metrics[
1331
+ "correction_min"
1332
+ ],
1333
+
1334
+ "blend_0.1":
1335
+ metrics[
1336
+ "blend_0.1"
1337
+ ],
1338
+
1339
+ "blend_0.2":
1340
+ metrics[
1341
+ "blend_0.2"
1342
+ ],
1343
+
1344
+ "blend_0.3":
1345
+ metrics[
1346
+ "blend_0.3"
1347
+ ],
1348
+
1349
+ "score":
1350
+ score,
1351
+ }
1352
+
1353
+ history.append(
1354
+ row
1355
+ )
1356
+
1357
+
1358
+ print(
1359
+ f"Epoch "
1360
+ f"{epoch:02d}/"
1361
+ f"{EPOCHS} | "
1362
+ f"loss="
1363
+ f"{row['train_loss']:.6f} | "
1364
+ f"rep="
1365
+ f"{row['rep_cos']:.6f} | "
1366
+ f"corr="
1367
+ f"{row['correction_cos']:.6f} | "
1368
+ f"main="
1369
+ f"{row['main_correction']:.6f} | "
1370
+ f"ood="
1371
+ f"{row['ood_correction']:.6f} | "
1372
+ f"min="
1373
+ f"{row['correction_min']:.6f}"
1374
+ )
1375
+
1376
+
1377
+ if score > best_score:
1378
+
1379
+ best_score = score
1380
+
1381
+ save_student(
1382
+ model,
1383
+ BEST_FILE,
1384
+ best_score,
1385
+ )
1386
+
1387
+
1388
+ # ============================================================
1389
+ # SAVE HISTORY
1390
+ # ============================================================
1391
+
1392
+ with open(
1393
+ CSV_FILE,
1394
+ "w",
1395
+ newline="",
1396
+ encoding="utf-8",
1397
+ ) as f:
1398
+
1399
+ writer = csv.DictWriter(
1400
+ f,
1401
+ fieldnames=list(
1402
+ history[0].keys()
1403
+ ),
1404
+ )
1405
+
1406
+ writer.writeheader()
1407
+
1408
+ writer.writerows(
1409
+ history
1410
+ )
1411
+
1412
+
1413
+ # ============================================================
1414
+ # LOAD BEST
1415
+ # ============================================================
1416
+
1417
+ best_weights = load_file(
1418
+ BEST_FILE,
1419
+ device="cpu",
1420
+ )
1421
+
1422
+ best_model = SemanticStudent(
1423
+ 5120,
1424
+ HIDDEN_DIM,
1425
+ 5120,
1426
+ )
1427
+
1428
+ with torch.no_grad():
1429
+
1430
+ best_model.fc1.weight.copy_(
1431
+ best_weights[
1432
+ "fc1.weight"
1433
+ ].float()
1434
+ )
1435
+
1436
+ best_model.fc1.bias.copy_(
1437
+ best_weights[
1438
+ "fc1.bias"
1439
+ ].float()
1440
+ )
1441
+
1442
+ best_model.fc2.weight.copy_(
1443
+ best_weights[
1444
+ "fc2.weight"
1445
+ ].float()
1446
+ )
1447
+
1448
+ best_model.fc2.bias.copy_(
1449
+ best_weights[
1450
+ "fc2.bias"
1451
+ ].float()
1452
+ )
1453
+
1454
+ best_model.fc3.weight.copy_(
1455
+ best_weights[
1456
+ "fc3.weight"
1457
+ ].float()
1458
+ )
1459
+
1460
+ best_model.fc3.bias.copy_(
1461
+ best_weights[
1462
+ "fc3.bias"
1463
+ ].float()
1464
+ )
1465
+
1466
+
1467
+ model = best_model.to(
1468
+ DEVICE
1469
+ )
1470
+
1471
+ model.eval()
1472
+
1473
+
1474
+ final_metrics, final_categories = (
1475
+ evaluate()
1476
+ )
1477
+
1478
+
1479
+ # ============================================================
1480
+ # REPORT
1481
+ # ============================================================
1482
+
1483
+ with open(
1484
+ REPORT_FILE,
1485
+ "w",
1486
+ encoding="utf-8",
1487
+ ) as f:
1488
+
1489
+ f.write(
1490
+ "DISTILLED STUDENT V3 FINAL REPORT\n"
1491
+ )
1492
+
1493
+ f.write(
1494
+ "=" * 100
1495
+ + "\n\n"
1496
+ )
1497
+
1498
+ f.write(
1499
+ "Training prompts: 500\n"
1500
+ )
1501
+
1502
+ f.write(
1503
+ "Validation prompts: 100\n"
1504
+ )
1505
+
1506
+ f.write(
1507
+ "Validation includes 70 original + 30 strict OOD prompts.\n\n"
1508
+ )
1509
+
1510
+ f.write(
1511
+ f"Representation cosine: "
1512
+ f"{final_metrics['rep_cos']:.6f}\n"
1513
+ )
1514
+
1515
+ f.write(
1516
+ f"Correction cosine: "
1517
+ f"{final_metrics['correction_cos']:.6f}\n"
1518
+ )
1519
+
1520
+ f.write(
1521
+ f"Main correction: "
1522
+ f"{final_metrics['main_correction']:.6f}\n"
1523
+ )
1524
+
1525
+ f.write(
1526
+ f"OOD correction: "
1527
+ f"{final_metrics['ood_correction']:.6f}\n"
1528
+ )
1529
+
1530
+ f.write(
1531
+ f"Correction minimum: "
1532
+ f"{final_metrics['correction_min']:.6f}\n\n"
1533
+ )
1534
+
1535
+ for alpha in ALPHAS:
1536
+
1537
+ f.write(
1538
+ f"Blend alpha "
1539
+ f"{alpha:.2f}: "
1540
+ f"{final_metrics[f'blend_{alpha}']:.6f}\n"
1541
+ )
1542
+
1543
+
1544
+ f.write(
1545
+ "\nCATEGORY CORRECTION\n"
1546
+ )
1547
+
1548
+ f.write(
1549
+ "-" * 100
1550
+ + "\n"
1551
+ )
1552
+
1553
+ for category in sorted(
1554
+ final_categories
1555
+ ):
1556
+
1557
+ f.write(
1558
+ f"{category:28s} "
1559
+ f"{final_categories[category]:.6f}\n"
1560
+ )
1561
+
1562
+
1563
+ f.write(
1564
+ "\nCheckpoint:\n"
1565
+ )
1566
+
1567
+ f.write(
1568
+ BEST_FILE
1569
+ + "\n"
1570
+ )
1571
+
1572
+
1573
+ # ============================================================
1574
+ # DONE
1575
+ # ============================================================
1576
+
1577
+ print()
1578
+ print("=" * 100)
1579
+ print("V3 FINAL")
1580
+ print("=" * 100)
1581
+
1582
+ print(
1583
+ "Representation cosine:",
1584
+ f"{final_metrics['rep_cos']:.6f}"
1585
+ )
1586
+
1587
+ print(
1588
+ "Correction cosine:",
1589
+ f"{final_metrics['correction_cos']:.6f}"
1590
+ )
1591
+
1592
+ print(
1593
+ "Main correction:",
1594
+ f"{final_metrics['main_correction']:.6f}"
1595
+ )
1596
+
1597
+ print(
1598
+ "OOD correction:",
1599
+ f"{final_metrics['ood_correction']:.6f}"
1600
+ )
1601
+
1602
+ print(
1603
+ "Correction minimum:",
1604
+ f"{final_metrics['correction_min']:.6f}"
1605
+ )
1606
+
1607
+ print()
1608
+ print(
1609
+ "Saved student:"
1610
+ )
1611
+
1612
+ print(
1613
+ BEST_FILE
1614
+ )
1615
+
1616
+ print()
1617
+ print(
1618
+ "Report:"
1619
+ )
1620
+
1621
+ print(
1622
+ REPORT_FILE
1623
+ )
1624
+
1625
+ print()
1626
+ print(
1627
+ "DONE"
1628
+ )
research/raw_scripts/train_embedding_bridge_test.py ADDED
@@ -0,0 +1,397 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import math
3
+ import torch
4
+
5
+ from safetensors import safe_open
6
+ from safetensors.torch import load_file, save_file
7
+
8
+
9
+ # =========================================================
10
+ # PATHS
11
+ # =========================================================
12
+
13
+ H3_EMB = r"E:\Models SD XL\UNET\Minimax_H3\h3_qwen3vl32b_embeddings_bf16.safetensors"
14
+
15
+ SENSENOVA = r"D:\ComfyUI_Python312\ComfyUI\models\unet\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors"
16
+
17
+ OUT = r"E:\Models SD XL\UNET\Minimax_H3\sensenova_to_h3_embedding_bridge_rank512.safetensors"
18
+
19
+
20
+ # =========================================================
21
+ # SETTINGS
22
+ # =========================================================
23
+
24
+ INPUT_DIM = 4096
25
+ OUTPUT_DIM = 5120
26
+ RANK = 512
27
+
28
+ BATCH_SIZE = 256
29
+ EPOCHS = 5
30
+ LR = 1e-3
31
+
32
+ VAL_SIZE = 12000
33
+
34
+ SEED = 1234
35
+
36
+
37
+ # =========================================================
38
+ # DEVICE
39
+ # =========================================================
40
+
41
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
42
+
43
+ print("=" * 80)
44
+ print("SenseNova -> MiniMax H3 embedding bridge test")
45
+ print("=" * 80)
46
+
47
+ print("Device:", device)
48
+
49
+ if device.type == "cuda":
50
+ print("GPU:", torch.cuda.get_device_name(0))
51
+ print(
52
+ "VRAM:",
53
+ round(torch.cuda.get_device_properties(0).total_memory / 1024**3, 2),
54
+ "GB"
55
+ )
56
+
57
+
58
+ # =========================================================
59
+ # LOAD H3 EMBEDDINGS
60
+ # =========================================================
61
+
62
+ print()
63
+ print("Loading H3 embeddings...")
64
+
65
+ h3_file = load_file(H3_EMB, device="cpu")
66
+
67
+ h3 = h3_file["model.embed_tokens.weight"]
68
+
69
+ print("H3 shape:", tuple(h3.shape))
70
+ print("H3 dtype:", h3.dtype)
71
+
72
+
73
+ # =========================================================
74
+ # LOAD SENSENOVA EMBEDDINGS ONLY
75
+ # =========================================================
76
+
77
+ print()
78
+ print("Loading SenseNova token embeddings...")
79
+
80
+ with safe_open(SENSENOVA, framework="pt", device="cpu") as f:
81
+
82
+ key = "language_model.model.embed_tokens.weight"
83
+
84
+ if key not in f.keys():
85
+ raise RuntimeError(
86
+ f"SenseNova embedding tensor not found: {key}"
87
+ )
88
+
89
+ sense = f.get_tensor(key)
90
+
91
+
92
+ print("SenseNova shape:", tuple(sense.shape))
93
+ print("SenseNova dtype:", sense.dtype)
94
+
95
+
96
+ # =========================================================
97
+ # VALIDATE
98
+ # =========================================================
99
+
100
+ if h3.shape[0] != sense.shape[0]:
101
+ raise RuntimeError(
102
+ f"Vocabulary mismatch: H3={h3.shape[0]}, SenseNova={sense.shape[0]}"
103
+ )
104
+
105
+ if sense.shape[1] != INPUT_DIM:
106
+ raise RuntimeError(
107
+ f"Unexpected SenseNova hidden size: {sense.shape[1]}"
108
+ )
109
+
110
+ if h3.shape[1] != OUTPUT_DIM:
111
+ raise RuntimeError(
112
+ f"Unexpected H3 hidden size: {h3.shape[1]}"
113
+ )
114
+
115
+
116
+ VOCAB = h3.shape[0]
117
+
118
+ print()
119
+ print("Vocabulary rows:", VOCAB)
120
+
121
+
122
+ # =========================================================
123
+ # NORMALIZATION
124
+ # =========================================================
125
+
126
+ # Convert storage to float32 only per batch later.
127
+ # Calculate global average norms using manageable chunks.
128
+
129
+ def average_norm(x, chunk=4096):
130
+
131
+ total = 0.0
132
+ count = 0
133
+
134
+ for start in range(0, x.shape[0], chunk):
135
+
136
+ end = min(start + chunk, x.shape[0])
137
+
138
+ b = x[start:end].float()
139
+
140
+ total += b.norm(dim=1).sum().item()
141
+ count += b.shape[0]
142
+
143
+ return total / count
144
+
145
+
146
+ print()
147
+ print("Calculating embedding norm statistics...")
148
+
149
+ sense_norm = average_norm(sense)
150
+ h3_norm = average_norm(h3)
151
+
152
+ print("SenseNova avg norm:", sense_norm)
153
+ print("H3 avg norm: ", h3_norm)
154
+
155
+
156
+ # =========================================================
157
+ # TRAIN / VALIDATION SPLIT
158
+ # =========================================================
159
+
160
+ g = torch.Generator()
161
+ g.manual_seed(SEED)
162
+
163
+ perm = torch.randperm(VOCAB, generator=g)
164
+
165
+ val_ids = perm[:VAL_SIZE]
166
+ train_ids = perm[VAL_SIZE:]
167
+
168
+ print()
169
+ print("Train rows:", len(train_ids))
170
+ print("Validation rows:", len(val_ids))
171
+
172
+
173
+ # =========================================================
174
+ # LOW-RANK PROJECTOR
175
+ #
176
+ # 4096 -> 512 -> 5120
177
+ # no activation: this is a low-rank linear mapping
178
+ # =========================================================
179
+
180
+ A = torch.nn.Linear(
181
+ INPUT_DIM,
182
+ RANK,
183
+ bias=False,
184
+ device=device,
185
+ dtype=torch.float32,
186
+ )
187
+
188
+ B = torch.nn.Linear(
189
+ RANK,
190
+ OUTPUT_DIM,
191
+ bias=False,
192
+ device=device,
193
+ dtype=torch.float32,
194
+ )
195
+
196
+
197
+ torch.nn.init.normal_(
198
+ A.weight,
199
+ std=1.0 / math.sqrt(INPUT_DIM)
200
+ )
201
+
202
+ torch.nn.init.zeros_(B.weight)
203
+
204
+
205
+ params = list(A.parameters()) + list(B.parameters())
206
+
207
+ optimizer = torch.optim.AdamW(
208
+ params,
209
+ lr=LR,
210
+ weight_decay=1e-4,
211
+ )
212
+
213
+
214
+ # =========================================================
215
+ # LOSS
216
+ # =========================================================
217
+
218
+ def bridge_forward(x):
219
+
220
+ return B(A(x))
221
+
222
+
223
+ def loss_fn(pred, target):
224
+
225
+ # Directional similarity.
226
+ cosine = 1.0 - torch.nn.functional.cosine_similarity(
227
+ pred,
228
+ target,
229
+ dim=-1
230
+ ).mean()
231
+
232
+ # Normalize magnitude sensitivity.
233
+ pred_n = torch.nn.functional.normalize(pred, dim=-1)
234
+ target_n = torch.nn.functional.normalize(target, dim=-1)
235
+
236
+ mse = torch.nn.functional.mse_loss(
237
+ pred_n,
238
+ target_n
239
+ )
240
+
241
+ return cosine + mse
242
+
243
+
244
+ # =========================================================
245
+ # EVALUATION
246
+ # =========================================================
247
+
248
+ @torch.no_grad()
249
+ def evaluate(ids):
250
+
251
+ A.eval()
252
+ B.eval()
253
+
254
+ cos_sum = 0.0
255
+ mse_sum = 0.0
256
+ count = 0
257
+
258
+ for start in range(0, len(ids), BATCH_SIZE):
259
+
260
+ idx = ids[start:start + BATCH_SIZE]
261
+
262
+ x = sense[idx].float().to(device)
263
+ y = h3[idx].float().to(device)
264
+
265
+ pred = bridge_forward(x)
266
+
267
+ cos = torch.nn.functional.cosine_similarity(
268
+ pred,
269
+ y,
270
+ dim=-1
271
+ )
272
+
273
+ mse = torch.nn.functional.mse_loss(
274
+ pred,
275
+ y,
276
+ reduction="none"
277
+ ).mean(dim=1)
278
+
279
+ cos_sum += cos.sum().item()
280
+ mse_sum += mse.sum().item()
281
+
282
+ count += len(idx)
283
+
284
+ return cos_sum / count, mse_sum / count
285
+
286
+
287
+ # =========================================================
288
+ # BASELINE
289
+ # =========================================================
290
+
291
+ print()
292
+ print("=" * 80)
293
+ print("TRAINING")
294
+ print("=" * 80)
295
+
296
+
297
+ # =========================================================
298
+ # TRAIN
299
+ # =========================================================
300
+
301
+ for epoch in range(EPOCHS):
302
+
303
+ A.train()
304
+ B.train()
305
+
306
+ shuffled = train_ids[
307
+ torch.randperm(len(train_ids), generator=g)
308
+ ]
309
+
310
+ running = 0.0
311
+ batches = 0
312
+
313
+ for start in range(0, len(shuffled), BATCH_SIZE):
314
+
315
+ idx = shuffled[start:start + BATCH_SIZE]
316
+
317
+ x = sense[idx].float().to(device)
318
+ y = h3[idx].float().to(device)
319
+
320
+ optimizer.zero_grad(set_to_none=True)
321
+
322
+ pred = bridge_forward(x)
323
+
324
+ loss = loss_fn(pred, y)
325
+
326
+ loss.backward()
327
+
328
+ torch.nn.utils.clip_grad_norm_(
329
+ params,
330
+ 1.0
331
+ )
332
+
333
+ optimizer.step()
334
+
335
+ running += loss.item()
336
+ batches += 1
337
+
338
+ if batches % 100 == 0:
339
+
340
+ print(
341
+ f"Epoch {epoch + 1}/{EPOCHS} | "
342
+ f"batch {batches} | "
343
+ f"loss {running / batches:.6f}"
344
+ )
345
+
346
+ val_cos, val_mse = evaluate(val_ids)
347
+
348
+ print()
349
+ print(
350
+ f"EPOCH {epoch + 1} COMPLETE | "
351
+ f"train_loss={running / batches:.6f} | "
352
+ f"val_cos={val_cos:.6f} | "
353
+ f"val_mse={val_mse:.8f}"
354
+ )
355
+ print()
356
+
357
+
358
+ # =========================================================
359
+ # FINAL EVALUATION
360
+ # =========================================================
361
+
362
+ val_cos, val_mse = evaluate(val_ids)
363
+
364
+ print("=" * 80)
365
+ print("FINAL")
366
+ print("=" * 80)
367
+
368
+ print("Validation cosine similarity:", val_cos)
369
+ print("Validation MSE:", val_mse)
370
+
371
+
372
+ # =========================================================
373
+ # SAVE
374
+ # =========================================================
375
+
376
+ save_file(
377
+ {
378
+ "proj_in.weight": A.weight.detach().cpu().to(torch.float16),
379
+ "proj_out.weight": B.weight.detach().cpu().to(torch.float16),
380
+ },
381
+ OUT,
382
+ metadata={
383
+ "source": "SenseNova-U1.5-8B-MoT",
384
+ "target": "MiniMax-H3-Qwen3VL32B",
385
+ "input_dim": str(INPUT_DIM),
386
+ "rank": str(RANK),
387
+ "output_dim": str(OUTPUT_DIM),
388
+ "validation_cosine": str(val_cos),
389
+ "validation_mse": str(val_mse),
390
+ }
391
+ )
392
+
393
+ print()
394
+ print("Saved bridge:")
395
+ print(OUT)
396
+ print()
397
+ print("DONE")
research/raw_scripts/train_hidden_bridge_screening.py ADDED
@@ -0,0 +1,859 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import csv
3
+ import json
4
+ import math
5
+ import random
6
+ import gc
7
+ import torch
8
+ import torch.nn.functional as F
9
+ from safetensors.torch import save_file
10
+
11
+
12
+ ROOT = r"E:\Models SD XL\UNET\Minimax_H3"
13
+
14
+ H3_FILE = os.path.join(
15
+ ROOT,
16
+ "h3_bridge_hidden_480.pt"
17
+ )
18
+
19
+ SN_FILE = os.path.join(
20
+ ROOT,
21
+ "sensenova_bridge_hidden_480.pt"
22
+ )
23
+
24
+ OUT_CSV = os.path.join(
25
+ ROOT,
26
+ "hidden_bridge_screening_rank128.csv"
27
+ )
28
+
29
+ OUT_TXT = os.path.join(
30
+ ROOT,
31
+ "hidden_bridge_screening_rank128.txt"
32
+ )
33
+
34
+ OUT_DIR = os.path.join(
35
+ ROOT,
36
+ "hidden_bridge_probes_rank128"
37
+ )
38
+
39
+ os.makedirs(OUT_DIR, exist_ok=True)
40
+
41
+
42
+ # ============================================================
43
+ # SETTINGS
44
+ # ============================================================
45
+
46
+ SN_LAYERS = [8, 16, 24, 32, 41]
47
+ H3_LAYERS = [8, 16, 24, 32, 40, 49]
48
+
49
+ RANK = 128
50
+
51
+ EPOCHS = 8
52
+ BATCH_SIZE = 512
53
+
54
+ LR = 2e-3
55
+ WEIGHT_DECAY = 1e-4
56
+
57
+ VAL_FRACTION = 0.20
58
+
59
+ SEED = 12345
60
+
61
+
62
+ # ============================================================
63
+ # DEVICE
64
+ # ============================================================
65
+
66
+ torch.manual_seed(SEED)
67
+ random.seed(SEED)
68
+
69
+ device = torch.device(
70
+ "cuda" if torch.cuda.is_available() else "cpu"
71
+ )
72
+
73
+ print("=" * 90)
74
+ print("SenseNova -> MiniMax H3")
75
+ print("HIDDEN-STATE BRIDGE SCREENING")
76
+ print("=" * 90)
77
+
78
+ print("Device:", device)
79
+
80
+ if device.type == "cuda":
81
+ print("GPU:", torch.cuda.get_device_name(0))
82
+ print(
83
+ "VRAM:",
84
+ round(
85
+ torch.cuda.get_device_properties(0).total_memory
86
+ / 1024**3,
87
+ 2
88
+ ),
89
+ "GB"
90
+ )
91
+
92
+
93
+ # ============================================================
94
+ # LOAD DATA
95
+ # ============================================================
96
+
97
+ print("\nLoading H3 dataset...")
98
+
99
+ h3_data = torch.load(
100
+ H3_FILE,
101
+ map_location="cpu",
102
+ weights_only=False
103
+ )
104
+
105
+ print("H3 prompts:", len(h3_data["results"]))
106
+
107
+ print("Loading SenseNova dataset...")
108
+
109
+ sn_data = torch.load(
110
+ SN_FILE,
111
+ map_location="cpu",
112
+ weights_only=False
113
+ )
114
+
115
+ print("SenseNova prompts:", len(sn_data["results"]))
116
+
117
+
118
+ h3_results = h3_data["results"]
119
+ sn_results = sn_data["results"]
120
+
121
+ if len(h3_results) != len(sn_results):
122
+ raise RuntimeError(
123
+ "H3 and SenseNova prompt counts differ."
124
+ )
125
+
126
+
127
+ # ============================================================
128
+ # VERIFY ALIGNMENT
129
+ # ============================================================
130
+
131
+ print("\nVerifying prompt/token alignment...")
132
+
133
+ for i, (h, s) in enumerate(
134
+ zip(h3_results, sn_results)
135
+ ):
136
+
137
+ if h["prompt"] != s["prompt"]:
138
+ raise RuntimeError(
139
+ f"Prompt mismatch at index {i}"
140
+ )
141
+
142
+ if not torch.equal(
143
+ h["input_ids"],
144
+ s["input_ids"]
145
+ ):
146
+ raise RuntimeError(
147
+ f"Token mismatch at index {i}\n"
148
+ f"{h['prompt']}"
149
+ )
150
+
151
+ print("All 480 prompts and token IDs align perfectly.")
152
+
153
+
154
+ # ============================================================
155
+ # TRAIN / VALIDATION SPLIT BY PROMPT
156
+ # ============================================================
157
+
158
+ N = len(h3_results)
159
+
160
+ indices = list(range(N))
161
+
162
+ rng = random.Random(SEED)
163
+ rng.shuffle(indices)
164
+
165
+ val_count = round(
166
+ N * VAL_FRACTION
167
+ )
168
+
169
+ val_prompt_ids = set(
170
+ indices[:val_count]
171
+ )
172
+
173
+ train_prompt_ids = [
174
+ x for x in indices
175
+ if x not in val_prompt_ids
176
+ ]
177
+
178
+ val_prompt_ids = sorted(val_prompt_ids)
179
+
180
+ print()
181
+ print("Train prompts:", len(train_prompt_ids))
182
+ print("Validation prompts:", len(val_prompt_ids))
183
+
184
+
185
+ # ============================================================
186
+ # CATEGORY DISTRIBUTION
187
+ # ============================================================
188
+
189
+ def category_counts(ids, results):
190
+
191
+ counts = {}
192
+
193
+ for i in ids:
194
+
195
+ category = results[i].get(
196
+ "category",
197
+ "unknown"
198
+ )
199
+
200
+ counts[category] = (
201
+ counts.get(category, 0) + 1
202
+ )
203
+
204
+ return counts
205
+
206
+
207
+ print("\nTrain categories:")
208
+ print(category_counts(train_prompt_ids, h3_results))
209
+
210
+ print("\nValidation categories:")
211
+ print(category_counts(val_prompt_ids, h3_results))
212
+
213
+
214
+ # ============================================================
215
+ # BUILD TOKEN DATASET FOR ONE PAIR
216
+ # ============================================================
217
+
218
+ def build_tokens(
219
+ sn_layer,
220
+ h3_layer,
221
+ prompt_ids
222
+ ):
223
+
224
+ xs = []
225
+ ys = []
226
+ cats = []
227
+
228
+ for i in prompt_ids:
229
+
230
+ sx = sn_results[i]["layers"][sn_layer]
231
+ hy = h3_results[i]["layers"][h3_layer]
232
+
233
+ if sx.shape[:2] != hy.shape[:2]:
234
+ raise RuntimeError(
235
+ f"Sequence mismatch at prompt {i}"
236
+ )
237
+
238
+ # [1,T,D] -> [T,D]
239
+ sx = sx.squeeze(0).float()
240
+ hy = hy.squeeze(0).float()
241
+
242
+ xs.append(sx)
243
+ ys.append(hy)
244
+
245
+ category = h3_results[i].get(
246
+ "category",
247
+ "unknown"
248
+ )
249
+
250
+ cats.extend(
251
+ [category] * sx.shape[0]
252
+ )
253
+
254
+ X = torch.cat(xs, dim=0)
255
+ Y = torch.cat(ys, dim=0)
256
+
257
+ return X, Y, cats
258
+
259
+
260
+ # ============================================================
261
+ # LOW-RANK PROJECTOR
262
+ # ============================================================
263
+
264
+ class LowRankBridge(torch.nn.Module):
265
+
266
+ def __init__(
267
+ self,
268
+ input_dim=4096,
269
+ rank=128,
270
+ output_dim=5120
271
+ ):
272
+
273
+ super().__init__()
274
+
275
+ self.down = torch.nn.Linear(
276
+ input_dim,
277
+ rank,
278
+ bias=False
279
+ )
280
+
281
+ self.up = torch.nn.Linear(
282
+ rank,
283
+ output_dim,
284
+ bias=False
285
+ )
286
+
287
+ torch.nn.init.normal_(
288
+ self.down.weight,
289
+ std=1.0 / math.sqrt(input_dim)
290
+ )
291
+
292
+ torch.nn.init.zeros_(
293
+ self.up.weight
294
+ )
295
+
296
+
297
+ def forward(self, x):
298
+
299
+ return self.up(
300
+ self.down(x)
301
+ )
302
+
303
+
304
+ # ============================================================
305
+ # NORMALIZATION
306
+ # ============================================================
307
+
308
+ def normalize_input(x):
309
+
310
+ # Per-token RMS normalization.
311
+ # Keeps direction/information but removes the huge
312
+ # layer-dependent magnitude differences we observed.
313
+
314
+ rms = torch.sqrt(
315
+ x.pow(2).mean(
316
+ dim=-1,
317
+ keepdim=True
318
+ ) + 1e-6
319
+ )
320
+
321
+ return x / rms
322
+
323
+
324
+ # ============================================================
325
+ # EVALUATE
326
+ # ============================================================
327
+
328
+ @torch.no_grad()
329
+ def evaluate(
330
+ model,
331
+ X,
332
+ Y,
333
+ categories
334
+ ):
335
+
336
+ model.eval()
337
+
338
+ all_cos = []
339
+ category_values = {}
340
+
341
+ for start in range(
342
+ 0,
343
+ X.shape[0],
344
+ BATCH_SIZE
345
+ ):
346
+
347
+ end = min(
348
+ start + BATCH_SIZE,
349
+ X.shape[0]
350
+ )
351
+
352
+ x = X[start:end].to(
353
+ device,
354
+ non_blocking=True
355
+ )
356
+
357
+ y = Y[start:end].to(
358
+ device,
359
+ non_blocking=True
360
+ )
361
+
362
+ x = normalize_input(x)
363
+
364
+ pred = model(x)
365
+
366
+ cos = F.cosine_similarity(
367
+ pred,
368
+ y,
369
+ dim=-1
370
+ )
371
+
372
+ cos_cpu = (
373
+ cos.detach()
374
+ .float()
375
+ .cpu()
376
+ )
377
+
378
+ all_cos.append(
379
+ cos_cpu
380
+ )
381
+
382
+ batch_categories = (
383
+ categories[start:end]
384
+ )
385
+
386
+ for c, value in zip(
387
+ batch_categories,
388
+ cos_cpu.tolist()
389
+ ):
390
+
391
+ category_values.setdefault(
392
+ c,
393
+ []
394
+ ).append(value)
395
+
396
+
397
+ values = torch.cat(
398
+ all_cos
399
+ )
400
+
401
+ category_scores = {
402
+ c: sum(v) / len(v)
403
+ for c, v in category_values.items()
404
+ }
405
+
406
+ return {
407
+ "mean": values.mean().item(),
408
+ "std": values.std().item(),
409
+ "min": values.min().item(),
410
+ "max": values.max().item(),
411
+ "categories": category_scores,
412
+ }
413
+
414
+
415
+ # ============================================================
416
+ # TRAIN ONE PAIR
417
+ # ============================================================
418
+
419
+ def train_pair(
420
+ sn_layer,
421
+ h3_layer
422
+ ):
423
+
424
+ print()
425
+ print("=" * 90)
426
+ print(
427
+ f"SN L{sn_layer:02d} "
428
+ f"-> H3 L{h3_layer:02d}"
429
+ )
430
+ print("=" * 90)
431
+
432
+ X_train, Y_train, train_cats = (
433
+ build_tokens(
434
+ sn_layer,
435
+ h3_layer,
436
+ train_prompt_ids
437
+ )
438
+ )
439
+
440
+ X_val, Y_val, val_cats = (
441
+ build_tokens(
442
+ sn_layer,
443
+ h3_layer,
444
+ val_prompt_ids
445
+ )
446
+ )
447
+
448
+ print(
449
+ "Train tokens:",
450
+ X_train.shape[0]
451
+ )
452
+
453
+ print(
454
+ "Validation tokens:",
455
+ X_val.shape[0]
456
+ )
457
+
458
+ model = LowRankBridge(
459
+ rank=RANK
460
+ ).to(device)
461
+
462
+ optimizer = torch.optim.AdamW(
463
+ model.parameters(),
464
+ lr=LR,
465
+ weight_decay=WEIGHT_DECAY
466
+ )
467
+
468
+ best_cos = -999.0
469
+ best_state = None
470
+ best_epoch = None
471
+
472
+ train_n = X_train.shape[0]
473
+
474
+ for epoch in range(
475
+ 1,
476
+ EPOCHS + 1
477
+ ):
478
+
479
+ model.train()
480
+
481
+ perm = torch.randperm(
482
+ train_n
483
+ )
484
+
485
+ loss_total = 0.0
486
+ batch_count = 0
487
+
488
+ for start in range(
489
+ 0,
490
+ train_n,
491
+ BATCH_SIZE
492
+ ):
493
+
494
+ idx = perm[
495
+ start:
496
+ start + BATCH_SIZE
497
+ ]
498
+
499
+ x = X_train[idx].to(
500
+ device,
501
+ non_blocking=True
502
+ )
503
+
504
+ y = Y_train[idx].to(
505
+ device,
506
+ non_blocking=True
507
+ )
508
+
509
+ x = normalize_input(x)
510
+
511
+ optimizer.zero_grad(
512
+ set_to_none=True
513
+ )
514
+
515
+ pred = model(x)
516
+
517
+ # Directional alignment is what matters
518
+ # for this screening.
519
+ cos = F.cosine_similarity(
520
+ pred,
521
+ y,
522
+ dim=-1
523
+ )
524
+
525
+ loss = 1.0 - cos.mean()
526
+
527
+ loss.backward()
528
+
529
+ torch.nn.utils.clip_grad_norm_(
530
+ model.parameters(),
531
+ 1.0
532
+ )
533
+
534
+ optimizer.step()
535
+
536
+ loss_total += loss.item()
537
+ batch_count += 1
538
+
539
+
540
+ val = evaluate(
541
+ model,
542
+ X_val,
543
+ Y_val,
544
+ val_cats
545
+ )
546
+
547
+ print(
548
+ f"epoch {epoch:02d}/{EPOCHS} | "
549
+ f"loss={loss_total / batch_count:.6f} | "
550
+ f"val_cos={val['mean']:.6f}"
551
+ )
552
+
553
+
554
+ if val["mean"] > best_cos:
555
+
556
+ best_cos = val["mean"]
557
+ best_epoch = epoch
558
+
559
+ best_state = {
560
+ k: v.detach()
561
+ .cpu()
562
+ .clone()
563
+ for k, v
564
+ in model.state_dict().items()
565
+ }
566
+
567
+
568
+ # Restore best validation checkpoint.
569
+ model.load_state_dict(
570
+ best_state
571
+ )
572
+
573
+ final = evaluate(
574
+ model,
575
+ X_val,
576
+ Y_val,
577
+ val_cats
578
+ )
579
+
580
+
581
+ # ========================================================
582
+ # SAVE PROBE
583
+ # ========================================================
584
+
585
+ probe_name = (
586
+ f"SN_L{sn_layer:02d}"
587
+ f"_to_H3_L{h3_layer:02d}"
588
+ f"_rank{RANK}.safetensors"
589
+ )
590
+
591
+ probe_path = os.path.join(
592
+ OUT_DIR,
593
+ probe_name
594
+ )
595
+
596
+ save_file(
597
+ {
598
+ "down.weight":
599
+ best_state[
600
+ "down.weight"
601
+ ].to(torch.float16),
602
+
603
+ "up.weight":
604
+ best_state[
605
+ "up.weight"
606
+ ].to(torch.float16),
607
+ },
608
+ probe_path,
609
+ metadata={
610
+ "sn_layer": str(sn_layer),
611
+ "h3_layer": str(h3_layer),
612
+ "rank": str(RANK),
613
+ "best_epoch": str(best_epoch),
614
+ "validation_cosine":
615
+ str(final["mean"]),
616
+ }
617
+ )
618
+
619
+
620
+ result = {
621
+ "sn_layer": sn_layer,
622
+ "h3_layer": h3_layer,
623
+ "rank": RANK,
624
+ "best_epoch": best_epoch,
625
+ "val_cosine": final["mean"],
626
+ "val_std": final["std"],
627
+ "val_min": final["min"],
628
+ "val_max": final["max"],
629
+ "probe_file": probe_name,
630
+ }
631
+
632
+ for category, score in sorted(
633
+ final["categories"].items()
634
+ ):
635
+
636
+ result[
637
+ f"cat_{category}"
638
+ ] = score
639
+
640
+
641
+ del model
642
+ del optimizer
643
+ del X_train
644
+ del Y_train
645
+ del X_val
646
+ del Y_val
647
+
648
+ gc.collect()
649
+
650
+ if device.type == "cuda":
651
+ torch.cuda.empty_cache()
652
+
653
+ return result
654
+
655
+
656
+ # ============================================================
657
+ # TRAIN ALL 30 PAIRS
658
+ # ============================================================
659
+
660
+ results = []
661
+
662
+ total_pairs = (
663
+ len(SN_LAYERS)
664
+ * len(H3_LAYERS)
665
+ )
666
+
667
+ pair_number = 0
668
+
669
+ for sn_layer in SN_LAYERS:
670
+
671
+ for h3_layer in H3_LAYERS:
672
+
673
+ pair_number += 1
674
+
675
+ print()
676
+ print(
677
+ f"PAIR {pair_number}/{total_pairs}"
678
+ )
679
+
680
+ result = train_pair(
681
+ sn_layer,
682
+ h3_layer
683
+ )
684
+
685
+ results.append(
686
+ result
687
+ )
688
+
689
+
690
+ # ============================================================
691
+ # SORT
692
+ # ============================================================
693
+
694
+ results.sort(
695
+ key=lambda x: x["val_cosine"],
696
+ reverse=True
697
+ )
698
+
699
+
700
+ # ============================================================
701
+ # SAVE CSV
702
+ # ============================================================
703
+
704
+ all_fields = set()
705
+
706
+ for r in results:
707
+ all_fields.update(r.keys())
708
+
709
+ base_fields = [
710
+ "sn_layer",
711
+ "h3_layer",
712
+ "rank",
713
+ "best_epoch",
714
+ "val_cosine",
715
+ "val_std",
716
+ "val_min",
717
+ "val_max",
718
+ "probe_file",
719
+ ]
720
+
721
+ category_fields = sorted(
722
+ f for f in all_fields
723
+ if f.startswith("cat_")
724
+ )
725
+
726
+ fields = (
727
+ base_fields
728
+ + category_fields
729
+ )
730
+
731
+ with open(
732
+ OUT_CSV,
733
+ "w",
734
+ newline="",
735
+ encoding="utf-8"
736
+ ) as f:
737
+
738
+ writer = csv.DictWriter(
739
+ f,
740
+ fieldnames=fields
741
+ )
742
+
743
+ writer.writeheader()
744
+
745
+ for r in results:
746
+ writer.writerow(r)
747
+
748
+
749
+ # ============================================================
750
+ # SAVE TXT
751
+ # ============================================================
752
+
753
+ with open(
754
+ OUT_TXT,
755
+ "w",
756
+ encoding="utf-8"
757
+ ) as f:
758
+
759
+ f.write(
760
+ "SenseNova -> MiniMax H3 hidden-state "
761
+ "bridge screening\n"
762
+ )
763
+
764
+ f.write("=" * 90 + "\n\n")
765
+
766
+ f.write(
767
+ f"Prompts: {N}\n"
768
+ f"Train prompts: {len(train_prompt_ids)}\n"
769
+ f"Validation prompts: {len(val_prompt_ids)}\n"
770
+ f"Rank: {RANK}\n"
771
+ f"Epochs: {EPOCHS}\n\n"
772
+ )
773
+
774
+ for rank_index, r in enumerate(
775
+ results,
776
+ 1
777
+ ):
778
+
779
+ f.write(
780
+ f"{rank_index:02d}. "
781
+ f"SN L{r['sn_layer']:02d} "
782
+ f"-> H3 L{r['h3_layer']:02d} | "
783
+ f"val_cos={r['val_cosine']:.6f} | "
784
+ f"epoch={r['best_epoch']}\n"
785
+ )
786
+
787
+ cats = [
788
+ f"{k[4:]}={v:.4f}"
789
+ for k, v in r.items()
790
+ if k.startswith("cat_")
791
+ ]
792
+
793
+ f.write(
794
+ " " +
795
+ ", ".join(cats) +
796
+ "\n"
797
+ )
798
+
799
+
800
+ # ============================================================
801
+ # PRINT RESULTS
802
+ # ============================================================
803
+
804
+ print()
805
+ print("=" * 90)
806
+ print("TOP 15 HIDDEN-STATE BRIDGES")
807
+ print("=" * 90)
808
+
809
+ for rank_index, r in enumerate(
810
+ results[:15],
811
+ 1
812
+ ):
813
+
814
+ print(
815
+ f"{rank_index:02d}. "
816
+ f"SN L{r['sn_layer']:02d} "
817
+ f"-> H3 L{r['h3_layer']:02d} | "
818
+ f"val_cos={r['val_cosine']:.6f} | "
819
+ f"best_epoch={r['best_epoch']}"
820
+ )
821
+
822
+
823
+ print()
824
+ print("=" * 90)
825
+ print("BEST H3 TARGET FOR EACH SENSENOVA LAYER")
826
+ print("=" * 90)
827
+
828
+ for sn_layer in SN_LAYERS:
829
+
830
+ candidates = [
831
+ x for x in results
832
+ if x["sn_layer"] == sn_layer
833
+ ]
834
+
835
+ best = max(
836
+ candidates,
837
+ key=lambda x: x["val_cosine"]
838
+ )
839
+
840
+ print(
841
+ f"SN L{sn_layer:02d} "
842
+ f"-> H3 L{best['h3_layer']:02d} | "
843
+ f"val_cos={best['val_cosine']:.6f}"
844
+ )
845
+
846
+
847
+ print()
848
+ print("=" * 90)
849
+ print("DONE")
850
+ print()
851
+ print("Report:")
852
+ print(OUT_TXT)
853
+ print()
854
+ print("CSV:")
855
+ print(OUT_CSV)
856
+ print()
857
+ print("Probes:")
858
+ print(OUT_DIR)
859
+ print("=" * 90)
research/reports/distillation_source_inspection.txt ADDED
@@ -0,0 +1,435 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ SenseNova -> MiniMax H3 Distillation Source Inspection
2
+ ====================================================================================================
3
+
4
+ ====================================================================================================
5
+ MINIMAX H3 DATA
6
+ ====================================================================================================
7
+ Path: E:\Models SD XL\UNET\Minimax_H3\h3_bridge_hidden_480.pt
8
+ Size: 323.15 MiB
9
+
10
+ TOP-LEVEL STRUCTURE
11
+ ----------------------------------------------------------------------------------------------------
12
+ dict keys=2
13
+ KEY: 'target_layers'
14
+ list len=6
15
+ ITEM [0]
16
+ int: 8
17
+ ITEM [1]
18
+ int: 16
19
+ ITEM [2]
20
+ int: 24
21
+ ITEM [3]
22
+ int: 32
23
+ ... (2 more)
24
+ KEY: 'results'
25
+ list len=440
26
+ ITEM [0]
27
+ dict keys=4
28
+ KEY: 'prompt'
29
+ ...
30
+ KEY: 'category'
31
+ ...
32
+ KEY: 'input_ids'
33
+ ...
34
+ KEY: 'layers'
35
+ ...
36
+ ITEM [1]
37
+ dict keys=4
38
+ KEY: 'prompt'
39
+ ...
40
+ KEY: 'category'
41
+ ...
42
+ KEY: 'input_ids'
43
+ ...
44
+ KEY: 'layers'
45
+ ...
46
+ ITEM [2]
47
+ dict keys=4
48
+ KEY: 'prompt'
49
+ ...
50
+ KEY: 'category'
51
+ ...
52
+ KEY: 'input_ids'
53
+ ...
54
+ KEY: 'layers'
55
+ ...
56
+ ITEM [3]
57
+ dict keys=4
58
+ KEY: 'prompt'
59
+ ...
60
+ KEY: 'category'
61
+ ...
62
+ KEY: 'input_ids'
63
+ ...
64
+ KEY: 'layers'
65
+ ...
66
+ ... (436 more)
67
+
68
+ Total tensors found recursively: 3080
69
+
70
+ MOST COMMON TENSOR SHAPES
71
+ ----------------------------------------------------------------------------------------------------
72
+ 606 x (1, 13, 5120)
73
+ 468 x (1, 14, 5120)
74
+ 294 x (1, 8, 5120)
75
+ 258 x (1, 15, 5120)
76
+ 204 x (1, 12, 5120)
77
+ 204 x (1, 9, 5120)
78
+ 156 x (1, 16, 5120)
79
+ 114 x (1, 10, 5120)
80
+ 114 x (1, 17, 5120)
81
+ 102 x (1, 11, 5120)
82
+ 101 x (1, 13)
83
+ 78 x (1, 14)
84
+ 49 x (1, 8)
85
+ 43 x (1, 15)
86
+ 42 x (1, 6, 5120)
87
+ 36 x (1, 19, 5120)
88
+ 34 x (1, 12)
89
+ 34 x (1, 9)
90
+ 30 x (1, 7, 5120)
91
+ 26 x (1, 16)
92
+
93
+ 4096-DIM CANDIDATES
94
+ ----------------------------------------------------------------------------------------------------
95
+ Count: 0
96
+
97
+ 5120-DIM CANDIDATES
98
+ ----------------------------------------------------------------------------------------------------
99
+ Count: 2640
100
+ root.results[0].layers.8 | (1, 14, 5120) | torch.float16
101
+ root.results[0].layers.16 | (1, 14, 5120) | torch.float16
102
+ root.results[0].layers.24 | (1, 14, 5120) | torch.float16
103
+ root.results[0].layers.32 | (1, 14, 5120) | torch.float16
104
+ root.results[0].layers.40 | (1, 14, 5120) | torch.float16
105
+ root.results[0].layers.49 | (1, 14, 5120) | torch.float16
106
+ root.results[1].layers.8 | (1, 14, 5120) | torch.float16
107
+ root.results[1].layers.16 | (1, 14, 5120) | torch.float16
108
+ root.results[1].layers.24 | (1, 14, 5120) | torch.float16
109
+ root.results[1].layers.32 | (1, 14, 5120) | torch.float16
110
+ root.results[1].layers.40 | (1, 14, 5120) | torch.float16
111
+ root.results[1].layers.49 | (1, 14, 5120) | torch.float16
112
+ root.results[2].layers.8 | (1, 8, 5120) | torch.float16
113
+ root.results[2].layers.16 | (1, 8, 5120) | torch.float16
114
+ root.results[2].layers.24 | (1, 8, 5120) | torch.float16
115
+ root.results[2].layers.32 | (1, 8, 5120) | torch.float16
116
+ root.results[2].layers.40 | (1, 8, 5120) | torch.float16
117
+ root.results[2].layers.49 | (1, 8, 5120) | torch.float16
118
+ root.results[3].layers.8 | (1, 13, 5120) | torch.float16
119
+ root.results[3].layers.16 | (1, 13, 5120) | torch.float16
120
+ root.results[3].layers.24 | (1, 13, 5120) | torch.float16
121
+ root.results[3].layers.32 | (1, 13, 5120) | torch.float16
122
+ root.results[3].layers.40 | (1, 13, 5120) | torch.float16
123
+ root.results[3].layers.49 | (1, 13, 5120) | torch.float16
124
+ root.results[4].layers.8 | (1, 8, 5120) | torch.float16
125
+ root.results[4].layers.16 | (1, 8, 5120) | torch.float16
126
+ root.results[4].layers.24 | (1, 8, 5120) | torch.float16
127
+ root.results[4].layers.32 | (1, 8, 5120) | torch.float16
128
+ root.results[4].layers.40 | (1, 8, 5120) | torch.float16
129
+ root.results[4].layers.49 | (1, 8, 5120) | torch.float16
130
+ ... 2610 more
131
+
132
+ PATHS CONTAINING LAYER 32 / L32
133
+ ----------------------------------------------------------------------------------------------------
134
+ Count: 530
135
+ root.results[0].layers.32 | (1, 14, 5120)
136
+ root.results[1].layers.32 | (1, 14, 5120)
137
+ root.results[2].layers.32 | (1, 8, 5120)
138
+ root.results[3].layers.32 | (1, 13, 5120)
139
+ root.results[4].layers.32 | (1, 8, 5120)
140
+ root.results[5].layers.32 | (1, 8, 5120)
141
+ root.results[6].layers.32 | (1, 14, 5120)
142
+ root.results[7].layers.32 | (1, 8, 5120)
143
+ root.results[8].layers.32 | (1, 6, 5120)
144
+ root.results[9].layers.32 | (1, 13, 5120)
145
+ root.results[10].layers.32 | (1, 8, 5120)
146
+ root.results[11].layers.32 | (1, 8, 5120)
147
+ root.results[12].layers.32 | (1, 15, 5120)
148
+ root.results[13].layers.32 | (1, 12, 5120)
149
+ root.results[14].layers.32 | (1, 12, 5120)
150
+ root.results[15].layers.32 | (1, 10, 5120)
151
+ root.results[16].layers.32 | (1, 17, 5120)
152
+ root.results[17].layers.32 | (1, 14, 5120)
153
+ root.results[18].layers.32 | (1, 14, 5120)
154
+ root.results[19].layers.32 | (1, 8, 5120)
155
+ root.results[20].layers.32 | (1, 15, 5120)
156
+ root.results[21].layers.32 | (1, 15, 5120)
157
+ root.results[22].layers.32 | (1, 17, 5120)
158
+ root.results[23].layers.32 | (1, 9, 5120)
159
+ root.results[24].layers.32 | (1, 13, 5120)
160
+ root.results[25].layers.32 | (1, 14, 5120)
161
+ root.results[26].layers.32 | (1, 8, 5120)
162
+ root.results[27].layers.32 | (1, 12, 5120)
163
+ root.results[28].layers.32 | (1, 12, 5120)
164
+ root.results[29].layers.32 | (1, 16, 5120)
165
+ root.results[30].layers.32 | (1, 13, 5120)
166
+ root.results[31].layers.32 | (1, 15, 5120)
167
+ root.results[32].input_ids | (1, 10)
168
+ root.results[32].layers.8 | (1, 10, 5120)
169
+ root.results[32].layers.16 | (1, 10, 5120)
170
+ root.results[32].layers.24 | (1, 10, 5120)
171
+ root.results[32].layers.32 | (1, 10, 5120)
172
+ root.results[32].layers.40 | (1, 10, 5120)
173
+ root.results[32].layers.49 | (1, 10, 5120)
174
+ root.results[33].layers.32 | (1, 13, 5120)
175
+
176
+ PATHS CONTAINING LAYER 49 / L49
177
+ ----------------------------------------------------------------------------------------------------
178
+ Count: 464
179
+ root.results[0].layers.49 | (1, 14, 5120)
180
+ root.results[1].layers.49 | (1, 14, 5120)
181
+ root.results[2].layers.49 | (1, 8, 5120)
182
+ root.results[3].layers.49 | (1, 13, 5120)
183
+ root.results[4].layers.49 | (1, 8, 5120)
184
+ root.results[5].layers.49 | (1, 8, 5120)
185
+ root.results[6].layers.49 | (1, 14, 5120)
186
+ root.results[7].layers.49 | (1, 8, 5120)
187
+ root.results[8].layers.49 | (1, 6, 5120)
188
+ root.results[9].layers.49 | (1, 13, 5120)
189
+ root.results[10].layers.49 | (1, 8, 5120)
190
+ root.results[11].layers.49 | (1, 8, 5120)
191
+ root.results[12].layers.49 | (1, 15, 5120)
192
+ root.results[13].layers.49 | (1, 12, 5120)
193
+ root.results[14].layers.49 | (1, 12, 5120)
194
+ root.results[15].layers.49 | (1, 10, 5120)
195
+ root.results[16].layers.49 | (1, 17, 5120)
196
+ root.results[17].layers.49 | (1, 14, 5120)
197
+ root.results[18].layers.49 | (1, 14, 5120)
198
+ root.results[19].layers.49 | (1, 8, 5120)
199
+ root.results[20].layers.49 | (1, 15, 5120)
200
+ root.results[21].layers.49 | (1, 15, 5120)
201
+ root.results[22].layers.49 | (1, 17, 5120)
202
+ root.results[23].layers.49 | (1, 9, 5120)
203
+ root.results[24].layers.49 | (1, 13, 5120)
204
+ root.results[25].layers.49 | (1, 14, 5120)
205
+ root.results[26].layers.49 | (1, 8, 5120)
206
+ root.results[27].layers.49 | (1, 12, 5120)
207
+ root.results[28].layers.49 | (1, 12, 5120)
208
+ root.results[29].layers.49 | (1, 16, 5120)
209
+ root.results[30].layers.49 | (1, 13, 5120)
210
+ root.results[31].layers.49 | (1, 15, 5120)
211
+ root.results[32].layers.49 | (1, 10, 5120)
212
+ root.results[33].layers.49 | (1, 13, 5120)
213
+ root.results[34].layers.49 | (1, 16, 5120)
214
+ root.results[35].layers.49 | (1, 13, 5120)
215
+ root.results[36].layers.49 | (1, 17, 5120)
216
+ root.results[37].layers.49 | (1, 16, 5120)
217
+ root.results[38].layers.49 | (1, 14, 5120)
218
+ root.results[39].layers.49 | (1, 13, 5120)
219
+
220
+ ====================================================================================================
221
+ SENSENOVA DATA
222
+ ====================================================================================================
223
+ Path: E:\Models SD XL\UNET\Minimax_H3\sensenova_bridge_hidden_480.pt
224
+ Size: 215.64 MiB
225
+
226
+ TOP-LEVEL STRUCTURE
227
+ ----------------------------------------------------------------------------------------------------
228
+ dict keys=3
229
+ KEY: 'target_layers'
230
+ list len=5
231
+ ITEM [0]
232
+ int: 8
233
+ ITEM [1]
234
+ int: 16
235
+ ITEM [2]
236
+ int: 24
237
+ ITEM [3]
238
+ int: 32
239
+ ... (1 more)
240
+ KEY: 'results'
241
+ list len=440
242
+ ITEM [0]
243
+ dict keys=4
244
+ KEY: 'prompt'
245
+ ...
246
+ KEY: 'category'
247
+ ...
248
+ KEY: 'input_ids'
249
+ ...
250
+ KEY: 'layers'
251
+ ...
252
+ ITEM [1]
253
+ dict keys=4
254
+ KEY: 'prompt'
255
+ ...
256
+ KEY: 'category'
257
+ ...
258
+ KEY: 'input_ids'
259
+ ...
260
+ KEY: 'layers'
261
+ ...
262
+ ITEM [2]
263
+ dict keys=4
264
+ KEY: 'prompt'
265
+ ...
266
+ KEY: 'category'
267
+ ...
268
+ KEY: 'input_ids'
269
+ ...
270
+ KEY: 'layers'
271
+ ...
272
+ ITEM [3]
273
+ dict keys=4
274
+ KEY: 'prompt'
275
+ ...
276
+ KEY: 'category'
277
+ ...
278
+ KEY: 'input_ids'
279
+ ...
280
+ KEY: 'layers'
281
+ ...
282
+ ... (436 more)
283
+ KEY: 'source_model'
284
+ str: 'D:\\ComfyUI_Python312\\ComfyUI\\models\\unet\\SenseNova-U1.5-8B-MoT-pruned-bf16.safetensors'
285
+
286
+ Total tensors found recursively: 2640
287
+
288
+ MOST COMMON TENSOR SHAPES
289
+ ----------------------------------------------------------------------------------------------------
290
+ 505 x (1, 13, 4096)
291
+ 390 x (1, 14, 4096)
292
+ 245 x (1, 8, 4096)
293
+ 215 x (1, 15, 4096)
294
+ 170 x (1, 12, 4096)
295
+ 170 x (1, 9, 4096)
296
+ 130 x (1, 16, 4096)
297
+ 101 x (1, 13)
298
+ 95 x (1, 10, 4096)
299
+ 95 x (1, 17, 4096)
300
+ 85 x (1, 11, 4096)
301
+ 78 x (1, 14)
302
+ 49 x (1, 8)
303
+ 43 x (1, 15)
304
+ 35 x (1, 6, 4096)
305
+ 34 x (1, 12)
306
+ 34 x (1, 9)
307
+ 30 x (1, 19, 4096)
308
+ 26 x (1, 16)
309
+ 25 x (1, 7, 4096)
310
+
311
+ 4096-DIM CANDIDATES
312
+ ----------------------------------------------------------------------------------------------------
313
+ Count: 2200
314
+ root.results[0].layers.8 | (1, 14, 4096) | torch.float16
315
+ root.results[0].layers.16 | (1, 14, 4096) | torch.float16
316
+ root.results[0].layers.24 | (1, 14, 4096) | torch.float16
317
+ root.results[0].layers.32 | (1, 14, 4096) | torch.float16
318
+ root.results[0].layers.41 | (1, 14, 4096) | torch.float16
319
+ root.results[1].layers.8 | (1, 14, 4096) | torch.float16
320
+ root.results[1].layers.16 | (1, 14, 4096) | torch.float16
321
+ root.results[1].layers.24 | (1, 14, 4096) | torch.float16
322
+ root.results[1].layers.32 | (1, 14, 4096) | torch.float16
323
+ root.results[1].layers.41 | (1, 14, 4096) | torch.float16
324
+ root.results[2].layers.8 | (1, 8, 4096) | torch.float16
325
+ root.results[2].layers.16 | (1, 8, 4096) | torch.float16
326
+ root.results[2].layers.24 | (1, 8, 4096) | torch.float16
327
+ root.results[2].layers.32 | (1, 8, 4096) | torch.float16
328
+ root.results[2].layers.41 | (1, 8, 4096) | torch.float16
329
+ root.results[3].layers.8 | (1, 13, 4096) | torch.float16
330
+ root.results[3].layers.16 | (1, 13, 4096) | torch.float16
331
+ root.results[3].layers.24 | (1, 13, 4096) | torch.float16
332
+ root.results[3].layers.32 | (1, 13, 4096) | torch.float16
333
+ root.results[3].layers.41 | (1, 13, 4096) | torch.float16
334
+ root.results[4].layers.8 | (1, 8, 4096) | torch.float16
335
+ root.results[4].layers.16 | (1, 8, 4096) | torch.float16
336
+ root.results[4].layers.24 | (1, 8, 4096) | torch.float16
337
+ root.results[4].layers.32 | (1, 8, 4096) | torch.float16
338
+ root.results[4].layers.41 | (1, 8, 4096) | torch.float16
339
+ root.results[5].layers.8 | (1, 8, 4096) | torch.float16
340
+ root.results[5].layers.16 | (1, 8, 4096) | torch.float16
341
+ root.results[5].layers.24 | (1, 8, 4096) | torch.float16
342
+ root.results[5].layers.32 | (1, 8, 4096) | torch.float16
343
+ root.results[5].layers.41 | (1, 8, 4096) | torch.float16
344
+ ... 2170 more
345
+
346
+ 5120-DIM CANDIDATES
347
+ ----------------------------------------------------------------------------------------------------
348
+ Count: 0
349
+
350
+ PATHS CONTAINING LAYER 32 / L32
351
+ ----------------------------------------------------------------------------------------------------
352
+ Count: 515
353
+ root.results[0].layers.32 | (1, 14, 4096)
354
+ root.results[1].layers.32 | (1, 14, 4096)
355
+ root.results[2].layers.32 | (1, 8, 4096)
356
+ root.results[3].layers.32 | (1, 13, 4096)
357
+ root.results[4].layers.32 | (1, 8, 4096)
358
+ root.results[5].layers.32 | (1, 8, 4096)
359
+ root.results[6].layers.32 | (1, 14, 4096)
360
+ root.results[7].layers.32 | (1, 8, 4096)
361
+ root.results[8].layers.32 | (1, 6, 4096)
362
+ root.results[9].layers.32 | (1, 13, 4096)
363
+ root.results[10].layers.32 | (1, 8, 4096)
364
+ root.results[11].layers.32 | (1, 8, 4096)
365
+ root.results[12].layers.32 | (1, 15, 4096)
366
+ root.results[13].layers.32 | (1, 12, 4096)
367
+ root.results[14].layers.32 | (1, 12, 4096)
368
+ root.results[15].layers.32 | (1, 10, 4096)
369
+ root.results[16].layers.32 | (1, 17, 4096)
370
+ root.results[17].layers.32 | (1, 14, 4096)
371
+ root.results[18].layers.32 | (1, 14, 4096)
372
+ root.results[19].layers.32 | (1, 8, 4096)
373
+ root.results[20].layers.32 | (1, 15, 4096)
374
+ root.results[21].layers.32 | (1, 15, 4096)
375
+ root.results[22].layers.32 | (1, 17, 4096)
376
+ root.results[23].layers.32 | (1, 9, 4096)
377
+ root.results[24].layers.32 | (1, 13, 4096)
378
+ root.results[25].layers.32 | (1, 14, 4096)
379
+ root.results[26].layers.32 | (1, 8, 4096)
380
+ root.results[27].layers.32 | (1, 12, 4096)
381
+ root.results[28].layers.32 | (1, 12, 4096)
382
+ root.results[29].layers.32 | (1, 16, 4096)
383
+ root.results[30].layers.32 | (1, 13, 4096)
384
+ root.results[31].layers.32 | (1, 15, 4096)
385
+ root.results[32].input_ids | (1, 10)
386
+ root.results[32].layers.8 | (1, 10, 4096)
387
+ root.results[32].layers.16 | (1, 10, 4096)
388
+ root.results[32].layers.24 | (1, 10, 4096)
389
+ root.results[32].layers.32 | (1, 10, 4096)
390
+ root.results[32].layers.41 | (1, 10, 4096)
391
+ root.results[33].layers.32 | (1, 13, 4096)
392
+ root.results[34].layers.32 | (1, 16, 4096)
393
+
394
+ PATHS CONTAINING LAYER 49 / L49
395
+ ----------------------------------------------------------------------------------------------------
396
+ Count: 24
397
+ root.results[49].input_ids | (1, 14)
398
+ root.results[49].layers.8 | (1, 14, 4096)
399
+ root.results[49].layers.16 | (1, 14, 4096)
400
+ root.results[49].layers.24 | (1, 14, 4096)
401
+ root.results[49].layers.32 | (1, 14, 4096)
402
+ root.results[49].layers.41 | (1, 14, 4096)
403
+ root.results[149].input_ids | (1, 14)
404
+ root.results[149].layers.8 | (1, 14, 4096)
405
+ root.results[149].layers.16 | (1, 14, 4096)
406
+ root.results[149].layers.24 | (1, 14, 4096)
407
+ root.results[149].layers.32 | (1, 14, 4096)
408
+ root.results[149].layers.41 | (1, 14, 4096)
409
+ root.results[249].input_ids | (1, 9)
410
+ root.results[249].layers.8 | (1, 9, 4096)
411
+ root.results[249].layers.16 | (1, 9, 4096)
412
+ root.results[249].layers.24 | (1, 9, 4096)
413
+ root.results[249].layers.32 | (1, 9, 4096)
414
+ root.results[249].layers.41 | (1, 9, 4096)
415
+ root.results[349].input_ids | (1, 13)
416
+ root.results[349].layers.8 | (1, 13, 4096)
417
+ root.results[349].layers.16 | (1, 13, 4096)
418
+ root.results[349].layers.24 | (1, 13, 4096)
419
+ root.results[349].layers.32 | (1, 13, 4096)
420
+ root.results[349].layers.41 | (1, 13, 4096)
421
+
422
+ ====================================================================================================
423
+ BRIDGE FILE
424
+ ====================================================================================================
425
+ Path: D:\ComfyUI_Krea2\ComfyUI\models\bridge\SN_L32_to_H3_L49_rank128.safetensors
426
+ down.weight (128, 4096) torch.float16
427
+ up.weight (5120, 128) torch.float16
428
+
429
+ ====================================================================================================
430
+ FINAL SUMMARY
431
+ ====================================================================================================
432
+ H3 5120-dim tensors found: 2640
433
+ SenseNova 4096-dim tensors found: 2200
434
+
435
+ The next script will use this report to identify the exact H3 input representation and SenseNova teacher target without guessing the .pt layout.
research/reports/distilled_adapter_screening.csv ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ candidate,source,input_dim,rank,best_val_token_cos,val_prompt_cos,val_prompt_min,val_prompt_max,model_file
2
+ H3_L40_to_teacher_delta,40,5120,128,0.5077614188194275,0.5062089238315821,0.3640197515487671,0.5760553479194641,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L40_to_teacher_delta_rank128.safetensors
3
+ H3_L49_to_teacher_delta,49,5120,128,0.48997873067855835,0.49219752959907054,0.38705527782440186,0.625289797782898,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L49_to_teacher_delta_rank128.safetensors
4
+ H3_L16_to_teacher_delta,16,5120,128,0.4790988266468048,0.479762102663517,0.3508469760417938,0.5512891411781311,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L16_to_teacher_delta_rank128.safetensors
5
+ H3_L32_to_teacher_delta,32,5120,128,0.46586617827415466,0.46708175875246527,0.34656330943107605,0.5676354169845581,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L32_to_teacher_delta_rank128.safetensors
6
+ H3_L24_to_teacher_delta,24,5120,128,0.45875978469848633,0.4611636396497488,0.35557711124420166,0.5570473074913025,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L24_to_teacher_delta_rank128.safetensors
7
+ H3_L8_to_teacher_delta,8,5120,128,0.4529753029346466,0.45625976473093033,0.35425642132759094,0.5591447949409485,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L8_to_teacher_delta_rank128.safetensors
8
+ H3_L32_L40_L49_to_teacher_delta,"[32, 40, 49]",15360,128,0.3979092836380005,0.40544212721288203,0.3122461140155792,0.5592644810676575,E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L32_L40_L49_to_teacher_delta_rank128.safetensors
research/reports/distilled_adapter_screening.txt ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ DISTILLED ADAPTER SCREENING
2
+ ====================================================================================================
3
+
4
+ 01. H3_L40_to_teacher_delta
5
+ source: 40
6
+ input_dim: 5120
7
+ rank: 128
8
+ token_cos: 0.507761
9
+ prompt_cos: 0.506209
10
+ prompt_min: 0.364020
11
+ prompt_max: 0.576055
12
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L40_to_teacher_delta_rank128.safetensors
13
+
14
+ 02. H3_L49_to_teacher_delta
15
+ source: 49
16
+ input_dim: 5120
17
+ rank: 128
18
+ token_cos: 0.489979
19
+ prompt_cos: 0.492198
20
+ prompt_min: 0.387055
21
+ prompt_max: 0.625290
22
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L49_to_teacher_delta_rank128.safetensors
23
+
24
+ 03. H3_L16_to_teacher_delta
25
+ source: 16
26
+ input_dim: 5120
27
+ rank: 128
28
+ token_cos: 0.479099
29
+ prompt_cos: 0.479762
30
+ prompt_min: 0.350847
31
+ prompt_max: 0.551289
32
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L16_to_teacher_delta_rank128.safetensors
33
+
34
+ 04. H3_L32_to_teacher_delta
35
+ source: 32
36
+ input_dim: 5120
37
+ rank: 128
38
+ token_cos: 0.465866
39
+ prompt_cos: 0.467082
40
+ prompt_min: 0.346563
41
+ prompt_max: 0.567635
42
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L32_to_teacher_delta_rank128.safetensors
43
+
44
+ 05. H3_L24_to_teacher_delta
45
+ source: 24
46
+ input_dim: 5120
47
+ rank: 128
48
+ token_cos: 0.458760
49
+ prompt_cos: 0.461164
50
+ prompt_min: 0.355577
51
+ prompt_max: 0.557047
52
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L24_to_teacher_delta_rank128.safetensors
53
+
54
+ 06. H3_L8_to_teacher_delta
55
+ source: 8
56
+ input_dim: 5120
57
+ rank: 128
58
+ token_cos: 0.452975
59
+ prompt_cos: 0.456260
60
+ prompt_min: 0.354256
61
+ prompt_max: 0.559145
62
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L8_to_teacher_delta_rank128.safetensors
63
+
64
+ 07. H3_L32_L40_L49_to_teacher_delta
65
+ source: [32, 40, 49]
66
+ input_dim: 15360
67
+ rank: 128
68
+ token_cos: 0.397909
69
+ prompt_cos: 0.405442
70
+ prompt_min: 0.312246
71
+ prompt_max: 0.559264
72
+ model: E:\Models SD XL\UNET\Minimax_H3\distilled_adapter_screening\H3_L32_L40_L49_to_teacher_delta_rank128.safetensors
73
+