Initial release of MiniMax H3 Semantic Bridge v1.0
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +4 -0
- LICENSE.md +27 -0
- MiniMaxH3_SemanticBridge_v1.safetensors +3 -0
- MiniMax_H3_Semantic_Bridge_v1.0.zip +3 -0
- NOTICE.txt +5 -0
- README.md +435 -0
- RESEARCH_ARTICLE.md +1017 -0
- SHA256SUMS.txt +65 -0
- UPSTREAM_LICENSES.md +23 -0
- examples/01_rooftop_train_chase/native_h3.mp4 +3 -0
- examples/01_rooftop_train_chase/prompt.txt +10 -0
- examples/01_rooftop_train_chase/semantic_bridge_alpha_0.15.mp4 +3 -0
- examples/02_glass_table_prompt_adherence/native_h3.mp4 +3 -0
- examples/02_glass_table_prompt_adherence/prompt.txt +13 -0
- examples/02_glass_table_prompt_adherence/semantic_bridge_alpha_0.15.mp4 +3 -0
- research/datasets/bridge_ood_prompts_160.json +642 -0
- research/datasets/bridge_prompts_480.json +1762 -0
- research/raw_scripts/compare_hidden_layers_bridge.py +439 -0
- research/raw_scripts/compare_minimax_sensenova.py +335 -0
- research/raw_scripts/compare_tokenizers.py +145 -0
- research/raw_scripts/evaluate_distilled_student_v2_ood.py +1343 -0
- research/raw_scripts/evaluate_hidden_bridge_ood.py +700 -0
- research/raw_scripts/extract_h3_bridge_dataset.py +272 -0
- research/raw_scripts/extract_h3_embeddings_full.py +188 -0
- research/raw_scripts/extract_h3_embeddings_test.py +139 -0
- research/raw_scripts/extract_h3_hidden_states.py +189 -0
- research/raw_scripts/extract_h3_hidden_states_fast.py +266 -0
- research/raw_scripts/extract_h3_ood_hidden.py +318 -0
- research/raw_scripts/extract_sensenova_bridge_dataset.py +425 -0
- research/raw_scripts/extract_sensenova_hidden_states.py +321 -0
- research/raw_scripts/extract_sensenova_hidden_states_local.py +421 -0
- research/raw_scripts/extract_sensenova_hidden_states_stream.py +352 -0
- research/raw_scripts/extract_sensenova_hidden_states_stream_v2.py +586 -0
- research/raw_scripts/extract_sensenova_hidden_states_stream_v3.py +546 -0
- research/raw_scripts/extract_sensenova_ood_hidden.py +633 -0
- research/raw_scripts/inspect_distillation_data.py +466 -0
- research/raw_scripts/inspect_h3_sensenova_bridge.py +274 -0
- research/raw_scripts/make_bridge_ood_prompts.py +391 -0
- research/raw_scripts/make_bridge_ood_prompts_v1.py +272 -0
- research/raw_scripts/make_bridge_prompts.py +255 -0
- research/raw_scripts/prepare_sensenova_h3_condition.py +383 -0
- research/raw_scripts/test_sensenova_mot_gen_bridge.py +698 -0
- research/raw_scripts/train_distilled_adapter_screening.py +1311 -0
- research/raw_scripts/train_distilled_student_v2.py +1802 -0
- research/raw_scripts/train_distilled_student_v3.py +1628 -0
- research/raw_scripts/train_embedding_bridge_test.py +397 -0
- research/raw_scripts/train_hidden_bridge_screening.py +859 -0
- research/reports/distillation_source_inspection.txt +435 -0
- research/reports/distilled_adapter_screening.csv +8 -0
- 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 |
+
|