hanjian.thu123 commited on
Commit
17a8581
·
0 Parent(s):

[update] app.py

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitignore +30 -0
  2. LICENSE +21 -0
  3. README.md +389 -0
  4. app.py +132 -0
  5. c2i_train_infer.py +378 -0
  6. environment.yaml +19 -0
  7. evaluation/gen_eval/_base_/datasets/coco_panoptic.py +59 -0
  8. evaluation/gen_eval/_base_/default_runtime.py +27 -0
  9. evaluation/gen_eval/evaluate_images.py +298 -0
  10. evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py +253 -0
  11. evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py +79 -0
  12. evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py +37 -0
  13. evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py +61 -0
  14. evaluation/gen_eval/prompts/create_prompts.py +183 -0
  15. evaluation/gen_eval/summary_scores.py +45 -0
  16. grn/__init__.py +0 -0
  17. grn/dataset/build.py +150 -0
  18. grn/dataset/dataset_joint_vi.py +687 -0
  19. grn/models/basic.py +256 -0
  20. grn/models/ema.py +23 -0
  21. grn/models/flex_attn_mask.py +67 -0
  22. grn/models/fused_op.py +27 -0
  23. grn/models/grn.py +754 -0
  24. grn/models/grn_c2i.py +399 -0
  25. grn/models/hbq_tokenizer.py +932 -0
  26. grn/models/init_param.py +33 -0
  27. grn/models/rope.py +191 -0
  28. grn/models/umt5/fsdp.py +42 -0
  29. grn/models/umt5/t5.py +514 -0
  30. grn/models/umt5/umt5_tokenizers.py +81 -0
  31. grn/schedules/__init__.py +6 -0
  32. grn/schedules/dynamic_resolution.py +99 -0
  33. grn/schedules/global_refine.py +220 -0
  34. grn/tokenizer/.gitignore +25 -0
  35. grn/tokenizer/sample.py +694 -0
  36. grn/tokenizer/train.py +561 -0
  37. grn/tokenizer/videovae/__init__.py +0 -0
  38. grn/tokenizer/videovae/evaluation/__init__.py +8 -0
  39. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore +1 -0
  40. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py +85 -0
  41. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py +72 -0
  42. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py +83 -0
  43. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py +138 -0
  44. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py +90 -0
  45. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py +137 -0
  46. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py +322 -0
  47. grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py +27 -0
  48. grn/tokenizer/videovae/evaluation/fid.py +62 -0
  49. grn/tokenizer/videovae/evaluation/fvd.py +150 -0
  50. grn/tokenizer/videovae/evaluation/inception.py +370 -0
.gitignore ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ checkpoints
2
+ __pycache__
3
+ imagenet_tiny
4
+ *.log
5
+ *.ckpt
6
+ data
7
+ wandb
8
+ .vscode
9
+ heun-steps*
10
+ log.txt
11
+ imagenet_openai_ref_images
12
+ *.jpg
13
+ *.npz
14
+ src
15
+ scan
16
+ *.txt
17
+ *.env
18
+ *.pem,
19
+ secrets.yaml
20
+ scripts
21
+ scripts/*/local*.sh
22
+ checkpoints_vision
23
+ tmp_videos
24
+ weights
25
+ .DS_Store
26
+ local
27
+ *.mp4
28
+ tmp
29
+ .git_bk
30
+ demo
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2026 MGenAI
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
README.md ADDED
@@ -0,0 +1,389 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # GRN: Generative Refinement Networks
2
+
3
+ [![arXiv](https://img.shields.io/badge/arXiv%20paper-2604.13030-b31b1b.svg)](https://arxiv.org/abs/2604.13030)
4
+ [![Homepage](https://img.shields.io/badge/🏠%20Homepage-GRN-green.svg)](https://bytedance.github.io/GRN/)
5
+ [![Models](https://img.shields.io/badge/🤗%20Hugging%20Face-Models-blue.svg)](https://huggingface.co/bytedance-research/GRN)
6
+ [![Demo](https://img.shields.io/badge/🤗%20Hugging%20Face-Demo-yellow.svg)](https://huggingface.co/spaces/hanjian/GRN)
7
+ [![License](https://img.shields.io/badge/License-MIT-blue.svg)](LICENSE)
8
+ [![GitHub stars](https://img.shields.io/github/stars/bytedance/GRN?style=social)](https://github.com/bytedance/GRN)
9
+
10
+ ---
11
+
12
+ ## 🔥 Updates!!
13
+ * June 3, 2026: 🍉 A toy image-video dataset is provided for GRN-T2I/GRN-T2V training and fine-tuning.
14
+ * May 23, 2026: 🌺 We release the training and evaluation code for HBQ tokenizer, enjoy~
15
+ * April 14, 2026: 🤗 Paper and code release
16
+
17
+ ## 📋 Table of Contents
18
+
19
+ - [🌟 Introduction](#-introduction)
20
+ - [✨ Gallery](#-gallery)
21
+ - [🚀 Demo](#-demo)
22
+ - [📦 Model Zoo](#-model-zoo)
23
+ - [🛠️ Installation](#️-installation)
24
+ - [🖼️ Class-to-Image](#️-class-to-image)
25
+ - [Dataset](#dataset)
26
+ - [Training](#training)
27
+ - [Evaluation](#evaluation)
28
+ - [🎨 Text-to-Image](#-text-to-image)
29
+ - [Data](#data)
30
+ - [Train](#train)
31
+ - [Inference](#inference)
32
+ - [🎬 Text-to-Video](#-text-to-video)
33
+ - [Data](#data-1)
34
+ - [Train](#train-1)
35
+ - [Inference](#inference-1)
36
+ - [📦 HBQ Tokenizer](#-hbq-tokenizer)
37
+ - [Data](#data-2)
38
+ - [Training](#training-1)
39
+ - [Evaluation](#evaluation-1)
40
+ - [📧 Contact](#-contact)
41
+ - [🤗 Acknowledgements](#-acknowledgements)
42
+ - [📝 Citation](#-citation)
43
+
44
+ ---
45
+
46
+ ## 🌟 Introduction
47
+
48
+ This is the official implementation of the paper **Generative Refinement Networks for Visual Synthesis**. Neither diffusion nor autoregressive — GRN is a third way. 🧠 Refines globally like an artist. ⚡ Generates adaptively by complexity. 🏆 New SOTA across image & video. The visual generation paradigm just got rewritten.
49
+
50
+ Diffusion models dominate visual generation but they allocate uniform computational effort to samples with varying levels of complexity. Autoregressive (AR) models are complexity-aware, as evidenced by their variable likelihoods, but suffer from lossy tokenization and error accumulation.
51
+
52
+ We introduce **Generative Refinement Networks (GRN)**, a new visual synthesis paradigm that addresses these issues:
53
+ - **Near-lossless tokenization** via Hierarchical Binary Quantization (HBQ)
54
+ - **Global refinement mechanism** that progressively perfects outputs like a human artist
55
+ - **Entropy-guided sampling** for complexity-aware, adaptive-step generation
56
+
57
+ GRN achieves state-of-the-art results on ImageNet reconstruction and class-conditional generation, and scales effectively to text-to-image and text-to-video tasks.
58
+
59
+ ---
60
+
61
+ <figure align="center">
62
+ <figcaption><strong><em>Generative Refinement Framework</em></strong></figcaption>
63
+ <img src="demo/framework.jpg" width="100%" alt="Framework">
64
+ </figure>
65
+
66
+ <p align="center">
67
+ Starting from a random token map, GRN randomly selects more predictions at each step and refines all input tokens. For example, compared to the second step, the third step filled six new tokens (<span style="color: rgb(220, 120, 117);">pink</span>), kept two tokens (<span style="color: rgb(88, 160, 227);">blue</span>), erased two tokens (<span style="color: rgb(240, 180, 40);">yellow</span>), and left six tokens blank (<span style="color: rgb(128, 138, 151);">gray</span>).
68
+ </p>
69
+
70
+ ---
71
+
72
+ ## ✨ Gallery
73
+
74
+ ### GRN-8B Text-to-Video Examples
75
+
76
+ <div align="center">
77
+ <table style="border-spacing: 6px; margin: auto;">
78
+ <tr>
79
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/6ce844dc-3185-4239-bcf1-d72ff20a3031" width="33%" autoplay muted loop playsinline></video></td>
80
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/1697066e-f00f-4e23-a55c-c6af5948c4af" width="33%" autoplay muted loop playsinline></video></td>
81
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/1023ea4d-d814-4be1-95f2-b1623de0f6bd" width="33%" autoplay muted loop playsinline></video></td>
82
+ </tr>
83
+ <tr>
84
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/6244dae4-f480-408a-ac3d-19e4d1ef0a2d" width="33%" autoplay muted loop playsinline></video></td>
85
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/5aefc8d2-bc99-48e4-bd1c-9b3077c9c35e" width="33%" autoplay muted loop playsinline></video></td>
86
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/014e8bb4-04a7-4fa4-a597-d0dfbcc23e02" width="33%" autoplay muted loop playsinline></video></td>
87
+ </tr>
88
+ <tr>
89
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/6bde1f2e-cebe-4f47-9eac-4fe817c3ebc7" width="33%" autoplay muted loop playsinline></video></td>
90
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/b9957300-fa98-411c-83d5-f972621245ad" width="33%" autoplay muted loop playsinline></video></td>
91
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/d07cef92-3eec-4e7c-93da-f6c6a8dc1658" width="33%" autoplay muted loop playsinline></video></td>
92
+ </tr>
93
+ </table>
94
+ </div>
95
+
96
+ ---
97
+
98
+ ### GRN-8B Image-to-Video Examples
99
+
100
+ <div align="center">
101
+ <table style="border-spacing: 6px; margin: auto;">
102
+ <tr>
103
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/527f94b0-4b04-4cbb-a86d-9f5ae05fab67" width="33%" autoplay muted loop playsinline></video></td>
104
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/0b63a9ed-2940-402f-8339-db0d05a09525" width="33%" autoplay muted loop playsinline></video></td>
105
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/8d108a5f-1414-43ca-af6b-8862640741e5" width="33%" autoplay muted loop playsinline></video></td>
106
+ </tr>
107
+ <tr>
108
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/64cd45a9-0c2f-4926-bcc0-b8a0a939ae54" width="33%" autoplay muted loop playsinline></video></td>
109
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/6c31c9e5-0742-4416-925c-16c39bc5a03a" width="33%" autoplay muted loop playsinline></video></td>
110
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/4e966b46-6107-4ffe-a24b-37dc3c8461dd" width="33%" autoplay muted loop playsinline></video></td>
111
+ </tr>
112
+ <tr>
113
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/56ce2dc1-3b64-4493-ab27-b2ba273c64ef" width="33%" autoplay muted loop playsinline></video></td>
114
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/ef45a3f4-8fb2-4bb5-885e-19645e5a0fb5" width="33%" autoplay muted loop playsinline></video></td>
115
+ <td style="padding: 2px;"><video src="https://github.com/user-attachments/assets/98e3fbb0-9a54-49e6-8cec-96a42d0634e6" width="33%" autoplay muted loop playsinline></video></td>
116
+ </tr>
117
+ </table>
118
+ </div>
119
+
120
+ ### GRN-2B Class-to-Image Examples
121
+ <figure align="center">
122
+ <!-- <figcaption><strong><em>GRN-2B Class-to-Image Examples</em></strong></figcaption> -->
123
+ <img src="demo/c2i_examples.jpg" width="100%" alt="Class-to-Image Examples">
124
+ </figure>
125
+
126
+ ### GRN-2B Text-to-Image Examples
127
+ <figure align="center">
128
+ <!-- <figcaption><strong><em>GRN-2B Text-to-Image Examples</em></strong></figcaption> -->
129
+ <img src="demo/t2i_examples.jpg" width="100%" alt="Text-to-Image Examples">
130
+ </figure>
131
+
132
+ ---
133
+
134
+ ## 🚀 Demo
135
+
136
+ ### 🖼️ Text-to-Image
137
+ Try our interactive Text-to-Image demo on 🤗 Hugging Face Space:
138
+
139
+ **[GRN T2I Demo](https://huggingface.co/spaces/hanjian/GRN)**
140
+
141
+ Experience the power of Generative Refinement Networks firsthand by generating images from text prompts directly in your browser!
142
+
143
+ ---
144
+
145
+ ### 🎬 Text-to-Video
146
+ Try our interactive Text-to-Video demo on Discord:
147
+
148
+ [![Discord](https://img.shields.io/badge/Discord-Join%20Server-5865F2?style=for-the-badge&logo=discord&logoColor=white)](http://opensource.bytedance.com/discord/invite)
149
+
150
+
151
+ <figure align="center">
152
+ <figcaption><strong><em>T2V Demo on Discord</em></strong></figcaption>
153
+ <img src="demo/t2v_demo.png" width="100%" alt="T2V Demo">
154
+ </figure>
155
+
156
+ ---
157
+
158
+ ## 📦 Model Zoo
159
+
160
+ | Model | Checkpoints |
161
+ |-------|:-----------:|
162
+ | **Tokenizers** | ✅ [ImageNet Tokenizer](https://huggingface.co/bytedance-research/GRN/blob/main/HBQ_image_tokenizer_16dim_M4.ckpt)<br>✅ [Joint Image/Video Tokenizer](https://huggingface.co/bytedance-research/GRN/blob/main/HBQ_tokenizer_64dim_M4.ckpt) |
163
+ | **GRN_ind_C2I** | ✅ [B](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_ind_B_ep599.pth)<br>⬜ L (TBD)<br>⬜ H (TBD)<br>⬜ G (TBD) |
164
+ | **GRN_bit_T2I** | ✅ [GRN_T2I](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_T2I_2B.pth) |
165
+ | **GRN_bit_T2V** | ✅ [GRN_T2V](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_T2V_2B.pth) |
166
+
167
+ ---
168
+
169
+ ## 🛠️ Installation
170
+
171
+ ### Step 1: Clone the repository
172
+ ```bash
173
+ git clone https://github.com/bytedance/GRN
174
+ cd GRN
175
+ ```
176
+
177
+ ### Step 2: Create conda environment
178
+ A suitable [conda](https://conda.io/) environment named `GRN` can be created and activated with:
179
+ ```bash
180
+ conda env create -f environment.yaml
181
+ conda activate GRN
182
+ ```
183
+
184
+ ### Troubleshooting
185
+ If you get `undefined symbol: iJIT_NotifyEvent` when importing `torch`, simply:
186
+ ```bash
187
+ pip uninstall torch
188
+ pip install torch==2.5.1 --index-url https://download.pytorch.org/whl/cu124
189
+ ```
190
+ Check this [issue](https://github.com/conda/conda/issues/13812#issuecomment-2071445372) for more details.
191
+
192
+ ---
193
+
194
+ ## 🖼️ Class-to-Image
195
+
196
+ ### Dataset
197
+ Download [ImageNet](http://image-net.org/download) dataset, and place it in your `IMAGENET_PATH`.
198
+
199
+ ### Training
200
+
201
+ All training scripts are located in `scripts/c2i/`. We suggest using 8x80GB GPUs for most models.
202
+
203
+ | Model | Training Script | GPUs Required |
204
+ |-------|:-------------:|:-------------:|
205
+ | GRN_ind_B | `bash scripts/c2i/train_GRN_ind_B.sh` | 8x80GB |
206
+ | GRN_bit_B | `bash scripts/c2i/train_GRN_bit_B.sh` | 8x80GB |
207
+ | GRN_ind_L | `bash scripts/c2i/train_GRN_ind_L.sh` | 8x80GB |
208
+ | GRN_ind_H | `bash scripts/c2i/train_GRN_ind_H.sh` | 16x80GB |
209
+ | GRN_ind_G | `bash scripts/c2i/train_GRN_ind_G.sh` | 32x80GB |
210
+
211
+ ### Evaluation
212
+
213
+ PyTorch pre-trained models are available [here](https://huggingface.co/bytedance-research/GRN/tree/main).
214
+
215
+ All evaluation scripts are located in `scripts/c2i/`. We suggest using 8x80GB vRAM GPUs.
216
+
217
+ | Model | Evaluation Script |
218
+ |-------|:--------------:|
219
+ | GRN_ind_B | `bash scripts/c2i/eval_GRN_ind_B.sh` |
220
+ | GRN_bit_B | `bash scripts/c2i/eval_GRN_bit_B.sh` |
221
+ | GRN_ind_L | `bash scripts/c2i/eval_GRN_ind_L.sh` |
222
+ | GRN_ind_H | `bash scripts/c2i/eval_GRN_ind_H.sh` |
223
+ | GRN_ind_G | `bash scripts/c2i/eval_GRN_ind_G.sh` |
224
+
225
+ We use [torch-fidelity](https://github.com/LTH14/torch-fidelity) to evaluate FID and IS against a reference image folder or statistics. We use the JiT's pre-computed reference stats under `grn/utils_c2i/fid_stats`.
226
+
227
+ ---
228
+
229
+ ## 🎨 Text-to-Image
230
+ ### Data
231
+ Refer to `data/toy_data/jsonls/000001/0001_0008_000000100.jsonl`
232
+ ```
233
+ {"image_path": "[image_path_1]", "long_caption": "xxx", "long_caption_type": "caption-InternVL2.0", "text": "", "short_caption_type": "blip2_caption", "width": 1080, "height": 1920}
234
+ {"image_path": "[image_path_2]", "long_caption": "xxx", "long_caption_type": "caption-InternVL2.0", "text": "", "short_caption_type": "blip2_caption", "width": 1080, "height": 1920}
235
+ ...
236
+ ```
237
+
238
+ ### Train
239
+ Run `bash scripts/train_GRN_ind_t2i.sh`
240
+
241
+ ### Inference
242
+
243
+ You can simply run `python3 t2i_infer.py` or use the following code:
244
+
245
+ ```python
246
+ from PIL import Image
247
+ from grn_pipeline import GRNPipeline
248
+
249
+ # Load pipeline
250
+ pipeline = GRNPipeline.from_pretrained(
251
+ hf_repo_id='bytedance-research/GRN',
252
+ task='T2I',
253
+ pn='1M',
254
+ device='cpu',
255
+ ).to('cuda')
256
+
257
+ # Generate one image
258
+ result = pipeline(
259
+ prompt="A cute cat playing in the garden",
260
+ guidance_scale=3.0,
261
+ temperature=1.1,
262
+ complexity_aware_Tmin=10,
263
+ complexity_aware_Tmax=50,
264
+ complexity_aware_k = 0,
265
+ complexity_aware_b = 50,
266
+ complexity_aware_wp = 5,
267
+ snr_shift = 1.,
268
+ h_div_w=1.,
269
+ content_type='image',
270
+ seed=42,
271
+ )
272
+ image = result.images[0]
273
+ image.save('./generated_image.jpg')
274
+ ```
275
+
276
+ ---
277
+
278
+ ## 🎬 Text-to-Video
279
+ ### Data
280
+ Refer to `data/toy_data/jsonls/000001/0001_0008_000000100.jsonl`
281
+ ```
282
+ {"video_path": "[video_path_1]", "begin_frame_id": xxx, "end_frame_id": xxx, "quality_prompt": "There is text in the video.", "fps": 25.0, "duration": 3.88, "width": 1280, "height": 720, "caption": [{"type": "short", "content": "[short_caption]"}, {"type": "medium", "content": "[medium_caption]"}, {"type": "long", "content": "[long_caption]"}]}
283
+ {"video_path": "[video_path_1]", "begin_frame_id": xxx, "end_frame_id": xxx, "quality_prompt": "The quality is very high!", "fps": 25.0, "duration": 3.88, "width": 1280, "height": 720, "caption": [{"type": "short", "content": "[short_caption]"}, {"type": "medium", "content": "[medium_caption]"}, {"type": "long", "content": "[long_caption]"}]}
284
+ ...
285
+ ```
286
+
287
+ ### Train
288
+ Run `bash scripts/train_GRN_ind_t2v.sh`
289
+
290
+ ### Inference
291
+
292
+ You can simply run `python3 t2v_infer.py` or use the following code:
293
+
294
+ ```python
295
+ from grn_pipeline import GRNPipeline
296
+
297
+ # Load pipeline
298
+ pipeline = GRNPipeline.from_pretrained(
299
+ hf_repo_id='bytedance-research/GRN',
300
+ task='T2V',
301
+ pn='0.41M',
302
+ device='cpu'
303
+ ).to('cuda')
304
+
305
+ # Generate one video
306
+ result = pipeline(
307
+ prompt="Two women demonstrate a makeup product, applying it with a sponge while smiling and engaging with the camera in a bright, clean setting.",
308
+ guidance_scale=4.0,
309
+ temperature=1.0,
310
+ complexity_aware_Tmin=10,
311
+ complexity_aware_Tmax=50,
312
+ complexity_aware_k = 0,
313
+ complexity_aware_b = 50,
314
+ complexity_aware_wp = 5,
315
+ snr_shift = 1.,
316
+ h_div_w=9/16,
317
+ duration=2.,
318
+ content_type='video',
319
+ seed=42,
320
+ )
321
+ video_file = result.videos[0]
322
+ ```
323
+
324
+ ---
325
+
326
+ ## 📦 HBQ Tokenizer
327
+
328
+ ### Data
329
+ Image Dataset, e.g., data_root/username/labels/imagenet/train.txt:
330
+ ```
331
+ [image_1_full_path]
332
+ [image_2_full_path]
333
+ [image_3_full_path]
334
+ ...
335
+ ```
336
+
337
+ Video Dataset, e.g., data_root/username/labels_hanjian/high-quality-video/horizontal_videos.txt
338
+ ```
339
+ [video_1_full_path]
340
+ [video_2_full_path]
341
+ [video_3_full_path]
342
+ ...
343
+ ```
344
+
345
+ ### Training
346
+ For example, set `latent_channels=16/64` and `quant_method=hierarchical_binary_quant_round_4` in `scripts/hbq_tokenizer_train.sh`, then run:
347
+ ```bash
348
+ cd grn/tokenizer
349
+ bash scripts/hbq_tokenizer_train.sh
350
+ ```
351
+
352
+ ### Evaluation
353
+ For example, set `latent_channels=16/64` and `quant_method=hierarchical_binary_quant_round_4` in `scripts/hbq_tokenizer_train.sh`, then run:
354
+ ```bash
355
+ cd grn/tokenizer
356
+ bash scripts/hbq_tokenizer_eval.sh
357
+ ```
358
+
359
+ ---
360
+
361
+ ## 📧 Contact
362
+
363
+ If you are interested in scaling GRN for image generation / image editing / video generation / video editing / unified model directions, please feel free to reach out!
364
+
365
+ **📧 Email:** [hanjian.thu123@bytedance.com](mailto:hanjian.thu123@bytedance.com)
366
+
367
+ ---
368
+
369
+ ## 🤗 Acknowledgements
370
+
371
+ - Thanks to [JiT](https://github.com/LTH14/JiT), [Infinity](https://github.com/FoundationVision/Infinity) and [InfinityStar](https://github.com/FoundationVision/InfinityStar) for their wonderful work and codebase!
372
+
373
+ ---
374
+
375
+ ## 📝 Citation
376
+
377
+ If you find our work useful, please consider citing:
378
+
379
+ ```bibtex
380
+ @misc{han2026grn,
381
+ title={Generative Refinement Networks for Visual Synthesis},
382
+ author={Jian Han and Jinlai Liu and Jiahuan Wang and Bingyue Peng and Zehuan Yuan},
383
+ year={2026},
384
+ eprint={2604.13030},
385
+ archivePrefix={arXiv},
386
+ primaryClass={cs.CV},
387
+ url={https://arxiv.org/abs/2604.13030},
388
+ }
389
+ ```
app.py ADDED
@@ -0,0 +1,132 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import traceback
4
+ import torch
5
+ import gradio as gr
6
+ import spaces
7
+
8
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
9
+
10
+ from grn_pipeline import GRNPipeline
11
+
12
+ # Global pipeline
13
+ pipe = None
14
+ device = "cuda" if torch.cuda.is_available() else "cpu"
15
+
16
+ def load_pipeline():
17
+ global pipe
18
+ print("Loading GRN pipeline...")
19
+ # 从 Hugging Face Hub 下载权重
20
+ pipe = GRNPipeline.from_pretrained(
21
+ hf_repo_id='bytedance-research/GRN',
22
+ task='T2I',
23
+ pn='1M',
24
+ model='GRN2b',
25
+ use_slow_attn=True,
26
+ device=device,
27
+ )
28
+ print("Pipeline loaded successfully!")
29
+ return pipe
30
+
31
+ # @spaces.GPU #[uncomment to use ZeroGPU]
32
+ @spaces.GPU(duration=40)
33
+ def generate(prompt, content_type="image", guidance_scale=3.0, temperature=1.0, seed=42, width=1024, height=1024):
34
+ global pipe
35
+ if pipe is None:
36
+ try:
37
+ pipe = load_pipeline()
38
+ except Exception as e:
39
+ print(f"Error loading pipeline: {e}")
40
+ traceback.print_exc()
41
+ return f"Error loading pipeline: {e}\n\n{traceback.format_exc()}"
42
+
43
+ try:
44
+ result = pipe(
45
+ prompt="<T2I>"+prompt,
46
+ guidance_scale=guidance_scale,
47
+ temperature=temperature,
48
+ complexity_aware_Tmin=10,
49
+ complexity_aware_Tmax=50,
50
+ complexity_aware_k = 0,
51
+ complexity_aware_b = 50,
52
+ complexity_aware_wp = 5,
53
+ snr_shift = 1.,
54
+ h_div_w=1.,
55
+ content_type=content_type,
56
+ seed=seed,
57
+ width=width,
58
+ height=height
59
+ )
60
+
61
+ if content_type == "image" and hasattr(result, 'images'):
62
+ return result.images[0]
63
+ elif content_type == "video" and hasattr(result, 'videos'):
64
+ return result.videos[0]
65
+ return f"Error: Invalid result from pipeline"
66
+ except Exception as e:
67
+ print(f"Error generating content: {e}")
68
+ traceback.print_exc()
69
+ return f"Error generating content: {e}\n\n{traceback.format_exc()}"
70
+
71
+ def create_demo():
72
+ with gr.Blocks(title="GRN: Generative Refinement Networks", theme=gr.themes.Soft()) as demo:
73
+ gr.Markdown("# GRN: Generative Refinement Networks")
74
+ gr.Markdown("Text-to-Image generation using GRN")
75
+
76
+ with gr.Row():
77
+ with gr.Column():
78
+ prompt_input = gr.Textbox(
79
+ label="Text Prompt",
80
+ placeholder="Enter your prompt here...",
81
+ value="A cute cat playing in the garden"
82
+ )
83
+
84
+ content_type = gr.Radio(
85
+ choices=["image"], # , "video"
86
+ value="image",
87
+ label="Content Type"
88
+ )
89
+
90
+ with gr.Accordion("Settings", open=True):
91
+ guidance_scale = gr.Slider(minimum=0, maximum=10, value=3.0, label="Guidance Scale")
92
+ temperature = gr.Slider(minimum=0.1, maximum=1.5, value=1.1, label="Temperature")
93
+ seed = gr.Number(value=42, label="Seed", precision=0)
94
+ width = gr.Number(value=1024, label="Width", precision=0)
95
+ height = gr.Number(value=1024, label="Height", precision=0)
96
+
97
+ generate_btn = gr.Button("Generate", variant="primary")
98
+
99
+ with gr.Column():
100
+ output = gr.Gallery(label="Output", show_label=True, elem_id="gallery", columns=1, height="auto", preview=True, object_fit="contain")
101
+
102
+ def generate_and_display(prompt, content_type, guidance_scale, temperature, seed, width, height):
103
+ result = generate(prompt, content_type, guidance_scale, temperature, seed, width, height)
104
+ if result:
105
+ return [result]
106
+ return []
107
+
108
+ generate_btn.click(
109
+ fn=generate_and_display,
110
+ inputs=[prompt_input, content_type, guidance_scale, temperature, seed, width, height],
111
+ outputs=output
112
+ )
113
+
114
+ gr.Examples(
115
+ examples=[
116
+ ["A majestic lion standing on a cliff at sunset", "image", 3.0, 1.0, 42, 1024, 1024],
117
+ ],
118
+ inputs=[prompt_input, content_type, guidance_scale, temperature, seed, width, height],
119
+ cache_examples=False
120
+ )
121
+
122
+ return demo
123
+
124
+ if __name__ == "__main__":
125
+ try:
126
+ load_pipeline()
127
+ except Exception as e:
128
+ print(f"Error loading pipeline: {e}")
129
+ traceback.print_exc()
130
+
131
+ demo = create_demo()
132
+ demo.launch()
c2i_train_infer.py ADDED
@@ -0,0 +1,378 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import datetime
3
+ import numpy as np
4
+ import os
5
+ import time
6
+ import functools
7
+ from pathlib import Path
8
+
9
+ import torch
10
+ import torch.backends.cudnn as cudnn
11
+ from torch.utils.tensorboard import SummaryWriter
12
+ import torchvision.transforms as transforms
13
+ import torchvision.datasets as datasets
14
+ from torch.distributed.device_mesh import init_device_mesh
15
+ from torch.distributed.fsdp import (
16
+ FullyShardedDataParallel as FSDP,
17
+ MixedPrecision,
18
+ BackwardPrefetch,
19
+ ShardingStrategy,
20
+ FullStateDictConfig,
21
+ StateDictType,
22
+ )
23
+ from torch.distributed.fsdp.wrap import (
24
+ transformer_auto_wrap_policy,
25
+ enable_wrap,
26
+ wrap,
27
+ )
28
+
29
+ from grn.utils_c2i.crop import center_crop_arr
30
+ import grn.utils_c2i.misc as misc
31
+
32
+ import copy
33
+ from grn.utils_c2i.engine import train_one_epoch, evaluate
34
+ from grn.utils import wandb_utils as wandb_utils
35
+
36
+ from grn.utils_c2i.denoiser import Denoiser
37
+ from grn.models.grn_c2i import GRNblock
38
+
39
+
40
+ def get_args_parser():
41
+ parser = argparse.ArgumentParser('GRN', add_help=False)
42
+
43
+ # architecture
44
+ parser.add_argument('--model', default='GRN_B', type=str, metavar='MODEL',
45
+ help='Name of the model to train')
46
+ parser.add_argument('--img_size', default=256, type=int, help='Image size')
47
+ parser.add_argument('--attn_dropout', type=float, default=0.0, help='Attention dropout rate')
48
+ parser.add_argument('--proj_dropout', type=float, default=0.0, help='Projection dropout rate')
49
+
50
+ # training
51
+ parser.add_argument('--epochs', default=200, type=int)
52
+ parser.add_argument('--warmup_epochs', type=int, default=5, metavar='N',
53
+ help='Epochs to warm up LR')
54
+ parser.add_argument('--batch_size', default=128, type=int,
55
+ help='Batch size per GPU (effective batch size = batch_size * # GPUs)')
56
+ parser.add_argument('--lr', type=float, default=None, metavar='LR',
57
+ help='Learning rate (absolute)')
58
+ parser.add_argument('--blr', type=float, default=5e-5, metavar='LR',
59
+ help='Base learning rate: absolute_lr = base_lr * total_batch_size / 256')
60
+ parser.add_argument('--min_lr', type=float, default=0., metavar='LR',
61
+ help='Minimum LR for cyclic schedulers that hit 0')
62
+ parser.add_argument('--lr_schedule', type=str, default='constant',
63
+ help='Learning rate schedule')
64
+ parser.add_argument('--weight_decay', type=float, default=0.0,
65
+ help='Weight decay (default: 0.0)')
66
+ parser.add_argument('--ema_decay1', type=float, default=0.9999,
67
+ help='The first ema to track. Use the first ema for sampling by default.')
68
+ parser.add_argument('--ema_decay2', type=float, default=0.9996,
69
+ help='The second ema to track')
70
+ parser.add_argument('--P_mean', default=-0.8, type=float)
71
+ parser.add_argument('--P_std', default=0.8, type=float)
72
+ parser.add_argument('--noise_scale', default=1.0, type=float)
73
+ parser.add_argument('--t_eps', default=5e-2, type=float)
74
+ parser.add_argument('--label_drop_prob', default=0.1, type=float)
75
+
76
+ parser.add_argument('--seed', default=0, type=int)
77
+ parser.add_argument('--start_epoch', default=0, type=int, metavar='N',
78
+ help='Starting epoch')
79
+ parser.add_argument('--num_workers', default=12, type=int)
80
+ parser.add_argument('--pin_mem', action='store_true',
81
+ help='Pin CPU memory in DataLoader for faster GPU transfers')
82
+ parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem')
83
+ parser.set_defaults(pin_mem=True)
84
+
85
+ # sampling
86
+ parser.add_argument('--sampling_method', default='heun', type=str,
87
+ help='ODE samping method')
88
+ parser.add_argument('--num_sampling_steps', default=50, type=int,
89
+ help='Sampling steps')
90
+ parser.add_argument('--cfg', default=1.0, type=float,
91
+ help='Classifier-free guidance factor')
92
+ parser.add_argument('--interval_min', default=0.0, type=float,
93
+ help='CFG interval min')
94
+ parser.add_argument('--interval_max', default=1.0, type=float,
95
+ help='CFG interval max')
96
+ parser.add_argument('--num_images', default=50000, type=int,
97
+ help='Number of images to generate')
98
+ parser.add_argument('--eval_freq', type=int, default=40,
99
+ help='Frequency (in epochs) for evaluation')
100
+ parser.add_argument('--online_eval', type=int, default=0, choices=[0,1],
101
+ help='Whether to evaluate the model online')
102
+ parser.add_argument('--evaluate_gen', action='store_true')
103
+ parser.add_argument('--gen_bsz', type=int, default=256,
104
+ help='Generation batch size')
105
+
106
+ # dataset
107
+ parser.add_argument('--data_path', default='./data/imagenet', type=str,
108
+ help='Path to the dataset')
109
+ parser.add_argument('--class_num', default=1000, type=int)
110
+
111
+ # checkpointing
112
+ parser.add_argument('--output_dir', default='./output_dir',
113
+ help='Directory to save outputs (empty for no saving)')
114
+ parser.add_argument('--resume', default='',
115
+ help='Folder that contains checkpoint to resume from')
116
+ parser.add_argument('--save_last_freq', type=int, default=5,
117
+ help='Frequency (in epochs) to save checkpoints')
118
+ parser.add_argument('--log_freq', default=100, type=int)
119
+ parser.add_argument('--device', default='cuda',
120
+ help='Device to use for training/testing')
121
+
122
+ # distributed training
123
+ parser.add_argument('--world_size', default=1, type=int,
124
+ help='Number of distributed processes')
125
+ parser.add_argument('--local_rank', default=-1, type=int)
126
+ parser.add_argument('--dist_on_itp', action='store_true')
127
+ parser.add_argument('--dist_url', default='env://',
128
+ help='URL used to set up distributed training')
129
+ parser.add_argument('--hbq_round', default=4, type=int,)
130
+ parser.add_argument('--in_channels', default=3, type=int,)
131
+ parser.add_argument('--method', default='GRN_ind', type=str, choices=['GRN_ind', 'GRN_bit'])
132
+ parser.add_argument('--vae_path', default='', type=str,)
133
+ parser.add_argument('--tau', default=1.0, type=float,)
134
+ parser.add_argument('--wandb', default=1, type=int, choices=[0,1])
135
+ parser.add_argument('--generation_dir', default='/tmp', type=str)
136
+ parser.add_argument('--clip_grad_norm', default=1., type=float)
137
+ parser.add_argument('--use_fsdp_train', default=0, type=int, choices=[0, 1])
138
+ parser.add_argument('--delete_images', default=1, type=int, choices=[0, 1])
139
+ parser.add_argument('--use_confidence_sampling', default=0, type=int, choices=[0, 1])
140
+ parser.add_argument('--inner_shard_degree', default=8, type=int)
141
+ parser.add_argument('--patch_size', default=1, type=int)
142
+ parser.add_argument('--convert_type', default='', type=str)
143
+ parser.add_argument('--mask_group_size', default=-1, type=int)
144
+ parser.add_argument('--grn_shift_factor', default=1., type=float)
145
+ parser.add_argument('--use_focal_loss', default=0, type=int, choices=[0, 1])
146
+ return parser
147
+
148
+
149
+ def main(args):
150
+ misc.init_distributed_mode(args)
151
+ print('Job directory:', os.path.dirname(os.path.realpath(__file__)))
152
+ print("Arguments:\n{}".format(args).replace(', ', ',\n'))
153
+
154
+ device = torch.device(args.device)
155
+
156
+ # Set seeds for reproducibility
157
+ seed = args.seed + misc.get_rank()
158
+ torch.manual_seed(seed)
159
+ np.random.seed(seed)
160
+
161
+ cudnn.benchmark = True
162
+
163
+ num_tasks = misc.get_world_size()
164
+ global_rank = misc.get_rank()
165
+
166
+ # Set up TensorBoard logging (only on main process)
167
+ if global_rank == 0 and args.output_dir is not None:
168
+ os.makedirs(args.output_dir, exist_ok=True)
169
+ log_writer = SummaryWriter(log_dir=args.output_dir)
170
+ if args.wandb:
171
+ entity = os.environ["EXP_NAME"]
172
+ project = os.environ["PROJECT"]
173
+ wandb_utils.wandb.init(project=project, name=entity, config={})
174
+ else:
175
+ log_writer = None
176
+
177
+ # Data augmentation transforms
178
+ transform_train = transforms.Compose([
179
+ transforms.Lambda(lambda img: center_crop_arr(img, args.img_size)),
180
+ transforms.RandomHorizontalFlip(),
181
+ transforms.PILToTensor()
182
+ ])
183
+
184
+ dataset_train = datasets.ImageFolder(os.path.join(args.data_path, 'train'), transform=transform_train)
185
+ print(dataset_train)
186
+
187
+ sampler_train = torch.utils.data.DistributedSampler(
188
+ dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True
189
+ )
190
+ print("Sampler_train =", sampler_train)
191
+
192
+ data_loader_train = torch.utils.data.DataLoader(
193
+ dataset_train, sampler=sampler_train,
194
+ batch_size=args.batch_size,
195
+ num_workers=args.num_workers,
196
+ pin_memory=args.pin_mem,
197
+ drop_last=True
198
+ )
199
+
200
+ torch._dynamo.config.cache_size_limit = 128
201
+ torch._dynamo.config.optimize_ddp = False
202
+
203
+ # Create denoiser
204
+ model = Denoiser(args)
205
+
206
+ # ininitalize vae
207
+ from grn.models.hbq_tokenizer import HBQ_Tokenizer
208
+ vae = HBQ_Tokenizer(args=args, latent_channels=16, encoder_out_type='feature_tanh')
209
+ vae.eval()
210
+ vae = vae.to('cuda')
211
+ for param in vae.parameters():
212
+ param.requires_grad = False
213
+ state_dict = torch.load(args.vae_path, map_location='cuda')
214
+ if 'ema' in state_dict:
215
+ print(f'Load ema vae weights')
216
+ state_dict = state_dict['ema']
217
+ else:
218
+ print(f'Load non ema vae weights')
219
+ state_dict = state_dict['vae']
220
+ print('Load vae: ', vae.load_state_dict(state_dict, assign=True))
221
+
222
+ print("Model =", model)
223
+ n_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
224
+ print("Number of trainable parameters: {:.6f}M".format(n_params / 1e6))
225
+
226
+ model.to(device)
227
+
228
+ eff_batch_size = args.batch_size * misc.get_world_size()
229
+ if args.lr is None: # only base_lr (blr) is specified
230
+ args.lr = args.blr * eff_batch_size / 256
231
+
232
+ print("Base lr: {:.2e}".format(args.lr * 256 / eff_batch_size))
233
+ print("Actual lr: {:.2e}".format(args.lr))
234
+ print("Effective batch size: %d" % eff_batch_size)
235
+
236
+ if args.use_fsdp_train:
237
+ auto_wrap_policy = functools.partial(
238
+ transformer_auto_wrap_policy,
239
+ transformer_layer_cls={GRNblock},
240
+ )
241
+ if args.inner_shard_degree > 0:
242
+ sharding_strategy = ShardingStrategy.HYBRID_SHARD
243
+ world_size = misc.get_world_size()
244
+ assert world_size % args.inner_shard_degree == 0
245
+ assert args.inner_shard_degree > 1 and args.inner_shard_degree <= world_size
246
+ device_mesh = init_device_mesh('cuda', (world_size // args.inner_shard_degree, args.inner_shard_degree))
247
+ else:
248
+ sharding_strategy = ShardingStrategy.FULL_SHARD
249
+ device_mesh = None
250
+ model = FSDP(
251
+ model,
252
+ auto_wrap_policy=auto_wrap_policy,
253
+ mixed_precision=MixedPrecision(
254
+ param_dtype=torch.bfloat16,
255
+ reduce_dtype=torch.bfloat16,
256
+ buffer_dtype=torch.bfloat16
257
+ ),
258
+ device_id=torch.cuda.current_device(),
259
+ sharding_strategy=sharding_strategy,
260
+ use_orig_params=True,
261
+ device_mesh=device_mesh,
262
+ )
263
+ model_without_ddp = model
264
+ else:
265
+ model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu])
266
+ model_without_ddp = model.module
267
+
268
+ # Set up optimizer with weight decay adjustment for bias and norm layers
269
+ param_groups = misc.add_weight_decay(model_without_ddp, args.weight_decay)
270
+ optimizer = torch.optim.AdamW(param_groups, lr=args.lr, betas=(0.9, 0.95))
271
+ print(optimizer)
272
+
273
+ # Resume from checkpoint if provided
274
+ # checkpoint_path = os.path.join(args.resume, "checkpoint-last.pth") if args.resume else None
275
+ checkpoint_path = args.resume if args.resume else None
276
+ if checkpoint_path and os.path.exists(checkpoint_path):
277
+ checkpoint = torch.load(checkpoint_path, map_location='cpu')
278
+ model_without_ddp.load_state_dict(checkpoint['model'])
279
+
280
+ if args.use_fsdp_train:
281
+ # For FSDP, load EMA state dict into model temporarily to set ema_params
282
+ model_without_ddp.load_state_dict(checkpoint['model_ema1'])
283
+ model_without_ddp.module.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
284
+
285
+ model_without_ddp.load_state_dict(checkpoint['model_ema2'])
286
+ model_without_ddp.module.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
287
+
288
+ # Restore model
289
+ model_without_ddp.load_state_dict(checkpoint['model'])
290
+ else:
291
+ ema_state_dict1 = checkpoint['model_ema1']
292
+ ema_state_dict2 = checkpoint['model_ema2']
293
+ model_without_ddp.ema_params1 = [ema_state_dict1[name].cuda() for name, _ in model_without_ddp.named_parameters()]
294
+ model_without_ddp.ema_params2 = [ema_state_dict2[name].cuda() for name, _ in model_without_ddp.named_parameters()]
295
+
296
+ print("Resumed checkpoint from", args.resume)
297
+
298
+ try:
299
+ if 'optimizer' in checkpoint and 'epoch' in checkpoint:
300
+ if args.use_fsdp_train:
301
+ opt_state = FSDP.optim_state_dict_to_load(
302
+ model_without_ddp, optimizer, checkpoint['optimizer']
303
+ )
304
+ optimizer.load_state_dict(opt_state)
305
+ else:
306
+ optimizer.load_state_dict(checkpoint['optimizer'])
307
+ print("Loaded optimizer & scaler state!")
308
+ except:
309
+ print("Failed to load optimizer & scaler state! Just load checkpoint.")
310
+ args.start_epoch = checkpoint['epoch'] + 1
311
+ del checkpoint
312
+ else:
313
+ if args.use_fsdp_train:
314
+ model_without_ddp.module.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
315
+ model_without_ddp.module.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
316
+ else:
317
+ model_without_ddp.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
318
+ model_without_ddp.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
319
+ print("Training from scratch")
320
+
321
+ # Evaluate generation
322
+ if args.evaluate_gen:
323
+ print("Evaluating checkpoint at {} epoch".format(args.start_epoch))
324
+ with torch.random.fork_rng():
325
+ torch.manual_seed(seed)
326
+ with torch.no_grad():
327
+ evaluate(model_without_ddp, args, args.start_epoch, batch_size=args.gen_bsz, log_writer=log_writer, vae=vae)
328
+ return
329
+
330
+ # Training loop
331
+ print(f"Start training for {args.epochs} epochs")
332
+ start_time = time.time()
333
+ for epoch in range(args.start_epoch, args.epochs):
334
+ if args.distributed:
335
+ data_loader_train.sampler.set_epoch(epoch)
336
+
337
+ train_one_epoch(model, model_without_ddp, data_loader_train, optimizer, device, epoch, log_writer=log_writer, args=args, vae=vae)
338
+
339
+ # Save checkpoint periodically
340
+ if epoch % args.save_last_freq == 0 or epoch + 1 == args.epochs:
341
+ if misc.is_main_process():
342
+ from grn.utils.safe_rm import safe_remove
343
+ safe_remove(f'{args.output_dir}/checkpoint-tmp_*.pth', args.output_dir)
344
+ misc.save_model(
345
+ args=args,
346
+ model_without_ddp=model_without_ddp,
347
+ optimizer=optimizer,
348
+ epoch=epoch,
349
+ epoch_name=f"tmp_{epoch}"
350
+ )
351
+
352
+ if epoch % 100 == 0 and epoch > 0:
353
+ misc.save_model(
354
+ args=args,
355
+ model_without_ddp=model_without_ddp,
356
+ optimizer=optimizer,
357
+ epoch=epoch
358
+ )
359
+
360
+ # Perform online evaluation at specified intervals
361
+ if args.online_eval and (epoch % args.eval_freq == 0 or epoch + 1 == args.epochs):
362
+ torch.cuda.empty_cache()
363
+ with torch.no_grad():
364
+ evaluate(model_without_ddp, args, epoch, batch_size=args.gen_bsz, log_writer=log_writer, vae=vae)
365
+ torch.cuda.empty_cache()
366
+
367
+ if misc.is_main_process() and log_writer is not None:
368
+ log_writer.flush()
369
+
370
+ total_time = time.time() - start_time
371
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
372
+ print('Training time:', total_time_str)
373
+
374
+
375
+ if __name__ == '__main__':
376
+ args = get_args_parser().parse_args()
377
+ Path(args.output_dir).mkdir(parents=True, exist_ok=True)
378
+ main(args)
environment.yaml ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ name: grn
2
+ channels:
3
+ - pytorch
4
+ - defaults
5
+ - nvidia
6
+ dependencies:
7
+ - python=3.10
8
+ - pip=22.3
9
+ - pytorch-cuda=12.4
10
+ - pytorch=2.5.1
11
+ - torchvision=0.20.1
12
+ - numpy=1.22
13
+ - pip:
14
+ - opencv-python==4.11.0.86
15
+ - timm==0.9.12
16
+ - tensorboard==2.10.0
17
+ - scipy==1.9.1
18
+ - einops==0.8.1
19
+ - gdown==5.2.0
evaluation/gen_eval/_base_/datasets/coco_panoptic.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # dataset settings
2
+ dataset_type = 'CocoPanopticDataset'
3
+ data_root = 'data/coco/'
4
+ img_norm_cfg = dict(
5
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
6
+ train_pipeline = [
7
+ dict(type='LoadImageFromFile'),
8
+ dict(
9
+ type='LoadPanopticAnnotations',
10
+ with_bbox=True,
11
+ with_mask=True,
12
+ with_seg=True),
13
+ dict(type='Resize', img_scale=(1333, 800), keep_ratio=True),
14
+ dict(type='RandomFlip', flip_ratio=0.5),
15
+ dict(type='Normalize', **img_norm_cfg),
16
+ dict(type='Pad', size_divisor=32),
17
+ dict(type='SegRescale', scale_factor=1 / 4),
18
+ dict(type='DefaultFormatBundle'),
19
+ dict(
20
+ type='Collect',
21
+ keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks', 'gt_semantic_seg']),
22
+ ]
23
+ test_pipeline = [
24
+ dict(type='LoadImageFromFile'),
25
+ dict(
26
+ type='MultiScaleFlipAug',
27
+ img_scale=(1333, 800),
28
+ flip=False,
29
+ transforms=[
30
+ dict(type='Resize', keep_ratio=True),
31
+ dict(type='RandomFlip'),
32
+ dict(type='Normalize', **img_norm_cfg),
33
+ dict(type='Pad', size_divisor=32),
34
+ dict(type='ImageToTensor', keys=['img']),
35
+ dict(type='Collect', keys=['img']),
36
+ ])
37
+ ]
38
+ data = dict(
39
+ samples_per_gpu=2,
40
+ workers_per_gpu=2,
41
+ train=dict(
42
+ type=dataset_type,
43
+ ann_file=data_root + 'annotations/panoptic_train2017.json',
44
+ img_prefix=data_root + 'train2017/',
45
+ seg_prefix=data_root + 'annotations/panoptic_train2017/',
46
+ pipeline=train_pipeline),
47
+ val=dict(
48
+ type=dataset_type,
49
+ ann_file=data_root + 'annotations/panoptic_val2017.json',
50
+ img_prefix=data_root + 'val2017/',
51
+ seg_prefix=data_root + 'annotations/panoptic_val2017/',
52
+ pipeline=test_pipeline),
53
+ test=dict(
54
+ type=dataset_type,
55
+ ann_file=data_root + 'annotations/panoptic_val2017.json',
56
+ img_prefix=data_root + 'val2017/',
57
+ seg_prefix=data_root + 'annotations/panoptic_val2017/',
58
+ pipeline=test_pipeline))
59
+ evaluation = dict(interval=1, metric=['PQ'])
evaluation/gen_eval/_base_/default_runtime.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ checkpoint_config = dict(interval=1)
2
+ # yapf:disable
3
+ log_config = dict(
4
+ interval=50,
5
+ hooks=[
6
+ dict(type='TextLoggerHook'),
7
+ # dict(type='TensorboardLoggerHook')
8
+ ])
9
+ # yapf:enable
10
+ custom_hooks = [dict(type='NumClassCheckHook')]
11
+
12
+ dist_params = dict(backend='nccl')
13
+ log_level = 'INFO'
14
+ load_from = None
15
+ resume_from = None
16
+ workflow = [('train', 1)]
17
+
18
+ # disable opencv multithreading to avoid system being overloaded
19
+ opencv_num_threads = 0
20
+ # set multi-process start method as `fork` to speed up the training
21
+ mp_start_method = 'fork'
22
+
23
+ # Default setting for scaling LR automatically
24
+ # - `enable` means enable scaling LR automatically
25
+ # or not by default.
26
+ # - `base_batch_size` = (8 GPUs) x (2 samples per GPU).
27
+ auto_scale_lr = dict(enable=False, base_batch_size=16)
evaluation/gen_eval/evaluate_images.py ADDED
@@ -0,0 +1,298 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Evaluate generated images using Mask2Former (or other object detector model)
3
+ """
4
+
5
+ import argparse
6
+ import json
7
+ import os
8
+ import re
9
+ import sys
10
+ import time
11
+ import os.path as osp
12
+
13
+ import warnings
14
+ warnings.filterwarnings("ignore")
15
+
16
+ import numpy as np
17
+ import tqdm
18
+ import pandas as pd
19
+ from PIL import Image, ImageOps
20
+ import torch
21
+ import mmdet
22
+ from mmdet.apis import inference_detector, init_detector
23
+
24
+ import open_clip
25
+ from clip_benchmark.metrics import zeroshot_classification as zsc
26
+ zsc.tqdm = lambda it, *args, **kwargs: it
27
+
28
+ # Get directory path
29
+
30
+ def parse_args():
31
+ parser = argparse.ArgumentParser()
32
+ parser.add_argument("imagedir", type=str)
33
+ parser.add_argument("--outfile", type=str, default="results.jsonl")
34
+ parser.add_argument("--model-config", type=str, default="")
35
+ parser.add_argument("--model-path", type=str, default="")
36
+ # Other arguments
37
+ parser.add_argument("--options", nargs="*", type=str, default=[])
38
+ args = parser.parse_args()
39
+ args.options = dict(opt.split("=", 1) for opt in args.options)
40
+ if args.model_config is None:
41
+ args.model_config = os.path.join(
42
+ os.path.dirname(mmdet.__file__),
43
+ "../configs/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py"
44
+ )
45
+ return args
46
+
47
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
48
+ assert DEVICE == "cuda"
49
+
50
+ def timed(fn):
51
+ def wrapper(*args, **kwargs):
52
+ startt = time.time()
53
+ result = fn(*args, **kwargs)
54
+ endt = time.time()
55
+ print(f'Function {fn.__name__!r} executed in {endt - startt:.3f}s', file=sys.stderr)
56
+ return result
57
+ return wrapper
58
+
59
+ # Load models
60
+
61
+ @timed
62
+ def load_models(args):
63
+ CONFIG_PATH = osp.abspath(args.model_config)
64
+ OBJECT_DETECTOR = args.options.get('model', "mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco")
65
+ CKPT_PATH = os.path.join(args.model_path, f"{OBJECT_DETECTOR}.pth")
66
+ object_detector = init_detector(CONFIG_PATH, CKPT_PATH, device=DEVICE)
67
+
68
+ clip_arch = args.options.get('clip_model', "ViT-L-14")
69
+ clip_model, _, transform = open_clip.create_model_and_transforms(clip_arch, pretrained="openai", device=DEVICE)
70
+ tokenizer = open_clip.get_tokenizer(clip_arch)
71
+
72
+ with open(os.path.join(os.path.dirname(__file__), "object_names.txt")) as cls_file:
73
+ classnames = [line.strip() for line in cls_file]
74
+
75
+ return object_detector, (clip_model, transform, tokenizer), classnames
76
+
77
+
78
+ COLORS = ["red", "orange", "yellow", "green", "blue", "purple", "pink", "brown", "black", "white"]
79
+ COLOR_CLASSIFIERS = {}
80
+
81
+ # Evaluation parts
82
+
83
+ class ImageCrops(torch.utils.data.Dataset):
84
+ def __init__(self, image: Image.Image, objects):
85
+ self._image = image.convert("RGB")
86
+ bgcolor = args.options.get('bgcolor', "#999")
87
+ if bgcolor == "original":
88
+ self._blank = self._image.copy()
89
+ else:
90
+ self._blank = Image.new("RGB", image.size, color=bgcolor)
91
+ self._objects = objects
92
+
93
+ def __len__(self):
94
+ return len(self._objects)
95
+
96
+ def __getitem__(self, index):
97
+ box, mask = self._objects[index]
98
+ if mask is not None:
99
+ assert tuple(self._image.size[::-1]) == tuple(mask.shape), (index, self._image.size[::-1], mask.shape)
100
+ image = Image.composite(self._image, self._blank, Image.fromarray(mask))
101
+ else:
102
+ image = self._image
103
+ if args.options.get('crop', '1') == '1':
104
+ image = image.crop(box[:4])
105
+ # if args.save:
106
+ # base_count = len(os.listdir(args.save))
107
+ # image.save(os.path.join(args.save, f"cropped_{base_count:05}.png"))
108
+ return (transform(image), 0)
109
+
110
+
111
+ def color_classification(image, bboxes, classname):
112
+ if classname not in COLOR_CLASSIFIERS:
113
+ COLOR_CLASSIFIERS[classname] = zsc.zero_shot_classifier(
114
+ clip_model, tokenizer, COLORS,
115
+ [
116
+ f"a photo of a {{c}} {classname}",
117
+ f"a photo of a {{c}}-colored {classname}",
118
+ f"a photo of a {{c}} object"
119
+ ],
120
+ DEVICE
121
+ )
122
+ clf = COLOR_CLASSIFIERS[classname]
123
+ dataloader = torch.utils.data.DataLoader(
124
+ ImageCrops(image, bboxes),
125
+ batch_size=16, num_workers=4
126
+ )
127
+ with torch.no_grad():
128
+ pred, _ = zsc.run_classification(clip_model, clf, dataloader, DEVICE)
129
+ return [COLORS[index.item()] for index in pred.argmax(1)]
130
+
131
+
132
+ def compute_iou(box_a, box_b):
133
+ area_fn = lambda box: max(box[2] - box[0] + 1, 0) * max(box[3] - box[1] + 1, 0)
134
+ i_area = area_fn([
135
+ max(box_a[0], box_b[0]), max(box_a[1], box_b[1]),
136
+ min(box_a[2], box_b[2]), min(box_a[3], box_b[3])
137
+ ])
138
+ u_area = area_fn(box_a) + area_fn(box_b) - i_area
139
+ return i_area / u_area if u_area else 0
140
+
141
+
142
+ def relative_position(obj_a, obj_b):
143
+ """Give position of A relative to B, factoring in object dimensions"""
144
+ boxes = np.array([obj_a[0], obj_b[0]])[:, :4].reshape(2, 2, 2)
145
+ center_a, center_b = boxes.mean(axis=-2)
146
+ dim_a, dim_b = np.abs(np.diff(boxes, axis=-2))[..., 0, :]
147
+ offset = center_a - center_b
148
+ #
149
+ revised_offset = np.maximum(np.abs(offset) - POSITION_THRESHOLD * (dim_a + dim_b), 0) * np.sign(offset)
150
+ if np.all(np.abs(revised_offset) < 1e-3):
151
+ return set()
152
+ #
153
+ dx, dy = revised_offset / np.linalg.norm(offset)
154
+ relations = set()
155
+ if dx < -0.5: relations.add("left of")
156
+ if dx > 0.5: relations.add("right of")
157
+ if dy < -0.5: relations.add("above")
158
+ if dy > 0.5: relations.add("below")
159
+ return relations
160
+
161
+
162
+ def evaluate(image, objects, metadata):
163
+ """
164
+ Evaluate given image using detected objects on the global metadata specifications.
165
+ Assumptions:
166
+ * Metadata combines 'include' clauses with AND, and 'exclude' clauses with OR
167
+ * All clauses are independent, i.e., duplicating a clause has no effect on the correctness
168
+ * CHANGED: Color and position will only be evaluated on the most confidently predicted objects;
169
+ therefore, objects are expected to appear in sorted order
170
+ """
171
+ correct = True
172
+ reason = []
173
+ matched_groups = []
174
+ # Check for expected objects
175
+ for req in metadata.get('include', []):
176
+ classname = req['class']
177
+ matched = True
178
+ found_objects = objects.get(classname, [])[:req['count']]
179
+ if len(found_objects) < req['count']:
180
+ correct = matched = False
181
+ reason.append(f"expected {classname}>={req['count']}, found {len(found_objects)}")
182
+ else:
183
+ if 'color' in req:
184
+ # Color check
185
+ colors = color_classification(image, found_objects, classname)
186
+ if colors.count(req['color']) < req['count']:
187
+ correct = matched = False
188
+ reason.append(
189
+ f"expected {req['color']} {classname}>={req['count']}, found " +
190
+ f"{colors.count(req['color'])} {req['color']}; and " +
191
+ ", ".join(f"{colors.count(c)} {c}" for c in COLORS if c in colors)
192
+ )
193
+ if 'position' in req and matched:
194
+ # Relative position check
195
+ expected_rel, target_group = req['position']
196
+ if matched_groups[target_group] is None:
197
+ correct = matched = False
198
+ reason.append(f"no target for {classname} to be {expected_rel}")
199
+ else:
200
+ for obj in found_objects:
201
+ for target_obj in matched_groups[target_group]:
202
+ true_rels = relative_position(obj, target_obj)
203
+ if expected_rel not in true_rels:
204
+ correct = matched = False
205
+ reason.append(
206
+ f"expected {classname} {expected_rel} target, found " +
207
+ f"{' and '.join(true_rels)} target"
208
+ )
209
+ break
210
+ if not matched:
211
+ break
212
+ if matched:
213
+ matched_groups.append(found_objects)
214
+ else:
215
+ matched_groups.append(None)
216
+ # Check for non-expected objects
217
+ for req in metadata.get('exclude', []):
218
+ classname = req['class']
219
+ if len(objects.get(classname, [])) >= req['count']:
220
+ correct = False
221
+ reason.append(f"expected {classname}<{req['count']}, found {len(objects[classname])}")
222
+ return correct, "\n".join(reason)
223
+
224
+
225
+ def evaluate_image(filepath, metadata):
226
+ result = inference_detector(object_detector, filepath)
227
+ bbox = result[0] if isinstance(result, tuple) else result
228
+ segm = result[1] if isinstance(result, tuple) and len(result) > 1 else None
229
+ image = ImageOps.exif_transpose(Image.open(filepath))
230
+ detected = {}
231
+ # Determine bounding boxes to keep
232
+ confidence_threshold = THRESHOLD if metadata['tag'] != "counting" else COUNTING_THRESHOLD
233
+ for index, classname in enumerate(classnames):
234
+ ordering = np.argsort(bbox[index][:, 4])[::-1]
235
+ ordering = ordering[bbox[index][ordering, 4] > confidence_threshold] # Threshold
236
+ ordering = ordering[:MAX_OBJECTS].tolist() # Limit number of detected objects per class
237
+ detected[classname] = []
238
+ while ordering:
239
+ max_obj = ordering.pop(0)
240
+ detected[classname].append((bbox[index][max_obj], None if segm is None else segm[index][max_obj]))
241
+ ordering = [
242
+ obj for obj in ordering
243
+ if NMS_THRESHOLD == 1 or compute_iou(bbox[index][max_obj], bbox[index][obj]) < NMS_THRESHOLD
244
+ ]
245
+ if not detected[classname]:
246
+ del detected[classname]
247
+ # Evaluate
248
+ is_correct, reason = evaluate(image, detected, metadata)
249
+ return {
250
+ 'filename': filepath,
251
+ 'tag': metadata['tag'],
252
+ 'prompt': metadata['prompt'],
253
+ 'correct': is_correct,
254
+ 'reason': reason,
255
+ 'metadata': json.dumps(metadata),
256
+ 'details': json.dumps({
257
+ key: [box.tolist() for box, _ in value]
258
+ for key, value in detected.items()
259
+ })
260
+ }
261
+
262
+
263
+ def main(args):
264
+ full_results = []
265
+ pbar = tqdm.tqdm(total=len(os.listdir(args.imagedir)))
266
+ for subfolder in os.listdir(args.imagedir):
267
+ pbar.update(1)
268
+ folderpath = os.path.join(args.imagedir, subfolder)
269
+ if not os.path.isdir(folderpath) or not subfolder.isdigit():
270
+ print('skip 269')
271
+ continue
272
+ with open(os.path.join(folderpath, "metadata.jsonl")) as fp:
273
+ metadata = json.load(fp)
274
+ # Evaluate each image
275
+ for imagename in os.listdir(os.path.join(folderpath, "samples")):
276
+ imagepath = os.path.join(folderpath, "samples", imagename)
277
+ if not os.path.isfile(imagepath) or not (re.match(r"\d+\.png", imagename) or re.match(r"\d+\.jpg", imagename)):
278
+ print('skip 276')
279
+ continue
280
+ result = evaluate_image(imagepath, metadata)
281
+ full_results.append(result)
282
+ # Save results
283
+ if os.path.dirname(args.outfile):
284
+ os.makedirs(os.path.dirname(args.outfile), exist_ok=True)
285
+ with open(args.outfile, "w") as fp:
286
+ pd.DataFrame(full_results).to_json(fp, orient="records", lines=True)
287
+
288
+
289
+ if __name__ == "__main__":
290
+ args = parse_args()
291
+ object_detector, (clip_model, transform, tokenizer), classnames = load_models(args)
292
+ THRESHOLD = float(args.options.get('threshold', 0.3))
293
+ COUNTING_THRESHOLD = float(args.options.get('counting_threshold', 0.9))
294
+ MAX_OBJECTS = int(args.options.get('max_objects', 16))
295
+ NMS_THRESHOLD = float(args.options.get('max_overlap', 1.0))
296
+ POSITION_THRESHOLD = float(args.options.get('position_threshold', 0.1))
297
+
298
+ main(args)
evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py ADDED
@@ -0,0 +1,253 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ _base_ = [
2
+ '../_base_/datasets/coco_panoptic.py', '../_base_/default_runtime.py'
3
+ ]
4
+ num_things_classes = 80
5
+ num_stuff_classes = 53
6
+ num_classes = num_things_classes + num_stuff_classes
7
+ model = dict(
8
+ type='Mask2Former',
9
+ backbone=dict(
10
+ type='ResNet',
11
+ depth=50,
12
+ num_stages=4,
13
+ out_indices=(0, 1, 2, 3),
14
+ frozen_stages=-1,
15
+ norm_cfg=dict(type='BN', requires_grad=False),
16
+ norm_eval=True,
17
+ style='pytorch',
18
+ init_cfg=dict(type='Pretrained', checkpoint='torchvision://resnet50')),
19
+ panoptic_head=dict(
20
+ type='Mask2FormerHead',
21
+ in_channels=[256, 512, 1024, 2048], # pass to pixel_decoder inside
22
+ strides=[4, 8, 16, 32],
23
+ feat_channels=256,
24
+ out_channels=256,
25
+ num_things_classes=num_things_classes,
26
+ num_stuff_classes=num_stuff_classes,
27
+ num_queries=100,
28
+ num_transformer_feat_level=3,
29
+ pixel_decoder=dict(
30
+ type='MSDeformAttnPixelDecoder',
31
+ num_outs=3,
32
+ norm_cfg=dict(type='GN', num_groups=32),
33
+ act_cfg=dict(type='ReLU'),
34
+ encoder=dict(
35
+ type='DetrTransformerEncoder',
36
+ num_layers=6,
37
+ transformerlayers=dict(
38
+ type='BaseTransformerLayer',
39
+ attn_cfgs=dict(
40
+ type='MultiScaleDeformableAttention',
41
+ embed_dims=256,
42
+ num_heads=8,
43
+ num_levels=3,
44
+ num_points=4,
45
+ im2col_step=64,
46
+ dropout=0.0,
47
+ batch_first=False,
48
+ norm_cfg=None,
49
+ init_cfg=None),
50
+ ffn_cfgs=dict(
51
+ type='FFN',
52
+ embed_dims=256,
53
+ feedforward_channels=1024,
54
+ num_fcs=2,
55
+ ffn_drop=0.0,
56
+ act_cfg=dict(type='ReLU', inplace=True)),
57
+ operation_order=('self_attn', 'norm', 'ffn', 'norm')),
58
+ init_cfg=None),
59
+ positional_encoding=dict(
60
+ type='SinePositionalEncoding', num_feats=128, normalize=True),
61
+ init_cfg=None),
62
+ enforce_decoder_input_project=False,
63
+ positional_encoding=dict(
64
+ type='SinePositionalEncoding', num_feats=128, normalize=True),
65
+ transformer_decoder=dict(
66
+ type='DetrTransformerDecoder',
67
+ return_intermediate=True,
68
+ num_layers=9,
69
+ transformerlayers=dict(
70
+ type='DetrTransformerDecoderLayer',
71
+ attn_cfgs=dict(
72
+ type='MultiheadAttention',
73
+ embed_dims=256,
74
+ num_heads=8,
75
+ attn_drop=0.0,
76
+ proj_drop=0.0,
77
+ dropout_layer=None,
78
+ batch_first=False),
79
+ ffn_cfgs=dict(
80
+ embed_dims=256,
81
+ feedforward_channels=2048,
82
+ num_fcs=2,
83
+ act_cfg=dict(type='ReLU', inplace=True),
84
+ ffn_drop=0.0,
85
+ dropout_layer=None,
86
+ add_identity=True),
87
+ feedforward_channels=2048,
88
+ operation_order=('cross_attn', 'norm', 'self_attn', 'norm',
89
+ 'ffn', 'norm')),
90
+ init_cfg=None),
91
+ loss_cls=dict(
92
+ type='CrossEntropyLoss',
93
+ use_sigmoid=False,
94
+ loss_weight=2.0,
95
+ reduction='mean',
96
+ class_weight=[1.0] * num_classes + [0.1]),
97
+ loss_mask=dict(
98
+ type='CrossEntropyLoss',
99
+ use_sigmoid=True,
100
+ reduction='mean',
101
+ loss_weight=5.0),
102
+ loss_dice=dict(
103
+ type='DiceLoss',
104
+ use_sigmoid=True,
105
+ activate=True,
106
+ reduction='mean',
107
+ naive_dice=True,
108
+ eps=1.0,
109
+ loss_weight=5.0)),
110
+ panoptic_fusion_head=dict(
111
+ type='MaskFormerFusionHead',
112
+ num_things_classes=num_things_classes,
113
+ num_stuff_classes=num_stuff_classes,
114
+ loss_panoptic=None,
115
+ init_cfg=None),
116
+ train_cfg=dict(
117
+ num_points=12544,
118
+ oversample_ratio=3.0,
119
+ importance_sample_ratio=0.75,
120
+ assigner=dict(
121
+ type='MaskHungarianAssigner',
122
+ cls_cost=dict(type='ClassificationCost', weight=2.0),
123
+ mask_cost=dict(
124
+ type='CrossEntropyLossCost', weight=5.0, use_sigmoid=True),
125
+ dice_cost=dict(
126
+ type='DiceCost', weight=5.0, pred_act=True, eps=1.0)),
127
+ sampler=dict(type='MaskPseudoSampler')),
128
+ test_cfg=dict(
129
+ panoptic_on=True,
130
+ # For now, the dataset does not support
131
+ # evaluating semantic segmentation metric.
132
+ semantic_on=False,
133
+ instance_on=True,
134
+ # max_per_image is for instance segmentation.
135
+ max_per_image=100,
136
+ iou_thr=0.8,
137
+ # In Mask2Former's panoptic postprocessing,
138
+ # it will filter mask area where score is less than 0.5 .
139
+ filter_low_score=True),
140
+ init_cfg=None)
141
+
142
+ # dataset settings
143
+ image_size = (1024, 1024)
144
+ img_norm_cfg = dict(
145
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
146
+ train_pipeline = [
147
+ dict(type='LoadImageFromFile', to_float32=True),
148
+ dict(
149
+ type='LoadPanopticAnnotations',
150
+ with_bbox=True,
151
+ with_mask=True,
152
+ with_seg=True),
153
+ dict(type='RandomFlip', flip_ratio=0.5),
154
+ # large scale jittering
155
+ dict(
156
+ type='Resize',
157
+ img_scale=image_size,
158
+ ratio_range=(0.1, 2.0),
159
+ multiscale_mode='range',
160
+ keep_ratio=True),
161
+ dict(
162
+ type='RandomCrop',
163
+ crop_size=image_size,
164
+ crop_type='absolute',
165
+ recompute_bbox=True,
166
+ allow_negative_crop=True),
167
+ dict(type='Normalize', **img_norm_cfg),
168
+ dict(type='Pad', size=image_size),
169
+ dict(type='DefaultFormatBundle', img_to_float=True),
170
+ dict(
171
+ type='Collect',
172
+ keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks', 'gt_semantic_seg']),
173
+ ]
174
+ test_pipeline = [
175
+ dict(type='LoadImageFromFile'),
176
+ dict(
177
+ type='MultiScaleFlipAug',
178
+ img_scale=(1333, 800),
179
+ flip=False,
180
+ transforms=[
181
+ dict(type='Resize', keep_ratio=True),
182
+ dict(type='RandomFlip'),
183
+ dict(type='Normalize', **img_norm_cfg),
184
+ dict(type='Pad', size_divisor=32),
185
+ dict(type='ImageToTensor', keys=['img']),
186
+ dict(type='Collect', keys=['img']),
187
+ ])
188
+ ]
189
+ data_root = 'data/coco/'
190
+ data = dict(
191
+ samples_per_gpu=2,
192
+ workers_per_gpu=2,
193
+ train=dict(pipeline=train_pipeline),
194
+ val=dict(
195
+ pipeline=test_pipeline,
196
+ ins_ann_file=data_root + 'annotations/instances_val2017.json',
197
+ ),
198
+ test=dict(
199
+ pipeline=test_pipeline,
200
+ ins_ann_file=data_root + 'annotations/instances_val2017.json',
201
+ ))
202
+
203
+ embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
204
+ # optimizer
205
+ optimizer = dict(
206
+ type='AdamW',
207
+ lr=0.0001,
208
+ weight_decay=0.05,
209
+ eps=1e-8,
210
+ betas=(0.9, 0.999),
211
+ paramwise_cfg=dict(
212
+ custom_keys={
213
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
214
+ 'query_embed': embed_multi,
215
+ 'query_feat': embed_multi,
216
+ 'level_embed': embed_multi,
217
+ },
218
+ norm_decay_mult=0.0))
219
+ optimizer_config = dict(grad_clip=dict(max_norm=0.01, norm_type=2))
220
+
221
+ # learning policy
222
+ lr_config = dict(
223
+ policy='step',
224
+ gamma=0.1,
225
+ by_epoch=False,
226
+ step=[327778, 355092],
227
+ warmup='linear',
228
+ warmup_by_epoch=False,
229
+ warmup_ratio=1.0, # no warmup
230
+ warmup_iters=10)
231
+
232
+ max_iters = 368750
233
+ runner = dict(type='IterBasedRunner', max_iters=max_iters)
234
+
235
+ log_config = dict(
236
+ interval=50,
237
+ hooks=[
238
+ dict(type='TextLoggerHook', by_epoch=False),
239
+ dict(type='TensorboardLoggerHook', by_epoch=False)
240
+ ])
241
+ interval = 5000
242
+ workflow = [('train', interval)]
243
+ checkpoint_config = dict(
244
+ by_epoch=False, interval=interval, save_last=True, max_keep_ckpts=3)
245
+
246
+ # Before 365001th iteration, we do evaluation every 5000 iterations.
247
+ # After 365000th iteration, we do evaluation every 368750 iterations,
248
+ # which means that we do evaluation at the end of training.
249
+ dynamic_intervals = [(max_iters // interval * interval + 1, max_iters)]
250
+ evaluation = dict(
251
+ interval=interval,
252
+ dynamic_intervals=dynamic_intervals,
253
+ metric=['PQ', 'bbox', 'segm'])
evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ _base_ = ['./mask2former_r50_lsj_8x2_50e_coco-panoptic.py']
2
+ num_things_classes = 80
3
+ num_stuff_classes = 0
4
+ num_classes = num_things_classes + num_stuff_classes
5
+ model = dict(
6
+ panoptic_head=dict(
7
+ num_things_classes=num_things_classes,
8
+ num_stuff_classes=num_stuff_classes,
9
+ loss_cls=dict(class_weight=[1.0] * num_classes + [0.1])),
10
+ panoptic_fusion_head=dict(
11
+ num_things_classes=num_things_classes,
12
+ num_stuff_classes=num_stuff_classes),
13
+ test_cfg=dict(panoptic_on=False))
14
+
15
+ # dataset settings
16
+ image_size = (1024, 1024)
17
+ img_norm_cfg = dict(
18
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
19
+ pad_cfg = dict(img=(128, 128, 128), masks=0, seg=255)
20
+ train_pipeline = [
21
+ dict(type='LoadImageFromFile', to_float32=True),
22
+ dict(type='LoadAnnotations', with_bbox=True, with_mask=True),
23
+ dict(type='RandomFlip', flip_ratio=0.5),
24
+ # large scale jittering
25
+ dict(
26
+ type='Resize',
27
+ img_scale=image_size,
28
+ ratio_range=(0.1, 2.0),
29
+ multiscale_mode='range',
30
+ keep_ratio=True),
31
+ dict(
32
+ type='RandomCrop',
33
+ crop_size=image_size,
34
+ crop_type='absolute',
35
+ recompute_bbox=True,
36
+ allow_negative_crop=True),
37
+ dict(type='FilterAnnotations', min_gt_bbox_wh=(1e-5, 1e-5), by_mask=True),
38
+ dict(type='Pad', size=image_size, pad_val=pad_cfg),
39
+ dict(type='Normalize', **img_norm_cfg),
40
+ dict(type='DefaultFormatBundle', img_to_float=True),
41
+ dict(type='Collect', keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks']),
42
+ ]
43
+ test_pipeline = [
44
+ dict(type='LoadImageFromFile'),
45
+ dict(
46
+ type='MultiScaleFlipAug',
47
+ img_scale=(1333, 800),
48
+ flip=False,
49
+ transforms=[
50
+ dict(type='Resize', keep_ratio=True),
51
+ dict(type='RandomFlip'),
52
+ dict(type='Pad', size_divisor=32, pad_val=pad_cfg),
53
+ dict(type='Normalize', **img_norm_cfg),
54
+ dict(type='ImageToTensor', keys=['img']),
55
+ dict(type='Collect', keys=['img']),
56
+ ])
57
+ ]
58
+ dataset_type = 'CocoDataset'
59
+ data_root = 'data/coco/'
60
+ data = dict(
61
+ _delete_=True,
62
+ samples_per_gpu=2,
63
+ workers_per_gpu=2,
64
+ train=dict(
65
+ type=dataset_type,
66
+ ann_file=data_root + 'annotations/instances_train2017.json',
67
+ img_prefix=data_root + 'train2017/',
68
+ pipeline=train_pipeline),
69
+ val=dict(
70
+ type=dataset_type,
71
+ ann_file=data_root + 'annotations/instances_val2017.json',
72
+ img_prefix=data_root + 'val2017/',
73
+ pipeline=test_pipeline),
74
+ test=dict(
75
+ type=dataset_type,
76
+ ann_file=data_root + 'annotations/instances_val2017.json',
77
+ img_prefix=data_root + 'val2017/',
78
+ pipeline=test_pipeline))
79
+ evaluation = dict(metric=['bbox', 'segm'])
evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ _base_ = ['./mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py']
2
+ pretrained = 'https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_small_patch4_window7_224.pth' # noqa
3
+
4
+ depths = [2, 2, 18, 2]
5
+ model = dict(
6
+ backbone=dict(
7
+ depths=depths, init_cfg=dict(type='Pretrained',
8
+ checkpoint=pretrained)))
9
+
10
+ # set all layers in backbone to lr_mult=0.1
11
+ # set all norm layers, position_embeding,
12
+ # query_embeding, level_embeding to decay_multi=0.0
13
+ backbone_norm_multi = dict(lr_mult=0.1, decay_mult=0.0)
14
+ backbone_embed_multi = dict(lr_mult=0.1, decay_mult=0.0)
15
+ embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
16
+ custom_keys = {
17
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
18
+ 'backbone.patch_embed.norm': backbone_norm_multi,
19
+ 'backbone.norm': backbone_norm_multi,
20
+ 'absolute_pos_embed': backbone_embed_multi,
21
+ 'relative_position_bias_table': backbone_embed_multi,
22
+ 'query_embed': embed_multi,
23
+ 'query_feat': embed_multi,
24
+ 'level_embed': embed_multi
25
+ }
26
+ custom_keys.update({
27
+ f'backbone.stages.{stage_id}.blocks.{block_id}.norm': backbone_norm_multi
28
+ for stage_id, num_blocks in enumerate(depths)
29
+ for block_id in range(num_blocks)
30
+ })
31
+ custom_keys.update({
32
+ f'backbone.stages.{stage_id}.downsample.norm': backbone_norm_multi
33
+ for stage_id in range(len(depths) - 1)
34
+ })
35
+ # optimizer
36
+ optimizer = dict(
37
+ paramwise_cfg=dict(custom_keys=custom_keys, norm_decay_mult=0.0))
evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ _base_ = ['./mask2former_r50_lsj_8x2_50e_coco.py']
2
+ pretrained = 'https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_tiny_patch4_window7_224.pth' # noqa
3
+ depths = [2, 2, 6, 2]
4
+ model = dict(
5
+ type='Mask2Former',
6
+ backbone=dict(
7
+ _delete_=True,
8
+ type='SwinTransformer',
9
+ embed_dims=96,
10
+ depths=depths,
11
+ num_heads=[3, 6, 12, 24],
12
+ window_size=7,
13
+ mlp_ratio=4,
14
+ qkv_bias=True,
15
+ qk_scale=None,
16
+ drop_rate=0.,
17
+ attn_drop_rate=0.,
18
+ drop_path_rate=0.3,
19
+ patch_norm=True,
20
+ out_indices=(0, 1, 2, 3),
21
+ with_cp=False,
22
+ convert_weights=True,
23
+ frozen_stages=-1,
24
+ init_cfg=dict(type='Pretrained', checkpoint=pretrained)),
25
+ panoptic_head=dict(
26
+ type='Mask2FormerHead', in_channels=[96, 192, 384, 768]),
27
+ init_cfg=None)
28
+
29
+ # set all layers in backbone to lr_mult=0.1
30
+ # set all norm layers, position_embeding,
31
+ # query_embeding, level_embeding to decay_multi=0.0
32
+ backbone_norm_multi = dict(lr_mult=0.1, decay_mult=0.0)
33
+ backbone_embed_multi = dict(lr_mult=0.1, decay_mult=0.0)
34
+ embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
35
+ custom_keys = {
36
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
37
+ 'backbone.patch_embed.norm': backbone_norm_multi,
38
+ 'backbone.norm': backbone_norm_multi,
39
+ 'absolute_pos_embed': backbone_embed_multi,
40
+ 'relative_position_bias_table': backbone_embed_multi,
41
+ 'query_embed': embed_multi,
42
+ 'query_feat': embed_multi,
43
+ 'level_embed': embed_multi
44
+ }
45
+ custom_keys.update({
46
+ f'backbone.stages.{stage_id}.blocks.{block_id}.norm': backbone_norm_multi
47
+ for stage_id, num_blocks in enumerate(depths)
48
+ for block_id in range(num_blocks)
49
+ })
50
+ custom_keys.update({
51
+ f'backbone.stages.{stage_id}.downsample.norm': backbone_norm_multi
52
+ for stage_id in range(len(depths) - 1)
53
+ })
54
+ # optimizer
55
+ optimizer = dict(
56
+ type='AdamW',
57
+ lr=0.0001,
58
+ weight_decay=0.05,
59
+ eps=1e-8,
60
+ betas=(0.9, 0.999),
61
+ paramwise_cfg=dict(custom_keys=custom_keys, norm_decay_mult=0.0))
evaluation/gen_eval/prompts/create_prompts.py ADDED
@@ -0,0 +1,183 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Generate prompts for evaluation
3
+ """
4
+
5
+ import argparse
6
+ import json
7
+ import os
8
+ import yaml
9
+
10
+ import numpy as np
11
+
12
+ # Load classnames
13
+
14
+ with open("object_names.txt") as cls_file:
15
+ classnames = [line.strip() for line in cls_file]
16
+
17
+ # Proper a vs an
18
+
19
+ def with_article(name: str):
20
+ if name[0] in "aeiou":
21
+ return f"an {name}"
22
+ return f"a {name}"
23
+
24
+ # Proper plural
25
+
26
+ def make_plural(name: str):
27
+ if name[-1] in "s":
28
+ return f"{name}es"
29
+ return f"{name}s"
30
+
31
+ # Generates single object samples
32
+
33
+ def generate_single_object_sample(rng: np.random.Generator, size: int = None):
34
+ TAG = "single_object"
35
+ if size > len(classnames):
36
+ size = len(classnames)
37
+ print(f"Not enough distinct classes, generating only {size} samples")
38
+ return_scalar = size is None
39
+ size = size or 1
40
+ idxs = rng.choice(len(classnames), size=size, replace=False)
41
+ samples = [dict(
42
+ tag=TAG,
43
+ include=[
44
+ {"class": classnames[idx], "count": 1}
45
+ ],
46
+ prompt=f"a photo of {with_article(classnames[idx])}"
47
+ ) for idx in idxs]
48
+ if return_scalar:
49
+ return samples[0]
50
+ return samples
51
+
52
+ # Generate two object samples
53
+
54
+ def generate_two_object_sample(rng: np.random.Generator):
55
+ TAG = "two_object"
56
+ idx_a, idx_b = rng.choice(len(classnames), size=2, replace=False)
57
+ return dict(
58
+ tag=TAG,
59
+ include=[
60
+ {"class": classnames[idx_a], "count": 1},
61
+ {"class": classnames[idx_b], "count": 1}
62
+ ],
63
+ prompt=f"a photo of {with_article(classnames[idx_a])} and {with_article(classnames[idx_b])}"
64
+ )
65
+
66
+ # Generate counting samples
67
+
68
+ numbers = ["zero", "one", "two", "three", "four", "five", "six", "seven", "eight", "nine", "ten"]
69
+
70
+ def generate_counting_sample(rng: np.random.Generator, max_count=4):
71
+ TAG = "counting"
72
+ idx = rng.choice(len(classnames))
73
+ num = int(rng.integers(2, max_count, endpoint=True))
74
+ return dict(
75
+ tag=TAG,
76
+ include=[
77
+ {"class": classnames[idx], "count": num}
78
+ ],
79
+ exclude=[
80
+ {"class": classnames[idx], "count": num + 1}
81
+ ],
82
+ prompt=f"a photo of {numbers[num]} {make_plural(classnames[idx])}"
83
+ )
84
+
85
+ # Generate color samples
86
+
87
+ colors = ["red", "orange", "yellow", "green", "blue", "purple", "pink", "brown", "black", "white"]
88
+
89
+ def generate_color_sample(rng: np.random.Generator):
90
+ TAG = "colors"
91
+ idx = rng.choice(len(classnames) - 1) + 1
92
+ idx = (idx + classnames.index("person")) % len(classnames) # No "[COLOR] person" prompts
93
+ color = colors[rng.choice(len(colors))]
94
+ return dict(
95
+ tag=TAG,
96
+ include=[
97
+ {"class": classnames[idx], "count": 1, "color": color}
98
+ ],
99
+ prompt=f"a photo of {with_article(color)} {classnames[idx]}"
100
+ )
101
+
102
+ # Generate position samples
103
+
104
+ positions = ["left of", "right of", "above", "below"]
105
+
106
+ def generate_position_sample(rng: np.random.Generator):
107
+ TAG = "position"
108
+ idx_a, idx_b = rng.choice(len(classnames), size=2, replace=False)
109
+ position = positions[rng.choice(len(positions))]
110
+ return dict(
111
+ tag=TAG,
112
+ include=[
113
+ {"class": classnames[idx_b], "count": 1},
114
+ {"class": classnames[idx_a], "count": 1, "position": (position, 0)}
115
+ ],
116
+ prompt=f"a photo of {with_article(classnames[idx_a])} {position} {with_article(classnames[idx_b])}"
117
+ )
118
+
119
+ # Generate color attribution samples
120
+
121
+ def generate_color_attribution_sample(rng: np.random.Generator):
122
+ TAG = "color_attr"
123
+ idxs = rng.choice(len(classnames) - 1, size=2, replace=False) + 1
124
+ idx_a, idx_b = (idxs + classnames.index("person")) % len(classnames) # No "[COLOR] person" prompts
125
+ cidx_a, cidx_b = rng.choice(len(colors), size=2, replace=False)
126
+ return dict(
127
+ tag=TAG,
128
+ include=[
129
+ {"class": classnames[idx_a], "count": 1, "color": colors[cidx_a]},
130
+ {"class": classnames[idx_b], "count": 1, "color": colors[cidx_b]}
131
+ ],
132
+ prompt=f"a photo of {with_article(colors[cidx_a])} {classnames[idx_a]} and {with_article(colors[cidx_b])} {classnames[idx_b]}"
133
+ )
134
+
135
+
136
+ # Generate evaluation suite
137
+
138
+ def generate_suite(rng: np.random.Generator, n: int = 100, output_path: str = ""):
139
+ samples = []
140
+ # Generate single object samples for all COCO classnames
141
+ samples.extend(generate_single_object_sample(rng, size=len(classnames)))
142
+ # Generate two object samples (~100)
143
+ for _ in range(n):
144
+ samples.append(generate_two_object_sample(rng))
145
+ # Generate counting samples
146
+ for _ in range(n):
147
+ samples.append(generate_counting_sample(rng, max_count=4))
148
+ # Generate color samples
149
+ for _ in range(n):
150
+ samples.append(generate_color_sample(rng))
151
+ # Generate position samples
152
+ for _ in range(n):
153
+ samples.append(generate_position_sample(rng))
154
+ # Generate color attribution samples
155
+ for _ in range(n):
156
+ samples.append(generate_color_attribution_sample(rng))
157
+ # De-duplicate
158
+ unique_samples, used_samples = [], set()
159
+ for sample in samples:
160
+ sample_text = yaml.safe_dump(sample)
161
+ if sample_text not in used_samples:
162
+ unique_samples.append(sample)
163
+ used_samples.add(sample_text)
164
+
165
+ # Write to files
166
+ os.makedirs(output_path, exist_ok=True)
167
+ with open(os.path.join(output_path, "generation_prompts.txt"), "w") as fp:
168
+ for sample in unique_samples:
169
+ print(sample['prompt'], file=fp)
170
+ with open(os.path.join(output_path, "evaluation_metadata.jsonl"), "w") as fp:
171
+ for sample in unique_samples:
172
+ print(json.dumps(sample), file=fp)
173
+
174
+
175
+ if __name__ == "__main__":
176
+ parser = argparse.ArgumentParser()
177
+ parser.add_argument("--seed", type=int, default=43, help="generation seed (default: 43)")
178
+ parser.add_argument("--num-prompts", "-n", type=int, default=100, help="number of prompts per task (default: 100)")
179
+ parser.add_argument("--output-path", "-o", type=str, default="prompts", help="output folder for prompts and metadata (default: 'prompts/')")
180
+ args = parser.parse_args()
181
+ rng = np.random.default_rng(args.seed)
182
+ generate_suite(rng, args.num_prompts, args.output_path)
183
+
evaluation/gen_eval/summary_scores.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Get results of evaluation
2
+
3
+ import argparse
4
+ import os
5
+
6
+ import numpy as np
7
+ import pandas as pd
8
+
9
+
10
+ parser = argparse.ArgumentParser()
11
+ parser.add_argument("filename", type=str)
12
+ args = parser.parse_args()
13
+
14
+ # Load classnames
15
+
16
+ with open(os.path.join(os.path.dirname(__file__), "object_names.txt")) as cls_file:
17
+ classnames = [line.strip() for line in cls_file]
18
+ cls_to_idx = {"_".join(cls.split()):idx for idx, cls in enumerate(classnames)}
19
+
20
+ # Load results
21
+
22
+ df = pd.read_json(args.filename, orient="records", lines=True)
23
+
24
+ # Measure overall success
25
+
26
+ print("Summary")
27
+ print("=======")
28
+ print(f"Total images: {len(df)}")
29
+ print(f"Total prompts: {len(df.groupby('metadata'))}")
30
+ print(f"% correct images: {df['correct'].mean():.2%}")
31
+ print(f"% correct prompts: {df.groupby('metadata')['correct'].any().mean():.2%}")
32
+ print()
33
+
34
+ # By group
35
+
36
+ task_scores = []
37
+
38
+ print("Task breakdown")
39
+ print("==============")
40
+ for tag, task_df in df.groupby('tag', sort=False):
41
+ task_scores.append(task_df['correct'].mean())
42
+ print(f"{tag:<16} = {task_df['correct'].mean():.2%} ({task_df['correct'].sum()} / {len(task_df)})")
43
+ print()
44
+
45
+ print(f"Overall score (avg. over tasks): {np.mean(task_scores):.5f}")
grn/__init__.py ADDED
File without changes
grn/dataset/build.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import datetime
2
+ import os
3
+ import os.path as osp
4
+ import random
5
+ import subprocess
6
+ from functools import partial
7
+ from typing import Optional
8
+ import time
9
+
10
+ import pytz
11
+
12
+ from grn.dataset.dataset_joint_vi import JointViDataset
13
+ from grn.utils_t2iv.sequence_parallel import SequenceParallelManager as sp_manager
14
+
15
+ try:
16
+ from grp import getgrgid
17
+ from pwd import getpwuid
18
+ except:
19
+ pass
20
+ import PIL.Image as PImage
21
+ from PIL import ImageFile
22
+ import numpy as np
23
+ from torchvision.transforms import transforms
24
+ from torchvision.transforms.functional import resize, to_tensor
25
+ import torch.distributed as tdist
26
+
27
+ from torchvision.transforms import InterpolationMode
28
+ bicubic = InterpolationMode.BICUBIC
29
+ lanczos = InterpolationMode.LANCZOS
30
+ PImage.MAX_IMAGE_PIXELS = (1024 * 1024 * 1024 // 4 // 3) * 5
31
+ ImageFile.LOAD_TRUNCATED_IMAGES = False
32
+
33
+
34
+ def time_str(fmt='[%m-%d %H:%M:%S]'):
35
+ return datetime.datetime.now(tz=pytz.timezone('Asia/Shanghai')).strftime(fmt)
36
+
37
+
38
+ def normalize_01_into_pm1(x): # normalize x from [0, 1] to [-1, 1] by (x*2) - 1
39
+ return x.add(x).add_(-1)
40
+
41
+
42
+ def denormalize_pm1_into_01(x): # denormalize x from [-1, 1] to [0, 1]
43
+ return x.add(1).mul_(0.5)
44
+
45
+
46
+ def center_crop_arr(pil_image, image_size):
47
+ """
48
+ Center cropping implementation from ADM.
49
+ https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
50
+ """
51
+ while min(*pil_image.size) >= 2 * image_size:
52
+ pil_image = pil_image.resize(
53
+ tuple(x // 2 for x in pil_image.size), resample=PImage.BOX
54
+ )
55
+
56
+ scale = image_size / min(*pil_image.size)
57
+ pil_image = pil_image.resize(
58
+ tuple(round(x * scale) for x in pil_image.size), resample=PImage.LANCZOS
59
+ )
60
+
61
+ arr = np.array(pil_image)
62
+ crop_y = (arr.shape[0] - image_size) // 2
63
+ crop_x = (arr.shape[1] - image_size) // 2
64
+ return PImage.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size])
65
+
66
+
67
+ class RandomResize:
68
+ def __init__(self, mid_reso, final_reso, interpolation):
69
+ ub = max(round((mid_reso + (mid_reso-final_reso) / 8) / 4) * 4, mid_reso)
70
+ self.reso_lb, self.reso_ub = final_reso, ub
71
+ self.interpolation = interpolation
72
+
73
+ def __call__(self, img):
74
+ return resize(img, size=random.randint(self.reso_lb, self.reso_ub), interpolation=self.interpolation)
75
+
76
+ def __repr__(self):
77
+ return f'RandomResize(reso=({self.reso_lb}, {self.reso_ub}), interpolation={self.interpolation})'
78
+
79
+
80
+ def print_aug(transform, label):
81
+ print(f'Transform {label} = ')
82
+ if hasattr(transform, 'transforms'):
83
+ for t in transform.transforms:
84
+ print(t)
85
+ else:
86
+ print(transform)
87
+ print('---------------------------\n')
88
+
89
+
90
+ def build_joint_dataset(
91
+ args,
92
+ meta_folders: str,
93
+ meta_folder_repeats: str,
94
+ max_caption_len: int,
95
+ short_prob=0.2,
96
+ load_vae_instead_of_image=False
97
+ ):
98
+ return JointViDataset(
99
+ meta_folders=meta_folders,
100
+ meta_folder_repeats=meta_folder_repeats,
101
+ max_caption_len=max_caption_len,
102
+ short_prob=short_prob,
103
+ load_vae_instead_of_image=load_vae_instead_of_image,
104
+ video_fps=args.video_fps,
105
+ num_frames=args.video_frames,
106
+ online_t5=args.online_t5,
107
+ num_replicas=sp_manager.get_sp_group_nums() if sp_manager.sp_on() else tdist.get_world_size(), # 1,
108
+ rank = sp_manager.get_sp_group_rank() if sp_manager.sp_on() else tdist.get_rank(),
109
+ dataloader_workers=args.workers,
110
+ enable_dynamic_length_prompt=args.enable_dynamic_length_prompt,
111
+ hdfs_mode=args.hdfs_mode,
112
+ dynamic_scale_schedule=args.dynamic_scale_schedule,
113
+ seed=args.seed,
114
+ other_args=args,
115
+ )
116
+
117
+
118
+ def pil_load(path: str, proposal_size):
119
+ with open(path, 'rb') as f:
120
+ img: PImage.Image = PImage.open(f)
121
+ w: int = img.width
122
+ h: int = img.height
123
+ sh: int = min(h, w)
124
+ if sh > proposal_size:
125
+ ratio: float = proposal_size / sh
126
+ w = round(ratio * w)
127
+ h = round(ratio * h)
128
+ img.draft('RGB', (w, h))
129
+ img = img.convert('RGB')
130
+ return img
131
+
132
+
133
+ def rewrite(im: PImage, file: str, info: str):
134
+ kw = dict(quality=100)
135
+ if file.lower().endswith('.tif') or file.lower().endswith('.tiff'):
136
+ kw['compression'] = 'none'
137
+ elif file.lower().endswith('.webp'):
138
+ kw['lossless'] = True
139
+
140
+ st = os.stat(file)
141
+ uname = getpwuid(st.st_uid).pw_name
142
+ gname = getgrgid(st.st_gid).gr_name
143
+ mode = oct(st.st_mode)[-3:]
144
+
145
+ local_file = osp.basename(file)
146
+ im.save(local_file, **kw)
147
+ print(f'************* <REWRITE: {info}> ************* @ {file}')
148
+ subprocess.call(['sudo', 'mv', local_file, file])
149
+ subprocess.call(['sudo', 'chown', f'{uname}:{gname}', file])
150
+ subprocess.call(['sudo', 'chmod', str(mode), file])
grn/dataset/dataset_joint_vi.py ADDED
@@ -0,0 +1,687 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import glob
2
+ import os
3
+ import pickle
4
+ import random
5
+ import re
6
+ import time
7
+ from functools import partial
8
+ from os import path as osp
9
+ from typing import List, Tuple, Union
10
+ import json
11
+ import itertools
12
+ import hashlib
13
+ import copy
14
+ import collections
15
+ import math
16
+
17
+ import tqdm
18
+ import numpy as np
19
+ import torch
20
+ import pandas as pd
21
+ from decord import VideoReader
22
+ from PIL import Image as PImage
23
+ from torch.nn import functional as F
24
+ from torchvision.transforms.functional import to_tensor, hflip
25
+ from torchvision.transforms import transforms, InterpolationMode
26
+ from torch.utils.data import Dataset, DataLoader
27
+ import torch.distributed as tdist
28
+ from PIL import Image
29
+ os.environ["TOKENIZERS_PARALLELISM"] = "false"
30
+
31
+ from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta
32
+ from grn.utils.video_decoder import EncodedVideoDecord
33
+ from grn.utils.compress_tokens import load_packed_tensor
34
+ from transformers import AutoTokenizer
35
+
36
+ def transform(pil_img, tgt_h, tgt_w):
37
+ width, height = pil_img.size
38
+ if width / height <= tgt_w / tgt_h:
39
+ resized_width = tgt_w
40
+ resized_height = int(tgt_w / (width / height))
41
+ else:
42
+ resized_height = tgt_h
43
+ resized_width = int((width / height) * tgt_h)
44
+ pil_img = pil_img.resize((resized_width, resized_height), resample=PImage.LANCZOS)
45
+ # crop the center out
46
+ arr = np.array(pil_img)
47
+ crop_y = (arr.shape[0] - tgt_h) // 2
48
+ crop_x = (arr.shape[1] - tgt_w) // 2
49
+ im = to_tensor(arr[crop_y: crop_y + tgt_h, crop_x: crop_x + tgt_w])
50
+ return im.add(im).add_(-1)
51
+
52
+ def normalize(x): # normalize x from [0, 1] to [-1, 1] by (x*2) - 1
53
+ return x.add(x).add_(-1)
54
+
55
+ def get_prompt_id(prompt):
56
+ md5 = hashlib.md5()
57
+ md5.update(prompt.encode('utf-8'))
58
+ prompt_id = md5.hexdigest()
59
+ return prompt_id
60
+
61
+ def prepend_motion_score(prompt, motion_score):
62
+ return f'<<<motion_score: {round(motion_score):.1f}>>> {prompt}'
63
+
64
+ class VideoReaderWrapper(VideoReader):
65
+ def __init__(self, *args, **kwargs):
66
+ super().__init__(*args, **kwargs)
67
+ self.seek(0)
68
+ def __getitem__(self, key):
69
+ frames = super().__getitem__(key)
70
+ self.seek(0)
71
+ return frames
72
+
73
+ class JointViDataset(Dataset):
74
+ def __init__(
75
+ self,
76
+ meta_folders: str = '',
77
+ meta_folder_repeats: str = '',
78
+ buffersize: int = 1000000 * 300,
79
+ seed: int = 0,
80
+ pn: str = '',
81
+ video_fps: int = 1,
82
+ num_replicas: int = 1,
83
+ rank: int = 0,
84
+ dataloader_workers: int = 2,
85
+ enable_dynamic_length_prompt: bool = True,
86
+ shuffle: bool = True,
87
+ short_prob: float = 0.2,
88
+ verbose=False,
89
+ temp_dir= "/dev/shm",
90
+ hdfs_mode='read',
91
+ other_args=None,
92
+ **kwargs,
93
+ ):
94
+ self.meta_folders = json.loads(meta_folders)
95
+ self.meta_folder_repeats = json.loads(meta_folder_repeats)
96
+ self.meta_folder_identifiers = json.loads(other_args.meta_folder_identifiers)
97
+ self.pn_list = json.loads(other_args.pn_list)
98
+ self.pn_probs = json.loads(other_args.pn_probs)
99
+ assert len(self.meta_folders) == len(self.meta_folder_repeats), f'{len(self.meta_folders)} != {len(self.meta_folder_repeats)}'
100
+ assert len(self.meta_folders) == len(self.meta_folder_identifiers), f'{len(self.meta_folders)} != {len(self.meta_folder_identifiers)}'
101
+ self.verbose = verbose
102
+ self.buffer_size = buffersize
103
+ self.num_replicas = num_replicas
104
+ self.rank = rank
105
+ self.worker_id = 0
106
+ self.global_worker_id = 0
107
+ self.short_prob = short_prob
108
+ self.dataloader_workers = max(1, dataloader_workers)
109
+ self.shuffle = shuffle
110
+ self.global_workers = self.num_replicas * self.dataloader_workers
111
+ self.seed = seed
112
+ self.text_tokenizer = other_args.text_tokenizer
113
+ self.feature_extraction = other_args.only_images4extract_feats # < 0 # no sequence packing, for feature extraction
114
+ self.epoch_generator = None
115
+ self.epoch_rank_generator = None
116
+ self.other_args = other_args
117
+ self.pair_input = other_args.pair_input
118
+ self.drop_long_video = other_args.drop_long_video
119
+ self.enable_dynamic_length_prompt = enable_dynamic_length_prompt
120
+ self.set_epoch_generator(other_args.epoch)
121
+ self.temporal_compress_rate = other_args.temporal_compress_rate
122
+ self.dynamic_resolution_h_w, self.h_div_w_templates = get_dynamic_resolution_meta(other_args.dynamic_scale_schedule, other_args.train_h_div_w_list, other_args.video_frames) # here video_frames is the max video frames
123
+ self.video_fps = video_fps
124
+ self.min_training_duration = (other_args.min_video_frames-1) // self.video_fps
125
+ self.max_training_duration = (other_args.video_frames-1) // self.video_fps
126
+ self.c2i = self.other_args.add_class_token > 0
127
+ if self.c2i:
128
+ self.c2i_transform = transforms.Compose([
129
+ transforms.RandomHorizontalFlip(),
130
+ transforms.Resize(288, interpolation=InterpolationMode.LANCZOS), # transforms.Resize: resize the shorter edge to mid_reso
131
+ transforms.RandomCrop((256, 256)),
132
+ transforms.ToTensor(),
133
+ normalize,
134
+ ])
135
+ else:
136
+ self.c2i_transform = None
137
+ self.print(f"{self.rank=} dataset {self.seed=}, {self.h_div_w_templates=} {self.min_training_duration=} {self.max_training_duration=}, cache_check_mode={self.other_args.cache_check_mode}")
138
+ self.token_cache_dir = other_args.token_cache_dir
139
+ self.use_vae_token_cache = other_args.use_vae_token_cache
140
+ self.allow_online_vae_feature_extraction = other_args.allow_online_vae_feature_extraction
141
+ self.use_text_token_cache = other_args.use_text_token_cache
142
+ self.max_video_frames = other_args.video_frames
143
+ self.cached_video_frames = other_args.cached_video_frames # cached max video frames
144
+ self.down_size_limit = other_args.down_size_limit
145
+ self.video_caption_type = other_args.video_caption_type
146
+ self.train_max_token_len = other_args.train_max_token_len
147
+ self.duration_resolution = other_args.duration_resolution
148
+ self.device = other_args.device
149
+ print(f'self.down_size_limit: {self.down_size_limit}')
150
+ self.hdfs_mode = hdfs_mode
151
+ self.max_text_len = other_args.tlen
152
+ self.temp_dir = temp_dir.rstrip("/")
153
+ self.mapped_duration2metas, self.mapped_duration2freqs = self.get_mapped_duration2metas()
154
+ self.batches = self.form_batches(self.mapped_duration2metas)
155
+ print(f'{num_replicas=}, {rank=}, {dataloader_workers=}, {len(self.batches)=}, {self.drop_long_video=} {self.max_text_len=} self.batches[:10]={self.batches[:10]}')
156
+
157
+ def print(self, string):
158
+ if self.feature_extraction:
159
+ print(string)
160
+ else:
161
+ print(string, force=True)
162
+
163
+ def get_captions_lens(self, captions):
164
+ if self.other_args.text_tokenizer_type == 'flan_t5':
165
+ tokens = self.other_args.text_tokenizer(text=captions, max_length=self.other_args.text_tokenizer.model_max_length, padding='max_length', truncation=True, return_tensors='pt')
166
+ mask = tokens.attention_mask.cuda(non_blocking=True)
167
+ lens: List[int] = mask.sum(dim=-1).tolist()
168
+ else: # umt5-xxl
169
+ ids, mask = self.other_args.text_tokenizer( captions, return_mask=True, add_special_tokens=True)
170
+ lens = mask.gt(0).sum(dim=1).tolist()
171
+ return lens
172
+
173
+ def get_video_caption(self, meta, mapped_duration):
174
+ if 'tarsier2_caption' not in meta:
175
+ caption = self.epoch_rank_generator.choice(meta['caption'])['content']
176
+ else:
177
+ caption_type = 'tarsier2_caption'
178
+ if ('MiniCPM_V_2_6_caption' in meta) and meta['MiniCPM_V_2_6_caption']:
179
+ caption_type = self.epoch_rank_generator.choice(['tarsier2_caption', 'MiniCPM_V_2_6_caption'])
180
+ caption = meta[caption_type]
181
+ if self.enable_dynamic_length_prompt and (self.epoch_rank_generator.random() < self.other_args.short_cap_prob):
182
+ caption = self.random_drop_sentences(caption, min_sentences=2)
183
+ if 'quality_prompt' in meta:
184
+ caption = caption + ' ' + meta['quality_prompt']
185
+ if meta['first_frame_condition']:
186
+ caption = '<I2V>' + caption
187
+ else:
188
+ caption = '<T2V>' + caption
189
+ assert caption
190
+ return caption
191
+
192
+ def get_image_caption(self, meta):
193
+ caption = meta['long_caption']
194
+ if not meta['long_caption']:
195
+ caption = meta['text']
196
+ else:
197
+ if self.epoch_rank_generator.random() < self.other_args.short_cap_prob:
198
+ if meta['text']:
199
+ caption = meta['text']
200
+ elif ('InternVL' in meta['long_caption_type']):
201
+ caption = self.random_drop_sentences(meta['long_caption'], min_sentences=2)
202
+ caption = '<T2I>' + caption
203
+ assert caption
204
+ return caption
205
+
206
+ def sample_pn(self, meta, pn_list, pn_probs):
207
+ if ('height' in meta) and ('width' in meta):
208
+ real_pn = meta['height'] * meta['width'] / 1000000
209
+ valid_pn_list, valid_pn_probs = [], []
210
+ for pn, pn_prob in zip(pn_list, pn_probs):
211
+ if real_pn > 0.7 * float(pn[:-1]):
212
+ valid_pn_list.append(pn)
213
+ valid_pn_probs.append(pn_prob)
214
+ if not len(valid_pn_list):
215
+ valid_pn_list = pn_list[:1]
216
+ valid_pn_probs = pn_probs[:1]
217
+ else:
218
+ valid_pn_list = pn_list
219
+ valid_pn_probs = pn_probs
220
+ return self.epoch_rank_generator.choice(valid_pn_list, p=valid_pn_probs)
221
+
222
+ def get_mapped_duration2metas(self):
223
+ part_filepaths = []
224
+ part_file2identifier = {}
225
+ for meta_folder, meta_folder_repeat, meta_folder_identifier in zip(self.meta_folders, self.meta_folder_repeats, self.meta_folder_identifiers):
226
+ tmp_part_filepaths = sorted(glob.glob(osp.join(meta_folder, '*/*.jsonl')))
227
+ self.epoch_generator.shuffle(tmp_part_filepaths)
228
+ if meta_folder_repeat > 1:
229
+ tmp_part_filepaths = tmp_part_filepaths * int(np.ceil(meta_folder_repeat))
230
+ tmp_part_filepaths = tmp_part_filepaths[:int(len(tmp_part_filepaths)*meta_folder_repeat)]
231
+ for file in tmp_part_filepaths: part_file2identifier[file] = meta_folder_identifier
232
+ part_filepaths.extend(tmp_part_filepaths)
233
+ self.epoch_generator.shuffle(part_filepaths)
234
+ self.print(f'{self.rank=} jsonls sample: {part_filepaths[:4]}')
235
+ if self.num_replicas > 1:
236
+ part_filepaths = part_filepaths[self.rank::self.num_replicas]
237
+
238
+ mapped_duration2metas = {}
239
+ pbar = tqdm.tqdm(total=len(part_filepaths))
240
+ total, corrupt = 0, 0
241
+ stop_read = False
242
+ rough_h_div_w = self.h_div_w_templates[np.argmin(np.abs((9/16-self.h_div_w_templates)))]
243
+ for part_filepath in part_filepaths:
244
+ file_quality_prompt = part_file2identifier[part_filepath]
245
+ if stop_read:
246
+ break
247
+ pbar.update(1)
248
+ try:
249
+ with open(part_filepath, encoding='utf-8') as f:
250
+ lines = f.readlines()
251
+ except Exception as e:
252
+ print(f'{part_filepath=} Error: {e}')
253
+ lines = []
254
+ for line in lines:
255
+ total += 1
256
+ try:
257
+ meta = json.loads(line)
258
+ except Exception as e:
259
+ print(e)
260
+ corrupt += 1
261
+ print(e, corrupt, total, corrupt/total)
262
+ continue
263
+ if file_quality_prompt: # override quality prompt
264
+ meta['quality_prompt'] = file_quality_prompt
265
+ if ('height' in meta) and ('width' in meta):
266
+ cur_h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((meta['height']/meta['width']-self.h_div_w_templates)))]
267
+ else:
268
+ cur_h_div_w_template = rough_h_div_w
269
+ if 'h_div_w' in meta:
270
+ del meta['h_div_w']
271
+ meta['first_frame_condition'] = False
272
+ meta['pn'] = self.sample_pn(meta, self.pn_list, self.pn_probs)
273
+ if 'video_path' in meta:
274
+ if self.epoch_rank_generator.random() < self.other_args.i2v_ratio:
275
+ meta['first_frame_condition'] = True
276
+ begin_frame_id, end_frame_id, fps = meta['begin_frame_id'], meta['end_frame_id'], meta['fps']
277
+ real_duration = (end_frame_id - begin_frame_id) / fps
278
+ mapped_duration = int(np.round(real_duration / self.duration_resolution)) * self.duration_resolution
279
+ if mapped_duration < self.min_training_duration:
280
+ continue
281
+ if mapped_duration > self.max_training_duration:
282
+ if self.drop_long_video:
283
+ continue
284
+ else:
285
+ mapped_duration = self.max_training_duration
286
+ if self.other_args.use_clipwise_caption:
287
+ meta['caption'] = [
288
+ meta['caption-InternVL2.0'],
289
+ self.get_video_caption(meta, mapped_duration),
290
+ ]
291
+ else:
292
+ meta['caption'] = [self.get_video_caption(meta, mapped_duration)]
293
+ sample_frames = int(mapped_duration * self.video_fps + 1)
294
+ pt = (sample_frames-1) // self.temporal_compress_rate + 1
295
+ scale_schedule = self.dynamic_resolution_h_w[cur_h_div_w_template][meta['pn']]['pt2scale_schedule'][pt]
296
+ meta['sample_frames'] = sample_frames
297
+ elif 'image_path' in meta:
298
+ mapped_duration = -1
299
+ scale_schedule = self.dynamic_resolution_h_w[cur_h_div_w_template][meta['pn']]['pt2scale_schedule'][1]
300
+ meta['caption'] = [self.get_image_caption(meta)]
301
+ # random set caption to "" for classifier-free guidance
302
+ # refer to: https://github.com/PixArt-alpha/PixArt-alpha/blob/master/train_scripts/train_diffusers.py#L67
303
+ for caption_ind in range(len(meta['caption'])):
304
+ if self.epoch_rank_generator.random() < self.other_args.drop_condition_prob:
305
+ meta['caption'][caption_ind] = ""
306
+ if mapped_duration not in mapped_duration2metas:
307
+ mapped_duration2metas[mapped_duration] = []
308
+
309
+ # get cum_text_visual_tokens
310
+ cum_visual_tokens = []
311
+ preserve_scale_inds = {}
312
+ assert len(scale_schedule) == len(self.other_args.video_scale_probs), f'{len(scale_schedule)=} {len(self.other_args.video_scale_probs)=}'
313
+ for scale_ind, scale in enumerate(scale_schedule):
314
+ if self.epoch_rank_generator.random() < self.other_args.video_scale_probs[scale_ind]:
315
+ preserve_scale_inds[scale_ind] = True
316
+ tokens_this_scale = np.array(scale).prod(-1) + self.other_args.add_scale_token
317
+ cum_visual_tokens.append(tokens_this_scale)
318
+ cum_visual_tokens = np.array(cum_visual_tokens).cumsum()
319
+ meta['cum_text_visual_tokens'] = cum_visual_tokens
320
+ meta['preserve_scale_inds'] = preserve_scale_inds
321
+
322
+ if self.other_args.cache_check_mode == 1: # check at the begining
323
+ if self.exists_cache_file(meta):
324
+ mapped_duration2metas[mapped_duration].append(meta)
325
+ elif self.other_args.cache_check_mode == -1: # select unexist, used for token cache
326
+ if not self.exists_cache_file(meta):
327
+ mapped_duration2metas[mapped_duration].append(meta)
328
+ else:
329
+ mapped_duration2metas[mapped_duration].append(meta)
330
+
331
+ total_metas = sum([len(item) for item in mapped_duration2metas.values()])
332
+ if (self.other_args.restrict_data_size > 0) and (total_metas > self.other_args.restrict_data_size / self.num_replicas):
333
+ stop_read = True
334
+ break
335
+
336
+ # set mapped_duration2freqs
337
+ mapped_duration2freqs = {}
338
+ for mapped_duration in sorted(mapped_duration2metas.keys()):
339
+ mapped_duration2freqs[mapped_duration] = len(mapped_duration2metas[mapped_duration])
340
+
341
+ for mapped_duration in mapped_duration2metas.keys():
342
+ freqs = mapped_duration2freqs[mapped_duration]
343
+ assert len(mapped_duration2metas[mapped_duration]) >= freqs
344
+ self.epoch_rank_generator.shuffle(mapped_duration2metas[mapped_duration])
345
+ mapped_duration2metas[mapped_duration] = mapped_duration2metas[mapped_duration][:freqs]
346
+ # append text tokens
347
+ skip_count_text_token = self.other_args.skip_count_text_token or self.other_args.add_class_token > 0
348
+ mapped_duration2metas[mapped_duration] = self.append_text_tokens(mapped_duration2metas[mapped_duration], skip_count_text_token=skip_count_text_token)
349
+
350
+ total_metas = sum([len(item) for item in mapped_duration2metas.values()])
351
+ for mapped_duration in sorted(mapped_duration2freqs.keys()):
352
+ freq = mapped_duration2freqs[mapped_duration]
353
+ proportion = freq / total_metas * 100
354
+ print(f'{mapped_duration=}, {freq=}, {proportion=:.1f}%')
355
+ return mapped_duration2metas, mapped_duration2freqs
356
+
357
+ def append_text_tokens(self, metas, skip_count_text_token=False, bucket_size=100):
358
+ t1 = time.time()
359
+ pbar = tqdm.tqdm(total=len(metas) // bucket_size + 1, desc='append text tokens')
360
+ valid_metas = []
361
+ for bucket_id in range(len(metas) // bucket_size + 1):
362
+ pbar.update(1)
363
+ start = bucket_id * bucket_size
364
+ end = min(start + bucket_size, len(metas))
365
+ if start >= end:
366
+ break
367
+ captions = []
368
+ caps_per_meta = []
369
+ for i in range(start, end):
370
+ captions.extend(metas[i]['caption'])
371
+ caps_per_meta.append(len(metas[i]['caption']))
372
+ assert len(captions), f'{len(captions)=}'
373
+ if skip_count_text_token:
374
+ lens = [0 for _ in range(len(captions))]
375
+ else:
376
+ lens = self.get_captions_lens(captions)
377
+ lens = np.clip(np.array(lens), a_min=0, a_max=self.max_text_len)
378
+ ptr = 0
379
+ for i in range(start, end):
380
+ text_tokens = sum(lens[ptr:ptr+caps_per_meta[i-start]])
381
+ ptr += caps_per_meta[i-start]
382
+ metas[i]['text_tokens'] = text_tokens
383
+ metas[i]['cum_text_visual_tokens'] = metas[i]['cum_text_visual_tokens'] + metas[i]['text_tokens']
384
+ metas[i]['text_visual_tokens'] = metas[i]['cum_text_visual_tokens'][-1]
385
+ if metas[i]['text_visual_tokens'] <= self.train_max_token_len * self.other_args.dense_ratio4seqpack:
386
+ valid_metas.append(metas[i])
387
+ t2 = time.time()
388
+ print(f'append text tokens: {t2-t1:.1f}s')
389
+ return valid_metas
390
+
391
+ def exists_cache_file(self, meta):
392
+ pn = meta['pn']
393
+ if 'image_path' in meta:
394
+ return osp.exists(self.get_image_cache_file(meta['image_path'], pn))
395
+ else:
396
+ if '/vdataset/clip' in meta['video_path']: # clip
397
+ cache_file = self.get_video_cache_file(meta['video_path'], 0, meta['end_frame_id']-meta['begin_frame_id'], self.video_fps, pn)
398
+ else:
399
+ cache_file = self.get_video_cache_file(meta['video_path'], meta['begin_frame_id'], meta['end_frame_id'], self.video_fps, pn)
400
+ return osp.exists(cache_file)
401
+
402
+ def form_batches(self, mapped_duration2metas):
403
+ examples = []
404
+ for mapped_duration in sorted(mapped_duration2metas.keys()):
405
+ for example_ind in range(len(mapped_duration2metas[mapped_duration])):
406
+ text_visual_tokens = mapped_duration2metas[mapped_duration][example_ind]['text_visual_tokens']
407
+ examples.append((mapped_duration, example_ind, text_visual_tokens))
408
+ examples = sorted(examples, key=lambda x: -x[2])
409
+ max_text_visual_tokens = examples[0][2] if len(examples) else 0
410
+ assert self.train_max_token_len >= max_text_visual_tokens, f'{self.train_max_token_len=} should >= {max_text_visual_tokens=}'
411
+ self.print(f'{self.rank=} {self.mapped_duration2freqs=} form_batches details: {self.rank=} examples={examples[:20]}')
412
+
413
+ st = time.time()
414
+ if self.feature_extraction or self.pair_input: # no sequence packing, for feature extraction or dpo training
415
+ batches = [[item[:2]] for item in examples]
416
+ else:
417
+ batches = []
418
+ left_ptr, right_ptr = 0, len(examples)-1
419
+ while left_ptr <= right_ptr:
420
+ tokens_remain = self.train_max_token_len
421
+ tmp_batch = []
422
+ while left_ptr <= right_ptr and (tokens_remain - examples[left_ptr][2] >= 0):
423
+ tokens_remain = tokens_remain - examples[left_ptr][2]
424
+ tmp_batch.append((left_ptr, examples[left_ptr][2]))
425
+ left_ptr += 1
426
+ while left_ptr <= right_ptr and (tokens_remain - examples[right_ptr][2] >= 0):
427
+ tokens_remain = tokens_remain - examples[right_ptr][2]
428
+ tmp_batch.append((right_ptr, examples[right_ptr][2]))
429
+ right_ptr -= 1
430
+ if len(tmp_batch):
431
+ tmp_batch = sorted(tmp_batch, key=lambda x: -x[1])
432
+ # total_tokens = sum(item[1] for item in tmp_batch)
433
+ # if total_tokens / self.train_max_token_len < 0.7:
434
+ # import pdb; pdb.set_trace()
435
+ batches.append([examples[ptr][:2] for (ptr, _) in tmp_batch])
436
+ if len(batches) % 1000 == 0:
437
+ print(f'form {len(batches)} batches, len(metas)={len(examples)}')
438
+ print(f'[data preprocess] form_batches done, got {len(batches)} batches, cost {time.time()-st:.2f}s')
439
+ self.epoch_rank_generator.shuffle(batches)
440
+ print(f'[data preprocess] shuffle batches done')
441
+ batch_num = len(batches)
442
+ try:
443
+ if self.num_replicas > 1:
444
+ batch_num = torch.tensor([batch_num], device=self.device)
445
+ if tdist.is_initialized():
446
+ tdist.all_reduce(batch_num, op=tdist.ReduceOp.MIN)
447
+ batch_num = batch_num.item()
448
+ except Exception as e:
449
+ print(e)
450
+ batches = batches[:batch_num]
451
+ print(f'[data preprocess] aligned batch number among gpus, got {batch_num} batches')
452
+ return batches
453
+
454
+ def set_epoch_generator(self, epoch):
455
+ self.epoch = epoch
456
+ self.epoch_generator = np.random.default_rng(self.seed + self.epoch)
457
+ self.epoch_rank_generator = np.random.default_rng(self.seed + self.epoch + self.rank)
458
+
459
+ def __getitem__(self, batch_ind_ptr):
460
+ try:
461
+ batch_info = self.batches[batch_ind_ptr%len(self.batches)]
462
+ batch_data = []
463
+ for (mapped_duration, example_ind) in batch_info:
464
+ ret = False
465
+ repeat_times = 0
466
+ mapped_duration_metas = self.mapped_duration2metas[mapped_duration]
467
+ while not ret:
468
+ example_ind = example_ind % len(mapped_duration_metas)
469
+ meta = mapped_duration_metas[example_ind]
470
+ if 'video_path' in meta:
471
+ if self.pair_input:
472
+ ret, model_input = self.prepare_pair_video_input(meta)
473
+ else:
474
+ ret, model_input = self.prepare_video_input(meta)
475
+ elif 'image_path' in meta:
476
+ if self.pair_input:
477
+ ret, model_input = self.prepare_pair_image_input(meta)
478
+ else:
479
+ ret, model_input = self.prepare_image_input(meta)
480
+ if ret:
481
+ if self.pair_input:
482
+ batch_data.extend(model_input)
483
+ else:
484
+ batch_data.append(model_input)
485
+ else: # Handle corrupt example in a batch, just try to read the next one
486
+ example_ind = example_ind + 1
487
+ repeat_times += 1
488
+ if repeat_times % 20 == 0: # Too many corrupt files, switch to another batch
489
+ self.print(f'Caution! I have repeat {repeat_times} times to read a video/image, but still failed to read it. {example_ind=} {meta=}')
490
+ return self.__getitem__(batch_ind_ptr+1)
491
+
492
+ images, raw_features_bcthw, feature_cache_files4images = [], [], []
493
+ text_feature_cache_files = []
494
+ addition_pn_images = {}
495
+ batch_data4images, batch_data4raw_features = [], []
496
+ for item in batch_data:
497
+ if item['raw_features_cthw'] is None:
498
+ images.append(item['img_T3HW'].permute(1,0,2,3)) # # tchw -> cthw
499
+ for key in item:
500
+ if key.startswith('img_T3HW_'):
501
+ if key not in addition_pn_images:
502
+ addition_pn_images[key] = []
503
+ addition_pn_images[key].append(item[key].permute(1,0,2,3))
504
+ feature_cache_files4images.append(item['feature_cache_file'])
505
+ batch_data4images.append(item)
506
+ else:
507
+ raw_features_bcthw.append(item['raw_features_cthw'])
508
+ batch_data4raw_features.append(item)
509
+ batch_data4images_raw_features = batch_data4images + batch_data4raw_features
510
+ captions = [item['text_input'] for item in batch_data4images_raw_features]
511
+ text_feature_cache_files = [item['text_feature_cache_file'] for item in batch_data4images_raw_features]
512
+ meta_list = [item['meta'] for item in batch_data4images_raw_features]
513
+ return {
514
+ 'captions': captions,
515
+ 'images': images,
516
+ 'addition_pn_images': addition_pn_images,
517
+ 'feature_cache_files4images': feature_cache_files4images,
518
+ 'raw_features_bcthw': raw_features_bcthw,
519
+ 'text_cond_tuple': None,
520
+ 'text_feature_cache_files': text_feature_cache_files,
521
+ 'meta_list': meta_list,
522
+ 'media': 'videos',
523
+ }
524
+ except Exception as e:
525
+ print(f'get item error: {e}')
526
+ return self.__getitem__(batch_ind_ptr+1)
527
+
528
+
529
+ def prepare_image_input(self, info) -> Tuple:
530
+ try:
531
+ img_path, text_input = osp.abspath(info['image_path']), info['caption']
532
+ img_T3HW, raw_features_cthw, feature_cache_file, text_features_lenxdim, text_feature_cache_file = [None] * 5
533
+ if self.use_vae_token_cache:
534
+ feature_cache_file = self.get_image_cache_file(img_path, info['pn'])
535
+ if osp.exists(feature_cache_file):
536
+ try:
537
+ raw_features_cthw = self.load_visual_token(feature_cache_file)
538
+ except Exception as e:
539
+ print(f'load cache file error: {e}')
540
+ os.remove(feature_cache_file)
541
+ if raw_features_cthw is None and (not self.allow_online_vae_feature_extraction):
542
+ return False, None
543
+ if raw_features_cthw is None:
544
+ with open(img_path, 'rb') as f:
545
+ img: PImage.Image = PImage.open(f)
546
+ w, h = img.size
547
+ h_div_w = h / w
548
+ h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((h_div_w-self.h_div_w_templates)))]
549
+ tgt_h, tgt_w = self.dynamic_resolution_h_w[h_div_w_template][info['pn']]['pixel']
550
+ img = img.convert('RGB')
551
+ if self.c2i:
552
+ img_T3HW = self.c2i_transform(img)
553
+ else:
554
+ img_T3HW = transform(img, tgt_h, tgt_w)
555
+ img_T3HW = img_T3HW.unsqueeze(0)
556
+ assert img_T3HW.shape == (1, 3, tgt_h, tgt_w)
557
+ data_item = {
558
+ 'text_input': text_input,
559
+ 'img_T3HW': img_T3HW,
560
+ 'raw_features_cthw': raw_features_cthw,
561
+ 'feature_cache_file': feature_cache_file,
562
+ 'text_features_lenxdim': text_features_lenxdim,
563
+ 'text_feature_cache_file': text_feature_cache_file,
564
+ 'meta': info,
565
+ }
566
+ return True, data_item
567
+ except Exception as e:
568
+ print(f'prepare_image_input error: {e}')
569
+ return False, None
570
+
571
+ def prepare_pair_image_input(self, info) -> Tuple:
572
+ pass
573
+
574
+ def prepare_pair_video_input(self, info) -> Tuple:
575
+ win_flag, win_data_item = self.prepare_video_input(copy.deepcopy(info))
576
+
577
+ info['video_path'] = info['lose_video_path']
578
+ lose_flag, lose_data_item = self.prepare_video_input(info)
579
+
580
+ flag = win_flag and lose_flag
581
+ return flag, [win_data_item, lose_data_item]
582
+
583
+ def load_visual_token(self, feature_cache_file):
584
+ raw_features_cthw = load_packed_tensor(feature_cache_file)
585
+ from grn.utils_t2iv.hbq_util_t2iv import bit_label2raw_feature
586
+ raw_features_cthw = bit_label2raw_feature(raw_features_cthw.unsqueeze(0), self.other_args.hbq_round)[0]
587
+ return raw_features_cthw
588
+
589
+ def prepare_video_input(self, info) -> Tuple:
590
+ filename, begin_frame_id, end_frame_id = (
591
+ info["video_path"],
592
+ info["begin_frame_id"],
593
+ info["end_frame_id"],
594
+ )
595
+
596
+ try:
597
+ img_T3HW, raw_features_cthw, feature_cache_file, text_features_lenxdim, text_feature_cache_file = None, None, None, None, None
598
+ img_T3HW_4additional_pn = {}
599
+ text_input = info['caption']
600
+ if '/vdataset/clip' in filename: # clip
601
+ begin_frame_id, end_frame_id = 0, end_frame_id - begin_frame_id
602
+ sample_frames = info['sample_frames']
603
+ tmp_local_path = ''
604
+ if self.use_vae_token_cache:
605
+ feature_cache_file = self.get_video_cache_file(info["video_path"], begin_frame_id, end_frame_id, self.video_fps, info['pn'])
606
+ if osp.exists(feature_cache_file):
607
+ try:
608
+ pt = (sample_frames-1) // self.temporal_compress_rate + 1
609
+ raw_features_cthw = self.load_visual_token(feature_cache_file)
610
+ assert raw_features_cthw.shape[1] >= pt, f'raw_features_cthw.shape[1] >= pt: {raw_features_cthw.shape[1]} vs {pt}'
611
+ if raw_features_cthw.shape[1] > pt:
612
+ raw_features_cthw = raw_features_cthw[:,:pt]
613
+ except Exception as e:
614
+ self.print(f'load video cache file error: {e}')
615
+ os.remove(feature_cache_file)
616
+ raw_features_cthw = None
617
+ if raw_features_cthw is None and (not self.allow_online_vae_feature_extraction):
618
+ return False, None
619
+ pn_list = [info['pn']]
620
+ if raw_features_cthw is None:
621
+ tmp_local_path = info["video_path"]
622
+ if not osp.exists(tmp_local_path):
623
+ return False, None
624
+ video = EncodedVideoDecord(tmp_local_path, os.path.basename(tmp_local_path), num_threads=0)
625
+ start_interval = max(0, begin_frame_id / video._fps)
626
+ end_interval = start_interval+(sample_frames-1)/self.video_fps
627
+ assert end_interval <= video.duration + 0.2, f'{end_interval=}, but {video.duration=}' # 0.2s margin
628
+ end_interval = min(end_interval, video.duration)
629
+ raw_video, _ = video.get_clip(start_interval, end_interval, sample_frames) # rgb order
630
+ h, w, _ = raw_video[0].shape
631
+ h_div_w = h / w
632
+ h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((h_div_w-self.h_div_w_templates)))]
633
+ tgt_h, tgt_w = self.dynamic_resolution_h_w[h_div_w_template][info['pn']]['pixel']
634
+
635
+ for pn in pn_list:
636
+ img_T3HW = [transform(Image.fromarray(frame).convert("RGB"), tgt_h, tgt_w) for frame in raw_video]
637
+ img_T3HW = torch.stack(img_T3HW, 0)
638
+ img_T3HW_4additional_pn[pn] = img_T3HW
639
+ del video
640
+ assert img_T3HW.shape[-3:] == (3, tgt_h, tgt_w)
641
+ data_item = {
642
+ 'text_input': text_input,
643
+ 'img_T3HW': img_T3HW_4additional_pn.get(info['pn'], None),
644
+ 'raw_features_cthw': raw_features_cthw,
645
+ 'feature_cache_file': feature_cache_file,
646
+ 'text_features_lenxdim': text_features_lenxdim,
647
+ 'text_feature_cache_file': text_feature_cache_file,
648
+ 'meta': info,
649
+ }
650
+ for pn in pn_list[1:]:
651
+ data_item.update({f'img_T3HW_{pn}': img_T3HW_4additional_pn.get(pn, None)})
652
+ return True, data_item
653
+ except Exception as e:
654
+ self.print(f'prepare_video_input error: {e}, info: {info}')
655
+ return False, None
656
+
657
+
658
+ @staticmethod
659
+ def collate_function(batch, online_t5: bool = False) -> None:
660
+ pass
661
+
662
+ def random_drop_sentences(self, caption, min_sentences):
663
+ elems = [item for item in caption.split('.') if item]
664
+ if len(elems) <= min_sentences:
665
+ return caption
666
+ sentences = self.epoch_rank_generator.integers(min_sentences, len(elems)+1)
667
+ return '.'.join(elems[:sentences]) + '.'
668
+
669
+ def __len__(self):
670
+ return len(self.batches) * self.other_args.loop_data_per_epoch
671
+
672
+ def get_image_cache_file(self, image_path, pn):
673
+ elems = image_path.split('/')
674
+ elems = [item for item in elems if item]
675
+ filename, ext = osp.splitext(elems[-1])
676
+ filename = get_prompt_id(filename)
677
+ save_filepath = osp.join(self.token_cache_dir, f'images_pn_{pn}', '/'.join(elems[4:-1]), f'{filename}.npz')
678
+ return save_filepath
679
+
680
+ def get_video_cache_file(self, video_path, begin_frame_id, end_frame_id, video_fps, pn):
681
+ elems = video_path.split('/')
682
+ elems = [item for item in elems if item]
683
+ filename, ext = osp.splitext(elems[-1])
684
+ filename = get_prompt_id(filename)
685
+ save_filepath = osp.join(self.token_cache_dir, f'pn_{pn}_sample_fps_{video_fps}', '/'.join(elems[4:-1]), f'{filename}_sf_{begin_frame_id}_ef_{end_frame_id}.npz')
686
+ return save_filepath
687
+
grn/models/basic.py ADDED
@@ -0,0 +1,256 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import os
3
+ from functools import partial
4
+ from typing import Optional, Tuple, Union
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ import numpy as np
10
+ from torch.utils.checkpoint import checkpoint
11
+ from torch.nn.functional import scaled_dot_product_attention as slow_attn # q, k, v: BHLc
12
+
13
+ from grn.models.rope import apply_rotary_emb
14
+ from grn.utils_t2iv.sequence_parallel import sp_all_to_all, SequenceParallelManager as sp_manager
15
+
16
+ try:
17
+ from flash_attn.cute import flash_attn_varlen_func
18
+ except:
19
+ from flash_attn import flash_attn_varlen_func
20
+
21
+ # Import flash_attn's fused ops
22
+ try:
23
+ from flash_attn.ops.rms_norm import rms_norm as rms_norm_impl
24
+ except ImportError:
25
+ def rms_norm_impl(x, weight, epsilon):
26
+ return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(epsilon))) * weight
27
+
28
+ def merge_states(states, splits, cfg):
29
+ """
30
+ pick key and value states for flash_attn_varlen_func
31
+ Args:
32
+ states: list of states
33
+ splits: list of split sizes
34
+ cfg: bool, use cfg or not"""
35
+ if cfg:
36
+ cond_len, uncond_len = 0, 0
37
+ cond_states, uncond_states = [], []
38
+ for stat_, split_ in zip(states, splits):
39
+ cond, uncond = torch.split(stat_, split_, dim=2)
40
+ cond_states.append(cond)
41
+ uncond_states.append(uncond)
42
+ cond_len += cond.shape[2]
43
+ uncond_len += uncond.shape[2]
44
+ return cond_states + uncond_states, [cond_len, uncond_len]
45
+ else:
46
+ cond_len = 0
47
+ for stat_ in states:
48
+ cond_len += stat_.shape[2]
49
+ return states, [cond_len]
50
+
51
+ class FastRMSNorm(nn.Module):
52
+ def __init__(self, C, eps=1e-6, elementwise_affine=True):
53
+ super().__init__()
54
+ self.C = C
55
+ self.eps = eps
56
+ self.elementwise_affine = elementwise_affine
57
+ if self.elementwise_affine:
58
+ self.weight = nn.Parameter(torch.ones(C))
59
+ else:
60
+ self.register_buffer('weight', torch.ones(C))
61
+
62
+ def forward(self, x):
63
+ src_type = x.dtype
64
+ return rms_norm_impl(x.float(), self.weight, epsilon=self.eps).to(src_type)
65
+
66
+ def extra_repr(self) -> str:
67
+ return f'C={self.C}, eps={self.eps:g}, elementwise_affine={self.elementwise_affine}'
68
+
69
+
70
+ class WanLayerNorm(nn.LayerNorm):
71
+
72
+ def __init__(self, dim, eps=1e-6, elementwise_affine=False):
73
+ super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
74
+
75
+ def forward(self, x):
76
+ r"""
77
+ Args:
78
+ x(Tensor): Shape [B, L, C]
79
+ """
80
+ return super().forward(x.float()).type_as(x)
81
+
82
+
83
+ class Qwen3MLP(nn.Module):
84
+ def __init__(self, hidden_size, intermediate_size):
85
+ super().__init__()
86
+ self.hidden_size = hidden_size
87
+ self.intermediate_size = intermediate_size
88
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
89
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
90
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
91
+ self.act_fn = nn.SiLU()
92
+
93
+ def forward(self, x):
94
+ down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
95
+ return down_proj
96
+
97
+ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
98
+ """
99
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
100
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
101
+ """
102
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
103
+ if n_rep == 1:
104
+ return hidden_states
105
+ hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
106
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
107
+
108
+ class SelfAttention(nn.Module):
109
+ def __init__(
110
+ self, embed_dim=768, num_heads=12, num_key_value_heads=-1,
111
+ use_flex_attn=False, qwen_qkvo_bias=False, **kwargs,
112
+ ):
113
+ """
114
+ :param embed_dim: model's width
115
+ :param num_heads: num heads of multi-head attention
116
+ :param proj_drop: always 0 for testing
117
+ :param tau: always 1
118
+ :param cos_attn: always True: during attention, q and k will be L2-normalized and scaled by a head-wise learnable parameter self.scale_mul_1H11
119
+ """
120
+ super().__init__()
121
+ assert embed_dim % num_heads == 0
122
+ assert num_key_value_heads == -1 or num_heads % num_key_value_heads == 0
123
+
124
+ self.num_heads, self.head_dim = num_heads, embed_dim // num_heads
125
+ self.num_key_value_heads = num_key_value_heads if num_key_value_heads > 0 else num_heads
126
+ self.q_proj = nn.Linear(embed_dim, self.num_heads*self.head_dim, bias=qwen_qkvo_bias)
127
+ self.k_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
128
+ self.v_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
129
+ self.o_proj = nn.Linear(self.num_heads*self.head_dim, embed_dim, bias=qwen_qkvo_bias)
130
+ self.q_norm = FastRMSNorm(self.head_dim)
131
+ self.k_norm = FastRMSNorm(self.head_dim)
132
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
133
+ self.scale = self.head_dim**-0.5
134
+
135
+ self.caching = False # kv caching: only used during inference
136
+ self.cached_k = {} # kv caching: only used during inference
137
+ self.cached_v = {} # kv caching: only used during inference
138
+ self.cached_split_cond_uncond = {} # only used during inference
139
+
140
+ self.use_flex_attn = use_flex_attn
141
+
142
+ def kv_caching(self, enable: bool): # kv caching: only used during inference
143
+ self.caching = enable
144
+ self.cached_k = {}
145
+ self.cached_v = {}
146
+ self.cached_split_cond_uncond = {}
147
+
148
+ # NOTE: attn_bias_or_two_vector is None during inference
149
+ def forward(self, x, cu_seqlens, max_seqlen, attn_bias_or_two_vector: Union[torch.Tensor, Tuple[torch.IntTensor, torch.IntTensor]], attn_fn=None, rope2d_freqs_grid=[], scale_ind=0, context_info=None, last_diffusion_step=True, ref_text_scale_inds=[], use_cfg=False, split_cond_uncond=[], **kwargs):
150
+ # x: fp32
151
+ B, L, C = x.shape
152
+ hidden_states = x
153
+ input_shape = hidden_states.shape[:-1]
154
+ hidden_shape = (*input_shape, -1, self.head_dim)
155
+ query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).contiguous()# batch, slen, heads, head_dim
156
+ key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).contiguous() # batch, slen, num_key_value_heads, head_dim
157
+ value_states = self.v_proj(hidden_states).view(hidden_shape).contiguous() # batch, slen, num_key_value_heads, head_dim
158
+
159
+ if sp_manager.sp_on():
160
+ # Headnum need to be sharded and L needs to be gathered
161
+ # [B, H, raw_L/sp, C] --> [B, H/sp, raw_L, C]
162
+ sdim = 1
163
+ gdim = 2
164
+ L = L * sp_manager.get_sp_size()
165
+ C = C // sp_manager.get_sp_size()
166
+ query_states = sp_all_to_all(query_states, sdim, gdim)
167
+ key_states = sp_all_to_all(key_states, sdim, gdim)
168
+ value_states = sp_all_to_all(value_states, sdim, gdim)
169
+
170
+ query_states, key_states = apply_rotary_emb(query_states, key_states, rope2d_freqs_grid)
171
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
172
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
173
+
174
+ if attn_bias_or_two_vector is None:
175
+ # fa4, flash_attn_func input/output should be (batch_size, seqlen, nheads, headdim)
176
+ from flash_attn.cute import flash_attn_varlen_func
177
+ attn_output = flash_attn_varlen_func(
178
+ q = query_states.squeeze(0),
179
+ k = key_states.squeeze(0),
180
+ v = value_states.squeeze(0),
181
+ cu_seqlens_q=cu_seqlens,
182
+ cu_seqlens_k=cu_seqlens,
183
+ max_seqlen_q=max_seqlen,
184
+ max_seqlen_k=max_seqlen,
185
+ softmax_scale=self.scale,
186
+ )
187
+ attn_output = attn_output[0].reshape(B, L, C).contiguous()
188
+ else:
189
+ # slow attn
190
+ attn_output = slow_attn(query=query_states.transpose(1, 2), key=key_states.transpose(1, 2), value=value_states.transpose(1, 2), scale=self.scale, attn_mask=attn_bias_or_two_vector, dropout_p=0).transpose(1, 2).reshape(B, L, C)
191
+
192
+ if sp_manager.sp_on():
193
+ # [B, raw_L, C/sp] --> [B, raw_L/sp, C]
194
+ sdim = 1
195
+ gdim = 2
196
+ attn_output = sp_all_to_all(attn_output, sdim, gdim)
197
+
198
+ attn_output = self.o_proj(attn_output)
199
+
200
+ return attn_output
201
+
202
+ class SelfAttnBlock(nn.Module):
203
+ def __init__(
204
+ self, embed_dim, num_heads, num_key_value_heads, mlp_ratio=4.,
205
+ use_flex_attn=False,
206
+ qwen_qkvo_bias=False, use_ada_layer_norm=False, **kwargs,
207
+ ):
208
+ super(SelfAttnBlock, self).__init__()
209
+ self.C = embed_dim
210
+ self.attn = SelfAttention(
211
+ embed_dim=embed_dim, num_heads=num_heads, num_key_value_heads=num_key_value_heads,
212
+ use_flex_attn=use_flex_attn, qwen_qkvo_bias=qwen_qkvo_bias, **kwargs,
213
+ )
214
+ self.mlp = Qwen3MLP(hidden_size=embed_dim, intermediate_size=round(embed_dim * mlp_ratio / 256) * 256)
215
+ self.use_ada_layer_norm = use_ada_layer_norm
216
+ if self.use_ada_layer_norm:
217
+ self.modulation = nn.Parameter(torch.randn(1, 6, embed_dim) / embed_dim**0.5)
218
+ self.input_layernorm = WanLayerNorm(embed_dim)
219
+ self.post_attention_layernorm = WanLayerNorm(embed_dim)
220
+ else:
221
+ self.input_layernorm = FastRMSNorm(embed_dim)
222
+ self.post_attention_layernorm = FastRMSNorm(embed_dim)
223
+
224
+ # NOTE: attn_bias_or_two_vector is None during inference
225
+ def forward(self, x, cu_seqlens, max_seqlen, e0, attn_bias_or_two_vector, attn_fn=None, rope2d_freqs_grid=[], scale_ind=0, context_info=None, last_diffusion_step=True, ref_text_scale_inds=[], use_cfg=False, split_cond_uncond=[], **kwargs):
226
+ # x: [B,L,C]
227
+ # e0: [B, L, 6, C]
228
+ if self.use_ada_layer_norm:
229
+ assert e0.dtype == torch.float32
230
+ e = e0
231
+ with torch.amp.autocast('cuda', dtype=torch.float32):
232
+ e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
233
+ residual = x
234
+ hidden_states = x
235
+ hidden_states = self.input_layernorm(hidden_states).float() * (1 + e[1].squeeze(2)) + e[0].squeeze(2)
236
+ hidden_states = self.attn(hidden_states, cu_seqlens, max_seqlen, attn_bias_or_two_vector, attn_fn, rope2d_freqs_grid, scale_ind, context_info, last_diffusion_step, ref_text_scale_inds, use_cfg, split_cond_uncond, **kwargs)
237
+ with torch.amp.autocast('cuda', dtype=torch.float32):
238
+ hidden_states = residual + hidden_states * e[2].squeeze(2)
239
+ # Fully Connected
240
+ residual = hidden_states
241
+ hidden_states = self.post_attention_layernorm(hidden_states).float() * (1 + e[4].squeeze(2)) + e[3].squeeze(2)
242
+ hidden_states = self.mlp(hidden_states)
243
+ with torch.amp.autocast('cuda', dtype=torch.float32):
244
+ hidden_states = residual + hidden_states * e[5].squeeze(2)
245
+ else:
246
+ residual = x
247
+ hidden_states = x
248
+ hidden_states = self.input_layernorm(hidden_states)
249
+ hidden_states = self.attn(hidden_states, cu_seqlens, max_seqlen, attn_bias_or_two_vector, attn_fn, rope2d_freqs_grid, scale_ind, context_info, last_diffusion_step, ref_text_scale_inds, use_cfg, split_cond_uncond, **kwargs)
250
+ hidden_states = residual + hidden_states
251
+ # Fully Connected
252
+ residual = hidden_states
253
+ hidden_states = self.post_attention_layernorm(hidden_states)
254
+ hidden_states = self.mlp(hidden_states)
255
+ hidden_states = residual + hidden_states
256
+ return hidden_states
grn/models/ema.py ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import torch
3
+ from collections import OrderedDict
4
+
5
+
6
+ def get_ema_model(model):
7
+ ema_model = copy.deepcopy(model)
8
+ ema_model.eval()
9
+ for param in ema_model.parameters():
10
+ param.requires_grad = False
11
+ return ema_model
12
+
13
+ @torch.no_grad()
14
+ def update_ema(ema_model, model, decay):
15
+ """
16
+ Step the EMA model towards the current model.
17
+ """
18
+ ema_params = OrderedDict(ema_model.named_parameters())
19
+ model_params = OrderedDict(model.named_parameters())
20
+
21
+ for name, param in model_params.items():
22
+ # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
23
+ ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
grn/models/flex_attn_mask.py ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from functools import partial
2
+ import torch
3
+ import numpy as np
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ from torch.nn.attention.flex_attention import flex_attention, create_block_mask
7
+
8
+
9
+ def _length_to_offsets(lengths, device):
10
+ offsets = [0]
11
+ offsets.extend(lengths)
12
+ offsets = torch.tensor(offsets, device=device, dtype=torch.int32)
13
+ offsets = torch.cumsum(offsets, dim=-1)
14
+ return offsets
15
+
16
+ def _offsets_to_doc_ids_tensor(offsets):
17
+ device = offsets.device
18
+ counts = offsets[1:] - offsets[:-1]
19
+ visual = torch.repeat_interleave(torch.arange(len(counts), device=device, dtype=torch.int32), counts)
20
+ return visual
21
+
22
+ def _generate_overall_mask(offsets, querysid_refsid):
23
+ document_id = _offsets_to_doc_ids_tensor(offsets) # to scale_ind
24
+ def overall_mask(b, h, q_idx, kv_idx):
25
+ querysid = document_id[q_idx]
26
+ kv_sid = document_id[kv_idx]
27
+ return querysid_refsid[querysid][kv_sid]
28
+ return overall_mask
29
+
30
+ def causal(b, h, q_idx, kv_idx):
31
+ return q_idx >= kv_idx
32
+
33
+ def build_flex_attn_func(
34
+ flex_attention,
35
+ seq_l,
36
+ prefix_lens,
37
+ args,
38
+ device,
39
+ batch_size,
40
+ heads,
41
+ pad_seq_len,
42
+ sequece_packing_scales,
43
+ super_scale_lengths,
44
+ super_querysid_super_refsid,
45
+ ):
46
+ """
47
+ Build a flex attn function for a given scale schedule.
48
+ Args:
49
+ flex_attention: compiled flex attention
50
+ seq_l: seq length
51
+ prefix_lens: valid text prefix lens, [bs]
52
+ args: arguments
53
+ device: device
54
+ batch_size: batch size
55
+ heads: heads
56
+ pad_seq_len: pad_seq_len
57
+ sequece_packing_scales: list of scale schedule
58
+ querysid_refsid: list of scale_pack_info
59
+ Returns:
60
+ attn_fn: flex attn function
61
+ """
62
+ assert sum(super_scale_lengths) == seq_l, f'{sum(super_scale_lengths)}!= {seq_l}'
63
+ offsets = _length_to_offsets(super_scale_lengths, device=device)
64
+ mask_mod = _generate_overall_mask(offsets, super_querysid_super_refsid)
65
+ block_mask = create_block_mask(mask_mod, B = batch_size, H = heads, Q_LEN = seq_l, KV_LEN = seq_l, device = device, _compile = True)
66
+ attn_fn = partial(flex_attention, block_mask=block_mask)
67
+ return attn_fn
grn/models/fused_op.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gc
2
+ from copy import deepcopy
3
+ from typing import Union
4
+
5
+ import torch
6
+ from torch import nn as nn
7
+ from torch.nn import functional as F
8
+
9
+
10
+ @torch.compile(fullgraph=True)
11
+ def fused_rms_norm(x: torch.Tensor, weight: nn.Parameter, eps: float):
12
+ x = x.float()
13
+ return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps))) * weight
14
+
15
+
16
+ @torch.compile(fullgraph=True)
17
+ def fused_ada_layer_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor):
18
+ x = x.float()
19
+ x = F.layer_norm(input=x, normalized_shape=(C,), weight=None, bias=None, eps=eps)
20
+ return x.mul(scale.add(1)).add_(shift)
21
+
22
+
23
+ @torch.compile(fullgraph=True)
24
+ def fused_ada_rms_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor):
25
+ x = x.float()
26
+ x = (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps)))
27
+ return x.mul(scale.add(1)).add_(shift)
grn/models/grn.py ADDED
@@ -0,0 +1,754 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import math
3
+ import time
4
+ from contextlib import nullcontext
5
+ from functools import partial
6
+ from typing import Any, Dict, List, Optional, Tuple, Union
7
+
8
+ import numpy as np
9
+ import torch
10
+ import torch.nn as nn
11
+ import torch.nn.functional as F
12
+ import torch.utils.checkpoint
13
+ import tqdm
14
+ from timm.models import register_model
15
+
16
+ import grn.utils_t2iv.dist as dist
17
+ from grn.models.basic import FastRMSNorm, SelfAttnBlock
18
+ from grn.models.rope import precompute_rope3d_freqs_grid
19
+ from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta
20
+ from grn.utils_t2iv.dist import for_visualize
21
+ from grn.utils_t2iv.hbq_util_t2iv import multiclass_labels2onehot_input
22
+ from grn.utils_t2iv.sequence_parallel import SequenceParallelManager as sp_manager
23
+ from grn.utils_t2iv.sequence_parallel import sp_gather_sequence_by_dim, sp_split_sequence_by_dim
24
+
25
+
26
+ class MultipleLayers(nn.Module):
27
+ """A sequential container for a chunk of multiple transformer blocks."""
28
+
29
+ def __init__(self, layers: List[nn.Module], num_blocks: int, start_index: int):
30
+ super().__init__()
31
+ self.module = nn.ModuleList([
32
+ layers[i] for i in range(start_index, start_index + num_blocks)
33
+ ])
34
+
35
+ def forward(
36
+ self, x, cu_seqlens, max_seqlen, e0: Optional[torch.Tensor],
37
+ attn_bias_or_two_vector: Optional[Any], attn_fn: Optional[Any] = None,
38
+ checkpointing_full_block: bool = False, rope2d_freqs_grid: Optional[torch.Tensor] = None,
39
+ scale_ind: Optional[Any] = None, context_info: Optional[Any] = None,
40
+ last_diffusion_step: bool = True, ref_text_scale_inds: Optional[List[Any]] = None,
41
+ use_cfg: bool = False, split_cond_uncond: Optional[List[Any]] = None
42
+ ) -> torch.Tensor:
43
+ h = x
44
+ for m in self.module:
45
+ if checkpointing_full_block:
46
+ h = torch.utils.checkpoint.checkpoint(
47
+ m, h, cu_seqlens, max_seqlen, e0, attn_bias_or_two_vector, attn_fn,
48
+ rope2d_freqs_grid, scale_ind, context_info,
49
+ last_diffusion_step, ref_text_scale_inds,
50
+ use_cfg, split_cond_uncond, use_reentrant=False
51
+ )
52
+ else:
53
+ h = m(
54
+ h, cu_seqlens, max_seqlen, e0, attn_bias_or_two_vector, attn_fn,
55
+ rope2d_freqs_grid, scale_ind, context_info,
56
+ last_diffusion_step, ref_text_scale_inds,
57
+ use_cfg, split_cond_uncond
58
+ )
59
+ return h
60
+
61
+
62
+ def sinusoidal_embedding_1d(dim: int, position: torch.Tensor) -> torch.Tensor:
63
+ """
64
+ Generate 1D sinusoidal embeddings.
65
+
66
+ Args:
67
+ dim (int): Embedding dimension (must be even).
68
+ position (torch.Tensor): Position tensor of shape [B, L].
69
+
70
+ Returns:
71
+ torch.Tensor: Embeddings of shape [B, L, dim].
72
+ """
73
+ if dim % 2 != 0:
74
+ raise ValueError(f"Embedding dimension must be even, got {dim}")
75
+
76
+ half = dim // 2
77
+ b, l = position.shape
78
+ position = position.reshape(-1).type(torch.float64)
79
+
80
+ sinusoid = torch.outer(
81
+ position,
82
+ torch.pow(10000, -torch.arange(half).to(position).div(half))
83
+ )
84
+ x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
85
+ return x.reshape(b, l, dim)
86
+
87
+
88
+ class TimestepEmbedder(nn.Module):
89
+ """Embeds scalar timesteps into vector representations."""
90
+
91
+ def __init__(self, hidden_size: int, frequency_embedding_size: int = 256):
92
+ super().__init__()
93
+ self.mlp = nn.Sequential(
94
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
95
+ nn.SiLU(),
96
+ nn.Linear(hidden_size, hidden_size, bias=True),
97
+ )
98
+ self.frequency_embedding_size = frequency_embedding_size
99
+
100
+ @staticmethod
101
+ def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor:
102
+ """Create sinusoidal timestep embeddings."""
103
+ half = dim // 2
104
+ freqs = torch.exp(
105
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
106
+ ).to(device=t.device)
107
+ args = t[:, None].float() * freqs[None]
108
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
109
+ if dim % 2:
110
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
111
+ return embedding
112
+
113
+ def forward(self, t: torch.Tensor) -> torch.Tensor:
114
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
115
+ return self.mlp(t_freq)
116
+
117
+
118
+ def bld_to_bthwd(item: torch.Tensor, patch_time: int, patch_height: int, patch_width: int, apply_spatial_patchify: bool = False) -> torch.Tensor:
119
+ """Reshape a sequence tensor to a spatial tensor."""
120
+ batch_size = item.shape[0]
121
+ return item.reshape(batch_size, patch_time, patch_height, patch_width, -1)
122
+
123
+
124
+ def build_attn_mask(seqlens, device):
125
+ attn_mask = torch.zeros((1, 1, sum(seqlens), sum(seqlens)), dtype=torch.bool, device=device)
126
+ q_start = 0
127
+ for i in range(len(seqlens)):
128
+ q_len = seqlens[i]
129
+ q_end = q_start + q_len
130
+ attn_mask[:, :, q_start:q_end, q_start:q_end] = True
131
+ q_start = q_end
132
+ return attn_mask
133
+
134
+
135
+ class FsqHead(nn.Module):
136
+ """Classification head for Finite Scalar Quantization (FSQ)."""
137
+
138
+ def __init__(self, hidden_dim: int, fsq_dim: int, fsq_lvl: int, use_ada_layer_norm: bool, eps: float = 1e-6):
139
+ super().__init__()
140
+ self.proj = nn.Linear(hidden_dim, fsq_dim * fsq_lvl)
141
+ self.norm = FastRMSNorm(hidden_dim)
142
+
143
+ def forward(self, x: torch.Tensor, e: Optional[torch.Tensor] = None) -> torch.Tensor:
144
+ with torch.amp.autocast('cuda', dtype=torch.float32):
145
+ return self.proj(self.norm(x))
146
+
147
+
148
+ class GRN(nn.Module):
149
+ def __init__(
150
+ self,
151
+ vae_local: Any,
152
+ arch: str = 'var',
153
+ qwen_qkvo_bias: bool = False,
154
+ text_channels: int = 0,
155
+ text_maxlen: int = 0,
156
+ embed_dim: int = 1024,
157
+ depth: int = 16,
158
+ num_key_value_heads: int = -1,
159
+ num_heads: int = 16,
160
+ mlp_ratio: float = 4.0,
161
+ drop_path_rate: float = 0.0,
162
+ norm_eps: float = 1e-6,
163
+ block_chunks: int = 1,
164
+ checkpointing: Optional[str] = None,
165
+ pad_to_multiplier: int = 0,
166
+ use_flex_attn: bool = False,
167
+ num_of_label_value: int = 2,
168
+ rope2d_normalized_by_hw: int = 0,
169
+ pn: Optional[str] = None,
170
+ video_frames: int = 1,
171
+ always_training_scales: int = 20,
172
+ apply_spatial_patchify: int = 0,
173
+ inference_mode: bool = False,
174
+ other_args: Optional[Any] = None,
175
+ **kwargs: Any,
176
+ ):
177
+ super().__init__()
178
+ # 1. Model Configuration
179
+ self.embed_dim = embed_dim
180
+ self.depth = depth
181
+ self.num_heads = num_heads
182
+ self.arch = arch
183
+ self.mlp_ratio = mlp_ratio
184
+ self.norm_eps = norm_eps
185
+ self.drop_path_rate = drop_path_rate
186
+ self.use_flex_attn = use_flex_attn
187
+ self.checkpointing = checkpointing
188
+ self.inference_mode = inference_mode
189
+ self.other_args = other_args
190
+
191
+ # 2. Embedding & Scale Configuration
192
+ self.vae_embed_dim = vae_local.codebook_dim
193
+ self.apply_spatial_patchify = apply_spatial_patchify
194
+ self.text_channels = text_channels
195
+ self.text_maxlen = text_maxlen
196
+ self.is_text_to_image = text_channels != 0
197
+
198
+ classifier_head_dim = other_args.detail_scale_dim
199
+ classifier_head_lvl = other_args.detail_num_lvl
200
+ hbq_round = other_args.hbq_round
201
+
202
+ if other_args.refine_mode in ['ar_discrete_GRN_ind']:
203
+ self.visual_embedding_in_dim = vae_local.codebook_dim * (2**hbq_round)
204
+ classifier_head_dim = vae_local.codebook_dim
205
+ elif other_args.refine_mode in ['ar_discrete_GRN_bit']:
206
+ self.visual_embedding_in_dim = hbq_round * vae_local.codebook_dim * 2
207
+ classifier_head_dim = hbq_round * vae_local.codebook_dim
208
+ else:
209
+ self.visual_embedding_in_dim = vae_local.codebook_dim
210
+
211
+ if self.apply_spatial_patchify:
212
+ self.visual_embedding_in_dim *= 4
213
+
214
+ # 3. Dynamic Resolution & Video Specifics
215
+ self.video_frames = video_frames
216
+ self.always_training_scales = always_training_scales
217
+ self.num_of_label_value = num_of_label_value
218
+ self.rope2d_normalized_by_hw = rope2d_normalized_by_hw
219
+
220
+ self.dynamic_resolution_h_w, self.h_div_w_templates = get_dynamic_resolution_meta(
221
+ other_args.dynamic_scale_schedule, other_args.train_h_div_w_list, other_args.video_frames
222
+ )
223
+ self.train_h_div_w_list = self.h_div_w_templates
224
+ print(f"train_h_div_w_list: {self.train_h_div_w_list}")
225
+
226
+ # 4. Utilities
227
+ self.entrophy_statistics = []
228
+ self.top_p, self.top_k = 1.0, 100
229
+ self.rng = torch.Generator(device=dist.get_device())
230
+ self.maybe_record_function = nullcontext
231
+ self.infer_ts = None
232
+
233
+ # 5. Model Components (Projections, Embeddings)
234
+ self.norm0_cond = nn.Identity()
235
+ self.text_proj = nn.Linear(self.text_channels, self.embed_dim)
236
+
237
+ if self.other_args.use_ada_layer_norm:
238
+ self.scale_or_time_dim = 256
239
+ self.scale_or_time_embedding = nn.Sequential(
240
+ nn.Linear(self.scale_or_time_dim, self.embed_dim), nn.SiLU(), nn.Linear(self.embed_dim, self.embed_dim),
241
+ )
242
+ self.scale_or_time_projection = nn.Sequential(nn.SiLU(), nn.Linear(self.embed_dim, self.embed_dim * 6))
243
+
244
+ tmp_h_div_w_template = self.train_h_div_w_list[0]
245
+
246
+ # RoPE grid initialization
247
+ with torch.amp.autocast('cuda', dtype=torch.float32):
248
+ self.rope2d_freqs_grid = precompute_rope3d_freqs_grid(
249
+ dim=self.embed_dim // self.num_heads,
250
+ rope2d_normalized_by_hw=self.rope2d_normalized_by_hw,
251
+ activated_h_div_w_templates=self.train_h_div_w_list,
252
+ max_scales=1010, # never used
253
+ max_frames=int(self.video_frames / other_args.temporal_compress_rate + 1),
254
+ max_height=1800 // 8,
255
+ max_width=1800 // 8,
256
+ text_maxlen=self.text_maxlen,
257
+ args=other_args,
258
+ )
259
+
260
+ self.word_embed = nn.Linear(self.visual_embedding_in_dim, self.embed_dim)
261
+ self.head = FsqHead(
262
+ hidden_dim=self.embed_dim,
263
+ fsq_dim=classifier_head_dim,
264
+ fsq_lvl=classifier_head_lvl,
265
+ use_ada_layer_norm=other_args.use_ada_layer_norm,
266
+ )
267
+
268
+ if other_args.add_scale_token > 0:
269
+ self.pt_embedder = TimestepEmbedder(self.embed_dim)
270
+
271
+ # 6. Transformer Blocks
272
+ self.attn_fn_compile_dict = {}
273
+ self.unregistered_blocks = []
274
+ for block_idx in range(depth):
275
+ block = SelfAttnBlock(
276
+ embed_dim=self.embed_dim,
277
+ num_heads=num_heads,
278
+ num_key_value_heads=num_key_value_heads,
279
+ mlp_ratio=mlp_ratio,
280
+ use_flex_attn=use_flex_attn,
281
+ qwen_qkvo_bias=qwen_qkvo_bias,
282
+ use_ada_layer_norm=other_args.use_ada_layer_norm,
283
+ )
284
+ self.unregistered_blocks.append(block)
285
+
286
+ self.num_block_chunks = block_chunks or 1
287
+ self.num_blocks_in_a_chunk = depth // self.num_block_chunks
288
+ assert self.num_blocks_in_a_chunk * self.num_block_chunks == depth, "Depth must be divisible by block_chunks"
289
+
290
+ self.block_chunks = nn.ModuleList([
291
+ MultipleLayers(self.unregistered_blocks, self.num_blocks_in_a_chunk, i * self.num_blocks_in_a_chunk)
292
+ for i in range(self.num_block_chunks)
293
+ ])
294
+
295
+ print(f" [Model Config] embed_dim={embed_dim}, num_heads={num_heads}, depth={depth}, "
296
+ f"mlp_ratio={mlp_ratio}, num_blocks_in_a_chunk={self.num_blocks_in_a_chunk}")
297
+ print(f" drop_path_rate={drop_path_rate:g}", end='\n\n', flush=True)
298
+
299
+ def get_loss_acc(
300
+ self,
301
+ hidden_states: torch.Tensor,
302
+ hidden_states_mask: Optional[torch.Tensor],
303
+ e: Optional[torch.Tensor],
304
+ sequence_packing_scales: List[List[Tuple[int, int, int]]],
305
+ gt: List[torch.Tensor],
306
+ other_info_by_scale: List[Dict[str, Any]],
307
+ return_last_hidden_states: bool
308
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
309
+ """
310
+ Calculate loss and accuracy for the predicted logits.
311
+
312
+ Args:
313
+ hidden_states: shaped (B, L, C)
314
+ hidden_states_mask: Optional mask for hidden states
315
+ e: scale or time embeddings
316
+ sequence_packing_scales: List of scales for sequence packing
317
+ gt: Ground truth labels
318
+ other_info_by_scale: Meta information for each scale
319
+ return_last_hidden_states: Whether to return the last hidden states
320
+
321
+ Returns:
322
+ Tuple of (logits_norm, loss_list, acc_list)
323
+ """
324
+ logits_norm = []
325
+ logits_full = self.head(hidden_states, e)
326
+ global_token_ptr, global_scale_ptr = 0, 0
327
+ loss_list, acc_list = [], []
328
+
329
+ for pack_scales in sequence_packing_scales:
330
+ for pt, ph, pw in pack_scales:
331
+ mul_pt_ph_pw = pt * ph * pw
332
+ cur_bits = other_info_by_scale[global_scale_ptr]['cur_bits']
333
+ cur_lvl = other_info_by_scale[global_scale_ptr]['cur_lvl']
334
+ predict_tokens = other_info_by_scale[global_scale_ptr]['predict_tokens']
335
+ all_tokens = other_info_by_scale[global_scale_ptr]['all_tokens']
336
+ logits = logits_full[:, global_token_ptr:global_token_ptr + predict_tokens]
337
+ logits = logits.reshape(hidden_states.shape[0], mul_pt_ph_pw, cur_bits, cur_lvl)
338
+ logits = logits.permute(0, 3, 1, 2) # [1, num_of_label_value, mul_pt_ph_pw, d]
339
+
340
+ logits_norm.append(logits.abs().mean())
341
+
342
+ # gt[global_scale_ptr]: [1, mul_pt_ph_pw, d]
343
+ loss_this_scale = F.cross_entropy(logits, gt[global_scale_ptr], reduction='none')[0] # [mul_pt_ph_pw, d]
344
+ acc_this_scale = (logits.argmax(1) == gt[global_scale_ptr]).float()[0] # [mul_pt_ph_pw, d]
345
+
346
+ loss_list.append(loss_this_scale.mean(-1))
347
+ acc_list.append(acc_this_scale.mean(-1))
348
+
349
+ global_scale_ptr += 1
350
+ global_token_ptr += all_tokens
351
+
352
+ loss_tensor = torch.cat(loss_list) if loss_list else torch.tensor([], device=hidden_states.device)
353
+ acc_tensor = torch.cat(acc_list) if acc_list else torch.tensor([], device=hidden_states.device)
354
+ logits_norm_tensor = torch.stack(logits_norm).mean() if logits_norm else torch.tensor(0.0, device=hidden_states.device)
355
+
356
+ return logits_norm_tensor, loss_tensor, acc_tensor
357
+
358
+ def get_logits_during_infer(self, hidden_states: torch.Tensor, e: Optional[torch.Tensor] = None) -> torch.Tensor:
359
+ """Get logits during inference."""
360
+ return self.head(hidden_states.float(), e)
361
+
362
+ def forward(
363
+ self,
364
+ label_B_or_BLT: Union[torch.LongTensor, Tuple[torch.FloatTensor, torch.IntTensor, int]],
365
+ x_BLC: torch.Tensor,
366
+ visual_rope_cache: Optional[List[torch.Tensor]] = None,
367
+ sequece_packing_scales: Optional[List[List[Tuple[int, int, int]]]] = None,
368
+ super_scale_lengths: Optional[List[int]] = None,
369
+ other_info_by_scale: Optional[List[Dict[str, Any]]] = None,
370
+ gt_BL: Optional[List[torch.Tensor]] = None,
371
+ x_BLC_mask: Optional[torch.Tensor] = None,
372
+ scale_or_time_ids: Optional[torch.Tensor] = None,
373
+ return_last_hidden_states: bool = False,
374
+ **kwargs: Any,
375
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, float]:
376
+ """
377
+ Forward pass for the GRN model.
378
+
379
+ Args:
380
+ label_B_or_BLT: Text conditions or labels
381
+ x_BLC: Input sequence hidden states
382
+ visual_rope_cache: Cache for visual RoPE embeddings
383
+ sequece_packing_scales: Scales for sequence packing
384
+ super_scale_lengths: Lengths of super scales
385
+ other_info_by_scale: Meta info for scales
386
+ gt_BL: Ground truth
387
+ x_BLC_mask: Mask for input sequence
388
+ scale_or_time_ids: IDs for scale or time embeddings
389
+ return_last_hidden_states: Whether to return last hidden states
390
+
391
+ Returns:
392
+ Tuple of (logits_norm, loss_list, acc_list, valid_sequence_ratio)
393
+ """
394
+ device = x_BLC[0].device
395
+
396
+ # [1. get input sequence x_BLC]
397
+ # word embedding
398
+ sub_L_list = [item.shape[1] for item in x_BLC]
399
+ cat_x_BLC = torch.cat(x_BLC, dim=1)
400
+ with torch.amp.autocast('cuda', dtype=torch.float32):
401
+ cat_x_BLC = self.word_embed(cat_x_BLC.float())
402
+ x_BLC = list(torch.split(cat_x_BLC, sub_L_list, dim=1))
403
+
404
+ # text tokens embedding
405
+ kv_compact, lens, cu_seqlens_k, max_seqlen_k, _ = label_B_or_BLT
406
+ with torch.amp.autocast('cuda', dtype=torch.float32):
407
+ kv_compact = self.text_proj(kv_compact).contiguous() # [sum(lens), C]
408
+ kv_compact_splits = torch.split(kv_compact, lens, dim=0)
409
+
410
+ # scale tokens embedding
411
+ scale_token_ids = torch.tensor([info["scale_token_id"] for info in other_info_by_scale], device=device)
412
+ with torch.amp.autocast("cuda", dtype=torch.float32):
413
+ pt_tokens = self.pt_embedder((scale_token_ids)) # [num_scales, C]
414
+
415
+ # construct final X_BLC input, [visual token, text token, scale token]
416
+ x_BLC_lists = []
417
+ for i in range(len(x_BLC)):
418
+ x_BLC_lists.extend([x_BLC[i], kv_compact_splits[i].unsqueeze(0), pt_tokens[i][None, None]])
419
+ x_BLC = torch.cat(x_BLC_lists, dim=1)
420
+
421
+ valid_sequence_ratio = x_BLC.shape[1] / self.other_args.train_max_token_len
422
+ attn_fn, attn_bias_or_two_vector = None, None
423
+
424
+ # calculate finalrope cache, [visual token, text token, scale token]
425
+ self.rope2d_freqs_grid['freqs_text'] = self.rope2d_freqs_grid['freqs_text'].to(x_BLC.device)
426
+ rope_cache_list = []
427
+ for i in range(len(visual_rope_cache)):
428
+ rope_cache_list.append(visual_rope_cache[i])
429
+ rope_cache_list.append(self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[i]])
430
+ rope_cache_list.append(self.rope2d_freqs_grid['freqs_text'][:,:,:,:,512:512+self.other_args.add_scale_token])
431
+ rope_cache = torch.cat(rope_cache_list, dim=4) # (2, 1, 1, 1, seq_len, head_dim / 2)
432
+ assert rope_cache.shape[4] == x_BLC.shape[1], f'{rope_cache.shape[4]} != {x_BLC.shape[1]}'
433
+ rope_cache = rope_cache[:,0].permute(0, 1, 3, 2, 4) # (2, 1, 1, 1, seq_len, head_dim / 2) -> (2, 1, 1, seq_len, head_dim / 2) -> (2, 1, seq_len, 1, head_dim / 2)
434
+
435
+ # calculate time or scale embeddings
436
+ if self.other_args.use_ada_layer_norm:
437
+ with torch.amp.autocast('cuda', dtype=torch.float32):
438
+ e = self.scale_or_time_embedding(sinusoidal_embedding_1d(self.scale_or_time_dim, scale_or_time_ids).float()) # [1, visual_seq_len,] -> [1, visual_seq_len, 256] -> [1, visual_seq_len, C]
439
+ if e.shape[1] < x_BLC.shape[1]:
440
+ e = F.pad(e, (0,0,0,x_BLC.shape[1]-e.shape[1]), 'constant', 0.) # [1, visual_seq_len, C] -> [1, L, C]
441
+ e0 = self.scale_or_time_projection(e).unflatten(2, (6, self.C)) # [1, L, C] -> [1, L, 6C] -> [1, L, 6, C]
442
+ assert e.dtype == torch.float32 and e0.dtype == torch.float32
443
+ else:
444
+ e, e0 = None, None
445
+
446
+ # [2. block loop]
447
+ checkpointing_full_block = self.checkpointing == 'full-block' and self.training
448
+
449
+ if sp_manager.sp_on():
450
+ # [B, raw_L, C] --> [B, raw_L/sp_size, C]
451
+ x_BLC = sp_split_sequence_by_dim(x_BLC, 1)
452
+
453
+ cu_seqlens = torch.tensor([0]+super_scale_lengths, device=device).cumsum(-1).to(torch.int32)
454
+ max_seqlen = max(super_scale_lengths)
455
+ for i, chunk in enumerate(self.block_chunks): # this path
456
+ x_BLC = chunk(x=x_BLC, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen, e0=e0, attn_bias_or_two_vector=attn_bias_or_two_vector, attn_fn=attn_fn, checkpointing_full_block=checkpointing_full_block, rope2d_freqs_grid=rope_cache)
457
+
458
+ if sp_manager.sp_on():
459
+ # [B, raw_L/sp_size, C] --> [B, raw_L, C]
460
+ x_BLC = sp_gather_sequence_by_dim(x_BLC, 1)
461
+
462
+ # [3. unpad the seqlen dim, and then get logits]
463
+ logits_norm, loss_list, acc_list = self.get_loss_acc(x_BLC, x_BLC_mask, e, sequece_packing_scales, gt_BL, other_info_by_scale, return_last_hidden_states)
464
+ return logits_norm, loss_list, acc_list, valid_sequence_ratio
465
+
466
+ def prepare_text_conditions(
467
+ self,
468
+ label_B_or_BLT: Tuple[torch.Tensor, ...],
469
+ negative_label_B_or_BLT: Optional[Tuple[torch.Tensor, ...]],
470
+ use_cfg: bool = False,
471
+ ) -> Tuple[torch.Tensor, List[int]]:
472
+ """Prepare text conditions for inference."""
473
+ kv_compact, lens, cu_seqlens_k, max_seqlen_k = label_B_or_BLT
474
+ if use_cfg:
475
+ kv_compact_un, lens_un, cu_seqlens_k_un, max_seqlen_k_un = negative_label_B_or_BLT
476
+ kv_compact = torch.cat((kv_compact, kv_compact_un), dim=0)
477
+ cu_seqlens_k = torch.cat((cu_seqlens_k, cu_seqlens_k_un[1:] + cu_seqlens_k[-1]), dim=0)
478
+ max_seqlen_k = max(max_seqlen_k, max_seqlen_k_un)
479
+ lens = lens + lens_un
480
+ kv_compact = self.text_proj(kv_compact).contiguous()
481
+ return kv_compact, lens
482
+
483
+ def embeds_codes2input(self, last_stage: torch.Tensor) -> torch.Tensor:
484
+ """Embed discrete codes into continuous input representations."""
485
+ last_stage = last_stage.reshape(*last_stage.shape[:2], -1) # [B, d, t*h*w] or [B, 4d, t*h*w]
486
+ last_stage = torch.permute(last_stage, [0, 2, 1]) # [B, t*h*w, d] or [B, t*h*w, 4d]
487
+ last_stage = self.word_embed(last_stage) # norm0_ve is Identity
488
+ return last_stage
489
+
490
+ @torch.no_grad()
491
+ def autoregressive_infer(
492
+ self,
493
+ vae: Optional[Any] = None,
494
+ scale_schedule: Optional[List[Tuple[int, int, int]]] = None,
495
+ label_B_or_BLT: Optional[List[Tuple[torch.Tensor, ...]]] = None,
496
+ negative_label_B_or_BLT: Optional[List[Tuple[torch.Tensor, ...]]] = None,
497
+ g_seed: Optional[int] = None,
498
+ cfg_list: Optional[List[float]] = None,
499
+ tau_list: Optional[List[float]] = None,
500
+ gt_leak: int = 0,
501
+ args: Optional[Any] = None,
502
+ get_visual_rope_embeds: Optional[Any] = None,
503
+ noise_list: Optional[List[torch.Tensor]] = None,
504
+ uncond_class_token_id: int = 1000,
505
+ first_frame_condition: bool = False,
506
+ **kwargs: Any,
507
+ ):
508
+ """Autoregressive inference loop for the GRN model."""
509
+ if cfg_list is None: cfg_list = []
510
+ if tau_list is None: tau_list = []
511
+
512
+ from grn.schedules.global_refine import shift_pt
513
+
514
+ rng = None
515
+ assert len(cfg_list) >= len(scale_schedule), "Not enough CFG values for scales"
516
+ assert len(tau_list) >= len(scale_schedule), "Not enough tau values for scales"
517
+
518
+ ret, idx_Bl_list = [], [] # current length, list of reconstructed images
519
+ for b in self.unregistered_blocks: b.attn.kv_caching(True)
520
+ total_steps = args.max_infer_steps
521
+ pbar = tqdm.tqdm(total=total_steps)
522
+ block_chunks = self.block_chunks if self.num_block_chunks > 1 else self.blocks
523
+ use_cfg = True
524
+ cfg_interval = float(args.cfg_type.split('_')[-1])
525
+ full_pt, ph, pw = scale_schedule[0]
526
+ if first_frame_condition:
527
+ pt = full_pt - 1
528
+ visual_rope_cache = get_visual_rope_embeds(self.rope2d_freqs_grid, (pt, ph, pw), 'cuda', args.mapped_h_div_w_template, t_offset=1)
529
+ else:
530
+ pt = full_pt
531
+ visual_rope_cache = get_visual_rope_embeds(self.rope2d_freqs_grid, (pt, ph, pw), 'cuda', args.mapped_h_div_w_template, t_offset=0)
532
+
533
+ # text tokens forward
534
+ self.rope2d_freqs_grid['freqs_text'] = self.rope2d_freqs_grid['freqs_text'].to('cuda')
535
+ prefix_tokens, lens = self.prepare_text_conditions(label_B_or_BLT[0], negative_label_B_or_BLT, use_cfg)
536
+ device = prefix_tokens.device
537
+ infer_device, infer_dtype = prefix_tokens.device, prefix_tokens.dtype
538
+ prefix_tokens = torch.split(prefix_tokens, lens, dim=0)
539
+ rope_cache_text_cond = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[0]]
540
+ rope_cache_text_uncond = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[1]]
541
+
542
+ if args.refine_mode in ['ar_discrete_GRN_bit']:
543
+ classes = 2
544
+ labels_shape = (1,args.detail_scale_dim*args.hbq_round,pt,ph,pw)
545
+ elif args.refine_mode in ['ar_discrete_GRN_index']:
546
+ classes = 2**args.hbq_round
547
+ labels_shape = (1,args.detail_scale_dim,pt,ph,pw)
548
+
549
+ mul_pt_ph_pw = pt * ph * pw
550
+ repeat_idx = -1
551
+ scale_token_rope_cache = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,512:512+args.add_scale_token]
552
+ if noise_list is not None:
553
+ absolute_gt_labels = noise_list[0].to('cuda').permute(0,2,3,4,1) # [B,d,t,h,w] -> [B,t,h,w,d]
554
+ assert len(scale_schedule) == 1
555
+ if first_frame_condition:
556
+ first_frame_labels = noise_list[0][:,:,:1] # [B,d,1,h,w]
557
+ first_frame_tokens_cond = self.embeds_codes2input(multiclass_labels2onehot_input(first_frame_labels, classes))
558
+ fist_frame_rope_cache = get_visual_rope_embeds(self.rope2d_freqs_grid, (1, ph, pw), device, args.mapped_h_div_w_template, t_offset=0)
559
+ visual_rope_cache = torch.cat((visual_rope_cache, fist_frame_rope_cache), dim=4)
560
+ tmp_seqlens = [mul_pt_ph_pw + ph * pw + lens[0] + args.add_scale_token, mul_pt_ph_pw + lens[1] + ph * pw + args.add_scale_token]
561
+ else:
562
+ tmp_seqlens = [mul_pt_ph_pw+lens[0]+args.add_scale_token, mul_pt_ph_pw+lens[1]+args.add_scale_token]
563
+
564
+ # [visual tokens, text tokens, pt tokens]
565
+ rope_cache = torch.cat([visual_rope_cache, rope_cache_text_cond, scale_token_rope_cache, visual_rope_cache, rope_cache_text_uncond, scale_token_rope_cache], dim=4) # (2, 1, 1, 1, seq_len, dim / 2)
566
+ rope_cache = rope_cache[:,0].permute(0, 1, 3, 2, 4) # (2, 1, 1, 1, seq_len, dim / 2) -> (2, 1, 1, seq_len, dim / 2) -> (2, 1, seq_len, 1, dim / 2)
567
+
568
+ cu_seqlens = torch.tensor([0]+tmp_seqlens, device=device).cumsum(-1).to(torch.int32)
569
+ max_seqlen = max(tmp_seqlens)
570
+
571
+ pure_rand_labels = torch.randint(low=0, high=classes, size=labels_shape, device=infer_device, dtype=infer_dtype)
572
+ mixed_xt = pure_rand_labels
573
+ next_pt = 0.
574
+ attn_mask = build_attn_mask(tmp_seqlens, device) if args.use_slow_attn else None
575
+ for cur_inner_round_si in range(args.max_infer_steps):
576
+ cur_pt = next_pt
577
+ is_last_step = np.abs(cur_pt - 1) < 0.02
578
+ if cur_inner_round_si == 0:
579
+ self.entrophy_statistics.append([])
580
+ repeat_idx += 1 # index scale tokens, very important
581
+ cfg = cfg_list[0] if cur_pt >= cfg_interval else 1.0
582
+ last_stage = self.embeds_codes2input(multiclass_labels2onehot_input(mixed_xt, classes))
583
+ pt_tokens = self.pt_embedder(torch.tensor([cur_pt], device=device)).unsqueeze(0)
584
+ # [visual tokens, text tokens, pt tokens]
585
+ if first_frame_condition:
586
+ last_stage_cond = torch.cat((last_stage, first_frame_tokens_cond, prefix_tokens[0].unsqueeze(0), pt_tokens), dim=1)
587
+ last_stage_uncond = torch.cat((last_stage, first_frame_tokens_cond, prefix_tokens[1].unsqueeze(0), pt_tokens), dim=1)
588
+ else:
589
+ last_stage_cond = torch.cat((last_stage, prefix_tokens[0].unsqueeze(0), pt_tokens), dim=1)
590
+ last_stage_uncond = torch.cat((last_stage, prefix_tokens[1].unsqueeze(0), pt_tokens), dim=1)
591
+ last_stage = torch.cat([last_stage_cond, last_stage_uncond], dim=1)
592
+
593
+ e, e0 = None, None
594
+ last_diffusion_step = False
595
+ for block_idx, b in enumerate(block_chunks):
596
+ last_stage = b(x=last_stage, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen, e0=e0, attn_bias_or_two_vector=attn_mask, attn_fn=None, rope2d_freqs_grid=rope_cache, last_diffusion_step=last_diffusion_step)
597
+ logits = self.get_logits_during_infer(last_stage, e=e)
598
+ tmp_bs, tmp_seq_len = logits.shape[:2]
599
+ logits = logits.reshape(tmp_bs, tmp_seq_len, -1, args.detail_num_lvl) # [B,thw+...,d,2]
600
+ pred_cond_logits = logits[:,:mul_pt_ph_pw] # [B,thw,d,2]
601
+ pred_uncond_logits = logits[:,tmp_seqlens[0]:tmp_seqlens[0]+mul_pt_ph_pw] # [B,thw,d,2]
602
+ pred_cond_probs = pred_cond_logits.softmax(-1) # [B,thw,d,2]
603
+ categories = pred_cond_logits.shape[-1]
604
+ entrophy = (-pred_cond_probs * torch.log2(pred_cond_probs)).sum(-1).mean().item() / np.log2(categories)
605
+
606
+ pt_unshift = (cur_inner_round_si + 1) / (args.complexity_aware_Tmax - 1)
607
+ pt_shift = shift_pt(min(1., pt_unshift), args.snr_shift)
608
+ next_pt = 1 - np.cos(np.pi/2*pt_shift)
609
+ next_pt = next_pt * 0.999
610
+
611
+ pred_cond_labels = torch.argmax(pred_cond_probs, dim=-1) # [B,thw,d]
612
+ pred_cond_labels = bld_to_bthwd(pred_cond_labels, pt, ph, pw)
613
+ if cfg != 1:
614
+ pred_cfg_logits = pred_uncond_logits + cfg * (pred_cond_logits - pred_uncond_logits)
615
+ else:
616
+ pred_cfg_logits = pred_cond_logits
617
+ pred_cfg_logits = pred_cfg_logits.mul(1/tau_list[0]) # [B,thw,d,2]
618
+ pred_cfg_probs = pred_cfg_logits.softmax(dim=-1) # [B,thw,d,2]
619
+ pred_cfg_labels = torch.argmax(pred_cfg_probs, dim=-1) # [B,thw,d]
620
+ pred_cfg_labels = bld_to_bthwd(pred_cfg_labels, pt, ph, pw) # [B,t,h,w,d]
621
+ pred_sample_labels = torch.multinomial(pred_cfg_probs.view(-1, args.detail_num_lvl), num_samples=1, replacement=True, generator=rng).view(tmp_bs, mul_pt_ph_pw, -1) # [B, thw,d]
622
+ pred_sample_probs = torch.gather(pred_cfg_probs, dim=3, index=pred_sample_labels.unsqueeze(-1)).squeeze(-1) # [B,thw,d]
623
+ pred_sample_probs = bld_to_bthwd(pred_sample_probs, pt, ph, pw) # [B,t,h,w,d]
624
+ pred_sample_labels = bld_to_bthwd(pred_sample_labels, pt, ph, pw) # [B,t,h,w,d]
625
+
626
+ assume_flip_ratio = (1 - cur_pt) / args.detail_num_lvl * 100. # different ratio between prediciton and input
627
+ pred_zero_ratio = (pred_cond_labels == 0).sum() / pred_cond_labels.numel() * 100.
628
+ pred_one_ratio = (pred_cond_labels == 1).sum() / pred_cond_labels.numel() * 100.
629
+ mixed_xt_Bthwd_01 = mixed_xt.clone().permute(0,2,3,4,1)
630
+ mixed_xt_Bthwd_01[mixed_xt_Bthwd_01<0] = 0
631
+ pred_cond_flip_ratio = (pred_cond_labels != mixed_xt_Bthwd_01).sum() / pred_cond_labels.numel() * 100.
632
+ pred_cfg_flip_ratio = (pred_cfg_labels != mixed_xt_Bthwd_01).sum() / pred_cfg_labels.numel() * 100.
633
+ pred_sample_flip_ratio = (pred_sample_labels != mixed_xt_Bthwd_01).sum() / pred_sample_labels.numel() * 100.
634
+ self.entrophy_statistics[-1].append({
635
+ 'cur_inner_round_si': cur_inner_round_si,
636
+ 'cur_pt': cur_pt,
637
+ # 'cur_tau': cur_tau,
638
+ # 'cur_cfg': cur_cfg,
639
+ 'entrophy': entrophy,
640
+ 'assume_flip_ratio': assume_flip_ratio,
641
+ 'pred_cond_flip_ratio': pred_cond_flip_ratio.item(),
642
+ 'pred_cfg_flip_ratio': pred_cfg_flip_ratio.item(),
643
+ 'pred_sample_flip_ratio': pred_sample_flip_ratio.item(),
644
+ 'pred_zero_ratio': pred_zero_ratio.item(),
645
+ 'pred_one_ratio': pred_one_ratio.item(),
646
+ 'meta': args.meta,
647
+ })
648
+ print(f'{repeat_idx=} {cur_inner_round_si=} {cur_pt=:.3f} {pred_sample_labels.shape=}')
649
+ print(f'{assume_flip_ratio=:.2f}% {pred_cond_flip_ratio=:.2f}% {pred_cfg_flip_ratio=:.2f}% {pred_sample_flip_ratio=:.2f}%')
650
+ if repeat_idx < gt_leak:
651
+ gt_labels = absolute_gt_labels
652
+ gt_flip_ratio = (gt_labels != mixed_xt_Bthwd_01).sum() / gt_labels.numel() * 100.
653
+ gt_flip_ratio = gt_flip_ratio.item()
654
+ pred_cond_acc = (gt_labels==pred_cond_labels).to(float).mean().item()
655
+ pred_cfg_acc = (gt_labels==pred_cfg_labels).to(float).mean().item()
656
+ pred_sample_acc = (gt_labels==pred_sample_labels).to(float).mean().item()
657
+ print(f'{repeat_idx=} {entrophy=:.4f} {pred_cond_acc=:.4f} {pred_cfg_acc=:.4f} {pred_sample_acc=:.4f}')
658
+ self.entrophy_statistics[-1][-1].update({
659
+ 'gt_flip_ratio': gt_flip_ratio,
660
+ 'pred_cond_acc': pred_cond_acc,
661
+ 'pred_cfg_acc': pred_cfg_acc,
662
+ 'pred_sample_acc': pred_sample_acc,
663
+ })
664
+ pred_sample_labels = gt_labels
665
+
666
+ pred_sample_labels = pred_sample_labels.permute(0,4,1,2,3) # [B,t,h,w,d] -> [B,d,t,h,w]
667
+ pred_sample_probs = pred_sample_probs.permute(0,4,1,2,3) # [B,t,h,w,d] -> [B,d,t,h,w]
668
+ use_predict_mask = torch.rand(pred_sample_labels.shape, device=device) < next_pt
669
+ mixed_xt = torch.where(use_predict_mask, pred_sample_labels, pure_rand_labels)
670
+ next_pt = use_predict_mask.float().mean().item()
671
+ pbar.update(1)
672
+ if is_last_step: break
673
+
674
+ if first_frame_condition:
675
+ pred_sample_labels = torch.cat((first_frame_labels, pred_sample_labels), dim=2)
676
+
677
+ if args.refine_mode == 'ar_discrete_GRN_ind':
678
+ from grn.utils_t2iv.hbq_util_t2iv import index_label2quant_features
679
+ approx_signal = index_label2quant_features(pred_sample_labels, hbq_round=args.hbq_round)
680
+ elif args.refine_mode == 'ar_discrete_GRN_bit':
681
+ from grn.utils_t2iv.hbq_util_t2iv import bit_label2raw_feature
682
+ approx_signal = bit_label2raw_feature(pred_sample_labels, hbq_round=args.hbq_round) # [B, hbq_round_mul_d, t, h, w] -> [B,d,t,h,w]
683
+ for b in self.unregistered_blocks: b.attn.kv_caching(False)
684
+ img = self.summed_codes2images(vae, approx_signal)
685
+ return ret, idx_Bl_list, img
686
+
687
+ def summed_codes2images(self, vae: Any, summed_codes: torch.Tensor) -> torch.Tensor:
688
+ """Decode summed codes into images using the VAE."""
689
+ t1 = time.time()
690
+ img = vae.decode(summed_codes, slice=True)
691
+ img = (img + 1) / 2
692
+ img = torch.clamp(img, 0, 1)
693
+ img = img.permute(0, 2, 3, 4, 1) # [bs, 3, t, h, w] -> [bs, t, h, w, 3]
694
+ img = img.mul_(255).to(torch.uint8).flip(dims=(4,))
695
+ print(f"Decode takes {time.time() - t1:.1f}s")
696
+ return img # bgr order
697
+
698
+ @for_visualize
699
+ def vis_key_params(self, ep: int) -> None:
700
+ return
701
+
702
+ def load_state_dict(self, state_dict: Dict[str, Any], strict: bool = False, assign: bool = False) -> Any:
703
+ return super().load_state_dict(state_dict=state_dict, strict=strict, assign=assign)
704
+
705
+ def special_init(self, **kwargs: Any) -> None:
706
+ """Apply special initialization to specific layers."""
707
+ std = 0.02
708
+ for name, module in self.named_modules():
709
+ if isinstance(module, nn.Linear):
710
+ module.weight.data.normal_(mean=0.0, std=std)
711
+ if module.bias is not None:
712
+ module.bias.data.zero_()
713
+ elif isinstance(module, nn.Embedding):
714
+ module.weight.data.normal_(mean=0.0, std=std)
715
+ if module.padding_idx is not None:
716
+ module.weight.data[module.padding_idx].zero_()
717
+
718
+ def extra_repr(self) -> str:
719
+ return f'drop_path_rate={self.drop_path_rate}'
720
+
721
+ def get_layer_id_and_scale_exp(self, para_name: str) -> Any:
722
+ raise NotImplementedError
723
+
724
+ TIMM_KEYS = {'img_size', 'pretrained', 'pretrained_cfg', 'pretrained_cfg_overlay', 'global_pool'}
725
+
726
+ @register_model
727
+ def GRN0b(depth: int = 4, block_chunks: int = 2, embed_dim: int = 512, num_heads: int = 4, num_key_value_heads: int = 4, drop_path_rate: float = 0.0, **kwargs: Any) -> GRN:
728
+ return GRN(
729
+ arch='qwen',
730
+ qwen_qkvo_bias=False,
731
+ depth=depth,
732
+ block_chunks=block_chunks,
733
+ embed_dim=embed_dim,
734
+ num_heads=num_heads,
735
+ num_key_value_heads=num_key_value_heads,
736
+ mlp_ratio=3.55,
737
+ drop_path_rate=drop_path_rate,
738
+ **{k: v for k, v in kwargs.items() if k not in TIMM_KEYS}
739
+ )
740
+
741
+ @register_model
742
+ def GRN2b(depth: int = 28, block_chunks: int = 7, embed_dim: int = 2304, num_heads: int = 18, num_key_value_heads: int = 18, drop_path_rate: float = 0.0, **kwargs: Any) -> GRN:
743
+ return GRN(
744
+ arch='qwen',
745
+ qwen_qkvo_bias=False,
746
+ depth=depth,
747
+ block_chunks=block_chunks,
748
+ embed_dim=embed_dim,
749
+ num_heads=num_heads,
750
+ num_key_value_heads=num_key_value_heads,
751
+ mlp_ratio=3.55,
752
+ drop_path_rate=drop_path_rate,
753
+ **{k: v for k, v in kwargs.items() if k not in TIMM_KEYS}
754
+ )
grn/models/grn_c2i.py ADDED
@@ -0,0 +1,399 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # References:
3
+ # SiT: https://github.com/willisma/SiT
4
+ # Lightning-DiT: https://github.com/hustvl/LightningDiT
5
+ # --------------------------------------------------------
6
+ import torch
7
+ import torch.nn as nn
8
+ import numpy as np
9
+ import math
10
+ import torch.nn.functional as F
11
+ from grn.utils_c2i.model_util import VisionRotaryEmbeddingFast, get_2d_sincos_pos_embed, RMSNorm
12
+ from grn.utils_c2i.hbq_util_c2i import multiclass_labels2onehot_input
13
+
14
+ def modulate(x, shift, scale):
15
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
16
+
17
+
18
+ class BottleneckPatchEmbed(nn.Module):
19
+ """ Image to Patch Embedding
20
+ """
21
+ def __init__(self, img_size=224, patch_size=16, in_chans=3, pca_dim=768, embed_dim=768, bias=True):
22
+ super().__init__()
23
+ img_size = (img_size, img_size)
24
+ patch_size = (patch_size, patch_size)
25
+ num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])
26
+ self.img_size = img_size
27
+ self.patch_size = patch_size
28
+ self.num_patches = num_patches
29
+
30
+ self.proj1 = nn.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
31
+ self.proj2 = nn.Conv2d(pca_dim, embed_dim, kernel_size=1, stride=1, bias=bias)
32
+
33
+ def forward(self, x):
34
+ B, C, H, W = x.shape
35
+ assert H == self.img_size[0] and W == self.img_size[1], \
36
+ f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})."
37
+ x = self.proj2(self.proj1(x)).flatten(2).transpose(1, 2)
38
+ return x
39
+
40
+
41
+ class TimestepEmbedder(nn.Module):
42
+ """
43
+ Embeds scalar timesteps into vector representations.
44
+ """
45
+ def __init__(self, hidden_size, frequency_embedding_size=256):
46
+ super().__init__()
47
+ self.mlp = nn.Sequential(
48
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
49
+ nn.SiLU(),
50
+ nn.Linear(hidden_size, hidden_size, bias=True),
51
+ )
52
+ self.frequency_embedding_size = frequency_embedding_size
53
+
54
+ @staticmethod
55
+ def timestep_embedding(t, dim, max_period=10000):
56
+ """
57
+ Create sinusoidal timestep embeddings.
58
+ :param t: a 1-D Tensor of N indices, one per batch element.
59
+ These may be fractional.
60
+ :param dim: the dimension of the output.
61
+ :param max_period: controls the minimum frequency of the embeddings.
62
+ :return: an (N, D) Tensor of positional embeddings.
63
+ """
64
+ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
65
+ half = dim // 2
66
+ freqs = torch.exp(
67
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
68
+ ).to(device=t.device)
69
+ args = t[:, None].float() * freqs[None]
70
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
71
+ if dim % 2:
72
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
73
+ return embedding
74
+
75
+ def forward(self, t):
76
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
77
+ t_emb = self.mlp(t_freq)
78
+ return t_emb
79
+
80
+
81
+ class LabelEmbedder(nn.Module):
82
+ """
83
+ Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
84
+ """
85
+ def __init__(self, num_classes, hidden_size):
86
+ super().__init__()
87
+ self.embedding_table = nn.Embedding(num_classes + 1, hidden_size)
88
+ self.num_classes = num_classes
89
+
90
+ def forward(self, labels):
91
+ embeddings = self.embedding_table(labels)
92
+ return embeddings
93
+
94
+
95
+ def scaled_dot_product_attention(query, key, value, dropout_p=0.0) -> torch.Tensor:
96
+ L, S = query.size(-2), key.size(-2)
97
+ scale_factor = 1 / math.sqrt(query.size(-1))
98
+ attn_bias = torch.zeros(query.size(0), 1, L, S, dtype=query.dtype).cuda()
99
+
100
+ with torch.cuda.amp.autocast(enabled=False):
101
+ attn_weight = query.float() @ key.float().transpose(-2, -1) * scale_factor
102
+ attn_weight += attn_bias
103
+ attn_weight = torch.softmax(attn_weight, dim=-1)
104
+ attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
105
+ return attn_weight @ value
106
+
107
+
108
+ class Attention(nn.Module):
109
+ def __init__(self, dim, num_heads=8, qkv_bias=True, qk_norm=True, attn_drop=0., proj_drop=0.):
110
+ super().__init__()
111
+ self.num_heads = num_heads
112
+ head_dim = dim // num_heads
113
+
114
+ self.q_norm = RMSNorm(head_dim) if qk_norm else nn.Identity()
115
+ self.k_norm = RMSNorm(head_dim) if qk_norm else nn.Identity()
116
+
117
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
118
+ self.attn_drop = nn.Dropout(attn_drop)
119
+ self.proj = nn.Linear(dim, dim)
120
+ self.proj_drop = nn.Dropout(proj_drop)
121
+
122
+ def forward(self, x, rope):
123
+ B, N, C = x.shape
124
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
125
+ q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
126
+
127
+ q = self.q_norm(q)
128
+ k = self.k_norm(k)
129
+
130
+ q = rope(q)
131
+ k = rope(k)
132
+
133
+ x = scaled_dot_product_attention(q, k, v, dropout_p=self.attn_drop.p if self.training else 0.)
134
+
135
+ x = x.transpose(1, 2).reshape(B, N, C)
136
+
137
+ x = self.proj(x)
138
+ x = self.proj_drop(x)
139
+ return x
140
+
141
+
142
+ class SwiGLUFFN(nn.Module):
143
+ def __init__(
144
+ self,
145
+ dim: int,
146
+ hidden_dim: int,
147
+ drop=0.0,
148
+ bias=True
149
+ ) -> None:
150
+ super().__init__()
151
+ hidden_dim = int(hidden_dim * 2 / 3)
152
+ self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=bias)
153
+ self.w3 = nn.Linear(hidden_dim, dim, bias=bias)
154
+ self.ffn_dropout = nn.Dropout(drop)
155
+
156
+ def forward(self, x):
157
+ x12 = self.w12(x)
158
+ x1, x2 = x12.chunk(2, dim=-1)
159
+ hidden = F.silu(x1) * x2
160
+ return self.w3(self.ffn_dropout(hidden))
161
+
162
+
163
+ class FinalLayer(nn.Module):
164
+ def __init__(self, hidden_size, patch_size, out_channels):
165
+ super().__init__()
166
+ self.norm_final = RMSNorm(hidden_size)
167
+ self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
168
+ self.adaLN_modulation = nn.Sequential(
169
+ nn.SiLU(),
170
+ nn.Linear(hidden_size, 2 * hidden_size, bias=True)
171
+ )
172
+
173
+ @torch.compile
174
+ def forward(self, x, c):
175
+ shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
176
+ x = modulate(self.norm_final(x), shift, scale)
177
+ x = self.linear(x)
178
+ return x
179
+
180
+
181
+ class GRNblock(nn.Module):
182
+ def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, attn_drop=0.0, proj_drop=0.0):
183
+ super().__init__()
184
+ self.norm1 = RMSNorm(hidden_size, eps=1e-6)
185
+ self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=True,
186
+ attn_drop=attn_drop, proj_drop=proj_drop)
187
+ self.norm2 = RMSNorm(hidden_size, eps=1e-6)
188
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
189
+ self.mlp = SwiGLUFFN(hidden_size, mlp_hidden_dim, drop=proj_drop)
190
+ self.adaLN_modulation = nn.Sequential(
191
+ nn.SiLU(),
192
+ nn.Linear(hidden_size, 6 * hidden_size, bias=True)
193
+ )
194
+
195
+ @torch.compile
196
+ def forward(self, x, c, feat_rope=None):
197
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1)
198
+ x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa), rope=feat_rope)
199
+ x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
200
+ return x
201
+
202
+
203
+ class GRN(nn.Module):
204
+ """
205
+ GRN image Transformer.
206
+ """
207
+ def __init__(
208
+ self,
209
+ input_size=256,
210
+ patch_size=16,
211
+ in_channels=3,
212
+ hidden_size=1024,
213
+ depth=24,
214
+ num_heads=16,
215
+ mlp_ratio=4.0,
216
+ attn_drop=0.0,
217
+ proj_drop=0.0,
218
+ num_classes=1000,
219
+ bottleneck_dim=128,
220
+ in_context_len=32,
221
+ in_context_start=8,
222
+ args=None,
223
+ ):
224
+ super().__init__()
225
+ self.in_channels = in_channels
226
+ self.out_channels = in_channels
227
+ self.patch_size = patch_size
228
+ self.num_heads = num_heads
229
+ self.hidden_size = hidden_size
230
+ self.input_size = input_size
231
+ self.in_context_len = in_context_len
232
+ self.in_context_start = in_context_start
233
+ self.num_classes = num_classes
234
+ self.args=args
235
+
236
+ # time and class embed
237
+ self.t_embedder = TimestepEmbedder(hidden_size)
238
+ self.y_embedder = LabelEmbedder(num_classes, hidden_size)
239
+
240
+ # linear embed
241
+ self.x_embedder = BottleneckPatchEmbed(input_size, patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
242
+
243
+ # use fixed sin-cos embedding
244
+ num_patches = self.x_embedder.num_patches
245
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False)
246
+
247
+ # in-context cls token
248
+ if self.in_context_len > 0:
249
+ self.in_context_posemb = nn.Parameter(torch.zeros(1, self.in_context_len, hidden_size), requires_grad=True)
250
+ torch.nn.init.normal_(self.in_context_posemb, std=.02)
251
+
252
+ # rope
253
+ half_head_dim = hidden_size // num_heads // 2
254
+ hw_seq_len = input_size // patch_size
255
+ self.feat_rope = VisionRotaryEmbeddingFast(
256
+ dim=half_head_dim,
257
+ pt_seq_len=hw_seq_len,
258
+ num_cls_token=0
259
+ )
260
+ self.feat_rope_incontext = VisionRotaryEmbeddingFast(
261
+ dim=half_head_dim,
262
+ pt_seq_len=hw_seq_len,
263
+ num_cls_token=self.in_context_len
264
+ )
265
+
266
+ # transformer
267
+ self.blocks = nn.ModuleList([
268
+ GRNblock(hidden_size, num_heads, mlp_ratio=mlp_ratio,
269
+ attn_drop=attn_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0,
270
+ proj_drop=proj_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0)
271
+ for i in range(depth)
272
+ ])
273
+
274
+ # linear predict
275
+ self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
276
+
277
+ self.initialize_weights()
278
+
279
+ def initialize_weights(self):
280
+ # Initialize transformer layers:
281
+ def _basic_init(module):
282
+ if isinstance(module, nn.Linear):
283
+ torch.nn.init.xavier_uniform_(module.weight)
284
+ if module.bias is not None:
285
+ nn.init.constant_(module.bias, 0)
286
+ self.apply(_basic_init)
287
+
288
+ # Initialize (and freeze) pos_embed by sin-cos embedding:
289
+ pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches ** 0.5))
290
+ self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
291
+
292
+ # Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
293
+ w1 = self.x_embedder.proj1.weight.data
294
+ nn.init.xavier_uniform_(w1.view([w1.shape[0], -1]))
295
+ w2 = self.x_embedder.proj2.weight.data
296
+ nn.init.xavier_uniform_(w2.view([w2.shape[0], -1]))
297
+ nn.init.constant_(self.x_embedder.proj2.bias, 0)
298
+
299
+ # Initialize label embedding table:
300
+ nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02)
301
+
302
+ nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
303
+ nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
304
+
305
+ # Zero-out adaLN modulation layers:
306
+ for block in self.blocks:
307
+ nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
308
+ nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
309
+
310
+ # Zero-out output layers:
311
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
312
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
313
+
314
+ nn.init.constant_(self.final_layer.linear.weight, 0)
315
+ nn.init.constant_(self.final_layer.linear.bias, 0)
316
+
317
+ def unpatchify(self, x, p):
318
+ """
319
+ x: (N, T, patch_size**2 * C)
320
+ imgs: (N, H, W, C)
321
+ """
322
+ c = self.out_channels
323
+ h = w = int(x.shape[1] ** 0.5)
324
+ assert h * w == x.shape[1]
325
+
326
+ x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
327
+ x = torch.einsum('nhwpqc->nchpwq', x)
328
+ imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
329
+ return imgs
330
+
331
+ def forward(self, x, t, y):
332
+ """
333
+ x: (N, C, H, W)
334
+ t: (N,)
335
+ y: (N,)
336
+ """
337
+ if self.args.method in ['GRN_ind']:
338
+ x = multiclass_labels2onehot_input(x, 2**self.args.hbq_round)
339
+ elif self.args.method in ['GRN_bit']:
340
+ x = multiclass_labels2onehot_input(x, 2)
341
+
342
+ # class and time embeddings
343
+ t_emb = self.t_embedder(t)
344
+ y_emb = self.y_embedder(y)
345
+ c = t_emb + y_emb
346
+
347
+ # forward
348
+ x = self.x_embedder(x)
349
+ x += self.pos_embed
350
+
351
+ for i, block in enumerate(self.blocks):
352
+ # in-context
353
+ if self.in_context_len > 0 and i == self.in_context_start:
354
+ in_context_tokens = y_emb.unsqueeze(1).repeat(1, self.in_context_len, 1)
355
+ in_context_tokens += self.in_context_posemb
356
+ x = torch.cat([in_context_tokens, x], dim=1)
357
+ x = block(x, c, self.feat_rope if i < self.in_context_start else self.feat_rope_incontext)
358
+
359
+ x = x[:, self.in_context_len:]
360
+
361
+ with torch.amp.autocast('cuda', dtype=torch.float32):
362
+ x = self.final_layer(x, c)
363
+
364
+ if self.args.method == 'GRN_ind':
365
+ B, h_mul_w, classes_mul_d = x.shape
366
+ classes = 2**self.args.hbq_round
367
+ h = w = int(np.round(math.sqrt(h_mul_w)))
368
+ output = x.reshape(B, h, w, classes, classes_mul_d//classes) # [B, h, w, classes, d]
369
+ output = output.permute(0, 3, 4, 1, 2) # [B, classes, d, h, w]
370
+ elif self.args.method == 'GRN_bit':
371
+ B, h_mul_w, classes_mul_d = x.shape
372
+ h = w = int(np.round(math.sqrt(h_mul_w)))
373
+ output = x.reshape(B, h, w, 2, classes_mul_d//2) # [B, h, w, 2, d]
374
+ output = output.permute(0, 3, 4, 1, 2) # [B, 2, d, h, w]
375
+ return output
376
+
377
+
378
+ def GRN_B(**kwargs):
379
+ return GRN(depth=12, hidden_size=768, num_heads=12,
380
+ bottleneck_dim=128, in_context_len=32, in_context_start=4, patch_size=1, **kwargs)
381
+
382
+ def GRN_L(**kwargs):
383
+ return GRN(depth=24, hidden_size=1024, num_heads=16,
384
+ bottleneck_dim=128, in_context_len=32, in_context_start=8, patch_size=1, **kwargs)
385
+
386
+ def GRN_H(**kwargs):
387
+ return GRN(depth=32, hidden_size=1280, num_heads=16,
388
+ bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=1, **kwargs)
389
+
390
+ def GRN_G(**kwargs):
391
+ return GRN(depth=40, hidden_size=1664, num_heads=16,
392
+ bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=1, **kwargs)
393
+
394
+ GRN_models = {
395
+ 'GRN_B': GRN_B,
396
+ 'GRN_L': GRN_L,
397
+ 'GRN_H': GRN_H,
398
+ 'GRN_G': GRN_G,
399
+ }
grn/models/hbq_tokenizer.py ADDED
@@ -0,0 +1,932 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ import os
3
+ import os.path as osp
4
+
5
+ import torch
6
+ import torch.cuda.amp as amp
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ import numpy as np
10
+ from einops import rearrange
11
+
12
+
13
+ CACHE_T = 2
14
+
15
+
16
+ class CausalConv3d(nn.Conv3d):
17
+ """
18
+ Causal 3d convolusion.
19
+ """
20
+
21
+ def __init__(self, *args, **kwargs):
22
+ super().__init__(*args, **kwargs)
23
+ self._padding = (
24
+ self.padding[2],
25
+ self.padding[2],
26
+ self.padding[1],
27
+ self.padding[1],
28
+ 2 * self.padding[0],
29
+ 0,
30
+ )
31
+ self.padding = (0, 0, 0)
32
+
33
+ def forward(self, x, cache_x=None):
34
+ padding = list(self._padding)
35
+ if cache_x is not None and self._padding[4] > 0:
36
+ cache_x = cache_x.to(x.device)
37
+ x = torch.cat([cache_x, x], dim=2)
38
+ padding[4] -= cache_x.shape[2]
39
+ x = F.pad(x, padding)
40
+
41
+ return super().forward(x)
42
+
43
+
44
+ class RMS_norm(nn.Module):
45
+
46
+ def __init__(self, dim, channel_first=True, images=True, bias=False):
47
+ super().__init__()
48
+ broadcastable_dims = (1, 1, 1) if not images else (1, 1)
49
+ shape = (dim, *broadcastable_dims) if channel_first else (dim,)
50
+
51
+ self.channel_first = channel_first
52
+ self.scale = dim**0.5
53
+ self.gamma = nn.Parameter(torch.ones(shape))
54
+ self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
55
+
56
+ def forward(self, x):
57
+ return (F.normalize(x, dim=(1 if self.channel_first else -1)) *
58
+ self.scale * self.gamma + self.bias)
59
+
60
+
61
+ class Upsample(nn.Upsample):
62
+
63
+ def forward(self, x):
64
+ """
65
+ Fix bfloat16 support for nearest neighbor interpolation.
66
+ """
67
+ return super().forward(x.float()).type_as(x)
68
+
69
+
70
+ class Resample(nn.Module):
71
+
72
+ def __init__(self, dim, mode):
73
+ assert mode in (
74
+ "none",
75
+ "upsample2d",
76
+ "upsample3d",
77
+ "downsample2d",
78
+ "downsample3d",
79
+ )
80
+ super().__init__()
81
+ self.dim = dim
82
+ self.mode = mode
83
+
84
+ # layers
85
+ if mode == "upsample2d":
86
+ self.resample = nn.Sequential(
87
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
88
+ nn.Conv2d(dim, dim, 3, padding=1),
89
+ )
90
+ elif mode == "upsample3d":
91
+ self.resample = nn.Sequential(
92
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
93
+ nn.Conv2d(dim, dim, 3, padding=1),
94
+ # nn.Conv2d(dim, dim//2, 3, padding=1)
95
+ )
96
+ self.time_conv = CausalConv3d(
97
+ dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
98
+ elif mode == "downsample2d":
99
+ self.resample = nn.Sequential(
100
+ nn.ZeroPad2d((0, 1, 0, 1)),
101
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
102
+ elif mode == "downsample3d":
103
+ self.resample = nn.Sequential(
104
+ nn.ZeroPad2d((0, 1, 0, 1)),
105
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
106
+ self.time_conv = CausalConv3d(
107
+ dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
108
+ else:
109
+ self.resample = nn.Identity()
110
+
111
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
112
+ b, c, t, h, w = x.size()
113
+ if self.mode == "upsample3d":
114
+ if feat_cache is not None:
115
+ idx = feat_idx[0]
116
+ if feat_cache[idx] is None:
117
+ feat_cache[idx] = "Rep"
118
+ feat_idx[0] += 1
119
+ else:
120
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
121
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
122
+ feat_cache[idx] != "Rep"):
123
+ # cache last frame of last two chunk
124
+ cache_x = torch.cat(
125
+ [
126
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
127
+ cache_x.device),
128
+ cache_x,
129
+ ],
130
+ dim=2,
131
+ )
132
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
133
+ feat_cache[idx] == "Rep"):
134
+ cache_x = torch.cat(
135
+ [
136
+ torch.zeros_like(cache_x).to(cache_x.device),
137
+ cache_x
138
+ ],
139
+ dim=2,
140
+ )
141
+ if feat_cache[idx] == "Rep":
142
+ x = self.time_conv(x)
143
+ else:
144
+ x = self.time_conv(x, feat_cache[idx])
145
+ feat_cache[idx] = cache_x
146
+ feat_idx[0] += 1
147
+ x = x.reshape(b, 2, c, t, h, w)
148
+ x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
149
+ 3)
150
+ x = x.reshape(b, c, t * 2, h, w)
151
+ t = x.shape[2]
152
+ x = rearrange(x, "b c t h w -> (b t) c h w")
153
+ x = self.resample(x)
154
+ x = rearrange(x, "(b t) c h w -> b c t h w", t=t) # this 4 lines do spatial down / up sample
155
+
156
+ if self.mode == "downsample3d":
157
+ if feat_cache is not None:
158
+ idx = feat_idx[0]
159
+ if feat_cache[idx] is None:
160
+ feat_cache[idx] = x.clone()
161
+ feat_idx[0] += 1
162
+ else:
163
+ cache_x = x[:, :, -1:, :, :].clone()
164
+ x = self.time_conv(
165
+ torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
166
+ feat_cache[idx] = cache_x
167
+ feat_idx[0] += 1
168
+ return x
169
+
170
+ def init_weight(self, conv):
171
+ conv_weight = conv.weight.detach().clone()
172
+ nn.init.zeros_(conv_weight)
173
+ c1, c2, t, h, w = conv_weight.size()
174
+ one_matrix = torch.eye(c1, c2)
175
+ init_matrix = one_matrix
176
+ nn.init.zeros_(conv_weight)
177
+ conv_weight.data[:, :, 1, 0, 0] = init_matrix # * 0.5
178
+ conv.weight = nn.Parameter(conv_weight)
179
+ nn.init.zeros_(conv.bias.data)
180
+
181
+ def init_weight2(self, conv):
182
+ conv_weight = conv.weight.data.detach().clone()
183
+ nn.init.zeros_(conv_weight)
184
+ c1, c2, t, h, w = conv_weight.size()
185
+ init_matrix = torch.eye(c1 // 2, c2)
186
+ conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
187
+ conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
188
+ conv.weight = nn.Parameter(conv_weight)
189
+ nn.init.zeros_(conv.bias.data)
190
+
191
+
192
+ class ResidualBlock(nn.Module):
193
+
194
+ def __init__(self, in_dim, out_dim, dropout=0.0):
195
+ super().__init__()
196
+ self.in_dim = in_dim
197
+ self.out_dim = out_dim
198
+
199
+ # layers
200
+ self.residual = nn.Sequential(
201
+ RMS_norm(in_dim, images=False),
202
+ nn.SiLU(),
203
+ CausalConv3d(in_dim, out_dim, 3, padding=1),
204
+ RMS_norm(out_dim, images=False),
205
+ nn.SiLU(),
206
+ nn.Dropout(dropout),
207
+ CausalConv3d(out_dim, out_dim, 3, padding=1),
208
+ )
209
+ self.shortcut = (
210
+ CausalConv3d(in_dim, out_dim, 1)
211
+ if in_dim != out_dim else nn.Identity())
212
+
213
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
214
+ h = self.shortcut(x)
215
+ for layer in self.residual:
216
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
217
+ idx = feat_idx[0]
218
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
219
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
220
+ # cache last frame of last two chunk
221
+ cache_x = torch.cat(
222
+ [
223
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
224
+ cache_x.device),
225
+ cache_x,
226
+ ],
227
+ dim=2,
228
+ )
229
+ x = layer(x, feat_cache[idx])
230
+ feat_cache[idx] = cache_x
231
+ feat_idx[0] += 1
232
+ else:
233
+ x = layer(x)
234
+ return x + h
235
+
236
+
237
+ class AttentionBlock(nn.Module):
238
+ """
239
+ Causal self-attention with a single head.
240
+ """
241
+
242
+ def __init__(self, dim):
243
+ super().__init__()
244
+ self.dim = dim
245
+
246
+ # layers
247
+ self.norm = RMS_norm(dim)
248
+ self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
249
+ self.proj = nn.Conv2d(dim, dim, 1)
250
+
251
+ # zero out the last layer params
252
+ nn.init.zeros_(self.proj.weight)
253
+
254
+ def forward(self, x):
255
+ identity = x
256
+ b, c, t, h, w = x.size()
257
+ x = rearrange(x, "b c t h w -> (b t) c h w")
258
+ x = self.norm(x)
259
+ # compute query, key, value
260
+ q, k, v = (
261
+ self.to_qkv(x).reshape(b * t, 1, c * 3,
262
+ -1).permute(0, 1, 3,
263
+ 2).contiguous().chunk(3, dim=-1))
264
+
265
+ # apply attention
266
+ x = F.scaled_dot_product_attention(
267
+ q,
268
+ k,
269
+ v,
270
+ )
271
+ x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
272
+
273
+ # output
274
+ x = self.proj(x)
275
+ x = rearrange(x, "(b t) c h w-> b c t h w", t=t)
276
+ return x + identity
277
+
278
+
279
+ def patchify(x, patch_size):
280
+ if patch_size == 1:
281
+ return x
282
+ if x.dim() == 4:
283
+ x = rearrange(
284
+ x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
285
+ elif x.dim() == 5:
286
+ x = rearrange(
287
+ x,
288
+ "b c f (h q) (w r) -> b (c r q) f h w",
289
+ q=patch_size,
290
+ r=patch_size,
291
+ )
292
+ else:
293
+ raise ValueError(f"Invalid input shape: {x.shape}")
294
+
295
+ return x
296
+
297
+
298
+ def unpatchify(x, patch_size):
299
+ if patch_size == 1:
300
+ return x
301
+
302
+ if x.dim() == 4:
303
+ x = rearrange(
304
+ x, "b (c r q) h w -> b c (h q) (w r)", q=patch_size, r=patch_size)
305
+ elif x.dim() == 5:
306
+ x = rearrange(
307
+ x,
308
+ "b (c r q) f h w -> b c f (h q) (w r)",
309
+ q=patch_size,
310
+ r=patch_size,
311
+ )
312
+ return x
313
+
314
+
315
+ class AvgDown3D(nn.Module):
316
+
317
+ def __init__(
318
+ self,
319
+ in_channels,
320
+ out_channels,
321
+ factor_t,
322
+ factor_s=1,
323
+ ):
324
+ super().__init__()
325
+ self.in_channels = in_channels
326
+ self.out_channels = out_channels
327
+ self.factor_t = factor_t
328
+ self.factor_s = factor_s
329
+ self.factor = self.factor_t * self.factor_s * self.factor_s
330
+
331
+ assert in_channels * self.factor % out_channels == 0
332
+ self.group_size = in_channels * self.factor // out_channels
333
+
334
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
335
+ pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
336
+ pad = (0, 0, 0, 0, pad_t, 0)
337
+ x = F.pad(x, pad)
338
+ B, C, T, H, W = x.shape
339
+ x = x.view(
340
+ B,
341
+ C,
342
+ T // self.factor_t,
343
+ self.factor_t,
344
+ H // self.factor_s,
345
+ self.factor_s,
346
+ W // self.factor_s,
347
+ self.factor_s,
348
+ )
349
+ x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
350
+ x = x.view(
351
+ B,
352
+ C * self.factor,
353
+ T // self.factor_t,
354
+ H // self.factor_s,
355
+ W // self.factor_s,
356
+ )
357
+ x = x.view(
358
+ B,
359
+ self.out_channels,
360
+ self.group_size,
361
+ T // self.factor_t,
362
+ H // self.factor_s,
363
+ W // self.factor_s,
364
+ )
365
+ x = x.mean(dim=2)
366
+ return x
367
+
368
+
369
+ class DupUp3D(nn.Module):
370
+
371
+ def __init__(
372
+ self,
373
+ in_channels: int,
374
+ out_channels: int,
375
+ factor_t,
376
+ factor_s=1,
377
+ ):
378
+ super().__init__()
379
+ self.in_channels = in_channels
380
+ self.out_channels = out_channels
381
+
382
+ self.factor_t = factor_t
383
+ self.factor_s = factor_s
384
+ self.factor = self.factor_t * self.factor_s * self.factor_s
385
+
386
+ assert out_channels * self.factor % in_channels == 0
387
+ self.repeats = out_channels * self.factor // in_channels
388
+
389
+ def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
390
+ x = x.repeat_interleave(self.repeats, dim=1)
391
+ x = x.view(
392
+ x.size(0),
393
+ self.out_channels,
394
+ self.factor_t,
395
+ self.factor_s,
396
+ self.factor_s,
397
+ x.size(2),
398
+ x.size(3),
399
+ x.size(4),
400
+ )
401
+ x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
402
+ x = x.view(
403
+ x.size(0),
404
+ self.out_channels,
405
+ x.size(2) * self.factor_t,
406
+ x.size(4) * self.factor_s,
407
+ x.size(6) * self.factor_s,
408
+ )
409
+ if first_chunk:
410
+ x = x[:, :, self.factor_t - 1:, :, :]
411
+ return x
412
+
413
+
414
+ class Down_ResidualBlock(nn.Module):
415
+
416
+ def __init__(self,
417
+ in_dim,
418
+ out_dim,
419
+ dropout,
420
+ mult,
421
+ temperal_downsample=False,
422
+ down_flag=False):
423
+ super().__init__()
424
+
425
+ # Shortcut path with downsample
426
+ self.avg_shortcut = AvgDown3D(
427
+ in_dim,
428
+ out_dim,
429
+ factor_t=2 if temperal_downsample else 1,
430
+ factor_s=2 if down_flag else 1,
431
+ )
432
+
433
+ # Main path with residual blocks and downsample
434
+ downsamples = []
435
+ for _ in range(mult): # mult=2, two block
436
+ downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
437
+ in_dim = out_dim
438
+
439
+ # Add the final downsample block
440
+ if down_flag:
441
+ mode = "downsample3d" if temperal_downsample else "downsample2d"
442
+ downsamples.append(Resample(out_dim, mode=mode))
443
+
444
+ self.downsamples = nn.Sequential(*downsamples)
445
+
446
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
447
+ x_copy = x.clone()
448
+ for module in self.downsamples:
449
+ x = module(x, feat_cache, feat_idx)
450
+
451
+ return x + self.avg_shortcut(x_copy)
452
+
453
+
454
+ class Up_ResidualBlock(nn.Module):
455
+
456
+ def __init__(self,
457
+ in_dim,
458
+ out_dim,
459
+ dropout,
460
+ mult,
461
+ temperal_upsample=False,
462
+ up_flag=False):
463
+ super().__init__()
464
+ # Shortcut path with upsample
465
+ if up_flag:
466
+ self.avg_shortcut = DupUp3D(
467
+ in_dim,
468
+ out_dim,
469
+ factor_t=2 if temperal_upsample else 1,
470
+ factor_s=2 if up_flag else 1,
471
+ )
472
+ else:
473
+ self.avg_shortcut = None
474
+
475
+ # Main path with residual blocks and upsample
476
+ upsamples = []
477
+ for _ in range(mult):
478
+ upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
479
+ in_dim = out_dim
480
+
481
+ # Add the final upsample block
482
+ if up_flag:
483
+ mode = "upsample3d" if temperal_upsample else "upsample2d"
484
+ upsamples.append(Resample(out_dim, mode=mode))
485
+
486
+ self.upsamples = nn.Sequential(*upsamples)
487
+
488
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
489
+ x_main = x.clone()
490
+ for module in self.upsamples:
491
+ x_main = module(x_main, feat_cache, feat_idx)
492
+ if self.avg_shortcut is not None:
493
+ x_shortcut = self.avg_shortcut(x, first_chunk)
494
+ return x_main + x_shortcut
495
+ else:
496
+ return x_main
497
+
498
+
499
+ class Encoder3d(nn.Module):
500
+
501
+ def __init__(
502
+ self,
503
+ dim=128,
504
+ z_dim=4,
505
+ dim_mult=[1, 2, 4, 4],
506
+ num_res_blocks=2,
507
+ attn_scales=[],
508
+ temperal_downsample=[True, True, False],
509
+ dropout=0.0,
510
+ ):
511
+ super().__init__()
512
+ self.dim = dim
513
+ self.z_dim = z_dim
514
+ self.dim_mult = dim_mult
515
+ self.num_res_blocks = num_res_blocks
516
+ self.attn_scales = attn_scales
517
+ self.temperal_downsample = temperal_downsample
518
+
519
+ # dimensions
520
+ dims = [dim * u for u in [1] + dim_mult] # [1,2,4,4] -> [1,1,2,4,4] -> [128,128,256,512,512]
521
+
522
+ # init block
523
+ self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
524
+
525
+ # downsample blocks
526
+ downsamples = []
527
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
528
+ t_down_flag = (
529
+ temperal_downsample[i]
530
+ if i < len(temperal_downsample) else False)
531
+ downsamples.append(
532
+ Down_ResidualBlock(
533
+ in_dim=in_dim,
534
+ out_dim=out_dim,
535
+ dropout=dropout,
536
+ mult=num_res_blocks,
537
+ temperal_downsample=t_down_flag,
538
+ down_flag=i != len(dim_mult) - 1,
539
+ ))
540
+ self.downsamples = nn.Sequential(*downsamples)
541
+
542
+ # middle blocks
543
+ self.middle = nn.Sequential(
544
+ ResidualBlock(out_dim, out_dim, dropout),
545
+ AttentionBlock(out_dim),
546
+ ResidualBlock(out_dim, out_dim, dropout),
547
+ )
548
+
549
+ # # output blocks
550
+ self.head = nn.Sequential(
551
+ RMS_norm(out_dim, images=False),
552
+ nn.SiLU(),
553
+ CausalConv3d(out_dim, z_dim, 3, padding=1),
554
+ )
555
+
556
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
557
+
558
+ if feat_cache is not None:
559
+ idx = feat_idx[0]
560
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
561
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
562
+ cache_x = torch.cat(
563
+ [
564
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
565
+ cache_x.device),
566
+ cache_x,
567
+ ],
568
+ dim=2,
569
+ )
570
+ x = self.conv1(x, feat_cache[idx])
571
+ feat_cache[idx] = cache_x
572
+ feat_idx[0] += 1
573
+ else:
574
+ x = self.conv1(x)
575
+
576
+ ## downsamples
577
+ for layer in self.downsamples:
578
+ if feat_cache is not None:
579
+ x = layer(x, feat_cache, feat_idx)
580
+ else:
581
+ x = layer(x)
582
+
583
+ ## middle
584
+ for layer in self.middle:
585
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
586
+ x = layer(x, feat_cache, feat_idx)
587
+ else:
588
+ x = layer(x)
589
+
590
+ ## head
591
+ for layer in self.head:
592
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
593
+ idx = feat_idx[0]
594
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
595
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
596
+ cache_x = torch.cat(
597
+ [
598
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
599
+ cache_x.device),
600
+ cache_x,
601
+ ],
602
+ dim=2,
603
+ )
604
+ x = layer(x, feat_cache[idx])
605
+ feat_cache[idx] = cache_x
606
+ feat_idx[0] += 1
607
+ else:
608
+ x = layer(x)
609
+
610
+ return x
611
+
612
+
613
+ class Decoder3d(nn.Module):
614
+
615
+ def __init__(
616
+ self,
617
+ dim=128,
618
+ z_dim=4,
619
+ dim_mult=[1, 2, 4, 4],
620
+ num_res_blocks=2,
621
+ attn_scales=[],
622
+ temperal_upsample=[False, True, True],
623
+ dropout=0.0,
624
+ ):
625
+ super().__init__()
626
+ self.dim = dim
627
+ self.z_dim = z_dim
628
+ self.dim_mult = dim_mult
629
+ self.num_res_blocks = num_res_blocks
630
+ self.attn_scales = attn_scales
631
+ self.temperal_upsample = temperal_upsample
632
+
633
+ # dimensions
634
+ dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
635
+ # scale = 1.0 / 2**(len(dim_mult) - 2)
636
+ # init block
637
+ self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
638
+
639
+ # middle blocks
640
+ self.middle = nn.Sequential(
641
+ ResidualBlock(dims[0], dims[0], dropout),
642
+ AttentionBlock(dims[0]),
643
+ ResidualBlock(dims[0], dims[0], dropout),
644
+ )
645
+
646
+ # upsample blocks
647
+ upsamples = []
648
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
649
+ t_up_flag = temperal_upsample[i] if i < len(
650
+ temperal_upsample) else False
651
+ upsamples.append(
652
+ Up_ResidualBlock(
653
+ in_dim=in_dim,
654
+ out_dim=out_dim,
655
+ dropout=dropout,
656
+ mult=num_res_blocks + 1,
657
+ temperal_upsample=t_up_flag,
658
+ up_flag=i != len(dim_mult) - 1,
659
+ ))
660
+ self.upsamples = nn.Sequential(*upsamples)
661
+
662
+ # output blocks
663
+ self.head = nn.Sequential(
664
+ RMS_norm(out_dim, images=False),
665
+ nn.SiLU(),
666
+ CausalConv3d(out_dim, 12, 3, padding=1),
667
+ )
668
+
669
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
670
+ if feat_cache is not None:
671
+ idx = feat_idx[0]
672
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
673
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
674
+ cache_x = torch.cat(
675
+ [
676
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
677
+ cache_x.device),
678
+ cache_x,
679
+ ],
680
+ dim=2,
681
+ )
682
+ x = self.conv1(x, feat_cache[idx])
683
+ feat_cache[idx] = cache_x
684
+ feat_idx[0] += 1
685
+ else:
686
+ x = self.conv1(x)
687
+
688
+ for layer in self.middle:
689
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
690
+ x = layer(x, feat_cache, feat_idx)
691
+ else:
692
+ x = layer(x)
693
+
694
+ ## upsamples
695
+ for layer in self.upsamples:
696
+ if feat_cache is not None:
697
+ x = layer(x, feat_cache, feat_idx, first_chunk)
698
+ else:
699
+ x = layer(x)
700
+
701
+ ## head
702
+ for layer in self.head:
703
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
704
+ idx = feat_idx[0]
705
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
706
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
707
+ cache_x = torch.cat(
708
+ [
709
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
710
+ cache_x.device),
711
+ cache_x,
712
+ ],
713
+ dim=2,
714
+ )
715
+ x = layer(x, feat_cache[idx])
716
+ feat_cache[idx] = cache_x
717
+ feat_idx[0] += 1
718
+ else:
719
+ x = layer(x)
720
+ return x
721
+
722
+
723
+ def count_conv3d(model):
724
+ count = 0
725
+ for m in model.modules():
726
+ if isinstance(m, CausalConv3d):
727
+ count += 1
728
+ return count
729
+
730
+
731
+ class WanVAE_(nn.Module):
732
+
733
+ def __init__(
734
+ self,
735
+ dim=160,
736
+ dec_dim=256,
737
+ z_dim=16,
738
+ dim_mult=[1, 2, 4, 4],
739
+ num_res_blocks=2,
740
+ attn_scales=[],
741
+ temperal_downsample=[True, True, False],
742
+ dropout=0.0,
743
+ ):
744
+ super().__init__()
745
+ self.dim = dim
746
+ self.z_dim = z_dim
747
+ self.dim_mult = dim_mult
748
+ self.num_res_blocks = num_res_blocks
749
+ self.attn_scales = attn_scales
750
+ self.temperal_downsample = temperal_downsample
751
+ self.temperal_upsample = temperal_downsample[::-1]
752
+
753
+ # modules
754
+ self.encoder = Encoder3d(
755
+ dim,
756
+ z_dim * 2,
757
+ dim_mult,
758
+ num_res_blocks,
759
+ attn_scales,
760
+ self.temperal_downsample,
761
+ dropout,
762
+ )
763
+ self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
764
+ self.conv2 = CausalConv3d(z_dim, z_dim, 1)
765
+ self.decoder = Decoder3d(
766
+ dec_dim,
767
+ z_dim,
768
+ dim_mult,
769
+ num_res_blocks,
770
+ attn_scales,
771
+ self.temperal_upsample,
772
+ dropout,
773
+ )
774
+
775
+ def forward(self, x, scale=[0, 1]):
776
+ mu = self.encode(x, scale)
777
+ x_recon = self.decode(mu, scale)
778
+ return x_recon, mu
779
+
780
+ def encode(self, x, scale=None):
781
+ self.clear_cache()
782
+ x = patchify(x, patch_size=2)
783
+ t = x.shape[2]
784
+ iter_ = 1 + (t - 1) // 4
785
+ for i in range(iter_):
786
+ self._enc_conv_idx = [0]
787
+ if i == 0:
788
+ out = self.encoder(
789
+ x[:, :, :1, :, :],
790
+ feat_cache=self._enc_feat_map,
791
+ feat_idx=self._enc_conv_idx,
792
+ )
793
+ else:
794
+ out_ = self.encoder(
795
+ x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
796
+ feat_cache=self._enc_feat_map,
797
+ feat_idx=self._enc_conv_idx,
798
+ )
799
+ out = torch.cat([out, out_], 2)
800
+ mu, log_var = self.conv1(out).chunk(2, dim=1)
801
+ if scale is not None:
802
+ if isinstance(scale[0], torch.Tensor):
803
+ mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
804
+ 1, self.z_dim, 1, 1, 1)
805
+ else:
806
+ mu = (mu - scale[0]) * scale[1]
807
+ self.clear_cache()
808
+ if self.encoder_out_type == 'feature_tanh':
809
+ mu = torch.tanh(mu)
810
+ elif self.encoder_out_type == 'feature':
811
+ pass
812
+ else:
813
+ raise ValueError(f'{self.encoder_out_type=} is not supported!')
814
+ return mu
815
+
816
+ def decode(self, z, scale=None, **kwargs):
817
+ self.clear_cache()
818
+ if scale is not None:
819
+ if isinstance(scale[0], torch.Tensor):
820
+ z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
821
+ 1, self.z_dim, 1, 1, 1)
822
+ else:
823
+ z = z / scale[1] + scale[0]
824
+ x = self.conv2(z)
825
+ iter_ = z.shape[2]
826
+ for i in range(iter_):
827
+ self._conv_idx = [0]
828
+ if i == 0:
829
+ out = self.decoder(
830
+ x[:, :, i:i + 1, :, :],
831
+ feat_cache=self._feat_map,
832
+ feat_idx=self._conv_idx,
833
+ first_chunk=True,
834
+ )
835
+ else:
836
+ out_ = self.decoder(
837
+ x[:, :, i:i + 1, :, :],
838
+ feat_cache=self._feat_map,
839
+ feat_idx=self._conv_idx,
840
+ )
841
+ out = torch.cat([out, out_], 2)
842
+ out = unpatchify(out, patch_size=2)
843
+ self.clear_cache()
844
+ return out
845
+
846
+ def reparameterize(self, mu, log_var):
847
+ std = torch.exp(0.5 * log_var)
848
+ eps = torch.randn_like(std)
849
+ return eps * std + mu
850
+
851
+ def sample(self, imgs, deterministic=False):
852
+ import pdb; pdb.set_trace()
853
+ mu, log_var = self.encode(imgs)
854
+ if deterministic:
855
+ return mu
856
+ std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
857
+ return mu + std * torch.randn_like(std)
858
+
859
+ def clear_cache(self):
860
+ self._conv_num = count_conv3d(self.decoder)
861
+ self._conv_idx = [0]
862
+ self._feat_map = [None] * self._conv_num
863
+ # cache encode
864
+ self._enc_conv_num = count_conv3d(self.encoder)
865
+ self._enc_conv_idx = [0]
866
+ self._enc_feat_map = [None] * self._enc_conv_num
867
+
868
+ class HBQ_Tokenizer(WanVAE_):
869
+ def __init__(
870
+ self,
871
+ args,
872
+ dim=160,
873
+ dec_dim=256,
874
+ latent_channels=16,
875
+ dim_mult=[1, 2, 4, 4],
876
+ num_res_blocks=2,
877
+ temperal_downsample=[0,1,1],
878
+ dropout=0.,
879
+ encoder_out_type='',
880
+ ):
881
+ super().__init__(
882
+ dim=dim,
883
+ dec_dim=dec_dim,
884
+ z_dim=latent_channels,
885
+ dim_mult=dim_mult,
886
+ num_res_blocks=num_res_blocks,
887
+ temperal_downsample=temperal_downsample,
888
+ dropout=dropout,
889
+ )
890
+ self.other_args = args
891
+ self.codebook_dim = latent_channels
892
+ self.encoder_out_type = encoder_out_type
893
+
894
+ def encode_for_raw_features(
895
+ self, x: torch.Tensor,
896
+ **kwargs,
897
+ ):
898
+ is_image = x.ndim == 4
899
+ if not is_image:
900
+ B, C, T, H, W = x.shape
901
+ else:
902
+ B, C, H, W = x.shape
903
+ T = 1
904
+ x = x.unsqueeze(2)
905
+ with torch.amp.autocast("cuda", dtype=torch.float):
906
+ z = self.encode(x)
907
+ return [z], None, None
908
+
909
+ def _video_vae(pretrained_path=None, z_dim=16, dim=160, device="cpu", **kwargs):
910
+ # params
911
+ cfg = dict(
912
+ dim=dim,
913
+ z_dim=z_dim,
914
+ dim_mult=[1, 2, 4, 4],
915
+ num_res_blocks=2,
916
+ attn_scales=[],
917
+ temperal_downsample=[True, True, True],
918
+ dropout=0.0,
919
+ )
920
+ cfg.update(**kwargs)
921
+
922
+ # init model
923
+ with torch.device("meta"):
924
+ model = WanVAE_(**cfg)
925
+
926
+ # load checkpoint
927
+ logging.info(f"loading {pretrained_path}")
928
+ model.load_state_dict(
929
+ torch.load(pretrained_path, map_location=device), assign=True)
930
+
931
+ return model
932
+
grn/models/init_param.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.nn as nn
2
+
3
+
4
+ def init_weights(model: nn.Module, conv_std_or_gain: float = 0.02, other_std: float = 0.02):
5
+ """
6
+ :param model: the model to be inited
7
+ :param conv_std_or_gain: how to init every conv layer `m`
8
+ > 0: nn.init.trunc_normal_(m.weight.data, std=conv_std_or_gain)
9
+ < 0: nn.init.xavier_normal_(m.weight.data, gain=-conv_std_or_gain)
10
+ :param other_std: how to init every linear layer or embedding layer
11
+ use nn.init.trunc_normal_(m.weight.data, std=other_std)
12
+ """
13
+ skip = abs(conv_std_or_gain) > 10
14
+ if skip: return
15
+ print(f'[init_weights] {type(model).__name__} with {"std" if conv_std_or_gain > 0 else "gain"}={abs(conv_std_or_gain):g}')
16
+ for m in model.modules():
17
+ if isinstance(m, nn.Linear):
18
+ nn.init.trunc_normal_(m.weight.data, std=other_std)
19
+ if m.bias is not None:
20
+ nn.init.constant_(m.bias.data, 0.)
21
+ elif isinstance(m, nn.Embedding):
22
+ nn.init.trunc_normal_(m.weight.data, std=other_std)
23
+ if m.padding_idx is not None:
24
+ m.weight.data[m.padding_idx].zero_()
25
+ elif isinstance(m, (nn.Conv1d, nn.Conv2d, nn.ConvTranspose1d, nn.ConvTranspose2d)):
26
+ nn.init.trunc_normal_(m.weight.data, std=conv_std_or_gain) if conv_std_or_gain > 0 else nn.init.xavier_normal_(m.weight.data, gain=-conv_std_or_gain) # todo: StyleSwin: (..., gain=.02)
27
+ if hasattr(m, 'bias') and m.bias is not None:
28
+ nn.init.constant_(m.bias.data, 0.)
29
+ elif isinstance(m, (nn.LayerNorm, nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm, nn.GroupNorm, nn.InstanceNorm1d, nn.InstanceNorm2d, nn.InstanceNorm3d)):
30
+ if m.bias is not None:
31
+ nn.init.constant_(m.bias.data, 0.)
32
+ if m.weight is not None:
33
+ nn.init.constant_(m.weight.data, 1.)
grn/models/rope.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import os
3
+ from functools import partial
4
+ from typing import Optional, Tuple, Union
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ import numpy as np
10
+ from timm.models.layers import DropPath, drop_path
11
+ from torch.utils.checkpoint import checkpoint
12
+
13
+
14
+ def precompute_rope2d_freqs_grid(dim, dynamic_resolution_h_w, rope2d_normalized_by_hw, pad_to_multiplier=1, max_height=2048 // 16, max_width=2048 // 16, base=10000.0, device=None, scaling_factor=1.0, activated_h_div_w_templates=[]):
15
+ # split the dimension into half, one for x and one for y
16
+ half_dim = dim // 2
17
+ inv_freq = 1.0 / (base ** (torch.arange(0, half_dim, 2, dtype=torch.int64).float().to(device) / half_dim)) # namely theta, 1 / (10000^(i/half_dim)), i=0,2,..., half_dim-2
18
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq)
19
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq)
20
+ t_height = t_height / scaling_factor
21
+ freqs_height = torch.outer(t_height, inv_freq) # (max_height, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2), namely y*theta
22
+ t_width = t_width / scaling_factor
23
+ freqs_width = torch.outer(t_width, inv_freq) # (max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2), namely x*theta
24
+ freqs_grid_map = torch.concat([
25
+ freqs_height[:, None, :].expand(-1, max_width, -1), # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2)
26
+ freqs_width[None, :, :].expand(max_height, -1, -1), # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2)
27
+ ], dim=-1) # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d))
28
+ freqs_grid_map = torch.stack([torch.cos(freqs_grid_map), torch.sin(freqs_grid_map)], dim=0)
29
+ # (2, max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d))
30
+
31
+ rope2d_freqs_grid = {}
32
+ for h_div_w in activated_h_div_w_templates:
33
+ assert h_div_w in dynamic_resolution_h_w, f'Unknown h_div_w: {h_div_w}'
34
+ scale_schedule = dynamic_resolution_h_w[h_div_w]['1M']['image_scales']
35
+ _, ph, pw = scale_schedule[-1]
36
+ max_edge_length = freqs_grid_map.shape[1]
37
+ if ph >= pw:
38
+ uph, upw = max_edge_length, int(max_edge_length / ph * pw)
39
+ else:
40
+ uph, upw = int(max_edge_length / pw * ph), max_edge_length
41
+ rope_cache_list = []
42
+ for (_, ph, pw) in scale_schedule:
43
+ ph_mul_pw = ph * pw
44
+ if rope2d_normalized_by_hw == 1: # downsample
45
+ rope_cache = F.interpolate(freqs_grid_map[:, :uph, :upw, :].permute([0,3,1,2]), size=(ph, pw), mode='bilinear', align_corners=True)
46
+ rope_cache = rope_cache.permute([0,2,3,1]) # (2, ph, pw, half_head_dim)
47
+ elif rope2d_normalized_by_hw == 2: # star stylee
48
+ _, uph, upw = scale_schedule[-1]
49
+ indices = torch.stack([
50
+ (torch.arange(ph) * (uph / ph)).reshape(ph, 1).expand(ph, pw),
51
+ (torch.arange(pw) * (upw / pw)).reshape(1, pw).expand(ph, pw),
52
+ ], dim=-1).round().int() # (ph, pw, 2)
53
+ indices = indices.reshape(-1, 2) # (ph*pw, 2)
54
+ rope_cache = freqs_grid_map[:, indices[:,0], indices[:,1], :] # (2, ph*pw, half_head_dim)
55
+ rope_cache = rope_cache.reshape(2, ph, pw, -1)
56
+ elif rope2d_normalized_by_hw == 0:
57
+ rope_cache = freqs_grid_map[:, :ph, :pw, :] # (2, ph, pw, half_head_dim)
58
+ else:
59
+ raise ValueError(f'Unknown rope2d_normalized_by_hw: {rope2d_normalized_by_hw}')
60
+ rope_cache_list.append(rope_cache.reshape(2, ph_mul_pw, -1))
61
+ cat_rope_cache = torch.cat(rope_cache_list, 1) # (2, seq_len, half_head_dim)
62
+ if cat_rope_cache.shape[1] % pad_to_multiplier:
63
+ pad = torch.zeros(2, pad_to_multiplier - cat_rope_cache.shape[1] % pad_to_multiplier, half_dim)
64
+ cat_rope_cache = torch.cat([cat_rope_cache, pad], dim=1)
65
+ cat_rope_cache = cat_rope_cache[:,None,None,None] # (2, 1, 1, 1, seq_len, half_dim)
66
+ for pn in dynamic_resolution_h_w[h_div_w]:
67
+ scale_schedule = dynamic_resolution_h_w[h_div_w][pn]['image_scales']
68
+ tmp_scale_schedule = [(1, h, w) for _, h, w in scale_schedule]
69
+ rope2d_freqs_grid[str(tuple(tmp_scale_schedule))] = cat_rope_cache
70
+ return rope2d_freqs_grid
71
+
72
+
73
+ def precompute_rope3d_freqs_grid(
74
+ dim,
75
+ rope2d_normalized_by_hw,
76
+ max_frames=128,
77
+ max_height=2048 // 8,
78
+ max_width=2048 // 8,
79
+ base=10000.0,
80
+ device=None,
81
+ activated_h_div_w_templates=[],
82
+ text_maxlen=0,
83
+ pn=None,
84
+ args=None,
85
+ **kwargs,
86
+ ):
87
+ # split the dimension into three parts, one for x, one for y, and one for t
88
+ print(f'[precompute_rope4d_freqs_grid: 3d]: start')
89
+ assert dim % 2 == 0, f'Only support dim % 2 == 0, but got dim={dim}'
90
+ dim_div_2 = dim // 2
91
+ num_of_freqs_former = dim_div_2 // 3
92
+ preserve_1d_length = 600
93
+ num_of_freqs_last = dim_div_2 - num_of_freqs_former * 2 # in some cases, dim_div_2 % 3 != 0. here tackle with these cases
94
+ inv_freq_former = 1.0 / (base ** (torch.arange(num_of_freqs_former, dtype=torch.int64).float().to(device) / num_of_freqs_former)) # namely theta, 1 / (10000^(i/dim_div_3)), i=0,2,..., dim_div_3-2, totally dim_div_3 / 2 elems
95
+ inv_freq_last = 1.0 / (base ** (torch.arange(num_of_freqs_last, dtype=torch.int64).float().to(device) / num_of_freqs_last))
96
+ t_frames = torch.arange(preserve_1d_length+max_frames, device=device, dtype=torch.int64).type_as(inv_freq_former)
97
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq_former)
98
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq_former)
99
+ freqs_frames = torch.outer(t_frames, inv_freq_former) # (max_frames, (dim_div_2 / 3)), namely x*theta
100
+ freqs_height = torch.outer(t_height, inv_freq_former) # (max_height, (dim_div_2 / 3), namely y*theta
101
+ freqs_width = torch.outer(t_width, inv_freq_last) # (max_width, (dim_div_2 / 3)), namely x*theta
102
+ freqs_frames = torch.stack([torch.cos(freqs_frames), torch.sin(freqs_frames)], dim=0)
103
+ freqs_height = torch.stack([torch.cos(freqs_height), torch.sin(freqs_height)], dim=0)
104
+ freqs_width = torch.stack([torch.cos(freqs_width), torch.sin(freqs_width)], dim=0)
105
+ tm = preserve_1d_length
106
+ rope_text_embeds = torch.cat([
107
+ freqs_frames[ :, :tm, None, None, :].expand(-1, -1, -1, -1, -1),
108
+ freqs_height[ :, None, :1, None, :].expand(-1, tm, -1, -1, -1),
109
+ freqs_width[ :, None, None, :1, :].expand(-1, tm, -1, -1, -1),
110
+ ], dim=-1) # (2, tm, 1, 1, dim_div_2)
111
+ rope_text_embeds = rope_text_embeds.reshape(2, 1, 1, 1, tm, dim_div_2)
112
+ rope2d_freqs_grid = {}
113
+ rope2d_freqs_grid['freqs_text'] = rope_text_embeds # (2, 1, 1, 1, preserve_1d_length, dim / 2)
114
+ rope2d_freqs_grid['freqs_frames'] = freqs_frames[:, tm:] # (2, max_frames, ceil(dim_div_2 / 4))
115
+ rope2d_freqs_grid['freqs_height'] = freqs_height # (2, max_height, ceil(dim_div_2 / 4))
116
+ rope2d_freqs_grid['freqs_width'] = freqs_width # (2, max_width, ceil(dim_div_2 / 4))
117
+ return rope2d_freqs_grid
118
+
119
+
120
+ def precompute_rope4d_freqs_grid(
121
+ dim,
122
+ rope2d_normalized_by_hw,
123
+ max_scales=128,
124
+ max_frames=128,
125
+ max_height=2048 // 8,
126
+ max_width=2048 // 8,
127
+ base=10000.0,
128
+ device=None,
129
+ activated_h_div_w_templates=[],
130
+ text_maxlen=0,
131
+ pn=None,
132
+ args=None,
133
+ **kwargs,
134
+ ):
135
+ # split the dimension into three parts, one for x, one for y, and one for t
136
+ print(f'[precompute_rope4d_freqs_grid: 4d]: start')
137
+ assert dim % 2 == 0, f'Only support dim % 2 == 0, but got dim={dim}'
138
+ dim_div_2 = dim // 2
139
+ num_of_freqs = int(np.ceil(dim_div_2 / 4))
140
+ inv_freq = 1.0 / (base ** (torch.arange(num_of_freqs, dtype=torch.int64).float().to(device) / num_of_freqs)) # namely theta, 1 / (10000^(i/dim_div_4)), i=0,2,..., dim_div_4-2, totally dim_div_4 / 2 elems
141
+ t_scales = torch.arange(text_maxlen+max_scales, device=device, dtype=torch.int64).type_as(inv_freq)
142
+ t_frames = torch.arange(max_frames, device=device, dtype=torch.int64).type_as(inv_freq)
143
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq)
144
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq)
145
+ freqs_scales = torch.outer(t_scales, inv_freq) # (text_maxlen+max_scales, ceil(dim_div_2 / 4)), namely x*theta
146
+ freqs_frames = torch.outer(t_frames, inv_freq) # (max_frames, ceil(dim_div_2 / 4)), namely x*theta
147
+ freqs_height = torch.outer(t_height, inv_freq) # (max_height, ceil(dim_div_2 / 4)), namely y*theta
148
+ freqs_width = torch.outer(t_width, inv_freq) # (max_width, ceil(dim_div_2 / 4)), namely x*theta
149
+ assert num_of_freqs*4==dim_div_2
150
+ freqs_scales = torch.stack([torch.cos(freqs_scales), torch.sin(freqs_scales)], dim=0)
151
+ freqs_frames = torch.stack([torch.cos(freqs_frames), torch.sin(freqs_frames)], dim=0)
152
+ freqs_height = torch.stack([torch.cos(freqs_height), torch.sin(freqs_height)], dim=0)
153
+ freqs_width = torch.stack([torch.cos(freqs_width), torch.sin(freqs_width)], dim=0)
154
+ tm = text_maxlen
155
+ rope_text_embeds = torch.cat([
156
+ freqs_scales[ :, :tm, None, None, None, :].expand(-1, -1, -1, -1, -1, -1),
157
+ freqs_frames[ :, None, :1, None, None, :].expand(-1, tm, -1, -1, -1, -1),
158
+ freqs_height[ :, None, None, :1, None, :].expand(-1, tm, -1, -1, -1, -1),
159
+ freqs_width[ :, None, None, None, :1, :].expand(-1, tm, -1, -1, -1, -1),
160
+ ], dim=-1) # (2, tm, 1, 1, 1, dim_div_2)
161
+ rope_text_embeds = rope_text_embeds.reshape(2, 1, 1, 1, tm, dim_div_2)
162
+ rope2d_freqs_grid = {}
163
+ rope2d_freqs_grid['freqs_text'] = rope_text_embeds # (2, 1, 1, 1, text_maxlen, dim / 2)
164
+ rope2d_freqs_grid['freqs_scales'] = freqs_scales[:, tm:] # (2, max_scales, ceil(dim_div_2 / 4))
165
+ rope2d_freqs_grid['freqs_frames'] = freqs_frames # (2, max_frames, ceil(dim_div_2 / 4))
166
+ rope2d_freqs_grid['freqs_height'] = freqs_height # (2, max_height, ceil(dim_div_2 / 4))
167
+ rope2d_freqs_grid['freqs_width'] = freqs_width # (2, max_width, ceil(dim_div_2 / 4))
168
+ return rope2d_freqs_grid
169
+
170
+ def apply_rotary_emb(q, k, rope_cache):
171
+ device_type = q.device.type
172
+ device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"
173
+ qk = [q, k]
174
+ rope_cache = rope_cache[:,0]
175
+ with torch.autocast(device_type=device_type, enabled=False):
176
+ for i in range(2):
177
+ qk[i] = qk[i].reshape(*qk[i].shape[:-1], -1, 2)
178
+ tmp1 = qk[i][..., 1] * rope_cache[1]
179
+ tmp2 = qk[i][..., 0] * rope_cache[1]
180
+ qk[i][..., 0].mul_(rope_cache[0]).sub_(tmp1)
181
+ qk[i][..., 1].mul_(rope_cache[0]).add_(tmp2)
182
+ qk[i] = qk[i].reshape(*qk[i].shape[:-2], -1)
183
+ q, k = qk
184
+ # qk = qk.reshape(*qk.shape[:-1], -1, 2) #(2, batch_size, heads, seq_len, half_head_dim, 2)
185
+ # qk = torch.stack([
186
+ # qk[...,0] * rope_cache[0] - qk[...,1] * rope_cache[1],
187
+ # qk[...,0] * rope_cache[1] + qk[...,1] * rope_cache[0],
188
+ # ], dim=-1) # (2, batch_size, heads, seq_len, half_head_dim, 2), here stack + reshape should not be concate
189
+ # qk = qk.reshape(*qk.shape[:-2], -1) #(2, batch_size, heads, seq_len, head_dim)
190
+ # q, k = qk.unbind(dim=0) # (batch_size, heads, seq_len, head_dim)
191
+ return q, k
grn/models/umt5/fsdp.py ADDED
@@ -0,0 +1,42 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gc
2
+ from functools import partial
3
+
4
+ import torch
5
+ from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
6
+ from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
7
+ from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
8
+ from torch.distributed.utils import _free_storage
9
+
10
+
11
+ def shard_model(
12
+ model,
13
+ device_id,
14
+ param_dtype=torch.bfloat16,
15
+ reduce_dtype=torch.float32,
16
+ buffer_dtype=torch.float32,
17
+ process_group=None,
18
+ sharding_strategy=ShardingStrategy.FULL_SHARD,
19
+ sync_module_states=True,
20
+ ):
21
+ model = FSDP(
22
+ module=model,
23
+ process_group=process_group,
24
+ sharding_strategy=sharding_strategy,
25
+ auto_wrap_policy=partial(
26
+ lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
27
+ mixed_precision=MixedPrecision(
28
+ param_dtype=param_dtype,
29
+ reduce_dtype=reduce_dtype,
30
+ buffer_dtype=buffer_dtype),
31
+ device_id=device_id,
32
+ sync_module_states=sync_module_states)
33
+ return model
34
+
35
+
36
+ def free_model(model):
37
+ for m in model.modules():
38
+ if isinstance(m, FSDP):
39
+ _free_storage(m._handle.flat_param.data)
40
+ del model
41
+ gc.collect()
42
+ torch.cuda.empty_cache()
grn/models/umt5/t5.py ADDED
@@ -0,0 +1,514 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ import math
3
+ from functools import partial
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ from grn.models.umt5.fsdp import shard_model
10
+ from grn.models.umt5.umt5_tokenizers import HuggingfaceTokenizer
11
+
12
+ __all__ = [
13
+ 'T5Model',
14
+ 'T5Encoder',
15
+ 'T5Decoder',
16
+ 'T5EncoderModel',
17
+ ]
18
+
19
+
20
+ def fp16_clamp(x):
21
+ if x.dtype == torch.float16 and torch.isinf(x).any():
22
+ clamp = torch.finfo(x.dtype).max - 1000
23
+ x = torch.clamp(x, min=-clamp, max=clamp)
24
+ return x
25
+
26
+
27
+ def init_weights(m):
28
+ if isinstance(m, T5LayerNorm):
29
+ nn.init.ones_(m.weight)
30
+ elif isinstance(m, T5Model):
31
+ nn.init.normal_(m.token_embedding.weight, std=1.0)
32
+ elif isinstance(m, T5FeedForward):
33
+ nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)
34
+ nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)
35
+ nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)
36
+ elif isinstance(m, T5Attention):
37
+ nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)
38
+ nn.init.normal_(m.k.weight, std=m.dim**-0.5)
39
+ nn.init.normal_(m.v.weight, std=m.dim**-0.5)
40
+ nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)
41
+ elif isinstance(m, T5RelativeEmbedding):
42
+ nn.init.normal_(
43
+ m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)
44
+
45
+
46
+ class GELU(nn.Module):
47
+
48
+ def forward(self, x):
49
+ return 0.5 * x * (1.0 + torch.tanh(
50
+ math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
51
+
52
+
53
+ class T5LayerNorm(nn.Module):
54
+
55
+ def __init__(self, dim, eps=1e-6):
56
+ super(T5LayerNorm, self).__init__()
57
+ self.dim = dim
58
+ self.eps = eps
59
+ self.weight = nn.Parameter(torch.ones(dim))
60
+
61
+ def forward(self, x):
62
+ x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +
63
+ self.eps)
64
+ if self.weight.dtype in [torch.float16, torch.bfloat16]:
65
+ x = x.type_as(self.weight)
66
+ return self.weight * x
67
+
68
+
69
+ class T5Attention(nn.Module):
70
+
71
+ def __init__(self, dim, dim_attn, num_heads, dropout=0.1):
72
+ assert dim_attn % num_heads == 0
73
+ super(T5Attention, self).__init__()
74
+ self.dim = dim
75
+ self.dim_attn = dim_attn
76
+ self.num_heads = num_heads
77
+ self.head_dim = dim_attn // num_heads
78
+
79
+ # layers
80
+ self.q = nn.Linear(dim, dim_attn, bias=False)
81
+ self.k = nn.Linear(dim, dim_attn, bias=False)
82
+ self.v = nn.Linear(dim, dim_attn, bias=False)
83
+ self.o = nn.Linear(dim_attn, dim, bias=False)
84
+ self.dropout = nn.Dropout(dropout)
85
+
86
+ def forward(self, x, context=None, mask=None, pos_bias=None):
87
+ """
88
+ x: [B, L1, C].
89
+ context: [B, L2, C] or None.
90
+ mask: [B, L2] or [B, L1, L2] or None.
91
+ """
92
+ # check inputs
93
+ context = x if context is None else context
94
+ b, n, c = x.size(0), self.num_heads, self.head_dim
95
+
96
+ # compute query, key, value
97
+ q = self.q(x).view(b, -1, n, c)
98
+ k = self.k(context).view(b, -1, n, c)
99
+ v = self.v(context).view(b, -1, n, c)
100
+
101
+ # attention bias
102
+ attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
103
+ if pos_bias is not None:
104
+ attn_bias += pos_bias
105
+ if mask is not None:
106
+ assert mask.ndim in [2, 3]
107
+ mask = mask.view(b, 1, 1,
108
+ -1) if mask.ndim == 2 else mask.unsqueeze(1)
109
+ attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
110
+
111
+ # compute attention (T5 does not use scaling)
112
+ attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias
113
+ attn = F.softmax(attn.float(), dim=-1).type_as(attn)
114
+ x = torch.einsum('bnij,bjnc->binc', attn, v)
115
+
116
+ # output
117
+ x = x.reshape(b, -1, n * c)
118
+ x = self.o(x)
119
+ x = self.dropout(x)
120
+ return x
121
+
122
+
123
+ class T5FeedForward(nn.Module):
124
+
125
+ def __init__(self, dim, dim_ffn, dropout=0.1):
126
+ super(T5FeedForward, self).__init__()
127
+ self.dim = dim
128
+ self.dim_ffn = dim_ffn
129
+
130
+ # layers
131
+ self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
132
+ self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
133
+ self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
134
+ self.dropout = nn.Dropout(dropout)
135
+
136
+ def forward(self, x):
137
+ x = self.fc1(x) * self.gate(x)
138
+ x = self.dropout(x)
139
+ x = self.fc2(x)
140
+ x = self.dropout(x)
141
+ return x
142
+
143
+
144
+ class T5SelfAttention(nn.Module):
145
+
146
+ def __init__(self,
147
+ dim,
148
+ dim_attn,
149
+ dim_ffn,
150
+ num_heads,
151
+ num_buckets,
152
+ shared_pos=True,
153
+ dropout=0.1):
154
+ super(T5SelfAttention, self).__init__()
155
+ self.dim = dim
156
+ self.dim_attn = dim_attn
157
+ self.dim_ffn = dim_ffn
158
+ self.num_heads = num_heads
159
+ self.num_buckets = num_buckets
160
+ self.shared_pos = shared_pos
161
+
162
+ # layers
163
+ self.norm1 = T5LayerNorm(dim)
164
+ self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
165
+ self.norm2 = T5LayerNorm(dim)
166
+ self.ffn = T5FeedForward(dim, dim_ffn, dropout)
167
+ self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
168
+ num_buckets, num_heads, bidirectional=True)
169
+
170
+ def forward(self, x, mask=None, pos_bias=None):
171
+ e = pos_bias if self.shared_pos else self.pos_embedding(
172
+ x.size(1), x.size(1))
173
+ x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
174
+ x = fp16_clamp(x + self.ffn(self.norm2(x)))
175
+ return x
176
+
177
+
178
+ class T5CrossAttention(nn.Module):
179
+
180
+ def __init__(self,
181
+ dim,
182
+ dim_attn,
183
+ dim_ffn,
184
+ num_heads,
185
+ num_buckets,
186
+ shared_pos=True,
187
+ dropout=0.1):
188
+ super(T5CrossAttention, self).__init__()
189
+ self.dim = dim
190
+ self.dim_attn = dim_attn
191
+ self.dim_ffn = dim_ffn
192
+ self.num_heads = num_heads
193
+ self.num_buckets = num_buckets
194
+ self.shared_pos = shared_pos
195
+
196
+ # layers
197
+ self.norm1 = T5LayerNorm(dim)
198
+ self.self_attn = T5Attention(dim, dim_attn, num_heads, dropout)
199
+ self.norm2 = T5LayerNorm(dim)
200
+ self.cross_attn = T5Attention(dim, dim_attn, num_heads, dropout)
201
+ self.norm3 = T5LayerNorm(dim)
202
+ self.ffn = T5FeedForward(dim, dim_ffn, dropout)
203
+ self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
204
+ num_buckets, num_heads, bidirectional=False)
205
+
206
+ def forward(self,
207
+ x,
208
+ mask=None,
209
+ encoder_states=None,
210
+ encoder_mask=None,
211
+ pos_bias=None):
212
+ e = pos_bias if self.shared_pos else self.pos_embedding(
213
+ x.size(1), x.size(1))
214
+ x = fp16_clamp(x + self.self_attn(self.norm1(x), mask=mask, pos_bias=e))
215
+ x = fp16_clamp(x + self.cross_attn(
216
+ self.norm2(x), context=encoder_states, mask=encoder_mask))
217
+ x = fp16_clamp(x + self.ffn(self.norm3(x)))
218
+ return x
219
+
220
+
221
+ class T5RelativeEmbedding(nn.Module):
222
+
223
+ def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
224
+ super(T5RelativeEmbedding, self).__init__()
225
+ self.num_buckets = num_buckets
226
+ self.num_heads = num_heads
227
+ self.bidirectional = bidirectional
228
+ self.max_dist = max_dist
229
+
230
+ # layers
231
+ self.embedding = nn.Embedding(num_buckets, num_heads)
232
+
233
+ def forward(self, lq, lk):
234
+ device = self.embedding.weight.device
235
+ # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
236
+ # torch.arange(lq).unsqueeze(1).to(device)
237
+ rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
238
+ torch.arange(lq, device=device).unsqueeze(1)
239
+ rel_pos = self._relative_position_bucket(rel_pos)
240
+ rel_pos_embeds = self.embedding(rel_pos)
241
+ rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
242
+ 0) # [1, N, Lq, Lk]
243
+ return rel_pos_embeds.contiguous()
244
+
245
+ def _relative_position_bucket(self, rel_pos):
246
+ # preprocess
247
+ if self.bidirectional:
248
+ num_buckets = self.num_buckets // 2
249
+ rel_buckets = (rel_pos > 0).long() * num_buckets
250
+ rel_pos = torch.abs(rel_pos)
251
+ else:
252
+ num_buckets = self.num_buckets
253
+ rel_buckets = 0
254
+ rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
255
+
256
+ # embeddings for small and large positions
257
+ max_exact = num_buckets // 2
258
+ rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
259
+ math.log(self.max_dist / max_exact) *
260
+ (num_buckets - max_exact)).long()
261
+ rel_pos_large = torch.min(
262
+ rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
263
+ rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
264
+ return rel_buckets
265
+
266
+
267
+ class T5Encoder(nn.Module):
268
+
269
+ def __init__(self,
270
+ vocab,
271
+ dim,
272
+ dim_attn,
273
+ dim_ffn,
274
+ num_heads,
275
+ num_layers,
276
+ num_buckets,
277
+ shared_pos=True,
278
+ dropout=0.1):
279
+ super(T5Encoder, self).__init__()
280
+ self.dim = dim
281
+ self.dim_attn = dim_attn
282
+ self.dim_ffn = dim_ffn
283
+ self.num_heads = num_heads
284
+ self.num_layers = num_layers
285
+ self.num_buckets = num_buckets
286
+ self.shared_pos = shared_pos
287
+
288
+ # layers
289
+ self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
290
+ else nn.Embedding(vocab, dim)
291
+ self.pos_embedding = T5RelativeEmbedding(
292
+ num_buckets, num_heads, bidirectional=True) if shared_pos else None
293
+ self.dropout = nn.Dropout(dropout)
294
+ self.blocks = nn.ModuleList([
295
+ T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
296
+ shared_pos, dropout) for _ in range(num_layers)
297
+ ])
298
+ self.norm = T5LayerNorm(dim)
299
+
300
+ # initialize weights
301
+ self.apply(init_weights)
302
+
303
+ def forward(self, ids, mask=None):
304
+ x = self.token_embedding(ids)
305
+ x = self.dropout(x)
306
+ e = self.pos_embedding(x.size(1),
307
+ x.size(1)) if self.shared_pos else None
308
+ for block in self.blocks:
309
+ x = block(x, mask, pos_bias=e)
310
+ x = self.norm(x)
311
+ x = self.dropout(x)
312
+ return x
313
+
314
+
315
+ class T5Decoder(nn.Module):
316
+
317
+ def __init__(self,
318
+ vocab,
319
+ dim,
320
+ dim_attn,
321
+ dim_ffn,
322
+ num_heads,
323
+ num_layers,
324
+ num_buckets,
325
+ shared_pos=True,
326
+ dropout=0.1):
327
+ super(T5Decoder, self).__init__()
328
+ self.dim = dim
329
+ self.dim_attn = dim_attn
330
+ self.dim_ffn = dim_ffn
331
+ self.num_heads = num_heads
332
+ self.num_layers = num_layers
333
+ self.num_buckets = num_buckets
334
+ self.shared_pos = shared_pos
335
+
336
+ # layers
337
+ self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
338
+ else nn.Embedding(vocab, dim)
339
+ self.pos_embedding = T5RelativeEmbedding(
340
+ num_buckets, num_heads, bidirectional=False) if shared_pos else None
341
+ self.dropout = nn.Dropout(dropout)
342
+ self.blocks = nn.ModuleList([
343
+ T5CrossAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
344
+ shared_pos, dropout) for _ in range(num_layers)
345
+ ])
346
+ self.norm = T5LayerNorm(dim)
347
+
348
+ # initialize weights
349
+ self.apply(init_weights)
350
+
351
+ def forward(self, ids, mask=None, encoder_states=None, encoder_mask=None):
352
+ b, s = ids.size()
353
+
354
+ # causal mask
355
+ if mask is None:
356
+ mask = torch.tril(torch.ones(1, s, s).to(ids.device))
357
+ elif mask.ndim == 2:
358
+ mask = torch.tril(mask.unsqueeze(1).expand(-1, s, -1))
359
+
360
+ # layers
361
+ x = self.token_embedding(ids)
362
+ x = self.dropout(x)
363
+ e = self.pos_embedding(x.size(1),
364
+ x.size(1)) if self.shared_pos else None
365
+ for block in self.blocks:
366
+ x = block(x, mask, encoder_states, encoder_mask, pos_bias=e)
367
+ x = self.norm(x)
368
+ x = self.dropout(x)
369
+ return x
370
+
371
+
372
+ class T5Model(nn.Module):
373
+
374
+ def __init__(self,
375
+ vocab_size,
376
+ dim,
377
+ dim_attn,
378
+ dim_ffn,
379
+ num_heads,
380
+ encoder_layers,
381
+ decoder_layers,
382
+ num_buckets,
383
+ shared_pos=True,
384
+ dropout=0.1):
385
+ super(T5Model, self).__init__()
386
+ self.vocab_size = vocab_size
387
+ self.dim = dim
388
+ self.dim_attn = dim_attn
389
+ self.dim_ffn = dim_ffn
390
+ self.num_heads = num_heads
391
+ self.encoder_layers = encoder_layers
392
+ self.decoder_layers = decoder_layers
393
+ self.num_buckets = num_buckets
394
+
395
+ # layers
396
+ self.token_embedding = nn.Embedding(vocab_size, dim)
397
+ self.encoder = T5Encoder(self.token_embedding, dim, dim_attn, dim_ffn,
398
+ num_heads, encoder_layers, num_buckets,
399
+ shared_pos, dropout)
400
+ self.decoder = T5Decoder(self.token_embedding, dim, dim_attn, dim_ffn,
401
+ num_heads, decoder_layers, num_buckets,
402
+ shared_pos, dropout)
403
+ self.head = nn.Linear(dim, vocab_size, bias=False)
404
+
405
+ # initialize weights
406
+ self.apply(init_weights)
407
+
408
+ def forward(self, encoder_ids, encoder_mask, decoder_ids, decoder_mask):
409
+ x = self.encoder(encoder_ids, encoder_mask)
410
+ x = self.decoder(decoder_ids, decoder_mask, x, encoder_mask)
411
+ x = self.head(x)
412
+ return x
413
+
414
+
415
+ def _t5(name,
416
+ encoder_only=False,
417
+ decoder_only=False,
418
+ return_tokenizer=False,
419
+ tokenizer_kwargs={},
420
+ dtype=torch.float32,
421
+ device='cpu',
422
+ **kwargs):
423
+ # sanity check
424
+ assert not (encoder_only and decoder_only)
425
+
426
+ # params
427
+ if encoder_only:
428
+ model_cls = T5Encoder
429
+ kwargs['vocab'] = kwargs.pop('vocab_size')
430
+ kwargs['num_layers'] = kwargs.pop('encoder_layers')
431
+ _ = kwargs.pop('decoder_layers')
432
+ elif decoder_only:
433
+ model_cls = T5Decoder
434
+ kwargs['vocab'] = kwargs.pop('vocab_size')
435
+ kwargs['num_layers'] = kwargs.pop('decoder_layers')
436
+ _ = kwargs.pop('encoder_layers')
437
+ else:
438
+ model_cls = T5Model
439
+
440
+ # init model
441
+ with torch.device(device):
442
+ model = model_cls(**kwargs)
443
+
444
+ # set device
445
+ model = model.to(dtype=dtype, device=device)
446
+
447
+ # init tokenizer
448
+ if return_tokenizer:
449
+ from .tokenizers import HuggingfaceTokenizer
450
+ tokenizer = HuggingfaceTokenizer(f'google/{name}', **tokenizer_kwargs)
451
+ return model, tokenizer
452
+ else:
453
+ return model
454
+
455
+
456
+ def umt5_xxl(**kwargs):
457
+ cfg = dict(
458
+ vocab_size=256384,
459
+ dim=4096,
460
+ dim_attn=4096,
461
+ dim_ffn=10240,
462
+ num_heads=64,
463
+ encoder_layers=24,
464
+ decoder_layers=24,
465
+ num_buckets=32,
466
+ shared_pos=False,
467
+ dropout=0.1)
468
+ cfg.update(**kwargs)
469
+ return _t5('umt5-xxl', **cfg)
470
+
471
+
472
+ class T5EncoderModel:
473
+
474
+ def __init__(
475
+ self,
476
+ text_len,
477
+ dtype=torch.bfloat16,
478
+ device=torch.cuda.current_device(),
479
+ checkpoint_path=None,
480
+ tokenizer_path=None,
481
+ enable_fsdp=False,
482
+ ):
483
+ self.text_len = text_len
484
+ self.dtype = dtype
485
+ self.device = device
486
+ self.checkpoint_path = checkpoint_path
487
+ self.tokenizer_path = tokenizer_path
488
+
489
+ # init model
490
+ model = umt5_xxl(
491
+ encoder_only=True,
492
+ return_tokenizer=False,
493
+ dtype=dtype,
494
+ device=device).eval().requires_grad_(False)
495
+ logging.info(f'loading {checkpoint_path}')
496
+ model.load_state_dict(torch.load(checkpoint_path, map_location='cpu'))
497
+ self.model = model
498
+ if enable_fsdp:
499
+ shard_fn = partial(shard_model, device_id=device)
500
+ self.model = shard_fn(self.model, sync_module_states=False)
501
+ else:
502
+ self.model.to(self.device)
503
+ # init tokenizer
504
+ self.tokenizer = HuggingfaceTokenizer(
505
+ name=tokenizer_path, seq_len=text_len, clean='whitespace')
506
+
507
+ def __call__(self, texts, device):
508
+ ids, mask = self.tokenizer(
509
+ texts, return_mask=True, add_special_tokens=True)
510
+ ids = ids.to(device)
511
+ mask = mask.to(device)
512
+ seq_lens = mask.gt(0).sum(dim=1).long()
513
+ context = self.model(ids, mask)
514
+ return [u[:v] for u, v in zip(context, seq_lens)]
grn/models/umt5/umt5_tokenizers.py ADDED
@@ -0,0 +1,81 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import html
2
+ import string
3
+
4
+ import ftfy
5
+ import regex as re
6
+ from transformers import AutoTokenizer
7
+
8
+ __all__ = ['HuggingfaceTokenizer']
9
+
10
+
11
+ def basic_clean(text):
12
+ text = ftfy.fix_text(text)
13
+ text = html.unescape(html.unescape(text))
14
+ return text.strip()
15
+
16
+
17
+ def whitespace_clean(text):
18
+ text = re.sub(r'\s+', ' ', text)
19
+ text = text.strip()
20
+ return text
21
+
22
+
23
+ def canonicalize(text, keep_punctuation_exact_string=None):
24
+ text = text.replace('_', ' ')
25
+ if keep_punctuation_exact_string:
26
+ text = keep_punctuation_exact_string.join(
27
+ part.translate(str.maketrans('', '', string.punctuation))
28
+ for part in text.split(keep_punctuation_exact_string))
29
+ else:
30
+ text = text.translate(str.maketrans('', '', string.punctuation))
31
+ text = text.lower()
32
+ text = re.sub(r'\s+', ' ', text)
33
+ return text.strip()
34
+
35
+
36
+ class HuggingfaceTokenizer:
37
+
38
+ def __init__(self, name, seq_len=None, clean=None, **kwargs):
39
+ assert clean in (None, 'whitespace', 'lower', 'canonicalize')
40
+ self.name = name
41
+ self.seq_len = seq_len
42
+ self.clean = clean
43
+
44
+ # init tokenizer
45
+ self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
46
+ self.vocab_size = self.tokenizer.vocab_size
47
+
48
+ def __call__(self, sequence, **kwargs):
49
+ return_mask = kwargs.pop('return_mask', False)
50
+
51
+ # arguments
52
+ _kwargs = {'return_tensors': 'pt'}
53
+ if self.seq_len is not None:
54
+ _kwargs.update({
55
+ 'padding': 'max_length',
56
+ 'truncation': True,
57
+ 'max_length': self.seq_len
58
+ })
59
+ _kwargs.update(**kwargs)
60
+
61
+ # tokenization
62
+ if isinstance(sequence, str):
63
+ sequence = [sequence]
64
+ if self.clean:
65
+ sequence = [self._clean(u) for u in sequence]
66
+ ids = self.tokenizer(sequence, **_kwargs)
67
+
68
+ # output
69
+ if return_mask:
70
+ return ids.input_ids, ids.attention_mask
71
+ else:
72
+ return ids.input_ids
73
+
74
+ def _clean(self, text):
75
+ if self.clean == 'whitespace':
76
+ text = whitespace_clean(basic_clean(text))
77
+ elif self.clean == 'lower':
78
+ text = whitespace_clean(basic_clean(text)).lower()
79
+ elif self.clean == 'canonicalize':
80
+ text = canonicalize(basic_clean(text))
81
+ return text
grn/schedules/__init__.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ def get_encode_decode_func(dynamic_scale_schedule):
2
+ if 'GRN_vae_stride16' in dynamic_scale_schedule:
3
+ from grn.schedules.global_refine import video_encode, video_decode, get_visual_rope_embeds, get_scale_pack_info
4
+ else:
5
+ raise NotImplementedError(f'{dynamic_scale_schedule} is unsupported')
6
+ return video_encode, video_decode, get_visual_rope_embeds, get_scale_pack_info
grn/schedules/dynamic_resolution.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import math
3
+ import copy
4
+
5
+ import tqdm
6
+ import numpy as np
7
+
8
+
9
+ def get_first_full_spatial_size_scale_index(vae_scale_schedule):
10
+ for si, (pt, ph, pw) in enumerate(vae_scale_schedule):
11
+ if vae_scale_schedule[si][-2:] == vae_scale_schedule[-1][-2:]:
12
+ return si
13
+
14
+ def get_full_spatial_size_scale_indices(vae_scale_schedule):
15
+ full_spatial_size_scale_indices = []
16
+ for si, (pt, ph, pw) in enumerate(vae_scale_schedule):
17
+ if vae_scale_schedule[si][-2:] == vae_scale_schedule[-1][-2:]:
18
+ full_spatial_size_scale_indices.append(si)
19
+ return full_spatial_size_scale_indices
20
+
21
+ def get_ratio2hws_pixels2scales(dynamic_scale_schedule, train_h_div_w_list, video_frames):
22
+ compressed_frames = video_frames // 4 + 1
23
+ if dynamic_scale_schedule in ['GRN_vae_stride16']:
24
+ assert type(train_h_div_w_list) is str
25
+ train_h_div_w_list = json.loads(train_h_div_w_list)
26
+ if len(train_h_div_w_list) == 0:
27
+ train_h_div_w_list = [3/1, 5/2, 2/1, 16/9, 3/2, 4/3, 116/100, 1, 100/116, 3/4, 2/3, 9/16, 1/2, 2/5, 1/3]
28
+ vae_stride = 16
29
+ if 'vae_stride' in dynamic_scale_schedule:
30
+ vae_stride = int(dynamic_scale_schedule.split('vae_stride')[-1])
31
+ dynamic_resolution_h_w = {}
32
+ for h_div_w in train_h_div_w_list:
33
+ ratio = int(h_div_w*1000)/1000
34
+ dynamic_resolution_h_w[ratio] = {}
35
+ for pn in ['0.06M', '0.25M', '0.41M', '0.92M', '1M', '2M']:
36
+ if pn == '0.06M': # 256x256, 192p
37
+ scale = 8
38
+ elif pn == '0.25M': # 512x512, 384p
39
+ scale = 16
40
+ elif pn == '0.41M': # 640x640, 480p
41
+ scale = 20
42
+ elif pn == '0.92M': # 960x960, 720p
43
+ scale = 30
44
+ elif pn == '1M': # 1024x1024, 768p
45
+ scale = 32
46
+ elif pn == '2M': # 1440x1440, 1080p
47
+ scale = 45
48
+ if vae_stride == 16:
49
+ scale = scale * 2
50
+ elif vae_stride == 32:
51
+ scale = scale * 1
52
+ else:
53
+ raise ValueError(f'vae_stride {vae_stride} is not supported')
54
+ area = scale * scale
55
+ pw_float = math.sqrt(area / h_div_w)
56
+ ph_float = pw_float * h_div_w
57
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
58
+ scales = [(ph,pw)]
59
+ pixel = (scales[-1][0] * vae_stride, scales[-1][1] * vae_stride)
60
+ dynamic_resolution_h_w[ratio][pn] = {
61
+ 'pixel': pixel,
62
+ 'scales': scales
63
+ }
64
+ for ratio in dynamic_resolution_h_w:
65
+ for pn in dynamic_resolution_h_w[ratio]:
66
+ base_scale_schedule = dynamic_resolution_h_w[ratio][pn]['scales']
67
+ scales_in_one_clip = len(base_scale_schedule)
68
+ dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'] = {}
69
+ for pt in range(1, compressed_frames+1, 1):
70
+ dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'][pt] = [(pt, h, w) for h, w in base_scale_schedule]
71
+ dynamic_resolution_h_w[ratio][pn]['image_scales'] = scales_in_one_clip
72
+ dynamic_resolution_h_w[ratio][pn]['scales_in_one_clip'] = scales_in_one_clip
73
+ dynamic_resolution_h_w[ratio][pn]['max_video_scales'] = len(dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'][compressed_frames])
74
+ del dynamic_resolution_h_w[ratio][pn]['scales']
75
+ else:
76
+ raise ValueError(f'dynamic_scale_schedule={dynamic_scale_schedule} not implemented')
77
+ return dynamic_resolution_h_w
78
+
79
+ def get_dynamic_resolution_meta(dynamic_scale_schedule, train_h_div_w_list, video_frames):
80
+ dynamic_resolution_h_w = get_ratio2hws_pixels2scales(dynamic_scale_schedule, train_h_div_w_list, video_frames)
81
+ h_div_w_templates = []
82
+ for h_div_w in dynamic_resolution_h_w.keys():
83
+ h_div_w_templates.append(h_div_w)
84
+ h_div_w_templates = np.array(h_div_w_templates)
85
+ return dynamic_resolution_h_w, h_div_w_templates
86
+
87
+ def get_h_div_w_template2indices(h_div_w_list, h_div_w_templates):
88
+ indices = list(range(len(h_div_w_list)))
89
+ h_div_w_template2indices = {}
90
+ pbar = tqdm.tqdm(total=len(indices), desc='get_h_div_w_template2indices...')
91
+ for h_div_w, index in zip(h_div_w_list, indices):
92
+ pbar.update(1)
93
+ nearest_h_div_w_template_ = h_div_w_templates[np.argmin(np.abs(h_div_w-h_div_w_templates))]
94
+ if nearest_h_div_w_template_ not in h_div_w_template2indices:
95
+ h_div_w_template2indices[nearest_h_div_w_template_] = []
96
+ h_div_w_template2indices[nearest_h_div_w_template_].append(index)
97
+ for h_div_w_template_, sub_indices in h_div_w_template2indices.items():
98
+ h_div_w_template2indices[h_div_w_template_] = np.array(sub_indices)
99
+ return h_div_w_template2indices
grn/schedules/global_refine.py ADDED
@@ -0,0 +1,220 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import math
4
+ import bisect
5
+
6
+ import numpy as np
7
+ import torch
8
+ import torch.nn.functional as F
9
+
10
+ from grn.utils_t2iv.hbq_util_t2iv import multiclass_labels2onehot_input
11
+
12
+ def get_scale_pack_info(**kwargs):
13
+ return
14
+
15
+ def flatten_two_level_list(two_level_list):
16
+ flatten_list = []
17
+ for item in two_level_list:
18
+ flatten_list.extend(item)
19
+ return flatten_list
20
+
21
+ def shift_pt(pt, alpha):
22
+ """shift pt (signal ratio) to lower one, recommand alpha=sqrt(height*width/256/256)"""
23
+ if alpha > 1000:
24
+ alpha = alpha - 1000
25
+ noise_pt = 1 - pt
26
+ noise_pt = alpha * noise_pt / (1+(alpha-1)*noise_pt) # shift noise_pt to higer one
27
+ pt = 1 - noise_pt
28
+ return pt
29
+
30
+ def video_encode(
31
+ vae,
32
+ inp_B3HW,
33
+ vae_features=None,
34
+ device='cuda',
35
+ args=None,
36
+ infer_mode=False,
37
+ rope2d_freqs_grid=None,
38
+ dynamic_resolution_h_w=None,
39
+ tokens_remain=9999999,
40
+ text_lens=[],
41
+ caption_nums=[],
42
+ rank_vary_generator=None,
43
+ vis_verbose=False,
44
+ meta_list=None,
45
+ **kwargs,
46
+ ):
47
+ if rank_vary_generator is not None:
48
+ numpy_generator = rank_vary_generator['numpy_generator']
49
+ torch_cuda_generator = rank_vary_generator['torch_cuda_generator']
50
+ else:
51
+ numpy_generator = np.random.default_rng()
52
+ torch_cuda_generator = torch.Generator(device='cuda')
53
+
54
+ if vae_features is None:
55
+ raw_features, _, _ = vae.encode_for_raw_features(inp_B3HW, scale_schedule=None, slice=True)
56
+ raw_features_list = [raw_features]
57
+ x_recon_raw = vae.decode(raw_features[-1], slice=True)
58
+ x_recon_raw = torch.clamp(x_recon_raw, min=-1, max=1)
59
+ print(f'raw_features[-1].shape: {raw_features[-1].shape}')
60
+ else:
61
+ raw_features_list = vae_features
62
+ # raw_features_list: list of [1,d,t,h,w]:
63
+ # import pdb; pdb.set_trace()
64
+ gt_all_bit_indices = []
65
+ pred_all_bit_indices = []
66
+ var_input_list = []
67
+ sequece_packing_scales = [] # with trunk
68
+ h_div_w_template_list = np.array(list(dynamic_resolution_h_w.keys()))
69
+ visual_rope_cache_list = []
70
+ other_info_by_scale = []
71
+ scale_lengths = []
72
+ with torch.amp.autocast('cuda', enabled = False):
73
+ for example_ind, raw_features in enumerate(raw_features_list):
74
+ meta = meta_list[example_ind]
75
+ gt_all_bit_indices.append([])
76
+ pred_all_bit_indices.append([])
77
+ var_input_list.append([])
78
+ visual_rope_cache_list.append([])
79
+ other_info_by_scale.append([])
80
+ B, C, T, H, W = raw_features[-1].shape
81
+ h_div_w = H / W
82
+ mapped_h_div_w_template = h_div_w_template_list[np.argmin(np.abs(h_div_w-h_div_w_template_list))]
83
+ pn = meta['pn']
84
+ if meta['first_frame_condition']:
85
+ scale_schedule = dynamic_resolution_h_w[mapped_h_div_w_template][pn]['pt2scale_schedule'][T-1]
86
+ else:
87
+ scale_schedule = dynamic_resolution_h_w[mapped_h_div_w_template][pn]['pt2scale_schedule'][T]
88
+ if not infer_mode:
89
+ next_tokens_remain = tokens_remain - T * H * W - args.add_scale_token - text_lens[example_ind]
90
+ if next_tokens_remain < 0:
91
+ break
92
+ tokens_remain = next_tokens_remain
93
+ scale_lengths.append(T * H * W + text_lens[example_ind] + args.add_scale_token)
94
+ preserve_scale_schedule = []
95
+ preserve_scale_schedule.append(scale_schedule[0])
96
+ target = raw_features[0]
97
+ if not infer_mode and args.log_norm_sigma > 0:
98
+ spt = torch.sigmoid(torch.randn(1, generator=torch_cuda_generator, device=target.device) * args.log_norm_sigma + args.log_norm_mean).item()
99
+ spt = shift_pt(spt, args.alpha)
100
+ else:
101
+ spt = shift_pt(numpy_generator.random(), args.alpha)
102
+
103
+ if args.refine_mode in ['ar_discrete_GRN_ind']:
104
+ from grn.utils_t2iv.hbq_util_t2iv import raw_feature2index_label
105
+ labels = raw_feature2index_label(target, hbq_round=args.hbq_round) # [B, hbq_round * d, t, h, w]
106
+ classes = 2**args.hbq_round
107
+ elif args.refine_mode in ['ar_discrete_GRN_bit']:
108
+ from grn.utils_t2iv.hbq_util_t2iv import raw_feature2bit_label
109
+ labels = raw_feature2bit_label(target, hbq_round=args.hbq_round) # [B, hbq_round * d, t, h, w]
110
+ classes = 2
111
+
112
+ random_labels = torch.randint(0, classes, size=labels.shape, generator=torch_cuda_generator, device=labels.device, dtype=labels.dtype) # random 0 or 1 labels
113
+ random_mask = torch.rand(size=labels.shape, generator=torch_cuda_generator, device=labels.device, dtype=target.dtype) < spt
114
+ mixed_xt = torch.where(random_mask, labels, random_labels) # [B, hbq_round * d, t, h, w]
115
+ precise_spt = random_mask.float().mean()
116
+ wandb_plot_index = min(9, int(precise_spt / 0.1)) # 0~9
117
+
118
+ # get visual rope
119
+ if not infer_mode:
120
+ if meta['first_frame_condition']:
121
+ visual_rope_cache_list[-1].append(get_visual_rope_embeds(rope2d_freqs_grid, scale_schedule[0], device, mapped_h_div_w_template, t_offset=1)) # (2, 1, 1, 1, pt*ph*pw, dim_div_2)
122
+ visual_rope_cache_list[-1].append(get_visual_rope_embeds(rope2d_freqs_grid, (1, H, W), device, mapped_h_div_w_template, t_offset=0))
123
+ visual_rope_cache_list[-1] = [torch.cat(visual_rope_cache_list[-1], dim=-2)]
124
+ else:
125
+ visual_rope_cache_list[-1].append(get_visual_rope_embeds(rope2d_freqs_grid, scale_schedule[0], device, mapped_h_div_w_template, t_offset=0)) # (2, 1, 1, 1, pt*ph*pw, dim_div_2)
126
+
127
+ # get visual tokens
128
+ visual_token_dim = mixed_xt.shape[1] * classes
129
+
130
+ if not infer_mode and meta['first_frame_condition']:
131
+ first_frame_labels = labels[:,:,:1] # [B, hbq_round * d, 1, h, w]
132
+ first_frame_tokens = multiclass_labels2onehot_input(first_frame_labels, classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
133
+ cur_visual_tokens = multiclass_labels2onehot_input(mixed_xt[:,:,1:], classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
134
+ cur_visual_tokens = torch.cat((cur_visual_tokens, first_frame_tokens), dim=1)
135
+ indices = labels[:,:,1:]
136
+ else:
137
+ cur_visual_tokens = multiclass_labels2onehot_input(mixed_xt, classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
138
+ indices = labels
139
+ indices = indices.type(torch.long).permute(0,2,3,4,1) # [B,d,t,h,w] -> [B,t,h,w,d]
140
+ gt_all_bit_indices[-1].append(indices)
141
+ var_input_list[-1].append(cur_visual_tokens)
142
+ other_info_by_scale[-1].append(
143
+ {
144
+ 'largest_scale': scale_schedule[-1],
145
+ 'wandb_plot_index': wandb_plot_index,
146
+ 'cur_bits': indices.shape[-1],
147
+ 'cur_lvl': args.detail_num_lvl,
148
+ 'scale_token_id': precise_spt,
149
+ 'predict_tokens': np.prod(scale_schedule[0]),
150
+ 'all_tokens': scale_lengths[-1] if len(scale_lengths) else -1,
151
+ 'first_frame_condition': meta['first_frame_condition'],
152
+ }
153
+ )
154
+ sequece_packing_scales.append(preserve_scale_schedule)
155
+
156
+ gt_all_bit_indices = flatten_two_level_list(gt_all_bit_indices)
157
+ pred_all_bit_indices = flatten_two_level_list(pred_all_bit_indices)
158
+ var_input_list = flatten_two_level_list(var_input_list)
159
+ visual_rope_cache_list = flatten_two_level_list(visual_rope_cache_list)
160
+ other_info_by_scale = flatten_two_level_list(other_info_by_scale)
161
+
162
+ if infer_mode:
163
+ return [labels, target], x_recon_raw, [target], None, None, None
164
+
165
+ gt_ms_idx_Bl = []
166
+ for item in gt_all_bit_indices:
167
+ _, tt, hh, ww, dd = item.shape
168
+ item = item.reshape(B, tt*hh*ww, dd)
169
+ gt_ms_idx_Bl.append(item)
170
+ gt_BLC = gt_ms_idx_Bl # torch.cat(gt_ms_idx_Bl, 1).contiguous().type(torch.long)
171
+ x_BLC = var_input_list
172
+ x_BLC_mask = None
173
+ scale_or_time_ids = None
174
+ return x_BLC, x_BLC_mask, scale_or_time_ids, gt_BLC, pred_all_bit_indices, visual_rope_cache_list, sequece_packing_scales, scale_lengths, other_info_by_scale
175
+
176
+ def video_decode(
177
+ vae,
178
+ all_indices,
179
+ scale_schedule,
180
+ label_type,
181
+ args=None,
182
+ noise_list=None,
183
+ trunc_scales=-1,
184
+ **kwargs,
185
+ ):
186
+ if trunc_scales < 0:
187
+ summed_codes = all_indices[-1]
188
+ else:
189
+ summed_codes = all_indices[trunc_scales-1]
190
+ x_recon = vae.decode(summed_codes, slice=True)
191
+ x_recon = torch.clamp(x_recon, min=-1, max=1)
192
+ x_recon_256 = None
193
+ return x_recon, x_recon_256
194
+
195
+ def get_visual_rope_embeds(rope2d_freqs_grid, scale_schedule, device=None, mapped_h_div_w_template=None, t_offset=0):
196
+ # freqs_frames: (2, max_frames, dim_div_2 / 3)
197
+ rope2d_freqs_grid['freqs_frames'] = rope2d_freqs_grid['freqs_frames'].to(device)
198
+ rope2d_freqs_grid['freqs_height'] = rope2d_freqs_grid['freqs_height'].to(device)
199
+ rope2d_freqs_grid['freqs_width'] = rope2d_freqs_grid['freqs_width'].to(device)
200
+ max_height = rope2d_freqs_grid['freqs_height'].shape[1]
201
+ max_width = rope2d_freqs_grid['freqs_width'].shape[1]
202
+ extreme_h_div_w = 3
203
+ assert mapped_h_div_w_template <= extreme_h_div_w
204
+ extreme_h = max_height
205
+ extreme_w = extreme_h / extreme_h_div_w
206
+ upw = np.sqrt(extreme_h * extreme_w / mapped_h_div_w_template)
207
+ uph = mapped_h_div_w_template * upw
208
+ uph, upw = int(uph), int(upw)
209
+ pt, ph, pw = scale_schedule
210
+ assert ph <= uph and pw <= upw
211
+ f_frames = rope2d_freqs_grid['freqs_frames'][:, t_offset:t_offset+pt]
212
+ f_height = rope2d_freqs_grid['freqs_height'][:, (torch.arange(ph) * (uph / ph)).round().int()]
213
+ f_width = rope2d_freqs_grid['freqs_width'][:, (torch.arange(pw) * (upw / pw)).round().int()]
214
+ rope_embeds = torch.cat([
215
+ f_frames[ :, :, None, None, :].expand(-1, -1, ph, pw, -1),
216
+ f_height[ :, None, :, None, :].expand(-1, pt,-1, pw, -1),
217
+ f_width[ :, None, None, :, :].expand(-1, pt,ph, -1, -1),
218
+ ], dim=-1) # (2, pt, ph, pw, dim_div_2)
219
+ rope_embeds = rope_embeds.reshape(2, 1, 1, 1, pt*ph*pw, -1) # (2, 1, 1, 1, pt*ph*pw, dim_div_2)
220
+ return rope_embeds
grn/tokenizer/.gitignore ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ **/__pycache__/
2
+ lightning_logs/
3
+ .ipynb_checkpoints/
4
+ *.egg-info
5
+ .pyc
6
+ results*
7
+ cmp_results*
8
+ logs*
9
+ recon.*
10
+ *.pt
11
+ *.npy
12
+ dataset/
13
+ **/span.log
14
+ cli/batch_post.py
15
+ model_arch*.txt
16
+ *.mp4
17
+ *.png
18
+ labels
19
+ video_vae_results
20
+ video_vae_results_bk
21
+ results
22
+ video_vae_results_full
23
+ bashrc_gpu_worker
24
+ hj_video_vae_results
25
+ wandb
grn/tokenizer/sample.py ADDED
@@ -0,0 +1,694 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import tqdm
3
+ import json
4
+ import re
5
+ import torch
6
+ import torch.nn.functional as F
7
+ import argparse
8
+ import time
9
+ import datetime
10
+ import numpy as np
11
+ import hashlib
12
+ import random
13
+ import torch.nn as nn
14
+ from torchvision.models.inception import inception_v3
15
+ from torch.profiler import record_function as torch_record_function
16
+ from contextlib import nullcontext
17
+ import lpips
18
+ import cv2
19
+ from einops import rearrange
20
+ from tqdm import tqdm
21
+ from PIL import Image
22
+ import os.path as osp
23
+ Image.MAX_IMAGE_PIXELS = None
24
+
25
+ from videovae.modules.commitments import DiagonalGaussianDistribution
26
+
27
+ import torch.distributed as dist
28
+ from torch.multiprocessing import spawn
29
+ from torch.nn.parallel import DistributedDataParallel as DDP
30
+
31
+ import imageio
32
+ import random
33
+ from skimage.metrics import peak_signal_noise_ratio as psnr_loss
34
+ from skimage.metrics import structural_similarity as ssim_loss
35
+
36
+ from videovae.data import VideoData
37
+ from videovae.utils.misc import save_video_grid, shift_dim, data_prefix_manager, rearranged_forward, seed_everything
38
+ from videovae.utils.init_models import init_cnn_from_image, load_cnn
39
+ from videovae.utils.arguments import MainArgs, add_model_specific_args, init_resolution
40
+ from videovae.evaluation import get_fvd_logits, frechet_distance, load_fvd_model
41
+ from videovae.evaluation import calculate_frechet_distance
42
+ from videovae.evaluation import InceptionV3
43
+ from videovae.evaluation import calculate_fvd, calculate_lpips, calculate_psnr, calculate_ssim
44
+
45
+ torch.set_num_threads(32)
46
+ os.environ["NCCL_DEBUG"] = "WARN"
47
+ os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'
48
+
49
+
50
+ def calculate_batch_codebook_usage_percentage(batch_encoding_indices,n_codes):
51
+ if isinstance(batch_encoding_indices, list):
52
+ all_indices = []
53
+ for one_encoding_indices in batch_encoding_indices:
54
+ all_indices.append(one_encoding_indices.flatten())
55
+ all_indices = torch.cat(all_indices, dim=0)
56
+ else:
57
+ # Flatten the batch of encoding indices into a single 1D tensor
58
+ all_indices = batch_encoding_indices.flatten()
59
+ all_indices = all_indices.detach().cpu()
60
+
61
+ # Obtain the total number of encoding indices in the batch to calculate percentages
62
+ total_indices = all_indices.numel()
63
+
64
+ # Initialize a tensor to store the percentage usage of each code
65
+ codebook_usage = torch.zeros(n_codes, dtype=torch.long)
66
+
67
+ # Count the number of occurrences of each index and get their frequency as percentages
68
+ unique_indices, counts = torch.unique(all_indices, return_counts=True)
69
+
70
+ # Populate the corresponding percentages in the codebook_usage_percentage tensor
71
+ codebook_usage[unique_indices.long()] = counts
72
+
73
+ return codebook_usage
74
+
75
+
76
+ def disabled_train(self, mode=True):
77
+ """Overwrite model.train with this function to make sure train/eval mode
78
+ does not change anymore."""
79
+ return self
80
+
81
+ def default_parse_args():
82
+ parser = argparse.ArgumentParser()
83
+ parser.add_argument('--vqgan_ckpt', type=str, default=None)
84
+ parser.add_argument('--sd_ckpt', type=str, default=None)
85
+ parser.add_argument('--use_frames', type=int, default=None)
86
+ parser.add_argument('--inference_type', type=str, choices=["image", "video", "video_concat"])
87
+ parser.add_argument('--save_prediction', action='store_true')
88
+ parser.add_argument('--save_dir', type=str, default="results")
89
+ parser.add_argument('--intermediate_tensor', action='store_true')
90
+ parser.add_argument('--save_z', action='store_true')
91
+ parser.add_argument('--save_frames', action='store_true')
92
+ parser.add_argument('--image_recon4video', action='store_true')
93
+ parser.add_argument('--junke_old', action='store_true')
94
+ parser.add_argument('--cal_norm', action='store_true')
95
+ parser.add_argument('--save_samples', type=str, default=None)
96
+ parser.add_argument('--device', type=str, default="cuda", choices=["cpu", "cuda"])
97
+ parser.add_argument('--noise_scale', type=float, default=0.0)
98
+ parser = MainArgs.add_main_args(parser)
99
+ parser = VideoData.add_data_specific_args(parser)
100
+ args, unknown = parser.parse_known_args()
101
+ args, parser, vae_model = add_model_specific_args(args, parser)
102
+ args = parser.parse_args()
103
+ return args, vae_model
104
+
105
+
106
+ def setup(rank, world_size):
107
+ os.environ['MASTER_ADDR'] = 'localhost'
108
+ os.environ['MASTER_PORT'] = str(12355+int(time.time())%1000)
109
+ # dist.init_process_group("nccl", rank=rank, world_size=world_size)
110
+ dist.init_process_group("nccl", rank=rank, world_size=world_size, timeout=datetime.timedelta(seconds=30 * 60))
111
+
112
+ def cleanup():
113
+ dist.destroy_process_group()
114
+
115
+ def main():
116
+ args, vae_model = default_parse_args()
117
+ assert len(args.dataset_list) == 1
118
+
119
+ # init data_prefix_manager
120
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
121
+ args.default_root_dir = data_prefix_manager(args.default_root_dir)
122
+ os.makedirs(args.default_root_dir, exist_ok=True)
123
+ print(args.default_root_dir)
124
+
125
+ # init intermediate_tensor_dir
126
+ if args.intermediate_tensor:
127
+ random.seed(time.time())
128
+ random_folder_name = hashlib.sha256(str(random.random()).encode('utf-8')).hexdigest()[:16]
129
+ args.intermediate_tensor_dir = os.path.join(args.default_root_dir, random_folder_name)
130
+ print(f"save temporal tensor to {args.intermediate_tensor_dir}")
131
+
132
+ seed_everything(seed=0, allow_tf32=True) # ALERT: allow_tf32=True may cause accumulate error in conv3d forward >
133
+
134
+ # init resolution
135
+ args.resolution = init_resolution(args.resolution, len(args.dataset_list))
136
+
137
+ # init profiler
138
+ def trace_handler(p):
139
+ p.export_chrome_trace(os.path.join(args.default_root_dir, f"trace_step_{p.step_num}_rank_{0}.json"))
140
+
141
+ tp = None
142
+ if args.turn_on_profiler:
143
+ tp = torch.profiler.profile(
144
+ activities=[
145
+ torch.profiler.ProfilerActivity.CPU,
146
+ torch.profiler.ProfilerActivity.CUDA,
147
+ ],
148
+ schedule=torch.profiler.schedule(
149
+ wait=args.profiler_scheduler_wait_steps,
150
+ warmup=3,
151
+ active=2,
152
+ repeat=1,
153
+ ),
154
+ with_stack=True,
155
+ record_shapes=True,
156
+ profile_memory=True,
157
+ on_trace_ready=trace_handler
158
+ )
159
+ tp.start()
160
+ record_function = torch_record_function
161
+ else:
162
+ record_function = nullcontext
163
+
164
+
165
+ vae = None
166
+ use_vae = None
167
+ num_codes = None
168
+ if args.vqgan_ckpt:
169
+ args.vqgan_ckpt = data_prefix_manager(args.vqgan_ckpt)
170
+ if args.tokenizer in ["hbq_tokenizer"]:
171
+ vae = vae_model(args)
172
+ state_dict = torch.load(args.vqgan_ckpt, map_location=torch.device("cpu"), weights_only=True)
173
+ new_state_dict = {}
174
+ for key in ['vae', 'ema']:
175
+ if (key not in state_dict) or (not state_dict[key]):
176
+ continue
177
+ if 'quantizer.scale_learnable_parameters' in state_dict[key]:
178
+ if len(state_dict[key]['quantizer.scale_learnable_parameters']) == 1:
179
+ state_dict[key]['quantizer.scale_learnable_parameters'] = state_dict[key]['quantizer.scale_learnable_parameters'].expand(4)
180
+ state_dict[key]['scale_learnable_parameters'] = state_dict[key]['quantizer.scale_learnable_parameters']
181
+ del state_dict[key]['quantizer.scale_learnable_parameters']
182
+ if 'z_mean' in state_dict[key]:
183
+ if state_dict[key]['z_mean'].shape != vae.z_mean.shape:
184
+ del state_dict[key]['z_mean']
185
+ del state_dict[key]['z_std']
186
+ new_state_dict[key] = state_dict[key]
187
+ slim_model_path = args.vqgan_ckpt.replace('/checkpoints/', f'/slim_{key}/')
188
+ if not osp.exists(slim_model_path):
189
+ os.makedirs(os.path.dirname(slim_model_path), exist_ok=True)
190
+ torch.save({key: state_dict[key]}, slim_model_path)
191
+ print(f'save to {slim_model_path}')
192
+
193
+ if args.ema == "yes":
194
+ print("testing ema weights")
195
+ print(vae.load_state_dict(new_state_dict["ema"], strict=False))
196
+ else:
197
+ print("testing non ema weights")
198
+ print(vae.load_state_dict(new_state_dict["vae"], strict=False))
199
+ for name, param in vae.named_parameters():
200
+ if name.startswith("scale_learnable_"):
201
+ try:
202
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
203
+ except:
204
+ print(f"{name}: {param[:32].cpu().detach().reshape(-1).tolist()}")
205
+ for name, param in vae.named_buffers():
206
+ if name.startswith("scale_learnable_"):
207
+ try:
208
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
209
+ except:
210
+ print(f"{name}: {param[:32].cpu().detach().reshape(-1).tolist()}")
211
+ if ("scale_wise_std_" in name) or ("scale_wise_mean_" in name):
212
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
213
+ if ('signal_' in name):
214
+ print(f"{name}: {param.cpu().detach().reshape(-1).tolist()}")
215
+ if args.tokenizer != 'hbq_tokenizer':
216
+ vae.enable_slicing()
217
+ # vae.enable_tiling()
218
+ else:
219
+ raise NotImplementedError
220
+
221
+ if args.inference_type == "video":
222
+ def extract_results(return_dict, world_size):
223
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = [], [], [], [], []
224
+ if args.intermediate_tensor:
225
+ for rank in range(world_size):
226
+ real_embeddings.append(return_dict[rank]['real_embeddings'])
227
+ fake_embeddings.append(return_dict[rank]['fake_embeddings'])
228
+ all_real_videos += return_dict[rank]['all_real_videos']
229
+ all_fake_videos += return_dict[rank]['all_fake_videos']
230
+ zs.append(return_dict[rank]['zs'])
231
+ real_embeddings = torch.cat(real_embeddings, 0).to('cuda:0')
232
+ fake_embeddings = torch.cat(fake_embeddings, 0).to('cuda:0')
233
+ zs = torch.cat(zs, 0).to('cuda:0')
234
+ else:
235
+ for rank in range(world_size):
236
+ real_embeddings.append(return_dict[rank]['real_embeddings'])
237
+ fake_embeddings.append(return_dict[rank]['fake_embeddings'])
238
+ all_real_videos.append(return_dict[rank]['all_real_videos'])
239
+ all_fake_videos.append(return_dict[rank]['all_fake_videos'])
240
+ zs.append(return_dict[rank]['zs'])
241
+ real_embeddings = torch.cat(real_embeddings, 0).to('cuda:0')
242
+ fake_embeddings = torch.cat(fake_embeddings, 0).to('cuda:0')
243
+ all_real_videos = torch.cat(all_real_videos, 0)
244
+ all_fake_videos = torch.cat(all_fake_videos, 0)
245
+ zs = torch.cat(zs, 0).to('cuda:0')
246
+ return real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs
247
+
248
+ def inference(mean=None, std=None, noise_scale=0):
249
+ world_size = torch.cuda.device_count()
250
+ manager = torch.multiprocessing.Manager()
251
+ return_dict = manager.dict()
252
+ ### multi-process
253
+ # try:
254
+ # spawn(inference_DDP, args=(world_size, args, vae_model, vae, record_function, tp, use_vae, num_codes, return_dict, mean, std, noise_scale), nprocs=world_size, join=True)
255
+ # except Exception as e:
256
+ # print(f"Error during spawn {e}")
257
+
258
+ ## single process
259
+ world_size = 1
260
+ inference_DDP(0, world_size, args, vae_model, vae, record_function, tp, use_vae, num_codes, return_dict, mean=mean, std=std, noise_scale=noise_scale)
261
+
262
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = extract_results(return_dict, world_size)
263
+ return real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs
264
+
265
+ def cal_std(zs):
266
+ dims_to_reduce = [i for i in range(zs.dim()) if i != 1]
267
+ total_std = zs.std().item()
268
+ _mean = zs.mean(dim=dims_to_reduce)
269
+ _std = zs.std(dim=dims_to_reduce)
270
+ return total_std, _mean, _std
271
+
272
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = inference()
273
+ if args.noise_scale > 0:
274
+ total_std, _mean, _std = cal_std(zs)
275
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = inference(mean=_mean, std=_std, noise_scale=args.noise_scale)
276
+
277
+ if args.save_samples:
278
+ torch.save(zs.cpu(), args.save_samples)
279
+
280
+ if args.cal_norm:
281
+ total_std, _mean, _std = cal_std(zs)
282
+ print(f"{total_std = } {_mean = } {_std = }")
283
+ if args.save_prediction:
284
+ fname = os.path.join(args.save_dir, args.dataset_list[0], "gt_recon", "mean_std.pth")
285
+ torch.save({'_mean': _mean, '_std': _std}, fname)
286
+
287
+ result_str = video_eval(real_embeddings, fake_embeddings, all_real_videos, all_fake_videos)
288
+ else:
289
+ world_size = 1 if args.debug else torch.cuda.device_count()
290
+ manager = torch.multiprocessing.Manager()
291
+ return_dict = manager.dict()
292
+
293
+ if args.debug:
294
+ inference_eval(0, world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict)
295
+ else:
296
+ spawn(inference_eval, args=(world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict), nprocs=world_size, join=True)
297
+
298
+ pred_xs, pred_recs, lpips_alex, lpips_vgg, ssim_value, psnr_value, num_iter, total_usage, total_usage_bit, total_num_token, all_bit_indices_cat = [], [], 0, 0, 0, 0, 0, 0, 0, 0, []
299
+ for rank in range(world_size):
300
+ pred_xs.append(return_dict[rank]['pred_xs'])
301
+ pred_recs.append(return_dict[rank]['pred_recs'])
302
+ lpips_alex += return_dict[rank]['lpips_alex']
303
+ lpips_vgg += return_dict[rank]['lpips_vgg']
304
+ ssim_value += return_dict[rank]['ssim_value']
305
+ psnr_value += return_dict[rank]['psnr_value']
306
+ num_iter += return_dict[rank]['num_iter']
307
+ total_usage += return_dict[rank]['total_usage']
308
+ pred_xs = np.concatenate(pred_xs, 0)
309
+ pred_recs = np.concatenate(pred_recs, 0)
310
+
311
+ result_str = image_eval(pred_xs, pred_recs, lpips_alex, lpips_vgg, ssim_value, psnr_value, num_iter, total_usage, num_codes, total_usage_bit, total_num_token)
312
+ # result_str = inference_eval(args, vae_model, vae, record_function, use_vae, num_codes)
313
+
314
+ print(f"noise scale = {args.noise_scale}")
315
+ print(result_str)
316
+ # save result_str to exp_dir
317
+ basename = os.path.basename(args.vqgan_ckpt)
318
+ match = re.search(r'model_step_(\d+)\.ckpt', basename)
319
+ iter_num = match.group(1) if match else None
320
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
321
+ ckpt_dir = os.path.dirname(data_prefix_manager(args.vqgan_ckpt))
322
+ use_frames = args.use_frames if args.use_frames else args.sequence_length
323
+ save_dir = os.path.join(ckpt_dir, "evaluation", args.dataset_list[0], f"{args.resolution[0][0]}_{args.resolution[0][1]}", f"{use_frames}")
324
+ os.makedirs(save_dir, exist_ok=True)
325
+ ema_suffix = "_ema" if args.ema == "yes" else ""
326
+ result_name = os.path.join(save_dir, f"result_{iter_num}{ema_suffix}.txt")
327
+ if (not args.save_prediction) and (args.noise_scale == 0):
328
+ with open(result_name, "w") as f:
329
+ f.write(result_str)
330
+ # print('Usage = %.2f'%((total_usage > 0.).sum() / num_codes))
331
+ if args.intermediate_tensor:
332
+ os.system(f"rm -rf {args.intermediate_tensor_dir}")
333
+
334
+ def add_noise(z, mean, std, noise_scale):
335
+ if noise_scale > 0:
336
+ mean = mean.view(1, mean.shape[0], 1, 1, 1).to(z.device)
337
+ std = std.view(1, std.shape[0], 1, 1, 1).to(z.device)
338
+ z = (z - mean) / std
339
+ noise = torch.randn(z.size()).to(z.device)
340
+ z = (z + noise * noise_scale) * std + mean
341
+ return z
342
+
343
+ def inference_DDP(rank, world_size, args, vae_model, vae, record_function, tp, use_vae, num_codes, return_dict, mean=None, std=None, noise_scale=0):
344
+ setup(rank, world_size)
345
+ # init data_prefix_manager
346
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
347
+
348
+ for param in vae.parameters():
349
+ param.requires_grad = False
350
+ vae = vae.eval()
351
+ vae = vae.to(f"cuda:{rank}")
352
+ # vae = torch.compile(vae)
353
+
354
+ save_dir = os.path.join(args.save_dir, args.dataset_list[0])
355
+ print('generating and saving video to %s...'%save_dir)
356
+ os.makedirs(save_dir, exist_ok=True)
357
+
358
+ data = VideoData(args)
359
+ loader = data.val_dataloader()
360
+
361
+ i3d = load_fvd_model(f"cuda:{rank}")
362
+
363
+ os.makedirs(os.path.join(save_dir, "gt"), exist_ok=True)
364
+ os.makedirs(os.path.join(save_dir, "recons"), exist_ok=True)
365
+
366
+ zs = []
367
+ real_embeddings = []
368
+ fake_embeddings = []
369
+
370
+ all_real_videos = []
371
+ all_fake_videos = []
372
+
373
+ num_videos = len(loader)
374
+ loader_iter = iter(loader)
375
+ progress_bar = tqdm(total=num_videos, desc=f"Testing {num_videos} batches")
376
+ for batch_idx in range(num_videos):
377
+ if args.turn_on_profiler and tp:
378
+ tp.step()
379
+ batch = next(loader_iter)
380
+ with torch.no_grad():
381
+ input_ = batch['video'] # B C T H W
382
+ B = input_.shape[0]
383
+ if args.tokenizer in ["hbq_tokenizer"]:
384
+ input_ = input_.to(f"cuda:{rank}").to(torch.bfloat16)
385
+ with torch.amp.autocast("cuda", dtype=torch.bfloat16):
386
+ x_raw, x_recons, z = vae(input_, 0, is_train=False)
387
+ batch['video'] = x_raw.to('cpu').to(torch.float32)
388
+ x_recons = x_recons.to(torch.float32)
389
+ else:
390
+ raise NotImplementedError
391
+
392
+ if args.tokenizer in ["icvivit", "sd"]:
393
+ x_recons = rearrange(x_recons, "(b t) c h w -> b c t h w", b=B)
394
+
395
+ real_videos = torch.clamp(batch['video'] / 2 + 0.5, 0, 1)
396
+ if args.junke_old:
397
+ fake_videos = torch.clamp(x_recons.detach().cpu() + 0.5, 0, 1)
398
+ else:
399
+ fake_videos = torch.clamp(x_recons.detach().cpu() / 2 + 0.5, 0, 1)
400
+
401
+ use_frames = args.use_frames if args.use_frames else args.sequence_length
402
+ if args.intermediate_tensor:
403
+ folder_name = os.path.join(args.intermediate_tensor_dir, f"{rank}_{batch_idx}")
404
+ os.makedirs(folder_name, exist_ok=True)
405
+ real_file = os.path.join(folder_name, "real_videos.pt")
406
+ fake_file = os.path.join(folder_name, "fake_videos.pt")
407
+ real_videos = real_videos[:,:,:use_frames,...]
408
+ fake_videos = fake_videos[:,:,:use_frames,...]
409
+ torch.save(real_videos.permute(0, 2, 1, 3, 4).squeeze(0), real_file)
410
+ torch.save(fake_videos.permute(0, 2, 1, 3, 4).squeeze(0), fake_file)
411
+ all_real_videos.append(real_file)
412
+ all_fake_videos.append(fake_file)
413
+ else:
414
+ real_videos = real_videos[:,:,:use_frames,...]
415
+ fake_videos = fake_videos[:,:,:use_frames,...]
416
+ all_real_videos.append(real_videos.clone())
417
+ all_fake_videos.append(fake_videos.clone())
418
+ if args.cal_norm or args.save_samples or args.noise_scale > 0:
419
+ zs.append(z)
420
+ real_embedding = get_fvd_logits(shift_dim(real_videos * 255, 1, -1).byte().data.numpy(), i3d=i3d, device=f"cuda:{rank}").cpu()
421
+ real_embeddings.append(real_embedding)
422
+ fake_embedding = get_fvd_logits(shift_dim(fake_videos * 255, 1, -1).byte().data.numpy(), i3d=i3d, device=f"cuda:{rank}").cpu()
423
+ fake_embeddings.append(fake_embedding)
424
+
425
+ if args.tokenizer in ['cvivit', "icvivit"] and not use_vae:
426
+ batch_codebook_usage = vq_output["batch_usage"]
427
+ total_usage += batch_codebook_usage
428
+
429
+ if args.save_prediction:
430
+ video = torch.cat([real_videos[:,:,:fake_videos.shape[2],:,:], fake_videos], dim=-1)
431
+ b, c, t, h, w = video.shape
432
+ video = video.permute(0, 2, 3, 4, 1).contiguous()
433
+ video = (video.squeeze().detach().cpu().numpy() * 255).astype('uint8')
434
+ os.makedirs(os.path.join(save_dir, "gt_recon"), exist_ok=True)
435
+ this_filename = batch["path"][0].split('/')[-1]
436
+ fname = os.path.join(save_dir, "gt_recon", this_filename)
437
+ import imageio
438
+ imageio.mimsave(fname, video, fps=15)
439
+
440
+ if args.save_z:
441
+ os.makedirs(os.path.join(save_dir, "gt_recon"), exist_ok=True)
442
+ this_filename = batch["path"][0].split('/')[-1].split(".")[0]
443
+ fname = os.path.join(save_dir, "gt_recon", this_filename+".pt")
444
+ torch.save(z, fname)
445
+
446
+ if args.save_frames:
447
+
448
+ def convert_to_uint8(image):
449
+ return (image.detach().cpu().numpy() * 255).astype(np.uint8)
450
+
451
+ # artifact_grid_size = 32
452
+ assert real_videos.shape == fake_videos.shape, f"shape of gt and predicted videos are not equal"
453
+ assert real_videos.shape[0] == fake_videos.shape[0] == 1, f"batch size must be 1, real_videos {real_videos.shape[0]}, fake_videos {fake_videos.shape[0]}"
454
+ _real_videos = real_videos.squeeze(0)
455
+ _fake_videos = fake_videos.squeeze(0)
456
+ # h, w = real_videos.shape[-2:]
457
+ # assert (h % artifact_grid_size == 0) and (w % artifact_grid_size == 0), f"height and width of video must be divisible by {artifact_grid_size}"
458
+
459
+ frame_num = _real_videos.shape[1]
460
+ for frame_idx in range(frame_num):
461
+ real_image = _real_videos[:,frame_idx,:,:]
462
+ fake_image = _fake_videos[:,frame_idx,:,:]
463
+ # most_different_top_left, max_difference = find_most_different_patch(real_image, fake_image, artifact_grid_size)
464
+ real_image_uint8 = convert_to_uint8(real_image)
465
+ predicted_image_uint8 = convert_to_uint8(fake_image)
466
+
467
+ real_image_bgr = cv2.cvtColor(real_image_uint8.transpose(1, 2, 0), cv2.COLOR_RGB2BGR)
468
+ predicted_image_bgr = cv2.cvtColor(predicted_image_uint8.transpose(1, 2, 0), cv2.COLOR_RGB2BGR)
469
+
470
+ concatenated_image = np.concatenate((real_image_bgr, predicted_image_bgr), axis=1)
471
+ fname = os.path.join(save_dir, "gt_recon", f"{this_filename}_{frame_idx}.png")
472
+ cv2.imwrite(fname, concatenated_image)
473
+ progress_bar.update(1)
474
+
475
+ real_embeddings = torch.cat(real_embeddings, 0)
476
+ fake_embeddings = torch.cat(fake_embeddings, 0)
477
+ zs = torch.cat(zs, 0) if len(zs) > 0 else torch.tensor([])
478
+ if args.intermediate_tensor:
479
+ temp_dict = {
480
+ 'real_embeddings':real_embeddings.cpu(),
481
+ 'fake_embeddings':fake_embeddings.cpu(),
482
+ 'all_real_videos':all_real_videos,
483
+ 'all_fake_videos':all_fake_videos,
484
+ 'zs': zs.cpu(),
485
+ }
486
+ else:
487
+ all_real_videos = torch.cat(all_real_videos, 0).permute(0, 2, 1, 3, 4)
488
+ all_fake_videos = torch.cat(all_fake_videos, 0).permute(0, 2, 1, 3, 4)
489
+ temp_dict = {
490
+ 'real_embeddings':real_embeddings.cpu(),
491
+ 'fake_embeddings':fake_embeddings.cpu(),
492
+ 'all_real_videos':all_real_videos.cpu(),
493
+ 'all_fake_videos':all_fake_videos.cpu(),
494
+ 'zs': zs.cpu(),
495
+ }
496
+ # if dist.is_initialized():
497
+ # dist.barrier()
498
+ return_dict[rank] = temp_dict
499
+ cleanup()
500
+
501
+
502
+ def video_eval(real_embeddings, fake_embeddings, all_real_videos, all_fake_videos):
503
+ fake_embeddings = fake_embeddings.to(torch.float64)
504
+ real_embeddings = real_embeddings.to(torch.float64)
505
+ FVD = frechet_distance(fake_embeddings, real_embeddings)
506
+ print(f"FVD: {FVD}") # can't wait to see this number :)
507
+ del real_embeddings, fake_embeddings
508
+
509
+ lpips = calculate_lpips(all_real_videos, all_fake_videos, device="cuda")["value"].values()
510
+ psnr = calculate_psnr(all_real_videos, all_fake_videos)["value"].values()
511
+ ssim = calculate_ssim(all_real_videos, all_fake_videos)["value"].values()
512
+ lpips = np.mean(np.stack(list(lpips)))
513
+ ssim = np.mean(np.stack(list(ssim)))
514
+ psnr = np.mean(np.stack(list(psnr)))
515
+
516
+ result_str = f"""
517
+ FVD = {FVD:.4f}
518
+ LPIPS = {lpips:.4f}
519
+ SSIM = {ssim:.4f}
520
+ PSNR = {psnr:.3f}
521
+ """
522
+ return result_str
523
+
524
+ def inference_eval(rank, world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict):
525
+ # Don't remove this setup!!! dist.init_process_group is important for building loader (data.distributed.DistributedSampler)
526
+ setup(rank, world_size)
527
+ # init data_prefix_manager
528
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
529
+
530
+ device = torch.device(f"cuda:{rank}")
531
+
532
+ for param in vae.parameters():
533
+ param.requires_grad = False
534
+ vae.to(device).eval()
535
+
536
+ save_dir = os.path.join(args.save_dir, args.dataset_list[0])
537
+ print('generating and saving video to %s...'%save_dir)
538
+ os.makedirs(save_dir, exist_ok=True)
539
+
540
+ data = VideoData(args)
541
+
542
+ loader = data.val_dataloader()
543
+
544
+ dims = 2048
545
+ block_idx = InceptionV3.BLOCK_INDEX_BY_DIM[dims]
546
+ inception_model = InceptionV3([block_idx]).to(device)
547
+ inception_model.eval()
548
+
549
+ loader_iter = iter(loader)
550
+
551
+ pred_xs = []
552
+ pred_recs = []
553
+ # LPIPS score related
554
+ loss_fn_alex = lpips.LPIPS(net='alex').to(device) # best forward scores
555
+ loss_fn_vgg = lpips.LPIPS(net='vgg').to(device) # closer to "traditional" perceptual loss, when used for optimization
556
+ lpips_alex = 0.0
557
+ lpips_vgg = 0.0
558
+
559
+ # SSIM score related
560
+ ssim_value = 0.0
561
+
562
+ # PSNR score related
563
+ psnr_value = 0.0
564
+
565
+ num_images = len(loader)
566
+ print(f"Testing {num_images} files")
567
+ num_iter = 0
568
+
569
+ total_usage = 0.0
570
+ total_usage_bit = 0.0
571
+ total_num_token = 0
572
+ for batch_idx in tqdm(range(num_images)):
573
+ batch = next(loader_iter)
574
+
575
+ with torch.no_grad():
576
+ x = batch['video']
577
+ if args.tokenizer in ["hbq_tokenizer"]:
578
+ x_raw, x_recons, z = vae(x.to(device), 0, is_train=False)
579
+ x_recons = x_recons.squeeze(-3).cpu()
580
+ else:
581
+ raise NotImplementedError
582
+
583
+ if args.image_recon4video:
584
+ # convert back to image format
585
+ x = x.squeeze(2)
586
+ x_recons = x_recons.squeeze(2)
587
+
588
+ if args.tokenizer in ["cvivit", "icvivit"] and not use_vae:
589
+
590
+ # encoding_indices = vq_output["encodings"].detach().cpu()
591
+ code_counts = calculate_batch_codebook_usage_percentage(vq_output["encodings"], num_codes)
592
+ total_counts += code_counts
593
+
594
+ batch_codebook_usage = vq_output["batch_usage"]
595
+ total_usage += batch_codebook_usage
596
+
597
+ paths = batch["path"]
598
+ assert len(paths) == x.shape[0]
599
+
600
+ for p, input_ori, recon_ori in zip(paths, x, x_recons):
601
+ if os.path.isabs(p):
602
+ p = "/".join(p.split("/")[6:])
603
+ assert not os.path.isabs(p), f"{p} should not be abspath"
604
+ path = os.path.join(save_dir, "input_recon", os.path.basename(p))
605
+ os.makedirs(os.path.split(path)[0], exist_ok=True)
606
+
607
+ input_ori = input_ori.unsqueeze(0).to(device)
608
+ input_ = (input_ori + 1) / 2 # [0, 1]
609
+
610
+ pred_x = inception_model(input_)[0]
611
+ pred_x = pred_x.squeeze(3).squeeze(2).cpu().numpy()
612
+
613
+ recon_ori = recon_ori.unsqueeze(0).to(device)
614
+ recon_ = (recon_ori + 1) / 2 # [0, 1]
615
+ # recon_ = recon_.permute(1, 2, 0).detach().cpu()
616
+ with torch.no_grad():
617
+ pred_rec = inception_model(recon_)[0]
618
+ pred_rec = pred_rec.squeeze(3).squeeze(2).cpu().numpy()
619
+ if args.save_prediction:
620
+ if input_.dim() == 4:
621
+ input_image = input_.squeeze(0)
622
+ if recon_.dim() == 4:
623
+ recon_image = recon_.squeeze(0)
624
+ input_recon = torch.cat([input_image, recon_image], dim=-1)
625
+ input_recon = Image.fromarray((torch.clamp(input_recon.permute(1, 2, 0).detach().cpu(), 0, 1).numpy() * 255).astype(np.uint8))
626
+ input_recon.save(path)
627
+
628
+ pred_xs.append(pred_x)
629
+ pred_recs.append(pred_rec)
630
+
631
+ # calculate lpips
632
+ with torch.no_grad():
633
+ lpips_alex += loss_fn_alex(input_ori, recon_ori).sum() # [-1, 1]
634
+ lpips_vgg += loss_fn_vgg(input_ori, recon_ori).sum() # [-1, 1]
635
+
636
+ #calculate PSNR and SSIM
637
+ rgb_restored = (recon_ * 255.0).permute(0, 2, 3, 1).to("cpu", dtype=torch.uint8).numpy()
638
+ rgb_gt = (input_ * 255.0).permute(0, 2, 3, 1).to("cpu", dtype=torch.uint8).numpy()
639
+ rgb_restored = rgb_restored.astype(np.float32) / 255.
640
+ rgb_gt = rgb_gt.astype(np.float32) / 255.
641
+ ssim_temp = 0
642
+ psnr_temp = 0
643
+ B, _, _, _ = rgb_restored.shape
644
+ for i in range(B):
645
+ rgb_restored_s, rgb_gt_s = rgb_restored[i], rgb_gt[i]
646
+ with torch.no_grad():
647
+ ssim_temp += ssim_loss(rgb_restored_s, rgb_gt_s, data_range=1.0, channel_axis=-1)
648
+ psnr_temp += psnr_loss(rgb_gt, rgb_restored)
649
+ ssim_value += ssim_temp / B
650
+ psnr_value += psnr_temp / B
651
+ num_iter += 1
652
+
653
+ pred_xs = np.concatenate(pred_xs, axis=0)
654
+ pred_recs = np.concatenate(pred_recs, axis=0)
655
+ temp_dict = {
656
+ 'pred_xs':pred_xs,
657
+ 'pred_recs':pred_recs,
658
+ 'lpips_alex':lpips_alex.cpu(),
659
+ 'lpips_vgg':lpips_vgg.cpu(),
660
+ 'ssim_value': ssim_value,
661
+ 'psnr_value': psnr_value,
662
+ 'num_iter': num_iter,
663
+ 'total_usage': total_usage,
664
+ 'total_usage_bit': total_usage_bit,
665
+ 'total_num_token': total_num_token,
666
+ }
667
+ return_dict[rank] = temp_dict
668
+
669
+ # if dist.is_initialized():
670
+ # dist.barrier()
671
+ cleanup()
672
+
673
+ def image_eval(pred_xs, pred_recs, lpips_alex, lpips_vgg, ssim_value, psnr_value, num_iter, total_usage, num_codes, total_usage_bit, total_num_token):
674
+ mu_x = np.mean(pred_xs, axis=0)
675
+ sigma_x = np.cov(pred_xs, rowvar=False)
676
+ mu_rec = np.mean(pred_recs, axis=0)
677
+ sigma_rec = np.cov(pred_recs, rowvar=False)
678
+
679
+ fid_value = calculate_frechet_distance(mu_x, sigma_x, mu_rec, sigma_rec)
680
+ lpips_alex_value = lpips_alex / num_iter
681
+ lpips_vgg_value = lpips_vgg / num_iter
682
+ ssim_value = ssim_value / num_iter
683
+ psnr_value = psnr_value / num_iter
684
+
685
+ result_str = f"""
686
+ FID = {fid_value:.4f}
687
+ LPIPS_VGG: {lpips_vgg_value.item():.4f}
688
+ LPIPS_ALEX: {lpips_alex_value.item():.4f}
689
+ SSIM: {ssim_value:.4f}
690
+ PSNR: {psnr_value:.3f}
691
+ """
692
+ return result_str
693
+ if __name__ == '__main__':
694
+ main()
grn/tokenizer/train.py ADDED
@@ -0,0 +1,561 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from json import load
2
+ import os
3
+ import argparse
4
+ import math
5
+ import glob
6
+ import time
7
+ import logging
8
+ from distutils.util import strtobool
9
+ from copy import deepcopy
10
+ import gc
11
+ gc.disable()
12
+ import os.path as osp
13
+
14
+ import torch
15
+ import torch.nn.functional as F
16
+ import torch.optim as optim
17
+ import torch.distributed as dist
18
+ from torch.profiler import record_function as torch_record_function
19
+ from contextlib import nullcontext
20
+ from torch.nn.parallel import DistributedDataParallel as DDP
21
+ from safetensors.torch import load_file
22
+ from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
23
+ from torch.distributed.fsdp import StateDictType, FullStateDictConfig
24
+
25
+ from videovae.utils.misc import data_prefix_manager, COLOR_BLUE, COLOR_RESET, is_torch_optim_sch
26
+ from videovae.utils.distributed import init_distributed_mode, reduce_losses, average_losses, _FSDP
27
+ from videovae.utils.ema import update_ema, requires_grad
28
+
29
+ from videovae.models.discriminator import ImageDiscriminator, VideoDiscriminator
30
+ from videovae.data import VideoData
31
+ from videovae.modules import build_lpips_model
32
+ from videovae.modules.loss import get_disc_loss, adopt_weight
33
+ from videovae.utils.misc import get_last_ckpt, seed_everything, print_gpu_usage, print_model_summary, version_checker
34
+ from videovae.utils.init_models import init_vae_only, init_vit_from_image, resume_from_ckpt, init_cnn_from_image, load_cnn
35
+ from videovae.utils.nan_detector import NanDetector
36
+ from videovae.utils.arguments import MainArgs, add_model_specific_args, init_args, format_args
37
+ from videovae.utils.scheduler import get_lambda
38
+ from videovae.utils.mfu import register_mfu_hook, get_mfu, get_tflops, get_tflops_dict
39
+
40
+ from videovae.utils.context_parallel import ContextParallelUtils as cp
41
+
42
+ def save_model(fsdp_model, rank, model_path, global_step):
43
+ # FSDP推荐用 state_dict_type=FULL_STATE_DICT 来保存
44
+ with FSDP.state_dict_type(fsdp_model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)):
45
+ state_dict = fsdp_model.state_dict()
46
+ os.makedirs(os.path.dirname(model_path), exist_ok=True)
47
+ torch.save({'vae': state_dict, 'step': global_step}, model_path)
48
+ print(f"模型已保存到 {model_path}")
49
+
50
+ # enable_timeline_sdk = strtobool(os.getenv("GenAI_USE_TIMELINE_SDK", "0"))
51
+ enable_timeline_sdk = False
52
+
53
+ def split_to_ranks(x):
54
+ bs = x.shape[0]
55
+ cp_size = cp.get_cp_size()
56
+ if cp_size > 1 and bs % cp_size == 0:
57
+ cp_rank = cp.get_cp_rank()
58
+ return x.chunk(cp_size, dim=0)[cp_rank]
59
+ else:
60
+ return x
61
+
62
+ if enable_timeline_sdk:
63
+ try:
64
+ import bytedance.ndtimeline as ndtimeline
65
+ except ImportError:
66
+ print(f"import vescale.ndtimeline failed, skipped")
67
+ enable_timeline_sdk = False
68
+
69
+ def init_data_scheduler(video_ranks_ratio: float = -1.0, cp_size: int = 1):
70
+ if video_ranks_ratio < 0:
71
+ return None,None
72
+
73
+ cp_size = max(1, cp_size)
74
+
75
+ rank = torch.distributed.get_rank()
76
+ world_size = torch.distributed.get_world_size()
77
+
78
+ video_ranks = list(range(int((world_size * video_ranks_ratio) // cp_size) * cp_size)) # align to cp_size for video
79
+ image_ranks = list(range(len(video_ranks), world_size))
80
+
81
+ print(f"[info] video_ranks: {video_ranks}, image_ranks: {image_ranks}")
82
+
83
+ if rank in image_ranks:
84
+ group = torch.distributed.new_group(image_ranks)
85
+ dataset_type_on_this_rank = "image"
86
+ else:
87
+ group = torch.distributed.new_group(video_ranks)
88
+ dataset_type_on_this_rank = "video"
89
+
90
+ return group, dataset_type_on_this_rank
91
+
92
+ def main():
93
+ parser = argparse.ArgumentParser()
94
+ parser = MainArgs.add_main_args(parser)
95
+ parser = VideoData.add_data_specific_args(parser)
96
+ args, unknown = parser.parse_known_args()
97
+ args, parser, vae_model = add_model_specific_args(args, parser)
98
+ args = parser.parse_args()
99
+
100
+ args = init_args(args) # post process args
101
+
102
+ # init data_prefix_manager
103
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
104
+ args.default_root_dir = data_prefix_manager(args.default_root_dir)
105
+
106
+ # Setup DDP:
107
+ init_distributed_mode(args)
108
+ rank = dist.get_rank()
109
+ world_size = dist.get_world_size()
110
+ device = rank % torch.cuda.device_count()
111
+ seed_everything(args.seed)
112
+ torch.cuda.set_device(device)
113
+
114
+ # init context parallel
115
+ cp_cfg = {"cp_size": args.context_parallel_size}
116
+ cp.initialize_context_parallel(cp_cfg)
117
+
118
+ ds_group, ds_type = init_data_scheduler(args.video_ranks_ratio, cp_size = args.context_parallel_size)
119
+
120
+ # Setup an experiment folder:
121
+ checkpoint_dir = f"{args.default_root_dir}/checkpoints" # Stores saved model checkpoints
122
+ os.makedirs(checkpoint_dir, exist_ok=True)
123
+ if rank == 0:
124
+ script_str = format_args(args)
125
+ with open(os.path.join(args.default_root_dir, "script.sh"), "w") as f:
126
+ f.write(script_str)
127
+ print(f"{COLOR_BLUE}Experiment directory created at {args.default_root_dir}{COLOR_RESET}")
128
+
129
+ import wandb
130
+ wandb_project = "HBQ_Tokenizer"
131
+ wandb.init(
132
+ project=wandb_project,
133
+ name=os.path.basename(os.path.normpath(args.default_root_dir)),
134
+ dir=args.default_root_dir,
135
+ config=args,
136
+ mode="offline" if args.debug else "online"
137
+ )
138
+
139
+ # init model
140
+ vae = vae_model(args).to(device)
141
+ if rank == 0:
142
+ model_arch_save_path = os.path.join(args.default_root_dir, "model_arch.txt")
143
+ print(f"{COLOR_BLUE}Logging model architecture at {model_arch_save_path}{COLOR_RESET}")
144
+ with open(model_arch_save_path, "w") as f:
145
+ f.write(str(vae))
146
+ image_disc = ImageDiscriminator(args).to(device)
147
+ video_disc = VideoDiscriminator(args).to(device)
148
+
149
+ # init optimizers and schedulers
150
+ if args.optim_type == "Adam":
151
+ vae_optim = torch.optim.Adam
152
+ elif args.optim_type == "AdamW":
153
+ vae_optim = torch.optim.AdamW
154
+ if args.disc_optim_type is None:
155
+ disc_optim = vae_optim
156
+ elif args.disc_optim_type == "rmsprop":
157
+ disc_optim = torch.optim.RMSprop
158
+
159
+ def get_param_groups(model):
160
+ decay = []
161
+ no_decay = []
162
+ for name, param in model.named_parameters():
163
+ if param.requires_grad:
164
+ if len(param.shape) == 1 or name.endswith(".bias") or ('scale_learnable_parameters' in name):
165
+ no_decay.append(param)
166
+ print(f'disable weight deacy for {name}')
167
+ else:
168
+ decay.append(param)
169
+ optimizer_grouped_parameters = [
170
+ {'params': decay, 'weight_decay': 0.01},
171
+ {'params': no_decay, 'weight_decay': 0.0}
172
+ ]
173
+ return optimizer_grouped_parameters
174
+
175
+ opt_vae = vae_optim(get_param_groups(vae), lr=args.lr, betas=(args.beta1, args.beta2))
176
+ if disc_optim == torch.optim.RMSprop:
177
+ opt_image_disc = disc_optim(image_disc.parameters(), lr=args.lr * args.dis_lr_multiplier)
178
+ opt_video_disc = disc_optim(video_disc.parameters(), lr=args.lr * args.dis_lr_multiplier)
179
+ else:
180
+ opt_image_disc = disc_optim(image_disc.parameters(), lr=args.lr * args.dis_lr_multiplier, betas=(args.beta1, args.beta2))
181
+ opt_video_disc = disc_optim(video_disc.parameters(), lr=args.lr * args.dis_lr_multiplier, betas=(args.beta1, args.beta2))
182
+
183
+ if args.scheduler == "no":
184
+ sch_vae, sch_image_disc, sch_video_disc = None, None, None
185
+ else:
186
+ lr_lambda = get_lambda(args)
187
+ sch_vae = optim.lr_scheduler.LambdaLR(opt_vae, lr_lambda)
188
+ sch_image_disc = optim.lr_scheduler.LambdaLR(opt_image_disc, lr_lambda)
189
+ sch_video_disc = optim.lr_scheduler.LambdaLR(opt_video_disc, lr_lambda)
190
+
191
+ ### ema
192
+ ema = None
193
+ if args.ema == "yes":
194
+ ema = deepcopy(vae).to(device) # Create an EMA of the model for use after training
195
+ requires_grad(ema, False)
196
+ print(f"EMA Parameters: {sum(p.numel() for p in ema.parameters()):,}")
197
+ update_ema(ema, vae, decay=0) # Ensure EMA is initialized with synced weights
198
+ ema.eval() # EMA model should always be in eval mode
199
+
200
+ model_optims = {
201
+ "vae" : vae,
202
+ "image_disc" : image_disc,
203
+ "video_disc" : video_disc,
204
+ "opt_vae" : opt_vae,
205
+ "opt_image_disc" : opt_image_disc,
206
+ "opt_video_disc" : opt_video_disc,
207
+ "sch_vae" : sch_vae,
208
+ "sch_image_disc" : sch_image_disc,
209
+ "sch_video_disc" : sch_video_disc,
210
+ "ema": ema,
211
+ }
212
+
213
+ ### Resume from checkpoint in default_root_dir or load pretrained weights if specified
214
+ ckpt_path = None
215
+ assert not args.default_root_dir is None # required argument
216
+ ckpt_path = get_last_ckpt(args.default_root_dir)
217
+ init_step = 0
218
+ if ckpt_path:
219
+ print(f"Resuming from {ckpt_path}")
220
+ state_dict = torch.load(ckpt_path, map_location="cpu")
221
+ model_optims, init_step = resume_from_ckpt(state_dict, model_optims, load_optims=args.zero<=0, remove_disc=args.remove_disc, ckpt_path=ckpt_path, args=args)
222
+ elif args.pretrained is not None:
223
+ args.pretrained = data_prefix_manager(args.pretrained)
224
+ # read weight
225
+ state_dict = torch.load(args.pretrained, map_location="cpu", weights_only=True)
226
+
227
+ if args.pretrained_ema == "yes":
228
+ state_dict["vae"] = state_dict["ema"] # replace vae weights with ema weight
229
+ # load model
230
+ if args.pretrained_mode == "weights":
231
+ model_optims, _ = resume_from_ckpt(state_dict, model_optims, load_optims=False, remove_disc=args.remove_disc, remove_enlarge_factors=args.remove_enlarge_factors, args=args) # load all models and optims
232
+ del state_dict
233
+ else:
234
+ raise NotImplementedError
235
+ print(f"Successfully loaded ckpt {args.pretrained}, pretrained_mode {args.pretrained_mode}")
236
+
237
+ # init dataloader
238
+ data = VideoData(args, ds_group = ds_group, ds_type = ds_type)
239
+ dataloaders = data.train_dataloader()
240
+ dataloader_iters = [iter(loader) for loader in dataloaders]
241
+ ### init epoch in resuming
242
+ dataloader_init_epoch = (
243
+ init_step if init_step > 0 # in case of resuming
244
+ else args.dataloader_init_epoch if args.dataloader_init_epoch > 0 # in case of fintuning
245
+ else 0
246
+ )
247
+ data_epochs = [dataloader_init_epoch for _ in dataloaders]
248
+ for idx in range(len(dataloaders)):
249
+ print(f"Reset the {idx}th dataloader as epoch {data_epochs[idx]}")
250
+ if hasattr(dataloaders[idx], "sampler"):
251
+ dataloaders[idx].sampler.set_epoch(data_epochs[idx])
252
+ else:
253
+ raise NotImplementedError
254
+
255
+ ### torch.compile after loading all weights
256
+ print_model_summary([vae, image_disc, video_disc])
257
+ if args.zero > 0:
258
+ from torch.distributed.fsdp import (
259
+ FullyShardedDataParallel as FSDP,
260
+ ShardingStrategy,
261
+ MixedPrecision,
262
+ )
263
+ def my_policy(
264
+ module: torch.nn.Module,
265
+ recurse: bool,
266
+ **kwargs,
267
+ ) -> bool:
268
+ return True
269
+ auto_wrap_policy = my_policy
270
+ vae = FSDP(
271
+ vae,
272
+ device_id=device,
273
+ sharding_strategy=ShardingStrategy.FULL_SHARD,
274
+ mixed_precision=None,
275
+ auto_wrap_policy=auto_wrap_policy,
276
+ use_orig_params=True,
277
+ sync_module_states=True,
278
+ limit_all_gathers=True,
279
+ device_mesh=None,
280
+ ).to(device)
281
+ # vae = _FSDP(vae, device, args.zero)
282
+ else:
283
+ vae = DDP(vae.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb, find_unused_parameters=True)
284
+
285
+ image_disc = DDP(image_disc.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb)
286
+ video_disc = DDP(video_disc.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb)
287
+
288
+ image_perceptual_model, video_perceptual_model = build_lpips_model(args)
289
+ image_perceptual_model = image_perceptual_model.to(device)
290
+ video_perceptual_model = video_perceptual_model.to(device)
291
+
292
+ if args.compile == "yes":
293
+ if args.vf_weight > 0 and args.vf_weight_approx < 0:
294
+ torch._functorch.config.donated_buffer = False # This backward function was compiled with non-empty donated buffers which requires create_graph=False and retain_graph=False
295
+ torch._dynamo.config.cache_size_limit = 256
296
+ torch._dynamo.config.accumulated_cache_size_limit = 4096
297
+ torch._dynamo.config.automatic_dynamic_shapes = True
298
+ torch._dynamo.config.suppress_errors = False
299
+ torch._dynamo.config.optimize_ddp = False if args.use_checkpoint else True
300
+
301
+ for k in model_optims:
302
+ if k != "ema" and model_optims[k] and not is_torch_optim_sch(model_optims[k]):
303
+ print(f"compiling model {k}")
304
+ if k == "vae":
305
+ model_optims[k].encoder.compile()#options={'fx_graph_cache':True})
306
+ model_optims[k].decoder.compile()#options={'fx_graph_cache':True})
307
+ else:
308
+ model_optims[k].compile()#options={'fx_graph_cache':True})
309
+ print(f"Successfully compiled all models")
310
+
311
+ disc_loss = get_disc_loss(args.disc_loss_type)
312
+
313
+ if enable_timeline_sdk:
314
+ version_checker("2.0.0", "3.0.0")
315
+ bmq_cluster = os.getenv('CUDA_TIMER_STREAM_KAFKA_CLUSTER', 'bmq_bigbang_3rd')
316
+ bmq_topic = os.getenv('CUDA_TIMER_STREAM_KAFKA_TOPIC', 'megatron_cuda_timer_tracing_original')
317
+ ndtimeline.init_ndtimers(
318
+ mode="fsdp",
319
+ mesh_shape=(world_size,),
320
+ world_size=world_size,
321
+ enable_streamer=True,
322
+ post_handlers=[ndtimeline.handlers.MQNDHandler(mq_sinks=[ndtimeline.handlers.format_mq_sink(bmq_cluster, bmq_topic)])],
323
+ )
324
+ ndtimeline.set_global_step(init_step)
325
+ print(f"init timeline successfully")
326
+
327
+ # init profiler
328
+ def trace_handler(p):
329
+ p.export_chrome_trace(os.path.join(args.default_root_dir, f"trace_step_{p.step_num}_rank_{rank}.json.gz"))
330
+
331
+ if args.turn_on_profiler:
332
+ print(f"start to init profiler")
333
+ tp = torch.profiler.profile(
334
+ activities=[
335
+ torch.profiler.ProfilerActivity.CPU,
336
+ torch.profiler.ProfilerActivity.CUDA,
337
+ ],
338
+ schedule=torch.profiler.schedule(
339
+ wait=args.profiler_scheduler_wait_steps,
340
+ warmup=3,
341
+ active=2,
342
+ repeat=1,
343
+ ),
344
+ with_stack=True,
345
+ record_shapes=True,
346
+ profile_memory=True,
347
+ on_trace_ready=trace_handler
348
+ )
349
+ tp.start()
350
+ record_function = torch_record_function
351
+ print(f"finish to init profiler")
352
+ else:
353
+ record_function = nullcontext
354
+
355
+ start_time = time.time()
356
+ cnt = 0
357
+ debug_root_dir = data_prefix_manager(f"debug/{cnt}")
358
+ while os.path.exists(debug_root_dir):
359
+ cnt += 1
360
+ debug_root_dir = data_prefix_manager(f"debug/{cnt}")
361
+ os.makedirs(debug_root_dir, exist_ok=True)
362
+ if cp.get_cp_rank() == 0:
363
+ cp_group = cp.get_cp_group_rank()
364
+ debug_f = open(os.path.join(debug_root_dir, f"rank_{rank}_cp_group_{cp_group}.txt"), "w")
365
+ else:
366
+ debug_f = None
367
+ for global_step in range(init_step, args.max_steps):
368
+ # if args.turn_on_profiler and tp:
369
+ # tp.step()
370
+ loss_dicts = []
371
+
372
+ if global_step == args.discriminator_iter_start - args.disc_pretrain_iter:
373
+ logging.info(f"discriminator begins pretraining ")
374
+ if global_step == args.discriminator_iter_start:
375
+ log_str = "add GAN loss into training"
376
+ if args.disc_pretrain_iter > 0:
377
+ log_str += ", discriminator ends pretraining"
378
+ logging.info(log_str)
379
+
380
+ for idx in range(len(dataloader_iters)):
381
+ try:
382
+ _batch = next(dataloader_iters[idx])
383
+ except StopIteration:
384
+ data_epochs[idx] += 1
385
+ print(f"Reset the {idx}th dataloader as epoch {data_epochs[idx]}")
386
+ dataloaders[idx].sampler.set_epoch(data_epochs[idx])
387
+ dataloader_iters[idx] = iter(dataloaders[idx]) # update dataloader iter
388
+ _batch = next(dataloader_iters[idx])
389
+ except Exception as e:
390
+ raise e
391
+ x = _batch["video"]
392
+
393
+ _type = _batch["type"][0]
394
+
395
+ if _type == "image" and ds_type is None:
396
+ x = split_to_ranks(x)
397
+
398
+ disc_factor = 1.
399
+ with NanDetector(vae) if args.enable_nan_detector else nullcontext():
400
+ with record_function("vae"):
401
+ if _type == "image":
402
+ x, x_recon, flat_frames, flat_frames_recon, vae_loss_dict, vae_log_dict = vae(x, disc_factor, image_disc=image_disc, image_perceptual_model=image_perceptual_model)
403
+ elif _type == "video":
404
+ if debug_f is not None:
405
+ debug_f.write(f'step {idx}, {_batch["path"]}\n')
406
+ x, x_recon, flat_frames, flat_frames_recon, vae_loss_dict, vae_log_dict = vae(
407
+ x, disc_factor,
408
+ image_disc=image_disc, video_disc=video_disc,
409
+ image_perceptual_model=image_perceptual_model,
410
+ video_perceptual_model=video_perceptual_model,
411
+ )
412
+ g_loss = sum(vae_loss_dict.values())
413
+ opt_vae.zero_grad()
414
+ g_loss.backward()
415
+ # print_gpu_usage("vae")
416
+ if args.max_grad_norm > 0:
417
+ torch.nn.utils.clip_grad_norm_(vae.parameters(), args.max_grad_norm)
418
+
419
+ # from pnp.utils import detect_anomalous_params
420
+ # detect_anomalous_params(g_loss, vae)
421
+
422
+ opt_vae.step()
423
+ opt_vae.zero_grad() # free memory
424
+ if args.ema == "yes":
425
+ update_ema(ema, vae.module)
426
+
427
+ with record_function("disc"):
428
+ disc_loss_dict = {}
429
+ # args.discriminator_iter_start=-1 args.disc_pretrain_iter=0
430
+ discloss = d_image_loss = d_video_loss = torch.tensor(0.).to(x.device)
431
+ ### enable pool warmup
432
+ for disc_step in range(args.disc_optim_steps):
433
+ require_optim = False
434
+ if _type == "image":
435
+ if args.image_disc_weight > 0:
436
+ require_optim = True
437
+ logits_image_real = image_disc(x, pool_name="real")
438
+ logits_image_fake = image_disc(x_recon.detach(), pool_name="fake")
439
+ d_image_loss = disc_loss(logits_image_real, logits_image_fake)
440
+ discloss = d_image_loss * args.image_disc_weight
441
+ disc_loss_dict["train/logits_image_real"] = logits_image_real.mean().detach()
442
+ disc_loss_dict["train/logits_image_fake"] = logits_image_fake.mean().detach()
443
+ disc_loss_dict["train/d_image_loss"] = discloss.detach()
444
+ opt_discs, sch_discs = [opt_image_disc], [sch_image_disc]
445
+ elif _type == "video":
446
+ if args.image_disc_weight > 0 and args.gan_image4video == "yes":
447
+ require_optim = True
448
+ logits_image_real = image_disc(flat_frames.detach(), pool_name="real")
449
+ logits_image_fake = image_disc(flat_frames_recon.detach(), pool_name="fake")
450
+ d_image_loss = disc_loss(logits_image_real, logits_image_fake)
451
+ disc_loss_dict["train/logits_image_real"] = logits_image_real.mean().detach()
452
+ disc_loss_dict["train/logits_image_fake"] = logits_image_fake.mean().detach()
453
+ disc_loss_dict["train/d_image_loss"] = (d_image_loss * args.image_disc_weight).detach()
454
+ if args.video_disc_weight > 0:
455
+ require_optim = True
456
+ logits_video_real = video_disc(x.detach(), pool_name="real")
457
+ logits_video_fake = video_disc(x_recon.detach(), pool_name="fake")
458
+ d_video_loss = disc_loss(logits_video_real, logits_video_fake)
459
+ disc_loss_dict["train/logits_video_real"] = logits_video_real.mean().detach()
460
+ disc_loss_dict["train/logits_video_fake"] = logits_video_fake.mean().detach()
461
+ disc_loss_dict["train/d_video_loss"] = (d_video_loss * args.video_disc_weight).detach()
462
+ discloss = d_image_loss * args.image_disc_weight + d_video_loss * args.video_disc_weight
463
+ opt_discs, sch_discs = [opt_image_disc, opt_video_disc], [sch_image_disc, sch_video_disc]
464
+ discloss = disc_factor * discloss
465
+
466
+ if require_optim:
467
+ for opt_disc in opt_discs:
468
+ opt_disc.zero_grad()
469
+ discloss.backward()
470
+ # print_gpu_usage("disc")
471
+ if args.max_grad_norm_disc > 0:
472
+ torch.nn.utils.clip_grad_norm_(image_disc.parameters(), args.max_grad_norm_disc)
473
+ torch.nn.utils.clip_grad_norm_(video_disc.parameters(), args.max_grad_norm_disc)
474
+
475
+ for opt_disc in opt_discs:
476
+ opt_disc.step()
477
+ for opt_disc in opt_discs:
478
+ opt_disc.zero_grad() # free memory
479
+
480
+ with record_function("loss"):
481
+ loss_dict = {**vae_loss_dict, **disc_loss_dict, **vae_log_dict}
482
+ if (global_step+1) % args.log_every == 0:
483
+ reduced_loss_dict = reduce_losses(loss_dict)
484
+ else:
485
+ reduced_loss_dict = {}
486
+ loss_dicts.append(reduced_loss_dict)
487
+
488
+ # update scheduler
489
+ if not sch_vae is None:
490
+ sch_vae.step()
491
+ for sch_disc in sch_discs:
492
+ if not sch_disc is None:
493
+ sch_disc.step()
494
+
495
+ # if enable_timeline_sdk:
496
+ # ndtimeline.inc_step()
497
+
498
+ if (global_step+1) % args.log_every == 0:
499
+ avg_loss_dict = average_losses(loss_dicts)
500
+ torch.cuda.synchronize()
501
+ end_time = time.time()
502
+ iter_speed = (end_time - start_time) / args.log_every
503
+
504
+ if args.mfu_logging == "yes":
505
+ tflops = get_tflops() / args.log_every
506
+ tflops_log_str = f"tflops={tflops:.1f}, "
507
+ tflops_dict = get_tflops_dict(args.log_every)
508
+ tflops_dict_log_str = f"tflops_Dict={tflops_dict}, "
509
+ mfu = get_mfu(iter_speed) / args.log_every
510
+ mfu_log_str = f"mfu={mfu:.3f}, "
511
+ else:
512
+ tflops_log_str = ""
513
+ tflops_dict_log_str = ""
514
+ mfu_log_str = ""
515
+
516
+ if rank == 0:
517
+ avg_loss_dict["lr"] = opt_vae.param_groups[0]['lr']
518
+ for key, value in avg_loss_dict.items():
519
+ wandb.log({key: value}, step=global_step)
520
+
521
+ recons_loss_sum, video_perceptual_loss_sum = 0., 0.
522
+ for key in avg_loss_dict:
523
+ if 'recon_loss' in key:
524
+ recons_loss_sum += avg_loss_dict[key]
525
+ if 'video_perceptual_loss' in key:
526
+ video_perceptual_loss_sum += avg_loss_dict[key]
527
+
528
+ print(f'global_step={global_step}, recon_loss={recons_loss_sum:.4f}, ' \
529
+ f'video_perceptual_loss={video_perceptual_loss_sum:.4f}, ' \
530
+ f'iter_speed={iter_speed:.2f}s, ' \
531
+ f'{mfu_log_str}' \
532
+ f'{tflops_log_str}' \
533
+ f'{tflops_dict_log_str}' \
534
+ )
535
+ start_time = time.time()
536
+ if enable_timeline_sdk:
537
+ ndtimeline.flush()
538
+
539
+ if (global_step+1) % args.ckpt_every == 0 and global_step != init_step:
540
+ checkpoint_path = os.path.join(checkpoint_dir, f'model_step_{global_step}.ckpt')
541
+ if args.zero > 0:
542
+ save_model(vae, rank, checkpoint_path, global_step)
543
+ else:
544
+ if rank == 0:
545
+ save_dict = {}
546
+ for k in model_optims:
547
+ model = model_optims[k]
548
+ save_dict[k] = None if model is None \
549
+ else model.module.state_dict() if hasattr(model, "module") \
550
+ else model.state_dict()
551
+ torch.save({
552
+ 'step': global_step,
553
+ **save_dict,
554
+ }, checkpoint_path)
555
+ print(f'Checkpoint saved at step {global_step}')
556
+
557
+ if (global_step+1) % args.manual_gc_interval == 0:
558
+ gc.collect()
559
+
560
+ if __name__ == '__main__':
561
+ main()
grn/tokenizer/videovae/__init__.py ADDED
File without changes
grn/tokenizer/videovae/evaluation/__init__.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ from .common_metrics_on_video_quality.calculate_fvd import calculate_fvd
2
+ from .common_metrics_on_video_quality.calculate_lpips import calculate_lpips
3
+ from .common_metrics_on_video_quality.calculate_psnr import calculate_psnr
4
+ from .common_metrics_on_video_quality.calculate_ssim import calculate_ssim
5
+
6
+ from .fvd import get_fvd_logits, frechet_distance, load_fvd_model
7
+ from .fid import calculate_frechet_distance
8
+ from .inception import InceptionV3
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ __pycache__
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from tqdm import tqdm
4
+
5
+ def trans(x):
6
+ # if greyscale images add channel
7
+ if x.shape[-3] == 1:
8
+ x = x.repeat(1, 1, 3, 1, 1)
9
+
10
+ # permute BTCHW -> BCTHW
11
+ x = x.permute(0, 2, 1, 3, 4)
12
+
13
+ return x
14
+
15
+ def calculate_fvd(videos1, videos2, device, method='styleganv'):
16
+
17
+ if method == 'styleganv':
18
+ from .fvd.styleganv.fvd import get_fvd_feats, frechet_distance, load_i3d_pretrained
19
+ elif method == 'videogpt':
20
+ from .fvd.videogpt.fvd import load_i3d_pretrained
21
+ from .fvd.videogpt.fvd import get_fvd_logits as get_fvd_feats
22
+ from .fvd.videogpt.fvd import frechet_distance
23
+
24
+ print("calculate_fvd...")
25
+
26
+ # videos [batch_size, timestamps, channel, h, w]
27
+
28
+ assert videos1.shape == videos2.shape
29
+
30
+ i3d = load_i3d_pretrained(device=device)
31
+ fvd_results = []
32
+
33
+ # support grayscale input, if grayscale -> channel*3
34
+ # BTCHW -> BCTHW
35
+ # videos -> [batch_size, channel, timestamps, h, w]
36
+
37
+ videos1 = trans(videos1)
38
+ videos2 = trans(videos2)
39
+
40
+ fvd_results = {}
41
+
42
+ # for calculate FVD, each clip_timestamp must >= 10
43
+ for clip_timestamp in tqdm(range(10, videos1.shape[-3]+1)):
44
+
45
+ # get a video clip
46
+ # videos_clip [batch_size, channel, timestamps[:clip], h, w]
47
+ videos_clip1 = videos1[:, :, : clip_timestamp]
48
+ videos_clip2 = videos2[:, :, : clip_timestamp]
49
+
50
+ # get FVD features
51
+ feats1 = get_fvd_feats(videos_clip1, i3d=i3d, device=device)
52
+ feats2 = get_fvd_feats(videos_clip2, i3d=i3d, device=device)
53
+
54
+ # calculate FVD when timestamps[:clip]
55
+ fvd_results[clip_timestamp] = frechet_distance(feats1, feats2)
56
+
57
+ result = {
58
+ "value": fvd_results,
59
+ "video_setting": videos1.shape,
60
+ "video_setting_name": "batch_size, channel, time, heigth, width",
61
+ }
62
+
63
+ return result
64
+
65
+ # test code / using example
66
+
67
+ def main():
68
+ NUMBER_OF_VIDEOS = 8
69
+ VIDEO_LENGTH = 50
70
+ CHANNEL = 3
71
+ SIZE = 64
72
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
73
+ videos2 = torch.ones(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
74
+ device = torch.device("cuda")
75
+ # device = torch.device("cpu")
76
+
77
+ import json
78
+ result = calculate_fvd(videos1, videos2, device, method='videogpt')
79
+ print(json.dumps(result, indent=4))
80
+
81
+ result = calculate_fvd(videos1, videos2, device, method='styleganv')
82
+ print(json.dumps(result, indent=4))
83
+
84
+ if __name__ == "__main__":
85
+ main()
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py ADDED
@@ -0,0 +1,72 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from tqdm import tqdm
4
+ import math
5
+
6
+ import torch
7
+ import lpips
8
+ from .utils import build_dataloader
9
+
10
+ spatial = True # Return a spatial map of perceptual distance.
11
+
12
+ # Linearly calibrated models (LPIPS)
13
+ loss_fn = lpips.LPIPS(net='vgg', spatial=spatial) # Can also set net = 'squeeze' or 'vgg'
14
+ # loss_fn = lpips.LPIPS(net='alex', spatial=spatial, lpips=False) # Can also set net = 'squeeze' or 'vgg'
15
+
16
+ def calculate_lpips(videos1, videos2, device):
17
+ # image should be RGB, IMPORTANT: normalized to [-1,1]
18
+ print("calculate_lpips...")
19
+
20
+ lpips_results = {}
21
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
22
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
23
+ # get a video [timestamps, channel, h, w]
24
+
25
+ assert video1.shape == video2.shape
26
+ video1 = video1.squeeze(0) * 2 - 1
27
+ video2 = video2.squeeze(0) * 2 - 1
28
+
29
+ for clip_timestamp in range(len(video1)):
30
+ # get a img
31
+ # img [timestamps[x], channel, h, w]
32
+ # img [channel, h, w] tensor
33
+
34
+ img1 = video1[clip_timestamp].unsqueeze(0).to(device)
35
+ img2 = video2[clip_timestamp].unsqueeze(0).to(device)
36
+
37
+ loss_fn.to(device)
38
+
39
+ # calculate lpips of a video
40
+ value = np.array(loss_fn.forward(img1, img2).mean().detach().cpu().tolist())
41
+ if clip_timestamp not in lpips_results:
42
+ lpips_results[clip_timestamp] = value
43
+ else:
44
+ lpips_results[clip_timestamp] += value
45
+
46
+ for clip_timestamp in range(len(video1)):
47
+ lpips_results[clip_timestamp] /= len(dataloader1)
48
+
49
+ result = {
50
+ "value": lpips_results,
51
+ }
52
+
53
+ return result
54
+
55
+ # test code / using example
56
+
57
+ def main():
58
+ NUMBER_OF_VIDEOS = 8
59
+ VIDEO_LENGTH = 50
60
+ CHANNEL = 3
61
+ SIZE = 64
62
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
63
+ videos2 = torch.ones(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
64
+ device = torch.device("cuda")
65
+ # device = torch.device("cpu")
66
+
67
+ import json
68
+ result = calculate_lpips(videos1, videos2, device)
69
+ print(json.dumps(result, indent=4))
70
+
71
+ if __name__ == "__main__":
72
+ main()
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from tqdm import tqdm
4
+ import math
5
+ from multiprocessing import Pool
6
+ from .utils import build_dataloader
7
+
8
+ def img_psnr(img1, img2):
9
+ # [0,1]
10
+ # compute mse
11
+ # mse = np.mean((img1-img2)**2)
12
+ mse = np.mean((img1 / 1.0 - img2 / 1.0) ** 2)
13
+ # compute psnr
14
+ if mse < 1e-10:
15
+ return 100
16
+ psnr = 20 * math.log10(1 / math.sqrt(mse))
17
+ return psnr
18
+
19
+ def trans(x):
20
+ return x
21
+
22
+ def process_video(video_pair):
23
+ video1, video2 = video_pair
24
+ psnr_results_of_a_video = []
25
+ for clip_timestamp in range(len(video1)):
26
+ # get a img
27
+ # img [timestamps[x], channel, h, w]
28
+ # img [channel, h, w] numpy
29
+
30
+ img1 = video1[clip_timestamp].numpy()
31
+ img2 = video2[clip_timestamp].numpy()
32
+
33
+ # calculate psnr of a video
34
+ psnr_results_of_a_video.append(img_psnr(img1, img2))
35
+
36
+ return psnr_results_of_a_video
37
+
38
+
39
+ def calculate_psnr(videos1, videos2):
40
+ print("calculate_psnr...")
41
+
42
+ # videos [batch_size, timestamps, channel, h, w]
43
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
44
+ psnr_results = []
45
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
46
+ video1 = video1.squeeze(0)
47
+ video2 = video2.squeeze(0)
48
+ video_pair = (video1, video2)
49
+ result = process_video(video_pair)
50
+ psnr_results.append(result)
51
+
52
+ psnr_results = np.array(psnr_results)
53
+
54
+ psnr = {}
55
+ psnr_std = {}
56
+
57
+ for clip_timestamp in range(len(video1)):
58
+ psnr[clip_timestamp] = np.mean(psnr_results[:,clip_timestamp])
59
+ psnr_std[clip_timestamp] = np.std(psnr_results[:,clip_timestamp])
60
+
61
+ result = {
62
+ "value": psnr,
63
+ "value_std": psnr_std,
64
+ }
65
+
66
+ return result
67
+
68
+ # test code / using example
69
+
70
+ def main():
71
+ NUMBER_OF_VIDEOS = 8
72
+ VIDEO_LENGTH = 50
73
+ CHANNEL = 3
74
+ SIZE = 64
75
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
76
+ videos2 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
77
+
78
+ import json
79
+ result = calculate_psnr(videos1, videos2)
80
+ print(json.dumps(result, indent=4))
81
+
82
+ if __name__ == "__main__":
83
+ main()
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py ADDED
@@ -0,0 +1,138 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from tqdm import tqdm
4
+ import cv2
5
+ from multiprocessing import Pool
6
+ import torch
7
+ import torch.nn.functional as F
8
+ from .utils import build_dataloader
9
+
10
+ def ssim(img1, img2):
11
+ _img1, _img2 = img1, img2
12
+
13
+ ### original implementation
14
+ # C1 = 0.01 ** 2
15
+ # C2 = 0.03 ** 2
16
+ # img1 = img1.astype(np.float64)
17
+ # img2 = img2.astype(np.float64)
18
+ # kernel = cv2.getGaussianKernel(11, 1.5)
19
+ # window = np.outer(kernel, kernel.transpose())
20
+ # mu1 = cv2.filter2D(img1, -1, window)[5:-5, 5:-5] # valid
21
+ # mu2 = cv2.filter2D(img2, -1, window)[5:-5, 5:-5]
22
+ # mu1_sq = mu1 ** 2
23
+ # mu2_sq = mu2 ** 2
24
+ # mu1_mu2 = mu1 * mu2
25
+ # sigma1_sq = cv2.filter2D(img1 ** 2, -1, window)[5:-5, 5:-5] - mu1_sq
26
+ # sigma2_sq = cv2.filter2D(img2 ** 2, -1, window)[5:-5, 5:-5] - mu2_sq
27
+ # sigma12 = cv2.filter2D(img1 * img2, -1, window)[5:-5, 5:-5] - mu1_mu2
28
+ # ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) *
29
+ # (sigma1_sq + sigma2_sq + C2))
30
+ # return ssim_map.mean()
31
+ # res1 = ssim_map.mean()
32
+
33
+ ### accelerated implementation
34
+ img1, img2 = torch.from_numpy(_img1), torch.from_numpy(_img2)
35
+ C1 = 0.01 ** 2
36
+ C2 = 0.03 ** 2
37
+
38
+ # Ensure data is float and move to GPU
39
+ img1 = img1.to(torch.float64).cuda()
40
+ img2 = img2.to(torch.float64).cuda()
41
+
42
+ # Gaussian kernel
43
+ kernel = torch.tensor(cv2.getGaussianKernel(11, 1.5)).to(torch.float64).cuda()
44
+ window = kernel @ kernel.t()
45
+ window = window.unsqueeze(0).unsqueeze(0).cuda()
46
+
47
+ mu1 = F.conv2d(img1.unsqueeze(0), window, padding=0, groups=1)
48
+ mu2 = F.conv2d(img2.unsqueeze(0), window, padding=0, groups=1)
49
+ mu1_sq = mu1.pow(2)
50
+ mu2_sq = mu2.pow(2)
51
+ mu1_mu2 = mu1 * mu2
52
+ sigma1_sq = F.conv2d(img1.unsqueeze(0) ** 2, window, padding=0, groups=1) - mu1_sq
53
+ sigma2_sq = F.conv2d(img2.unsqueeze(0) ** 2, window, padding=0, groups=1) - mu2_sq
54
+ sigma12 = F.conv2d(img1.unsqueeze(0) * img2.unsqueeze(0), window, padding=0, groups=1) - mu1_mu2
55
+
56
+ ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) *
57
+ (sigma1_sq + sigma2_sq + C2))
58
+ res2 = ssim_map.mean().item()
59
+ # print(res1-res2)
60
+ return res2
61
+ # return ssim_map.mean().item()
62
+
63
+
64
+ def calculate_ssim_function(img1, img2):
65
+ # [0,1]
66
+ # ssim is the only metric extremely sensitive to gray being compared to b/w
67
+ if not img1.shape == img2.shape:
68
+ raise ValueError('Input images must have the same dimensions.')
69
+ if img1.ndim == 2:
70
+ return ssim(img1, img2)
71
+ elif img1.ndim == 3:
72
+ if img1.shape[0] == 3:
73
+ ssims = []
74
+ for i in range(3):
75
+ ssims.append(ssim(img1[i], img2[i]))
76
+ return np.array(ssims).mean()
77
+ elif img1.shape[0] == 1:
78
+ return ssim(np.squeeze(img1), np.squeeze(img2))
79
+ else:
80
+ raise ValueError('Wrong input image dimensions.')
81
+
82
+ def trans(x):
83
+ return x
84
+
85
+ def process_video(video_pair):
86
+ video1, video2 = video_pair
87
+ ssim_results_of_a_video = []
88
+ for clip_timestamp in range(len(video1)):
89
+ img1 = video1[clip_timestamp].numpy()
90
+ img2 = video2[clip_timestamp].numpy()
91
+ ssim_results_of_a_video.append(calculate_ssim_function(img1, img2))
92
+ return ssim_results_of_a_video
93
+
94
+ def calculate_ssim(videos1, videos2):
95
+ print("calculate_ssim...")
96
+
97
+ ssim_results = []
98
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
99
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
100
+ video1 = video1.squeeze(0)
101
+ video2 = video2.squeeze(0)
102
+ video_pair = (video1, video2)
103
+ result = process_video(video_pair)
104
+ ssim_results.append(result)
105
+
106
+ ssim_results = np.array(ssim_results)
107
+
108
+ ssim = {}
109
+ ssim_std = {}
110
+
111
+ for clip_timestamp in range(len(video1)):
112
+ ssim[clip_timestamp] = np.mean(ssim_results[:,clip_timestamp])
113
+ ssim_std[clip_timestamp] = np.std(ssim_results[:,clip_timestamp])
114
+
115
+ result = {
116
+ "value": ssim,
117
+ "value_std": ssim_std,
118
+ }
119
+
120
+ return result
121
+
122
+ # test code / using example
123
+
124
+ def main():
125
+ NUMBER_OF_VIDEOS = 8
126
+ VIDEO_LENGTH = 50
127
+ CHANNEL = 3
128
+ SIZE = 64
129
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
130
+ videos2 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
131
+ device = torch.device("cuda")
132
+
133
+ import json
134
+ result = calculate_ssim(videos1, videos2)
135
+ print(json.dumps(result, indent=4))
136
+
137
+ if __name__ == "__main__":
138
+ main()
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import os
3
+ import math
4
+ import torch.nn.functional as F
5
+
6
+ # https://github.com/universome/fvd-comparison
7
+
8
+
9
+ def load_i3d_pretrained(device=torch.device('cpu')):
10
+ i3D_WEIGHTS_URL = "https://www.dropbox.com/s/ge9e5ujwgetktms/i3d_torchscript.pt"
11
+ filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'i3d_torchscript.pt')
12
+ print(filepath)
13
+ if not os.path.exists(filepath):
14
+ print(f"preparing for download {i3D_WEIGHTS_URL}, you can download it by yourself.")
15
+ os.system(f"wget {i3D_WEIGHTS_URL} -O {filepath}")
16
+ i3d = torch.jit.load(filepath).eval().to(device)
17
+ i3d = torch.nn.DataParallel(i3d)
18
+ return i3d
19
+
20
+
21
+ def get_feats(videos, detector, device, bs=10):
22
+ # videos : torch.tensor BCTHW [0, 1]
23
+ detector_kwargs = dict(rescale=False, resize=False, return_features=True) # Return raw features before the softmax layer.
24
+ feats = np.empty((0, 400))
25
+ with torch.no_grad():
26
+ for i in range((len(videos)-1)//bs + 1):
27
+ feats = np.vstack([feats, detector(torch.stack([preprocess_single(video) for video in videos[i*bs:(i+1)*bs]]).to(device), **detector_kwargs).detach().cpu().numpy()])
28
+ return feats
29
+
30
+
31
+ def get_fvd_feats(videos, i3d, device, bs=10):
32
+ # videos in [0, 1] as torch tensor BCTHW
33
+ # videos = [preprocess_single(video) for video in videos]
34
+ embeddings = get_feats(videos, i3d, device, bs)
35
+ return embeddings
36
+
37
+
38
+ def preprocess_single(video, resolution=224, sequence_length=None):
39
+ # video: CTHW, [0, 1]
40
+ c, t, h, w = video.shape
41
+
42
+ # temporal crop
43
+ if sequence_length is not None:
44
+ assert sequence_length <= t
45
+ video = video[:, :sequence_length]
46
+
47
+ # scale shorter side to resolution
48
+ scale = resolution / min(h, w)
49
+ if h < w:
50
+ target_size = (resolution, math.ceil(w * scale))
51
+ else:
52
+ target_size = (math.ceil(h * scale), resolution)
53
+ video = F.interpolate(video, size=target_size, mode='bilinear', align_corners=False)
54
+
55
+ # center crop
56
+ c, t, h, w = video.shape
57
+ w_start = (w - resolution) // 2
58
+ h_start = (h - resolution) // 2
59
+ video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
60
+
61
+ # [0, 1] -> [-1, 1]
62
+ video = (video - 0.5) * 2
63
+
64
+ return video.contiguous()
65
+
66
+
67
+ """
68
+ Copy-pasted from https://github.com/cvpr2022-stylegan-v/stylegan-v/blob/main/src/metrics/frechet_video_distance.py
69
+ """
70
+ from typing import Tuple
71
+ from scipy.linalg import sqrtm
72
+ import numpy as np
73
+
74
+
75
+ def compute_stats(feats: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
76
+ mu = feats.mean(axis=0) # [d]
77
+ sigma = np.cov(feats, rowvar=False) # [d, d]
78
+ return mu, sigma
79
+
80
+
81
+ def frechet_distance(feats_fake: np.ndarray, feats_real: np.ndarray) -> float:
82
+ mu_gen, sigma_gen = compute_stats(feats_fake)
83
+ mu_real, sigma_real = compute_stats(feats_real)
84
+ m = np.square(mu_gen - mu_real).sum()
85
+ if feats_fake.shape[0]>1:
86
+ s, _ = sqrtm(np.dot(sigma_gen, sigma_real), disp=False) # pylint: disable=no-member
87
+ fid = np.real(m + np.trace(sigma_gen + sigma_real - s * 2))
88
+ else:
89
+ fid = np.real(m)
90
+ return float(fid)
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import os
3
+ import math
4
+ import torch.nn.functional as F
5
+ import numpy as np
6
+ import einops
7
+
8
+ def load_i3d_pretrained(device=torch.device('cpu')):
9
+ i3D_WEIGHTS_URL = "https://onedrive.live.com/download?cid=78EEF3EB6AE7DBCB&resid=78EEF3EB6AE7DBCB%21199&authkey=AApKdFHPXzWLNyI"
10
+ filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'i3d_pretrained_400.pt')
11
+ print(filepath)
12
+ if not os.path.exists(filepath):
13
+ print(f"preparing for download {i3D_WEIGHTS_URL}, you can download it by yourself.")
14
+ os.system(f"wget {i3D_WEIGHTS_URL} -O {filepath}")
15
+ from .pytorch_i3d import InceptionI3d
16
+ i3d = InceptionI3d(400, in_channels=3).eval().to(device)
17
+ i3d.load_state_dict(torch.load(filepath, map_location=device, weights_only=True))
18
+ i3d = torch.nn.DataParallel(i3d)
19
+ return i3d
20
+
21
+ def preprocess_single(video, resolution, sequence_length=None):
22
+ # video: THWC, {0, ..., 255}
23
+ video = video.permute(0, 3, 1, 2).float() / 255. # TCHW
24
+ t, c, h, w = video.shape
25
+
26
+ # temporal crop
27
+ if sequence_length is not None:
28
+ assert sequence_length <= t
29
+ video = video[:sequence_length]
30
+
31
+ # scale shorter side to resolution
32
+ scale = resolution / min(h, w)
33
+ if h < w:
34
+ target_size = (resolution, math.ceil(w * scale))
35
+ else:
36
+ target_size = (math.ceil(h * scale), resolution)
37
+ video = F.interpolate(video, size=target_size, mode='bilinear',
38
+ align_corners=False)
39
+
40
+ # center crop
41
+ t, c, h, w = video.shape
42
+ w_start = (w - resolution) // 2
43
+ h_start = (h - resolution) // 2
44
+ video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
45
+ video = video.permute(1, 0, 2, 3).contiguous() # CTHW
46
+
47
+ video -= 0.5
48
+
49
+ return video
50
+
51
+ def preprocess(videos, target_resolution=224):
52
+ # we should tras videos in [0-1] [b c t h w] as th.float
53
+ # -> videos in {0, ..., 255} [b t h w c] as np.uint8 array
54
+ videos = einops.rearrange(videos, 'b c t h w -> b t h w c')
55
+ videos = (videos*255).numpy().astype(np.uint8)
56
+
57
+ b, t, h, w, c = videos.shape
58
+ videos = torch.from_numpy(videos)
59
+ videos = torch.stack([preprocess_single(video, target_resolution) for video in videos])
60
+ return videos * 2 # [-0.5, 0.5] -> [-1, 1]
61
+
62
+ def get_fvd_logits(videos, i3d, device, bs=10):
63
+ videos = preprocess(videos)
64
+ embeddings = get_logits(i3d, videos, device, bs=10)
65
+ return embeddings
66
+
67
+ # https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L161
68
+ def _symmetric_matrix_square_root(mat, eps=1e-10):
69
+ u, s, v = torch.svd(mat)
70
+ si = torch.where(s < eps, s, torch.sqrt(s))
71
+ return torch.matmul(torch.matmul(u, torch.diag(si)), v.t())
72
+
73
+ # https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L400
74
+ def trace_sqrt_product(sigma, sigma_v):
75
+ sqrt_sigma = _symmetric_matrix_square_root(sigma)
76
+ sqrt_a_sigmav_a = torch.matmul(sqrt_sigma, torch.matmul(sigma_v, sqrt_sigma))
77
+ return torch.trace(_symmetric_matrix_square_root(sqrt_a_sigmav_a))
78
+
79
+ # https://discuss.pytorch.org/t/covariance-and-gradient-support/16217/2
80
+ def cov(m, rowvar=False):
81
+ '''Estimate a covariance matrix given data.
82
+
83
+ Covariance indicates the level to which two variables vary together.
84
+ If we examine N-dimensional samples, `X = [x_1, x_2, ... x_N]^T`,
85
+ then the covariance matrix element `C_{ij}` is the covariance of
86
+ `x_i` and `x_j`. The element `C_{ii}` is the variance of `x_i`.
87
+
88
+ Args:
89
+ m: A 1-D or 2-D array containing multiple variables and observations.
90
+ Each row of `m` represents a variable, and each column a single
91
+ observation of all those variables.
92
+ rowvar: If `rowvar` is True, then each row represents a
93
+ variable, with observations in the columns. Otherwise, the
94
+ relationship is transposed: each column represents a variable,
95
+ while the rows contain observations.
96
+
97
+ Returns:
98
+ The covariance matrix of the variables.
99
+ '''
100
+ if m.dim() > 2:
101
+ raise ValueError('m has more than 2 dimensions')
102
+ if m.dim() < 2:
103
+ m = m.view(1, -1)
104
+ if not rowvar and m.size(0) != 1:
105
+ m = m.t()
106
+
107
+ fact = 1.0 / (m.size(1) - 1) # unbiased estimate
108
+ m -= torch.mean(m, dim=1, keepdim=True)
109
+ mt = m.t() # if complex: mt = m.t().conj()
110
+ return fact * m.matmul(mt).squeeze()
111
+
112
+
113
+ def frechet_distance(x1, x2):
114
+ x1 = x1.flatten(start_dim=1)
115
+ x2 = x2.flatten(start_dim=1)
116
+ m, m_w = x1.mean(dim=0), x2.mean(dim=0)
117
+ sigma, sigma_w = cov(x1, rowvar=False), cov(x2, rowvar=False)
118
+ mean = torch.sum((m - m_w) ** 2)
119
+ if x1.shape[0]>1:
120
+ sqrt_trace_component = trace_sqrt_product(sigma, sigma_w)
121
+ trace = torch.trace(sigma + sigma_w) - 2.0 * sqrt_trace_component
122
+ fd = trace + mean
123
+ else:
124
+ fd = np.real(mean)
125
+ return float(fd)
126
+
127
+
128
+ def get_logits(i3d, videos, device, bs=10):
129
+ # assert videos.shape[0] % 16 == 0
130
+ with torch.no_grad():
131
+ logits = []
132
+ for i in range(0, videos.shape[0], bs):
133
+ batch = videos[i:i + bs].to(device)
134
+ # logits.append(i3d.module.extract_features(batch)) # wrong
135
+ logits.append(i3d(batch)) # right
136
+ logits = torch.cat(logits, dim=0)
137
+ return logits
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py ADDED
@@ -0,0 +1,322 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Original code from https://github.com/piergiaj/pytorch-i3d
2
+ import torch
3
+ import torch.nn as nn
4
+ import torch.nn.functional as F
5
+ import numpy as np
6
+
7
+ class MaxPool3dSamePadding(nn.MaxPool3d):
8
+
9
+ def compute_pad(self, dim, s):
10
+ if s % self.stride[dim] == 0:
11
+ return max(self.kernel_size[dim] - self.stride[dim], 0)
12
+ else:
13
+ return max(self.kernel_size[dim] - (s % self.stride[dim]), 0)
14
+
15
+ def forward(self, x):
16
+ # compute 'same' padding
17
+ (batch, channel, t, h, w) = x.size()
18
+ out_t = np.ceil(float(t) / float(self.stride[0]))
19
+ out_h = np.ceil(float(h) / float(self.stride[1]))
20
+ out_w = np.ceil(float(w) / float(self.stride[2]))
21
+ pad_t = self.compute_pad(0, t)
22
+ pad_h = self.compute_pad(1, h)
23
+ pad_w = self.compute_pad(2, w)
24
+
25
+ pad_t_f = pad_t // 2
26
+ pad_t_b = pad_t - pad_t_f
27
+ pad_h_f = pad_h // 2
28
+ pad_h_b = pad_h - pad_h_f
29
+ pad_w_f = pad_w // 2
30
+ pad_w_b = pad_w - pad_w_f
31
+
32
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
33
+ x = F.pad(x, pad)
34
+ return super(MaxPool3dSamePadding, self).forward(x)
35
+
36
+
37
+ class Unit3D(nn.Module):
38
+
39
+ def __init__(self, in_channels,
40
+ output_channels,
41
+ kernel_shape=(1, 1, 1),
42
+ stride=(1, 1, 1),
43
+ padding=0,
44
+ activation_fn=F.relu,
45
+ use_batch_norm=True,
46
+ use_bias=False,
47
+ name='unit_3d'):
48
+
49
+ """Initializes Unit3D module."""
50
+ super(Unit3D, self).__init__()
51
+
52
+ self._output_channels = output_channels
53
+ self._kernel_shape = kernel_shape
54
+ self._stride = stride
55
+ self._use_batch_norm = use_batch_norm
56
+ self._activation_fn = activation_fn
57
+ self._use_bias = use_bias
58
+ self.name = name
59
+ self.padding = padding
60
+
61
+ self.conv3d = nn.Conv3d(in_channels=in_channels,
62
+ out_channels=self._output_channels,
63
+ kernel_size=self._kernel_shape,
64
+ stride=self._stride,
65
+ padding=0, # we always want padding to be 0 here. We will dynamically pad based on input size in forward function
66
+ bias=self._use_bias)
67
+
68
+ if self._use_batch_norm:
69
+ self.bn = nn.BatchNorm3d(self._output_channels, eps=1e-5, momentum=0.001)
70
+
71
+ def compute_pad(self, dim, s):
72
+ if s % self._stride[dim] == 0:
73
+ return max(self._kernel_shape[dim] - self._stride[dim], 0)
74
+ else:
75
+ return max(self._kernel_shape[dim] - (s % self._stride[dim]), 0)
76
+
77
+
78
+ def forward(self, x):
79
+ # compute 'same' padding
80
+ (batch, channel, t, h, w) = x.size()
81
+ out_t = np.ceil(float(t) / float(self._stride[0]))
82
+ out_h = np.ceil(float(h) / float(self._stride[1]))
83
+ out_w = np.ceil(float(w) / float(self._stride[2]))
84
+ pad_t = self.compute_pad(0, t)
85
+ pad_h = self.compute_pad(1, h)
86
+ pad_w = self.compute_pad(2, w)
87
+
88
+ pad_t_f = pad_t // 2
89
+ pad_t_b = pad_t - pad_t_f
90
+ pad_h_f = pad_h // 2
91
+ pad_h_b = pad_h - pad_h_f
92
+ pad_w_f = pad_w // 2
93
+ pad_w_b = pad_w - pad_w_f
94
+
95
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
96
+ x = F.pad(x, pad)
97
+
98
+ x = self.conv3d(x)
99
+ if self._use_batch_norm:
100
+ x = self.bn(x)
101
+ if self._activation_fn is not None:
102
+ x = self._activation_fn(x)
103
+ return x
104
+
105
+
106
+
107
+ class InceptionModule(nn.Module):
108
+ def __init__(self, in_channels, out_channels, name):
109
+ super(InceptionModule, self).__init__()
110
+
111
+ self.b0 = Unit3D(in_channels=in_channels, output_channels=out_channels[0], kernel_shape=[1, 1, 1], padding=0,
112
+ name=name+'/Branch_0/Conv3d_0a_1x1')
113
+ self.b1a = Unit3D(in_channels=in_channels, output_channels=out_channels[1], kernel_shape=[1, 1, 1], padding=0,
114
+ name=name+'/Branch_1/Conv3d_0a_1x1')
115
+ self.b1b = Unit3D(in_channels=out_channels[1], output_channels=out_channels[2], kernel_shape=[3, 3, 3],
116
+ name=name+'/Branch_1/Conv3d_0b_3x3')
117
+ self.b2a = Unit3D(in_channels=in_channels, output_channels=out_channels[3], kernel_shape=[1, 1, 1], padding=0,
118
+ name=name+'/Branch_2/Conv3d_0a_1x1')
119
+ self.b2b = Unit3D(in_channels=out_channels[3], output_channels=out_channels[4], kernel_shape=[3, 3, 3],
120
+ name=name+'/Branch_2/Conv3d_0b_3x3')
121
+ self.b3a = MaxPool3dSamePadding(kernel_size=[3, 3, 3],
122
+ stride=(1, 1, 1), padding=0)
123
+ self.b3b = Unit3D(in_channels=in_channels, output_channels=out_channels[5], kernel_shape=[1, 1, 1], padding=0,
124
+ name=name+'/Branch_3/Conv3d_0b_1x1')
125
+ self.name = name
126
+
127
+ def forward(self, x):
128
+ b0 = self.b0(x)
129
+ b1 = self.b1b(self.b1a(x))
130
+ b2 = self.b2b(self.b2a(x))
131
+ b3 = self.b3b(self.b3a(x))
132
+ return torch.cat([b0,b1,b2,b3], dim=1)
133
+
134
+
135
+ class InceptionI3d(nn.Module):
136
+ """Inception-v1 I3D architecture.
137
+ The model is introduced in:
138
+ Quo Vadis, Action Recognition? A New Model and the Kinetics Dataset
139
+ Joao Carreira, Andrew Zisserman
140
+ https://arxiv.org/pdf/1705.07750v1.pdf.
141
+ See also the Inception architecture, introduced in:
142
+ Going deeper with convolutions
143
+ Christian Szegedy, Wei Liu, Yangqing Jia, Pierre Sermanet, Scott Reed,
144
+ Dragomir Anguelov, Dumitru Erhan, Vincent Vanhoucke, Andrew Rabinovich.
145
+ http://arxiv.org/pdf/1409.4842v1.pdf.
146
+ """
147
+
148
+ # Endpoints of the model in order. During construction, all the endpoints up
149
+ # to a designated `final_endpoint` are returned in a dictionary as the
150
+ # second return value.
151
+ VALID_ENDPOINTS = (
152
+ 'Conv3d_1a_7x7',
153
+ 'MaxPool3d_2a_3x3',
154
+ 'Conv3d_2b_1x1',
155
+ 'Conv3d_2c_3x3',
156
+ 'MaxPool3d_3a_3x3',
157
+ 'Mixed_3b',
158
+ 'Mixed_3c',
159
+ 'MaxPool3d_4a_3x3',
160
+ 'Mixed_4b',
161
+ 'Mixed_4c',
162
+ 'Mixed_4d',
163
+ 'Mixed_4e',
164
+ 'Mixed_4f',
165
+ 'MaxPool3d_5a_2x2',
166
+ 'Mixed_5b',
167
+ 'Mixed_5c',
168
+ 'Logits',
169
+ 'Predictions',
170
+ )
171
+
172
+ def __init__(self, num_classes=400, spatial_squeeze=True,
173
+ final_endpoint='Logits', name='inception_i3d', in_channels=3, dropout_keep_prob=0.5):
174
+ """Initializes I3D model instance.
175
+ Args:
176
+ num_classes: The number of outputs in the logit layer (default 400, which
177
+ matches the Kinetics dataset).
178
+ spatial_squeeze: Whether to squeeze the spatial dimensions for the logits
179
+ before returning (default True).
180
+ final_endpoint: The model contains many possible endpoints.
181
+ `final_endpoint` specifies the last endpoint for the model to be built
182
+ up to. In addition to the output at `final_endpoint`, all the outputs
183
+ at endpoints up to `final_endpoint` will also be returned, in a
184
+ dictionary. `final_endpoint` must be one of
185
+ InceptionI3d.VALID_ENDPOINTS (default 'Logits').
186
+ name: A string (optional). The name of this module.
187
+ Raises:
188
+ ValueError: if `final_endpoint` is not recognized.
189
+ """
190
+
191
+ if final_endpoint not in self.VALID_ENDPOINTS:
192
+ raise ValueError('Unknown final endpoint %s' % final_endpoint)
193
+
194
+ super(InceptionI3d, self).__init__()
195
+ self._num_classes = num_classes
196
+ self._spatial_squeeze = spatial_squeeze
197
+ self._final_endpoint = final_endpoint
198
+ self.logits = None
199
+
200
+ if self._final_endpoint not in self.VALID_ENDPOINTS:
201
+ raise ValueError('Unknown final endpoint %s' % self._final_endpoint)
202
+
203
+ self.end_points = {}
204
+ end_point = 'Conv3d_1a_7x7'
205
+ self.end_points[end_point] = Unit3D(in_channels=in_channels, output_channels=64, kernel_shape=[7, 7, 7],
206
+ stride=(2, 2, 2), padding=(3,3,3), name=name+end_point)
207
+ if self._final_endpoint == end_point: return
208
+
209
+ end_point = 'MaxPool3d_2a_3x3'
210
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
211
+ padding=0)
212
+ if self._final_endpoint == end_point: return
213
+
214
+ end_point = 'Conv3d_2b_1x1'
215
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=64, kernel_shape=[1, 1, 1], padding=0,
216
+ name=name+end_point)
217
+ if self._final_endpoint == end_point: return
218
+
219
+ end_point = 'Conv3d_2c_3x3'
220
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=192, kernel_shape=[3, 3, 3], padding=1,
221
+ name=name+end_point)
222
+ if self._final_endpoint == end_point: return
223
+
224
+ end_point = 'MaxPool3d_3a_3x3'
225
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
226
+ padding=0)
227
+ if self._final_endpoint == end_point: return
228
+
229
+ end_point = 'Mixed_3b'
230
+ self.end_points[end_point] = InceptionModule(192, [64,96,128,16,32,32], name+end_point)
231
+ if self._final_endpoint == end_point: return
232
+
233
+ end_point = 'Mixed_3c'
234
+ self.end_points[end_point] = InceptionModule(256, [128,128,192,32,96,64], name+end_point)
235
+ if self._final_endpoint == end_point: return
236
+
237
+ end_point = 'MaxPool3d_4a_3x3'
238
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[3, 3, 3], stride=(2, 2, 2),
239
+ padding=0)
240
+ if self._final_endpoint == end_point: return
241
+
242
+ end_point = 'Mixed_4b'
243
+ self.end_points[end_point] = InceptionModule(128+192+96+64, [192,96,208,16,48,64], name+end_point)
244
+ if self._final_endpoint == end_point: return
245
+
246
+ end_point = 'Mixed_4c'
247
+ self.end_points[end_point] = InceptionModule(192+208+48+64, [160,112,224,24,64,64], name+end_point)
248
+ if self._final_endpoint == end_point: return
249
+
250
+ end_point = 'Mixed_4d'
251
+ self.end_points[end_point] = InceptionModule(160+224+64+64, [128,128,256,24,64,64], name+end_point)
252
+ if self._final_endpoint == end_point: return
253
+
254
+ end_point = 'Mixed_4e'
255
+ self.end_points[end_point] = InceptionModule(128+256+64+64, [112,144,288,32,64,64], name+end_point)
256
+ if self._final_endpoint == end_point: return
257
+
258
+ end_point = 'Mixed_4f'
259
+ self.end_points[end_point] = InceptionModule(112+288+64+64, [256,160,320,32,128,128], name+end_point)
260
+ if self._final_endpoint == end_point: return
261
+
262
+ end_point = 'MaxPool3d_5a_2x2'
263
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[2, 2, 2], stride=(2, 2, 2),
264
+ padding=0)
265
+ if self._final_endpoint == end_point: return
266
+
267
+ end_point = 'Mixed_5b'
268
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [256,160,320,32,128,128], name+end_point)
269
+ if self._final_endpoint == end_point: return
270
+
271
+ end_point = 'Mixed_5c'
272
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [384,192,384,48,128,128], name+end_point)
273
+ if self._final_endpoint == end_point: return
274
+
275
+ end_point = 'Logits'
276
+ self.avg_pool = nn.AvgPool3d(kernel_size=[2, 7, 7],
277
+ stride=(1, 1, 1))
278
+ self.dropout = nn.Dropout(dropout_keep_prob)
279
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
280
+ kernel_shape=[1, 1, 1],
281
+ padding=0,
282
+ activation_fn=None,
283
+ use_batch_norm=False,
284
+ use_bias=True,
285
+ name='logits')
286
+
287
+ self.build()
288
+
289
+
290
+ def replace_logits(self, num_classes):
291
+ self._num_classes = num_classes
292
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
293
+ kernel_shape=[1, 1, 1],
294
+ padding=0,
295
+ activation_fn=None,
296
+ use_batch_norm=False,
297
+ use_bias=True,
298
+ name='logits')
299
+
300
+
301
+ def build(self):
302
+ for k in self.end_points.keys():
303
+ self.add_module(k, self.end_points[k])
304
+
305
+ def forward(self, x):
306
+ for end_point in self.VALID_ENDPOINTS:
307
+ if end_point in self.end_points:
308
+ x = self._modules[end_point](x) # use _modules to work with dataparallel
309
+
310
+ x = self.logits(self.dropout(self.avg_pool(x)))
311
+ if self._spatial_squeeze:
312
+ logits = x.squeeze(3).squeeze(3)
313
+ logits = logits.mean(dim=2)
314
+ # logits is batch X time X classes, which is what we want to work with
315
+ return logits
316
+
317
+
318
+ def extract_features(self, x):
319
+ for end_point in self.VALID_ENDPOINTS:
320
+ if end_point in self.end_points:
321
+ x = self._modules[end_point](x)
322
+ return self.avg_pool(x)
grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from torch.utils.data import Dataset, DataLoader
3
+
4
+ class VideoDataset(Dataset):
5
+ def __init__(self, videos):
6
+ self.videos = videos
7
+
8
+ def __len__(self):
9
+ return len(self.videos) if type(self.videos) == list else self.videos.shape[0]
10
+
11
+ def __getitem__(self, idx):
12
+ video = self.videos[idx]
13
+ if isinstance(video, str):
14
+ video = torch.load(video, weights_only=True)
15
+ return video
16
+
17
+ def build_dataloader(videos1, videos2):
18
+ dataset1 = VideoDataset(videos1)
19
+ dataset2 = VideoDataset(videos2)
20
+ assert len(dataset1) == len(dataset2)
21
+
22
+ dataloader1 = DataLoader(dataset1, batch_size=1, num_workers=8, shuffle=False)
23
+ dataloader2 = DataLoader(dataset2, batch_size=1, num_workers=8, shuffle=False)
24
+
25
+ return dataloader1, dataloader2
26
+
27
+
grn/tokenizer/videovae/evaluation/fid.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from scipy import linalg
3
+
4
+
5
+ def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
6
+ """Numpy implementation of the Frechet Distance.
7
+ The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
8
+ and X_2 ~ N(mu_2, C_2) is
9
+ d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
10
+
11
+ Stable version by Dougal J. Sutherland.
12
+
13
+ Params:
14
+ -- mu1 : Numpy array containing the activations of a layer of the
15
+ inception net (like returned by the function 'get_predictions')
16
+ for generated samples.
17
+ -- mu2 : The sample mean over activations, precalculated on an
18
+ representative data set.
19
+ -- sigma1: The covariance matrix over activations for generated samples.
20
+ -- sigma2: The covariance matrix over activations, precalculated on an
21
+ representative data set.
22
+
23
+ Returns:
24
+ -- : The Frechet Distance.
25
+ """
26
+
27
+ mu1 = np.atleast_1d(mu1)
28
+ mu2 = np.atleast_1d(mu2)
29
+
30
+ sigma1 = np.atleast_2d(sigma1)
31
+ sigma2 = np.atleast_2d(sigma2)
32
+
33
+ assert (
34
+ mu1.shape == mu2.shape
35
+ ), "Training and test mean vectors have different lengths"
36
+ assert (
37
+ sigma1.shape == sigma2.shape
38
+ ), "Training and test covariances have different dimensions"
39
+
40
+ diff = mu1 - mu2
41
+
42
+ # Product might be almost singular
43
+ covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
44
+ if not np.isfinite(covmean).all():
45
+ msg = (
46
+ "fid calculation produces singular product; "
47
+ "adding %s to diagonal of cov estimates"
48
+ ) % eps
49
+ print(msg)
50
+ offset = np.eye(sigma1.shape[0]) * eps
51
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
52
+
53
+ # Numerical error might give slight imaginary component
54
+ if np.iscomplexobj(covmean):
55
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
56
+ m = np.max(np.abs(covmean.imag))
57
+ raise ValueError("Imaginary component {}".format(m))
58
+ covmean = covmean.real
59
+
60
+ tr_covmean = np.trace(covmean)
61
+
62
+ return diff.dot(diff) + np.trace(sigma1) + np.trace(sigma2) - 2 * tr_covmean
grn/tokenizer/videovae/evaluation/fvd.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ from email.policy import strict
3
+ import numpy as np
4
+
5
+ import torch
6
+ import torch.nn.functional as F
7
+ import torch.utils.data as data
8
+
9
+ from .pytorch_i3d import InceptionI3d
10
+ import os
11
+ from videovae.utils.misc import data_prefix_manager
12
+
13
+ from sklearn.metrics.pairwise import polynomial_kernel
14
+
15
+ MAX_BATCH = 16
16
+ FVD_SAMPLE_SIZE = 2048
17
+ TARGET_RESOLUTION = (224, 224)
18
+
19
+ def preprocess(videos, target_resolution):
20
+ # videos in {0, ..., 255} as np.uint8 array
21
+ b, t, h, w, c = videos.shape
22
+ all_frames = torch.FloatTensor(videos).flatten(end_dim=1) # (b * t, h, w, c)
23
+
24
+ all_frames = all_frames.permute(0, 3, 1, 2).contiguous() # (b * t, c, h, w)
25
+ resized_videos = F.interpolate(all_frames, size=target_resolution,
26
+ mode='bilinear', align_corners=False)
27
+ resized_videos = resized_videos.view(b, t, c, *target_resolution)
28
+ output_videos = resized_videos.transpose(1, 2).contiguous() # (b, c, t, *)
29
+ scaled_videos = 2. * output_videos / 255. - 1 # [-1, 1]
30
+ return scaled_videos
31
+
32
+ def get_fvd_logits(videos, i3d, device):
33
+ videos = preprocess(videos, TARGET_RESOLUTION)
34
+ embeddings = get_logits(i3d, videos, device)
35
+ return embeddings
36
+
37
+ def load_fvd_model(device="cpu"):
38
+ i3d = InceptionI3d(400, in_channels=3).to(device)
39
+ i3d_path = data_prefix_manager('checkpoints/i3d_pretrained_400.pt')
40
+ i3d.load_state_dict(torch.load(i3d_path, map_location=device, weights_only=True))
41
+ i3d.eval()
42
+ return i3d
43
+
44
+
45
+ def load_i3d_perceptual():
46
+ i3d = InceptionI3d(400, in_channels=3)
47
+ current_dir = os.path.dirname(os.path.abspath(__file__))
48
+ i3d_path = os.path.join(current_dir, 'i3d_pretrained_400.pt')
49
+ i3d.load_state_dict(torch.load(i3d_path, map_location=torch.device("cpu"), weights_only=True), strict=False)
50
+ for param in i3d.parameters():
51
+ param.requires_grad = False
52
+
53
+ return i3d
54
+
55
+ # https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L161
56
+ def _symmetric_matrix_square_root(mat, eps=1e-10):
57
+ u, s, v = torch.svd(mat)
58
+ si = torch.where(s < eps, s, torch.sqrt(s))
59
+ return torch.matmul(torch.matmul(u, torch.diag(si)), v.t())
60
+
61
+ # https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L400
62
+ def trace_sqrt_product(sigma, sigma_v):
63
+ sqrt_sigma = _symmetric_matrix_square_root(sigma)
64
+ sqrt_a_sigmav_a = torch.matmul(sqrt_sigma, torch.matmul(sigma_v, sqrt_sigma))
65
+ return torch.trace(_symmetric_matrix_square_root(sqrt_a_sigmav_a))
66
+
67
+ # https://discuss.pytorch.org/t/covariance-and-gradient-support/16217/2
68
+ def cov(m, rowvar=False):
69
+ '''Estimate a covariance matrix given data.
70
+
71
+ Covariance indicates the level to which two variables vary together.
72
+ If we examine N-dimensional samples, `X = [x_1, x_2, ... x_N]^T`,
73
+ then the covariance matrix element `C_{ij}` is the covariance of
74
+ `x_i` and `x_j`. The element `C_{ii}` is the variance of `x_i`.
75
+
76
+ Args:
77
+ m: A 1-D or 2-D array containing multiple variables and observations.
78
+ Each row of `m` represents a variable, and each column a single
79
+ observation of all those variables.
80
+ rowvar: If `rowvar` is True, then each row represents a
81
+ variable, with observations in the columns. Otherwise, the
82
+ relationship is transposed: each column represents a variable,
83
+ while the rows contain observations.
84
+
85
+ Returns:
86
+ The covariance matrix of the variables.
87
+ '''
88
+ if m.dim() > 2:
89
+ raise ValueError('m has more than 2 dimensions')
90
+ if m.dim() < 2:
91
+ m = m.view(1, -1)
92
+ if not rowvar and m.size(0) != 1:
93
+ m = m.t()
94
+
95
+ fact = 1.0 / (m.size(1) - 1) # unbiased estimate
96
+ m_center = m - torch.mean(m, dim=1, keepdim=True)
97
+ mt = m_center.t() # if complex: mt = m.t().conj()
98
+ return fact * m_center.matmul(mt).squeeze()
99
+
100
+
101
+ def frechet_distance(x1, x2):
102
+ x1 = x1.flatten(start_dim=1)
103
+ x2 = x2.flatten(start_dim=1)
104
+ m, m_w = x1.mean(dim=0), x2.mean(dim=0)
105
+ sigma, sigma_w = cov(x1, rowvar=False), cov(x2, rowvar=False)
106
+
107
+ sqrt_trace_component = trace_sqrt_product(sigma, sigma_w)
108
+ trace = torch.trace(sigma + sigma_w) - 2.0 * sqrt_trace_component
109
+
110
+ mean = torch.sum((m - m_w) ** 2)
111
+ fd = trace + mean
112
+ return fd
113
+
114
+
115
+ def polynomial_mmd(X, Y):
116
+ m = X.shape[0]
117
+ n = Y.shape[0]
118
+ # compute kernels
119
+ K_XX = polynomial_kernel(X)
120
+ K_YY = polynomial_kernel(Y)
121
+ K_XY = polynomial_kernel(X, Y)
122
+ # compute mmd distance
123
+ K_XX_sum = (K_XX.sum() - np.diagonal(K_XX).sum()) / (m * (m - 1))
124
+ K_YY_sum = (K_YY.sum() - np.diagonal(K_YY).sum()) / (n * (n - 1))
125
+ K_XY_sum = K_XY.sum() / (m * n)
126
+ mmd = K_XX_sum + K_YY_sum - 2 * K_XY_sum
127
+ return mmd
128
+
129
+
130
+
131
+ def get_logits(i3d, videos, device):
132
+ # assert videos.shape[0] % MAX_BATCH == 0
133
+ with torch.no_grad():
134
+ logits = []
135
+ for i in range(0, videos.shape[0], MAX_BATCH):
136
+ batch = videos[i:i + MAX_BATCH].to(device)
137
+ logits.append(i3d(batch))
138
+ logits = torch.cat(logits, dim=0)
139
+ return logits
140
+
141
+
142
+ def compute_fvd(real, samples, i3d, device=torch.device('cpu')):
143
+ # i3d.to(device)
144
+ # real, samples are (N, T, H, W, C) numpy arrays in np.uint8
145
+ real, samples = preprocess(real, (224, 224)), preprocess(samples, (224, 224))
146
+ first_embed = get_logits(i3d, real, device)
147
+ second_embed = get_logits(i3d, samples, device)
148
+
149
+ return frechet_distance(first_embed, second_embed)
150
+
grn/tokenizer/videovae/evaluation/inception.py ADDED
@@ -0,0 +1,370 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ from torchvision import models
5
+ from scipy import linalg
6
+ import numpy as np
7
+
8
+ try:
9
+ from torchvision.models.utils import load_state_dict_from_url
10
+ except ImportError:
11
+ from torch.utils.model_zoo import load_url as load_state_dict_from_url
12
+
13
+ # Inception weights ported to Pytorch from
14
+ # http://download.tensorflow.org/models/image/imagenet/inception-2015-12-05.tgz
15
+ FID_WEIGHTS_URL = 'https://github.com/mseitzer/pytorch-fid/releases/download/fid_weights/pt_inception-2015-12-05-6726825d.pth'
16
+
17
+ FID_WEIGHTS_PATH = "../../pretrained/inception/pt_inception-2015-12-05-6726825d.pth"
18
+
19
+ def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
20
+ """Numpy implementation of the Frechet Distance.
21
+ The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
22
+ and X_2 ~ N(mu_2, C_2) is
23
+ d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
24
+
25
+ Stable version by Dougal J. Sutherland.
26
+
27
+ Params:
28
+ -- mu1 : Numpy array containing the activations of a layer of the
29
+ inception net (like returned by the function 'get_predictions')
30
+ for generated samples.
31
+ -- mu2 : The sample mean over activations, precalculated on an
32
+ representative data set.
33
+ -- sigma1: The covariance matrix over activations for generated samples.
34
+ -- sigma2: The covariance matrix over activations, precalculated on an
35
+ representative data set.
36
+
37
+ Returns:
38
+ -- : The Frechet Distance.
39
+ """
40
+
41
+ mu1 = np.atleast_1d(mu1)
42
+ mu2 = np.atleast_1d(mu2)
43
+
44
+ sigma1 = np.atleast_2d(sigma1)
45
+ sigma2 = np.atleast_2d(sigma2)
46
+
47
+ assert mu1.shape == mu2.shape, \
48
+ 'Training and test mean vectors have different lengths'
49
+ assert sigma1.shape == sigma2.shape, \
50
+ 'Training and test covariances have different dimensions'
51
+
52
+ diff = mu1 - mu2
53
+
54
+ # Product might be almost singular
55
+ covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
56
+ if not np.isfinite(covmean).all():
57
+ msg = ('fid calculation produces singular product; '
58
+ 'adding %s to diagonal of cov estimates') % eps
59
+ print(msg)
60
+ offset = np.eye(sigma1.shape[0]) * eps
61
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
62
+
63
+ # Numerical error might give slight imaginary component
64
+ if np.iscomplexobj(covmean):
65
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
66
+ m = np.max(np.abs(covmean.imag))
67
+ raise ValueError('Imaginary component {}'.format(m))
68
+ covmean = covmean.real
69
+
70
+ tr_covmean = np.trace(covmean)
71
+
72
+ return (diff.dot(diff) + np.trace(sigma1) +
73
+ np.trace(sigma2) - 2 * tr_covmean)
74
+
75
+ class InceptionV3(nn.Module):
76
+ """Pretrained InceptionV3 network returning feature maps"""
77
+
78
+ # Index of default block of inception to return,
79
+ # corresponds to output of final average pooling
80
+ DEFAULT_BLOCK_INDEX = 3
81
+
82
+ # Maps feature dimensionality to their output blocks indices
83
+ BLOCK_INDEX_BY_DIM = {
84
+ 64: 0, # First max pooling features
85
+ 192: 1, # Second max pooling featurs
86
+ 768: 2, # Pre-aux classifier features
87
+ 2048: 3 # Final average pooling features
88
+ }
89
+
90
+ def __init__(self,
91
+ output_blocks=[DEFAULT_BLOCK_INDEX],
92
+ resize_input=True,
93
+ normalize_input=True,
94
+ requires_grad=False,
95
+ use_fid_inception=True):
96
+ """Build pretrained InceptionV3
97
+
98
+ Parameters
99
+ ----------
100
+ output_blocks : list of int
101
+ Indices of blocks to return features of. Possible values are:
102
+ - 0: corresponds to output of first max pooling
103
+ - 1: corresponds to output of second max pooling
104
+ - 2: corresponds to output which is fed to aux classifier
105
+ - 3: corresponds to output of final average pooling
106
+ resize_input : bool
107
+ If true, bilinearly resizes input to width and height 299 before
108
+ feeding input to model. As the network without fully connected
109
+ layers is fully convolutional, it should be able to handle inputs
110
+ of arbitrary size, so resizing might not be strictly needed
111
+ normalize_input : bool
112
+ If true, scales the input from range (0, 1) to the range the
113
+ pretrained Inception network expects, namely (-1, 1)
114
+ requires_grad : bool
115
+ If true, parameters of the model require gradients. Possibly useful
116
+ for finetuning the network
117
+ use_fid_inception : bool
118
+ If true, uses the pretrained Inception model used in Tensorflow's
119
+ FID implementation. If false, uses the pretrained Inception model
120
+ available in torchvision. The FID Inception model has different
121
+ weights and a slightly different structure from torchvision's
122
+ Inception model. If you want to compute FID scores, you are
123
+ strongly advised to set this parameter to true to get comparable
124
+ results.
125
+ """
126
+ super(InceptionV3, self).__init__()
127
+
128
+ self.resize_input = resize_input
129
+ self.normalize_input = normalize_input
130
+ self.output_blocks = sorted(output_blocks)
131
+ self.last_needed_block = max(output_blocks)
132
+
133
+ assert self.last_needed_block <= 3, \
134
+ 'Last possible output block index is 3'
135
+
136
+ self.blocks = nn.ModuleList()
137
+
138
+ if use_fid_inception:
139
+ inception = fid_inception_v3()
140
+ else:
141
+ inception = models.inception_v3(pretrained=True)
142
+
143
+ # Block 0: input to maxpool1
144
+ block0 = [
145
+ inception.Conv2d_1a_3x3,
146
+ inception.Conv2d_2a_3x3,
147
+ inception.Conv2d_2b_3x3,
148
+ nn.MaxPool2d(kernel_size=3, stride=2)
149
+ ]
150
+ self.blocks.append(nn.Sequential(*block0))
151
+
152
+ # Block 1: maxpool1 to maxpool2
153
+ if self.last_needed_block >= 1:
154
+ block1 = [
155
+ inception.Conv2d_3b_1x1,
156
+ inception.Conv2d_4a_3x3,
157
+ nn.MaxPool2d(kernel_size=3, stride=2)
158
+ ]
159
+ self.blocks.append(nn.Sequential(*block1))
160
+
161
+ # Block 2: maxpool2 to aux classifier
162
+ if self.last_needed_block >= 2:
163
+ block2 = [
164
+ inception.Mixed_5b,
165
+ inception.Mixed_5c,
166
+ inception.Mixed_5d,
167
+ inception.Mixed_6a,
168
+ inception.Mixed_6b,
169
+ inception.Mixed_6c,
170
+ inception.Mixed_6d,
171
+ inception.Mixed_6e,
172
+ ]
173
+ self.blocks.append(nn.Sequential(*block2))
174
+
175
+ # Block 3: aux classifier to final avgpool
176
+ if self.last_needed_block >= 3:
177
+ block3 = [
178
+ inception.Mixed_7a,
179
+ inception.Mixed_7b,
180
+ inception.Mixed_7c,
181
+ nn.AdaptiveAvgPool2d(output_size=(1, 1))
182
+ ]
183
+ self.blocks.append(nn.Sequential(*block3))
184
+
185
+ for param in self.parameters():
186
+ param.requires_grad = requires_grad
187
+
188
+ def forward(self, inp):
189
+ """Get Inception feature maps
190
+
191
+ Parameters
192
+ ----------
193
+ inp : torch.autograd.Variable
194
+ Input tensor of shape Bx3xHxW. Values are expected to be in
195
+ range (0, 1)
196
+
197
+ Returns
198
+ -------
199
+ List of torch.autograd.Variable, corresponding to the selected output
200
+ block, sorted ascending by index
201
+ """
202
+ outp = []
203
+ x = inp
204
+
205
+ if self.resize_input:
206
+ x = F.interpolate(x,
207
+ size=(299, 299),
208
+ mode='bilinear',
209
+ align_corners=False)
210
+
211
+ if self.normalize_input:
212
+ x = 2 * x - 1 # Scale from range (0, 1) to range (-1, 1)
213
+
214
+ for idx, block in enumerate(self.blocks):
215
+ x = block(x)
216
+ if idx in self.output_blocks:
217
+ outp.append(x)
218
+
219
+ if idx == self.last_needed_block:
220
+ break
221
+
222
+ return outp
223
+
224
+
225
+ def fid_inception_v3():
226
+ """Build pretrained Inception model for FID computation
227
+
228
+ The Inception model for FID computation uses a different set of weights
229
+ and has a slightly different structure than torchvision's Inception.
230
+
231
+ This method first constructs torchvision's Inception and then patches the
232
+ necessary parts that are different in the FID Inception model.
233
+ """
234
+ inception = models.inception_v3(num_classes=1008,
235
+ aux_logits=False,
236
+ pretrained=False)
237
+ inception.Mixed_5b = FIDInceptionA(192, pool_features=32)
238
+ inception.Mixed_5c = FIDInceptionA(256, pool_features=64)
239
+ inception.Mixed_5d = FIDInceptionA(288, pool_features=64)
240
+ inception.Mixed_6b = FIDInceptionC(768, channels_7x7=128)
241
+ inception.Mixed_6c = FIDInceptionC(768, channels_7x7=160)
242
+ inception.Mixed_6d = FIDInceptionC(768, channels_7x7=160)
243
+ inception.Mixed_6e = FIDInceptionC(768, channels_7x7=192)
244
+ inception.Mixed_7b = FIDInceptionE_1(1280)
245
+ inception.Mixed_7c = FIDInceptionE_2(2048)
246
+
247
+ state_dict = load_state_dict_from_url(FID_WEIGHTS_URL, progress=True)
248
+ # state_dict = torch.load(FID_WEIGHTS_PATH)
249
+ inception.load_state_dict(state_dict)
250
+ return inception
251
+
252
+
253
+ class FIDInceptionA(models.inception.InceptionA):
254
+ """InceptionA block patched for FID computation"""
255
+ def __init__(self, in_channels, pool_features):
256
+ super(FIDInceptionA, self).__init__(in_channels, pool_features)
257
+
258
+ def forward(self, x):
259
+ branch1x1 = self.branch1x1(x)
260
+
261
+ branch5x5 = self.branch5x5_1(x)
262
+ branch5x5 = self.branch5x5_2(branch5x5)
263
+
264
+ branch3x3dbl = self.branch3x3dbl_1(x)
265
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
266
+ branch3x3dbl = self.branch3x3dbl_3(branch3x3dbl)
267
+
268
+ # Patch: Tensorflow's average pool does not use the padded zero's in
269
+ # its average calculation
270
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
271
+ count_include_pad=False)
272
+ branch_pool = self.branch_pool(branch_pool)
273
+
274
+ outputs = [branch1x1, branch5x5, branch3x3dbl, branch_pool]
275
+ return torch.cat(outputs, 1)
276
+
277
+
278
+ class FIDInceptionC(models.inception.InceptionC):
279
+ """InceptionC block patched for FID computation"""
280
+ def __init__(self, in_channels, channels_7x7):
281
+ super(FIDInceptionC, self).__init__(in_channels, channels_7x7)
282
+
283
+ def forward(self, x):
284
+ branch1x1 = self.branch1x1(x)
285
+
286
+ branch7x7 = self.branch7x7_1(x)
287
+ branch7x7 = self.branch7x7_2(branch7x7)
288
+ branch7x7 = self.branch7x7_3(branch7x7)
289
+
290
+ branch7x7dbl = self.branch7x7dbl_1(x)
291
+ branch7x7dbl = self.branch7x7dbl_2(branch7x7dbl)
292
+ branch7x7dbl = self.branch7x7dbl_3(branch7x7dbl)
293
+ branch7x7dbl = self.branch7x7dbl_4(branch7x7dbl)
294
+ branch7x7dbl = self.branch7x7dbl_5(branch7x7dbl)
295
+
296
+ # Patch: Tensorflow's average pool does not use the padded zero's in
297
+ # its average calculation
298
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
299
+ count_include_pad=False)
300
+ branch_pool = self.branch_pool(branch_pool)
301
+
302
+ outputs = [branch1x1, branch7x7, branch7x7dbl, branch_pool]
303
+ return torch.cat(outputs, 1)
304
+
305
+
306
+ class FIDInceptionE_1(models.inception.InceptionE):
307
+ """First InceptionE block patched for FID computation"""
308
+ def __init__(self, in_channels):
309
+ super(FIDInceptionE_1, self).__init__(in_channels)
310
+
311
+ def forward(self, x):
312
+ branch1x1 = self.branch1x1(x)
313
+
314
+ branch3x3 = self.branch3x3_1(x)
315
+ branch3x3 = [
316
+ self.branch3x3_2a(branch3x3),
317
+ self.branch3x3_2b(branch3x3),
318
+ ]
319
+ branch3x3 = torch.cat(branch3x3, 1)
320
+
321
+ branch3x3dbl = self.branch3x3dbl_1(x)
322
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
323
+ branch3x3dbl = [
324
+ self.branch3x3dbl_3a(branch3x3dbl),
325
+ self.branch3x3dbl_3b(branch3x3dbl),
326
+ ]
327
+ branch3x3dbl = torch.cat(branch3x3dbl, 1)
328
+
329
+ # Patch: Tensorflow's average pool does not use the padded zero's in
330
+ # its average calculation
331
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
332
+ count_include_pad=False)
333
+ branch_pool = self.branch_pool(branch_pool)
334
+
335
+ outputs = [branch1x1, branch3x3, branch3x3dbl, branch_pool]
336
+ return torch.cat(outputs, 1)
337
+
338
+
339
+ class FIDInceptionE_2(models.inception.InceptionE):
340
+ """Second InceptionE block patched for FID computation"""
341
+ def __init__(self, in_channels):
342
+ super(FIDInceptionE_2, self).__init__(in_channels)
343
+
344
+ def forward(self, x):
345
+ branch1x1 = self.branch1x1(x)
346
+
347
+ branch3x3 = self.branch3x3_1(x)
348
+ branch3x3 = [
349
+ self.branch3x3_2a(branch3x3),
350
+ self.branch3x3_2b(branch3x3),
351
+ ]
352
+ branch3x3 = torch.cat(branch3x3, 1)
353
+
354
+ branch3x3dbl = self.branch3x3dbl_1(x)
355
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
356
+ branch3x3dbl = [
357
+ self.branch3x3dbl_3a(branch3x3dbl),
358
+ self.branch3x3dbl_3b(branch3x3dbl),
359
+ ]
360
+ branch3x3dbl = torch.cat(branch3x3dbl, 1)
361
+
362
+ # Patch: The FID Inception model uses max pooling instead of average
363
+ # pooling. This is likely an error in this specific Inception
364
+ # implementation, as other Inception models use average pooling here
365
+ # (which matches the description in the paper).
366
+ branch_pool = F.max_pool2d(x, kernel_size=3, stride=1, padding=1)
367
+ branch_pool = self.branch_pool(branch_pool)
368
+
369
+ outputs = [branch1x1, branch3x3, branch3x3dbl, branch_pool]
370
+ return torch.cat(outputs, 1)