Spaces:
Running on Zero
Running on Zero
hanjian.thu123 commited on
Commit ·
17a8581
0
Parent(s):
[update] app.py
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitignore +30 -0
- LICENSE +21 -0
- README.md +389 -0
- app.py +132 -0
- c2i_train_infer.py +378 -0
- environment.yaml +19 -0
- evaluation/gen_eval/_base_/datasets/coco_panoptic.py +59 -0
- evaluation/gen_eval/_base_/default_runtime.py +27 -0
- evaluation/gen_eval/evaluate_images.py +298 -0
- evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py +253 -0
- evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py +79 -0
- evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py +37 -0
- evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py +61 -0
- evaluation/gen_eval/prompts/create_prompts.py +183 -0
- evaluation/gen_eval/summary_scores.py +45 -0
- grn/__init__.py +0 -0
- grn/dataset/build.py +150 -0
- grn/dataset/dataset_joint_vi.py +687 -0
- grn/models/basic.py +256 -0
- grn/models/ema.py +23 -0
- grn/models/flex_attn_mask.py +67 -0
- grn/models/fused_op.py +27 -0
- grn/models/grn.py +754 -0
- grn/models/grn_c2i.py +399 -0
- grn/models/hbq_tokenizer.py +932 -0
- grn/models/init_param.py +33 -0
- grn/models/rope.py +191 -0
- grn/models/umt5/fsdp.py +42 -0
- grn/models/umt5/t5.py +514 -0
- grn/models/umt5/umt5_tokenizers.py +81 -0
- grn/schedules/__init__.py +6 -0
- grn/schedules/dynamic_resolution.py +99 -0
- grn/schedules/global_refine.py +220 -0
- grn/tokenizer/.gitignore +25 -0
- grn/tokenizer/sample.py +694 -0
- grn/tokenizer/train.py +561 -0
- grn/tokenizer/videovae/__init__.py +0 -0
- grn/tokenizer/videovae/evaluation/__init__.py +8 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore +1 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py +85 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py +72 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py +83 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py +138 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py +90 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py +137 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py +322 -0
- grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py +27 -0
- grn/tokenizer/videovae/evaluation/fid.py +62 -0
- grn/tokenizer/videovae/evaluation/fvd.py +150 -0
- 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 |
+
[](https://arxiv.org/abs/2604.13030)
|
| 4 |
+
[](https://bytedance.github.io/GRN/)
|
| 5 |
+
[](https://huggingface.co/bytedance-research/GRN)
|
| 6 |
+
[](https://huggingface.co/spaces/hanjian/GRN)
|
| 7 |
+
[](LICENSE)
|
| 8 |
+
[](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 |
+
[](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)
|