diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000000000000000000000000000000000000..1fde3227b1fc185f27e91edc4255417d5a9bea91
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,30 @@
+checkpoints
+__pycache__
+imagenet_tiny
+*.log
+*.ckpt
+data
+wandb
+.vscode
+heun-steps*
+log.txt
+imagenet_openai_ref_images
+*.jpg
+*.npz
+src
+scan
+*.txt
+*.env
+*.pem,
+secrets.yaml
+scripts
+scripts/*/local*.sh
+checkpoints_vision
+tmp_videos
+weights
+.DS_Store
+local
+*.mp4
+tmp
+.git_bk
+demo
diff --git a/LICENSE b/LICENSE
new file mode 100644
index 0000000000000000000000000000000000000000..173947b9c50c290492730165d9c71cedb33fa632
--- /dev/null
+++ b/LICENSE
@@ -0,0 +1,21 @@
+MIT License
+
+Copyright (c) 2026 MGenAI
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+SOFTWARE.
diff --git a/README.md b/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..9f87a2aac44e084327324eb3f8bd0cc57316526d
--- /dev/null
+++ b/README.md
@@ -0,0 +1,389 @@
+# GRN: Generative Refinement Networks
+
+[](https://arxiv.org/abs/2604.13030)
+[](https://bytedance.github.io/GRN/)
+[](https://huggingface.co/bytedance-research/GRN)
+[](https://huggingface.co/spaces/hanjian/GRN)
+[](LICENSE)
+[](https://github.com/bytedance/GRN)
+
+---
+
+## 🔥 Updates!!
+* June 3, 2026: 🍉 A toy image-video dataset is provided for GRN-T2I/GRN-T2V training and fine-tuning.
+* May 23, 2026: 🌺 We release the training and evaluation code for HBQ tokenizer, enjoy~
+* April 14, 2026: 🤗 Paper and code release
+
+## 📋 Table of Contents
+
+- [🌟 Introduction](#-introduction)
+- [✨ Gallery](#-gallery)
+- [🚀 Demo](#-demo)
+- [📦 Model Zoo](#-model-zoo)
+- [🛠️ Installation](#️-installation)
+- [🖼️ Class-to-Image](#️-class-to-image)
+ - [Dataset](#dataset)
+ - [Training](#training)
+ - [Evaluation](#evaluation)
+- [🎨 Text-to-Image](#-text-to-image)
+ - [Data](#data)
+ - [Train](#train)
+ - [Inference](#inference)
+- [🎬 Text-to-Video](#-text-to-video)
+ - [Data](#data-1)
+ - [Train](#train-1)
+ - [Inference](#inference-1)
+- [📦 HBQ Tokenizer](#-hbq-tokenizer)
+ - [Data](#data-2)
+ - [Training](#training-1)
+ - [Evaluation](#evaluation-1)
+- [📧 Contact](#-contact)
+- [🤗 Acknowledgements](#-acknowledgements)
+- [📝 Citation](#-citation)
+
+---
+
+## 🌟 Introduction
+
+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.
+
+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.
+
+We introduce **Generative Refinement Networks (GRN)**, a new visual synthesis paradigm that addresses these issues:
+- **Near-lossless tokenization** via Hierarchical Binary Quantization (HBQ)
+- **Global refinement mechanism** that progressively perfects outputs like a human artist
+- **Entropy-guided sampling** for complexity-aware, adaptive-step generation
+
+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.
+
+---
+
+
+ Generative Refinement Framework
+
+
+
+
+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 (pink ), kept two tokens (blue ), erased two tokens (yellow ), and left six tokens blank (gray ).
+
+
+---
+
+## ✨ Gallery
+
+### GRN-8B Text-to-Video Examples
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+---
+
+### GRN-8B Image-to-Video Examples
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+### GRN-2B Class-to-Image Examples
+
+
+
+
+
+### GRN-2B Text-to-Image Examples
+
+
+
+
+
+---
+
+## 🚀 Demo
+
+### 🖼️ Text-to-Image
+Try our interactive Text-to-Image demo on 🤗 Hugging Face Space:
+
+**[GRN T2I Demo](https://huggingface.co/spaces/hanjian/GRN)**
+
+Experience the power of Generative Refinement Networks firsthand by generating images from text prompts directly in your browser!
+
+---
+
+### 🎬 Text-to-Video
+Try our interactive Text-to-Video demo on Discord:
+
+[](http://opensource.bytedance.com/discord/invite)
+
+
+
+ T2V Demo on Discord
+
+
+
+---
+
+## 📦 Model Zoo
+
+| Model | Checkpoints |
+|-------|:-----------:|
+| **Tokenizers** | ✅ [ImageNet Tokenizer](https://huggingface.co/bytedance-research/GRN/blob/main/HBQ_image_tokenizer_16dim_M4.ckpt) ✅ [Joint Image/Video Tokenizer](https://huggingface.co/bytedance-research/GRN/blob/main/HBQ_tokenizer_64dim_M4.ckpt) |
+| **GRN_ind_C2I** | ✅ [B](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_ind_B_ep599.pth) ⬜ L (TBD) ⬜ H (TBD) ⬜ G (TBD) |
+| **GRN_bit_T2I** | ✅ [GRN_T2I](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_T2I_2B.pth) |
+| **GRN_bit_T2V** | ✅ [GRN_T2V](https://huggingface.co/bytedance-research/GRN/blob/main/GRN_T2V_2B.pth) |
+
+---
+
+## 🛠️ Installation
+
+### Step 1: Clone the repository
+```bash
+git clone https://github.com/bytedance/GRN
+cd GRN
+```
+
+### Step 2: Create conda environment
+A suitable [conda](https://conda.io/) environment named `GRN` can be created and activated with:
+```bash
+conda env create -f environment.yaml
+conda activate GRN
+```
+
+### Troubleshooting
+If you get `undefined symbol: iJIT_NotifyEvent` when importing `torch`, simply:
+```bash
+pip uninstall torch
+pip install torch==2.5.1 --index-url https://download.pytorch.org/whl/cu124
+```
+Check this [issue](https://github.com/conda/conda/issues/13812#issuecomment-2071445372) for more details.
+
+---
+
+## 🖼️ Class-to-Image
+
+### Dataset
+Download [ImageNet](http://image-net.org/download) dataset, and place it in your `IMAGENET_PATH`.
+
+### Training
+
+All training scripts are located in `scripts/c2i/`. We suggest using 8x80GB GPUs for most models.
+
+| Model | Training Script | GPUs Required |
+|-------|:-------------:|:-------------:|
+| GRN_ind_B | `bash scripts/c2i/train_GRN_ind_B.sh` | 8x80GB |
+| GRN_bit_B | `bash scripts/c2i/train_GRN_bit_B.sh` | 8x80GB |
+| GRN_ind_L | `bash scripts/c2i/train_GRN_ind_L.sh` | 8x80GB |
+| GRN_ind_H | `bash scripts/c2i/train_GRN_ind_H.sh` | 16x80GB |
+| GRN_ind_G | `bash scripts/c2i/train_GRN_ind_G.sh` | 32x80GB |
+
+### Evaluation
+
+PyTorch pre-trained models are available [here](https://huggingface.co/bytedance-research/GRN/tree/main).
+
+All evaluation scripts are located in `scripts/c2i/`. We suggest using 8x80GB vRAM GPUs.
+
+| Model | Evaluation Script |
+|-------|:--------------:|
+| GRN_ind_B | `bash scripts/c2i/eval_GRN_ind_B.sh` |
+| GRN_bit_B | `bash scripts/c2i/eval_GRN_bit_B.sh` |
+| GRN_ind_L | `bash scripts/c2i/eval_GRN_ind_L.sh` |
+| GRN_ind_H | `bash scripts/c2i/eval_GRN_ind_H.sh` |
+| GRN_ind_G | `bash scripts/c2i/eval_GRN_ind_G.sh` |
+
+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`.
+
+---
+
+## 🎨 Text-to-Image
+### Data
+Refer to `data/toy_data/jsonls/000001/0001_0008_000000100.jsonl`
+```
+{"image_path": "[image_path_1]", "long_caption": "xxx", "long_caption_type": "caption-InternVL2.0", "text": "", "short_caption_type": "blip2_caption", "width": 1080, "height": 1920}
+{"image_path": "[image_path_2]", "long_caption": "xxx", "long_caption_type": "caption-InternVL2.0", "text": "", "short_caption_type": "blip2_caption", "width": 1080, "height": 1920}
+...
+```
+
+### Train
+Run `bash scripts/train_GRN_ind_t2i.sh`
+
+### Inference
+
+You can simply run `python3 t2i_infer.py` or use the following code:
+
+```python
+from PIL import Image
+from grn_pipeline import GRNPipeline
+
+# Load pipeline
+pipeline = GRNPipeline.from_pretrained(
+ hf_repo_id='bytedance-research/GRN',
+ task='T2I',
+ pn='1M',
+ device='cpu',
+).to('cuda')
+
+# Generate one image
+result = pipeline(
+ prompt="A cute cat playing in the garden",
+ guidance_scale=3.0,
+ temperature=1.1,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ snr_shift = 1.,
+ h_div_w=1.,
+ content_type='image',
+ seed=42,
+)
+image = result.images[0]
+image.save('./generated_image.jpg')
+```
+
+---
+
+## 🎬 Text-to-Video
+### Data
+Refer to `data/toy_data/jsonls/000001/0001_0008_000000100.jsonl`
+```
+{"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]"}]}
+{"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]"}]}
+...
+```
+
+### Train
+Run `bash scripts/train_GRN_ind_t2v.sh`
+
+### Inference
+
+You can simply run `python3 t2v_infer.py` or use the following code:
+
+```python
+from grn_pipeline import GRNPipeline
+
+# Load pipeline
+pipeline = GRNPipeline.from_pretrained(
+ hf_repo_id='bytedance-research/GRN',
+ task='T2V',
+ pn='0.41M',
+ device='cpu'
+).to('cuda')
+
+# Generate one video
+result = pipeline(
+ prompt="Two women demonstrate a makeup product, applying it with a sponge while smiling and engaging with the camera in a bright, clean setting.",
+ guidance_scale=4.0,
+ temperature=1.0,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ snr_shift = 1.,
+ h_div_w=9/16,
+ duration=2.,
+ content_type='video',
+ seed=42,
+)
+video_file = result.videos[0]
+```
+
+---
+
+## 📦 HBQ Tokenizer
+
+### Data
+Image Dataset, e.g., data_root/username/labels/imagenet/train.txt:
+```
+[image_1_full_path]
+[image_2_full_path]
+[image_3_full_path]
+...
+```
+
+Video Dataset, e.g., data_root/username/labels_hanjian/high-quality-video/horizontal_videos.txt
+```
+[video_1_full_path]
+[video_2_full_path]
+[video_3_full_path]
+...
+```
+
+### Training
+For example, set `latent_channels=16/64` and `quant_method=hierarchical_binary_quant_round_4` in `scripts/hbq_tokenizer_train.sh`, then run:
+```bash
+cd grn/tokenizer
+bash scripts/hbq_tokenizer_train.sh
+```
+
+### Evaluation
+For example, set `latent_channels=16/64` and `quant_method=hierarchical_binary_quant_round_4` in `scripts/hbq_tokenizer_train.sh`, then run:
+```bash
+cd grn/tokenizer
+bash scripts/hbq_tokenizer_eval.sh
+```
+
+---
+
+## 📧 Contact
+
+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!
+
+**📧 Email:** [hanjian.thu123@bytedance.com](mailto:hanjian.thu123@bytedance.com)
+
+---
+
+## 🤗 Acknowledgements
+
+- 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!
+
+---
+
+## 📝 Citation
+
+If you find our work useful, please consider citing:
+
+```bibtex
+@misc{han2026grn,
+ title={Generative Refinement Networks for Visual Synthesis},
+ author={Jian Han and Jinlai Liu and Jiahuan Wang and Bingyue Peng and Zehuan Yuan},
+ year={2026},
+ eprint={2604.13030},
+ archivePrefix={arXiv},
+ primaryClass={cs.CV},
+ url={https://arxiv.org/abs/2604.13030},
+}
+```
diff --git a/app.py b/app.py
new file mode 100644
index 0000000000000000000000000000000000000000..be872ef8026c16b721a06d163c17fa400600251f
--- /dev/null
+++ b/app.py
@@ -0,0 +1,132 @@
+import os
+import sys
+import traceback
+import torch
+import gradio as gr
+import spaces
+
+sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
+
+from grn_pipeline import GRNPipeline
+
+# Global pipeline
+pipe = None
+device = "cuda" if torch.cuda.is_available() else "cpu"
+
+def load_pipeline():
+ global pipe
+ print("Loading GRN pipeline...")
+ # 从 Hugging Face Hub 下载权重
+ pipe = GRNPipeline.from_pretrained(
+ hf_repo_id='bytedance-research/GRN',
+ task='T2I',
+ pn='1M',
+ model='GRN2b',
+ use_slow_attn=True,
+ device=device,
+ )
+ print("Pipeline loaded successfully!")
+ return pipe
+
+# @spaces.GPU #[uncomment to use ZeroGPU]
+@spaces.GPU(duration=40)
+def generate(prompt, content_type="image", guidance_scale=3.0, temperature=1.0, seed=42, width=1024, height=1024):
+ global pipe
+ if pipe is None:
+ try:
+ pipe = load_pipeline()
+ except Exception as e:
+ print(f"Error loading pipeline: {e}")
+ traceback.print_exc()
+ return f"Error loading pipeline: {e}\n\n{traceback.format_exc()}"
+
+ try:
+ result = pipe(
+ prompt=""+prompt,
+ guidance_scale=guidance_scale,
+ temperature=temperature,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ snr_shift = 1.,
+ h_div_w=1.,
+ content_type=content_type,
+ seed=seed,
+ width=width,
+ height=height
+ )
+
+ if content_type == "image" and hasattr(result, 'images'):
+ return result.images[0]
+ elif content_type == "video" and hasattr(result, 'videos'):
+ return result.videos[0]
+ return f"Error: Invalid result from pipeline"
+ except Exception as e:
+ print(f"Error generating content: {e}")
+ traceback.print_exc()
+ return f"Error generating content: {e}\n\n{traceback.format_exc()}"
+
+def create_demo():
+ with gr.Blocks(title="GRN: Generative Refinement Networks", theme=gr.themes.Soft()) as demo:
+ gr.Markdown("# GRN: Generative Refinement Networks")
+ gr.Markdown("Text-to-Image generation using GRN")
+
+ with gr.Row():
+ with gr.Column():
+ prompt_input = gr.Textbox(
+ label="Text Prompt",
+ placeholder="Enter your prompt here...",
+ value="A cute cat playing in the garden"
+ )
+
+ content_type = gr.Radio(
+ choices=["image"], # , "video"
+ value="image",
+ label="Content Type"
+ )
+
+ with gr.Accordion("Settings", open=True):
+ guidance_scale = gr.Slider(minimum=0, maximum=10, value=3.0, label="Guidance Scale")
+ temperature = gr.Slider(minimum=0.1, maximum=1.5, value=1.1, label="Temperature")
+ seed = gr.Number(value=42, label="Seed", precision=0)
+ width = gr.Number(value=1024, label="Width", precision=0)
+ height = gr.Number(value=1024, label="Height", precision=0)
+
+ generate_btn = gr.Button("Generate", variant="primary")
+
+ with gr.Column():
+ output = gr.Gallery(label="Output", show_label=True, elem_id="gallery", columns=1, height="auto", preview=True, object_fit="contain")
+
+ def generate_and_display(prompt, content_type, guidance_scale, temperature, seed, width, height):
+ result = generate(prompt, content_type, guidance_scale, temperature, seed, width, height)
+ if result:
+ return [result]
+ return []
+
+ generate_btn.click(
+ fn=generate_and_display,
+ inputs=[prompt_input, content_type, guidance_scale, temperature, seed, width, height],
+ outputs=output
+ )
+
+ gr.Examples(
+ examples=[
+ ["A majestic lion standing on a cliff at sunset", "image", 3.0, 1.0, 42, 1024, 1024],
+ ],
+ inputs=[prompt_input, content_type, guidance_scale, temperature, seed, width, height],
+ cache_examples=False
+ )
+
+ return demo
+
+if __name__ == "__main__":
+ try:
+ load_pipeline()
+ except Exception as e:
+ print(f"Error loading pipeline: {e}")
+ traceback.print_exc()
+
+ demo = create_demo()
+ demo.launch()
\ No newline at end of file
diff --git a/c2i_train_infer.py b/c2i_train_infer.py
new file mode 100644
index 0000000000000000000000000000000000000000..7401a43963e679d49e89f981d19fb188be39d958
--- /dev/null
+++ b/c2i_train_infer.py
@@ -0,0 +1,378 @@
+import argparse
+import datetime
+import numpy as np
+import os
+import time
+import functools
+from pathlib import Path
+
+import torch
+import torch.backends.cudnn as cudnn
+from torch.utils.tensorboard import SummaryWriter
+import torchvision.transforms as transforms
+import torchvision.datasets as datasets
+from torch.distributed.device_mesh import init_device_mesh
+from torch.distributed.fsdp import (
+ FullyShardedDataParallel as FSDP,
+ MixedPrecision,
+ BackwardPrefetch,
+ ShardingStrategy,
+ FullStateDictConfig,
+ StateDictType,
+)
+from torch.distributed.fsdp.wrap import (
+ transformer_auto_wrap_policy,
+ enable_wrap,
+ wrap,
+)
+
+from grn.utils_c2i.crop import center_crop_arr
+import grn.utils_c2i.misc as misc
+
+import copy
+from grn.utils_c2i.engine import train_one_epoch, evaluate
+from grn.utils import wandb_utils as wandb_utils
+
+from grn.utils_c2i.denoiser import Denoiser
+from grn.models.grn_c2i import GRNblock
+
+
+def get_args_parser():
+ parser = argparse.ArgumentParser('GRN', add_help=False)
+
+ # architecture
+ parser.add_argument('--model', default='GRN_B', type=str, metavar='MODEL',
+ help='Name of the model to train')
+ parser.add_argument('--img_size', default=256, type=int, help='Image size')
+ parser.add_argument('--attn_dropout', type=float, default=0.0, help='Attention dropout rate')
+ parser.add_argument('--proj_dropout', type=float, default=0.0, help='Projection dropout rate')
+
+ # training
+ parser.add_argument('--epochs', default=200, type=int)
+ parser.add_argument('--warmup_epochs', type=int, default=5, metavar='N',
+ help='Epochs to warm up LR')
+ parser.add_argument('--batch_size', default=128, type=int,
+ help='Batch size per GPU (effective batch size = batch_size * # GPUs)')
+ parser.add_argument('--lr', type=float, default=None, metavar='LR',
+ help='Learning rate (absolute)')
+ parser.add_argument('--blr', type=float, default=5e-5, metavar='LR',
+ help='Base learning rate: absolute_lr = base_lr * total_batch_size / 256')
+ parser.add_argument('--min_lr', type=float, default=0., metavar='LR',
+ help='Minimum LR for cyclic schedulers that hit 0')
+ parser.add_argument('--lr_schedule', type=str, default='constant',
+ help='Learning rate schedule')
+ parser.add_argument('--weight_decay', type=float, default=0.0,
+ help='Weight decay (default: 0.0)')
+ parser.add_argument('--ema_decay1', type=float, default=0.9999,
+ help='The first ema to track. Use the first ema for sampling by default.')
+ parser.add_argument('--ema_decay2', type=float, default=0.9996,
+ help='The second ema to track')
+ parser.add_argument('--P_mean', default=-0.8, type=float)
+ parser.add_argument('--P_std', default=0.8, type=float)
+ parser.add_argument('--noise_scale', default=1.0, type=float)
+ parser.add_argument('--t_eps', default=5e-2, type=float)
+ parser.add_argument('--label_drop_prob', default=0.1, type=float)
+
+ parser.add_argument('--seed', default=0, type=int)
+ parser.add_argument('--start_epoch', default=0, type=int, metavar='N',
+ help='Starting epoch')
+ parser.add_argument('--num_workers', default=12, type=int)
+ parser.add_argument('--pin_mem', action='store_true',
+ help='Pin CPU memory in DataLoader for faster GPU transfers')
+ parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem')
+ parser.set_defaults(pin_mem=True)
+
+ # sampling
+ parser.add_argument('--sampling_method', default='heun', type=str,
+ help='ODE samping method')
+ parser.add_argument('--num_sampling_steps', default=50, type=int,
+ help='Sampling steps')
+ parser.add_argument('--cfg', default=1.0, type=float,
+ help='Classifier-free guidance factor')
+ parser.add_argument('--interval_min', default=0.0, type=float,
+ help='CFG interval min')
+ parser.add_argument('--interval_max', default=1.0, type=float,
+ help='CFG interval max')
+ parser.add_argument('--num_images', default=50000, type=int,
+ help='Number of images to generate')
+ parser.add_argument('--eval_freq', type=int, default=40,
+ help='Frequency (in epochs) for evaluation')
+ parser.add_argument('--online_eval', type=int, default=0, choices=[0,1],
+ help='Whether to evaluate the model online')
+ parser.add_argument('--evaluate_gen', action='store_true')
+ parser.add_argument('--gen_bsz', type=int, default=256,
+ help='Generation batch size')
+
+ # dataset
+ parser.add_argument('--data_path', default='./data/imagenet', type=str,
+ help='Path to the dataset')
+ parser.add_argument('--class_num', default=1000, type=int)
+
+ # checkpointing
+ parser.add_argument('--output_dir', default='./output_dir',
+ help='Directory to save outputs (empty for no saving)')
+ parser.add_argument('--resume', default='',
+ help='Folder that contains checkpoint to resume from')
+ parser.add_argument('--save_last_freq', type=int, default=5,
+ help='Frequency (in epochs) to save checkpoints')
+ parser.add_argument('--log_freq', default=100, type=int)
+ parser.add_argument('--device', default='cuda',
+ help='Device to use for training/testing')
+
+ # distributed training
+ parser.add_argument('--world_size', default=1, type=int,
+ help='Number of distributed processes')
+ parser.add_argument('--local_rank', default=-1, type=int)
+ parser.add_argument('--dist_on_itp', action='store_true')
+ parser.add_argument('--dist_url', default='env://',
+ help='URL used to set up distributed training')
+ parser.add_argument('--hbq_round', default=4, type=int,)
+ parser.add_argument('--in_channels', default=3, type=int,)
+ parser.add_argument('--method', default='GRN_ind', type=str, choices=['GRN_ind', 'GRN_bit'])
+ parser.add_argument('--vae_path', default='', type=str,)
+ parser.add_argument('--tau', default=1.0, type=float,)
+ parser.add_argument('--wandb', default=1, type=int, choices=[0,1])
+ parser.add_argument('--generation_dir', default='/tmp', type=str)
+ parser.add_argument('--clip_grad_norm', default=1., type=float)
+ parser.add_argument('--use_fsdp_train', default=0, type=int, choices=[0, 1])
+ parser.add_argument('--delete_images', default=1, type=int, choices=[0, 1])
+ parser.add_argument('--use_confidence_sampling', default=0, type=int, choices=[0, 1])
+ parser.add_argument('--inner_shard_degree', default=8, type=int)
+ parser.add_argument('--patch_size', default=1, type=int)
+ parser.add_argument('--convert_type', default='', type=str)
+ parser.add_argument('--mask_group_size', default=-1, type=int)
+ parser.add_argument('--grn_shift_factor', default=1., type=float)
+ parser.add_argument('--use_focal_loss', default=0, type=int, choices=[0, 1])
+ return parser
+
+
+def main(args):
+ misc.init_distributed_mode(args)
+ print('Job directory:', os.path.dirname(os.path.realpath(__file__)))
+ print("Arguments:\n{}".format(args).replace(', ', ',\n'))
+
+ device = torch.device(args.device)
+
+ # Set seeds for reproducibility
+ seed = args.seed + misc.get_rank()
+ torch.manual_seed(seed)
+ np.random.seed(seed)
+
+ cudnn.benchmark = True
+
+ num_tasks = misc.get_world_size()
+ global_rank = misc.get_rank()
+
+ # Set up TensorBoard logging (only on main process)
+ if global_rank == 0 and args.output_dir is not None:
+ os.makedirs(args.output_dir, exist_ok=True)
+ log_writer = SummaryWriter(log_dir=args.output_dir)
+ if args.wandb:
+ entity = os.environ["EXP_NAME"]
+ project = os.environ["PROJECT"]
+ wandb_utils.wandb.init(project=project, name=entity, config={})
+ else:
+ log_writer = None
+
+ # Data augmentation transforms
+ transform_train = transforms.Compose([
+ transforms.Lambda(lambda img: center_crop_arr(img, args.img_size)),
+ transforms.RandomHorizontalFlip(),
+ transforms.PILToTensor()
+ ])
+
+ dataset_train = datasets.ImageFolder(os.path.join(args.data_path, 'train'), transform=transform_train)
+ print(dataset_train)
+
+ sampler_train = torch.utils.data.DistributedSampler(
+ dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True
+ )
+ print("Sampler_train =", sampler_train)
+
+ data_loader_train = torch.utils.data.DataLoader(
+ dataset_train, sampler=sampler_train,
+ batch_size=args.batch_size,
+ num_workers=args.num_workers,
+ pin_memory=args.pin_mem,
+ drop_last=True
+ )
+
+ torch._dynamo.config.cache_size_limit = 128
+ torch._dynamo.config.optimize_ddp = False
+
+ # Create denoiser
+ model = Denoiser(args)
+
+ # ininitalize vae
+ from grn.models.hbq_tokenizer import HBQ_Tokenizer
+ vae = HBQ_Tokenizer(args=args, latent_channels=16, encoder_out_type='feature_tanh')
+ vae.eval()
+ vae = vae.to('cuda')
+ for param in vae.parameters():
+ param.requires_grad = False
+ state_dict = torch.load(args.vae_path, map_location='cuda')
+ if 'ema' in state_dict:
+ print(f'Load ema vae weights')
+ state_dict = state_dict['ema']
+ else:
+ print(f'Load non ema vae weights')
+ state_dict = state_dict['vae']
+ print('Load vae: ', vae.load_state_dict(state_dict, assign=True))
+
+ print("Model =", model)
+ n_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
+ print("Number of trainable parameters: {:.6f}M".format(n_params / 1e6))
+
+ model.to(device)
+
+ eff_batch_size = args.batch_size * misc.get_world_size()
+ if args.lr is None: # only base_lr (blr) is specified
+ args.lr = args.blr * eff_batch_size / 256
+
+ print("Base lr: {:.2e}".format(args.lr * 256 / eff_batch_size))
+ print("Actual lr: {:.2e}".format(args.lr))
+ print("Effective batch size: %d" % eff_batch_size)
+
+ if args.use_fsdp_train:
+ auto_wrap_policy = functools.partial(
+ transformer_auto_wrap_policy,
+ transformer_layer_cls={GRNblock},
+ )
+ if args.inner_shard_degree > 0:
+ sharding_strategy = ShardingStrategy.HYBRID_SHARD
+ world_size = misc.get_world_size()
+ assert world_size % args.inner_shard_degree == 0
+ assert args.inner_shard_degree > 1 and args.inner_shard_degree <= world_size
+ device_mesh = init_device_mesh('cuda', (world_size // args.inner_shard_degree, args.inner_shard_degree))
+ else:
+ sharding_strategy = ShardingStrategy.FULL_SHARD
+ device_mesh = None
+ model = FSDP(
+ model,
+ auto_wrap_policy=auto_wrap_policy,
+ mixed_precision=MixedPrecision(
+ param_dtype=torch.bfloat16,
+ reduce_dtype=torch.bfloat16,
+ buffer_dtype=torch.bfloat16
+ ),
+ device_id=torch.cuda.current_device(),
+ sharding_strategy=sharding_strategy,
+ use_orig_params=True,
+ device_mesh=device_mesh,
+ )
+ model_without_ddp = model
+ else:
+ model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu])
+ model_without_ddp = model.module
+
+ # Set up optimizer with weight decay adjustment for bias and norm layers
+ param_groups = misc.add_weight_decay(model_without_ddp, args.weight_decay)
+ optimizer = torch.optim.AdamW(param_groups, lr=args.lr, betas=(0.9, 0.95))
+ print(optimizer)
+
+ # Resume from checkpoint if provided
+ # checkpoint_path = os.path.join(args.resume, "checkpoint-last.pth") if args.resume else None
+ checkpoint_path = args.resume if args.resume else None
+ if checkpoint_path and os.path.exists(checkpoint_path):
+ checkpoint = torch.load(checkpoint_path, map_location='cpu')
+ model_without_ddp.load_state_dict(checkpoint['model'])
+
+ if args.use_fsdp_train:
+ # For FSDP, load EMA state dict into model temporarily to set ema_params
+ model_without_ddp.load_state_dict(checkpoint['model_ema1'])
+ model_without_ddp.module.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
+
+ model_without_ddp.load_state_dict(checkpoint['model_ema2'])
+ model_without_ddp.module.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
+
+ # Restore model
+ model_without_ddp.load_state_dict(checkpoint['model'])
+ else:
+ ema_state_dict1 = checkpoint['model_ema1']
+ ema_state_dict2 = checkpoint['model_ema2']
+ model_without_ddp.ema_params1 = [ema_state_dict1[name].cuda() for name, _ in model_without_ddp.named_parameters()]
+ model_without_ddp.ema_params2 = [ema_state_dict2[name].cuda() for name, _ in model_without_ddp.named_parameters()]
+
+ print("Resumed checkpoint from", args.resume)
+
+ try:
+ if 'optimizer' in checkpoint and 'epoch' in checkpoint:
+ if args.use_fsdp_train:
+ opt_state = FSDP.optim_state_dict_to_load(
+ model_without_ddp, optimizer, checkpoint['optimizer']
+ )
+ optimizer.load_state_dict(opt_state)
+ else:
+ optimizer.load_state_dict(checkpoint['optimizer'])
+ print("Loaded optimizer & scaler state!")
+ except:
+ print("Failed to load optimizer & scaler state! Just load checkpoint.")
+ args.start_epoch = checkpoint['epoch'] + 1
+ del checkpoint
+ else:
+ if args.use_fsdp_train:
+ model_without_ddp.module.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
+ model_without_ddp.module.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
+ else:
+ model_without_ddp.ema_params1 = [p.detach().clone() for p in model_without_ddp.parameters()]
+ model_without_ddp.ema_params2 = [p.detach().clone() for p in model_without_ddp.parameters()]
+ print("Training from scratch")
+
+ # Evaluate generation
+ if args.evaluate_gen:
+ print("Evaluating checkpoint at {} epoch".format(args.start_epoch))
+ with torch.random.fork_rng():
+ torch.manual_seed(seed)
+ with torch.no_grad():
+ evaluate(model_without_ddp, args, args.start_epoch, batch_size=args.gen_bsz, log_writer=log_writer, vae=vae)
+ return
+
+ # Training loop
+ print(f"Start training for {args.epochs} epochs")
+ start_time = time.time()
+ for epoch in range(args.start_epoch, args.epochs):
+ if args.distributed:
+ data_loader_train.sampler.set_epoch(epoch)
+
+ train_one_epoch(model, model_without_ddp, data_loader_train, optimizer, device, epoch, log_writer=log_writer, args=args, vae=vae)
+
+ # Save checkpoint periodically
+ if epoch % args.save_last_freq == 0 or epoch + 1 == args.epochs:
+ if misc.is_main_process():
+ from grn.utils.safe_rm import safe_remove
+ safe_remove(f'{args.output_dir}/checkpoint-tmp_*.pth', args.output_dir)
+ misc.save_model(
+ args=args,
+ model_without_ddp=model_without_ddp,
+ optimizer=optimizer,
+ epoch=epoch,
+ epoch_name=f"tmp_{epoch}"
+ )
+
+ if epoch % 100 == 0 and epoch > 0:
+ misc.save_model(
+ args=args,
+ model_without_ddp=model_without_ddp,
+ optimizer=optimizer,
+ epoch=epoch
+ )
+
+ # Perform online evaluation at specified intervals
+ if args.online_eval and (epoch % args.eval_freq == 0 or epoch + 1 == args.epochs):
+ torch.cuda.empty_cache()
+ with torch.no_grad():
+ evaluate(model_without_ddp, args, epoch, batch_size=args.gen_bsz, log_writer=log_writer, vae=vae)
+ torch.cuda.empty_cache()
+
+ if misc.is_main_process() and log_writer is not None:
+ log_writer.flush()
+
+ total_time = time.time() - start_time
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
+ print('Training time:', total_time_str)
+
+
+if __name__ == '__main__':
+ args = get_args_parser().parse_args()
+ Path(args.output_dir).mkdir(parents=True, exist_ok=True)
+ main(args)
diff --git a/environment.yaml b/environment.yaml
new file mode 100644
index 0000000000000000000000000000000000000000..4c8710f3ac82c01912a92126b09006564faedb96
--- /dev/null
+++ b/environment.yaml
@@ -0,0 +1,19 @@
+name: grn
+channels:
+ - pytorch
+ - defaults
+ - nvidia
+dependencies:
+ - python=3.10
+ - pip=22.3
+ - pytorch-cuda=12.4
+ - pytorch=2.5.1
+ - torchvision=0.20.1
+ - numpy=1.22
+ - pip:
+ - opencv-python==4.11.0.86
+ - timm==0.9.12
+ - tensorboard==2.10.0
+ - scipy==1.9.1
+ - einops==0.8.1
+ - gdown==5.2.0
diff --git a/evaluation/gen_eval/_base_/datasets/coco_panoptic.py b/evaluation/gen_eval/_base_/datasets/coco_panoptic.py
new file mode 100644
index 0000000000000000000000000000000000000000..dbade7c0ac20141806b93f0ea7b5ca26d748246e
--- /dev/null
+++ b/evaluation/gen_eval/_base_/datasets/coco_panoptic.py
@@ -0,0 +1,59 @@
+# dataset settings
+dataset_type = 'CocoPanopticDataset'
+data_root = 'data/coco/'
+img_norm_cfg = dict(
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
+train_pipeline = [
+ dict(type='LoadImageFromFile'),
+ dict(
+ type='LoadPanopticAnnotations',
+ with_bbox=True,
+ with_mask=True,
+ with_seg=True),
+ dict(type='Resize', img_scale=(1333, 800), keep_ratio=True),
+ dict(type='RandomFlip', flip_ratio=0.5),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='Pad', size_divisor=32),
+ dict(type='SegRescale', scale_factor=1 / 4),
+ dict(type='DefaultFormatBundle'),
+ dict(
+ type='Collect',
+ keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks', 'gt_semantic_seg']),
+]
+test_pipeline = [
+ dict(type='LoadImageFromFile'),
+ dict(
+ type='MultiScaleFlipAug',
+ img_scale=(1333, 800),
+ flip=False,
+ transforms=[
+ dict(type='Resize', keep_ratio=True),
+ dict(type='RandomFlip'),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='Pad', size_divisor=32),
+ dict(type='ImageToTensor', keys=['img']),
+ dict(type='Collect', keys=['img']),
+ ])
+]
+data = dict(
+ samples_per_gpu=2,
+ workers_per_gpu=2,
+ train=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/panoptic_train2017.json',
+ img_prefix=data_root + 'train2017/',
+ seg_prefix=data_root + 'annotations/panoptic_train2017/',
+ pipeline=train_pipeline),
+ val=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/panoptic_val2017.json',
+ img_prefix=data_root + 'val2017/',
+ seg_prefix=data_root + 'annotations/panoptic_val2017/',
+ pipeline=test_pipeline),
+ test=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/panoptic_val2017.json',
+ img_prefix=data_root + 'val2017/',
+ seg_prefix=data_root + 'annotations/panoptic_val2017/',
+ pipeline=test_pipeline))
+evaluation = dict(interval=1, metric=['PQ'])
diff --git a/evaluation/gen_eval/_base_/default_runtime.py b/evaluation/gen_eval/_base_/default_runtime.py
new file mode 100644
index 0000000000000000000000000000000000000000..5b0b1452c0a625e331be7b1e6c5cf341cc91ff64
--- /dev/null
+++ b/evaluation/gen_eval/_base_/default_runtime.py
@@ -0,0 +1,27 @@
+checkpoint_config = dict(interval=1)
+# yapf:disable
+log_config = dict(
+ interval=50,
+ hooks=[
+ dict(type='TextLoggerHook'),
+ # dict(type='TensorboardLoggerHook')
+ ])
+# yapf:enable
+custom_hooks = [dict(type='NumClassCheckHook')]
+
+dist_params = dict(backend='nccl')
+log_level = 'INFO'
+load_from = None
+resume_from = None
+workflow = [('train', 1)]
+
+# disable opencv multithreading to avoid system being overloaded
+opencv_num_threads = 0
+# set multi-process start method as `fork` to speed up the training
+mp_start_method = 'fork'
+
+# Default setting for scaling LR automatically
+# - `enable` means enable scaling LR automatically
+# or not by default.
+# - `base_batch_size` = (8 GPUs) x (2 samples per GPU).
+auto_scale_lr = dict(enable=False, base_batch_size=16)
diff --git a/evaluation/gen_eval/evaluate_images.py b/evaluation/gen_eval/evaluate_images.py
new file mode 100755
index 0000000000000000000000000000000000000000..0a7fee4757c48cdbc5b977aa73fca4ab26b6146d
--- /dev/null
+++ b/evaluation/gen_eval/evaluate_images.py
@@ -0,0 +1,298 @@
+"""
+Evaluate generated images using Mask2Former (or other object detector model)
+"""
+
+import argparse
+import json
+import os
+import re
+import sys
+import time
+import os.path as osp
+
+import warnings
+warnings.filterwarnings("ignore")
+
+import numpy as np
+import tqdm
+import pandas as pd
+from PIL import Image, ImageOps
+import torch
+import mmdet
+from mmdet.apis import inference_detector, init_detector
+
+import open_clip
+from clip_benchmark.metrics import zeroshot_classification as zsc
+zsc.tqdm = lambda it, *args, **kwargs: it
+
+# Get directory path
+
+def parse_args():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("imagedir", type=str)
+ parser.add_argument("--outfile", type=str, default="results.jsonl")
+ parser.add_argument("--model-config", type=str, default="")
+ parser.add_argument("--model-path", type=str, default="")
+ # Other arguments
+ parser.add_argument("--options", nargs="*", type=str, default=[])
+ args = parser.parse_args()
+ args.options = dict(opt.split("=", 1) for opt in args.options)
+ if args.model_config is None:
+ args.model_config = os.path.join(
+ os.path.dirname(mmdet.__file__),
+ "../configs/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py"
+ )
+ return args
+
+DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
+assert DEVICE == "cuda"
+
+def timed(fn):
+ def wrapper(*args, **kwargs):
+ startt = time.time()
+ result = fn(*args, **kwargs)
+ endt = time.time()
+ print(f'Function {fn.__name__!r} executed in {endt - startt:.3f}s', file=sys.stderr)
+ return result
+ return wrapper
+
+# Load models
+
+@timed
+def load_models(args):
+ CONFIG_PATH = osp.abspath(args.model_config)
+ OBJECT_DETECTOR = args.options.get('model', "mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco")
+ CKPT_PATH = os.path.join(args.model_path, f"{OBJECT_DETECTOR}.pth")
+ object_detector = init_detector(CONFIG_PATH, CKPT_PATH, device=DEVICE)
+
+ clip_arch = args.options.get('clip_model', "ViT-L-14")
+ clip_model, _, transform = open_clip.create_model_and_transforms(clip_arch, pretrained="openai", device=DEVICE)
+ tokenizer = open_clip.get_tokenizer(clip_arch)
+
+ with open(os.path.join(os.path.dirname(__file__), "object_names.txt")) as cls_file:
+ classnames = [line.strip() for line in cls_file]
+
+ return object_detector, (clip_model, transform, tokenizer), classnames
+
+
+COLORS = ["red", "orange", "yellow", "green", "blue", "purple", "pink", "brown", "black", "white"]
+COLOR_CLASSIFIERS = {}
+
+# Evaluation parts
+
+class ImageCrops(torch.utils.data.Dataset):
+ def __init__(self, image: Image.Image, objects):
+ self._image = image.convert("RGB")
+ bgcolor = args.options.get('bgcolor', "#999")
+ if bgcolor == "original":
+ self._blank = self._image.copy()
+ else:
+ self._blank = Image.new("RGB", image.size, color=bgcolor)
+ self._objects = objects
+
+ def __len__(self):
+ return len(self._objects)
+
+ def __getitem__(self, index):
+ box, mask = self._objects[index]
+ if mask is not None:
+ assert tuple(self._image.size[::-1]) == tuple(mask.shape), (index, self._image.size[::-1], mask.shape)
+ image = Image.composite(self._image, self._blank, Image.fromarray(mask))
+ else:
+ image = self._image
+ if args.options.get('crop', '1') == '1':
+ image = image.crop(box[:4])
+ # if args.save:
+ # base_count = len(os.listdir(args.save))
+ # image.save(os.path.join(args.save, f"cropped_{base_count:05}.png"))
+ return (transform(image), 0)
+
+
+def color_classification(image, bboxes, classname):
+ if classname not in COLOR_CLASSIFIERS:
+ COLOR_CLASSIFIERS[classname] = zsc.zero_shot_classifier(
+ clip_model, tokenizer, COLORS,
+ [
+ f"a photo of a {{c}} {classname}",
+ f"a photo of a {{c}}-colored {classname}",
+ f"a photo of a {{c}} object"
+ ],
+ DEVICE
+ )
+ clf = COLOR_CLASSIFIERS[classname]
+ dataloader = torch.utils.data.DataLoader(
+ ImageCrops(image, bboxes),
+ batch_size=16, num_workers=4
+ )
+ with torch.no_grad():
+ pred, _ = zsc.run_classification(clip_model, clf, dataloader, DEVICE)
+ return [COLORS[index.item()] for index in pred.argmax(1)]
+
+
+def compute_iou(box_a, box_b):
+ area_fn = lambda box: max(box[2] - box[0] + 1, 0) * max(box[3] - box[1] + 1, 0)
+ i_area = area_fn([
+ max(box_a[0], box_b[0]), max(box_a[1], box_b[1]),
+ min(box_a[2], box_b[2]), min(box_a[3], box_b[3])
+ ])
+ u_area = area_fn(box_a) + area_fn(box_b) - i_area
+ return i_area / u_area if u_area else 0
+
+
+def relative_position(obj_a, obj_b):
+ """Give position of A relative to B, factoring in object dimensions"""
+ boxes = np.array([obj_a[0], obj_b[0]])[:, :4].reshape(2, 2, 2)
+ center_a, center_b = boxes.mean(axis=-2)
+ dim_a, dim_b = np.abs(np.diff(boxes, axis=-2))[..., 0, :]
+ offset = center_a - center_b
+ #
+ revised_offset = np.maximum(np.abs(offset) - POSITION_THRESHOLD * (dim_a + dim_b), 0) * np.sign(offset)
+ if np.all(np.abs(revised_offset) < 1e-3):
+ return set()
+ #
+ dx, dy = revised_offset / np.linalg.norm(offset)
+ relations = set()
+ if dx < -0.5: relations.add("left of")
+ if dx > 0.5: relations.add("right of")
+ if dy < -0.5: relations.add("above")
+ if dy > 0.5: relations.add("below")
+ return relations
+
+
+def evaluate(image, objects, metadata):
+ """
+ Evaluate given image using detected objects on the global metadata specifications.
+ Assumptions:
+ * Metadata combines 'include' clauses with AND, and 'exclude' clauses with OR
+ * All clauses are independent, i.e., duplicating a clause has no effect on the correctness
+ * CHANGED: Color and position will only be evaluated on the most confidently predicted objects;
+ therefore, objects are expected to appear in sorted order
+ """
+ correct = True
+ reason = []
+ matched_groups = []
+ # Check for expected objects
+ for req in metadata.get('include', []):
+ classname = req['class']
+ matched = True
+ found_objects = objects.get(classname, [])[:req['count']]
+ if len(found_objects) < req['count']:
+ correct = matched = False
+ reason.append(f"expected {classname}>={req['count']}, found {len(found_objects)}")
+ else:
+ if 'color' in req:
+ # Color check
+ colors = color_classification(image, found_objects, classname)
+ if colors.count(req['color']) < req['count']:
+ correct = matched = False
+ reason.append(
+ f"expected {req['color']} {classname}>={req['count']}, found " +
+ f"{colors.count(req['color'])} {req['color']}; and " +
+ ", ".join(f"{colors.count(c)} {c}" for c in COLORS if c in colors)
+ )
+ if 'position' in req and matched:
+ # Relative position check
+ expected_rel, target_group = req['position']
+ if matched_groups[target_group] is None:
+ correct = matched = False
+ reason.append(f"no target for {classname} to be {expected_rel}")
+ else:
+ for obj in found_objects:
+ for target_obj in matched_groups[target_group]:
+ true_rels = relative_position(obj, target_obj)
+ if expected_rel not in true_rels:
+ correct = matched = False
+ reason.append(
+ f"expected {classname} {expected_rel} target, found " +
+ f"{' and '.join(true_rels)} target"
+ )
+ break
+ if not matched:
+ break
+ if matched:
+ matched_groups.append(found_objects)
+ else:
+ matched_groups.append(None)
+ # Check for non-expected objects
+ for req in metadata.get('exclude', []):
+ classname = req['class']
+ if len(objects.get(classname, [])) >= req['count']:
+ correct = False
+ reason.append(f"expected {classname}<{req['count']}, found {len(objects[classname])}")
+ return correct, "\n".join(reason)
+
+
+def evaluate_image(filepath, metadata):
+ result = inference_detector(object_detector, filepath)
+ bbox = result[0] if isinstance(result, tuple) else result
+ segm = result[1] if isinstance(result, tuple) and len(result) > 1 else None
+ image = ImageOps.exif_transpose(Image.open(filepath))
+ detected = {}
+ # Determine bounding boxes to keep
+ confidence_threshold = THRESHOLD if metadata['tag'] != "counting" else COUNTING_THRESHOLD
+ for index, classname in enumerate(classnames):
+ ordering = np.argsort(bbox[index][:, 4])[::-1]
+ ordering = ordering[bbox[index][ordering, 4] > confidence_threshold] # Threshold
+ ordering = ordering[:MAX_OBJECTS].tolist() # Limit number of detected objects per class
+ detected[classname] = []
+ while ordering:
+ max_obj = ordering.pop(0)
+ detected[classname].append((bbox[index][max_obj], None if segm is None else segm[index][max_obj]))
+ ordering = [
+ obj for obj in ordering
+ if NMS_THRESHOLD == 1 or compute_iou(bbox[index][max_obj], bbox[index][obj]) < NMS_THRESHOLD
+ ]
+ if not detected[classname]:
+ del detected[classname]
+ # Evaluate
+ is_correct, reason = evaluate(image, detected, metadata)
+ return {
+ 'filename': filepath,
+ 'tag': metadata['tag'],
+ 'prompt': metadata['prompt'],
+ 'correct': is_correct,
+ 'reason': reason,
+ 'metadata': json.dumps(metadata),
+ 'details': json.dumps({
+ key: [box.tolist() for box, _ in value]
+ for key, value in detected.items()
+ })
+ }
+
+
+def main(args):
+ full_results = []
+ pbar = tqdm.tqdm(total=len(os.listdir(args.imagedir)))
+ for subfolder in os.listdir(args.imagedir):
+ pbar.update(1)
+ folderpath = os.path.join(args.imagedir, subfolder)
+ if not os.path.isdir(folderpath) or not subfolder.isdigit():
+ print('skip 269')
+ continue
+ with open(os.path.join(folderpath, "metadata.jsonl")) as fp:
+ metadata = json.load(fp)
+ # Evaluate each image
+ for imagename in os.listdir(os.path.join(folderpath, "samples")):
+ imagepath = os.path.join(folderpath, "samples", imagename)
+ if not os.path.isfile(imagepath) or not (re.match(r"\d+\.png", imagename) or re.match(r"\d+\.jpg", imagename)):
+ print('skip 276')
+ continue
+ result = evaluate_image(imagepath, metadata)
+ full_results.append(result)
+ # Save results
+ if os.path.dirname(args.outfile):
+ os.makedirs(os.path.dirname(args.outfile), exist_ok=True)
+ with open(args.outfile, "w") as fp:
+ pd.DataFrame(full_results).to_json(fp, orient="records", lines=True)
+
+
+if __name__ == "__main__":
+ args = parse_args()
+ object_detector, (clip_model, transform, tokenizer), classnames = load_models(args)
+ THRESHOLD = float(args.options.get('threshold', 0.3))
+ COUNTING_THRESHOLD = float(args.options.get('counting_threshold', 0.9))
+ MAX_OBJECTS = int(args.options.get('max_objects', 16))
+ NMS_THRESHOLD = float(args.options.get('max_overlap', 1.0))
+ POSITION_THRESHOLD = float(args.options.get('position_threshold', 0.1))
+
+ main(args)
diff --git a/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py b/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py
new file mode 100644
index 0000000000000000000000000000000000000000..2c23625e139391f6341bb7a4826b13803c80c6b2
--- /dev/null
+++ b/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco-panoptic.py
@@ -0,0 +1,253 @@
+_base_ = [
+ '../_base_/datasets/coco_panoptic.py', '../_base_/default_runtime.py'
+]
+num_things_classes = 80
+num_stuff_classes = 53
+num_classes = num_things_classes + num_stuff_classes
+model = dict(
+ type='Mask2Former',
+ backbone=dict(
+ type='ResNet',
+ depth=50,
+ num_stages=4,
+ out_indices=(0, 1, 2, 3),
+ frozen_stages=-1,
+ norm_cfg=dict(type='BN', requires_grad=False),
+ norm_eval=True,
+ style='pytorch',
+ init_cfg=dict(type='Pretrained', checkpoint='torchvision://resnet50')),
+ panoptic_head=dict(
+ type='Mask2FormerHead',
+ in_channels=[256, 512, 1024, 2048], # pass to pixel_decoder inside
+ strides=[4, 8, 16, 32],
+ feat_channels=256,
+ out_channels=256,
+ num_things_classes=num_things_classes,
+ num_stuff_classes=num_stuff_classes,
+ num_queries=100,
+ num_transformer_feat_level=3,
+ pixel_decoder=dict(
+ type='MSDeformAttnPixelDecoder',
+ num_outs=3,
+ norm_cfg=dict(type='GN', num_groups=32),
+ act_cfg=dict(type='ReLU'),
+ encoder=dict(
+ type='DetrTransformerEncoder',
+ num_layers=6,
+ transformerlayers=dict(
+ type='BaseTransformerLayer',
+ attn_cfgs=dict(
+ type='MultiScaleDeformableAttention',
+ embed_dims=256,
+ num_heads=8,
+ num_levels=3,
+ num_points=4,
+ im2col_step=64,
+ dropout=0.0,
+ batch_first=False,
+ norm_cfg=None,
+ init_cfg=None),
+ ffn_cfgs=dict(
+ type='FFN',
+ embed_dims=256,
+ feedforward_channels=1024,
+ num_fcs=2,
+ ffn_drop=0.0,
+ act_cfg=dict(type='ReLU', inplace=True)),
+ operation_order=('self_attn', 'norm', 'ffn', 'norm')),
+ init_cfg=None),
+ positional_encoding=dict(
+ type='SinePositionalEncoding', num_feats=128, normalize=True),
+ init_cfg=None),
+ enforce_decoder_input_project=False,
+ positional_encoding=dict(
+ type='SinePositionalEncoding', num_feats=128, normalize=True),
+ transformer_decoder=dict(
+ type='DetrTransformerDecoder',
+ return_intermediate=True,
+ num_layers=9,
+ transformerlayers=dict(
+ type='DetrTransformerDecoderLayer',
+ attn_cfgs=dict(
+ type='MultiheadAttention',
+ embed_dims=256,
+ num_heads=8,
+ attn_drop=0.0,
+ proj_drop=0.0,
+ dropout_layer=None,
+ batch_first=False),
+ ffn_cfgs=dict(
+ embed_dims=256,
+ feedforward_channels=2048,
+ num_fcs=2,
+ act_cfg=dict(type='ReLU', inplace=True),
+ ffn_drop=0.0,
+ dropout_layer=None,
+ add_identity=True),
+ feedforward_channels=2048,
+ operation_order=('cross_attn', 'norm', 'self_attn', 'norm',
+ 'ffn', 'norm')),
+ init_cfg=None),
+ loss_cls=dict(
+ type='CrossEntropyLoss',
+ use_sigmoid=False,
+ loss_weight=2.0,
+ reduction='mean',
+ class_weight=[1.0] * num_classes + [0.1]),
+ loss_mask=dict(
+ type='CrossEntropyLoss',
+ use_sigmoid=True,
+ reduction='mean',
+ loss_weight=5.0),
+ loss_dice=dict(
+ type='DiceLoss',
+ use_sigmoid=True,
+ activate=True,
+ reduction='mean',
+ naive_dice=True,
+ eps=1.0,
+ loss_weight=5.0)),
+ panoptic_fusion_head=dict(
+ type='MaskFormerFusionHead',
+ num_things_classes=num_things_classes,
+ num_stuff_classes=num_stuff_classes,
+ loss_panoptic=None,
+ init_cfg=None),
+ train_cfg=dict(
+ num_points=12544,
+ oversample_ratio=3.0,
+ importance_sample_ratio=0.75,
+ assigner=dict(
+ type='MaskHungarianAssigner',
+ cls_cost=dict(type='ClassificationCost', weight=2.0),
+ mask_cost=dict(
+ type='CrossEntropyLossCost', weight=5.0, use_sigmoid=True),
+ dice_cost=dict(
+ type='DiceCost', weight=5.0, pred_act=True, eps=1.0)),
+ sampler=dict(type='MaskPseudoSampler')),
+ test_cfg=dict(
+ panoptic_on=True,
+ # For now, the dataset does not support
+ # evaluating semantic segmentation metric.
+ semantic_on=False,
+ instance_on=True,
+ # max_per_image is for instance segmentation.
+ max_per_image=100,
+ iou_thr=0.8,
+ # In Mask2Former's panoptic postprocessing,
+ # it will filter mask area where score is less than 0.5 .
+ filter_low_score=True),
+ init_cfg=None)
+
+# dataset settings
+image_size = (1024, 1024)
+img_norm_cfg = dict(
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
+train_pipeline = [
+ dict(type='LoadImageFromFile', to_float32=True),
+ dict(
+ type='LoadPanopticAnnotations',
+ with_bbox=True,
+ with_mask=True,
+ with_seg=True),
+ dict(type='RandomFlip', flip_ratio=0.5),
+ # large scale jittering
+ dict(
+ type='Resize',
+ img_scale=image_size,
+ ratio_range=(0.1, 2.0),
+ multiscale_mode='range',
+ keep_ratio=True),
+ dict(
+ type='RandomCrop',
+ crop_size=image_size,
+ crop_type='absolute',
+ recompute_bbox=True,
+ allow_negative_crop=True),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='Pad', size=image_size),
+ dict(type='DefaultFormatBundle', img_to_float=True),
+ dict(
+ type='Collect',
+ keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks', 'gt_semantic_seg']),
+]
+test_pipeline = [
+ dict(type='LoadImageFromFile'),
+ dict(
+ type='MultiScaleFlipAug',
+ img_scale=(1333, 800),
+ flip=False,
+ transforms=[
+ dict(type='Resize', keep_ratio=True),
+ dict(type='RandomFlip'),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='Pad', size_divisor=32),
+ dict(type='ImageToTensor', keys=['img']),
+ dict(type='Collect', keys=['img']),
+ ])
+]
+data_root = 'data/coco/'
+data = dict(
+ samples_per_gpu=2,
+ workers_per_gpu=2,
+ train=dict(pipeline=train_pipeline),
+ val=dict(
+ pipeline=test_pipeline,
+ ins_ann_file=data_root + 'annotations/instances_val2017.json',
+ ),
+ test=dict(
+ pipeline=test_pipeline,
+ ins_ann_file=data_root + 'annotations/instances_val2017.json',
+ ))
+
+embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
+# optimizer
+optimizer = dict(
+ type='AdamW',
+ lr=0.0001,
+ weight_decay=0.05,
+ eps=1e-8,
+ betas=(0.9, 0.999),
+ paramwise_cfg=dict(
+ custom_keys={
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
+ 'query_embed': embed_multi,
+ 'query_feat': embed_multi,
+ 'level_embed': embed_multi,
+ },
+ norm_decay_mult=0.0))
+optimizer_config = dict(grad_clip=dict(max_norm=0.01, norm_type=2))
+
+# learning policy
+lr_config = dict(
+ policy='step',
+ gamma=0.1,
+ by_epoch=False,
+ step=[327778, 355092],
+ warmup='linear',
+ warmup_by_epoch=False,
+ warmup_ratio=1.0, # no warmup
+ warmup_iters=10)
+
+max_iters = 368750
+runner = dict(type='IterBasedRunner', max_iters=max_iters)
+
+log_config = dict(
+ interval=50,
+ hooks=[
+ dict(type='TextLoggerHook', by_epoch=False),
+ dict(type='TensorboardLoggerHook', by_epoch=False)
+ ])
+interval = 5000
+workflow = [('train', interval)]
+checkpoint_config = dict(
+ by_epoch=False, interval=interval, save_last=True, max_keep_ckpts=3)
+
+# Before 365001th iteration, we do evaluation every 5000 iterations.
+# After 365000th iteration, we do evaluation every 368750 iterations,
+# which means that we do evaluation at the end of training.
+dynamic_intervals = [(max_iters // interval * interval + 1, max_iters)]
+evaluation = dict(
+ interval=interval,
+ dynamic_intervals=dynamic_intervals,
+ metric=['PQ', 'bbox', 'segm'])
diff --git a/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py b/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py
new file mode 100644
index 0000000000000000000000000000000000000000..eca6135ba7cb1c28eeb5bf5f031176dc77952630
--- /dev/null
+++ b/evaluation/gen_eval/mask2former/mask2former_r50_lsj_8x2_50e_coco.py
@@ -0,0 +1,79 @@
+_base_ = ['./mask2former_r50_lsj_8x2_50e_coco-panoptic.py']
+num_things_classes = 80
+num_stuff_classes = 0
+num_classes = num_things_classes + num_stuff_classes
+model = dict(
+ panoptic_head=dict(
+ num_things_classes=num_things_classes,
+ num_stuff_classes=num_stuff_classes,
+ loss_cls=dict(class_weight=[1.0] * num_classes + [0.1])),
+ panoptic_fusion_head=dict(
+ num_things_classes=num_things_classes,
+ num_stuff_classes=num_stuff_classes),
+ test_cfg=dict(panoptic_on=False))
+
+# dataset settings
+image_size = (1024, 1024)
+img_norm_cfg = dict(
+ mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True)
+pad_cfg = dict(img=(128, 128, 128), masks=0, seg=255)
+train_pipeline = [
+ dict(type='LoadImageFromFile', to_float32=True),
+ dict(type='LoadAnnotations', with_bbox=True, with_mask=True),
+ dict(type='RandomFlip', flip_ratio=0.5),
+ # large scale jittering
+ dict(
+ type='Resize',
+ img_scale=image_size,
+ ratio_range=(0.1, 2.0),
+ multiscale_mode='range',
+ keep_ratio=True),
+ dict(
+ type='RandomCrop',
+ crop_size=image_size,
+ crop_type='absolute',
+ recompute_bbox=True,
+ allow_negative_crop=True),
+ dict(type='FilterAnnotations', min_gt_bbox_wh=(1e-5, 1e-5), by_mask=True),
+ dict(type='Pad', size=image_size, pad_val=pad_cfg),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='DefaultFormatBundle', img_to_float=True),
+ dict(type='Collect', keys=['img', 'gt_bboxes', 'gt_labels', 'gt_masks']),
+]
+test_pipeline = [
+ dict(type='LoadImageFromFile'),
+ dict(
+ type='MultiScaleFlipAug',
+ img_scale=(1333, 800),
+ flip=False,
+ transforms=[
+ dict(type='Resize', keep_ratio=True),
+ dict(type='RandomFlip'),
+ dict(type='Pad', size_divisor=32, pad_val=pad_cfg),
+ dict(type='Normalize', **img_norm_cfg),
+ dict(type='ImageToTensor', keys=['img']),
+ dict(type='Collect', keys=['img']),
+ ])
+]
+dataset_type = 'CocoDataset'
+data_root = 'data/coco/'
+data = dict(
+ _delete_=True,
+ samples_per_gpu=2,
+ workers_per_gpu=2,
+ train=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/instances_train2017.json',
+ img_prefix=data_root + 'train2017/',
+ pipeline=train_pipeline),
+ val=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/instances_val2017.json',
+ img_prefix=data_root + 'val2017/',
+ pipeline=test_pipeline),
+ test=dict(
+ type=dataset_type,
+ ann_file=data_root + 'annotations/instances_val2017.json',
+ img_prefix=data_root + 'val2017/',
+ pipeline=test_pipeline))
+evaluation = dict(metric=['bbox', 'segm'])
diff --git a/evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py b/evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py
new file mode 100644
index 0000000000000000000000000000000000000000..7b1b05abafe6133eb79b1537dad08d9d9f205deb
--- /dev/null
+++ b/evaluation/gen_eval/mask2former/mask2former_swin-s-p4-w7-224_lsj_8x2_50e_coco.py
@@ -0,0 +1,37 @@
+_base_ = ['./mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py']
+pretrained = 'https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_small_patch4_window7_224.pth' # noqa
+
+depths = [2, 2, 18, 2]
+model = dict(
+ backbone=dict(
+ depths=depths, init_cfg=dict(type='Pretrained',
+ checkpoint=pretrained)))
+
+# set all layers in backbone to lr_mult=0.1
+# set all norm layers, position_embeding,
+# query_embeding, level_embeding to decay_multi=0.0
+backbone_norm_multi = dict(lr_mult=0.1, decay_mult=0.0)
+backbone_embed_multi = dict(lr_mult=0.1, decay_mult=0.0)
+embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
+custom_keys = {
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
+ 'backbone.patch_embed.norm': backbone_norm_multi,
+ 'backbone.norm': backbone_norm_multi,
+ 'absolute_pos_embed': backbone_embed_multi,
+ 'relative_position_bias_table': backbone_embed_multi,
+ 'query_embed': embed_multi,
+ 'query_feat': embed_multi,
+ 'level_embed': embed_multi
+}
+custom_keys.update({
+ f'backbone.stages.{stage_id}.blocks.{block_id}.norm': backbone_norm_multi
+ for stage_id, num_blocks in enumerate(depths)
+ for block_id in range(num_blocks)
+})
+custom_keys.update({
+ f'backbone.stages.{stage_id}.downsample.norm': backbone_norm_multi
+ for stage_id in range(len(depths) - 1)
+})
+# optimizer
+optimizer = dict(
+ paramwise_cfg=dict(custom_keys=custom_keys, norm_decay_mult=0.0))
diff --git a/evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py b/evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py
new file mode 100644
index 0000000000000000000000000000000000000000..0ccbe91c683de4e272125b24f348ffc080134f50
--- /dev/null
+++ b/evaluation/gen_eval/mask2former/mask2former_swin-t-p4-w7-224_lsj_8x2_50e_coco.py
@@ -0,0 +1,61 @@
+_base_ = ['./mask2former_r50_lsj_8x2_50e_coco.py']
+pretrained = 'https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_tiny_patch4_window7_224.pth' # noqa
+depths = [2, 2, 6, 2]
+model = dict(
+ type='Mask2Former',
+ backbone=dict(
+ _delete_=True,
+ type='SwinTransformer',
+ embed_dims=96,
+ depths=depths,
+ num_heads=[3, 6, 12, 24],
+ window_size=7,
+ mlp_ratio=4,
+ qkv_bias=True,
+ qk_scale=None,
+ drop_rate=0.,
+ attn_drop_rate=0.,
+ drop_path_rate=0.3,
+ patch_norm=True,
+ out_indices=(0, 1, 2, 3),
+ with_cp=False,
+ convert_weights=True,
+ frozen_stages=-1,
+ init_cfg=dict(type='Pretrained', checkpoint=pretrained)),
+ panoptic_head=dict(
+ type='Mask2FormerHead', in_channels=[96, 192, 384, 768]),
+ init_cfg=None)
+
+# set all layers in backbone to lr_mult=0.1
+# set all norm layers, position_embeding,
+# query_embeding, level_embeding to decay_multi=0.0
+backbone_norm_multi = dict(lr_mult=0.1, decay_mult=0.0)
+backbone_embed_multi = dict(lr_mult=0.1, decay_mult=0.0)
+embed_multi = dict(lr_mult=1.0, decay_mult=0.0)
+custom_keys = {
+ 'backbone': dict(lr_mult=0.1, decay_mult=1.0),
+ 'backbone.patch_embed.norm': backbone_norm_multi,
+ 'backbone.norm': backbone_norm_multi,
+ 'absolute_pos_embed': backbone_embed_multi,
+ 'relative_position_bias_table': backbone_embed_multi,
+ 'query_embed': embed_multi,
+ 'query_feat': embed_multi,
+ 'level_embed': embed_multi
+}
+custom_keys.update({
+ f'backbone.stages.{stage_id}.blocks.{block_id}.norm': backbone_norm_multi
+ for stage_id, num_blocks in enumerate(depths)
+ for block_id in range(num_blocks)
+})
+custom_keys.update({
+ f'backbone.stages.{stage_id}.downsample.norm': backbone_norm_multi
+ for stage_id in range(len(depths) - 1)
+})
+# optimizer
+optimizer = dict(
+ type='AdamW',
+ lr=0.0001,
+ weight_decay=0.05,
+ eps=1e-8,
+ betas=(0.9, 0.999),
+ paramwise_cfg=dict(custom_keys=custom_keys, norm_decay_mult=0.0))
diff --git a/evaluation/gen_eval/prompts/create_prompts.py b/evaluation/gen_eval/prompts/create_prompts.py
new file mode 100644
index 0000000000000000000000000000000000000000..5c7412e818157319a65608d1da1657fc80be3749
--- /dev/null
+++ b/evaluation/gen_eval/prompts/create_prompts.py
@@ -0,0 +1,183 @@
+"""
+Generate prompts for evaluation
+"""
+
+import argparse
+import json
+import os
+import yaml
+
+import numpy as np
+
+# Load classnames
+
+with open("object_names.txt") as cls_file:
+ classnames = [line.strip() for line in cls_file]
+
+# Proper a vs an
+
+def with_article(name: str):
+ if name[0] in "aeiou":
+ return f"an {name}"
+ return f"a {name}"
+
+# Proper plural
+
+def make_plural(name: str):
+ if name[-1] in "s":
+ return f"{name}es"
+ return f"{name}s"
+
+# Generates single object samples
+
+def generate_single_object_sample(rng: np.random.Generator, size: int = None):
+ TAG = "single_object"
+ if size > len(classnames):
+ size = len(classnames)
+ print(f"Not enough distinct classes, generating only {size} samples")
+ return_scalar = size is None
+ size = size or 1
+ idxs = rng.choice(len(classnames), size=size, replace=False)
+ samples = [dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx], "count": 1}
+ ],
+ prompt=f"a photo of {with_article(classnames[idx])}"
+ ) for idx in idxs]
+ if return_scalar:
+ return samples[0]
+ return samples
+
+# Generate two object samples
+
+def generate_two_object_sample(rng: np.random.Generator):
+ TAG = "two_object"
+ idx_a, idx_b = rng.choice(len(classnames), size=2, replace=False)
+ return dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx_a], "count": 1},
+ {"class": classnames[idx_b], "count": 1}
+ ],
+ prompt=f"a photo of {with_article(classnames[idx_a])} and {with_article(classnames[idx_b])}"
+ )
+
+# Generate counting samples
+
+numbers = ["zero", "one", "two", "three", "four", "five", "six", "seven", "eight", "nine", "ten"]
+
+def generate_counting_sample(rng: np.random.Generator, max_count=4):
+ TAG = "counting"
+ idx = rng.choice(len(classnames))
+ num = int(rng.integers(2, max_count, endpoint=True))
+ return dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx], "count": num}
+ ],
+ exclude=[
+ {"class": classnames[idx], "count": num + 1}
+ ],
+ prompt=f"a photo of {numbers[num]} {make_plural(classnames[idx])}"
+ )
+
+# Generate color samples
+
+colors = ["red", "orange", "yellow", "green", "blue", "purple", "pink", "brown", "black", "white"]
+
+def generate_color_sample(rng: np.random.Generator):
+ TAG = "colors"
+ idx = rng.choice(len(classnames) - 1) + 1
+ idx = (idx + classnames.index("person")) % len(classnames) # No "[COLOR] person" prompts
+ color = colors[rng.choice(len(colors))]
+ return dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx], "count": 1, "color": color}
+ ],
+ prompt=f"a photo of {with_article(color)} {classnames[idx]}"
+ )
+
+# Generate position samples
+
+positions = ["left of", "right of", "above", "below"]
+
+def generate_position_sample(rng: np.random.Generator):
+ TAG = "position"
+ idx_a, idx_b = rng.choice(len(classnames), size=2, replace=False)
+ position = positions[rng.choice(len(positions))]
+ return dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx_b], "count": 1},
+ {"class": classnames[idx_a], "count": 1, "position": (position, 0)}
+ ],
+ prompt=f"a photo of {with_article(classnames[idx_a])} {position} {with_article(classnames[idx_b])}"
+ )
+
+# Generate color attribution samples
+
+def generate_color_attribution_sample(rng: np.random.Generator):
+ TAG = "color_attr"
+ idxs = rng.choice(len(classnames) - 1, size=2, replace=False) + 1
+ idx_a, idx_b = (idxs + classnames.index("person")) % len(classnames) # No "[COLOR] person" prompts
+ cidx_a, cidx_b = rng.choice(len(colors), size=2, replace=False)
+ return dict(
+ tag=TAG,
+ include=[
+ {"class": classnames[idx_a], "count": 1, "color": colors[cidx_a]},
+ {"class": classnames[idx_b], "count": 1, "color": colors[cidx_b]}
+ ],
+ prompt=f"a photo of {with_article(colors[cidx_a])} {classnames[idx_a]} and {with_article(colors[cidx_b])} {classnames[idx_b]}"
+ )
+
+
+# Generate evaluation suite
+
+def generate_suite(rng: np.random.Generator, n: int = 100, output_path: str = ""):
+ samples = []
+ # Generate single object samples for all COCO classnames
+ samples.extend(generate_single_object_sample(rng, size=len(classnames)))
+ # Generate two object samples (~100)
+ for _ in range(n):
+ samples.append(generate_two_object_sample(rng))
+ # Generate counting samples
+ for _ in range(n):
+ samples.append(generate_counting_sample(rng, max_count=4))
+ # Generate color samples
+ for _ in range(n):
+ samples.append(generate_color_sample(rng))
+ # Generate position samples
+ for _ in range(n):
+ samples.append(generate_position_sample(rng))
+ # Generate color attribution samples
+ for _ in range(n):
+ samples.append(generate_color_attribution_sample(rng))
+ # De-duplicate
+ unique_samples, used_samples = [], set()
+ for sample in samples:
+ sample_text = yaml.safe_dump(sample)
+ if sample_text not in used_samples:
+ unique_samples.append(sample)
+ used_samples.add(sample_text)
+
+ # Write to files
+ os.makedirs(output_path, exist_ok=True)
+ with open(os.path.join(output_path, "generation_prompts.txt"), "w") as fp:
+ for sample in unique_samples:
+ print(sample['prompt'], file=fp)
+ with open(os.path.join(output_path, "evaluation_metadata.jsonl"), "w") as fp:
+ for sample in unique_samples:
+ print(json.dumps(sample), file=fp)
+
+
+if __name__ == "__main__":
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--seed", type=int, default=43, help="generation seed (default: 43)")
+ parser.add_argument("--num-prompts", "-n", type=int, default=100, help="number of prompts per task (default: 100)")
+ parser.add_argument("--output-path", "-o", type=str, default="prompts", help="output folder for prompts and metadata (default: 'prompts/')")
+ args = parser.parse_args()
+ rng = np.random.default_rng(args.seed)
+ generate_suite(rng, args.num_prompts, args.output_path)
+
diff --git a/evaluation/gen_eval/summary_scores.py b/evaluation/gen_eval/summary_scores.py
new file mode 100644
index 0000000000000000000000000000000000000000..6de7b08a6c8a28df60c39002b720196afbab58fd
--- /dev/null
+++ b/evaluation/gen_eval/summary_scores.py
@@ -0,0 +1,45 @@
+# Get results of evaluation
+
+import argparse
+import os
+
+import numpy as np
+import pandas as pd
+
+
+parser = argparse.ArgumentParser()
+parser.add_argument("filename", type=str)
+args = parser.parse_args()
+
+# Load classnames
+
+with open(os.path.join(os.path.dirname(__file__), "object_names.txt")) as cls_file:
+ classnames = [line.strip() for line in cls_file]
+ cls_to_idx = {"_".join(cls.split()):idx for idx, cls in enumerate(classnames)}
+
+# Load results
+
+df = pd.read_json(args.filename, orient="records", lines=True)
+
+# Measure overall success
+
+print("Summary")
+print("=======")
+print(f"Total images: {len(df)}")
+print(f"Total prompts: {len(df.groupby('metadata'))}")
+print(f"% correct images: {df['correct'].mean():.2%}")
+print(f"% correct prompts: {df.groupby('metadata')['correct'].any().mean():.2%}")
+print()
+
+# By group
+
+task_scores = []
+
+print("Task breakdown")
+print("==============")
+for tag, task_df in df.groupby('tag', sort=False):
+ task_scores.append(task_df['correct'].mean())
+ print(f"{tag:<16} = {task_df['correct'].mean():.2%} ({task_df['correct'].sum()} / {len(task_df)})")
+print()
+
+print(f"Overall score (avg. over tasks): {np.mean(task_scores):.5f}")
\ No newline at end of file
diff --git a/grn/__init__.py b/grn/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/grn/dataset/build.py b/grn/dataset/build.py
new file mode 100644
index 0000000000000000000000000000000000000000..3b432b4f6ba4dc5557bff44fca21c9a1c7dcf022
--- /dev/null
+++ b/grn/dataset/build.py
@@ -0,0 +1,150 @@
+import datetime
+import os
+import os.path as osp
+import random
+import subprocess
+from functools import partial
+from typing import Optional
+import time
+
+import pytz
+
+from grn.dataset.dataset_joint_vi import JointViDataset
+from grn.utils_t2iv.sequence_parallel import SequenceParallelManager as sp_manager
+
+try:
+ from grp import getgrgid
+ from pwd import getpwuid
+except:
+ pass
+import PIL.Image as PImage
+from PIL import ImageFile
+import numpy as np
+from torchvision.transforms import transforms
+from torchvision.transforms.functional import resize, to_tensor
+import torch.distributed as tdist
+
+from torchvision.transforms import InterpolationMode
+bicubic = InterpolationMode.BICUBIC
+lanczos = InterpolationMode.LANCZOS
+PImage.MAX_IMAGE_PIXELS = (1024 * 1024 * 1024 // 4 // 3) * 5
+ImageFile.LOAD_TRUNCATED_IMAGES = False
+
+
+def time_str(fmt='[%m-%d %H:%M:%S]'):
+ return datetime.datetime.now(tz=pytz.timezone('Asia/Shanghai')).strftime(fmt)
+
+
+def normalize_01_into_pm1(x): # normalize x from [0, 1] to [-1, 1] by (x*2) - 1
+ return x.add(x).add_(-1)
+
+
+def denormalize_pm1_into_01(x): # denormalize x from [-1, 1] to [0, 1]
+ return x.add(1).mul_(0.5)
+
+
+def center_crop_arr(pil_image, image_size):
+ """
+ Center cropping implementation from ADM.
+ https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
+ """
+ while min(*pil_image.size) >= 2 * image_size:
+ pil_image = pil_image.resize(
+ tuple(x // 2 for x in pil_image.size), resample=PImage.BOX
+ )
+
+ scale = image_size / min(*pil_image.size)
+ pil_image = pil_image.resize(
+ tuple(round(x * scale) for x in pil_image.size), resample=PImage.LANCZOS
+ )
+
+ arr = np.array(pil_image)
+ crop_y = (arr.shape[0] - image_size) // 2
+ crop_x = (arr.shape[1] - image_size) // 2
+ return PImage.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size])
+
+
+class RandomResize:
+ def __init__(self, mid_reso, final_reso, interpolation):
+ ub = max(round((mid_reso + (mid_reso-final_reso) / 8) / 4) * 4, mid_reso)
+ self.reso_lb, self.reso_ub = final_reso, ub
+ self.interpolation = interpolation
+
+ def __call__(self, img):
+ return resize(img, size=random.randint(self.reso_lb, self.reso_ub), interpolation=self.interpolation)
+
+ def __repr__(self):
+ return f'RandomResize(reso=({self.reso_lb}, {self.reso_ub}), interpolation={self.interpolation})'
+
+
+def print_aug(transform, label):
+ print(f'Transform {label} = ')
+ if hasattr(transform, 'transforms'):
+ for t in transform.transforms:
+ print(t)
+ else:
+ print(transform)
+ print('---------------------------\n')
+
+
+def build_joint_dataset(
+ args,
+ meta_folders: str,
+ meta_folder_repeats: str,
+ max_caption_len: int,
+ short_prob=0.2,
+ load_vae_instead_of_image=False
+):
+ return JointViDataset(
+ meta_folders=meta_folders,
+ meta_folder_repeats=meta_folder_repeats,
+ max_caption_len=max_caption_len,
+ short_prob=short_prob,
+ load_vae_instead_of_image=load_vae_instead_of_image,
+ video_fps=args.video_fps,
+ num_frames=args.video_frames,
+ online_t5=args.online_t5,
+ num_replicas=sp_manager.get_sp_group_nums() if sp_manager.sp_on() else tdist.get_world_size(), # 1,
+ rank = sp_manager.get_sp_group_rank() if sp_manager.sp_on() else tdist.get_rank(),
+ dataloader_workers=args.workers,
+ enable_dynamic_length_prompt=args.enable_dynamic_length_prompt,
+ hdfs_mode=args.hdfs_mode,
+ dynamic_scale_schedule=args.dynamic_scale_schedule,
+ seed=args.seed,
+ other_args=args,
+ )
+
+
+def pil_load(path: str, proposal_size):
+ with open(path, 'rb') as f:
+ img: PImage.Image = PImage.open(f)
+ w: int = img.width
+ h: int = img.height
+ sh: int = min(h, w)
+ if sh > proposal_size:
+ ratio: float = proposal_size / sh
+ w = round(ratio * w)
+ h = round(ratio * h)
+ img.draft('RGB', (w, h))
+ img = img.convert('RGB')
+ return img
+
+
+def rewrite(im: PImage, file: str, info: str):
+ kw = dict(quality=100)
+ if file.lower().endswith('.tif') or file.lower().endswith('.tiff'):
+ kw['compression'] = 'none'
+ elif file.lower().endswith('.webp'):
+ kw['lossless'] = True
+
+ st = os.stat(file)
+ uname = getpwuid(st.st_uid).pw_name
+ gname = getgrgid(st.st_gid).gr_name
+ mode = oct(st.st_mode)[-3:]
+
+ local_file = osp.basename(file)
+ im.save(local_file, **kw)
+ print(f'************* ************* @ {file}')
+ subprocess.call(['sudo', 'mv', local_file, file])
+ subprocess.call(['sudo', 'chown', f'{uname}:{gname}', file])
+ subprocess.call(['sudo', 'chmod', str(mode), file])
diff --git a/grn/dataset/dataset_joint_vi.py b/grn/dataset/dataset_joint_vi.py
new file mode 100644
index 0000000000000000000000000000000000000000..3906506b37329b986b9a0438c6441254850f3001
--- /dev/null
+++ b/grn/dataset/dataset_joint_vi.py
@@ -0,0 +1,687 @@
+import glob
+import os
+import pickle
+import random
+import re
+import time
+from functools import partial
+from os import path as osp
+from typing import List, Tuple, Union
+import json
+import itertools
+import hashlib
+import copy
+import collections
+import math
+
+import tqdm
+import numpy as np
+import torch
+import pandas as pd
+from decord import VideoReader
+from PIL import Image as PImage
+from torch.nn import functional as F
+from torchvision.transforms.functional import to_tensor, hflip
+from torchvision.transforms import transforms, InterpolationMode
+from torch.utils.data import Dataset, DataLoader
+import torch.distributed as tdist
+from PIL import Image
+os.environ["TOKENIZERS_PARALLELISM"] = "false"
+
+from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta
+from grn.utils.video_decoder import EncodedVideoDecord
+from grn.utils.compress_tokens import load_packed_tensor
+from transformers import AutoTokenizer
+
+def transform(pil_img, tgt_h, tgt_w):
+ width, height = pil_img.size
+ if width / height <= tgt_w / tgt_h:
+ resized_width = tgt_w
+ resized_height = int(tgt_w / (width / height))
+ else:
+ resized_height = tgt_h
+ resized_width = int((width / height) * tgt_h)
+ pil_img = pil_img.resize((resized_width, resized_height), resample=PImage.LANCZOS)
+ # crop the center out
+ arr = np.array(pil_img)
+ crop_y = (arr.shape[0] - tgt_h) // 2
+ crop_x = (arr.shape[1] - tgt_w) // 2
+ im = to_tensor(arr[crop_y: crop_y + tgt_h, crop_x: crop_x + tgt_w])
+ return im.add(im).add_(-1)
+
+def normalize(x): # normalize x from [0, 1] to [-1, 1] by (x*2) - 1
+ return x.add(x).add_(-1)
+
+def get_prompt_id(prompt):
+ md5 = hashlib.md5()
+ md5.update(prompt.encode('utf-8'))
+ prompt_id = md5.hexdigest()
+ return prompt_id
+
+def prepend_motion_score(prompt, motion_score):
+ return f'<<>> {prompt}'
+
+class VideoReaderWrapper(VideoReader):
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.seek(0)
+ def __getitem__(self, key):
+ frames = super().__getitem__(key)
+ self.seek(0)
+ return frames
+
+class JointViDataset(Dataset):
+ def __init__(
+ self,
+ meta_folders: str = '',
+ meta_folder_repeats: str = '',
+ buffersize: int = 1000000 * 300,
+ seed: int = 0,
+ pn: str = '',
+ video_fps: int = 1,
+ num_replicas: int = 1,
+ rank: int = 0,
+ dataloader_workers: int = 2,
+ enable_dynamic_length_prompt: bool = True,
+ shuffle: bool = True,
+ short_prob: float = 0.2,
+ verbose=False,
+ temp_dir= "/dev/shm",
+ hdfs_mode='read',
+ other_args=None,
+ **kwargs,
+ ):
+ self.meta_folders = json.loads(meta_folders)
+ self.meta_folder_repeats = json.loads(meta_folder_repeats)
+ self.meta_folder_identifiers = json.loads(other_args.meta_folder_identifiers)
+ self.pn_list = json.loads(other_args.pn_list)
+ self.pn_probs = json.loads(other_args.pn_probs)
+ assert len(self.meta_folders) == len(self.meta_folder_repeats), f'{len(self.meta_folders)} != {len(self.meta_folder_repeats)}'
+ assert len(self.meta_folders) == len(self.meta_folder_identifiers), f'{len(self.meta_folders)} != {len(self.meta_folder_identifiers)}'
+ self.verbose = verbose
+ self.buffer_size = buffersize
+ self.num_replicas = num_replicas
+ self.rank = rank
+ self.worker_id = 0
+ self.global_worker_id = 0
+ self.short_prob = short_prob
+ self.dataloader_workers = max(1, dataloader_workers)
+ self.shuffle = shuffle
+ self.global_workers = self.num_replicas * self.dataloader_workers
+ self.seed = seed
+ self.text_tokenizer = other_args.text_tokenizer
+ self.feature_extraction = other_args.only_images4extract_feats # < 0 # no sequence packing, for feature extraction
+ self.epoch_generator = None
+ self.epoch_rank_generator = None
+ self.other_args = other_args
+ self.pair_input = other_args.pair_input
+ self.drop_long_video = other_args.drop_long_video
+ self.enable_dynamic_length_prompt = enable_dynamic_length_prompt
+ self.set_epoch_generator(other_args.epoch)
+ self.temporal_compress_rate = other_args.temporal_compress_rate
+ 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
+ self.video_fps = video_fps
+ self.min_training_duration = (other_args.min_video_frames-1) // self.video_fps
+ self.max_training_duration = (other_args.video_frames-1) // self.video_fps
+ self.c2i = self.other_args.add_class_token > 0
+ if self.c2i:
+ self.c2i_transform = transforms.Compose([
+ transforms.RandomHorizontalFlip(),
+ transforms.Resize(288, interpolation=InterpolationMode.LANCZOS), # transforms.Resize: resize the shorter edge to mid_reso
+ transforms.RandomCrop((256, 256)),
+ transforms.ToTensor(),
+ normalize,
+ ])
+ else:
+ self.c2i_transform = None
+ 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}")
+ self.token_cache_dir = other_args.token_cache_dir
+ self.use_vae_token_cache = other_args.use_vae_token_cache
+ self.allow_online_vae_feature_extraction = other_args.allow_online_vae_feature_extraction
+ self.use_text_token_cache = other_args.use_text_token_cache
+ self.max_video_frames = other_args.video_frames
+ self.cached_video_frames = other_args.cached_video_frames # cached max video frames
+ self.down_size_limit = other_args.down_size_limit
+ self.video_caption_type = other_args.video_caption_type
+ self.train_max_token_len = other_args.train_max_token_len
+ self.duration_resolution = other_args.duration_resolution
+ self.device = other_args.device
+ print(f'self.down_size_limit: {self.down_size_limit}')
+ self.hdfs_mode = hdfs_mode
+ self.max_text_len = other_args.tlen
+ self.temp_dir = temp_dir.rstrip("/")
+ self.mapped_duration2metas, self.mapped_duration2freqs = self.get_mapped_duration2metas()
+ self.batches = self.form_batches(self.mapped_duration2metas)
+ print(f'{num_replicas=}, {rank=}, {dataloader_workers=}, {len(self.batches)=}, {self.drop_long_video=} {self.max_text_len=} self.batches[:10]={self.batches[:10]}')
+
+ def print(self, string):
+ if self.feature_extraction:
+ print(string)
+ else:
+ print(string, force=True)
+
+ def get_captions_lens(self, captions):
+ if self.other_args.text_tokenizer_type == 'flan_t5':
+ 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')
+ mask = tokens.attention_mask.cuda(non_blocking=True)
+ lens: List[int] = mask.sum(dim=-1).tolist()
+ else: # umt5-xxl
+ ids, mask = self.other_args.text_tokenizer( captions, return_mask=True, add_special_tokens=True)
+ lens = mask.gt(0).sum(dim=1).tolist()
+ return lens
+
+ def get_video_caption(self, meta, mapped_duration):
+ if 'tarsier2_caption' not in meta:
+ caption = self.epoch_rank_generator.choice(meta['caption'])['content']
+ else:
+ caption_type = 'tarsier2_caption'
+ if ('MiniCPM_V_2_6_caption' in meta) and meta['MiniCPM_V_2_6_caption']:
+ caption_type = self.epoch_rank_generator.choice(['tarsier2_caption', 'MiniCPM_V_2_6_caption'])
+ caption = meta[caption_type]
+ if self.enable_dynamic_length_prompt and (self.epoch_rank_generator.random() < self.other_args.short_cap_prob):
+ caption = self.random_drop_sentences(caption, min_sentences=2)
+ if 'quality_prompt' in meta:
+ caption = caption + ' ' + meta['quality_prompt']
+ if meta['first_frame_condition']:
+ caption = '' + caption
+ else:
+ caption = '' + caption
+ assert caption
+ return caption
+
+ def get_image_caption(self, meta):
+ caption = meta['long_caption']
+ if not meta['long_caption']:
+ caption = meta['text']
+ else:
+ if self.epoch_rank_generator.random() < self.other_args.short_cap_prob:
+ if meta['text']:
+ caption = meta['text']
+ elif ('InternVL' in meta['long_caption_type']):
+ caption = self.random_drop_sentences(meta['long_caption'], min_sentences=2)
+ caption = '' + caption
+ assert caption
+ return caption
+
+ def sample_pn(self, meta, pn_list, pn_probs):
+ if ('height' in meta) and ('width' in meta):
+ real_pn = meta['height'] * meta['width'] / 1000000
+ valid_pn_list, valid_pn_probs = [], []
+ for pn, pn_prob in zip(pn_list, pn_probs):
+ if real_pn > 0.7 * float(pn[:-1]):
+ valid_pn_list.append(pn)
+ valid_pn_probs.append(pn_prob)
+ if not len(valid_pn_list):
+ valid_pn_list = pn_list[:1]
+ valid_pn_probs = pn_probs[:1]
+ else:
+ valid_pn_list = pn_list
+ valid_pn_probs = pn_probs
+ return self.epoch_rank_generator.choice(valid_pn_list, p=valid_pn_probs)
+
+ def get_mapped_duration2metas(self):
+ part_filepaths = []
+ part_file2identifier = {}
+ for meta_folder, meta_folder_repeat, meta_folder_identifier in zip(self.meta_folders, self.meta_folder_repeats, self.meta_folder_identifiers):
+ tmp_part_filepaths = sorted(glob.glob(osp.join(meta_folder, '*/*.jsonl')))
+ self.epoch_generator.shuffle(tmp_part_filepaths)
+ if meta_folder_repeat > 1:
+ tmp_part_filepaths = tmp_part_filepaths * int(np.ceil(meta_folder_repeat))
+ tmp_part_filepaths = tmp_part_filepaths[:int(len(tmp_part_filepaths)*meta_folder_repeat)]
+ for file in tmp_part_filepaths: part_file2identifier[file] = meta_folder_identifier
+ part_filepaths.extend(tmp_part_filepaths)
+ self.epoch_generator.shuffle(part_filepaths)
+ self.print(f'{self.rank=} jsonls sample: {part_filepaths[:4]}')
+ if self.num_replicas > 1:
+ part_filepaths = part_filepaths[self.rank::self.num_replicas]
+
+ mapped_duration2metas = {}
+ pbar = tqdm.tqdm(total=len(part_filepaths))
+ total, corrupt = 0, 0
+ stop_read = False
+ rough_h_div_w = self.h_div_w_templates[np.argmin(np.abs((9/16-self.h_div_w_templates)))]
+ for part_filepath in part_filepaths:
+ file_quality_prompt = part_file2identifier[part_filepath]
+ if stop_read:
+ break
+ pbar.update(1)
+ try:
+ with open(part_filepath, encoding='utf-8') as f:
+ lines = f.readlines()
+ except Exception as e:
+ print(f'{part_filepath=} Error: {e}')
+ lines = []
+ for line in lines:
+ total += 1
+ try:
+ meta = json.loads(line)
+ except Exception as e:
+ print(e)
+ corrupt += 1
+ print(e, corrupt, total, corrupt/total)
+ continue
+ if file_quality_prompt: # override quality prompt
+ meta['quality_prompt'] = file_quality_prompt
+ if ('height' in meta) and ('width' in meta):
+ cur_h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((meta['height']/meta['width']-self.h_div_w_templates)))]
+ else:
+ cur_h_div_w_template = rough_h_div_w
+ if 'h_div_w' in meta:
+ del meta['h_div_w']
+ meta['first_frame_condition'] = False
+ meta['pn'] = self.sample_pn(meta, self.pn_list, self.pn_probs)
+ if 'video_path' in meta:
+ if self.epoch_rank_generator.random() < self.other_args.i2v_ratio:
+ meta['first_frame_condition'] = True
+ begin_frame_id, end_frame_id, fps = meta['begin_frame_id'], meta['end_frame_id'], meta['fps']
+ real_duration = (end_frame_id - begin_frame_id) / fps
+ mapped_duration = int(np.round(real_duration / self.duration_resolution)) * self.duration_resolution
+ if mapped_duration < self.min_training_duration:
+ continue
+ if mapped_duration > self.max_training_duration:
+ if self.drop_long_video:
+ continue
+ else:
+ mapped_duration = self.max_training_duration
+ if self.other_args.use_clipwise_caption:
+ meta['caption'] = [
+ meta['caption-InternVL2.0'],
+ self.get_video_caption(meta, mapped_duration),
+ ]
+ else:
+ meta['caption'] = [self.get_video_caption(meta, mapped_duration)]
+ sample_frames = int(mapped_duration * self.video_fps + 1)
+ pt = (sample_frames-1) // self.temporal_compress_rate + 1
+ scale_schedule = self.dynamic_resolution_h_w[cur_h_div_w_template][meta['pn']]['pt2scale_schedule'][pt]
+ meta['sample_frames'] = sample_frames
+ elif 'image_path' in meta:
+ mapped_duration = -1
+ scale_schedule = self.dynamic_resolution_h_w[cur_h_div_w_template][meta['pn']]['pt2scale_schedule'][1]
+ meta['caption'] = [self.get_image_caption(meta)]
+ # random set caption to "" for classifier-free guidance
+ # refer to: https://github.com/PixArt-alpha/PixArt-alpha/blob/master/train_scripts/train_diffusers.py#L67
+ for caption_ind in range(len(meta['caption'])):
+ if self.epoch_rank_generator.random() < self.other_args.drop_condition_prob:
+ meta['caption'][caption_ind] = ""
+ if mapped_duration not in mapped_duration2metas:
+ mapped_duration2metas[mapped_duration] = []
+
+ # get cum_text_visual_tokens
+ cum_visual_tokens = []
+ preserve_scale_inds = {}
+ assert len(scale_schedule) == len(self.other_args.video_scale_probs), f'{len(scale_schedule)=} {len(self.other_args.video_scale_probs)=}'
+ for scale_ind, scale in enumerate(scale_schedule):
+ if self.epoch_rank_generator.random() < self.other_args.video_scale_probs[scale_ind]:
+ preserve_scale_inds[scale_ind] = True
+ tokens_this_scale = np.array(scale).prod(-1) + self.other_args.add_scale_token
+ cum_visual_tokens.append(tokens_this_scale)
+ cum_visual_tokens = np.array(cum_visual_tokens).cumsum()
+ meta['cum_text_visual_tokens'] = cum_visual_tokens
+ meta['preserve_scale_inds'] = preserve_scale_inds
+
+ if self.other_args.cache_check_mode == 1: # check at the begining
+ if self.exists_cache_file(meta):
+ mapped_duration2metas[mapped_duration].append(meta)
+ elif self.other_args.cache_check_mode == -1: # select unexist, used for token cache
+ if not self.exists_cache_file(meta):
+ mapped_duration2metas[mapped_duration].append(meta)
+ else:
+ mapped_duration2metas[mapped_duration].append(meta)
+
+ total_metas = sum([len(item) for item in mapped_duration2metas.values()])
+ if (self.other_args.restrict_data_size > 0) and (total_metas > self.other_args.restrict_data_size / self.num_replicas):
+ stop_read = True
+ break
+
+ # set mapped_duration2freqs
+ mapped_duration2freqs = {}
+ for mapped_duration in sorted(mapped_duration2metas.keys()):
+ mapped_duration2freqs[mapped_duration] = len(mapped_duration2metas[mapped_duration])
+
+ for mapped_duration in mapped_duration2metas.keys():
+ freqs = mapped_duration2freqs[mapped_duration]
+ assert len(mapped_duration2metas[mapped_duration]) >= freqs
+ self.epoch_rank_generator.shuffle(mapped_duration2metas[mapped_duration])
+ mapped_duration2metas[mapped_duration] = mapped_duration2metas[mapped_duration][:freqs]
+ # append text tokens
+ skip_count_text_token = self.other_args.skip_count_text_token or self.other_args.add_class_token > 0
+ mapped_duration2metas[mapped_duration] = self.append_text_tokens(mapped_duration2metas[mapped_duration], skip_count_text_token=skip_count_text_token)
+
+ total_metas = sum([len(item) for item in mapped_duration2metas.values()])
+ for mapped_duration in sorted(mapped_duration2freqs.keys()):
+ freq = mapped_duration2freqs[mapped_duration]
+ proportion = freq / total_metas * 100
+ print(f'{mapped_duration=}, {freq=}, {proportion=:.1f}%')
+ return mapped_duration2metas, mapped_duration2freqs
+
+ def append_text_tokens(self, metas, skip_count_text_token=False, bucket_size=100):
+ t1 = time.time()
+ pbar = tqdm.tqdm(total=len(metas) // bucket_size + 1, desc='append text tokens')
+ valid_metas = []
+ for bucket_id in range(len(metas) // bucket_size + 1):
+ pbar.update(1)
+ start = bucket_id * bucket_size
+ end = min(start + bucket_size, len(metas))
+ if start >= end:
+ break
+ captions = []
+ caps_per_meta = []
+ for i in range(start, end):
+ captions.extend(metas[i]['caption'])
+ caps_per_meta.append(len(metas[i]['caption']))
+ assert len(captions), f'{len(captions)=}'
+ if skip_count_text_token:
+ lens = [0 for _ in range(len(captions))]
+ else:
+ lens = self.get_captions_lens(captions)
+ lens = np.clip(np.array(lens), a_min=0, a_max=self.max_text_len)
+ ptr = 0
+ for i in range(start, end):
+ text_tokens = sum(lens[ptr:ptr+caps_per_meta[i-start]])
+ ptr += caps_per_meta[i-start]
+ metas[i]['text_tokens'] = text_tokens
+ metas[i]['cum_text_visual_tokens'] = metas[i]['cum_text_visual_tokens'] + metas[i]['text_tokens']
+ metas[i]['text_visual_tokens'] = metas[i]['cum_text_visual_tokens'][-1]
+ if metas[i]['text_visual_tokens'] <= self.train_max_token_len * self.other_args.dense_ratio4seqpack:
+ valid_metas.append(metas[i])
+ t2 = time.time()
+ print(f'append text tokens: {t2-t1:.1f}s')
+ return valid_metas
+
+ def exists_cache_file(self, meta):
+ pn = meta['pn']
+ if 'image_path' in meta:
+ return osp.exists(self.get_image_cache_file(meta['image_path'], pn))
+ else:
+ if '/vdataset/clip' in meta['video_path']: # clip
+ cache_file = self.get_video_cache_file(meta['video_path'], 0, meta['end_frame_id']-meta['begin_frame_id'], self.video_fps, pn)
+ else:
+ cache_file = self.get_video_cache_file(meta['video_path'], meta['begin_frame_id'], meta['end_frame_id'], self.video_fps, pn)
+ return osp.exists(cache_file)
+
+ def form_batches(self, mapped_duration2metas):
+ examples = []
+ for mapped_duration in sorted(mapped_duration2metas.keys()):
+ for example_ind in range(len(mapped_duration2metas[mapped_duration])):
+ text_visual_tokens = mapped_duration2metas[mapped_duration][example_ind]['text_visual_tokens']
+ examples.append((mapped_duration, example_ind, text_visual_tokens))
+ examples = sorted(examples, key=lambda x: -x[2])
+ max_text_visual_tokens = examples[0][2] if len(examples) else 0
+ assert self.train_max_token_len >= max_text_visual_tokens, f'{self.train_max_token_len=} should >= {max_text_visual_tokens=}'
+ self.print(f'{self.rank=} {self.mapped_duration2freqs=} form_batches details: {self.rank=} examples={examples[:20]}')
+
+ st = time.time()
+ if self.feature_extraction or self.pair_input: # no sequence packing, for feature extraction or dpo training
+ batches = [[item[:2]] for item in examples]
+ else:
+ batches = []
+ left_ptr, right_ptr = 0, len(examples)-1
+ while left_ptr <= right_ptr:
+ tokens_remain = self.train_max_token_len
+ tmp_batch = []
+ while left_ptr <= right_ptr and (tokens_remain - examples[left_ptr][2] >= 0):
+ tokens_remain = tokens_remain - examples[left_ptr][2]
+ tmp_batch.append((left_ptr, examples[left_ptr][2]))
+ left_ptr += 1
+ while left_ptr <= right_ptr and (tokens_remain - examples[right_ptr][2] >= 0):
+ tokens_remain = tokens_remain - examples[right_ptr][2]
+ tmp_batch.append((right_ptr, examples[right_ptr][2]))
+ right_ptr -= 1
+ if len(tmp_batch):
+ tmp_batch = sorted(tmp_batch, key=lambda x: -x[1])
+ # total_tokens = sum(item[1] for item in tmp_batch)
+ # if total_tokens / self.train_max_token_len < 0.7:
+ # import pdb; pdb.set_trace()
+ batches.append([examples[ptr][:2] for (ptr, _) in tmp_batch])
+ if len(batches) % 1000 == 0:
+ print(f'form {len(batches)} batches, len(metas)={len(examples)}')
+ print(f'[data preprocess] form_batches done, got {len(batches)} batches, cost {time.time()-st:.2f}s')
+ self.epoch_rank_generator.shuffle(batches)
+ print(f'[data preprocess] shuffle batches done')
+ batch_num = len(batches)
+ try:
+ if self.num_replicas > 1:
+ batch_num = torch.tensor([batch_num], device=self.device)
+ if tdist.is_initialized():
+ tdist.all_reduce(batch_num, op=tdist.ReduceOp.MIN)
+ batch_num = batch_num.item()
+ except Exception as e:
+ print(e)
+ batches = batches[:batch_num]
+ print(f'[data preprocess] aligned batch number among gpus, got {batch_num} batches')
+ return batches
+
+ def set_epoch_generator(self, epoch):
+ self.epoch = epoch
+ self.epoch_generator = np.random.default_rng(self.seed + self.epoch)
+ self.epoch_rank_generator = np.random.default_rng(self.seed + self.epoch + self.rank)
+
+ def __getitem__(self, batch_ind_ptr):
+ try:
+ batch_info = self.batches[batch_ind_ptr%len(self.batches)]
+ batch_data = []
+ for (mapped_duration, example_ind) in batch_info:
+ ret = False
+ repeat_times = 0
+ mapped_duration_metas = self.mapped_duration2metas[mapped_duration]
+ while not ret:
+ example_ind = example_ind % len(mapped_duration_metas)
+ meta = mapped_duration_metas[example_ind]
+ if 'video_path' in meta:
+ if self.pair_input:
+ ret, model_input = self.prepare_pair_video_input(meta)
+ else:
+ ret, model_input = self.prepare_video_input(meta)
+ elif 'image_path' in meta:
+ if self.pair_input:
+ ret, model_input = self.prepare_pair_image_input(meta)
+ else:
+ ret, model_input = self.prepare_image_input(meta)
+ if ret:
+ if self.pair_input:
+ batch_data.extend(model_input)
+ else:
+ batch_data.append(model_input)
+ else: # Handle corrupt example in a batch, just try to read the next one
+ example_ind = example_ind + 1
+ repeat_times += 1
+ if repeat_times % 20 == 0: # Too many corrupt files, switch to another batch
+ self.print(f'Caution! I have repeat {repeat_times} times to read a video/image, but still failed to read it. {example_ind=} {meta=}')
+ return self.__getitem__(batch_ind_ptr+1)
+
+ images, raw_features_bcthw, feature_cache_files4images = [], [], []
+ text_feature_cache_files = []
+ addition_pn_images = {}
+ batch_data4images, batch_data4raw_features = [], []
+ for item in batch_data:
+ if item['raw_features_cthw'] is None:
+ images.append(item['img_T3HW'].permute(1,0,2,3)) # # tchw -> cthw
+ for key in item:
+ if key.startswith('img_T3HW_'):
+ if key not in addition_pn_images:
+ addition_pn_images[key] = []
+ addition_pn_images[key].append(item[key].permute(1,0,2,3))
+ feature_cache_files4images.append(item['feature_cache_file'])
+ batch_data4images.append(item)
+ else:
+ raw_features_bcthw.append(item['raw_features_cthw'])
+ batch_data4raw_features.append(item)
+ batch_data4images_raw_features = batch_data4images + batch_data4raw_features
+ captions = [item['text_input'] for item in batch_data4images_raw_features]
+ text_feature_cache_files = [item['text_feature_cache_file'] for item in batch_data4images_raw_features]
+ meta_list = [item['meta'] for item in batch_data4images_raw_features]
+ return {
+ 'captions': captions,
+ 'images': images,
+ 'addition_pn_images': addition_pn_images,
+ 'feature_cache_files4images': feature_cache_files4images,
+ 'raw_features_bcthw': raw_features_bcthw,
+ 'text_cond_tuple': None,
+ 'text_feature_cache_files': text_feature_cache_files,
+ 'meta_list': meta_list,
+ 'media': 'videos',
+ }
+ except Exception as e:
+ print(f'get item error: {e}')
+ return self.__getitem__(batch_ind_ptr+1)
+
+
+ def prepare_image_input(self, info) -> Tuple:
+ try:
+ img_path, text_input = osp.abspath(info['image_path']), info['caption']
+ img_T3HW, raw_features_cthw, feature_cache_file, text_features_lenxdim, text_feature_cache_file = [None] * 5
+ if self.use_vae_token_cache:
+ feature_cache_file = self.get_image_cache_file(img_path, info['pn'])
+ if osp.exists(feature_cache_file):
+ try:
+ raw_features_cthw = self.load_visual_token(feature_cache_file)
+ except Exception as e:
+ print(f'load cache file error: {e}')
+ os.remove(feature_cache_file)
+ if raw_features_cthw is None and (not self.allow_online_vae_feature_extraction):
+ return False, None
+ if raw_features_cthw is None:
+ with open(img_path, 'rb') as f:
+ img: PImage.Image = PImage.open(f)
+ w, h = img.size
+ h_div_w = h / w
+ h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((h_div_w-self.h_div_w_templates)))]
+ tgt_h, tgt_w = self.dynamic_resolution_h_w[h_div_w_template][info['pn']]['pixel']
+ img = img.convert('RGB')
+ if self.c2i:
+ img_T3HW = self.c2i_transform(img)
+ else:
+ img_T3HW = transform(img, tgt_h, tgt_w)
+ img_T3HW = img_T3HW.unsqueeze(0)
+ assert img_T3HW.shape == (1, 3, tgt_h, tgt_w)
+ data_item = {
+ 'text_input': text_input,
+ 'img_T3HW': img_T3HW,
+ 'raw_features_cthw': raw_features_cthw,
+ 'feature_cache_file': feature_cache_file,
+ 'text_features_lenxdim': text_features_lenxdim,
+ 'text_feature_cache_file': text_feature_cache_file,
+ 'meta': info,
+ }
+ return True, data_item
+ except Exception as e:
+ print(f'prepare_image_input error: {e}')
+ return False, None
+
+ def prepare_pair_image_input(self, info) -> Tuple:
+ pass
+
+ def prepare_pair_video_input(self, info) -> Tuple:
+ win_flag, win_data_item = self.prepare_video_input(copy.deepcopy(info))
+
+ info['video_path'] = info['lose_video_path']
+ lose_flag, lose_data_item = self.prepare_video_input(info)
+
+ flag = win_flag and lose_flag
+ return flag, [win_data_item, lose_data_item]
+
+ def load_visual_token(self, feature_cache_file):
+ raw_features_cthw = load_packed_tensor(feature_cache_file)
+ from grn.utils_t2iv.hbq_util_t2iv import bit_label2raw_feature
+ raw_features_cthw = bit_label2raw_feature(raw_features_cthw.unsqueeze(0), self.other_args.hbq_round)[0]
+ return raw_features_cthw
+
+ def prepare_video_input(self, info) -> Tuple:
+ filename, begin_frame_id, end_frame_id = (
+ info["video_path"],
+ info["begin_frame_id"],
+ info["end_frame_id"],
+ )
+
+ try:
+ img_T3HW, raw_features_cthw, feature_cache_file, text_features_lenxdim, text_feature_cache_file = None, None, None, None, None
+ img_T3HW_4additional_pn = {}
+ text_input = info['caption']
+ if '/vdataset/clip' in filename: # clip
+ begin_frame_id, end_frame_id = 0, end_frame_id - begin_frame_id
+ sample_frames = info['sample_frames']
+ tmp_local_path = ''
+ if self.use_vae_token_cache:
+ feature_cache_file = self.get_video_cache_file(info["video_path"], begin_frame_id, end_frame_id, self.video_fps, info['pn'])
+ if osp.exists(feature_cache_file):
+ try:
+ pt = (sample_frames-1) // self.temporal_compress_rate + 1
+ raw_features_cthw = self.load_visual_token(feature_cache_file)
+ assert raw_features_cthw.shape[1] >= pt, f'raw_features_cthw.shape[1] >= pt: {raw_features_cthw.shape[1]} vs {pt}'
+ if raw_features_cthw.shape[1] > pt:
+ raw_features_cthw = raw_features_cthw[:,:pt]
+ except Exception as e:
+ self.print(f'load video cache file error: {e}')
+ os.remove(feature_cache_file)
+ raw_features_cthw = None
+ if raw_features_cthw is None and (not self.allow_online_vae_feature_extraction):
+ return False, None
+ pn_list = [info['pn']]
+ if raw_features_cthw is None:
+ tmp_local_path = info["video_path"]
+ if not osp.exists(tmp_local_path):
+ return False, None
+ video = EncodedVideoDecord(tmp_local_path, os.path.basename(tmp_local_path), num_threads=0)
+ start_interval = max(0, begin_frame_id / video._fps)
+ end_interval = start_interval+(sample_frames-1)/self.video_fps
+ assert end_interval <= video.duration + 0.2, f'{end_interval=}, but {video.duration=}' # 0.2s margin
+ end_interval = min(end_interval, video.duration)
+ raw_video, _ = video.get_clip(start_interval, end_interval, sample_frames) # rgb order
+ h, w, _ = raw_video[0].shape
+ h_div_w = h / w
+ h_div_w_template = self.h_div_w_templates[np.argmin(np.abs((h_div_w-self.h_div_w_templates)))]
+ tgt_h, tgt_w = self.dynamic_resolution_h_w[h_div_w_template][info['pn']]['pixel']
+
+ for pn in pn_list:
+ img_T3HW = [transform(Image.fromarray(frame).convert("RGB"), tgt_h, tgt_w) for frame in raw_video]
+ img_T3HW = torch.stack(img_T3HW, 0)
+ img_T3HW_4additional_pn[pn] = img_T3HW
+ del video
+ assert img_T3HW.shape[-3:] == (3, tgt_h, tgt_w)
+ data_item = {
+ 'text_input': text_input,
+ 'img_T3HW': img_T3HW_4additional_pn.get(info['pn'], None),
+ 'raw_features_cthw': raw_features_cthw,
+ 'feature_cache_file': feature_cache_file,
+ 'text_features_lenxdim': text_features_lenxdim,
+ 'text_feature_cache_file': text_feature_cache_file,
+ 'meta': info,
+ }
+ for pn in pn_list[1:]:
+ data_item.update({f'img_T3HW_{pn}': img_T3HW_4additional_pn.get(pn, None)})
+ return True, data_item
+ except Exception as e:
+ self.print(f'prepare_video_input error: {e}, info: {info}')
+ return False, None
+
+
+ @staticmethod
+ def collate_function(batch, online_t5: bool = False) -> None:
+ pass
+
+ def random_drop_sentences(self, caption, min_sentences):
+ elems = [item for item in caption.split('.') if item]
+ if len(elems) <= min_sentences:
+ return caption
+ sentences = self.epoch_rank_generator.integers(min_sentences, len(elems)+1)
+ return '.'.join(elems[:sentences]) + '.'
+
+ def __len__(self):
+ return len(self.batches) * self.other_args.loop_data_per_epoch
+
+ def get_image_cache_file(self, image_path, pn):
+ elems = image_path.split('/')
+ elems = [item for item in elems if item]
+ filename, ext = osp.splitext(elems[-1])
+ filename = get_prompt_id(filename)
+ save_filepath = osp.join(self.token_cache_dir, f'images_pn_{pn}', '/'.join(elems[4:-1]), f'{filename}.npz')
+ return save_filepath
+
+ def get_video_cache_file(self, video_path, begin_frame_id, end_frame_id, video_fps, pn):
+ elems = video_path.split('/')
+ elems = [item for item in elems if item]
+ filename, ext = osp.splitext(elems[-1])
+ filename = get_prompt_id(filename)
+ 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')
+ return save_filepath
+
\ No newline at end of file
diff --git a/grn/models/basic.py b/grn/models/basic.py
new file mode 100644
index 0000000000000000000000000000000000000000..96f74e50c28c82138277e07406de2d706d146a3f
--- /dev/null
+++ b/grn/models/basic.py
@@ -0,0 +1,256 @@
+import math
+import os
+from functools import partial
+from typing import Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+from torch.utils.checkpoint import checkpoint
+from torch.nn.functional import scaled_dot_product_attention as slow_attn # q, k, v: BHLc
+
+from grn.models.rope import apply_rotary_emb
+from grn.utils_t2iv.sequence_parallel import sp_all_to_all, SequenceParallelManager as sp_manager
+
+try:
+ from flash_attn.cute import flash_attn_varlen_func
+except:
+ from flash_attn import flash_attn_varlen_func
+
+# Import flash_attn's fused ops
+try:
+ from flash_attn.ops.rms_norm import rms_norm as rms_norm_impl
+except ImportError:
+ def rms_norm_impl(x, weight, epsilon):
+ return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(epsilon))) * weight
+
+def merge_states(states, splits, cfg):
+ """
+ pick key and value states for flash_attn_varlen_func
+ Args:
+ states: list of states
+ splits: list of split sizes
+ cfg: bool, use cfg or not"""
+ if cfg:
+ cond_len, uncond_len = 0, 0
+ cond_states, uncond_states = [], []
+ for stat_, split_ in zip(states, splits):
+ cond, uncond = torch.split(stat_, split_, dim=2)
+ cond_states.append(cond)
+ uncond_states.append(uncond)
+ cond_len += cond.shape[2]
+ uncond_len += uncond.shape[2]
+ return cond_states + uncond_states, [cond_len, uncond_len]
+ else:
+ cond_len = 0
+ for stat_ in states:
+ cond_len += stat_.shape[2]
+ return states, [cond_len]
+
+class FastRMSNorm(nn.Module):
+ def __init__(self, C, eps=1e-6, elementwise_affine=True):
+ super().__init__()
+ self.C = C
+ self.eps = eps
+ self.elementwise_affine = elementwise_affine
+ if self.elementwise_affine:
+ self.weight = nn.Parameter(torch.ones(C))
+ else:
+ self.register_buffer('weight', torch.ones(C))
+
+ def forward(self, x):
+ src_type = x.dtype
+ return rms_norm_impl(x.float(), self.weight, epsilon=self.eps).to(src_type)
+
+ def extra_repr(self) -> str:
+ return f'C={self.C}, eps={self.eps:g}, elementwise_affine={self.elementwise_affine}'
+
+
+class WanLayerNorm(nn.LayerNorm):
+
+ def __init__(self, dim, eps=1e-6, elementwise_affine=False):
+ super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
+
+ def forward(self, x):
+ r"""
+ Args:
+ x(Tensor): Shape [B, L, C]
+ """
+ return super().forward(x.float()).type_as(x)
+
+
+class Qwen3MLP(nn.Module):
+ def __init__(self, hidden_size, intermediate_size):
+ super().__init__()
+ self.hidden_size = hidden_size
+ self.intermediate_size = intermediate_size
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
+ self.act_fn = nn.SiLU()
+
+ def forward(self, x):
+ down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
+ return down_proj
+
+def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
+ """
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
+ """
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
+ if n_rep == 1:
+ return hidden_states
+ hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
+
+class SelfAttention(nn.Module):
+ def __init__(
+ self, embed_dim=768, num_heads=12, num_key_value_heads=-1,
+ use_flex_attn=False, qwen_qkvo_bias=False, **kwargs,
+ ):
+ """
+ :param embed_dim: model's width
+ :param num_heads: num heads of multi-head attention
+ :param proj_drop: always 0 for testing
+ :param tau: always 1
+ :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
+ """
+ super().__init__()
+ assert embed_dim % num_heads == 0
+ assert num_key_value_heads == -1 or num_heads % num_key_value_heads == 0
+
+ self.num_heads, self.head_dim = num_heads, embed_dim // num_heads
+ self.num_key_value_heads = num_key_value_heads if num_key_value_heads > 0 else num_heads
+ self.q_proj = nn.Linear(embed_dim, self.num_heads*self.head_dim, bias=qwen_qkvo_bias)
+ self.k_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
+ self.v_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
+ self.o_proj = nn.Linear(self.num_heads*self.head_dim, embed_dim, bias=qwen_qkvo_bias)
+ self.q_norm = FastRMSNorm(self.head_dim)
+ self.k_norm = FastRMSNorm(self.head_dim)
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
+ self.scale = self.head_dim**-0.5
+
+ self.caching = False # kv caching: only used during inference
+ self.cached_k = {} # kv caching: only used during inference
+ self.cached_v = {} # kv caching: only used during inference
+ self.cached_split_cond_uncond = {} # only used during inference
+
+ self.use_flex_attn = use_flex_attn
+
+ def kv_caching(self, enable: bool): # kv caching: only used during inference
+ self.caching = enable
+ self.cached_k = {}
+ self.cached_v = {}
+ self.cached_split_cond_uncond = {}
+
+ # NOTE: attn_bias_or_two_vector is None during inference
+ 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):
+ # x: fp32
+ B, L, C = x.shape
+ hidden_states = x
+ input_shape = hidden_states.shape[:-1]
+ hidden_shape = (*input_shape, -1, self.head_dim)
+ query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).contiguous()# batch, slen, heads, head_dim
+ key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).contiguous() # batch, slen, num_key_value_heads, head_dim
+ value_states = self.v_proj(hidden_states).view(hidden_shape).contiguous() # batch, slen, num_key_value_heads, head_dim
+
+ if sp_manager.sp_on():
+ # Headnum need to be sharded and L needs to be gathered
+ # [B, H, raw_L/sp, C] --> [B, H/sp, raw_L, C]
+ sdim = 1
+ gdim = 2
+ L = L * sp_manager.get_sp_size()
+ C = C // sp_manager.get_sp_size()
+ query_states = sp_all_to_all(query_states, sdim, gdim)
+ key_states = sp_all_to_all(key_states, sdim, gdim)
+ value_states = sp_all_to_all(value_states, sdim, gdim)
+
+ query_states, key_states = apply_rotary_emb(query_states, key_states, rope2d_freqs_grid)
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
+
+ if attn_bias_or_two_vector is None:
+ # fa4, flash_attn_func input/output should be (batch_size, seqlen, nheads, headdim)
+ from flash_attn.cute import flash_attn_varlen_func
+ attn_output = flash_attn_varlen_func(
+ q = query_states.squeeze(0),
+ k = key_states.squeeze(0),
+ v = value_states.squeeze(0),
+ cu_seqlens_q=cu_seqlens,
+ cu_seqlens_k=cu_seqlens,
+ max_seqlen_q=max_seqlen,
+ max_seqlen_k=max_seqlen,
+ softmax_scale=self.scale,
+ )
+ attn_output = attn_output[0].reshape(B, L, C).contiguous()
+ else:
+ # slow attn
+ 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)
+
+ if sp_manager.sp_on():
+ # [B, raw_L, C/sp] --> [B, raw_L/sp, C]
+ sdim = 1
+ gdim = 2
+ attn_output = sp_all_to_all(attn_output, sdim, gdim)
+
+ attn_output = self.o_proj(attn_output)
+
+ return attn_output
+
+class SelfAttnBlock(nn.Module):
+ def __init__(
+ self, embed_dim, num_heads, num_key_value_heads, mlp_ratio=4.,
+ use_flex_attn=False,
+ qwen_qkvo_bias=False, use_ada_layer_norm=False, **kwargs,
+ ):
+ super(SelfAttnBlock, self).__init__()
+ self.C = embed_dim
+ self.attn = SelfAttention(
+ embed_dim=embed_dim, num_heads=num_heads, num_key_value_heads=num_key_value_heads,
+ use_flex_attn=use_flex_attn, qwen_qkvo_bias=qwen_qkvo_bias, **kwargs,
+ )
+ self.mlp = Qwen3MLP(hidden_size=embed_dim, intermediate_size=round(embed_dim * mlp_ratio / 256) * 256)
+ self.use_ada_layer_norm = use_ada_layer_norm
+ if self.use_ada_layer_norm:
+ self.modulation = nn.Parameter(torch.randn(1, 6, embed_dim) / embed_dim**0.5)
+ self.input_layernorm = WanLayerNorm(embed_dim)
+ self.post_attention_layernorm = WanLayerNorm(embed_dim)
+ else:
+ self.input_layernorm = FastRMSNorm(embed_dim)
+ self.post_attention_layernorm = FastRMSNorm(embed_dim)
+
+ # NOTE: attn_bias_or_two_vector is None during inference
+ 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):
+ # x: [B,L,C]
+ # e0: [B, L, 6, C]
+ if self.use_ada_layer_norm:
+ assert e0.dtype == torch.float32
+ e = e0
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
+ residual = x
+ hidden_states = x
+ hidden_states = self.input_layernorm(hidden_states).float() * (1 + e[1].squeeze(2)) + e[0].squeeze(2)
+ 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)
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ hidden_states = residual + hidden_states * e[2].squeeze(2)
+ # Fully Connected
+ residual = hidden_states
+ hidden_states = self.post_attention_layernorm(hidden_states).float() * (1 + e[4].squeeze(2)) + e[3].squeeze(2)
+ hidden_states = self.mlp(hidden_states)
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ hidden_states = residual + hidden_states * e[5].squeeze(2)
+ else:
+ residual = x
+ hidden_states = x
+ hidden_states = self.input_layernorm(hidden_states)
+ 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)
+ hidden_states = residual + hidden_states
+ # Fully Connected
+ residual = hidden_states
+ hidden_states = self.post_attention_layernorm(hidden_states)
+ hidden_states = self.mlp(hidden_states)
+ hidden_states = residual + hidden_states
+ return hidden_states
diff --git a/grn/models/ema.py b/grn/models/ema.py
new file mode 100644
index 0000000000000000000000000000000000000000..0d6ad39fb319d3b88fade2ad420c07f0c3898410
--- /dev/null
+++ b/grn/models/ema.py
@@ -0,0 +1,23 @@
+import copy
+import torch
+from collections import OrderedDict
+
+
+def get_ema_model(model):
+ ema_model = copy.deepcopy(model)
+ ema_model.eval()
+ for param in ema_model.parameters():
+ param.requires_grad = False
+ return ema_model
+
+@torch.no_grad()
+def update_ema(ema_model, model, decay):
+ """
+ Step the EMA model towards the current model.
+ """
+ ema_params = OrderedDict(ema_model.named_parameters())
+ model_params = OrderedDict(model.named_parameters())
+
+ for name, param in model_params.items():
+ # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
+ ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
diff --git a/grn/models/flex_attn_mask.py b/grn/models/flex_attn_mask.py
new file mode 100644
index 0000000000000000000000000000000000000000..289fe354ef3d8dfb1e1f9356cca120ed9185012f
--- /dev/null
+++ b/grn/models/flex_attn_mask.py
@@ -0,0 +1,67 @@
+from functools import partial
+import torch
+import numpy as np
+import torch.nn as nn
+import torch.nn.functional as F
+from torch.nn.attention.flex_attention import flex_attention, create_block_mask
+
+
+def _length_to_offsets(lengths, device):
+ offsets = [0]
+ offsets.extend(lengths)
+ offsets = torch.tensor(offsets, device=device, dtype=torch.int32)
+ offsets = torch.cumsum(offsets, dim=-1)
+ return offsets
+
+def _offsets_to_doc_ids_tensor(offsets):
+ device = offsets.device
+ counts = offsets[1:] - offsets[:-1]
+ visual = torch.repeat_interleave(torch.arange(len(counts), device=device, dtype=torch.int32), counts)
+ return visual
+
+def _generate_overall_mask(offsets, querysid_refsid):
+ document_id = _offsets_to_doc_ids_tensor(offsets) # to scale_ind
+ def overall_mask(b, h, q_idx, kv_idx):
+ querysid = document_id[q_idx]
+ kv_sid = document_id[kv_idx]
+ return querysid_refsid[querysid][kv_sid]
+ return overall_mask
+
+def causal(b, h, q_idx, kv_idx):
+ return q_idx >= kv_idx
+
+def build_flex_attn_func(
+ flex_attention,
+ seq_l,
+ prefix_lens,
+ args,
+ device,
+ batch_size,
+ heads,
+ pad_seq_len,
+ sequece_packing_scales,
+ super_scale_lengths,
+ super_querysid_super_refsid,
+):
+ """
+ Build a flex attn function for a given scale schedule.
+ Args:
+ flex_attention: compiled flex attention
+ seq_l: seq length
+ prefix_lens: valid text prefix lens, [bs]
+ args: arguments
+ device: device
+ batch_size: batch size
+ heads: heads
+ pad_seq_len: pad_seq_len
+ sequece_packing_scales: list of scale schedule
+ querysid_refsid: list of scale_pack_info
+ Returns:
+ attn_fn: flex attn function
+ """
+ assert sum(super_scale_lengths) == seq_l, f'{sum(super_scale_lengths)}!= {seq_l}'
+ offsets = _length_to_offsets(super_scale_lengths, device=device)
+ mask_mod = _generate_overall_mask(offsets, super_querysid_super_refsid)
+ block_mask = create_block_mask(mask_mod, B = batch_size, H = heads, Q_LEN = seq_l, KV_LEN = seq_l, device = device, _compile = True)
+ attn_fn = partial(flex_attention, block_mask=block_mask)
+ return attn_fn
diff --git a/grn/models/fused_op.py b/grn/models/fused_op.py
new file mode 100644
index 0000000000000000000000000000000000000000..94b5a3091d8d252301a2361425ebfa7373708f13
--- /dev/null
+++ b/grn/models/fused_op.py
@@ -0,0 +1,27 @@
+import gc
+from copy import deepcopy
+from typing import Union
+
+import torch
+from torch import nn as nn
+from torch.nn import functional as F
+
+
+@torch.compile(fullgraph=True)
+def fused_rms_norm(x: torch.Tensor, weight: nn.Parameter, eps: float):
+ x = x.float()
+ return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps))) * weight
+
+
+@torch.compile(fullgraph=True)
+def fused_ada_layer_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor):
+ x = x.float()
+ x = F.layer_norm(input=x, normalized_shape=(C,), weight=None, bias=None, eps=eps)
+ return x.mul(scale.add(1)).add_(shift)
+
+
+@torch.compile(fullgraph=True)
+def fused_ada_rms_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor):
+ x = x.float()
+ x = (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps)))
+ return x.mul(scale.add(1)).add_(shift)
diff --git a/grn/models/grn.py b/grn/models/grn.py
new file mode 100644
index 0000000000000000000000000000000000000000..90a692328f53d9bb6e61de479ca801a0d3719167
--- /dev/null
+++ b/grn/models/grn.py
@@ -0,0 +1,754 @@
+import json
+import math
+import time
+from contextlib import nullcontext
+from functools import partial
+from typing import Any, Dict, List, Optional, Tuple, Union
+
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.utils.checkpoint
+import tqdm
+from timm.models import register_model
+
+import grn.utils_t2iv.dist as dist
+from grn.models.basic import FastRMSNorm, SelfAttnBlock
+from grn.models.rope import precompute_rope3d_freqs_grid
+from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta
+from grn.utils_t2iv.dist import for_visualize
+from grn.utils_t2iv.hbq_util_t2iv import multiclass_labels2onehot_input
+from grn.utils_t2iv.sequence_parallel import SequenceParallelManager as sp_manager
+from grn.utils_t2iv.sequence_parallel import sp_gather_sequence_by_dim, sp_split_sequence_by_dim
+
+
+class MultipleLayers(nn.Module):
+ """A sequential container for a chunk of multiple transformer blocks."""
+
+ def __init__(self, layers: List[nn.Module], num_blocks: int, start_index: int):
+ super().__init__()
+ self.module = nn.ModuleList([
+ layers[i] for i in range(start_index, start_index + num_blocks)
+ ])
+
+ def forward(
+ self, x, cu_seqlens, max_seqlen, e0: Optional[torch.Tensor],
+ attn_bias_or_two_vector: Optional[Any], attn_fn: Optional[Any] = None,
+ checkpointing_full_block: bool = False, rope2d_freqs_grid: Optional[torch.Tensor] = None,
+ scale_ind: Optional[Any] = None, context_info: Optional[Any] = None,
+ last_diffusion_step: bool = True, ref_text_scale_inds: Optional[List[Any]] = None,
+ use_cfg: bool = False, split_cond_uncond: Optional[List[Any]] = None
+ ) -> torch.Tensor:
+ h = x
+ for m in self.module:
+ if checkpointing_full_block:
+ h = torch.utils.checkpoint.checkpoint(
+ m, h, cu_seqlens, max_seqlen, e0, 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, use_reentrant=False
+ )
+ else:
+ h = m(
+ h, cu_seqlens, max_seqlen, e0, 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
+ )
+ return h
+
+
+def sinusoidal_embedding_1d(dim: int, position: torch.Tensor) -> torch.Tensor:
+ """
+ Generate 1D sinusoidal embeddings.
+
+ Args:
+ dim (int): Embedding dimension (must be even).
+ position (torch.Tensor): Position tensor of shape [B, L].
+
+ Returns:
+ torch.Tensor: Embeddings of shape [B, L, dim].
+ """
+ if dim % 2 != 0:
+ raise ValueError(f"Embedding dimension must be even, got {dim}")
+
+ half = dim // 2
+ b, l = position.shape
+ position = position.reshape(-1).type(torch.float64)
+
+ sinusoid = torch.outer(
+ position,
+ torch.pow(10000, -torch.arange(half).to(position).div(half))
+ )
+ x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
+ return x.reshape(b, l, dim)
+
+
+class TimestepEmbedder(nn.Module):
+ """Embeds scalar timesteps into vector representations."""
+
+ def __init__(self, hidden_size: int, frequency_embedding_size: int = 256):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
+ nn.SiLU(),
+ nn.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.frequency_embedding_size = frequency_embedding_size
+
+ @staticmethod
+ def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor:
+ """Create sinusoidal timestep embeddings."""
+ half = dim // 2
+ freqs = torch.exp(
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
+ ).to(device=t.device)
+ args = t[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ return embedding
+
+ def forward(self, t: torch.Tensor) -> torch.Tensor:
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
+ return self.mlp(t_freq)
+
+
+def bld_to_bthwd(item: torch.Tensor, patch_time: int, patch_height: int, patch_width: int, apply_spatial_patchify: bool = False) -> torch.Tensor:
+ """Reshape a sequence tensor to a spatial tensor."""
+ batch_size = item.shape[0]
+ return item.reshape(batch_size, patch_time, patch_height, patch_width, -1)
+
+
+def build_attn_mask(seqlens, device):
+ attn_mask = torch.zeros((1, 1, sum(seqlens), sum(seqlens)), dtype=torch.bool, device=device)
+ q_start = 0
+ for i in range(len(seqlens)):
+ q_len = seqlens[i]
+ q_end = q_start + q_len
+ attn_mask[:, :, q_start:q_end, q_start:q_end] = True
+ q_start = q_end
+ return attn_mask
+
+
+class FsqHead(nn.Module):
+ """Classification head for Finite Scalar Quantization (FSQ)."""
+
+ def __init__(self, hidden_dim: int, fsq_dim: int, fsq_lvl: int, use_ada_layer_norm: bool, eps: float = 1e-6):
+ super().__init__()
+ self.proj = nn.Linear(hidden_dim, fsq_dim * fsq_lvl)
+ self.norm = FastRMSNorm(hidden_dim)
+
+ def forward(self, x: torch.Tensor, e: Optional[torch.Tensor] = None) -> torch.Tensor:
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ return self.proj(self.norm(x))
+
+
+class GRN(nn.Module):
+ def __init__(
+ self,
+ vae_local: Any,
+ arch: str = 'var',
+ qwen_qkvo_bias: bool = False,
+ text_channels: int = 0,
+ text_maxlen: int = 0,
+ embed_dim: int = 1024,
+ depth: int = 16,
+ num_key_value_heads: int = -1,
+ num_heads: int = 16,
+ mlp_ratio: float = 4.0,
+ drop_path_rate: float = 0.0,
+ norm_eps: float = 1e-6,
+ block_chunks: int = 1,
+ checkpointing: Optional[str] = None,
+ pad_to_multiplier: int = 0,
+ use_flex_attn: bool = False,
+ num_of_label_value: int = 2,
+ rope2d_normalized_by_hw: int = 0,
+ pn: Optional[str] = None,
+ video_frames: int = 1,
+ always_training_scales: int = 20,
+ apply_spatial_patchify: int = 0,
+ inference_mode: bool = False,
+ other_args: Optional[Any] = None,
+ **kwargs: Any,
+ ):
+ super().__init__()
+ # 1. Model Configuration
+ self.embed_dim = embed_dim
+ self.depth = depth
+ self.num_heads = num_heads
+ self.arch = arch
+ self.mlp_ratio = mlp_ratio
+ self.norm_eps = norm_eps
+ self.drop_path_rate = drop_path_rate
+ self.use_flex_attn = use_flex_attn
+ self.checkpointing = checkpointing
+ self.inference_mode = inference_mode
+ self.other_args = other_args
+
+ # 2. Embedding & Scale Configuration
+ self.vae_embed_dim = vae_local.codebook_dim
+ self.apply_spatial_patchify = apply_spatial_patchify
+ self.text_channels = text_channels
+ self.text_maxlen = text_maxlen
+ self.is_text_to_image = text_channels != 0
+
+ classifier_head_dim = other_args.detail_scale_dim
+ classifier_head_lvl = other_args.detail_num_lvl
+ hbq_round = other_args.hbq_round
+
+ if other_args.refine_mode in ['ar_discrete_GRN_ind']:
+ self.visual_embedding_in_dim = vae_local.codebook_dim * (2**hbq_round)
+ classifier_head_dim = vae_local.codebook_dim
+ elif other_args.refine_mode in ['ar_discrete_GRN_bit']:
+ self.visual_embedding_in_dim = hbq_round * vae_local.codebook_dim * 2
+ classifier_head_dim = hbq_round * vae_local.codebook_dim
+ else:
+ self.visual_embedding_in_dim = vae_local.codebook_dim
+
+ if self.apply_spatial_patchify:
+ self.visual_embedding_in_dim *= 4
+
+ # 3. Dynamic Resolution & Video Specifics
+ self.video_frames = video_frames
+ self.always_training_scales = always_training_scales
+ self.num_of_label_value = num_of_label_value
+ self.rope2d_normalized_by_hw = rope2d_normalized_by_hw
+
+ 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
+ )
+ self.train_h_div_w_list = self.h_div_w_templates
+ print(f"train_h_div_w_list: {self.train_h_div_w_list}")
+
+ # 4. Utilities
+ self.entrophy_statistics = []
+ self.top_p, self.top_k = 1.0, 100
+ self.rng = torch.Generator(device=dist.get_device())
+ self.maybe_record_function = nullcontext
+ self.infer_ts = None
+
+ # 5. Model Components (Projections, Embeddings)
+ self.norm0_cond = nn.Identity()
+ self.text_proj = nn.Linear(self.text_channels, self.embed_dim)
+
+ if self.other_args.use_ada_layer_norm:
+ self.scale_or_time_dim = 256
+ self.scale_or_time_embedding = nn.Sequential(
+ nn.Linear(self.scale_or_time_dim, self.embed_dim), nn.SiLU(), nn.Linear(self.embed_dim, self.embed_dim),
+ )
+ self.scale_or_time_projection = nn.Sequential(nn.SiLU(), nn.Linear(self.embed_dim, self.embed_dim * 6))
+
+ tmp_h_div_w_template = self.train_h_div_w_list[0]
+
+ # RoPE grid initialization
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ self.rope2d_freqs_grid = precompute_rope3d_freqs_grid(
+ dim=self.embed_dim // self.num_heads,
+ rope2d_normalized_by_hw=self.rope2d_normalized_by_hw,
+ activated_h_div_w_templates=self.train_h_div_w_list,
+ max_scales=1010, # never used
+ max_frames=int(self.video_frames / other_args.temporal_compress_rate + 1),
+ max_height=1800 // 8,
+ max_width=1800 // 8,
+ text_maxlen=self.text_maxlen,
+ args=other_args,
+ )
+
+ self.word_embed = nn.Linear(self.visual_embedding_in_dim, self.embed_dim)
+ self.head = FsqHead(
+ hidden_dim=self.embed_dim,
+ fsq_dim=classifier_head_dim,
+ fsq_lvl=classifier_head_lvl,
+ use_ada_layer_norm=other_args.use_ada_layer_norm,
+ )
+
+ if other_args.add_scale_token > 0:
+ self.pt_embedder = TimestepEmbedder(self.embed_dim)
+
+ # 6. Transformer Blocks
+ self.attn_fn_compile_dict = {}
+ self.unregistered_blocks = []
+ for block_idx in range(depth):
+ block = SelfAttnBlock(
+ embed_dim=self.embed_dim,
+ num_heads=num_heads,
+ num_key_value_heads=num_key_value_heads,
+ mlp_ratio=mlp_ratio,
+ use_flex_attn=use_flex_attn,
+ qwen_qkvo_bias=qwen_qkvo_bias,
+ use_ada_layer_norm=other_args.use_ada_layer_norm,
+ )
+ self.unregistered_blocks.append(block)
+
+ self.num_block_chunks = block_chunks or 1
+ self.num_blocks_in_a_chunk = depth // self.num_block_chunks
+ assert self.num_blocks_in_a_chunk * self.num_block_chunks == depth, "Depth must be divisible by block_chunks"
+
+ self.block_chunks = nn.ModuleList([
+ MultipleLayers(self.unregistered_blocks, self.num_blocks_in_a_chunk, i * self.num_blocks_in_a_chunk)
+ for i in range(self.num_block_chunks)
+ ])
+
+ print(f" [Model Config] embed_dim={embed_dim}, num_heads={num_heads}, depth={depth}, "
+ f"mlp_ratio={mlp_ratio}, num_blocks_in_a_chunk={self.num_blocks_in_a_chunk}")
+ print(f" drop_path_rate={drop_path_rate:g}", end='\n\n', flush=True)
+
+ def get_loss_acc(
+ self,
+ hidden_states: torch.Tensor,
+ hidden_states_mask: Optional[torch.Tensor],
+ e: Optional[torch.Tensor],
+ sequence_packing_scales: List[List[Tuple[int, int, int]]],
+ gt: List[torch.Tensor],
+ other_info_by_scale: List[Dict[str, Any]],
+ return_last_hidden_states: bool
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ """
+ Calculate loss and accuracy for the predicted logits.
+
+ Args:
+ hidden_states: shaped (B, L, C)
+ hidden_states_mask: Optional mask for hidden states
+ e: scale or time embeddings
+ sequence_packing_scales: List of scales for sequence packing
+ gt: Ground truth labels
+ other_info_by_scale: Meta information for each scale
+ return_last_hidden_states: Whether to return the last hidden states
+
+ Returns:
+ Tuple of (logits_norm, loss_list, acc_list)
+ """
+ logits_norm = []
+ logits_full = self.head(hidden_states, e)
+ global_token_ptr, global_scale_ptr = 0, 0
+ loss_list, acc_list = [], []
+
+ for pack_scales in sequence_packing_scales:
+ for pt, ph, pw in pack_scales:
+ mul_pt_ph_pw = pt * ph * pw
+ cur_bits = other_info_by_scale[global_scale_ptr]['cur_bits']
+ cur_lvl = other_info_by_scale[global_scale_ptr]['cur_lvl']
+ predict_tokens = other_info_by_scale[global_scale_ptr]['predict_tokens']
+ all_tokens = other_info_by_scale[global_scale_ptr]['all_tokens']
+ logits = logits_full[:, global_token_ptr:global_token_ptr + predict_tokens]
+ logits = logits.reshape(hidden_states.shape[0], mul_pt_ph_pw, cur_bits, cur_lvl)
+ logits = logits.permute(0, 3, 1, 2) # [1, num_of_label_value, mul_pt_ph_pw, d]
+
+ logits_norm.append(logits.abs().mean())
+
+ # gt[global_scale_ptr]: [1, mul_pt_ph_pw, d]
+ loss_this_scale = F.cross_entropy(logits, gt[global_scale_ptr], reduction='none')[0] # [mul_pt_ph_pw, d]
+ acc_this_scale = (logits.argmax(1) == gt[global_scale_ptr]).float()[0] # [mul_pt_ph_pw, d]
+
+ loss_list.append(loss_this_scale.mean(-1))
+ acc_list.append(acc_this_scale.mean(-1))
+
+ global_scale_ptr += 1
+ global_token_ptr += all_tokens
+
+ loss_tensor = torch.cat(loss_list) if loss_list else torch.tensor([], device=hidden_states.device)
+ acc_tensor = torch.cat(acc_list) if acc_list else torch.tensor([], device=hidden_states.device)
+ logits_norm_tensor = torch.stack(logits_norm).mean() if logits_norm else torch.tensor(0.0, device=hidden_states.device)
+
+ return logits_norm_tensor, loss_tensor, acc_tensor
+
+ def get_logits_during_infer(self, hidden_states: torch.Tensor, e: Optional[torch.Tensor] = None) -> torch.Tensor:
+ """Get logits during inference."""
+ return self.head(hidden_states.float(), e)
+
+ def forward(
+ self,
+ label_B_or_BLT: Union[torch.LongTensor, Tuple[torch.FloatTensor, torch.IntTensor, int]],
+ x_BLC: torch.Tensor,
+ visual_rope_cache: Optional[List[torch.Tensor]] = None,
+ sequece_packing_scales: Optional[List[List[Tuple[int, int, int]]]] = None,
+ super_scale_lengths: Optional[List[int]] = None,
+ other_info_by_scale: Optional[List[Dict[str, Any]]] = None,
+ gt_BL: Optional[List[torch.Tensor]] = None,
+ x_BLC_mask: Optional[torch.Tensor] = None,
+ scale_or_time_ids: Optional[torch.Tensor] = None,
+ return_last_hidden_states: bool = False,
+ **kwargs: Any,
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, float]:
+ """
+ Forward pass for the GRN model.
+
+ Args:
+ label_B_or_BLT: Text conditions or labels
+ x_BLC: Input sequence hidden states
+ visual_rope_cache: Cache for visual RoPE embeddings
+ sequece_packing_scales: Scales for sequence packing
+ super_scale_lengths: Lengths of super scales
+ other_info_by_scale: Meta info for scales
+ gt_BL: Ground truth
+ x_BLC_mask: Mask for input sequence
+ scale_or_time_ids: IDs for scale or time embeddings
+ return_last_hidden_states: Whether to return last hidden states
+
+ Returns:
+ Tuple of (logits_norm, loss_list, acc_list, valid_sequence_ratio)
+ """
+ device = x_BLC[0].device
+
+ # [1. get input sequence x_BLC]
+ # word embedding
+ sub_L_list = [item.shape[1] for item in x_BLC]
+ cat_x_BLC = torch.cat(x_BLC, dim=1)
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ cat_x_BLC = self.word_embed(cat_x_BLC.float())
+ x_BLC = list(torch.split(cat_x_BLC, sub_L_list, dim=1))
+
+ # text tokens embedding
+ kv_compact, lens, cu_seqlens_k, max_seqlen_k, _ = label_B_or_BLT
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ kv_compact = self.text_proj(kv_compact).contiguous() # [sum(lens), C]
+ kv_compact_splits = torch.split(kv_compact, lens, dim=0)
+
+ # scale tokens embedding
+ scale_token_ids = torch.tensor([info["scale_token_id"] for info in other_info_by_scale], device=device)
+ with torch.amp.autocast("cuda", dtype=torch.float32):
+ pt_tokens = self.pt_embedder((scale_token_ids)) # [num_scales, C]
+
+ # construct final X_BLC input, [visual token, text token, scale token]
+ x_BLC_lists = []
+ for i in range(len(x_BLC)):
+ x_BLC_lists.extend([x_BLC[i], kv_compact_splits[i].unsqueeze(0), pt_tokens[i][None, None]])
+ x_BLC = torch.cat(x_BLC_lists, dim=1)
+
+ valid_sequence_ratio = x_BLC.shape[1] / self.other_args.train_max_token_len
+ attn_fn, attn_bias_or_two_vector = None, None
+
+ # calculate finalrope cache, [visual token, text token, scale token]
+ self.rope2d_freqs_grid['freqs_text'] = self.rope2d_freqs_grid['freqs_text'].to(x_BLC.device)
+ rope_cache_list = []
+ for i in range(len(visual_rope_cache)):
+ rope_cache_list.append(visual_rope_cache[i])
+ rope_cache_list.append(self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[i]])
+ rope_cache_list.append(self.rope2d_freqs_grid['freqs_text'][:,:,:,:,512:512+self.other_args.add_scale_token])
+ rope_cache = torch.cat(rope_cache_list, dim=4) # (2, 1, 1, 1, seq_len, head_dim / 2)
+ assert rope_cache.shape[4] == x_BLC.shape[1], f'{rope_cache.shape[4]} != {x_BLC.shape[1]}'
+ 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)
+
+ # calculate time or scale embeddings
+ if self.other_args.use_ada_layer_norm:
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ 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]
+ if e.shape[1] < x_BLC.shape[1]:
+ e = F.pad(e, (0,0,0,x_BLC.shape[1]-e.shape[1]), 'constant', 0.) # [1, visual_seq_len, C] -> [1, L, C]
+ e0 = self.scale_or_time_projection(e).unflatten(2, (6, self.C)) # [1, L, C] -> [1, L, 6C] -> [1, L, 6, C]
+ assert e.dtype == torch.float32 and e0.dtype == torch.float32
+ else:
+ e, e0 = None, None
+
+ # [2. block loop]
+ checkpointing_full_block = self.checkpointing == 'full-block' and self.training
+
+ if sp_manager.sp_on():
+ # [B, raw_L, C] --> [B, raw_L/sp_size, C]
+ x_BLC = sp_split_sequence_by_dim(x_BLC, 1)
+
+ cu_seqlens = torch.tensor([0]+super_scale_lengths, device=device).cumsum(-1).to(torch.int32)
+ max_seqlen = max(super_scale_lengths)
+ for i, chunk in enumerate(self.block_chunks): # this path
+ 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)
+
+ if sp_manager.sp_on():
+ # [B, raw_L/sp_size, C] --> [B, raw_L, C]
+ x_BLC = sp_gather_sequence_by_dim(x_BLC, 1)
+
+ # [3. unpad the seqlen dim, and then get logits]
+ 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)
+ return logits_norm, loss_list, acc_list, valid_sequence_ratio
+
+ def prepare_text_conditions(
+ self,
+ label_B_or_BLT: Tuple[torch.Tensor, ...],
+ negative_label_B_or_BLT: Optional[Tuple[torch.Tensor, ...]],
+ use_cfg: bool = False,
+ ) -> Tuple[torch.Tensor, List[int]]:
+ """Prepare text conditions for inference."""
+ kv_compact, lens, cu_seqlens_k, max_seqlen_k = label_B_or_BLT
+ if use_cfg:
+ kv_compact_un, lens_un, cu_seqlens_k_un, max_seqlen_k_un = negative_label_B_or_BLT
+ kv_compact = torch.cat((kv_compact, kv_compact_un), dim=0)
+ cu_seqlens_k = torch.cat((cu_seqlens_k, cu_seqlens_k_un[1:] + cu_seqlens_k[-1]), dim=0)
+ max_seqlen_k = max(max_seqlen_k, max_seqlen_k_un)
+ lens = lens + lens_un
+ kv_compact = self.text_proj(kv_compact).contiguous()
+ return kv_compact, lens
+
+ def embeds_codes2input(self, last_stage: torch.Tensor) -> torch.Tensor:
+ """Embed discrete codes into continuous input representations."""
+ last_stage = last_stage.reshape(*last_stage.shape[:2], -1) # [B, d, t*h*w] or [B, 4d, t*h*w]
+ last_stage = torch.permute(last_stage, [0, 2, 1]) # [B, t*h*w, d] or [B, t*h*w, 4d]
+ last_stage = self.word_embed(last_stage) # norm0_ve is Identity
+ return last_stage
+
+ @torch.no_grad()
+ def autoregressive_infer(
+ self,
+ vae: Optional[Any] = None,
+ scale_schedule: Optional[List[Tuple[int, int, int]]] = None,
+ label_B_or_BLT: Optional[List[Tuple[torch.Tensor, ...]]] = None,
+ negative_label_B_or_BLT: Optional[List[Tuple[torch.Tensor, ...]]] = None,
+ g_seed: Optional[int] = None,
+ cfg_list: Optional[List[float]] = None,
+ tau_list: Optional[List[float]] = None,
+ gt_leak: int = 0,
+ args: Optional[Any] = None,
+ get_visual_rope_embeds: Optional[Any] = None,
+ noise_list: Optional[List[torch.Tensor]] = None,
+ uncond_class_token_id: int = 1000,
+ first_frame_condition: bool = False,
+ **kwargs: Any,
+ ):
+ """Autoregressive inference loop for the GRN model."""
+ if cfg_list is None: cfg_list = []
+ if tau_list is None: tau_list = []
+
+ from grn.schedules.global_refine import shift_pt
+
+ rng = None
+ assert len(cfg_list) >= len(scale_schedule), "Not enough CFG values for scales"
+ assert len(tau_list) >= len(scale_schedule), "Not enough tau values for scales"
+
+ ret, idx_Bl_list = [], [] # current length, list of reconstructed images
+ for b in self.unregistered_blocks: b.attn.kv_caching(True)
+ total_steps = args.max_infer_steps
+ pbar = tqdm.tqdm(total=total_steps)
+ block_chunks = self.block_chunks if self.num_block_chunks > 1 else self.blocks
+ use_cfg = True
+ cfg_interval = float(args.cfg_type.split('_')[-1])
+ full_pt, ph, pw = scale_schedule[0]
+ if first_frame_condition:
+ pt = full_pt - 1
+ visual_rope_cache = get_visual_rope_embeds(self.rope2d_freqs_grid, (pt, ph, pw), 'cuda', args.mapped_h_div_w_template, t_offset=1)
+ else:
+ pt = full_pt
+ visual_rope_cache = get_visual_rope_embeds(self.rope2d_freqs_grid, (pt, ph, pw), 'cuda', args.mapped_h_div_w_template, t_offset=0)
+
+ # text tokens forward
+ self.rope2d_freqs_grid['freqs_text'] = self.rope2d_freqs_grid['freqs_text'].to('cuda')
+ prefix_tokens, lens = self.prepare_text_conditions(label_B_or_BLT[0], negative_label_B_or_BLT, use_cfg)
+ device = prefix_tokens.device
+ infer_device, infer_dtype = prefix_tokens.device, prefix_tokens.dtype
+ prefix_tokens = torch.split(prefix_tokens, lens, dim=0)
+ rope_cache_text_cond = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[0]]
+ rope_cache_text_uncond = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,:lens[1]]
+
+ if args.refine_mode in ['ar_discrete_GRN_bit']:
+ classes = 2
+ labels_shape = (1,args.detail_scale_dim*args.hbq_round,pt,ph,pw)
+ elif args.refine_mode in ['ar_discrete_GRN_index']:
+ classes = 2**args.hbq_round
+ labels_shape = (1,args.detail_scale_dim,pt,ph,pw)
+
+ mul_pt_ph_pw = pt * ph * pw
+ repeat_idx = -1
+ scale_token_rope_cache = self.rope2d_freqs_grid['freqs_text'][:,:,:,:,512:512+args.add_scale_token]
+ if noise_list is not None:
+ absolute_gt_labels = noise_list[0].to('cuda').permute(0,2,3,4,1) # [B,d,t,h,w] -> [B,t,h,w,d]
+ assert len(scale_schedule) == 1
+ if first_frame_condition:
+ first_frame_labels = noise_list[0][:,:,:1] # [B,d,1,h,w]
+ first_frame_tokens_cond = self.embeds_codes2input(multiclass_labels2onehot_input(first_frame_labels, classes))
+ 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)
+ visual_rope_cache = torch.cat((visual_rope_cache, fist_frame_rope_cache), dim=4)
+ 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]
+ else:
+ tmp_seqlens = [mul_pt_ph_pw+lens[0]+args.add_scale_token, mul_pt_ph_pw+lens[1]+args.add_scale_token]
+
+ # [visual tokens, text tokens, pt tokens]
+ 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)
+ 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)
+
+ cu_seqlens = torch.tensor([0]+tmp_seqlens, device=device).cumsum(-1).to(torch.int32)
+ max_seqlen = max(tmp_seqlens)
+
+ pure_rand_labels = torch.randint(low=0, high=classes, size=labels_shape, device=infer_device, dtype=infer_dtype)
+ mixed_xt = pure_rand_labels
+ next_pt = 0.
+ attn_mask = build_attn_mask(tmp_seqlens, device) if args.use_slow_attn else None
+ for cur_inner_round_si in range(args.max_infer_steps):
+ cur_pt = next_pt
+ is_last_step = np.abs(cur_pt - 1) < 0.02
+ if cur_inner_round_si == 0:
+ self.entrophy_statistics.append([])
+ repeat_idx += 1 # index scale tokens, very important
+ cfg = cfg_list[0] if cur_pt >= cfg_interval else 1.0
+ last_stage = self.embeds_codes2input(multiclass_labels2onehot_input(mixed_xt, classes))
+ pt_tokens = self.pt_embedder(torch.tensor([cur_pt], device=device)).unsqueeze(0)
+ # [visual tokens, text tokens, pt tokens]
+ if first_frame_condition:
+ last_stage_cond = torch.cat((last_stage, first_frame_tokens_cond, prefix_tokens[0].unsqueeze(0), pt_tokens), dim=1)
+ last_stage_uncond = torch.cat((last_stage, first_frame_tokens_cond, prefix_tokens[1].unsqueeze(0), pt_tokens), dim=1)
+ else:
+ last_stage_cond = torch.cat((last_stage, prefix_tokens[0].unsqueeze(0), pt_tokens), dim=1)
+ last_stage_uncond = torch.cat((last_stage, prefix_tokens[1].unsqueeze(0), pt_tokens), dim=1)
+ last_stage = torch.cat([last_stage_cond, last_stage_uncond], dim=1)
+
+ e, e0 = None, None
+ last_diffusion_step = False
+ for block_idx, b in enumerate(block_chunks):
+ 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)
+ logits = self.get_logits_during_infer(last_stage, e=e)
+ tmp_bs, tmp_seq_len = logits.shape[:2]
+ logits = logits.reshape(tmp_bs, tmp_seq_len, -1, args.detail_num_lvl) # [B,thw+...,d,2]
+ pred_cond_logits = logits[:,:mul_pt_ph_pw] # [B,thw,d,2]
+ pred_uncond_logits = logits[:,tmp_seqlens[0]:tmp_seqlens[0]+mul_pt_ph_pw] # [B,thw,d,2]
+ pred_cond_probs = pred_cond_logits.softmax(-1) # [B,thw,d,2]
+ categories = pred_cond_logits.shape[-1]
+ entrophy = (-pred_cond_probs * torch.log2(pred_cond_probs)).sum(-1).mean().item() / np.log2(categories)
+
+ pt_unshift = (cur_inner_round_si + 1) / (args.complexity_aware_Tmax - 1)
+ pt_shift = shift_pt(min(1., pt_unshift), args.snr_shift)
+ next_pt = 1 - np.cos(np.pi/2*pt_shift)
+ next_pt = next_pt * 0.999
+
+ pred_cond_labels = torch.argmax(pred_cond_probs, dim=-1) # [B,thw,d]
+ pred_cond_labels = bld_to_bthwd(pred_cond_labels, pt, ph, pw)
+ if cfg != 1:
+ pred_cfg_logits = pred_uncond_logits + cfg * (pred_cond_logits - pred_uncond_logits)
+ else:
+ pred_cfg_logits = pred_cond_logits
+ pred_cfg_logits = pred_cfg_logits.mul(1/tau_list[0]) # [B,thw,d,2]
+ pred_cfg_probs = pred_cfg_logits.softmax(dim=-1) # [B,thw,d,2]
+ pred_cfg_labels = torch.argmax(pred_cfg_probs, dim=-1) # [B,thw,d]
+ pred_cfg_labels = bld_to_bthwd(pred_cfg_labels, pt, ph, pw) # [B,t,h,w,d]
+ 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]
+ pred_sample_probs = torch.gather(pred_cfg_probs, dim=3, index=pred_sample_labels.unsqueeze(-1)).squeeze(-1) # [B,thw,d]
+ pred_sample_probs = bld_to_bthwd(pred_sample_probs, pt, ph, pw) # [B,t,h,w,d]
+ pred_sample_labels = bld_to_bthwd(pred_sample_labels, pt, ph, pw) # [B,t,h,w,d]
+
+ assume_flip_ratio = (1 - cur_pt) / args.detail_num_lvl * 100. # different ratio between prediciton and input
+ pred_zero_ratio = (pred_cond_labels == 0).sum() / pred_cond_labels.numel() * 100.
+ pred_one_ratio = (pred_cond_labels == 1).sum() / pred_cond_labels.numel() * 100.
+ mixed_xt_Bthwd_01 = mixed_xt.clone().permute(0,2,3,4,1)
+ mixed_xt_Bthwd_01[mixed_xt_Bthwd_01<0] = 0
+ pred_cond_flip_ratio = (pred_cond_labels != mixed_xt_Bthwd_01).sum() / pred_cond_labels.numel() * 100.
+ pred_cfg_flip_ratio = (pred_cfg_labels != mixed_xt_Bthwd_01).sum() / pred_cfg_labels.numel() * 100.
+ pred_sample_flip_ratio = (pred_sample_labels != mixed_xt_Bthwd_01).sum() / pred_sample_labels.numel() * 100.
+ self.entrophy_statistics[-1].append({
+ 'cur_inner_round_si': cur_inner_round_si,
+ 'cur_pt': cur_pt,
+ # 'cur_tau': cur_tau,
+ # 'cur_cfg': cur_cfg,
+ 'entrophy': entrophy,
+ 'assume_flip_ratio': assume_flip_ratio,
+ 'pred_cond_flip_ratio': pred_cond_flip_ratio.item(),
+ 'pred_cfg_flip_ratio': pred_cfg_flip_ratio.item(),
+ 'pred_sample_flip_ratio': pred_sample_flip_ratio.item(),
+ 'pred_zero_ratio': pred_zero_ratio.item(),
+ 'pred_one_ratio': pred_one_ratio.item(),
+ 'meta': args.meta,
+ })
+ print(f'{repeat_idx=} {cur_inner_round_si=} {cur_pt=:.3f} {pred_sample_labels.shape=}')
+ print(f'{assume_flip_ratio=:.2f}% {pred_cond_flip_ratio=:.2f}% {pred_cfg_flip_ratio=:.2f}% {pred_sample_flip_ratio=:.2f}%')
+ if repeat_idx < gt_leak:
+ gt_labels = absolute_gt_labels
+ gt_flip_ratio = (gt_labels != mixed_xt_Bthwd_01).sum() / gt_labels.numel() * 100.
+ gt_flip_ratio = gt_flip_ratio.item()
+ pred_cond_acc = (gt_labels==pred_cond_labels).to(float).mean().item()
+ pred_cfg_acc = (gt_labels==pred_cfg_labels).to(float).mean().item()
+ pred_sample_acc = (gt_labels==pred_sample_labels).to(float).mean().item()
+ print(f'{repeat_idx=} {entrophy=:.4f} {pred_cond_acc=:.4f} {pred_cfg_acc=:.4f} {pred_sample_acc=:.4f}')
+ self.entrophy_statistics[-1][-1].update({
+ 'gt_flip_ratio': gt_flip_ratio,
+ 'pred_cond_acc': pred_cond_acc,
+ 'pred_cfg_acc': pred_cfg_acc,
+ 'pred_sample_acc': pred_sample_acc,
+ })
+ pred_sample_labels = gt_labels
+
+ pred_sample_labels = pred_sample_labels.permute(0,4,1,2,3) # [B,t,h,w,d] -> [B,d,t,h,w]
+ pred_sample_probs = pred_sample_probs.permute(0,4,1,2,3) # [B,t,h,w,d] -> [B,d,t,h,w]
+ use_predict_mask = torch.rand(pred_sample_labels.shape, device=device) < next_pt
+ mixed_xt = torch.where(use_predict_mask, pred_sample_labels, pure_rand_labels)
+ next_pt = use_predict_mask.float().mean().item()
+ pbar.update(1)
+ if is_last_step: break
+
+ if first_frame_condition:
+ pred_sample_labels = torch.cat((first_frame_labels, pred_sample_labels), dim=2)
+
+ if args.refine_mode == 'ar_discrete_GRN_ind':
+ from grn.utils_t2iv.hbq_util_t2iv import index_label2quant_features
+ approx_signal = index_label2quant_features(pred_sample_labels, hbq_round=args.hbq_round)
+ elif args.refine_mode == 'ar_discrete_GRN_bit':
+ from grn.utils_t2iv.hbq_util_t2iv import bit_label2raw_feature
+ 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]
+ for b in self.unregistered_blocks: b.attn.kv_caching(False)
+ img = self.summed_codes2images(vae, approx_signal)
+ return ret, idx_Bl_list, img
+
+ def summed_codes2images(self, vae: Any, summed_codes: torch.Tensor) -> torch.Tensor:
+ """Decode summed codes into images using the VAE."""
+ t1 = time.time()
+ img = vae.decode(summed_codes, slice=True)
+ img = (img + 1) / 2
+ img = torch.clamp(img, 0, 1)
+ img = img.permute(0, 2, 3, 4, 1) # [bs, 3, t, h, w] -> [bs, t, h, w, 3]
+ img = img.mul_(255).to(torch.uint8).flip(dims=(4,))
+ print(f"Decode takes {time.time() - t1:.1f}s")
+ return img # bgr order
+
+ @for_visualize
+ def vis_key_params(self, ep: int) -> None:
+ return
+
+ def load_state_dict(self, state_dict: Dict[str, Any], strict: bool = False, assign: bool = False) -> Any:
+ return super().load_state_dict(state_dict=state_dict, strict=strict, assign=assign)
+
+ def special_init(self, **kwargs: Any) -> None:
+ """Apply special initialization to specific layers."""
+ std = 0.02
+ for name, module in self.named_modules():
+ if isinstance(module, nn.Linear):
+ module.weight.data.normal_(mean=0.0, std=std)
+ if module.bias is not None:
+ module.bias.data.zero_()
+ elif isinstance(module, nn.Embedding):
+ module.weight.data.normal_(mean=0.0, std=std)
+ if module.padding_idx is not None:
+ module.weight.data[module.padding_idx].zero_()
+
+ def extra_repr(self) -> str:
+ return f'drop_path_rate={self.drop_path_rate}'
+
+ def get_layer_id_and_scale_exp(self, para_name: str) -> Any:
+ raise NotImplementedError
+
+TIMM_KEYS = {'img_size', 'pretrained', 'pretrained_cfg', 'pretrained_cfg_overlay', 'global_pool'}
+
+@register_model
+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:
+ return GRN(
+ arch='qwen',
+ qwen_qkvo_bias=False,
+ depth=depth,
+ block_chunks=block_chunks,
+ embed_dim=embed_dim,
+ num_heads=num_heads,
+ num_key_value_heads=num_key_value_heads,
+ mlp_ratio=3.55,
+ drop_path_rate=drop_path_rate,
+ **{k: v for k, v in kwargs.items() if k not in TIMM_KEYS}
+ )
+
+@register_model
+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:
+ return GRN(
+ arch='qwen',
+ qwen_qkvo_bias=False,
+ depth=depth,
+ block_chunks=block_chunks,
+ embed_dim=embed_dim,
+ num_heads=num_heads,
+ num_key_value_heads=num_key_value_heads,
+ mlp_ratio=3.55,
+ drop_path_rate=drop_path_rate,
+ **{k: v for k, v in kwargs.items() if k not in TIMM_KEYS}
+ )
\ No newline at end of file
diff --git a/grn/models/grn_c2i.py b/grn/models/grn_c2i.py
new file mode 100644
index 0000000000000000000000000000000000000000..9d0232bd948702701a9f1214075e98d327a908b0
--- /dev/null
+++ b/grn/models/grn_c2i.py
@@ -0,0 +1,399 @@
+# --------------------------------------------------------
+# References:
+# SiT: https://github.com/willisma/SiT
+# Lightning-DiT: https://github.com/hustvl/LightningDiT
+# --------------------------------------------------------
+import torch
+import torch.nn as nn
+import numpy as np
+import math
+import torch.nn.functional as F
+from grn.utils_c2i.model_util import VisionRotaryEmbeddingFast, get_2d_sincos_pos_embed, RMSNorm
+from grn.utils_c2i.hbq_util_c2i import multiclass_labels2onehot_input
+
+def modulate(x, shift, scale):
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
+
+
+class BottleneckPatchEmbed(nn.Module):
+ """ Image to Patch Embedding
+ """
+ def __init__(self, img_size=224, patch_size=16, in_chans=3, pca_dim=768, embed_dim=768, bias=True):
+ super().__init__()
+ img_size = (img_size, img_size)
+ patch_size = (patch_size, patch_size)
+ num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])
+ self.img_size = img_size
+ self.patch_size = patch_size
+ self.num_patches = num_patches
+
+ self.proj1 = nn.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
+ self.proj2 = nn.Conv2d(pca_dim, embed_dim, kernel_size=1, stride=1, bias=bias)
+
+ def forward(self, x):
+ B, C, H, W = x.shape
+ assert H == self.img_size[0] and W == self.img_size[1], \
+ f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})."
+ x = self.proj2(self.proj1(x)).flatten(2).transpose(1, 2)
+ return x
+
+
+class TimestepEmbedder(nn.Module):
+ """
+ Embeds scalar timesteps into vector representations.
+ """
+ def __init__(self, hidden_size, frequency_embedding_size=256):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
+ nn.SiLU(),
+ nn.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.frequency_embedding_size = frequency_embedding_size
+
+ @staticmethod
+ def timestep_embedding(t, dim, max_period=10000):
+ """
+ Create sinusoidal timestep embeddings.
+ :param t: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param dim: the dimension of the output.
+ :param max_period: controls the minimum frequency of the embeddings.
+ :return: an (N, D) Tensor of positional embeddings.
+ """
+ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
+ half = dim // 2
+ freqs = torch.exp(
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
+ ).to(device=t.device)
+ args = t[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ return embedding
+
+ def forward(self, t):
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
+ t_emb = self.mlp(t_freq)
+ return t_emb
+
+
+class LabelEmbedder(nn.Module):
+ """
+ Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
+ """
+ def __init__(self, num_classes, hidden_size):
+ super().__init__()
+ self.embedding_table = nn.Embedding(num_classes + 1, hidden_size)
+ self.num_classes = num_classes
+
+ def forward(self, labels):
+ embeddings = self.embedding_table(labels)
+ return embeddings
+
+
+def scaled_dot_product_attention(query, key, value, dropout_p=0.0) -> torch.Tensor:
+ L, S = query.size(-2), key.size(-2)
+ scale_factor = 1 / math.sqrt(query.size(-1))
+ attn_bias = torch.zeros(query.size(0), 1, L, S, dtype=query.dtype).cuda()
+
+ with torch.cuda.amp.autocast(enabled=False):
+ attn_weight = query.float() @ key.float().transpose(-2, -1) * scale_factor
+ attn_weight += attn_bias
+ attn_weight = torch.softmax(attn_weight, dim=-1)
+ attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
+ return attn_weight @ value
+
+
+class Attention(nn.Module):
+ def __init__(self, dim, num_heads=8, qkv_bias=True, qk_norm=True, attn_drop=0., proj_drop=0.):
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+
+ self.q_norm = RMSNorm(head_dim) if qk_norm else nn.Identity()
+ self.k_norm = RMSNorm(head_dim) if qk_norm else nn.Identity()
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ def forward(self, x, rope):
+ B, N, C = x.shape
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
+
+ q = self.q_norm(q)
+ k = self.k_norm(k)
+
+ q = rope(q)
+ k = rope(k)
+
+ x = scaled_dot_product_attention(q, k, v, dropout_p=self.attn_drop.p if self.training else 0.)
+
+ x = x.transpose(1, 2).reshape(B, N, C)
+
+ x = self.proj(x)
+ x = self.proj_drop(x)
+ return x
+
+
+class SwiGLUFFN(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ hidden_dim: int,
+ drop=0.0,
+ bias=True
+ ) -> None:
+ super().__init__()
+ hidden_dim = int(hidden_dim * 2 / 3)
+ self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=bias)
+ self.w3 = nn.Linear(hidden_dim, dim, bias=bias)
+ self.ffn_dropout = nn.Dropout(drop)
+
+ def forward(self, x):
+ x12 = self.w12(x)
+ x1, x2 = x12.chunk(2, dim=-1)
+ hidden = F.silu(x1) * x2
+ return self.w3(self.ffn_dropout(hidden))
+
+
+class FinalLayer(nn.Module):
+ def __init__(self, hidden_size, patch_size, out_channels):
+ super().__init__()
+ self.norm_final = RMSNorm(hidden_size)
+ self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(hidden_size, 2 * hidden_size, bias=True)
+ )
+
+ @torch.compile
+ def forward(self, x, c):
+ shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
+ x = modulate(self.norm_final(x), shift, scale)
+ x = self.linear(x)
+ return x
+
+
+class GRNblock(nn.Module):
+ def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, attn_drop=0.0, proj_drop=0.0):
+ super().__init__()
+ self.norm1 = RMSNorm(hidden_size, eps=1e-6)
+ self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=True,
+ attn_drop=attn_drop, proj_drop=proj_drop)
+ self.norm2 = RMSNorm(hidden_size, eps=1e-6)
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ self.mlp = SwiGLUFFN(hidden_size, mlp_hidden_dim, drop=proj_drop)
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(hidden_size, 6 * hidden_size, bias=True)
+ )
+
+ @torch.compile
+ def forward(self, x, c, feat_rope=None):
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1)
+ x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa), rope=feat_rope)
+ x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
+ return x
+
+
+class GRN(nn.Module):
+ """
+ GRN image Transformer.
+ """
+ def __init__(
+ self,
+ input_size=256,
+ patch_size=16,
+ in_channels=3,
+ hidden_size=1024,
+ depth=24,
+ num_heads=16,
+ mlp_ratio=4.0,
+ attn_drop=0.0,
+ proj_drop=0.0,
+ num_classes=1000,
+ bottleneck_dim=128,
+ in_context_len=32,
+ in_context_start=8,
+ args=None,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = in_channels
+ self.patch_size = patch_size
+ self.num_heads = num_heads
+ self.hidden_size = hidden_size
+ self.input_size = input_size
+ self.in_context_len = in_context_len
+ self.in_context_start = in_context_start
+ self.num_classes = num_classes
+ self.args=args
+
+ # time and class embed
+ self.t_embedder = TimestepEmbedder(hidden_size)
+ self.y_embedder = LabelEmbedder(num_classes, hidden_size)
+
+ # linear embed
+ self.x_embedder = BottleneckPatchEmbed(input_size, patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
+
+ # use fixed sin-cos embedding
+ num_patches = self.x_embedder.num_patches
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False)
+
+ # in-context cls token
+ if self.in_context_len > 0:
+ self.in_context_posemb = nn.Parameter(torch.zeros(1, self.in_context_len, hidden_size), requires_grad=True)
+ torch.nn.init.normal_(self.in_context_posemb, std=.02)
+
+ # rope
+ half_head_dim = hidden_size // num_heads // 2
+ hw_seq_len = input_size // patch_size
+ self.feat_rope = VisionRotaryEmbeddingFast(
+ dim=half_head_dim,
+ pt_seq_len=hw_seq_len,
+ num_cls_token=0
+ )
+ self.feat_rope_incontext = VisionRotaryEmbeddingFast(
+ dim=half_head_dim,
+ pt_seq_len=hw_seq_len,
+ num_cls_token=self.in_context_len
+ )
+
+ # transformer
+ self.blocks = nn.ModuleList([
+ GRNblock(hidden_size, num_heads, mlp_ratio=mlp_ratio,
+ attn_drop=attn_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0,
+ proj_drop=proj_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0)
+ for i in range(depth)
+ ])
+
+ # linear predict
+ self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
+
+ self.initialize_weights()
+
+ def initialize_weights(self):
+ # Initialize transformer layers:
+ def _basic_init(module):
+ if isinstance(module, nn.Linear):
+ torch.nn.init.xavier_uniform_(module.weight)
+ if module.bias is not None:
+ nn.init.constant_(module.bias, 0)
+ self.apply(_basic_init)
+
+ # Initialize (and freeze) pos_embed by sin-cos embedding:
+ pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches ** 0.5))
+ self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
+
+ # Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
+ w1 = self.x_embedder.proj1.weight.data
+ nn.init.xavier_uniform_(w1.view([w1.shape[0], -1]))
+ w2 = self.x_embedder.proj2.weight.data
+ nn.init.xavier_uniform_(w2.view([w2.shape[0], -1]))
+ nn.init.constant_(self.x_embedder.proj2.bias, 0)
+
+ # Initialize label embedding table:
+ nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02)
+
+ nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
+ nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
+
+ # Zero-out adaLN modulation layers:
+ for block in self.blocks:
+ nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
+ nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
+
+ # Zero-out output layers:
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
+
+ nn.init.constant_(self.final_layer.linear.weight, 0)
+ nn.init.constant_(self.final_layer.linear.bias, 0)
+
+ def unpatchify(self, x, p):
+ """
+ x: (N, T, patch_size**2 * C)
+ imgs: (N, H, W, C)
+ """
+ c = self.out_channels
+ h = w = int(x.shape[1] ** 0.5)
+ assert h * w == x.shape[1]
+
+ x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
+ x = torch.einsum('nhwpqc->nchpwq', x)
+ imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
+ return imgs
+
+ def forward(self, x, t, y):
+ """
+ x: (N, C, H, W)
+ t: (N,)
+ y: (N,)
+ """
+ if self.args.method in ['GRN_ind']:
+ x = multiclass_labels2onehot_input(x, 2**self.args.hbq_round)
+ elif self.args.method in ['GRN_bit']:
+ x = multiclass_labels2onehot_input(x, 2)
+
+ # class and time embeddings
+ t_emb = self.t_embedder(t)
+ y_emb = self.y_embedder(y)
+ c = t_emb + y_emb
+
+ # forward
+ x = self.x_embedder(x)
+ x += self.pos_embed
+
+ for i, block in enumerate(self.blocks):
+ # in-context
+ if self.in_context_len > 0 and i == self.in_context_start:
+ in_context_tokens = y_emb.unsqueeze(1).repeat(1, self.in_context_len, 1)
+ in_context_tokens += self.in_context_posemb
+ x = torch.cat([in_context_tokens, x], dim=1)
+ x = block(x, c, self.feat_rope if i < self.in_context_start else self.feat_rope_incontext)
+
+ x = x[:, self.in_context_len:]
+
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ x = self.final_layer(x, c)
+
+ if self.args.method == 'GRN_ind':
+ B, h_mul_w, classes_mul_d = x.shape
+ classes = 2**self.args.hbq_round
+ h = w = int(np.round(math.sqrt(h_mul_w)))
+ output = x.reshape(B, h, w, classes, classes_mul_d//classes) # [B, h, w, classes, d]
+ output = output.permute(0, 3, 4, 1, 2) # [B, classes, d, h, w]
+ elif self.args.method == 'GRN_bit':
+ B, h_mul_w, classes_mul_d = x.shape
+ h = w = int(np.round(math.sqrt(h_mul_w)))
+ output = x.reshape(B, h, w, 2, classes_mul_d//2) # [B, h, w, 2, d]
+ output = output.permute(0, 3, 4, 1, 2) # [B, 2, d, h, w]
+ return output
+
+
+def GRN_B(**kwargs):
+ return GRN(depth=12, hidden_size=768, num_heads=12,
+ bottleneck_dim=128, in_context_len=32, in_context_start=4, patch_size=1, **kwargs)
+
+def GRN_L(**kwargs):
+ return GRN(depth=24, hidden_size=1024, num_heads=16,
+ bottleneck_dim=128, in_context_len=32, in_context_start=8, patch_size=1, **kwargs)
+
+def GRN_H(**kwargs):
+ return GRN(depth=32, hidden_size=1280, num_heads=16,
+ bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=1, **kwargs)
+
+def GRN_G(**kwargs):
+ return GRN(depth=40, hidden_size=1664, num_heads=16,
+ bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=1, **kwargs)
+
+GRN_models = {
+ 'GRN_B': GRN_B,
+ 'GRN_L': GRN_L,
+ 'GRN_H': GRN_H,
+ 'GRN_G': GRN_G,
+}
diff --git a/grn/models/hbq_tokenizer.py b/grn/models/hbq_tokenizer.py
new file mode 100644
index 0000000000000000000000000000000000000000..f356fb783eec3d5e6301eadfb4e463266b8130eb
--- /dev/null
+++ b/grn/models/hbq_tokenizer.py
@@ -0,0 +1,932 @@
+import logging
+import os
+import os.path as osp
+
+import torch
+import torch.cuda.amp as amp
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+from einops import rearrange
+
+
+CACHE_T = 2
+
+
+class CausalConv3d(nn.Conv3d):
+ """
+ Causal 3d convolusion.
+ """
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self._padding = (
+ self.padding[2],
+ self.padding[2],
+ self.padding[1],
+ self.padding[1],
+ 2 * self.padding[0],
+ 0,
+ )
+ self.padding = (0, 0, 0)
+
+ def forward(self, x, cache_x=None):
+ padding = list(self._padding)
+ if cache_x is not None and self._padding[4] > 0:
+ cache_x = cache_x.to(x.device)
+ x = torch.cat([cache_x, x], dim=2)
+ padding[4] -= cache_x.shape[2]
+ x = F.pad(x, padding)
+
+ return super().forward(x)
+
+
+class RMS_norm(nn.Module):
+
+ def __init__(self, dim, channel_first=True, images=True, bias=False):
+ super().__init__()
+ broadcastable_dims = (1, 1, 1) if not images else (1, 1)
+ shape = (dim, *broadcastable_dims) if channel_first else (dim,)
+
+ self.channel_first = channel_first
+ self.scale = dim**0.5
+ self.gamma = nn.Parameter(torch.ones(shape))
+ self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
+
+ def forward(self, x):
+ return (F.normalize(x, dim=(1 if self.channel_first else -1)) *
+ self.scale * self.gamma + self.bias)
+
+
+class Upsample(nn.Upsample):
+
+ def forward(self, x):
+ """
+ Fix bfloat16 support for nearest neighbor interpolation.
+ """
+ return super().forward(x.float()).type_as(x)
+
+
+class Resample(nn.Module):
+
+ def __init__(self, dim, mode):
+ assert mode in (
+ "none",
+ "upsample2d",
+ "upsample3d",
+ "downsample2d",
+ "downsample3d",
+ )
+ super().__init__()
+ self.dim = dim
+ self.mode = mode
+
+ # layers
+ if mode == "upsample2d":
+ self.resample = nn.Sequential(
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
+ nn.Conv2d(dim, dim, 3, padding=1),
+ )
+ elif mode == "upsample3d":
+ self.resample = nn.Sequential(
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
+ nn.Conv2d(dim, dim, 3, padding=1),
+ # nn.Conv2d(dim, dim//2, 3, padding=1)
+ )
+ self.time_conv = CausalConv3d(
+ dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
+ elif mode == "downsample2d":
+ self.resample = nn.Sequential(
+ nn.ZeroPad2d((0, 1, 0, 1)),
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
+ elif mode == "downsample3d":
+ self.resample = nn.Sequential(
+ nn.ZeroPad2d((0, 1, 0, 1)),
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
+ self.time_conv = CausalConv3d(
+ dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
+ else:
+ self.resample = nn.Identity()
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ b, c, t, h, w = x.size()
+ if self.mode == "upsample3d":
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ if feat_cache[idx] is None:
+ feat_cache[idx] = "Rep"
+ feat_idx[0] += 1
+ else:
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
+ feat_cache[idx] != "Rep"):
+ # cache last frame of last two chunk
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
+ feat_cache[idx] == "Rep"):
+ cache_x = torch.cat(
+ [
+ torch.zeros_like(cache_x).to(cache_x.device),
+ cache_x
+ ],
+ dim=2,
+ )
+ if feat_cache[idx] == "Rep":
+ x = self.time_conv(x)
+ else:
+ x = self.time_conv(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ x = x.reshape(b, 2, c, t, h, w)
+ x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
+ 3)
+ x = x.reshape(b, c, t * 2, h, w)
+ t = x.shape[2]
+ x = rearrange(x, "b c t h w -> (b t) c h w")
+ x = self.resample(x)
+ x = rearrange(x, "(b t) c h w -> b c t h w", t=t) # this 4 lines do spatial down / up sample
+
+ if self.mode == "downsample3d":
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ if feat_cache[idx] is None:
+ feat_cache[idx] = x.clone()
+ feat_idx[0] += 1
+ else:
+ cache_x = x[:, :, -1:, :, :].clone()
+ x = self.time_conv(
+ torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ return x
+
+ def init_weight(self, conv):
+ conv_weight = conv.weight.detach().clone()
+ nn.init.zeros_(conv_weight)
+ c1, c2, t, h, w = conv_weight.size()
+ one_matrix = torch.eye(c1, c2)
+ init_matrix = one_matrix
+ nn.init.zeros_(conv_weight)
+ conv_weight.data[:, :, 1, 0, 0] = init_matrix # * 0.5
+ conv.weight = nn.Parameter(conv_weight)
+ nn.init.zeros_(conv.bias.data)
+
+ def init_weight2(self, conv):
+ conv_weight = conv.weight.data.detach().clone()
+ nn.init.zeros_(conv_weight)
+ c1, c2, t, h, w = conv_weight.size()
+ init_matrix = torch.eye(c1 // 2, c2)
+ conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
+ conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
+ conv.weight = nn.Parameter(conv_weight)
+ nn.init.zeros_(conv.bias.data)
+
+
+class ResidualBlock(nn.Module):
+
+ def __init__(self, in_dim, out_dim, dropout=0.0):
+ super().__init__()
+ self.in_dim = in_dim
+ self.out_dim = out_dim
+
+ # layers
+ self.residual = nn.Sequential(
+ RMS_norm(in_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(in_dim, out_dim, 3, padding=1),
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ nn.Dropout(dropout),
+ CausalConv3d(out_dim, out_dim, 3, padding=1),
+ )
+ self.shortcut = (
+ CausalConv3d(in_dim, out_dim, 1)
+ if in_dim != out_dim else nn.Identity())
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ h = self.shortcut(x)
+ for layer in self.residual:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ # cache last frame of last two chunk
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+ return x + h
+
+
+class AttentionBlock(nn.Module):
+ """
+ Causal self-attention with a single head.
+ """
+
+ def __init__(self, dim):
+ super().__init__()
+ self.dim = dim
+
+ # layers
+ self.norm = RMS_norm(dim)
+ self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
+ self.proj = nn.Conv2d(dim, dim, 1)
+
+ # zero out the last layer params
+ nn.init.zeros_(self.proj.weight)
+
+ def forward(self, x):
+ identity = x
+ b, c, t, h, w = x.size()
+ x = rearrange(x, "b c t h w -> (b t) c h w")
+ x = self.norm(x)
+ # compute query, key, value
+ q, k, v = (
+ self.to_qkv(x).reshape(b * t, 1, c * 3,
+ -1).permute(0, 1, 3,
+ 2).contiguous().chunk(3, dim=-1))
+
+ # apply attention
+ x = F.scaled_dot_product_attention(
+ q,
+ k,
+ v,
+ )
+ x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
+
+ # output
+ x = self.proj(x)
+ x = rearrange(x, "(b t) c h w-> b c t h w", t=t)
+ return x + identity
+
+
+def patchify(x, patch_size):
+ if patch_size == 1:
+ return x
+ if x.dim() == 4:
+ x = rearrange(
+ x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
+ elif x.dim() == 5:
+ x = rearrange(
+ x,
+ "b c f (h q) (w r) -> b (c r q) f h w",
+ q=patch_size,
+ r=patch_size,
+ )
+ else:
+ raise ValueError(f"Invalid input shape: {x.shape}")
+
+ return x
+
+
+def unpatchify(x, patch_size):
+ if patch_size == 1:
+ return x
+
+ if x.dim() == 4:
+ x = rearrange(
+ x, "b (c r q) h w -> b c (h q) (w r)", q=patch_size, r=patch_size)
+ elif x.dim() == 5:
+ x = rearrange(
+ x,
+ "b (c r q) f h w -> b c f (h q) (w r)",
+ q=patch_size,
+ r=patch_size,
+ )
+ return x
+
+
+class AvgDown3D(nn.Module):
+
+ def __init__(
+ self,
+ in_channels,
+ out_channels,
+ factor_t,
+ factor_s=1,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+ self.factor_t = factor_t
+ self.factor_s = factor_s
+ self.factor = self.factor_t * self.factor_s * self.factor_s
+
+ assert in_channels * self.factor % out_channels == 0
+ self.group_size = in_channels * self.factor // out_channels
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
+ pad = (0, 0, 0, 0, pad_t, 0)
+ x = F.pad(x, pad)
+ B, C, T, H, W = x.shape
+ x = x.view(
+ B,
+ C,
+ T // self.factor_t,
+ self.factor_t,
+ H // self.factor_s,
+ self.factor_s,
+ W // self.factor_s,
+ self.factor_s,
+ )
+ x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
+ x = x.view(
+ B,
+ C * self.factor,
+ T // self.factor_t,
+ H // self.factor_s,
+ W // self.factor_s,
+ )
+ x = x.view(
+ B,
+ self.out_channels,
+ self.group_size,
+ T // self.factor_t,
+ H // self.factor_s,
+ W // self.factor_s,
+ )
+ x = x.mean(dim=2)
+ return x
+
+
+class DupUp3D(nn.Module):
+
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ factor_t,
+ factor_s=1,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+
+ self.factor_t = factor_t
+ self.factor_s = factor_s
+ self.factor = self.factor_t * self.factor_s * self.factor_s
+
+ assert out_channels * self.factor % in_channels == 0
+ self.repeats = out_channels * self.factor // in_channels
+
+ def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
+ x = x.repeat_interleave(self.repeats, dim=1)
+ x = x.view(
+ x.size(0),
+ self.out_channels,
+ self.factor_t,
+ self.factor_s,
+ self.factor_s,
+ x.size(2),
+ x.size(3),
+ x.size(4),
+ )
+ x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
+ x = x.view(
+ x.size(0),
+ self.out_channels,
+ x.size(2) * self.factor_t,
+ x.size(4) * self.factor_s,
+ x.size(6) * self.factor_s,
+ )
+ if first_chunk:
+ x = x[:, :, self.factor_t - 1:, :, :]
+ return x
+
+
+class Down_ResidualBlock(nn.Module):
+
+ def __init__(self,
+ in_dim,
+ out_dim,
+ dropout,
+ mult,
+ temperal_downsample=False,
+ down_flag=False):
+ super().__init__()
+
+ # Shortcut path with downsample
+ self.avg_shortcut = AvgDown3D(
+ in_dim,
+ out_dim,
+ factor_t=2 if temperal_downsample else 1,
+ factor_s=2 if down_flag else 1,
+ )
+
+ # Main path with residual blocks and downsample
+ downsamples = []
+ for _ in range(mult): # mult=2, two block
+ downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
+ in_dim = out_dim
+
+ # Add the final downsample block
+ if down_flag:
+ mode = "downsample3d" if temperal_downsample else "downsample2d"
+ downsamples.append(Resample(out_dim, mode=mode))
+
+ self.downsamples = nn.Sequential(*downsamples)
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ x_copy = x.clone()
+ for module in self.downsamples:
+ x = module(x, feat_cache, feat_idx)
+
+ return x + self.avg_shortcut(x_copy)
+
+
+class Up_ResidualBlock(nn.Module):
+
+ def __init__(self,
+ in_dim,
+ out_dim,
+ dropout,
+ mult,
+ temperal_upsample=False,
+ up_flag=False):
+ super().__init__()
+ # Shortcut path with upsample
+ if up_flag:
+ self.avg_shortcut = DupUp3D(
+ in_dim,
+ out_dim,
+ factor_t=2 if temperal_upsample else 1,
+ factor_s=2 if up_flag else 1,
+ )
+ else:
+ self.avg_shortcut = None
+
+ # Main path with residual blocks and upsample
+ upsamples = []
+ for _ in range(mult):
+ upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
+ in_dim = out_dim
+
+ # Add the final upsample block
+ if up_flag:
+ mode = "upsample3d" if temperal_upsample else "upsample2d"
+ upsamples.append(Resample(out_dim, mode=mode))
+
+ self.upsamples = nn.Sequential(*upsamples)
+
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
+ x_main = x.clone()
+ for module in self.upsamples:
+ x_main = module(x_main, feat_cache, feat_idx)
+ if self.avg_shortcut is not None:
+ x_shortcut = self.avg_shortcut(x, first_chunk)
+ return x_main + x_shortcut
+ else:
+ return x_main
+
+
+class Encoder3d(nn.Module):
+
+ def __init__(
+ self,
+ dim=128,
+ z_dim=4,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, False],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_downsample = temperal_downsample
+
+ # dimensions
+ dims = [dim * u for u in [1] + dim_mult] # [1,2,4,4] -> [1,1,2,4,4] -> [128,128,256,512,512]
+
+ # init block
+ self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
+
+ # downsample blocks
+ downsamples = []
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
+ t_down_flag = (
+ temperal_downsample[i]
+ if i < len(temperal_downsample) else False)
+ downsamples.append(
+ Down_ResidualBlock(
+ in_dim=in_dim,
+ out_dim=out_dim,
+ dropout=dropout,
+ mult=num_res_blocks,
+ temperal_downsample=t_down_flag,
+ down_flag=i != len(dim_mult) - 1,
+ ))
+ self.downsamples = nn.Sequential(*downsamples)
+
+ # middle blocks
+ self.middle = nn.Sequential(
+ ResidualBlock(out_dim, out_dim, dropout),
+ AttentionBlock(out_dim),
+ ResidualBlock(out_dim, out_dim, dropout),
+ )
+
+ # # output blocks
+ self.head = nn.Sequential(
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(out_dim, z_dim, 3, padding=1),
+ )
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = self.conv1(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = self.conv1(x)
+
+ ## downsamples
+ for layer in self.downsamples:
+ if feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## middle
+ for layer in self.middle:
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## head
+ for layer in self.head:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+
+ return x
+
+
+class Decoder3d(nn.Module):
+
+ def __init__(
+ self,
+ dim=128,
+ z_dim=4,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_upsample=[False, True, True],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_upsample = temperal_upsample
+
+ # dimensions
+ dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
+ # scale = 1.0 / 2**(len(dim_mult) - 2)
+ # init block
+ self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
+
+ # middle blocks
+ self.middle = nn.Sequential(
+ ResidualBlock(dims[0], dims[0], dropout),
+ AttentionBlock(dims[0]),
+ ResidualBlock(dims[0], dims[0], dropout),
+ )
+
+ # upsample blocks
+ upsamples = []
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
+ t_up_flag = temperal_upsample[i] if i < len(
+ temperal_upsample) else False
+ upsamples.append(
+ Up_ResidualBlock(
+ in_dim=in_dim,
+ out_dim=out_dim,
+ dropout=dropout,
+ mult=num_res_blocks + 1,
+ temperal_upsample=t_up_flag,
+ up_flag=i != len(dim_mult) - 1,
+ ))
+ self.upsamples = nn.Sequential(*upsamples)
+
+ # output blocks
+ self.head = nn.Sequential(
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(out_dim, 12, 3, padding=1),
+ )
+
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = self.conv1(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = self.conv1(x)
+
+ for layer in self.middle:
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## upsamples
+ for layer in self.upsamples:
+ if feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx, first_chunk)
+ else:
+ x = layer(x)
+
+ ## head
+ for layer in self.head:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+ return x
+
+
+def count_conv3d(model):
+ count = 0
+ for m in model.modules():
+ if isinstance(m, CausalConv3d):
+ count += 1
+ return count
+
+
+class WanVAE_(nn.Module):
+
+ def __init__(
+ self,
+ dim=160,
+ dec_dim=256,
+ z_dim=16,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, False],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_downsample = temperal_downsample
+ self.temperal_upsample = temperal_downsample[::-1]
+
+ # modules
+ self.encoder = Encoder3d(
+ dim,
+ z_dim * 2,
+ dim_mult,
+ num_res_blocks,
+ attn_scales,
+ self.temperal_downsample,
+ dropout,
+ )
+ self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
+ self.conv2 = CausalConv3d(z_dim, z_dim, 1)
+ self.decoder = Decoder3d(
+ dec_dim,
+ z_dim,
+ dim_mult,
+ num_res_blocks,
+ attn_scales,
+ self.temperal_upsample,
+ dropout,
+ )
+
+ def forward(self, x, scale=[0, 1]):
+ mu = self.encode(x, scale)
+ x_recon = self.decode(mu, scale)
+ return x_recon, mu
+
+ def encode(self, x, scale=None):
+ self.clear_cache()
+ x = patchify(x, patch_size=2)
+ t = x.shape[2]
+ iter_ = 1 + (t - 1) // 4
+ for i in range(iter_):
+ self._enc_conv_idx = [0]
+ if i == 0:
+ out = self.encoder(
+ x[:, :, :1, :, :],
+ feat_cache=self._enc_feat_map,
+ feat_idx=self._enc_conv_idx,
+ )
+ else:
+ out_ = self.encoder(
+ x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
+ feat_cache=self._enc_feat_map,
+ feat_idx=self._enc_conv_idx,
+ )
+ out = torch.cat([out, out_], 2)
+ mu, log_var = self.conv1(out).chunk(2, dim=1)
+ if scale is not None:
+ if isinstance(scale[0], torch.Tensor):
+ mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
+ 1, self.z_dim, 1, 1, 1)
+ else:
+ mu = (mu - scale[0]) * scale[1]
+ self.clear_cache()
+ if self.encoder_out_type == 'feature_tanh':
+ mu = torch.tanh(mu)
+ elif self.encoder_out_type == 'feature':
+ pass
+ else:
+ raise ValueError(f'{self.encoder_out_type=} is not supported!')
+ return mu
+
+ def decode(self, z, scale=None, **kwargs):
+ self.clear_cache()
+ if scale is not None:
+ if isinstance(scale[0], torch.Tensor):
+ z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
+ 1, self.z_dim, 1, 1, 1)
+ else:
+ z = z / scale[1] + scale[0]
+ x = self.conv2(z)
+ iter_ = z.shape[2]
+ for i in range(iter_):
+ self._conv_idx = [0]
+ if i == 0:
+ out = self.decoder(
+ x[:, :, i:i + 1, :, :],
+ feat_cache=self._feat_map,
+ feat_idx=self._conv_idx,
+ first_chunk=True,
+ )
+ else:
+ out_ = self.decoder(
+ x[:, :, i:i + 1, :, :],
+ feat_cache=self._feat_map,
+ feat_idx=self._conv_idx,
+ )
+ out = torch.cat([out, out_], 2)
+ out = unpatchify(out, patch_size=2)
+ self.clear_cache()
+ return out
+
+ def reparameterize(self, mu, log_var):
+ std = torch.exp(0.5 * log_var)
+ eps = torch.randn_like(std)
+ return eps * std + mu
+
+ def sample(self, imgs, deterministic=False):
+ import pdb; pdb.set_trace()
+ mu, log_var = self.encode(imgs)
+ if deterministic:
+ return mu
+ std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
+ return mu + std * torch.randn_like(std)
+
+ def clear_cache(self):
+ self._conv_num = count_conv3d(self.decoder)
+ self._conv_idx = [0]
+ self._feat_map = [None] * self._conv_num
+ # cache encode
+ self._enc_conv_num = count_conv3d(self.encoder)
+ self._enc_conv_idx = [0]
+ self._enc_feat_map = [None] * self._enc_conv_num
+
+class HBQ_Tokenizer(WanVAE_):
+ def __init__(
+ self,
+ args,
+ dim=160,
+ dec_dim=256,
+ latent_channels=16,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ temperal_downsample=[0,1,1],
+ dropout=0.,
+ encoder_out_type='',
+ ):
+ super().__init__(
+ dim=dim,
+ dec_dim=dec_dim,
+ z_dim=latent_channels,
+ dim_mult=dim_mult,
+ num_res_blocks=num_res_blocks,
+ temperal_downsample=temperal_downsample,
+ dropout=dropout,
+ )
+ self.other_args = args
+ self.codebook_dim = latent_channels
+ self.encoder_out_type = encoder_out_type
+
+ def encode_for_raw_features(
+ self, x: torch.Tensor,
+ **kwargs,
+ ):
+ is_image = x.ndim == 4
+ if not is_image:
+ B, C, T, H, W = x.shape
+ else:
+ B, C, H, W = x.shape
+ T = 1
+ x = x.unsqueeze(2)
+ with torch.amp.autocast("cuda", dtype=torch.float):
+ z = self.encode(x)
+ return [z], None, None
+
+def _video_vae(pretrained_path=None, z_dim=16, dim=160, device="cpu", **kwargs):
+ # params
+ cfg = dict(
+ dim=dim,
+ z_dim=z_dim,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, True],
+ dropout=0.0,
+ )
+ cfg.update(**kwargs)
+
+ # init model
+ with torch.device("meta"):
+ model = WanVAE_(**cfg)
+
+ # load checkpoint
+ logging.info(f"loading {pretrained_path}")
+ model.load_state_dict(
+ torch.load(pretrained_path, map_location=device), assign=True)
+
+ return model
+
diff --git a/grn/models/init_param.py b/grn/models/init_param.py
new file mode 100644
index 0000000000000000000000000000000000000000..2854318d5f56fbad6fb6e14b21e603475be6ccb4
--- /dev/null
+++ b/grn/models/init_param.py
@@ -0,0 +1,33 @@
+import torch.nn as nn
+
+
+def init_weights(model: nn.Module, conv_std_or_gain: float = 0.02, other_std: float = 0.02):
+ """
+ :param model: the model to be inited
+ :param conv_std_or_gain: how to init every conv layer `m`
+ > 0: nn.init.trunc_normal_(m.weight.data, std=conv_std_or_gain)
+ < 0: nn.init.xavier_normal_(m.weight.data, gain=-conv_std_or_gain)
+ :param other_std: how to init every linear layer or embedding layer
+ use nn.init.trunc_normal_(m.weight.data, std=other_std)
+ """
+ skip = abs(conv_std_or_gain) > 10
+ if skip: return
+ print(f'[init_weights] {type(model).__name__} with {"std" if conv_std_or_gain > 0 else "gain"}={abs(conv_std_or_gain):g}')
+ for m in model.modules():
+ if isinstance(m, nn.Linear):
+ nn.init.trunc_normal_(m.weight.data, std=other_std)
+ if m.bias is not None:
+ nn.init.constant_(m.bias.data, 0.)
+ elif isinstance(m, nn.Embedding):
+ nn.init.trunc_normal_(m.weight.data, std=other_std)
+ if m.padding_idx is not None:
+ m.weight.data[m.padding_idx].zero_()
+ elif isinstance(m, (nn.Conv1d, nn.Conv2d, nn.ConvTranspose1d, nn.ConvTranspose2d)):
+ 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)
+ if hasattr(m, 'bias') and m.bias is not None:
+ nn.init.constant_(m.bias.data, 0.)
+ elif isinstance(m, (nn.LayerNorm, nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, nn.SyncBatchNorm, nn.GroupNorm, nn.InstanceNorm1d, nn.InstanceNorm2d, nn.InstanceNorm3d)):
+ if m.bias is not None:
+ nn.init.constant_(m.bias.data, 0.)
+ if m.weight is not None:
+ nn.init.constant_(m.weight.data, 1.)
diff --git a/grn/models/rope.py b/grn/models/rope.py
new file mode 100644
index 0000000000000000000000000000000000000000..fda111baa1e893e9c8263842f617f7956cbbbc69
--- /dev/null
+++ b/grn/models/rope.py
@@ -0,0 +1,191 @@
+import math
+import os
+from functools import partial
+from typing import Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+from timm.models.layers import DropPath, drop_path
+from torch.utils.checkpoint import checkpoint
+
+
+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=[]):
+ # split the dimension into half, one for x and one for y
+ half_dim = dim // 2
+ 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
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq)
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq)
+ t_height = t_height / scaling_factor
+ freqs_height = torch.outer(t_height, inv_freq) # (max_height, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2), namely y*theta
+ t_width = t_width / scaling_factor
+ freqs_width = torch.outer(t_width, inv_freq) # (max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2), namely x*theta
+ freqs_grid_map = torch.concat([
+ freqs_height[:, None, :].expand(-1, max_width, -1), # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2)
+ freqs_width[None, :, :].expand(max_height, -1, -1), # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d) / 2)
+ ], dim=-1) # (max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d))
+ freqs_grid_map = torch.stack([torch.cos(freqs_grid_map), torch.sin(freqs_grid_map)], dim=0)
+ # (2, max_height, max_width, dim / (1 for 1d, 2 for 2d, 3 for 3d))
+
+ rope2d_freqs_grid = {}
+ for h_div_w in activated_h_div_w_templates:
+ assert h_div_w in dynamic_resolution_h_w, f'Unknown h_div_w: {h_div_w}'
+ scale_schedule = dynamic_resolution_h_w[h_div_w]['1M']['image_scales']
+ _, ph, pw = scale_schedule[-1]
+ max_edge_length = freqs_grid_map.shape[1]
+ if ph >= pw:
+ uph, upw = max_edge_length, int(max_edge_length / ph * pw)
+ else:
+ uph, upw = int(max_edge_length / pw * ph), max_edge_length
+ rope_cache_list = []
+ for (_, ph, pw) in scale_schedule:
+ ph_mul_pw = ph * pw
+ if rope2d_normalized_by_hw == 1: # downsample
+ rope_cache = F.interpolate(freqs_grid_map[:, :uph, :upw, :].permute([0,3,1,2]), size=(ph, pw), mode='bilinear', align_corners=True)
+ rope_cache = rope_cache.permute([0,2,3,1]) # (2, ph, pw, half_head_dim)
+ elif rope2d_normalized_by_hw == 2: # star stylee
+ _, uph, upw = scale_schedule[-1]
+ indices = torch.stack([
+ (torch.arange(ph) * (uph / ph)).reshape(ph, 1).expand(ph, pw),
+ (torch.arange(pw) * (upw / pw)).reshape(1, pw).expand(ph, pw),
+ ], dim=-1).round().int() # (ph, pw, 2)
+ indices = indices.reshape(-1, 2) # (ph*pw, 2)
+ rope_cache = freqs_grid_map[:, indices[:,0], indices[:,1], :] # (2, ph*pw, half_head_dim)
+ rope_cache = rope_cache.reshape(2, ph, pw, -1)
+ elif rope2d_normalized_by_hw == 0:
+ rope_cache = freqs_grid_map[:, :ph, :pw, :] # (2, ph, pw, half_head_dim)
+ else:
+ raise ValueError(f'Unknown rope2d_normalized_by_hw: {rope2d_normalized_by_hw}')
+ rope_cache_list.append(rope_cache.reshape(2, ph_mul_pw, -1))
+ cat_rope_cache = torch.cat(rope_cache_list, 1) # (2, seq_len, half_head_dim)
+ if cat_rope_cache.shape[1] % pad_to_multiplier:
+ pad = torch.zeros(2, pad_to_multiplier - cat_rope_cache.shape[1] % pad_to_multiplier, half_dim)
+ cat_rope_cache = torch.cat([cat_rope_cache, pad], dim=1)
+ cat_rope_cache = cat_rope_cache[:,None,None,None] # (2, 1, 1, 1, seq_len, half_dim)
+ for pn in dynamic_resolution_h_w[h_div_w]:
+ scale_schedule = dynamic_resolution_h_w[h_div_w][pn]['image_scales']
+ tmp_scale_schedule = [(1, h, w) for _, h, w in scale_schedule]
+ rope2d_freqs_grid[str(tuple(tmp_scale_schedule))] = cat_rope_cache
+ return rope2d_freqs_grid
+
+
+def precompute_rope3d_freqs_grid(
+ dim,
+ rope2d_normalized_by_hw,
+ max_frames=128,
+ max_height=2048 // 8,
+ max_width=2048 // 8,
+ base=10000.0,
+ device=None,
+ activated_h_div_w_templates=[],
+ text_maxlen=0,
+ pn=None,
+ args=None,
+ **kwargs,
+):
+ # split the dimension into three parts, one for x, one for y, and one for t
+ print(f'[precompute_rope4d_freqs_grid: 3d]: start')
+ assert dim % 2 == 0, f'Only support dim % 2 == 0, but got dim={dim}'
+ dim_div_2 = dim // 2
+ num_of_freqs_former = dim_div_2 // 3
+ preserve_1d_length = 600
+ 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
+ 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
+ inv_freq_last = 1.0 / (base ** (torch.arange(num_of_freqs_last, dtype=torch.int64).float().to(device) / num_of_freqs_last))
+ t_frames = torch.arange(preserve_1d_length+max_frames, device=device, dtype=torch.int64).type_as(inv_freq_former)
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq_former)
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq_former)
+ freqs_frames = torch.outer(t_frames, inv_freq_former) # (max_frames, (dim_div_2 / 3)), namely x*theta
+ freqs_height = torch.outer(t_height, inv_freq_former) # (max_height, (dim_div_2 / 3), namely y*theta
+ freqs_width = torch.outer(t_width, inv_freq_last) # (max_width, (dim_div_2 / 3)), namely x*theta
+ freqs_frames = torch.stack([torch.cos(freqs_frames), torch.sin(freqs_frames)], dim=0)
+ freqs_height = torch.stack([torch.cos(freqs_height), torch.sin(freqs_height)], dim=0)
+ freqs_width = torch.stack([torch.cos(freqs_width), torch.sin(freqs_width)], dim=0)
+ tm = preserve_1d_length
+ rope_text_embeds = torch.cat([
+ freqs_frames[ :, :tm, None, None, :].expand(-1, -1, -1, -1, -1),
+ freqs_height[ :, None, :1, None, :].expand(-1, tm, -1, -1, -1),
+ freqs_width[ :, None, None, :1, :].expand(-1, tm, -1, -1, -1),
+ ], dim=-1) # (2, tm, 1, 1, dim_div_2)
+ rope_text_embeds = rope_text_embeds.reshape(2, 1, 1, 1, tm, dim_div_2)
+ rope2d_freqs_grid = {}
+ rope2d_freqs_grid['freqs_text'] = rope_text_embeds # (2, 1, 1, 1, preserve_1d_length, dim / 2)
+ rope2d_freqs_grid['freqs_frames'] = freqs_frames[:, tm:] # (2, max_frames, ceil(dim_div_2 / 4))
+ rope2d_freqs_grid['freqs_height'] = freqs_height # (2, max_height, ceil(dim_div_2 / 4))
+ rope2d_freqs_grid['freqs_width'] = freqs_width # (2, max_width, ceil(dim_div_2 / 4))
+ return rope2d_freqs_grid
+
+
+def precompute_rope4d_freqs_grid(
+ dim,
+ rope2d_normalized_by_hw,
+ max_scales=128,
+ max_frames=128,
+ max_height=2048 // 8,
+ max_width=2048 // 8,
+ base=10000.0,
+ device=None,
+ activated_h_div_w_templates=[],
+ text_maxlen=0,
+ pn=None,
+ args=None,
+ **kwargs,
+):
+ # split the dimension into three parts, one for x, one for y, and one for t
+ print(f'[precompute_rope4d_freqs_grid: 4d]: start')
+ assert dim % 2 == 0, f'Only support dim % 2 == 0, but got dim={dim}'
+ dim_div_2 = dim // 2
+ num_of_freqs = int(np.ceil(dim_div_2 / 4))
+ 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
+ t_scales = torch.arange(text_maxlen+max_scales, device=device, dtype=torch.int64).type_as(inv_freq)
+ t_frames = torch.arange(max_frames, device=device, dtype=torch.int64).type_as(inv_freq)
+ t_height = torch.arange(max_height, device=device, dtype=torch.int64).type_as(inv_freq)
+ t_width = torch.arange(max_width, device=device, dtype=torch.int64).type_as(inv_freq)
+ freqs_scales = torch.outer(t_scales, inv_freq) # (text_maxlen+max_scales, ceil(dim_div_2 / 4)), namely x*theta
+ freqs_frames = torch.outer(t_frames, inv_freq) # (max_frames, ceil(dim_div_2 / 4)), namely x*theta
+ freqs_height = torch.outer(t_height, inv_freq) # (max_height, ceil(dim_div_2 / 4)), namely y*theta
+ freqs_width = torch.outer(t_width, inv_freq) # (max_width, ceil(dim_div_2 / 4)), namely x*theta
+ assert num_of_freqs*4==dim_div_2
+ freqs_scales = torch.stack([torch.cos(freqs_scales), torch.sin(freqs_scales)], dim=0)
+ freqs_frames = torch.stack([torch.cos(freqs_frames), torch.sin(freqs_frames)], dim=0)
+ freqs_height = torch.stack([torch.cos(freqs_height), torch.sin(freqs_height)], dim=0)
+ freqs_width = torch.stack([torch.cos(freqs_width), torch.sin(freqs_width)], dim=0)
+ tm = text_maxlen
+ rope_text_embeds = torch.cat([
+ freqs_scales[ :, :tm, None, None, None, :].expand(-1, -1, -1, -1, -1, -1),
+ freqs_frames[ :, None, :1, None, None, :].expand(-1, tm, -1, -1, -1, -1),
+ freqs_height[ :, None, None, :1, None, :].expand(-1, tm, -1, -1, -1, -1),
+ freqs_width[ :, None, None, None, :1, :].expand(-1, tm, -1, -1, -1, -1),
+ ], dim=-1) # (2, tm, 1, 1, 1, dim_div_2)
+ rope_text_embeds = rope_text_embeds.reshape(2, 1, 1, 1, tm, dim_div_2)
+ rope2d_freqs_grid = {}
+ rope2d_freqs_grid['freqs_text'] = rope_text_embeds # (2, 1, 1, 1, text_maxlen, dim / 2)
+ rope2d_freqs_grid['freqs_scales'] = freqs_scales[:, tm:] # (2, max_scales, ceil(dim_div_2 / 4))
+ rope2d_freqs_grid['freqs_frames'] = freqs_frames # (2, max_frames, ceil(dim_div_2 / 4))
+ rope2d_freqs_grid['freqs_height'] = freqs_height # (2, max_height, ceil(dim_div_2 / 4))
+ rope2d_freqs_grid['freqs_width'] = freqs_width # (2, max_width, ceil(dim_div_2 / 4))
+ return rope2d_freqs_grid
+
+def apply_rotary_emb(q, k, rope_cache):
+ device_type = q.device.type
+ device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"
+ qk = [q, k]
+ rope_cache = rope_cache[:,0]
+ with torch.autocast(device_type=device_type, enabled=False):
+ for i in range(2):
+ qk[i] = qk[i].reshape(*qk[i].shape[:-1], -1, 2)
+ tmp1 = qk[i][..., 1] * rope_cache[1]
+ tmp2 = qk[i][..., 0] * rope_cache[1]
+ qk[i][..., 0].mul_(rope_cache[0]).sub_(tmp1)
+ qk[i][..., 1].mul_(rope_cache[0]).add_(tmp2)
+ qk[i] = qk[i].reshape(*qk[i].shape[:-2], -1)
+ q, k = qk
+ # qk = qk.reshape(*qk.shape[:-1], -1, 2) #(2, batch_size, heads, seq_len, half_head_dim, 2)
+ # qk = torch.stack([
+ # qk[...,0] * rope_cache[0] - qk[...,1] * rope_cache[1],
+ # qk[...,0] * rope_cache[1] + qk[...,1] * rope_cache[0],
+ # ], dim=-1) # (2, batch_size, heads, seq_len, half_head_dim, 2), here stack + reshape should not be concate
+ # qk = qk.reshape(*qk.shape[:-2], -1) #(2, batch_size, heads, seq_len, head_dim)
+ # q, k = qk.unbind(dim=0) # (batch_size, heads, seq_len, head_dim)
+ return q, k
\ No newline at end of file
diff --git a/grn/models/umt5/fsdp.py b/grn/models/umt5/fsdp.py
new file mode 100644
index 0000000000000000000000000000000000000000..8935e84e3eff9a1ed0c329fd7d4eef35f0f41276
--- /dev/null
+++ b/grn/models/umt5/fsdp.py
@@ -0,0 +1,42 @@
+import gc
+from functools import partial
+
+import torch
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
+from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
+from torch.distributed.utils import _free_storage
+
+
+def shard_model(
+ model,
+ device_id,
+ param_dtype=torch.bfloat16,
+ reduce_dtype=torch.float32,
+ buffer_dtype=torch.float32,
+ process_group=None,
+ sharding_strategy=ShardingStrategy.FULL_SHARD,
+ sync_module_states=True,
+):
+ model = FSDP(
+ module=model,
+ process_group=process_group,
+ sharding_strategy=sharding_strategy,
+ auto_wrap_policy=partial(
+ lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
+ mixed_precision=MixedPrecision(
+ param_dtype=param_dtype,
+ reduce_dtype=reduce_dtype,
+ buffer_dtype=buffer_dtype),
+ device_id=device_id,
+ sync_module_states=sync_module_states)
+ return model
+
+
+def free_model(model):
+ for m in model.modules():
+ if isinstance(m, FSDP):
+ _free_storage(m._handle.flat_param.data)
+ del model
+ gc.collect()
+ torch.cuda.empty_cache()
\ No newline at end of file
diff --git a/grn/models/umt5/t5.py b/grn/models/umt5/t5.py
new file mode 100644
index 0000000000000000000000000000000000000000..6d522e10206fe64836108debdae7c6bdb49e87b8
--- /dev/null
+++ b/grn/models/umt5/t5.py
@@ -0,0 +1,514 @@
+import logging
+import math
+from functools import partial
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from grn.models.umt5.fsdp import shard_model
+from grn.models.umt5.umt5_tokenizers import HuggingfaceTokenizer
+
+__all__ = [
+ 'T5Model',
+ 'T5Encoder',
+ 'T5Decoder',
+ 'T5EncoderModel',
+]
+
+
+def fp16_clamp(x):
+ if x.dtype == torch.float16 and torch.isinf(x).any():
+ clamp = torch.finfo(x.dtype).max - 1000
+ x = torch.clamp(x, min=-clamp, max=clamp)
+ return x
+
+
+def init_weights(m):
+ if isinstance(m, T5LayerNorm):
+ nn.init.ones_(m.weight)
+ elif isinstance(m, T5Model):
+ nn.init.normal_(m.token_embedding.weight, std=1.0)
+ elif isinstance(m, T5FeedForward):
+ nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)
+ nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)
+ nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)
+ elif isinstance(m, T5Attention):
+ nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)
+ nn.init.normal_(m.k.weight, std=m.dim**-0.5)
+ nn.init.normal_(m.v.weight, std=m.dim**-0.5)
+ nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)
+ elif isinstance(m, T5RelativeEmbedding):
+ nn.init.normal_(
+ m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)
+
+
+class GELU(nn.Module):
+
+ def forward(self, x):
+ return 0.5 * x * (1.0 + torch.tanh(
+ math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
+
+
+class T5LayerNorm(nn.Module):
+
+ def __init__(self, dim, eps=1e-6):
+ super(T5LayerNorm, self).__init__()
+ self.dim = dim
+ self.eps = eps
+ self.weight = nn.Parameter(torch.ones(dim))
+
+ def forward(self, x):
+ x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +
+ self.eps)
+ if self.weight.dtype in [torch.float16, torch.bfloat16]:
+ x = x.type_as(self.weight)
+ return self.weight * x
+
+
+class T5Attention(nn.Module):
+
+ def __init__(self, dim, dim_attn, num_heads, dropout=0.1):
+ assert dim_attn % num_heads == 0
+ super(T5Attention, self).__init__()
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.num_heads = num_heads
+ self.head_dim = dim_attn // num_heads
+
+ # layers
+ self.q = nn.Linear(dim, dim_attn, bias=False)
+ self.k = nn.Linear(dim, dim_attn, bias=False)
+ self.v = nn.Linear(dim, dim_attn, bias=False)
+ self.o = nn.Linear(dim_attn, dim, bias=False)
+ self.dropout = nn.Dropout(dropout)
+
+ def forward(self, x, context=None, mask=None, pos_bias=None):
+ """
+ x: [B, L1, C].
+ context: [B, L2, C] or None.
+ mask: [B, L2] or [B, L1, L2] or None.
+ """
+ # check inputs
+ context = x if context is None else context
+ b, n, c = x.size(0), self.num_heads, self.head_dim
+
+ # compute query, key, value
+ q = self.q(x).view(b, -1, n, c)
+ k = self.k(context).view(b, -1, n, c)
+ v = self.v(context).view(b, -1, n, c)
+
+ # attention bias
+ attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
+ if pos_bias is not None:
+ attn_bias += pos_bias
+ if mask is not None:
+ assert mask.ndim in [2, 3]
+ mask = mask.view(b, 1, 1,
+ -1) if mask.ndim == 2 else mask.unsqueeze(1)
+ attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
+
+ # compute attention (T5 does not use scaling)
+ attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias
+ attn = F.softmax(attn.float(), dim=-1).type_as(attn)
+ x = torch.einsum('bnij,bjnc->binc', attn, v)
+
+ # output
+ x = x.reshape(b, -1, n * c)
+ x = self.o(x)
+ x = self.dropout(x)
+ return x
+
+
+class T5FeedForward(nn.Module):
+
+ def __init__(self, dim, dim_ffn, dropout=0.1):
+ super(T5FeedForward, self).__init__()
+ self.dim = dim
+ self.dim_ffn = dim_ffn
+
+ # layers
+ self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
+ self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
+ self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
+ self.dropout = nn.Dropout(dropout)
+
+ def forward(self, x):
+ x = self.fc1(x) * self.gate(x)
+ x = self.dropout(x)
+ x = self.fc2(x)
+ x = self.dropout(x)
+ return x
+
+
+class T5SelfAttention(nn.Module):
+
+ def __init__(self,
+ dim,
+ dim_attn,
+ dim_ffn,
+ num_heads,
+ num_buckets,
+ shared_pos=True,
+ dropout=0.1):
+ super(T5SelfAttention, self).__init__()
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.dim_ffn = dim_ffn
+ self.num_heads = num_heads
+ self.num_buckets = num_buckets
+ self.shared_pos = shared_pos
+
+ # layers
+ self.norm1 = T5LayerNorm(dim)
+ self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
+ self.norm2 = T5LayerNorm(dim)
+ self.ffn = T5FeedForward(dim, dim_ffn, dropout)
+ self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
+ num_buckets, num_heads, bidirectional=True)
+
+ def forward(self, x, mask=None, pos_bias=None):
+ e = pos_bias if self.shared_pos else self.pos_embedding(
+ x.size(1), x.size(1))
+ x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
+ x = fp16_clamp(x + self.ffn(self.norm2(x)))
+ return x
+
+
+class T5CrossAttention(nn.Module):
+
+ def __init__(self,
+ dim,
+ dim_attn,
+ dim_ffn,
+ num_heads,
+ num_buckets,
+ shared_pos=True,
+ dropout=0.1):
+ super(T5CrossAttention, self).__init__()
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.dim_ffn = dim_ffn
+ self.num_heads = num_heads
+ self.num_buckets = num_buckets
+ self.shared_pos = shared_pos
+
+ # layers
+ self.norm1 = T5LayerNorm(dim)
+ self.self_attn = T5Attention(dim, dim_attn, num_heads, dropout)
+ self.norm2 = T5LayerNorm(dim)
+ self.cross_attn = T5Attention(dim, dim_attn, num_heads, dropout)
+ self.norm3 = T5LayerNorm(dim)
+ self.ffn = T5FeedForward(dim, dim_ffn, dropout)
+ self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
+ num_buckets, num_heads, bidirectional=False)
+
+ def forward(self,
+ x,
+ mask=None,
+ encoder_states=None,
+ encoder_mask=None,
+ pos_bias=None):
+ e = pos_bias if self.shared_pos else self.pos_embedding(
+ x.size(1), x.size(1))
+ x = fp16_clamp(x + self.self_attn(self.norm1(x), mask=mask, pos_bias=e))
+ x = fp16_clamp(x + self.cross_attn(
+ self.norm2(x), context=encoder_states, mask=encoder_mask))
+ x = fp16_clamp(x + self.ffn(self.norm3(x)))
+ return x
+
+
+class T5RelativeEmbedding(nn.Module):
+
+ def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
+ super(T5RelativeEmbedding, self).__init__()
+ self.num_buckets = num_buckets
+ self.num_heads = num_heads
+ self.bidirectional = bidirectional
+ self.max_dist = max_dist
+
+ # layers
+ self.embedding = nn.Embedding(num_buckets, num_heads)
+
+ def forward(self, lq, lk):
+ device = self.embedding.weight.device
+ # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
+ # torch.arange(lq).unsqueeze(1).to(device)
+ rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
+ torch.arange(lq, device=device).unsqueeze(1)
+ rel_pos = self._relative_position_bucket(rel_pos)
+ rel_pos_embeds = self.embedding(rel_pos)
+ rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
+ 0) # [1, N, Lq, Lk]
+ return rel_pos_embeds.contiguous()
+
+ def _relative_position_bucket(self, rel_pos):
+ # preprocess
+ if self.bidirectional:
+ num_buckets = self.num_buckets // 2
+ rel_buckets = (rel_pos > 0).long() * num_buckets
+ rel_pos = torch.abs(rel_pos)
+ else:
+ num_buckets = self.num_buckets
+ rel_buckets = 0
+ rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
+
+ # embeddings for small and large positions
+ max_exact = num_buckets // 2
+ rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
+ math.log(self.max_dist / max_exact) *
+ (num_buckets - max_exact)).long()
+ rel_pos_large = torch.min(
+ rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
+ rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
+ return rel_buckets
+
+
+class T5Encoder(nn.Module):
+
+ def __init__(self,
+ vocab,
+ dim,
+ dim_attn,
+ dim_ffn,
+ num_heads,
+ num_layers,
+ num_buckets,
+ shared_pos=True,
+ dropout=0.1):
+ super(T5Encoder, self).__init__()
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.dim_ffn = dim_ffn
+ self.num_heads = num_heads
+ self.num_layers = num_layers
+ self.num_buckets = num_buckets
+ self.shared_pos = shared_pos
+
+ # layers
+ self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
+ else nn.Embedding(vocab, dim)
+ self.pos_embedding = T5RelativeEmbedding(
+ num_buckets, num_heads, bidirectional=True) if shared_pos else None
+ self.dropout = nn.Dropout(dropout)
+ self.blocks = nn.ModuleList([
+ T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
+ shared_pos, dropout) for _ in range(num_layers)
+ ])
+ self.norm = T5LayerNorm(dim)
+
+ # initialize weights
+ self.apply(init_weights)
+
+ def forward(self, ids, mask=None):
+ x = self.token_embedding(ids)
+ x = self.dropout(x)
+ e = self.pos_embedding(x.size(1),
+ x.size(1)) if self.shared_pos else None
+ for block in self.blocks:
+ x = block(x, mask, pos_bias=e)
+ x = self.norm(x)
+ x = self.dropout(x)
+ return x
+
+
+class T5Decoder(nn.Module):
+
+ def __init__(self,
+ vocab,
+ dim,
+ dim_attn,
+ dim_ffn,
+ num_heads,
+ num_layers,
+ num_buckets,
+ shared_pos=True,
+ dropout=0.1):
+ super(T5Decoder, self).__init__()
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.dim_ffn = dim_ffn
+ self.num_heads = num_heads
+ self.num_layers = num_layers
+ self.num_buckets = num_buckets
+ self.shared_pos = shared_pos
+
+ # layers
+ self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
+ else nn.Embedding(vocab, dim)
+ self.pos_embedding = T5RelativeEmbedding(
+ num_buckets, num_heads, bidirectional=False) if shared_pos else None
+ self.dropout = nn.Dropout(dropout)
+ self.blocks = nn.ModuleList([
+ T5CrossAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
+ shared_pos, dropout) for _ in range(num_layers)
+ ])
+ self.norm = T5LayerNorm(dim)
+
+ # initialize weights
+ self.apply(init_weights)
+
+ def forward(self, ids, mask=None, encoder_states=None, encoder_mask=None):
+ b, s = ids.size()
+
+ # causal mask
+ if mask is None:
+ mask = torch.tril(torch.ones(1, s, s).to(ids.device))
+ elif mask.ndim == 2:
+ mask = torch.tril(mask.unsqueeze(1).expand(-1, s, -1))
+
+ # layers
+ x = self.token_embedding(ids)
+ x = self.dropout(x)
+ e = self.pos_embedding(x.size(1),
+ x.size(1)) if self.shared_pos else None
+ for block in self.blocks:
+ x = block(x, mask, encoder_states, encoder_mask, pos_bias=e)
+ x = self.norm(x)
+ x = self.dropout(x)
+ return x
+
+
+class T5Model(nn.Module):
+
+ def __init__(self,
+ vocab_size,
+ dim,
+ dim_attn,
+ dim_ffn,
+ num_heads,
+ encoder_layers,
+ decoder_layers,
+ num_buckets,
+ shared_pos=True,
+ dropout=0.1):
+ super(T5Model, self).__init__()
+ self.vocab_size = vocab_size
+ self.dim = dim
+ self.dim_attn = dim_attn
+ self.dim_ffn = dim_ffn
+ self.num_heads = num_heads
+ self.encoder_layers = encoder_layers
+ self.decoder_layers = decoder_layers
+ self.num_buckets = num_buckets
+
+ # layers
+ self.token_embedding = nn.Embedding(vocab_size, dim)
+ self.encoder = T5Encoder(self.token_embedding, dim, dim_attn, dim_ffn,
+ num_heads, encoder_layers, num_buckets,
+ shared_pos, dropout)
+ self.decoder = T5Decoder(self.token_embedding, dim, dim_attn, dim_ffn,
+ num_heads, decoder_layers, num_buckets,
+ shared_pos, dropout)
+ self.head = nn.Linear(dim, vocab_size, bias=False)
+
+ # initialize weights
+ self.apply(init_weights)
+
+ def forward(self, encoder_ids, encoder_mask, decoder_ids, decoder_mask):
+ x = self.encoder(encoder_ids, encoder_mask)
+ x = self.decoder(decoder_ids, decoder_mask, x, encoder_mask)
+ x = self.head(x)
+ return x
+
+
+def _t5(name,
+ encoder_only=False,
+ decoder_only=False,
+ return_tokenizer=False,
+ tokenizer_kwargs={},
+ dtype=torch.float32,
+ device='cpu',
+ **kwargs):
+ # sanity check
+ assert not (encoder_only and decoder_only)
+
+ # params
+ if encoder_only:
+ model_cls = T5Encoder
+ kwargs['vocab'] = kwargs.pop('vocab_size')
+ kwargs['num_layers'] = kwargs.pop('encoder_layers')
+ _ = kwargs.pop('decoder_layers')
+ elif decoder_only:
+ model_cls = T5Decoder
+ kwargs['vocab'] = kwargs.pop('vocab_size')
+ kwargs['num_layers'] = kwargs.pop('decoder_layers')
+ _ = kwargs.pop('encoder_layers')
+ else:
+ model_cls = T5Model
+
+ # init model
+ with torch.device(device):
+ model = model_cls(**kwargs)
+
+ # set device
+ model = model.to(dtype=dtype, device=device)
+
+ # init tokenizer
+ if return_tokenizer:
+ from .tokenizers import HuggingfaceTokenizer
+ tokenizer = HuggingfaceTokenizer(f'google/{name}', **tokenizer_kwargs)
+ return model, tokenizer
+ else:
+ return model
+
+
+def umt5_xxl(**kwargs):
+ cfg = dict(
+ vocab_size=256384,
+ dim=4096,
+ dim_attn=4096,
+ dim_ffn=10240,
+ num_heads=64,
+ encoder_layers=24,
+ decoder_layers=24,
+ num_buckets=32,
+ shared_pos=False,
+ dropout=0.1)
+ cfg.update(**kwargs)
+ return _t5('umt5-xxl', **cfg)
+
+
+class T5EncoderModel:
+
+ def __init__(
+ self,
+ text_len,
+ dtype=torch.bfloat16,
+ device=torch.cuda.current_device(),
+ checkpoint_path=None,
+ tokenizer_path=None,
+ enable_fsdp=False,
+ ):
+ self.text_len = text_len
+ self.dtype = dtype
+ self.device = device
+ self.checkpoint_path = checkpoint_path
+ self.tokenizer_path = tokenizer_path
+
+ # init model
+ model = umt5_xxl(
+ encoder_only=True,
+ return_tokenizer=False,
+ dtype=dtype,
+ device=device).eval().requires_grad_(False)
+ logging.info(f'loading {checkpoint_path}')
+ model.load_state_dict(torch.load(checkpoint_path, map_location='cpu'))
+ self.model = model
+ if enable_fsdp:
+ shard_fn = partial(shard_model, device_id=device)
+ self.model = shard_fn(self.model, sync_module_states=False)
+ else:
+ self.model.to(self.device)
+ # init tokenizer
+ self.tokenizer = HuggingfaceTokenizer(
+ name=tokenizer_path, seq_len=text_len, clean='whitespace')
+
+ def __call__(self, texts, device):
+ ids, mask = self.tokenizer(
+ texts, return_mask=True, add_special_tokens=True)
+ ids = ids.to(device)
+ mask = mask.to(device)
+ seq_lens = mask.gt(0).sum(dim=1).long()
+ context = self.model(ids, mask)
+ return [u[:v] for u, v in zip(context, seq_lens)]
diff --git a/grn/models/umt5/umt5_tokenizers.py b/grn/models/umt5/umt5_tokenizers.py
new file mode 100644
index 0000000000000000000000000000000000000000..a69972adf2711f73e9b5c3b59dcf0264fa7742a2
--- /dev/null
+++ b/grn/models/umt5/umt5_tokenizers.py
@@ -0,0 +1,81 @@
+import html
+import string
+
+import ftfy
+import regex as re
+from transformers import AutoTokenizer
+
+__all__ = ['HuggingfaceTokenizer']
+
+
+def basic_clean(text):
+ text = ftfy.fix_text(text)
+ text = html.unescape(html.unescape(text))
+ return text.strip()
+
+
+def whitespace_clean(text):
+ text = re.sub(r'\s+', ' ', text)
+ text = text.strip()
+ return text
+
+
+def canonicalize(text, keep_punctuation_exact_string=None):
+ text = text.replace('_', ' ')
+ if keep_punctuation_exact_string:
+ text = keep_punctuation_exact_string.join(
+ part.translate(str.maketrans('', '', string.punctuation))
+ for part in text.split(keep_punctuation_exact_string))
+ else:
+ text = text.translate(str.maketrans('', '', string.punctuation))
+ text = text.lower()
+ text = re.sub(r'\s+', ' ', text)
+ return text.strip()
+
+
+class HuggingfaceTokenizer:
+
+ def __init__(self, name, seq_len=None, clean=None, **kwargs):
+ assert clean in (None, 'whitespace', 'lower', 'canonicalize')
+ self.name = name
+ self.seq_len = seq_len
+ self.clean = clean
+
+ # init tokenizer
+ self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
+ self.vocab_size = self.tokenizer.vocab_size
+
+ def __call__(self, sequence, **kwargs):
+ return_mask = kwargs.pop('return_mask', False)
+
+ # arguments
+ _kwargs = {'return_tensors': 'pt'}
+ if self.seq_len is not None:
+ _kwargs.update({
+ 'padding': 'max_length',
+ 'truncation': True,
+ 'max_length': self.seq_len
+ })
+ _kwargs.update(**kwargs)
+
+ # tokenization
+ if isinstance(sequence, str):
+ sequence = [sequence]
+ if self.clean:
+ sequence = [self._clean(u) for u in sequence]
+ ids = self.tokenizer(sequence, **_kwargs)
+
+ # output
+ if return_mask:
+ return ids.input_ids, ids.attention_mask
+ else:
+ return ids.input_ids
+
+ def _clean(self, text):
+ if self.clean == 'whitespace':
+ text = whitespace_clean(basic_clean(text))
+ elif self.clean == 'lower':
+ text = whitespace_clean(basic_clean(text)).lower()
+ elif self.clean == 'canonicalize':
+ text = canonicalize(basic_clean(text))
+ return text
diff --git a/grn/schedules/__init__.py b/grn/schedules/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..1b814f56d21e54be0ed54b5d956a9ffe9c0b7d6b
--- /dev/null
+++ b/grn/schedules/__init__.py
@@ -0,0 +1,6 @@
+def get_encode_decode_func(dynamic_scale_schedule):
+ if 'GRN_vae_stride16' in dynamic_scale_schedule:
+ from grn.schedules.global_refine import video_encode, video_decode, get_visual_rope_embeds, get_scale_pack_info
+ else:
+ raise NotImplementedError(f'{dynamic_scale_schedule} is unsupported')
+ return video_encode, video_decode, get_visual_rope_embeds, get_scale_pack_info
diff --git a/grn/schedules/dynamic_resolution.py b/grn/schedules/dynamic_resolution.py
new file mode 100644
index 0000000000000000000000000000000000000000..729b04571e838784e48d598b296e2e02bf060b83
--- /dev/null
+++ b/grn/schedules/dynamic_resolution.py
@@ -0,0 +1,99 @@
+import json
+import math
+import copy
+
+import tqdm
+import numpy as np
+
+
+def get_first_full_spatial_size_scale_index(vae_scale_schedule):
+ for si, (pt, ph, pw) in enumerate(vae_scale_schedule):
+ if vae_scale_schedule[si][-2:] == vae_scale_schedule[-1][-2:]:
+ return si
+
+def get_full_spatial_size_scale_indices(vae_scale_schedule):
+ full_spatial_size_scale_indices = []
+ for si, (pt, ph, pw) in enumerate(vae_scale_schedule):
+ if vae_scale_schedule[si][-2:] == vae_scale_schedule[-1][-2:]:
+ full_spatial_size_scale_indices.append(si)
+ return full_spatial_size_scale_indices
+
+def get_ratio2hws_pixels2scales(dynamic_scale_schedule, train_h_div_w_list, video_frames):
+ compressed_frames = video_frames // 4 + 1
+ if dynamic_scale_schedule in ['GRN_vae_stride16']:
+ assert type(train_h_div_w_list) is str
+ train_h_div_w_list = json.loads(train_h_div_w_list)
+ if len(train_h_div_w_list) == 0:
+ 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]
+ vae_stride = 16
+ if 'vae_stride' in dynamic_scale_schedule:
+ vae_stride = int(dynamic_scale_schedule.split('vae_stride')[-1])
+ dynamic_resolution_h_w = {}
+ for h_div_w in train_h_div_w_list:
+ ratio = int(h_div_w*1000)/1000
+ dynamic_resolution_h_w[ratio] = {}
+ for pn in ['0.06M', '0.25M', '0.41M', '0.92M', '1M', '2M']:
+ if pn == '0.06M': # 256x256, 192p
+ scale = 8
+ elif pn == '0.25M': # 512x512, 384p
+ scale = 16
+ elif pn == '0.41M': # 640x640, 480p
+ scale = 20
+ elif pn == '0.92M': # 960x960, 720p
+ scale = 30
+ elif pn == '1M': # 1024x1024, 768p
+ scale = 32
+ elif pn == '2M': # 1440x1440, 1080p
+ scale = 45
+ if vae_stride == 16:
+ scale = scale * 2
+ elif vae_stride == 32:
+ scale = scale * 1
+ else:
+ raise ValueError(f'vae_stride {vae_stride} is not supported')
+ area = scale * scale
+ pw_float = math.sqrt(area / h_div_w)
+ ph_float = pw_float * h_div_w
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
+ scales = [(ph,pw)]
+ pixel = (scales[-1][0] * vae_stride, scales[-1][1] * vae_stride)
+ dynamic_resolution_h_w[ratio][pn] = {
+ 'pixel': pixel,
+ 'scales': scales
+ }
+ for ratio in dynamic_resolution_h_w:
+ for pn in dynamic_resolution_h_w[ratio]:
+ base_scale_schedule = dynamic_resolution_h_w[ratio][pn]['scales']
+ scales_in_one_clip = len(base_scale_schedule)
+ dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'] = {}
+ for pt in range(1, compressed_frames+1, 1):
+ dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'][pt] = [(pt, h, w) for h, w in base_scale_schedule]
+ dynamic_resolution_h_w[ratio][pn]['image_scales'] = scales_in_one_clip
+ dynamic_resolution_h_w[ratio][pn]['scales_in_one_clip'] = scales_in_one_clip
+ dynamic_resolution_h_w[ratio][pn]['max_video_scales'] = len(dynamic_resolution_h_w[ratio][pn]['pt2scale_schedule'][compressed_frames])
+ del dynamic_resolution_h_w[ratio][pn]['scales']
+ else:
+ raise ValueError(f'dynamic_scale_schedule={dynamic_scale_schedule} not implemented')
+ return dynamic_resolution_h_w
+
+def get_dynamic_resolution_meta(dynamic_scale_schedule, train_h_div_w_list, video_frames):
+ dynamic_resolution_h_w = get_ratio2hws_pixels2scales(dynamic_scale_schedule, train_h_div_w_list, video_frames)
+ h_div_w_templates = []
+ for h_div_w in dynamic_resolution_h_w.keys():
+ h_div_w_templates.append(h_div_w)
+ h_div_w_templates = np.array(h_div_w_templates)
+ return dynamic_resolution_h_w, h_div_w_templates
+
+def get_h_div_w_template2indices(h_div_w_list, h_div_w_templates):
+ indices = list(range(len(h_div_w_list)))
+ h_div_w_template2indices = {}
+ pbar = tqdm.tqdm(total=len(indices), desc='get_h_div_w_template2indices...')
+ for h_div_w, index in zip(h_div_w_list, indices):
+ pbar.update(1)
+ nearest_h_div_w_template_ = h_div_w_templates[np.argmin(np.abs(h_div_w-h_div_w_templates))]
+ if nearest_h_div_w_template_ not in h_div_w_template2indices:
+ h_div_w_template2indices[nearest_h_div_w_template_] = []
+ h_div_w_template2indices[nearest_h_div_w_template_].append(index)
+ for h_div_w_template_, sub_indices in h_div_w_template2indices.items():
+ h_div_w_template2indices[h_div_w_template_] = np.array(sub_indices)
+ return h_div_w_template2indices
diff --git a/grn/schedules/global_refine.py b/grn/schedules/global_refine.py
new file mode 100644
index 0000000000000000000000000000000000000000..920af5853b72deb87ff01512bb8ff968b8d4ebfe
--- /dev/null
+++ b/grn/schedules/global_refine.py
@@ -0,0 +1,220 @@
+import os
+import json
+import math
+import bisect
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+from grn.utils_t2iv.hbq_util_t2iv import multiclass_labels2onehot_input
+
+def get_scale_pack_info(**kwargs):
+ return
+
+def flatten_two_level_list(two_level_list):
+ flatten_list = []
+ for item in two_level_list:
+ flatten_list.extend(item)
+ return flatten_list
+
+def shift_pt(pt, alpha):
+ """shift pt (signal ratio) to lower one, recommand alpha=sqrt(height*width/256/256)"""
+ if alpha > 1000:
+ alpha = alpha - 1000
+ noise_pt = 1 - pt
+ noise_pt = alpha * noise_pt / (1+(alpha-1)*noise_pt) # shift noise_pt to higer one
+ pt = 1 - noise_pt
+ return pt
+
+def video_encode(
+ vae,
+ inp_B3HW,
+ vae_features=None,
+ device='cuda',
+ args=None,
+ infer_mode=False,
+ rope2d_freqs_grid=None,
+ dynamic_resolution_h_w=None,
+ tokens_remain=9999999,
+ text_lens=[],
+ caption_nums=[],
+ rank_vary_generator=None,
+ vis_verbose=False,
+ meta_list=None,
+ **kwargs,
+):
+ if rank_vary_generator is not None:
+ numpy_generator = rank_vary_generator['numpy_generator']
+ torch_cuda_generator = rank_vary_generator['torch_cuda_generator']
+ else:
+ numpy_generator = np.random.default_rng()
+ torch_cuda_generator = torch.Generator(device='cuda')
+
+ if vae_features is None:
+ raw_features, _, _ = vae.encode_for_raw_features(inp_B3HW, scale_schedule=None, slice=True)
+ raw_features_list = [raw_features]
+ x_recon_raw = vae.decode(raw_features[-1], slice=True)
+ x_recon_raw = torch.clamp(x_recon_raw, min=-1, max=1)
+ print(f'raw_features[-1].shape: {raw_features[-1].shape}')
+ else:
+ raw_features_list = vae_features
+ # raw_features_list: list of [1,d,t,h,w]:
+ # import pdb; pdb.set_trace()
+ gt_all_bit_indices = []
+ pred_all_bit_indices = []
+ var_input_list = []
+ sequece_packing_scales = [] # with trunk
+ h_div_w_template_list = np.array(list(dynamic_resolution_h_w.keys()))
+ visual_rope_cache_list = []
+ other_info_by_scale = []
+ scale_lengths = []
+ with torch.amp.autocast('cuda', enabled = False):
+ for example_ind, raw_features in enumerate(raw_features_list):
+ meta = meta_list[example_ind]
+ gt_all_bit_indices.append([])
+ pred_all_bit_indices.append([])
+ var_input_list.append([])
+ visual_rope_cache_list.append([])
+ other_info_by_scale.append([])
+ B, C, T, H, W = raw_features[-1].shape
+ h_div_w = H / W
+ mapped_h_div_w_template = h_div_w_template_list[np.argmin(np.abs(h_div_w-h_div_w_template_list))]
+ pn = meta['pn']
+ if meta['first_frame_condition']:
+ scale_schedule = dynamic_resolution_h_w[mapped_h_div_w_template][pn]['pt2scale_schedule'][T-1]
+ else:
+ scale_schedule = dynamic_resolution_h_w[mapped_h_div_w_template][pn]['pt2scale_schedule'][T]
+ if not infer_mode:
+ next_tokens_remain = tokens_remain - T * H * W - args.add_scale_token - text_lens[example_ind]
+ if next_tokens_remain < 0:
+ break
+ tokens_remain = next_tokens_remain
+ scale_lengths.append(T * H * W + text_lens[example_ind] + args.add_scale_token)
+ preserve_scale_schedule = []
+ preserve_scale_schedule.append(scale_schedule[0])
+ target = raw_features[0]
+ if not infer_mode and args.log_norm_sigma > 0:
+ spt = torch.sigmoid(torch.randn(1, generator=torch_cuda_generator, device=target.device) * args.log_norm_sigma + args.log_norm_mean).item()
+ spt = shift_pt(spt, args.alpha)
+ else:
+ spt = shift_pt(numpy_generator.random(), args.alpha)
+
+ if args.refine_mode in ['ar_discrete_GRN_ind']:
+ from grn.utils_t2iv.hbq_util_t2iv import raw_feature2index_label
+ labels = raw_feature2index_label(target, hbq_round=args.hbq_round) # [B, hbq_round * d, t, h, w]
+ classes = 2**args.hbq_round
+ elif args.refine_mode in ['ar_discrete_GRN_bit']:
+ from grn.utils_t2iv.hbq_util_t2iv import raw_feature2bit_label
+ labels = raw_feature2bit_label(target, hbq_round=args.hbq_round) # [B, hbq_round * d, t, h, w]
+ classes = 2
+
+ random_labels = torch.randint(0, classes, size=labels.shape, generator=torch_cuda_generator, device=labels.device, dtype=labels.dtype) # random 0 or 1 labels
+ random_mask = torch.rand(size=labels.shape, generator=torch_cuda_generator, device=labels.device, dtype=target.dtype) < spt
+ mixed_xt = torch.where(random_mask, labels, random_labels) # [B, hbq_round * d, t, h, w]
+ precise_spt = random_mask.float().mean()
+ wandb_plot_index = min(9, int(precise_spt / 0.1)) # 0~9
+
+ # get visual rope
+ if not infer_mode:
+ if meta['first_frame_condition']:
+ 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)
+ 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))
+ visual_rope_cache_list[-1] = [torch.cat(visual_rope_cache_list[-1], dim=-2)]
+ else:
+ 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)
+
+ # get visual tokens
+ visual_token_dim = mixed_xt.shape[1] * classes
+
+ if not infer_mode and meta['first_frame_condition']:
+ first_frame_labels = labels[:,:,:1] # [B, hbq_round * d, 1, h, w]
+ first_frame_tokens = multiclass_labels2onehot_input(first_frame_labels, classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
+ cur_visual_tokens = multiclass_labels2onehot_input(mixed_xt[:,:,1:], classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
+ cur_visual_tokens = torch.cat((cur_visual_tokens, first_frame_tokens), dim=1)
+ indices = labels[:,:,1:]
+ else:
+ cur_visual_tokens = multiclass_labels2onehot_input(mixed_xt, classes).reshape(1, visual_token_dim, -1).permute(0, 2, 1)
+ indices = labels
+ indices = indices.type(torch.long).permute(0,2,3,4,1) # [B,d,t,h,w] -> [B,t,h,w,d]
+ gt_all_bit_indices[-1].append(indices)
+ var_input_list[-1].append(cur_visual_tokens)
+ other_info_by_scale[-1].append(
+ {
+ 'largest_scale': scale_schedule[-1],
+ 'wandb_plot_index': wandb_plot_index,
+ 'cur_bits': indices.shape[-1],
+ 'cur_lvl': args.detail_num_lvl,
+ 'scale_token_id': precise_spt,
+ 'predict_tokens': np.prod(scale_schedule[0]),
+ 'all_tokens': scale_lengths[-1] if len(scale_lengths) else -1,
+ 'first_frame_condition': meta['first_frame_condition'],
+ }
+ )
+ sequece_packing_scales.append(preserve_scale_schedule)
+
+ gt_all_bit_indices = flatten_two_level_list(gt_all_bit_indices)
+ pred_all_bit_indices = flatten_two_level_list(pred_all_bit_indices)
+ var_input_list = flatten_two_level_list(var_input_list)
+ visual_rope_cache_list = flatten_two_level_list(visual_rope_cache_list)
+ other_info_by_scale = flatten_two_level_list(other_info_by_scale)
+
+ if infer_mode:
+ return [labels, target], x_recon_raw, [target], None, None, None
+
+ gt_ms_idx_Bl = []
+ for item in gt_all_bit_indices:
+ _, tt, hh, ww, dd = item.shape
+ item = item.reshape(B, tt*hh*ww, dd)
+ gt_ms_idx_Bl.append(item)
+ gt_BLC = gt_ms_idx_Bl # torch.cat(gt_ms_idx_Bl, 1).contiguous().type(torch.long)
+ x_BLC = var_input_list
+ x_BLC_mask = None
+ scale_or_time_ids = None
+ 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
+
+def video_decode(
+ vae,
+ all_indices,
+ scale_schedule,
+ label_type,
+ args=None,
+ noise_list=None,
+ trunc_scales=-1,
+ **kwargs,
+):
+ if trunc_scales < 0:
+ summed_codes = all_indices[-1]
+ else:
+ summed_codes = all_indices[trunc_scales-1]
+ x_recon = vae.decode(summed_codes, slice=True)
+ x_recon = torch.clamp(x_recon, min=-1, max=1)
+ x_recon_256 = None
+ return x_recon, x_recon_256
+
+def get_visual_rope_embeds(rope2d_freqs_grid, scale_schedule, device=None, mapped_h_div_w_template=None, t_offset=0):
+ # freqs_frames: (2, max_frames, dim_div_2 / 3)
+ rope2d_freqs_grid['freqs_frames'] = rope2d_freqs_grid['freqs_frames'].to(device)
+ rope2d_freqs_grid['freqs_height'] = rope2d_freqs_grid['freqs_height'].to(device)
+ rope2d_freqs_grid['freqs_width'] = rope2d_freqs_grid['freqs_width'].to(device)
+ max_height = rope2d_freqs_grid['freqs_height'].shape[1]
+ max_width = rope2d_freqs_grid['freqs_width'].shape[1]
+ extreme_h_div_w = 3
+ assert mapped_h_div_w_template <= extreme_h_div_w
+ extreme_h = max_height
+ extreme_w = extreme_h / extreme_h_div_w
+ upw = np.sqrt(extreme_h * extreme_w / mapped_h_div_w_template)
+ uph = mapped_h_div_w_template * upw
+ uph, upw = int(uph), int(upw)
+ pt, ph, pw = scale_schedule
+ assert ph <= uph and pw <= upw
+ f_frames = rope2d_freqs_grid['freqs_frames'][:, t_offset:t_offset+pt]
+ f_height = rope2d_freqs_grid['freqs_height'][:, (torch.arange(ph) * (uph / ph)).round().int()]
+ f_width = rope2d_freqs_grid['freqs_width'][:, (torch.arange(pw) * (upw / pw)).round().int()]
+ rope_embeds = torch.cat([
+ f_frames[ :, :, None, None, :].expand(-1, -1, ph, pw, -1),
+ f_height[ :, None, :, None, :].expand(-1, pt,-1, pw, -1),
+ f_width[ :, None, None, :, :].expand(-1, pt,ph, -1, -1),
+ ], dim=-1) # (2, pt, ph, pw, dim_div_2)
+ rope_embeds = rope_embeds.reshape(2, 1, 1, 1, pt*ph*pw, -1) # (2, 1, 1, 1, pt*ph*pw, dim_div_2)
+ return rope_embeds
diff --git a/grn/tokenizer/.gitignore b/grn/tokenizer/.gitignore
new file mode 100644
index 0000000000000000000000000000000000000000..ef82a2fa3501a42c2c90a235a8517d7010544e9a
--- /dev/null
+++ b/grn/tokenizer/.gitignore
@@ -0,0 +1,25 @@
+**/__pycache__/
+lightning_logs/
+.ipynb_checkpoints/
+*.egg-info
+.pyc
+results*
+cmp_results*
+logs*
+recon.*
+*.pt
+*.npy
+dataset/
+**/span.log
+cli/batch_post.py
+model_arch*.txt
+*.mp4
+*.png
+labels
+video_vae_results
+video_vae_results_bk
+results
+video_vae_results_full
+bashrc_gpu_worker
+hj_video_vae_results
+wandb
diff --git a/grn/tokenizer/sample.py b/grn/tokenizer/sample.py
new file mode 100644
index 0000000000000000000000000000000000000000..7873c78c573ae3cf668736c35d6fc941d0c7b830
--- /dev/null
+++ b/grn/tokenizer/sample.py
@@ -0,0 +1,694 @@
+import os
+import tqdm
+import json
+import re
+import torch
+import torch.nn.functional as F
+import argparse
+import time
+import datetime
+import numpy as np
+import hashlib
+import random
+import torch.nn as nn
+from torchvision.models.inception import inception_v3
+from torch.profiler import record_function as torch_record_function
+from contextlib import nullcontext
+import lpips
+import cv2
+from einops import rearrange
+from tqdm import tqdm
+from PIL import Image
+import os.path as osp
+Image.MAX_IMAGE_PIXELS = None
+
+from videovae.modules.commitments import DiagonalGaussianDistribution
+
+import torch.distributed as dist
+from torch.multiprocessing import spawn
+from torch.nn.parallel import DistributedDataParallel as DDP
+
+import imageio
+import random
+from skimage.metrics import peak_signal_noise_ratio as psnr_loss
+from skimage.metrics import structural_similarity as ssim_loss
+
+from videovae.data import VideoData
+from videovae.utils.misc import save_video_grid, shift_dim, data_prefix_manager, rearranged_forward, seed_everything
+from videovae.utils.init_models import init_cnn_from_image, load_cnn
+from videovae.utils.arguments import MainArgs, add_model_specific_args, init_resolution
+from videovae.evaluation import get_fvd_logits, frechet_distance, load_fvd_model
+from videovae.evaluation import calculate_frechet_distance
+from videovae.evaluation import InceptionV3
+from videovae.evaluation import calculate_fvd, calculate_lpips, calculate_psnr, calculate_ssim
+
+torch.set_num_threads(32)
+os.environ["NCCL_DEBUG"] = "WARN"
+os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'
+
+
+def calculate_batch_codebook_usage_percentage(batch_encoding_indices,n_codes):
+ if isinstance(batch_encoding_indices, list):
+ all_indices = []
+ for one_encoding_indices in batch_encoding_indices:
+ all_indices.append(one_encoding_indices.flatten())
+ all_indices = torch.cat(all_indices, dim=0)
+ else:
+ # Flatten the batch of encoding indices into a single 1D tensor
+ all_indices = batch_encoding_indices.flatten()
+ all_indices = all_indices.detach().cpu()
+
+ # Obtain the total number of encoding indices in the batch to calculate percentages
+ total_indices = all_indices.numel()
+
+ # Initialize a tensor to store the percentage usage of each code
+ codebook_usage = torch.zeros(n_codes, dtype=torch.long)
+
+ # Count the number of occurrences of each index and get their frequency as percentages
+ unique_indices, counts = torch.unique(all_indices, return_counts=True)
+
+ # Populate the corresponding percentages in the codebook_usage_percentage tensor
+ codebook_usage[unique_indices.long()] = counts
+
+ return codebook_usage
+
+
+def disabled_train(self, mode=True):
+ """Overwrite model.train with this function to make sure train/eval mode
+ does not change anymore."""
+ return self
+
+def default_parse_args():
+ parser = argparse.ArgumentParser()
+ parser.add_argument('--vqgan_ckpt', type=str, default=None)
+ parser.add_argument('--sd_ckpt', type=str, default=None)
+ parser.add_argument('--use_frames', type=int, default=None)
+ parser.add_argument('--inference_type', type=str, choices=["image", "video", "video_concat"])
+ parser.add_argument('--save_prediction', action='store_true')
+ parser.add_argument('--save_dir', type=str, default="results")
+ parser.add_argument('--intermediate_tensor', action='store_true')
+ parser.add_argument('--save_z', action='store_true')
+ parser.add_argument('--save_frames', action='store_true')
+ parser.add_argument('--image_recon4video', action='store_true')
+ parser.add_argument('--junke_old', action='store_true')
+ parser.add_argument('--cal_norm', action='store_true')
+ parser.add_argument('--save_samples', type=str, default=None)
+ parser.add_argument('--device', type=str, default="cuda", choices=["cpu", "cuda"])
+ parser.add_argument('--noise_scale', type=float, default=0.0)
+ parser = MainArgs.add_main_args(parser)
+ parser = VideoData.add_data_specific_args(parser)
+ args, unknown = parser.parse_known_args()
+ args, parser, vae_model = add_model_specific_args(args, parser)
+ args = parser.parse_args()
+ return args, vae_model
+
+
+def setup(rank, world_size):
+ os.environ['MASTER_ADDR'] = 'localhost'
+ os.environ['MASTER_PORT'] = str(12355+int(time.time())%1000)
+ # dist.init_process_group("nccl", rank=rank, world_size=world_size)
+ dist.init_process_group("nccl", rank=rank, world_size=world_size, timeout=datetime.timedelta(seconds=30 * 60))
+
+def cleanup():
+ dist.destroy_process_group()
+
+def main():
+ args, vae_model = default_parse_args()
+ assert len(args.dataset_list) == 1
+
+ # init data_prefix_manager
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
+ args.default_root_dir = data_prefix_manager(args.default_root_dir)
+ os.makedirs(args.default_root_dir, exist_ok=True)
+ print(args.default_root_dir)
+
+ # init intermediate_tensor_dir
+ if args.intermediate_tensor:
+ random.seed(time.time())
+ random_folder_name = hashlib.sha256(str(random.random()).encode('utf-8')).hexdigest()[:16]
+ args.intermediate_tensor_dir = os.path.join(args.default_root_dir, random_folder_name)
+ print(f"save temporal tensor to {args.intermediate_tensor_dir}")
+
+ seed_everything(seed=0, allow_tf32=True) # ALERT: allow_tf32=True may cause accumulate error in conv3d forward >
+
+ # init resolution
+ args.resolution = init_resolution(args.resolution, len(args.dataset_list))
+
+ # init profiler
+ def trace_handler(p):
+ p.export_chrome_trace(os.path.join(args.default_root_dir, f"trace_step_{p.step_num}_rank_{0}.json"))
+
+ tp = None
+ if args.turn_on_profiler:
+ tp = torch.profiler.profile(
+ activities=[
+ torch.profiler.ProfilerActivity.CPU,
+ torch.profiler.ProfilerActivity.CUDA,
+ ],
+ schedule=torch.profiler.schedule(
+ wait=args.profiler_scheduler_wait_steps,
+ warmup=3,
+ active=2,
+ repeat=1,
+ ),
+ with_stack=True,
+ record_shapes=True,
+ profile_memory=True,
+ on_trace_ready=trace_handler
+ )
+ tp.start()
+ record_function = torch_record_function
+ else:
+ record_function = nullcontext
+
+
+ vae = None
+ use_vae = None
+ num_codes = None
+ if args.vqgan_ckpt:
+ args.vqgan_ckpt = data_prefix_manager(args.vqgan_ckpt)
+ if args.tokenizer in ["hbq_tokenizer"]:
+ vae = vae_model(args)
+ state_dict = torch.load(args.vqgan_ckpt, map_location=torch.device("cpu"), weights_only=True)
+ new_state_dict = {}
+ for key in ['vae', 'ema']:
+ if (key not in state_dict) or (not state_dict[key]):
+ continue
+ if 'quantizer.scale_learnable_parameters' in state_dict[key]:
+ if len(state_dict[key]['quantizer.scale_learnable_parameters']) == 1:
+ state_dict[key]['quantizer.scale_learnable_parameters'] = state_dict[key]['quantizer.scale_learnable_parameters'].expand(4)
+ state_dict[key]['scale_learnable_parameters'] = state_dict[key]['quantizer.scale_learnable_parameters']
+ del state_dict[key]['quantizer.scale_learnable_parameters']
+ if 'z_mean' in state_dict[key]:
+ if state_dict[key]['z_mean'].shape != vae.z_mean.shape:
+ del state_dict[key]['z_mean']
+ del state_dict[key]['z_std']
+ new_state_dict[key] = state_dict[key]
+ slim_model_path = args.vqgan_ckpt.replace('/checkpoints/', f'/slim_{key}/')
+ if not osp.exists(slim_model_path):
+ os.makedirs(os.path.dirname(slim_model_path), exist_ok=True)
+ torch.save({key: state_dict[key]}, slim_model_path)
+ print(f'save to {slim_model_path}')
+
+ if args.ema == "yes":
+ print("testing ema weights")
+ print(vae.load_state_dict(new_state_dict["ema"], strict=False))
+ else:
+ print("testing non ema weights")
+ print(vae.load_state_dict(new_state_dict["vae"], strict=False))
+ for name, param in vae.named_parameters():
+ if name.startswith("scale_learnable_"):
+ try:
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
+ except:
+ print(f"{name}: {param[:32].cpu().detach().reshape(-1).tolist()}")
+ for name, param in vae.named_buffers():
+ if name.startswith("scale_learnable_"):
+ try:
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
+ except:
+ print(f"{name}: {param[:32].cpu().detach().reshape(-1).tolist()}")
+ if ("scale_wise_std_" in name) or ("scale_wise_mean_" in name):
+ print(f"{name}: {param[:32,0,0].cpu().detach().reshape(-1).tolist()}")
+ if ('signal_' in name):
+ print(f"{name}: {param.cpu().detach().reshape(-1).tolist()}")
+ if args.tokenizer != 'hbq_tokenizer':
+ vae.enable_slicing()
+ # vae.enable_tiling()
+ else:
+ raise NotImplementedError
+
+ if args.inference_type == "video":
+ def extract_results(return_dict, world_size):
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = [], [], [], [], []
+ if args.intermediate_tensor:
+ for rank in range(world_size):
+ real_embeddings.append(return_dict[rank]['real_embeddings'])
+ fake_embeddings.append(return_dict[rank]['fake_embeddings'])
+ all_real_videos += return_dict[rank]['all_real_videos']
+ all_fake_videos += return_dict[rank]['all_fake_videos']
+ zs.append(return_dict[rank]['zs'])
+ real_embeddings = torch.cat(real_embeddings, 0).to('cuda:0')
+ fake_embeddings = torch.cat(fake_embeddings, 0).to('cuda:0')
+ zs = torch.cat(zs, 0).to('cuda:0')
+ else:
+ for rank in range(world_size):
+ real_embeddings.append(return_dict[rank]['real_embeddings'])
+ fake_embeddings.append(return_dict[rank]['fake_embeddings'])
+ all_real_videos.append(return_dict[rank]['all_real_videos'])
+ all_fake_videos.append(return_dict[rank]['all_fake_videos'])
+ zs.append(return_dict[rank]['zs'])
+ real_embeddings = torch.cat(real_embeddings, 0).to('cuda:0')
+ fake_embeddings = torch.cat(fake_embeddings, 0).to('cuda:0')
+ all_real_videos = torch.cat(all_real_videos, 0)
+ all_fake_videos = torch.cat(all_fake_videos, 0)
+ zs = torch.cat(zs, 0).to('cuda:0')
+ return real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs
+
+ def inference(mean=None, std=None, noise_scale=0):
+ world_size = torch.cuda.device_count()
+ manager = torch.multiprocessing.Manager()
+ return_dict = manager.dict()
+ ### multi-process
+ # try:
+ # 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)
+ # except Exception as e:
+ # print(f"Error during spawn {e}")
+
+ ## single process
+ world_size = 1
+ 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)
+
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = extract_results(return_dict, world_size)
+ return real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs
+
+ def cal_std(zs):
+ dims_to_reduce = [i for i in range(zs.dim()) if i != 1]
+ total_std = zs.std().item()
+ _mean = zs.mean(dim=dims_to_reduce)
+ _std = zs.std(dim=dims_to_reduce)
+ return total_std, _mean, _std
+
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = inference()
+ if args.noise_scale > 0:
+ total_std, _mean, _std = cal_std(zs)
+ real_embeddings, fake_embeddings, all_real_videos, all_fake_videos, zs = inference(mean=_mean, std=_std, noise_scale=args.noise_scale)
+
+ if args.save_samples:
+ torch.save(zs.cpu(), args.save_samples)
+
+ if args.cal_norm:
+ total_std, _mean, _std = cal_std(zs)
+ print(f"{total_std = } {_mean = } {_std = }")
+ if args.save_prediction:
+ fname = os.path.join(args.save_dir, args.dataset_list[0], "gt_recon", "mean_std.pth")
+ torch.save({'_mean': _mean, '_std': _std}, fname)
+
+ result_str = video_eval(real_embeddings, fake_embeddings, all_real_videos, all_fake_videos)
+ else:
+ world_size = 1 if args.debug else torch.cuda.device_count()
+ manager = torch.multiprocessing.Manager()
+ return_dict = manager.dict()
+
+ if args.debug:
+ inference_eval(0, world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict)
+ else:
+ spawn(inference_eval, args=(world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict), nprocs=world_size, join=True)
+
+ 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, []
+ for rank in range(world_size):
+ pred_xs.append(return_dict[rank]['pred_xs'])
+ pred_recs.append(return_dict[rank]['pred_recs'])
+ lpips_alex += return_dict[rank]['lpips_alex']
+ lpips_vgg += return_dict[rank]['lpips_vgg']
+ ssim_value += return_dict[rank]['ssim_value']
+ psnr_value += return_dict[rank]['psnr_value']
+ num_iter += return_dict[rank]['num_iter']
+ total_usage += return_dict[rank]['total_usage']
+ pred_xs = np.concatenate(pred_xs, 0)
+ pred_recs = np.concatenate(pred_recs, 0)
+
+ 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)
+ # result_str = inference_eval(args, vae_model, vae, record_function, use_vae, num_codes)
+
+ print(f"noise scale = {args.noise_scale}")
+ print(result_str)
+ # save result_str to exp_dir
+ basename = os.path.basename(args.vqgan_ckpt)
+ match = re.search(r'model_step_(\d+)\.ckpt', basename)
+ iter_num = match.group(1) if match else None
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
+ ckpt_dir = os.path.dirname(data_prefix_manager(args.vqgan_ckpt))
+ use_frames = args.use_frames if args.use_frames else args.sequence_length
+ save_dir = os.path.join(ckpt_dir, "evaluation", args.dataset_list[0], f"{args.resolution[0][0]}_{args.resolution[0][1]}", f"{use_frames}")
+ os.makedirs(save_dir, exist_ok=True)
+ ema_suffix = "_ema" if args.ema == "yes" else ""
+ result_name = os.path.join(save_dir, f"result_{iter_num}{ema_suffix}.txt")
+ if (not args.save_prediction) and (args.noise_scale == 0):
+ with open(result_name, "w") as f:
+ f.write(result_str)
+ # print('Usage = %.2f'%((total_usage > 0.).sum() / num_codes))
+ if args.intermediate_tensor:
+ os.system(f"rm -rf {args.intermediate_tensor_dir}")
+
+def add_noise(z, mean, std, noise_scale):
+ if noise_scale > 0:
+ mean = mean.view(1, mean.shape[0], 1, 1, 1).to(z.device)
+ std = std.view(1, std.shape[0], 1, 1, 1).to(z.device)
+ z = (z - mean) / std
+ noise = torch.randn(z.size()).to(z.device)
+ z = (z + noise * noise_scale) * std + mean
+ return z
+
+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):
+ setup(rank, world_size)
+ # init data_prefix_manager
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
+
+ for param in vae.parameters():
+ param.requires_grad = False
+ vae = vae.eval()
+ vae = vae.to(f"cuda:{rank}")
+ # vae = torch.compile(vae)
+
+ save_dir = os.path.join(args.save_dir, args.dataset_list[0])
+ print('generating and saving video to %s...'%save_dir)
+ os.makedirs(save_dir, exist_ok=True)
+
+ data = VideoData(args)
+ loader = data.val_dataloader()
+
+ i3d = load_fvd_model(f"cuda:{rank}")
+
+ os.makedirs(os.path.join(save_dir, "gt"), exist_ok=True)
+ os.makedirs(os.path.join(save_dir, "recons"), exist_ok=True)
+
+ zs = []
+ real_embeddings = []
+ fake_embeddings = []
+
+ all_real_videos = []
+ all_fake_videos = []
+
+ num_videos = len(loader)
+ loader_iter = iter(loader)
+ progress_bar = tqdm(total=num_videos, desc=f"Testing {num_videos} batches")
+ for batch_idx in range(num_videos):
+ if args.turn_on_profiler and tp:
+ tp.step()
+ batch = next(loader_iter)
+ with torch.no_grad():
+ input_ = batch['video'] # B C T H W
+ B = input_.shape[0]
+ if args.tokenizer in ["hbq_tokenizer"]:
+ input_ = input_.to(f"cuda:{rank}").to(torch.bfloat16)
+ with torch.amp.autocast("cuda", dtype=torch.bfloat16):
+ x_raw, x_recons, z = vae(input_, 0, is_train=False)
+ batch['video'] = x_raw.to('cpu').to(torch.float32)
+ x_recons = x_recons.to(torch.float32)
+ else:
+ raise NotImplementedError
+
+ if args.tokenizer in ["icvivit", "sd"]:
+ x_recons = rearrange(x_recons, "(b t) c h w -> b c t h w", b=B)
+
+ real_videos = torch.clamp(batch['video'] / 2 + 0.5, 0, 1)
+ if args.junke_old:
+ fake_videos = torch.clamp(x_recons.detach().cpu() + 0.5, 0, 1)
+ else:
+ fake_videos = torch.clamp(x_recons.detach().cpu() / 2 + 0.5, 0, 1)
+
+ use_frames = args.use_frames if args.use_frames else args.sequence_length
+ if args.intermediate_tensor:
+ folder_name = os.path.join(args.intermediate_tensor_dir, f"{rank}_{batch_idx}")
+ os.makedirs(folder_name, exist_ok=True)
+ real_file = os.path.join(folder_name, "real_videos.pt")
+ fake_file = os.path.join(folder_name, "fake_videos.pt")
+ real_videos = real_videos[:,:,:use_frames,...]
+ fake_videos = fake_videos[:,:,:use_frames,...]
+ torch.save(real_videos.permute(0, 2, 1, 3, 4).squeeze(0), real_file)
+ torch.save(fake_videos.permute(0, 2, 1, 3, 4).squeeze(0), fake_file)
+ all_real_videos.append(real_file)
+ all_fake_videos.append(fake_file)
+ else:
+ real_videos = real_videos[:,:,:use_frames,...]
+ fake_videos = fake_videos[:,:,:use_frames,...]
+ all_real_videos.append(real_videos.clone())
+ all_fake_videos.append(fake_videos.clone())
+ if args.cal_norm or args.save_samples or args.noise_scale > 0:
+ zs.append(z)
+ real_embedding = get_fvd_logits(shift_dim(real_videos * 255, 1, -1).byte().data.numpy(), i3d=i3d, device=f"cuda:{rank}").cpu()
+ real_embeddings.append(real_embedding)
+ fake_embedding = get_fvd_logits(shift_dim(fake_videos * 255, 1, -1).byte().data.numpy(), i3d=i3d, device=f"cuda:{rank}").cpu()
+ fake_embeddings.append(fake_embedding)
+
+ if args.tokenizer in ['cvivit', "icvivit"] and not use_vae:
+ batch_codebook_usage = vq_output["batch_usage"]
+ total_usage += batch_codebook_usage
+
+ if args.save_prediction:
+ video = torch.cat([real_videos[:,:,:fake_videos.shape[2],:,:], fake_videos], dim=-1)
+ b, c, t, h, w = video.shape
+ video = video.permute(0, 2, 3, 4, 1).contiguous()
+ video = (video.squeeze().detach().cpu().numpy() * 255).astype('uint8')
+ os.makedirs(os.path.join(save_dir, "gt_recon"), exist_ok=True)
+ this_filename = batch["path"][0].split('/')[-1]
+ fname = os.path.join(save_dir, "gt_recon", this_filename)
+ import imageio
+ imageio.mimsave(fname, video, fps=15)
+
+ if args.save_z:
+ os.makedirs(os.path.join(save_dir, "gt_recon"), exist_ok=True)
+ this_filename = batch["path"][0].split('/')[-1].split(".")[0]
+ fname = os.path.join(save_dir, "gt_recon", this_filename+".pt")
+ torch.save(z, fname)
+
+ if args.save_frames:
+
+ def convert_to_uint8(image):
+ return (image.detach().cpu().numpy() * 255).astype(np.uint8)
+
+ # artifact_grid_size = 32
+ assert real_videos.shape == fake_videos.shape, f"shape of gt and predicted videos are not equal"
+ 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]}"
+ _real_videos = real_videos.squeeze(0)
+ _fake_videos = fake_videos.squeeze(0)
+ # h, w = real_videos.shape[-2:]
+ # 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}"
+
+ frame_num = _real_videos.shape[1]
+ for frame_idx in range(frame_num):
+ real_image = _real_videos[:,frame_idx,:,:]
+ fake_image = _fake_videos[:,frame_idx,:,:]
+ # most_different_top_left, max_difference = find_most_different_patch(real_image, fake_image, artifact_grid_size)
+ real_image_uint8 = convert_to_uint8(real_image)
+ predicted_image_uint8 = convert_to_uint8(fake_image)
+
+ real_image_bgr = cv2.cvtColor(real_image_uint8.transpose(1, 2, 0), cv2.COLOR_RGB2BGR)
+ predicted_image_bgr = cv2.cvtColor(predicted_image_uint8.transpose(1, 2, 0), cv2.COLOR_RGB2BGR)
+
+ concatenated_image = np.concatenate((real_image_bgr, predicted_image_bgr), axis=1)
+ fname = os.path.join(save_dir, "gt_recon", f"{this_filename}_{frame_idx}.png")
+ cv2.imwrite(fname, concatenated_image)
+ progress_bar.update(1)
+
+ real_embeddings = torch.cat(real_embeddings, 0)
+ fake_embeddings = torch.cat(fake_embeddings, 0)
+ zs = torch.cat(zs, 0) if len(zs) > 0 else torch.tensor([])
+ if args.intermediate_tensor:
+ temp_dict = {
+ 'real_embeddings':real_embeddings.cpu(),
+ 'fake_embeddings':fake_embeddings.cpu(),
+ 'all_real_videos':all_real_videos,
+ 'all_fake_videos':all_fake_videos,
+ 'zs': zs.cpu(),
+ }
+ else:
+ all_real_videos = torch.cat(all_real_videos, 0).permute(0, 2, 1, 3, 4)
+ all_fake_videos = torch.cat(all_fake_videos, 0).permute(0, 2, 1, 3, 4)
+ temp_dict = {
+ 'real_embeddings':real_embeddings.cpu(),
+ 'fake_embeddings':fake_embeddings.cpu(),
+ 'all_real_videos':all_real_videos.cpu(),
+ 'all_fake_videos':all_fake_videos.cpu(),
+ 'zs': zs.cpu(),
+ }
+ # if dist.is_initialized():
+ # dist.barrier()
+ return_dict[rank] = temp_dict
+ cleanup()
+
+
+def video_eval(real_embeddings, fake_embeddings, all_real_videos, all_fake_videos):
+ fake_embeddings = fake_embeddings.to(torch.float64)
+ real_embeddings = real_embeddings.to(torch.float64)
+ FVD = frechet_distance(fake_embeddings, real_embeddings)
+ print(f"FVD: {FVD}") # can't wait to see this number :)
+ del real_embeddings, fake_embeddings
+
+ lpips = calculate_lpips(all_real_videos, all_fake_videos, device="cuda")["value"].values()
+ psnr = calculate_psnr(all_real_videos, all_fake_videos)["value"].values()
+ ssim = calculate_ssim(all_real_videos, all_fake_videos)["value"].values()
+ lpips = np.mean(np.stack(list(lpips)))
+ ssim = np.mean(np.stack(list(ssim)))
+ psnr = np.mean(np.stack(list(psnr)))
+
+ result_str = f"""
+ FVD = {FVD:.4f}
+ LPIPS = {lpips:.4f}
+ SSIM = {ssim:.4f}
+ PSNR = {psnr:.3f}
+ """
+ return result_str
+
+def inference_eval(rank, world_size, args, vae_model, vae, record_function, use_vae, num_codes, return_dict):
+ # Don't remove this setup!!! dist.init_process_group is important for building loader (data.distributed.DistributedSampler)
+ setup(rank, world_size)
+ # init data_prefix_manager
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
+
+ device = torch.device(f"cuda:{rank}")
+
+ for param in vae.parameters():
+ param.requires_grad = False
+ vae.to(device).eval()
+
+ save_dir = os.path.join(args.save_dir, args.dataset_list[0])
+ print('generating and saving video to %s...'%save_dir)
+ os.makedirs(save_dir, exist_ok=True)
+
+ data = VideoData(args)
+
+ loader = data.val_dataloader()
+
+ dims = 2048
+ block_idx = InceptionV3.BLOCK_INDEX_BY_DIM[dims]
+ inception_model = InceptionV3([block_idx]).to(device)
+ inception_model.eval()
+
+ loader_iter = iter(loader)
+
+ pred_xs = []
+ pred_recs = []
+ # LPIPS score related
+ loss_fn_alex = lpips.LPIPS(net='alex').to(device) # best forward scores
+ loss_fn_vgg = lpips.LPIPS(net='vgg').to(device) # closer to "traditional" perceptual loss, when used for optimization
+ lpips_alex = 0.0
+ lpips_vgg = 0.0
+
+ # SSIM score related
+ ssim_value = 0.0
+
+ # PSNR score related
+ psnr_value = 0.0
+
+ num_images = len(loader)
+ print(f"Testing {num_images} files")
+ num_iter = 0
+
+ total_usage = 0.0
+ total_usage_bit = 0.0
+ total_num_token = 0
+ for batch_idx in tqdm(range(num_images)):
+ batch = next(loader_iter)
+
+ with torch.no_grad():
+ x = batch['video']
+ if args.tokenizer in ["hbq_tokenizer"]:
+ x_raw, x_recons, z = vae(x.to(device), 0, is_train=False)
+ x_recons = x_recons.squeeze(-3).cpu()
+ else:
+ raise NotImplementedError
+
+ if args.image_recon4video:
+ # convert back to image format
+ x = x.squeeze(2)
+ x_recons = x_recons.squeeze(2)
+
+ if args.tokenizer in ["cvivit", "icvivit"] and not use_vae:
+
+ # encoding_indices = vq_output["encodings"].detach().cpu()
+ code_counts = calculate_batch_codebook_usage_percentage(vq_output["encodings"], num_codes)
+ total_counts += code_counts
+
+ batch_codebook_usage = vq_output["batch_usage"]
+ total_usage += batch_codebook_usage
+
+ paths = batch["path"]
+ assert len(paths) == x.shape[0]
+
+ for p, input_ori, recon_ori in zip(paths, x, x_recons):
+ if os.path.isabs(p):
+ p = "/".join(p.split("/")[6:])
+ assert not os.path.isabs(p), f"{p} should not be abspath"
+ path = os.path.join(save_dir, "input_recon", os.path.basename(p))
+ os.makedirs(os.path.split(path)[0], exist_ok=True)
+
+ input_ori = input_ori.unsqueeze(0).to(device)
+ input_ = (input_ori + 1) / 2 # [0, 1]
+
+ pred_x = inception_model(input_)[0]
+ pred_x = pred_x.squeeze(3).squeeze(2).cpu().numpy()
+
+ recon_ori = recon_ori.unsqueeze(0).to(device)
+ recon_ = (recon_ori + 1) / 2 # [0, 1]
+ # recon_ = recon_.permute(1, 2, 0).detach().cpu()
+ with torch.no_grad():
+ pred_rec = inception_model(recon_)[0]
+ pred_rec = pred_rec.squeeze(3).squeeze(2).cpu().numpy()
+ if args.save_prediction:
+ if input_.dim() == 4:
+ input_image = input_.squeeze(0)
+ if recon_.dim() == 4:
+ recon_image = recon_.squeeze(0)
+ input_recon = torch.cat([input_image, recon_image], dim=-1)
+ input_recon = Image.fromarray((torch.clamp(input_recon.permute(1, 2, 0).detach().cpu(), 0, 1).numpy() * 255).astype(np.uint8))
+ input_recon.save(path)
+
+ pred_xs.append(pred_x)
+ pred_recs.append(pred_rec)
+
+ # calculate lpips
+ with torch.no_grad():
+ lpips_alex += loss_fn_alex(input_ori, recon_ori).sum() # [-1, 1]
+ lpips_vgg += loss_fn_vgg(input_ori, recon_ori).sum() # [-1, 1]
+
+ #calculate PSNR and SSIM
+ rgb_restored = (recon_ * 255.0).permute(0, 2, 3, 1).to("cpu", dtype=torch.uint8).numpy()
+ rgb_gt = (input_ * 255.0).permute(0, 2, 3, 1).to("cpu", dtype=torch.uint8).numpy()
+ rgb_restored = rgb_restored.astype(np.float32) / 255.
+ rgb_gt = rgb_gt.astype(np.float32) / 255.
+ ssim_temp = 0
+ psnr_temp = 0
+ B, _, _, _ = rgb_restored.shape
+ for i in range(B):
+ rgb_restored_s, rgb_gt_s = rgb_restored[i], rgb_gt[i]
+ with torch.no_grad():
+ ssim_temp += ssim_loss(rgb_restored_s, rgb_gt_s, data_range=1.0, channel_axis=-1)
+ psnr_temp += psnr_loss(rgb_gt, rgb_restored)
+ ssim_value += ssim_temp / B
+ psnr_value += psnr_temp / B
+ num_iter += 1
+
+ pred_xs = np.concatenate(pred_xs, axis=0)
+ pred_recs = np.concatenate(pred_recs, axis=0)
+ temp_dict = {
+ 'pred_xs':pred_xs,
+ 'pred_recs':pred_recs,
+ 'lpips_alex':lpips_alex.cpu(),
+ 'lpips_vgg':lpips_vgg.cpu(),
+ 'ssim_value': ssim_value,
+ 'psnr_value': psnr_value,
+ 'num_iter': num_iter,
+ 'total_usage': total_usage,
+ 'total_usage_bit': total_usage_bit,
+ 'total_num_token': total_num_token,
+ }
+ return_dict[rank] = temp_dict
+
+ # if dist.is_initialized():
+ # dist.barrier()
+ cleanup()
+
+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):
+ mu_x = np.mean(pred_xs, axis=0)
+ sigma_x = np.cov(pred_xs, rowvar=False)
+ mu_rec = np.mean(pred_recs, axis=0)
+ sigma_rec = np.cov(pred_recs, rowvar=False)
+
+ fid_value = calculate_frechet_distance(mu_x, sigma_x, mu_rec, sigma_rec)
+ lpips_alex_value = lpips_alex / num_iter
+ lpips_vgg_value = lpips_vgg / num_iter
+ ssim_value = ssim_value / num_iter
+ psnr_value = psnr_value / num_iter
+
+ result_str = f"""
+ FID = {fid_value:.4f}
+ LPIPS_VGG: {lpips_vgg_value.item():.4f}
+ LPIPS_ALEX: {lpips_alex_value.item():.4f}
+ SSIM: {ssim_value:.4f}
+ PSNR: {psnr_value:.3f}
+ """
+ return result_str
+if __name__ == '__main__':
+ main()
\ No newline at end of file
diff --git a/grn/tokenizer/train.py b/grn/tokenizer/train.py
new file mode 100644
index 0000000000000000000000000000000000000000..08bf4961dc922aa55a1a874273155f1caf751ea5
--- /dev/null
+++ b/grn/tokenizer/train.py
@@ -0,0 +1,561 @@
+from json import load
+import os
+import argparse
+import math
+import glob
+import time
+import logging
+from distutils.util import strtobool
+from copy import deepcopy
+import gc
+gc.disable()
+import os.path as osp
+
+import torch
+import torch.nn.functional as F
+import torch.optim as optim
+import torch.distributed as dist
+from torch.profiler import record_function as torch_record_function
+from contextlib import nullcontext
+from torch.nn.parallel import DistributedDataParallel as DDP
+from safetensors.torch import load_file
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.fsdp import StateDictType, FullStateDictConfig
+
+from videovae.utils.misc import data_prefix_manager, COLOR_BLUE, COLOR_RESET, is_torch_optim_sch
+from videovae.utils.distributed import init_distributed_mode, reduce_losses, average_losses, _FSDP
+from videovae.utils.ema import update_ema, requires_grad
+
+from videovae.models.discriminator import ImageDiscriminator, VideoDiscriminator
+from videovae.data import VideoData
+from videovae.modules import build_lpips_model
+from videovae.modules.loss import get_disc_loss, adopt_weight
+from videovae.utils.misc import get_last_ckpt, seed_everything, print_gpu_usage, print_model_summary, version_checker
+from videovae.utils.init_models import init_vae_only, init_vit_from_image, resume_from_ckpt, init_cnn_from_image, load_cnn
+from videovae.utils.nan_detector import NanDetector
+from videovae.utils.arguments import MainArgs, add_model_specific_args, init_args, format_args
+from videovae.utils.scheduler import get_lambda
+from videovae.utils.mfu import register_mfu_hook, get_mfu, get_tflops, get_tflops_dict
+
+from videovae.utils.context_parallel import ContextParallelUtils as cp
+
+def save_model(fsdp_model, rank, model_path, global_step):
+ # FSDP推荐用 state_dict_type=FULL_STATE_DICT 来保存
+ with FSDP.state_dict_type(fsdp_model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)):
+ state_dict = fsdp_model.state_dict()
+ os.makedirs(os.path.dirname(model_path), exist_ok=True)
+ torch.save({'vae': state_dict, 'step': global_step}, model_path)
+ print(f"模型已保存到 {model_path}")
+
+# enable_timeline_sdk = strtobool(os.getenv("GenAI_USE_TIMELINE_SDK", "0"))
+enable_timeline_sdk = False
+
+def split_to_ranks(x):
+ bs = x.shape[0]
+ cp_size = cp.get_cp_size()
+ if cp_size > 1 and bs % cp_size == 0:
+ cp_rank = cp.get_cp_rank()
+ return x.chunk(cp_size, dim=0)[cp_rank]
+ else:
+ return x
+
+if enable_timeline_sdk:
+ try:
+ import bytedance.ndtimeline as ndtimeline
+ except ImportError:
+ print(f"import vescale.ndtimeline failed, skipped")
+ enable_timeline_sdk = False
+
+def init_data_scheduler(video_ranks_ratio: float = -1.0, cp_size: int = 1):
+ if video_ranks_ratio < 0:
+ return None,None
+
+ cp_size = max(1, cp_size)
+
+ rank = torch.distributed.get_rank()
+ world_size = torch.distributed.get_world_size()
+
+ video_ranks = list(range(int((world_size * video_ranks_ratio) // cp_size) * cp_size)) # align to cp_size for video
+ image_ranks = list(range(len(video_ranks), world_size))
+
+ print(f"[info] video_ranks: {video_ranks}, image_ranks: {image_ranks}")
+
+ if rank in image_ranks:
+ group = torch.distributed.new_group(image_ranks)
+ dataset_type_on_this_rank = "image"
+ else:
+ group = torch.distributed.new_group(video_ranks)
+ dataset_type_on_this_rank = "video"
+
+ return group, dataset_type_on_this_rank
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser = MainArgs.add_main_args(parser)
+ parser = VideoData.add_data_specific_args(parser)
+ args, unknown = parser.parse_known_args()
+ args, parser, vae_model = add_model_specific_args(args, parser)
+ args = parser.parse_args()
+
+ args = init_args(args) # post process args
+
+ # init data_prefix_manager
+ data_prefix_manager.set_data_root(args.data_root, username=args.username)
+ args.default_root_dir = data_prefix_manager(args.default_root_dir)
+
+ # Setup DDP:
+ init_distributed_mode(args)
+ rank = dist.get_rank()
+ world_size = dist.get_world_size()
+ device = rank % torch.cuda.device_count()
+ seed_everything(args.seed)
+ torch.cuda.set_device(device)
+
+ # init context parallel
+ cp_cfg = {"cp_size": args.context_parallel_size}
+ cp.initialize_context_parallel(cp_cfg)
+
+ ds_group, ds_type = init_data_scheduler(args.video_ranks_ratio, cp_size = args.context_parallel_size)
+
+ # Setup an experiment folder:
+ checkpoint_dir = f"{args.default_root_dir}/checkpoints" # Stores saved model checkpoints
+ os.makedirs(checkpoint_dir, exist_ok=True)
+ if rank == 0:
+ script_str = format_args(args)
+ with open(os.path.join(args.default_root_dir, "script.sh"), "w") as f:
+ f.write(script_str)
+ print(f"{COLOR_BLUE}Experiment directory created at {args.default_root_dir}{COLOR_RESET}")
+
+ import wandb
+ wandb_project = "HBQ_Tokenizer"
+ wandb.init(
+ project=wandb_project,
+ name=os.path.basename(os.path.normpath(args.default_root_dir)),
+ dir=args.default_root_dir,
+ config=args,
+ mode="offline" if args.debug else "online"
+ )
+
+ # init model
+ vae = vae_model(args).to(device)
+ if rank == 0:
+ model_arch_save_path = os.path.join(args.default_root_dir, "model_arch.txt")
+ print(f"{COLOR_BLUE}Logging model architecture at {model_arch_save_path}{COLOR_RESET}")
+ with open(model_arch_save_path, "w") as f:
+ f.write(str(vae))
+ image_disc = ImageDiscriminator(args).to(device)
+ video_disc = VideoDiscriminator(args).to(device)
+
+ # init optimizers and schedulers
+ if args.optim_type == "Adam":
+ vae_optim = torch.optim.Adam
+ elif args.optim_type == "AdamW":
+ vae_optim = torch.optim.AdamW
+ if args.disc_optim_type is None:
+ disc_optim = vae_optim
+ elif args.disc_optim_type == "rmsprop":
+ disc_optim = torch.optim.RMSprop
+
+ def get_param_groups(model):
+ decay = []
+ no_decay = []
+ for name, param in model.named_parameters():
+ if param.requires_grad:
+ if len(param.shape) == 1 or name.endswith(".bias") or ('scale_learnable_parameters' in name):
+ no_decay.append(param)
+ print(f'disable weight deacy for {name}')
+ else:
+ decay.append(param)
+ optimizer_grouped_parameters = [
+ {'params': decay, 'weight_decay': 0.01},
+ {'params': no_decay, 'weight_decay': 0.0}
+ ]
+ return optimizer_grouped_parameters
+
+ opt_vae = vae_optim(get_param_groups(vae), lr=args.lr, betas=(args.beta1, args.beta2))
+ if disc_optim == torch.optim.RMSprop:
+ opt_image_disc = disc_optim(image_disc.parameters(), lr=args.lr * args.dis_lr_multiplier)
+ opt_video_disc = disc_optim(video_disc.parameters(), lr=args.lr * args.dis_lr_multiplier)
+ else:
+ opt_image_disc = disc_optim(image_disc.parameters(), lr=args.lr * args.dis_lr_multiplier, betas=(args.beta1, args.beta2))
+ opt_video_disc = disc_optim(video_disc.parameters(), lr=args.lr * args.dis_lr_multiplier, betas=(args.beta1, args.beta2))
+
+ if args.scheduler == "no":
+ sch_vae, sch_image_disc, sch_video_disc = None, None, None
+ else:
+ lr_lambda = get_lambda(args)
+ sch_vae = optim.lr_scheduler.LambdaLR(opt_vae, lr_lambda)
+ sch_image_disc = optim.lr_scheduler.LambdaLR(opt_image_disc, lr_lambda)
+ sch_video_disc = optim.lr_scheduler.LambdaLR(opt_video_disc, lr_lambda)
+
+ ### ema
+ ema = None
+ if args.ema == "yes":
+ ema = deepcopy(vae).to(device) # Create an EMA of the model for use after training
+ requires_grad(ema, False)
+ print(f"EMA Parameters: {sum(p.numel() for p in ema.parameters()):,}")
+ update_ema(ema, vae, decay=0) # Ensure EMA is initialized with synced weights
+ ema.eval() # EMA model should always be in eval mode
+
+ model_optims = {
+ "vae" : vae,
+ "image_disc" : image_disc,
+ "video_disc" : video_disc,
+ "opt_vae" : opt_vae,
+ "opt_image_disc" : opt_image_disc,
+ "opt_video_disc" : opt_video_disc,
+ "sch_vae" : sch_vae,
+ "sch_image_disc" : sch_image_disc,
+ "sch_video_disc" : sch_video_disc,
+ "ema": ema,
+ }
+
+ ### Resume from checkpoint in default_root_dir or load pretrained weights if specified
+ ckpt_path = None
+ assert not args.default_root_dir is None # required argument
+ ckpt_path = get_last_ckpt(args.default_root_dir)
+ init_step = 0
+ if ckpt_path:
+ print(f"Resuming from {ckpt_path}")
+ state_dict = torch.load(ckpt_path, map_location="cpu")
+ 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)
+ elif args.pretrained is not None:
+ args.pretrained = data_prefix_manager(args.pretrained)
+ # read weight
+ state_dict = torch.load(args.pretrained, map_location="cpu", weights_only=True)
+
+ if args.pretrained_ema == "yes":
+ state_dict["vae"] = state_dict["ema"] # replace vae weights with ema weight
+ # load model
+ if args.pretrained_mode == "weights":
+ 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
+ del state_dict
+ else:
+ raise NotImplementedError
+ print(f"Successfully loaded ckpt {args.pretrained}, pretrained_mode {args.pretrained_mode}")
+
+ # init dataloader
+ data = VideoData(args, ds_group = ds_group, ds_type = ds_type)
+ dataloaders = data.train_dataloader()
+ dataloader_iters = [iter(loader) for loader in dataloaders]
+ ### init epoch in resuming
+ dataloader_init_epoch = (
+ init_step if init_step > 0 # in case of resuming
+ else args.dataloader_init_epoch if args.dataloader_init_epoch > 0 # in case of fintuning
+ else 0
+ )
+ data_epochs = [dataloader_init_epoch for _ in dataloaders]
+ for idx in range(len(dataloaders)):
+ print(f"Reset the {idx}th dataloader as epoch {data_epochs[idx]}")
+ if hasattr(dataloaders[idx], "sampler"):
+ dataloaders[idx].sampler.set_epoch(data_epochs[idx])
+ else:
+ raise NotImplementedError
+
+ ### torch.compile after loading all weights
+ print_model_summary([vae, image_disc, video_disc])
+ if args.zero > 0:
+ from torch.distributed.fsdp import (
+ FullyShardedDataParallel as FSDP,
+ ShardingStrategy,
+ MixedPrecision,
+ )
+ def my_policy(
+ module: torch.nn.Module,
+ recurse: bool,
+ **kwargs,
+ ) -> bool:
+ return True
+ auto_wrap_policy = my_policy
+ vae = FSDP(
+ vae,
+ device_id=device,
+ sharding_strategy=ShardingStrategy.FULL_SHARD,
+ mixed_precision=None,
+ auto_wrap_policy=auto_wrap_policy,
+ use_orig_params=True,
+ sync_module_states=True,
+ limit_all_gathers=True,
+ device_mesh=None,
+ ).to(device)
+ # vae = _FSDP(vae, device, args.zero)
+ else:
+ vae = DDP(vae.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb, find_unused_parameters=True)
+
+ image_disc = DDP(image_disc.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb)
+ video_disc = DDP(video_disc.to(device), device_ids=[args.gpu], bucket_cap_mb=args.bucket_cap_mb)
+
+ image_perceptual_model, video_perceptual_model = build_lpips_model(args)
+ image_perceptual_model = image_perceptual_model.to(device)
+ video_perceptual_model = video_perceptual_model.to(device)
+
+ if args.compile == "yes":
+ if args.vf_weight > 0 and args.vf_weight_approx < 0:
+ 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
+ torch._dynamo.config.cache_size_limit = 256
+ torch._dynamo.config.accumulated_cache_size_limit = 4096
+ torch._dynamo.config.automatic_dynamic_shapes = True
+ torch._dynamo.config.suppress_errors = False
+ torch._dynamo.config.optimize_ddp = False if args.use_checkpoint else True
+
+ for k in model_optims:
+ if k != "ema" and model_optims[k] and not is_torch_optim_sch(model_optims[k]):
+ print(f"compiling model {k}")
+ if k == "vae":
+ model_optims[k].encoder.compile()#options={'fx_graph_cache':True})
+ model_optims[k].decoder.compile()#options={'fx_graph_cache':True})
+ else:
+ model_optims[k].compile()#options={'fx_graph_cache':True})
+ print(f"Successfully compiled all models")
+
+ disc_loss = get_disc_loss(args.disc_loss_type)
+
+ if enable_timeline_sdk:
+ version_checker("2.0.0", "3.0.0")
+ bmq_cluster = os.getenv('CUDA_TIMER_STREAM_KAFKA_CLUSTER', 'bmq_bigbang_3rd')
+ bmq_topic = os.getenv('CUDA_TIMER_STREAM_KAFKA_TOPIC', 'megatron_cuda_timer_tracing_original')
+ ndtimeline.init_ndtimers(
+ mode="fsdp",
+ mesh_shape=(world_size,),
+ world_size=world_size,
+ enable_streamer=True,
+ post_handlers=[ndtimeline.handlers.MQNDHandler(mq_sinks=[ndtimeline.handlers.format_mq_sink(bmq_cluster, bmq_topic)])],
+ )
+ ndtimeline.set_global_step(init_step)
+ print(f"init timeline successfully")
+
+ # init profiler
+ def trace_handler(p):
+ p.export_chrome_trace(os.path.join(args.default_root_dir, f"trace_step_{p.step_num}_rank_{rank}.json.gz"))
+
+ if args.turn_on_profiler:
+ print(f"start to init profiler")
+ tp = torch.profiler.profile(
+ activities=[
+ torch.profiler.ProfilerActivity.CPU,
+ torch.profiler.ProfilerActivity.CUDA,
+ ],
+ schedule=torch.profiler.schedule(
+ wait=args.profiler_scheduler_wait_steps,
+ warmup=3,
+ active=2,
+ repeat=1,
+ ),
+ with_stack=True,
+ record_shapes=True,
+ profile_memory=True,
+ on_trace_ready=trace_handler
+ )
+ tp.start()
+ record_function = torch_record_function
+ print(f"finish to init profiler")
+ else:
+ record_function = nullcontext
+
+ start_time = time.time()
+ cnt = 0
+ debug_root_dir = data_prefix_manager(f"debug/{cnt}")
+ while os.path.exists(debug_root_dir):
+ cnt += 1
+ debug_root_dir = data_prefix_manager(f"debug/{cnt}")
+ os.makedirs(debug_root_dir, exist_ok=True)
+ if cp.get_cp_rank() == 0:
+ cp_group = cp.get_cp_group_rank()
+ debug_f = open(os.path.join(debug_root_dir, f"rank_{rank}_cp_group_{cp_group}.txt"), "w")
+ else:
+ debug_f = None
+ for global_step in range(init_step, args.max_steps):
+ # if args.turn_on_profiler and tp:
+ # tp.step()
+ loss_dicts = []
+
+ if global_step == args.discriminator_iter_start - args.disc_pretrain_iter:
+ logging.info(f"discriminator begins pretraining ")
+ if global_step == args.discriminator_iter_start:
+ log_str = "add GAN loss into training"
+ if args.disc_pretrain_iter > 0:
+ log_str += ", discriminator ends pretraining"
+ logging.info(log_str)
+
+ for idx in range(len(dataloader_iters)):
+ try:
+ _batch = next(dataloader_iters[idx])
+ except StopIteration:
+ data_epochs[idx] += 1
+ print(f"Reset the {idx}th dataloader as epoch {data_epochs[idx]}")
+ dataloaders[idx].sampler.set_epoch(data_epochs[idx])
+ dataloader_iters[idx] = iter(dataloaders[idx]) # update dataloader iter
+ _batch = next(dataloader_iters[idx])
+ except Exception as e:
+ raise e
+ x = _batch["video"]
+
+ _type = _batch["type"][0]
+
+ if _type == "image" and ds_type is None:
+ x = split_to_ranks(x)
+
+ disc_factor = 1.
+ with NanDetector(vae) if args.enable_nan_detector else nullcontext():
+ with record_function("vae"):
+ if _type == "image":
+ 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)
+ elif _type == "video":
+ if debug_f is not None:
+ debug_f.write(f'step {idx}, {_batch["path"]}\n')
+ x, x_recon, flat_frames, flat_frames_recon, vae_loss_dict, vae_log_dict = vae(
+ x, disc_factor,
+ image_disc=image_disc, video_disc=video_disc,
+ image_perceptual_model=image_perceptual_model,
+ video_perceptual_model=video_perceptual_model,
+ )
+ g_loss = sum(vae_loss_dict.values())
+ opt_vae.zero_grad()
+ g_loss.backward()
+ # print_gpu_usage("vae")
+ if args.max_grad_norm > 0:
+ torch.nn.utils.clip_grad_norm_(vae.parameters(), args.max_grad_norm)
+
+ # from pnp.utils import detect_anomalous_params
+ # detect_anomalous_params(g_loss, vae)
+
+ opt_vae.step()
+ opt_vae.zero_grad() # free memory
+ if args.ema == "yes":
+ update_ema(ema, vae.module)
+
+ with record_function("disc"):
+ disc_loss_dict = {}
+ # args.discriminator_iter_start=-1 args.disc_pretrain_iter=0
+ discloss = d_image_loss = d_video_loss = torch.tensor(0.).to(x.device)
+ ### enable pool warmup
+ for disc_step in range(args.disc_optim_steps):
+ require_optim = False
+ if _type == "image":
+ if args.image_disc_weight > 0:
+ require_optim = True
+ logits_image_real = image_disc(x, pool_name="real")
+ logits_image_fake = image_disc(x_recon.detach(), pool_name="fake")
+ d_image_loss = disc_loss(logits_image_real, logits_image_fake)
+ discloss = d_image_loss * args.image_disc_weight
+ disc_loss_dict["train/logits_image_real"] = logits_image_real.mean().detach()
+ disc_loss_dict["train/logits_image_fake"] = logits_image_fake.mean().detach()
+ disc_loss_dict["train/d_image_loss"] = discloss.detach()
+ opt_discs, sch_discs = [opt_image_disc], [sch_image_disc]
+ elif _type == "video":
+ if args.image_disc_weight > 0 and args.gan_image4video == "yes":
+ require_optim = True
+ logits_image_real = image_disc(flat_frames.detach(), pool_name="real")
+ logits_image_fake = image_disc(flat_frames_recon.detach(), pool_name="fake")
+ d_image_loss = disc_loss(logits_image_real, logits_image_fake)
+ disc_loss_dict["train/logits_image_real"] = logits_image_real.mean().detach()
+ disc_loss_dict["train/logits_image_fake"] = logits_image_fake.mean().detach()
+ disc_loss_dict["train/d_image_loss"] = (d_image_loss * args.image_disc_weight).detach()
+ if args.video_disc_weight > 0:
+ require_optim = True
+ logits_video_real = video_disc(x.detach(), pool_name="real")
+ logits_video_fake = video_disc(x_recon.detach(), pool_name="fake")
+ d_video_loss = disc_loss(logits_video_real, logits_video_fake)
+ disc_loss_dict["train/logits_video_real"] = logits_video_real.mean().detach()
+ disc_loss_dict["train/logits_video_fake"] = logits_video_fake.mean().detach()
+ disc_loss_dict["train/d_video_loss"] = (d_video_loss * args.video_disc_weight).detach()
+ discloss = d_image_loss * args.image_disc_weight + d_video_loss * args.video_disc_weight
+ opt_discs, sch_discs = [opt_image_disc, opt_video_disc], [sch_image_disc, sch_video_disc]
+ discloss = disc_factor * discloss
+
+ if require_optim:
+ for opt_disc in opt_discs:
+ opt_disc.zero_grad()
+ discloss.backward()
+ # print_gpu_usage("disc")
+ if args.max_grad_norm_disc > 0:
+ torch.nn.utils.clip_grad_norm_(image_disc.parameters(), args.max_grad_norm_disc)
+ torch.nn.utils.clip_grad_norm_(video_disc.parameters(), args.max_grad_norm_disc)
+
+ for opt_disc in opt_discs:
+ opt_disc.step()
+ for opt_disc in opt_discs:
+ opt_disc.zero_grad() # free memory
+
+ with record_function("loss"):
+ loss_dict = {**vae_loss_dict, **disc_loss_dict, **vae_log_dict}
+ if (global_step+1) % args.log_every == 0:
+ reduced_loss_dict = reduce_losses(loss_dict)
+ else:
+ reduced_loss_dict = {}
+ loss_dicts.append(reduced_loss_dict)
+
+ # update scheduler
+ if not sch_vae is None:
+ sch_vae.step()
+ for sch_disc in sch_discs:
+ if not sch_disc is None:
+ sch_disc.step()
+
+ # if enable_timeline_sdk:
+ # ndtimeline.inc_step()
+
+ if (global_step+1) % args.log_every == 0:
+ avg_loss_dict = average_losses(loss_dicts)
+ torch.cuda.synchronize()
+ end_time = time.time()
+ iter_speed = (end_time - start_time) / args.log_every
+
+ if args.mfu_logging == "yes":
+ tflops = get_tflops() / args.log_every
+ tflops_log_str = f"tflops={tflops:.1f}, "
+ tflops_dict = get_tflops_dict(args.log_every)
+ tflops_dict_log_str = f"tflops_Dict={tflops_dict}, "
+ mfu = get_mfu(iter_speed) / args.log_every
+ mfu_log_str = f"mfu={mfu:.3f}, "
+ else:
+ tflops_log_str = ""
+ tflops_dict_log_str = ""
+ mfu_log_str = ""
+
+ if rank == 0:
+ avg_loss_dict["lr"] = opt_vae.param_groups[0]['lr']
+ for key, value in avg_loss_dict.items():
+ wandb.log({key: value}, step=global_step)
+
+ recons_loss_sum, video_perceptual_loss_sum = 0., 0.
+ for key in avg_loss_dict:
+ if 'recon_loss' in key:
+ recons_loss_sum += avg_loss_dict[key]
+ if 'video_perceptual_loss' in key:
+ video_perceptual_loss_sum += avg_loss_dict[key]
+
+ print(f'global_step={global_step}, recon_loss={recons_loss_sum:.4f}, ' \
+ f'video_perceptual_loss={video_perceptual_loss_sum:.4f}, ' \
+ f'iter_speed={iter_speed:.2f}s, ' \
+ f'{mfu_log_str}' \
+ f'{tflops_log_str}' \
+ f'{tflops_dict_log_str}' \
+ )
+ start_time = time.time()
+ if enable_timeline_sdk:
+ ndtimeline.flush()
+
+ if (global_step+1) % args.ckpt_every == 0 and global_step != init_step:
+ checkpoint_path = os.path.join(checkpoint_dir, f'model_step_{global_step}.ckpt')
+ if args.zero > 0:
+ save_model(vae, rank, checkpoint_path, global_step)
+ else:
+ if rank == 0:
+ save_dict = {}
+ for k in model_optims:
+ model = model_optims[k]
+ save_dict[k] = None if model is None \
+ else model.module.state_dict() if hasattr(model, "module") \
+ else model.state_dict()
+ torch.save({
+ 'step': global_step,
+ **save_dict,
+ }, checkpoint_path)
+ print(f'Checkpoint saved at step {global_step}')
+
+ if (global_step+1) % args.manual_gc_interval == 0:
+ gc.collect()
+
+if __name__ == '__main__':
+ main()
diff --git a/grn/tokenizer/videovae/__init__.py b/grn/tokenizer/videovae/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/grn/tokenizer/videovae/evaluation/__init__.py b/grn/tokenizer/videovae/evaluation/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..c10859408ed3340de87004aea0c524a2266b8785
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/__init__.py
@@ -0,0 +1,8 @@
+from .common_metrics_on_video_quality.calculate_fvd import calculate_fvd
+from .common_metrics_on_video_quality.calculate_lpips import calculate_lpips
+from .common_metrics_on_video_quality.calculate_psnr import calculate_psnr
+from .common_metrics_on_video_quality.calculate_ssim import calculate_ssim
+
+from .fvd import get_fvd_logits, frechet_distance, load_fvd_model
+from .fid import calculate_frechet_distance
+from .inception import InceptionV3
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore
new file mode 100644
index 0000000000000000000000000000000000000000..ed8ebf583f771da9150c35db3955987b7d757904
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/.gitignore
@@ -0,0 +1 @@
+__pycache__
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py
new file mode 100644
index 0000000000000000000000000000000000000000..04d2a200f2d1ae529ccacc8bac5f676d12db3fa4
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_fvd.py
@@ -0,0 +1,85 @@
+import numpy as np
+import torch
+from tqdm import tqdm
+
+def trans(x):
+ # if greyscale images add channel
+ if x.shape[-3] == 1:
+ x = x.repeat(1, 1, 3, 1, 1)
+
+ # permute BTCHW -> BCTHW
+ x = x.permute(0, 2, 1, 3, 4)
+
+ return x
+
+def calculate_fvd(videos1, videos2, device, method='styleganv'):
+
+ if method == 'styleganv':
+ from .fvd.styleganv.fvd import get_fvd_feats, frechet_distance, load_i3d_pretrained
+ elif method == 'videogpt':
+ from .fvd.videogpt.fvd import load_i3d_pretrained
+ from .fvd.videogpt.fvd import get_fvd_logits as get_fvd_feats
+ from .fvd.videogpt.fvd import frechet_distance
+
+ print("calculate_fvd...")
+
+ # videos [batch_size, timestamps, channel, h, w]
+
+ assert videos1.shape == videos2.shape
+
+ i3d = load_i3d_pretrained(device=device)
+ fvd_results = []
+
+ # support grayscale input, if grayscale -> channel*3
+ # BTCHW -> BCTHW
+ # videos -> [batch_size, channel, timestamps, h, w]
+
+ videos1 = trans(videos1)
+ videos2 = trans(videos2)
+
+ fvd_results = {}
+
+ # for calculate FVD, each clip_timestamp must >= 10
+ for clip_timestamp in tqdm(range(10, videos1.shape[-3]+1)):
+
+ # get a video clip
+ # videos_clip [batch_size, channel, timestamps[:clip], h, w]
+ videos_clip1 = videos1[:, :, : clip_timestamp]
+ videos_clip2 = videos2[:, :, : clip_timestamp]
+
+ # get FVD features
+ feats1 = get_fvd_feats(videos_clip1, i3d=i3d, device=device)
+ feats2 = get_fvd_feats(videos_clip2, i3d=i3d, device=device)
+
+ # calculate FVD when timestamps[:clip]
+ fvd_results[clip_timestamp] = frechet_distance(feats1, feats2)
+
+ result = {
+ "value": fvd_results,
+ "video_setting": videos1.shape,
+ "video_setting_name": "batch_size, channel, time, heigth, width",
+ }
+
+ return result
+
+# test code / using example
+
+def main():
+ NUMBER_OF_VIDEOS = 8
+ VIDEO_LENGTH = 50
+ CHANNEL = 3
+ SIZE = 64
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ videos2 = torch.ones(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ device = torch.device("cuda")
+ # device = torch.device("cpu")
+
+ import json
+ result = calculate_fvd(videos1, videos2, device, method='videogpt')
+ print(json.dumps(result, indent=4))
+
+ result = calculate_fvd(videos1, videos2, device, method='styleganv')
+ print(json.dumps(result, indent=4))
+
+if __name__ == "__main__":
+ main()
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py
new file mode 100644
index 0000000000000000000000000000000000000000..ae37b08b9176ca1b4a92442f520666f06bd82e9e
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_lpips.py
@@ -0,0 +1,72 @@
+import numpy as np
+import torch
+from tqdm import tqdm
+import math
+
+import torch
+import lpips
+from .utils import build_dataloader
+
+spatial = True # Return a spatial map of perceptual distance.
+
+# Linearly calibrated models (LPIPS)
+loss_fn = lpips.LPIPS(net='vgg', spatial=spatial) # Can also set net = 'squeeze' or 'vgg'
+# loss_fn = lpips.LPIPS(net='alex', spatial=spatial, lpips=False) # Can also set net = 'squeeze' or 'vgg'
+
+def calculate_lpips(videos1, videos2, device):
+ # image should be RGB, IMPORTANT: normalized to [-1,1]
+ print("calculate_lpips...")
+
+ lpips_results = {}
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
+ # get a video [timestamps, channel, h, w]
+
+ assert video1.shape == video2.shape
+ video1 = video1.squeeze(0) * 2 - 1
+ video2 = video2.squeeze(0) * 2 - 1
+
+ for clip_timestamp in range(len(video1)):
+ # get a img
+ # img [timestamps[x], channel, h, w]
+ # img [channel, h, w] tensor
+
+ img1 = video1[clip_timestamp].unsqueeze(0).to(device)
+ img2 = video2[clip_timestamp].unsqueeze(0).to(device)
+
+ loss_fn.to(device)
+
+ # calculate lpips of a video
+ value = np.array(loss_fn.forward(img1, img2).mean().detach().cpu().tolist())
+ if clip_timestamp not in lpips_results:
+ lpips_results[clip_timestamp] = value
+ else:
+ lpips_results[clip_timestamp] += value
+
+ for clip_timestamp in range(len(video1)):
+ lpips_results[clip_timestamp] /= len(dataloader1)
+
+ result = {
+ "value": lpips_results,
+ }
+
+ return result
+
+# test code / using example
+
+def main():
+ NUMBER_OF_VIDEOS = 8
+ VIDEO_LENGTH = 50
+ CHANNEL = 3
+ SIZE = 64
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ videos2 = torch.ones(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ device = torch.device("cuda")
+ # device = torch.device("cpu")
+
+ import json
+ result = calculate_lpips(videos1, videos2, device)
+ print(json.dumps(result, indent=4))
+
+if __name__ == "__main__":
+ main()
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py
new file mode 100644
index 0000000000000000000000000000000000000000..1e9fdf1ce09bd32cafc37ad127c15f50764f3358
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_psnr.py
@@ -0,0 +1,83 @@
+import numpy as np
+import torch
+from tqdm import tqdm
+import math
+from multiprocessing import Pool
+from .utils import build_dataloader
+
+def img_psnr(img1, img2):
+ # [0,1]
+ # compute mse
+ # mse = np.mean((img1-img2)**2)
+ mse = np.mean((img1 / 1.0 - img2 / 1.0) ** 2)
+ # compute psnr
+ if mse < 1e-10:
+ return 100
+ psnr = 20 * math.log10(1 / math.sqrt(mse))
+ return psnr
+
+def trans(x):
+ return x
+
+def process_video(video_pair):
+ video1, video2 = video_pair
+ psnr_results_of_a_video = []
+ for clip_timestamp in range(len(video1)):
+ # get a img
+ # img [timestamps[x], channel, h, w]
+ # img [channel, h, w] numpy
+
+ img1 = video1[clip_timestamp].numpy()
+ img2 = video2[clip_timestamp].numpy()
+
+ # calculate psnr of a video
+ psnr_results_of_a_video.append(img_psnr(img1, img2))
+
+ return psnr_results_of_a_video
+
+
+def calculate_psnr(videos1, videos2):
+ print("calculate_psnr...")
+
+ # videos [batch_size, timestamps, channel, h, w]
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
+ psnr_results = []
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
+ video1 = video1.squeeze(0)
+ video2 = video2.squeeze(0)
+ video_pair = (video1, video2)
+ result = process_video(video_pair)
+ psnr_results.append(result)
+
+ psnr_results = np.array(psnr_results)
+
+ psnr = {}
+ psnr_std = {}
+
+ for clip_timestamp in range(len(video1)):
+ psnr[clip_timestamp] = np.mean(psnr_results[:,clip_timestamp])
+ psnr_std[clip_timestamp] = np.std(psnr_results[:,clip_timestamp])
+
+ result = {
+ "value": psnr,
+ "value_std": psnr_std,
+ }
+
+ return result
+
+# test code / using example
+
+def main():
+ NUMBER_OF_VIDEOS = 8
+ VIDEO_LENGTH = 50
+ CHANNEL = 3
+ SIZE = 64
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ videos2 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+
+ import json
+ result = calculate_psnr(videos1, videos2)
+ print(json.dumps(result, indent=4))
+
+if __name__ == "__main__":
+ main()
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py
new file mode 100644
index 0000000000000000000000000000000000000000..6aaad70c1da1af458f0cb5305e92c4edcb808e05
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/calculate_ssim.py
@@ -0,0 +1,138 @@
+import numpy as np
+import torch
+from tqdm import tqdm
+import cv2
+from multiprocessing import Pool
+import torch
+import torch.nn.functional as F
+from .utils import build_dataloader
+
+def ssim(img1, img2):
+ _img1, _img2 = img1, img2
+
+ ### original implementation
+ # C1 = 0.01 ** 2
+ # C2 = 0.03 ** 2
+ # img1 = img1.astype(np.float64)
+ # img2 = img2.astype(np.float64)
+ # kernel = cv2.getGaussianKernel(11, 1.5)
+ # window = np.outer(kernel, kernel.transpose())
+ # mu1 = cv2.filter2D(img1, -1, window)[5:-5, 5:-5] # valid
+ # mu2 = cv2.filter2D(img2, -1, window)[5:-5, 5:-5]
+ # mu1_sq = mu1 ** 2
+ # mu2_sq = mu2 ** 2
+ # mu1_mu2 = mu1 * mu2
+ # sigma1_sq = cv2.filter2D(img1 ** 2, -1, window)[5:-5, 5:-5] - mu1_sq
+ # sigma2_sq = cv2.filter2D(img2 ** 2, -1, window)[5:-5, 5:-5] - mu2_sq
+ # sigma12 = cv2.filter2D(img1 * img2, -1, window)[5:-5, 5:-5] - mu1_mu2
+ # ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) *
+ # (sigma1_sq + sigma2_sq + C2))
+ # return ssim_map.mean()
+ # res1 = ssim_map.mean()
+
+ ### accelerated implementation
+ img1, img2 = torch.from_numpy(_img1), torch.from_numpy(_img2)
+ C1 = 0.01 ** 2
+ C2 = 0.03 ** 2
+
+ # Ensure data is float and move to GPU
+ img1 = img1.to(torch.float64).cuda()
+ img2 = img2.to(torch.float64).cuda()
+
+ # Gaussian kernel
+ kernel = torch.tensor(cv2.getGaussianKernel(11, 1.5)).to(torch.float64).cuda()
+ window = kernel @ kernel.t()
+ window = window.unsqueeze(0).unsqueeze(0).cuda()
+
+ mu1 = F.conv2d(img1.unsqueeze(0), window, padding=0, groups=1)
+ mu2 = F.conv2d(img2.unsqueeze(0), window, padding=0, groups=1)
+ mu1_sq = mu1.pow(2)
+ mu2_sq = mu2.pow(2)
+ mu1_mu2 = mu1 * mu2
+ sigma1_sq = F.conv2d(img1.unsqueeze(0) ** 2, window, padding=0, groups=1) - mu1_sq
+ sigma2_sq = F.conv2d(img2.unsqueeze(0) ** 2, window, padding=0, groups=1) - mu2_sq
+ sigma12 = F.conv2d(img1.unsqueeze(0) * img2.unsqueeze(0), window, padding=0, groups=1) - mu1_mu2
+
+ ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) *
+ (sigma1_sq + sigma2_sq + C2))
+ res2 = ssim_map.mean().item()
+ # print(res1-res2)
+ return res2
+ # return ssim_map.mean().item()
+
+
+def calculate_ssim_function(img1, img2):
+ # [0,1]
+ # ssim is the only metric extremely sensitive to gray being compared to b/w
+ if not img1.shape == img2.shape:
+ raise ValueError('Input images must have the same dimensions.')
+ if img1.ndim == 2:
+ return ssim(img1, img2)
+ elif img1.ndim == 3:
+ if img1.shape[0] == 3:
+ ssims = []
+ for i in range(3):
+ ssims.append(ssim(img1[i], img2[i]))
+ return np.array(ssims).mean()
+ elif img1.shape[0] == 1:
+ return ssim(np.squeeze(img1), np.squeeze(img2))
+ else:
+ raise ValueError('Wrong input image dimensions.')
+
+def trans(x):
+ return x
+
+def process_video(video_pair):
+ video1, video2 = video_pair
+ ssim_results_of_a_video = []
+ for clip_timestamp in range(len(video1)):
+ img1 = video1[clip_timestamp].numpy()
+ img2 = video2[clip_timestamp].numpy()
+ ssim_results_of_a_video.append(calculate_ssim_function(img1, img2))
+ return ssim_results_of_a_video
+
+def calculate_ssim(videos1, videos2):
+ print("calculate_ssim...")
+
+ ssim_results = []
+ dataloader1, dataloader2 = build_dataloader(videos1, videos2)
+ for video1, video2 in tqdm(zip(dataloader1, dataloader2), total=len(dataloader1)):
+ video1 = video1.squeeze(0)
+ video2 = video2.squeeze(0)
+ video_pair = (video1, video2)
+ result = process_video(video_pair)
+ ssim_results.append(result)
+
+ ssim_results = np.array(ssim_results)
+
+ ssim = {}
+ ssim_std = {}
+
+ for clip_timestamp in range(len(video1)):
+ ssim[clip_timestamp] = np.mean(ssim_results[:,clip_timestamp])
+ ssim_std[clip_timestamp] = np.std(ssim_results[:,clip_timestamp])
+
+ result = {
+ "value": ssim,
+ "value_std": ssim_std,
+ }
+
+ return result
+
+# test code / using example
+
+def main():
+ NUMBER_OF_VIDEOS = 8
+ VIDEO_LENGTH = 50
+ CHANNEL = 3
+ SIZE = 64
+ videos1 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ videos2 = torch.zeros(NUMBER_OF_VIDEOS, VIDEO_LENGTH, CHANNEL, SIZE, SIZE, requires_grad=False)
+ device = torch.device("cuda")
+
+ import json
+ result = calculate_ssim(videos1, videos2)
+ print(json.dumps(result, indent=4))
+
+if __name__ == "__main__":
+ main()
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py
new file mode 100644
index 0000000000000000000000000000000000000000..3043a2a4a4c4fc48ca97aedba074c2c2379685e4
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/styleganv/fvd.py
@@ -0,0 +1,90 @@
+import torch
+import os
+import math
+import torch.nn.functional as F
+
+# https://github.com/universome/fvd-comparison
+
+
+def load_i3d_pretrained(device=torch.device('cpu')):
+ i3D_WEIGHTS_URL = "https://www.dropbox.com/s/ge9e5ujwgetktms/i3d_torchscript.pt"
+ filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'i3d_torchscript.pt')
+ print(filepath)
+ if not os.path.exists(filepath):
+ print(f"preparing for download {i3D_WEIGHTS_URL}, you can download it by yourself.")
+ os.system(f"wget {i3D_WEIGHTS_URL} -O {filepath}")
+ i3d = torch.jit.load(filepath).eval().to(device)
+ i3d = torch.nn.DataParallel(i3d)
+ return i3d
+
+
+def get_feats(videos, detector, device, bs=10):
+ # videos : torch.tensor BCTHW [0, 1]
+ detector_kwargs = dict(rescale=False, resize=False, return_features=True) # Return raw features before the softmax layer.
+ feats = np.empty((0, 400))
+ with torch.no_grad():
+ for i in range((len(videos)-1)//bs + 1):
+ 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()])
+ return feats
+
+
+def get_fvd_feats(videos, i3d, device, bs=10):
+ # videos in [0, 1] as torch tensor BCTHW
+ # videos = [preprocess_single(video) for video in videos]
+ embeddings = get_feats(videos, i3d, device, bs)
+ return embeddings
+
+
+def preprocess_single(video, resolution=224, sequence_length=None):
+ # video: CTHW, [0, 1]
+ c, t, h, w = video.shape
+
+ # temporal crop
+ if sequence_length is not None:
+ assert sequence_length <= t
+ video = video[:, :sequence_length]
+
+ # scale shorter side to resolution
+ scale = resolution / min(h, w)
+ if h < w:
+ target_size = (resolution, math.ceil(w * scale))
+ else:
+ target_size = (math.ceil(h * scale), resolution)
+ video = F.interpolate(video, size=target_size, mode='bilinear', align_corners=False)
+
+ # center crop
+ c, t, h, w = video.shape
+ w_start = (w - resolution) // 2
+ h_start = (h - resolution) // 2
+ video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
+
+ # [0, 1] -> [-1, 1]
+ video = (video - 0.5) * 2
+
+ return video.contiguous()
+
+
+"""
+Copy-pasted from https://github.com/cvpr2022-stylegan-v/stylegan-v/blob/main/src/metrics/frechet_video_distance.py
+"""
+from typing import Tuple
+from scipy.linalg import sqrtm
+import numpy as np
+
+
+def compute_stats(feats: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
+ mu = feats.mean(axis=0) # [d]
+ sigma = np.cov(feats, rowvar=False) # [d, d]
+ return mu, sigma
+
+
+def frechet_distance(feats_fake: np.ndarray, feats_real: np.ndarray) -> float:
+ mu_gen, sigma_gen = compute_stats(feats_fake)
+ mu_real, sigma_real = compute_stats(feats_real)
+ m = np.square(mu_gen - mu_real).sum()
+ if feats_fake.shape[0]>1:
+ s, _ = sqrtm(np.dot(sigma_gen, sigma_real), disp=False) # pylint: disable=no-member
+ fid = np.real(m + np.trace(sigma_gen + sigma_real - s * 2))
+ else:
+ fid = np.real(m)
+ return float(fid)
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py
new file mode 100644
index 0000000000000000000000000000000000000000..f90d4dbc2baeaf126ff4a385ccedf419f69b7d85
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/fvd.py
@@ -0,0 +1,137 @@
+import torch
+import os
+import math
+import torch.nn.functional as F
+import numpy as np
+import einops
+
+def load_i3d_pretrained(device=torch.device('cpu')):
+ i3D_WEIGHTS_URL = "https://onedrive.live.com/download?cid=78EEF3EB6AE7DBCB&resid=78EEF3EB6AE7DBCB%21199&authkey=AApKdFHPXzWLNyI"
+ filepath = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'i3d_pretrained_400.pt')
+ print(filepath)
+ if not os.path.exists(filepath):
+ print(f"preparing for download {i3D_WEIGHTS_URL}, you can download it by yourself.")
+ os.system(f"wget {i3D_WEIGHTS_URL} -O {filepath}")
+ from .pytorch_i3d import InceptionI3d
+ i3d = InceptionI3d(400, in_channels=3).eval().to(device)
+ i3d.load_state_dict(torch.load(filepath, map_location=device, weights_only=True))
+ i3d = torch.nn.DataParallel(i3d)
+ return i3d
+
+def preprocess_single(video, resolution, sequence_length=None):
+ # video: THWC, {0, ..., 255}
+ video = video.permute(0, 3, 1, 2).float() / 255. # TCHW
+ t, c, h, w = video.shape
+
+ # temporal crop
+ if sequence_length is not None:
+ assert sequence_length <= t
+ video = video[:sequence_length]
+
+ # scale shorter side to resolution
+ scale = resolution / min(h, w)
+ if h < w:
+ target_size = (resolution, math.ceil(w * scale))
+ else:
+ target_size = (math.ceil(h * scale), resolution)
+ video = F.interpolate(video, size=target_size, mode='bilinear',
+ align_corners=False)
+
+ # center crop
+ t, c, h, w = video.shape
+ w_start = (w - resolution) // 2
+ h_start = (h - resolution) // 2
+ video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
+ video = video.permute(1, 0, 2, 3).contiguous() # CTHW
+
+ video -= 0.5
+
+ return video
+
+def preprocess(videos, target_resolution=224):
+ # we should tras videos in [0-1] [b c t h w] as th.float
+ # -> videos in {0, ..., 255} [b t h w c] as np.uint8 array
+ videos = einops.rearrange(videos, 'b c t h w -> b t h w c')
+ videos = (videos*255).numpy().astype(np.uint8)
+
+ b, t, h, w, c = videos.shape
+ videos = torch.from_numpy(videos)
+ videos = torch.stack([preprocess_single(video, target_resolution) for video in videos])
+ return videos * 2 # [-0.5, 0.5] -> [-1, 1]
+
+def get_fvd_logits(videos, i3d, device, bs=10):
+ videos = preprocess(videos)
+ embeddings = get_logits(i3d, videos, device, bs=10)
+ return embeddings
+
+# https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L161
+def _symmetric_matrix_square_root(mat, eps=1e-10):
+ u, s, v = torch.svd(mat)
+ si = torch.where(s < eps, s, torch.sqrt(s))
+ return torch.matmul(torch.matmul(u, torch.diag(si)), v.t())
+
+# https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L400
+def trace_sqrt_product(sigma, sigma_v):
+ sqrt_sigma = _symmetric_matrix_square_root(sigma)
+ sqrt_a_sigmav_a = torch.matmul(sqrt_sigma, torch.matmul(sigma_v, sqrt_sigma))
+ return torch.trace(_symmetric_matrix_square_root(sqrt_a_sigmav_a))
+
+# https://discuss.pytorch.org/t/covariance-and-gradient-support/16217/2
+def cov(m, rowvar=False):
+ '''Estimate a covariance matrix given data.
+
+ Covariance indicates the level to which two variables vary together.
+ If we examine N-dimensional samples, `X = [x_1, x_2, ... x_N]^T`,
+ then the covariance matrix element `C_{ij}` is the covariance of
+ `x_i` and `x_j`. The element `C_{ii}` is the variance of `x_i`.
+
+ Args:
+ m: A 1-D or 2-D array containing multiple variables and observations.
+ Each row of `m` represents a variable, and each column a single
+ observation of all those variables.
+ rowvar: If `rowvar` is True, then each row represents a
+ variable, with observations in the columns. Otherwise, the
+ relationship is transposed: each column represents a variable,
+ while the rows contain observations.
+
+ Returns:
+ The covariance matrix of the variables.
+ '''
+ if m.dim() > 2:
+ raise ValueError('m has more than 2 dimensions')
+ if m.dim() < 2:
+ m = m.view(1, -1)
+ if not rowvar and m.size(0) != 1:
+ m = m.t()
+
+ fact = 1.0 / (m.size(1) - 1) # unbiased estimate
+ m -= torch.mean(m, dim=1, keepdim=True)
+ mt = m.t() # if complex: mt = m.t().conj()
+ return fact * m.matmul(mt).squeeze()
+
+
+def frechet_distance(x1, x2):
+ x1 = x1.flatten(start_dim=1)
+ x2 = x2.flatten(start_dim=1)
+ m, m_w = x1.mean(dim=0), x2.mean(dim=0)
+ sigma, sigma_w = cov(x1, rowvar=False), cov(x2, rowvar=False)
+ mean = torch.sum((m - m_w) ** 2)
+ if x1.shape[0]>1:
+ sqrt_trace_component = trace_sqrt_product(sigma, sigma_w)
+ trace = torch.trace(sigma + sigma_w) - 2.0 * sqrt_trace_component
+ fd = trace + mean
+ else:
+ fd = np.real(mean)
+ return float(fd)
+
+
+def get_logits(i3d, videos, device, bs=10):
+ # assert videos.shape[0] % 16 == 0
+ with torch.no_grad():
+ logits = []
+ for i in range(0, videos.shape[0], bs):
+ batch = videos[i:i + bs].to(device)
+ # logits.append(i3d.module.extract_features(batch)) # wrong
+ logits.append(i3d(batch)) # right
+ logits = torch.cat(logits, dim=0)
+ return logits
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py
new file mode 100644
index 0000000000000000000000000000000000000000..58a16cdd797e5a0e9d711bb0ba14281486d44d03
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/fvd/videogpt/pytorch_i3d.py
@@ -0,0 +1,322 @@
+# Original code from https://github.com/piergiaj/pytorch-i3d
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+
+class MaxPool3dSamePadding(nn.MaxPool3d):
+
+ def compute_pad(self, dim, s):
+ if s % self.stride[dim] == 0:
+ return max(self.kernel_size[dim] - self.stride[dim], 0)
+ else:
+ return max(self.kernel_size[dim] - (s % self.stride[dim]), 0)
+
+ def forward(self, x):
+ # compute 'same' padding
+ (batch, channel, t, h, w) = x.size()
+ out_t = np.ceil(float(t) / float(self.stride[0]))
+ out_h = np.ceil(float(h) / float(self.stride[1]))
+ out_w = np.ceil(float(w) / float(self.stride[2]))
+ pad_t = self.compute_pad(0, t)
+ pad_h = self.compute_pad(1, h)
+ pad_w = self.compute_pad(2, w)
+
+ pad_t_f = pad_t // 2
+ pad_t_b = pad_t - pad_t_f
+ pad_h_f = pad_h // 2
+ pad_h_b = pad_h - pad_h_f
+ pad_w_f = pad_w // 2
+ pad_w_b = pad_w - pad_w_f
+
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
+ x = F.pad(x, pad)
+ return super(MaxPool3dSamePadding, self).forward(x)
+
+
+class Unit3D(nn.Module):
+
+ def __init__(self, in_channels,
+ output_channels,
+ kernel_shape=(1, 1, 1),
+ stride=(1, 1, 1),
+ padding=0,
+ activation_fn=F.relu,
+ use_batch_norm=True,
+ use_bias=False,
+ name='unit_3d'):
+
+ """Initializes Unit3D module."""
+ super(Unit3D, self).__init__()
+
+ self._output_channels = output_channels
+ self._kernel_shape = kernel_shape
+ self._stride = stride
+ self._use_batch_norm = use_batch_norm
+ self._activation_fn = activation_fn
+ self._use_bias = use_bias
+ self.name = name
+ self.padding = padding
+
+ self.conv3d = nn.Conv3d(in_channels=in_channels,
+ out_channels=self._output_channels,
+ kernel_size=self._kernel_shape,
+ stride=self._stride,
+ padding=0, # we always want padding to be 0 here. We will dynamically pad based on input size in forward function
+ bias=self._use_bias)
+
+ if self._use_batch_norm:
+ self.bn = nn.BatchNorm3d(self._output_channels, eps=1e-5, momentum=0.001)
+
+ def compute_pad(self, dim, s):
+ if s % self._stride[dim] == 0:
+ return max(self._kernel_shape[dim] - self._stride[dim], 0)
+ else:
+ return max(self._kernel_shape[dim] - (s % self._stride[dim]), 0)
+
+
+ def forward(self, x):
+ # compute 'same' padding
+ (batch, channel, t, h, w) = x.size()
+ out_t = np.ceil(float(t) / float(self._stride[0]))
+ out_h = np.ceil(float(h) / float(self._stride[1]))
+ out_w = np.ceil(float(w) / float(self._stride[2]))
+ pad_t = self.compute_pad(0, t)
+ pad_h = self.compute_pad(1, h)
+ pad_w = self.compute_pad(2, w)
+
+ pad_t_f = pad_t // 2
+ pad_t_b = pad_t - pad_t_f
+ pad_h_f = pad_h // 2
+ pad_h_b = pad_h - pad_h_f
+ pad_w_f = pad_w // 2
+ pad_w_b = pad_w - pad_w_f
+
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
+ x = F.pad(x, pad)
+
+ x = self.conv3d(x)
+ if self._use_batch_norm:
+ x = self.bn(x)
+ if self._activation_fn is not None:
+ x = self._activation_fn(x)
+ return x
+
+
+
+class InceptionModule(nn.Module):
+ def __init__(self, in_channels, out_channels, name):
+ super(InceptionModule, self).__init__()
+
+ self.b0 = Unit3D(in_channels=in_channels, output_channels=out_channels[0], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_0/Conv3d_0a_1x1')
+ self.b1a = Unit3D(in_channels=in_channels, output_channels=out_channels[1], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_1/Conv3d_0a_1x1')
+ self.b1b = Unit3D(in_channels=out_channels[1], output_channels=out_channels[2], kernel_shape=[3, 3, 3],
+ name=name+'/Branch_1/Conv3d_0b_3x3')
+ self.b2a = Unit3D(in_channels=in_channels, output_channels=out_channels[3], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_2/Conv3d_0a_1x1')
+ self.b2b = Unit3D(in_channels=out_channels[3], output_channels=out_channels[4], kernel_shape=[3, 3, 3],
+ name=name+'/Branch_2/Conv3d_0b_3x3')
+ self.b3a = MaxPool3dSamePadding(kernel_size=[3, 3, 3],
+ stride=(1, 1, 1), padding=0)
+ self.b3b = Unit3D(in_channels=in_channels, output_channels=out_channels[5], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_3/Conv3d_0b_1x1')
+ self.name = name
+
+ def forward(self, x):
+ b0 = self.b0(x)
+ b1 = self.b1b(self.b1a(x))
+ b2 = self.b2b(self.b2a(x))
+ b3 = self.b3b(self.b3a(x))
+ return torch.cat([b0,b1,b2,b3], dim=1)
+
+
+class InceptionI3d(nn.Module):
+ """Inception-v1 I3D architecture.
+ The model is introduced in:
+ Quo Vadis, Action Recognition? A New Model and the Kinetics Dataset
+ Joao Carreira, Andrew Zisserman
+ https://arxiv.org/pdf/1705.07750v1.pdf.
+ See also the Inception architecture, introduced in:
+ Going deeper with convolutions
+ Christian Szegedy, Wei Liu, Yangqing Jia, Pierre Sermanet, Scott Reed,
+ Dragomir Anguelov, Dumitru Erhan, Vincent Vanhoucke, Andrew Rabinovich.
+ http://arxiv.org/pdf/1409.4842v1.pdf.
+ """
+
+ # Endpoints of the model in order. During construction, all the endpoints up
+ # to a designated `final_endpoint` are returned in a dictionary as the
+ # second return value.
+ VALID_ENDPOINTS = (
+ 'Conv3d_1a_7x7',
+ 'MaxPool3d_2a_3x3',
+ 'Conv3d_2b_1x1',
+ 'Conv3d_2c_3x3',
+ 'MaxPool3d_3a_3x3',
+ 'Mixed_3b',
+ 'Mixed_3c',
+ 'MaxPool3d_4a_3x3',
+ 'Mixed_4b',
+ 'Mixed_4c',
+ 'Mixed_4d',
+ 'Mixed_4e',
+ 'Mixed_4f',
+ 'MaxPool3d_5a_2x2',
+ 'Mixed_5b',
+ 'Mixed_5c',
+ 'Logits',
+ 'Predictions',
+ )
+
+ def __init__(self, num_classes=400, spatial_squeeze=True,
+ final_endpoint='Logits', name='inception_i3d', in_channels=3, dropout_keep_prob=0.5):
+ """Initializes I3D model instance.
+ Args:
+ num_classes: The number of outputs in the logit layer (default 400, which
+ matches the Kinetics dataset).
+ spatial_squeeze: Whether to squeeze the spatial dimensions for the logits
+ before returning (default True).
+ final_endpoint: The model contains many possible endpoints.
+ `final_endpoint` specifies the last endpoint for the model to be built
+ up to. In addition to the output at `final_endpoint`, all the outputs
+ at endpoints up to `final_endpoint` will also be returned, in a
+ dictionary. `final_endpoint` must be one of
+ InceptionI3d.VALID_ENDPOINTS (default 'Logits').
+ name: A string (optional). The name of this module.
+ Raises:
+ ValueError: if `final_endpoint` is not recognized.
+ """
+
+ if final_endpoint not in self.VALID_ENDPOINTS:
+ raise ValueError('Unknown final endpoint %s' % final_endpoint)
+
+ super(InceptionI3d, self).__init__()
+ self._num_classes = num_classes
+ self._spatial_squeeze = spatial_squeeze
+ self._final_endpoint = final_endpoint
+ self.logits = None
+
+ if self._final_endpoint not in self.VALID_ENDPOINTS:
+ raise ValueError('Unknown final endpoint %s' % self._final_endpoint)
+
+ self.end_points = {}
+ end_point = 'Conv3d_1a_7x7'
+ self.end_points[end_point] = Unit3D(in_channels=in_channels, output_channels=64, kernel_shape=[7, 7, 7],
+ stride=(2, 2, 2), padding=(3,3,3), name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_2a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Conv3d_2b_1x1'
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=64, kernel_shape=[1, 1, 1], padding=0,
+ name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Conv3d_2c_3x3'
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=192, kernel_shape=[3, 3, 3], padding=1,
+ name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_3a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_3b'
+ self.end_points[end_point] = InceptionModule(192, [64,96,128,16,32,32], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_3c'
+ self.end_points[end_point] = InceptionModule(256, [128,128,192,32,96,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_4a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[3, 3, 3], stride=(2, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4b'
+ self.end_points[end_point] = InceptionModule(128+192+96+64, [192,96,208,16,48,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4c'
+ self.end_points[end_point] = InceptionModule(192+208+48+64, [160,112,224,24,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4d'
+ self.end_points[end_point] = InceptionModule(160+224+64+64, [128,128,256,24,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4e'
+ self.end_points[end_point] = InceptionModule(128+256+64+64, [112,144,288,32,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4f'
+ self.end_points[end_point] = InceptionModule(112+288+64+64, [256,160,320,32,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_5a_2x2'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[2, 2, 2], stride=(2, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_5b'
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [256,160,320,32,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_5c'
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [384,192,384,48,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Logits'
+ self.avg_pool = nn.AvgPool3d(kernel_size=[2, 7, 7],
+ stride=(1, 1, 1))
+ self.dropout = nn.Dropout(dropout_keep_prob)
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
+ kernel_shape=[1, 1, 1],
+ padding=0,
+ activation_fn=None,
+ use_batch_norm=False,
+ use_bias=True,
+ name='logits')
+
+ self.build()
+
+
+ def replace_logits(self, num_classes):
+ self._num_classes = num_classes
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
+ kernel_shape=[1, 1, 1],
+ padding=0,
+ activation_fn=None,
+ use_batch_norm=False,
+ use_bias=True,
+ name='logits')
+
+
+ def build(self):
+ for k in self.end_points.keys():
+ self.add_module(k, self.end_points[k])
+
+ def forward(self, x):
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x) # use _modules to work with dataparallel
+
+ x = self.logits(self.dropout(self.avg_pool(x)))
+ if self._spatial_squeeze:
+ logits = x.squeeze(3).squeeze(3)
+ logits = logits.mean(dim=2)
+ # logits is batch X time X classes, which is what we want to work with
+ return logits
+
+
+ def extract_features(self, x):
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x)
+ return self.avg_pool(x)
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..9f25788f1af258ce7bc0f66ee2f02561ccd96588
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/common_metrics_on_video_quality/utils.py
@@ -0,0 +1,27 @@
+import torch
+from torch.utils.data import Dataset, DataLoader
+
+class VideoDataset(Dataset):
+ def __init__(self, videos):
+ self.videos = videos
+
+ def __len__(self):
+ return len(self.videos) if type(self.videos) == list else self.videos.shape[0]
+
+ def __getitem__(self, idx):
+ video = self.videos[idx]
+ if isinstance(video, str):
+ video = torch.load(video, weights_only=True)
+ return video
+
+def build_dataloader(videos1, videos2):
+ dataset1 = VideoDataset(videos1)
+ dataset2 = VideoDataset(videos2)
+ assert len(dataset1) == len(dataset2)
+
+ dataloader1 = DataLoader(dataset1, batch_size=1, num_workers=8, shuffle=False)
+ dataloader2 = DataLoader(dataset2, batch_size=1, num_workers=8, shuffle=False)
+
+ return dataloader1, dataloader2
+
+
diff --git a/grn/tokenizer/videovae/evaluation/fid.py b/grn/tokenizer/videovae/evaluation/fid.py
new file mode 100644
index 0000000000000000000000000000000000000000..ee8c33b4ff437f3056bfba056411612e79bd8716
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/fid.py
@@ -0,0 +1,62 @@
+import numpy as np
+from scipy import linalg
+
+
+def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
+ """Numpy implementation of the Frechet Distance.
+ The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
+ and X_2 ~ N(mu_2, C_2) is
+ d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
+
+ Stable version by Dougal J. Sutherland.
+
+ Params:
+ -- mu1 : Numpy array containing the activations of a layer of the
+ inception net (like returned by the function 'get_predictions')
+ for generated samples.
+ -- mu2 : The sample mean over activations, precalculated on an
+ representative data set.
+ -- sigma1: The covariance matrix over activations for generated samples.
+ -- sigma2: The covariance matrix over activations, precalculated on an
+ representative data set.
+
+ Returns:
+ -- : The Frechet Distance.
+ """
+
+ mu1 = np.atleast_1d(mu1)
+ mu2 = np.atleast_1d(mu2)
+
+ sigma1 = np.atleast_2d(sigma1)
+ sigma2 = np.atleast_2d(sigma2)
+
+ assert (
+ mu1.shape == mu2.shape
+ ), "Training and test mean vectors have different lengths"
+ assert (
+ sigma1.shape == sigma2.shape
+ ), "Training and test covariances have different dimensions"
+
+ diff = mu1 - mu2
+
+ # Product might be almost singular
+ covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
+ if not np.isfinite(covmean).all():
+ msg = (
+ "fid calculation produces singular product; "
+ "adding %s to diagonal of cov estimates"
+ ) % eps
+ print(msg)
+ offset = np.eye(sigma1.shape[0]) * eps
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
+
+ # Numerical error might give slight imaginary component
+ if np.iscomplexobj(covmean):
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
+ m = np.max(np.abs(covmean.imag))
+ raise ValueError("Imaginary component {}".format(m))
+ covmean = covmean.real
+
+ tr_covmean = np.trace(covmean)
+
+ return diff.dot(diff) + np.trace(sigma1) + np.trace(sigma2) - 2 * tr_covmean
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/fvd.py b/grn/tokenizer/videovae/evaluation/fvd.py
new file mode 100644
index 0000000000000000000000000000000000000000..c699940f58ba44da49ceba8c7441dc571a594fd8
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/fvd.py
@@ -0,0 +1,150 @@
+import argparse
+from email.policy import strict
+import numpy as np
+
+import torch
+import torch.nn.functional as F
+import torch.utils.data as data
+
+from .pytorch_i3d import InceptionI3d
+import os
+from videovae.utils.misc import data_prefix_manager
+
+from sklearn.metrics.pairwise import polynomial_kernel
+
+MAX_BATCH = 16
+FVD_SAMPLE_SIZE = 2048
+TARGET_RESOLUTION = (224, 224)
+
+def preprocess(videos, target_resolution):
+ # videos in {0, ..., 255} as np.uint8 array
+ b, t, h, w, c = videos.shape
+ all_frames = torch.FloatTensor(videos).flatten(end_dim=1) # (b * t, h, w, c)
+
+ all_frames = all_frames.permute(0, 3, 1, 2).contiguous() # (b * t, c, h, w)
+ resized_videos = F.interpolate(all_frames, size=target_resolution,
+ mode='bilinear', align_corners=False)
+ resized_videos = resized_videos.view(b, t, c, *target_resolution)
+ output_videos = resized_videos.transpose(1, 2).contiguous() # (b, c, t, *)
+ scaled_videos = 2. * output_videos / 255. - 1 # [-1, 1]
+ return scaled_videos
+
+def get_fvd_logits(videos, i3d, device):
+ videos = preprocess(videos, TARGET_RESOLUTION)
+ embeddings = get_logits(i3d, videos, device)
+ return embeddings
+
+def load_fvd_model(device="cpu"):
+ i3d = InceptionI3d(400, in_channels=3).to(device)
+ i3d_path = data_prefix_manager('checkpoints/i3d_pretrained_400.pt')
+ i3d.load_state_dict(torch.load(i3d_path, map_location=device, weights_only=True))
+ i3d.eval()
+ return i3d
+
+
+def load_i3d_perceptual():
+ i3d = InceptionI3d(400, in_channels=3)
+ current_dir = os.path.dirname(os.path.abspath(__file__))
+ i3d_path = os.path.join(current_dir, 'i3d_pretrained_400.pt')
+ i3d.load_state_dict(torch.load(i3d_path, map_location=torch.device("cpu"), weights_only=True), strict=False)
+ for param in i3d.parameters():
+ param.requires_grad = False
+
+ return i3d
+
+# https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L161
+def _symmetric_matrix_square_root(mat, eps=1e-10):
+ u, s, v = torch.svd(mat)
+ si = torch.where(s < eps, s, torch.sqrt(s))
+ return torch.matmul(torch.matmul(u, torch.diag(si)), v.t())
+
+# https://github.com/tensorflow/gan/blob/de4b8da3853058ea380a6152bd3bd454013bf619/tensorflow_gan/python/eval/classifier_metrics.py#L400
+def trace_sqrt_product(sigma, sigma_v):
+ sqrt_sigma = _symmetric_matrix_square_root(sigma)
+ sqrt_a_sigmav_a = torch.matmul(sqrt_sigma, torch.matmul(sigma_v, sqrt_sigma))
+ return torch.trace(_symmetric_matrix_square_root(sqrt_a_sigmav_a))
+
+# https://discuss.pytorch.org/t/covariance-and-gradient-support/16217/2
+def cov(m, rowvar=False):
+ '''Estimate a covariance matrix given data.
+
+ Covariance indicates the level to which two variables vary together.
+ If we examine N-dimensional samples, `X = [x_1, x_2, ... x_N]^T`,
+ then the covariance matrix element `C_{ij}` is the covariance of
+ `x_i` and `x_j`. The element `C_{ii}` is the variance of `x_i`.
+
+ Args:
+ m: A 1-D or 2-D array containing multiple variables and observations.
+ Each row of `m` represents a variable, and each column a single
+ observation of all those variables.
+ rowvar: If `rowvar` is True, then each row represents a
+ variable, with observations in the columns. Otherwise, the
+ relationship is transposed: each column represents a variable,
+ while the rows contain observations.
+
+ Returns:
+ The covariance matrix of the variables.
+ '''
+ if m.dim() > 2:
+ raise ValueError('m has more than 2 dimensions')
+ if m.dim() < 2:
+ m = m.view(1, -1)
+ if not rowvar and m.size(0) != 1:
+ m = m.t()
+
+ fact = 1.0 / (m.size(1) - 1) # unbiased estimate
+ m_center = m - torch.mean(m, dim=1, keepdim=True)
+ mt = m_center.t() # if complex: mt = m.t().conj()
+ return fact * m_center.matmul(mt).squeeze()
+
+
+def frechet_distance(x1, x2):
+ x1 = x1.flatten(start_dim=1)
+ x2 = x2.flatten(start_dim=1)
+ m, m_w = x1.mean(dim=0), x2.mean(dim=0)
+ sigma, sigma_w = cov(x1, rowvar=False), cov(x2, rowvar=False)
+
+ sqrt_trace_component = trace_sqrt_product(sigma, sigma_w)
+ trace = torch.trace(sigma + sigma_w) - 2.0 * sqrt_trace_component
+
+ mean = torch.sum((m - m_w) ** 2)
+ fd = trace + mean
+ return fd
+
+
+def polynomial_mmd(X, Y):
+ m = X.shape[0]
+ n = Y.shape[0]
+ # compute kernels
+ K_XX = polynomial_kernel(X)
+ K_YY = polynomial_kernel(Y)
+ K_XY = polynomial_kernel(X, Y)
+ # compute mmd distance
+ K_XX_sum = (K_XX.sum() - np.diagonal(K_XX).sum()) / (m * (m - 1))
+ K_YY_sum = (K_YY.sum() - np.diagonal(K_YY).sum()) / (n * (n - 1))
+ K_XY_sum = K_XY.sum() / (m * n)
+ mmd = K_XX_sum + K_YY_sum - 2 * K_XY_sum
+ return mmd
+
+
+
+def get_logits(i3d, videos, device):
+ # assert videos.shape[0] % MAX_BATCH == 0
+ with torch.no_grad():
+ logits = []
+ for i in range(0, videos.shape[0], MAX_BATCH):
+ batch = videos[i:i + MAX_BATCH].to(device)
+ logits.append(i3d(batch))
+ logits = torch.cat(logits, dim=0)
+ return logits
+
+
+def compute_fvd(real, samples, i3d, device=torch.device('cpu')):
+ # i3d.to(device)
+ # real, samples are (N, T, H, W, C) numpy arrays in np.uint8
+ real, samples = preprocess(real, (224, 224)), preprocess(samples, (224, 224))
+ first_embed = get_logits(i3d, real, device)
+ second_embed = get_logits(i3d, samples, device)
+
+ return frechet_distance(first_embed, second_embed)
+
diff --git a/grn/tokenizer/videovae/evaluation/inception.py b/grn/tokenizer/videovae/evaluation/inception.py
new file mode 100644
index 0000000000000000000000000000000000000000..cb29d5d04a82d45022570bc6d6f4643d74417249
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/inception.py
@@ -0,0 +1,370 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torchvision import models
+from scipy import linalg
+import numpy as np
+
+try:
+ from torchvision.models.utils import load_state_dict_from_url
+except ImportError:
+ from torch.utils.model_zoo import load_url as load_state_dict_from_url
+
+# Inception weights ported to Pytorch from
+# http://download.tensorflow.org/models/image/imagenet/inception-2015-12-05.tgz
+FID_WEIGHTS_URL = 'https://github.com/mseitzer/pytorch-fid/releases/download/fid_weights/pt_inception-2015-12-05-6726825d.pth'
+
+FID_WEIGHTS_PATH = "../../pretrained/inception/pt_inception-2015-12-05-6726825d.pth"
+
+def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
+ """Numpy implementation of the Frechet Distance.
+ The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
+ and X_2 ~ N(mu_2, C_2) is
+ d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
+
+ Stable version by Dougal J. Sutherland.
+
+ Params:
+ -- mu1 : Numpy array containing the activations of a layer of the
+ inception net (like returned by the function 'get_predictions')
+ for generated samples.
+ -- mu2 : The sample mean over activations, precalculated on an
+ representative data set.
+ -- sigma1: The covariance matrix over activations for generated samples.
+ -- sigma2: The covariance matrix over activations, precalculated on an
+ representative data set.
+
+ Returns:
+ -- : The Frechet Distance.
+ """
+
+ mu1 = np.atleast_1d(mu1)
+ mu2 = np.atleast_1d(mu2)
+
+ sigma1 = np.atleast_2d(sigma1)
+ sigma2 = np.atleast_2d(sigma2)
+
+ assert mu1.shape == mu2.shape, \
+ 'Training and test mean vectors have different lengths'
+ assert sigma1.shape == sigma2.shape, \
+ 'Training and test covariances have different dimensions'
+
+ diff = mu1 - mu2
+
+ # Product might be almost singular
+ covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
+ if not np.isfinite(covmean).all():
+ msg = ('fid calculation produces singular product; '
+ 'adding %s to diagonal of cov estimates') % eps
+ print(msg)
+ offset = np.eye(sigma1.shape[0]) * eps
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
+
+ # Numerical error might give slight imaginary component
+ if np.iscomplexobj(covmean):
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
+ m = np.max(np.abs(covmean.imag))
+ raise ValueError('Imaginary component {}'.format(m))
+ covmean = covmean.real
+
+ tr_covmean = np.trace(covmean)
+
+ return (diff.dot(diff) + np.trace(sigma1) +
+ np.trace(sigma2) - 2 * tr_covmean)
+
+class InceptionV3(nn.Module):
+ """Pretrained InceptionV3 network returning feature maps"""
+
+ # Index of default block of inception to return,
+ # corresponds to output of final average pooling
+ DEFAULT_BLOCK_INDEX = 3
+
+ # Maps feature dimensionality to their output blocks indices
+ BLOCK_INDEX_BY_DIM = {
+ 64: 0, # First max pooling features
+ 192: 1, # Second max pooling featurs
+ 768: 2, # Pre-aux classifier features
+ 2048: 3 # Final average pooling features
+ }
+
+ def __init__(self,
+ output_blocks=[DEFAULT_BLOCK_INDEX],
+ resize_input=True,
+ normalize_input=True,
+ requires_grad=False,
+ use_fid_inception=True):
+ """Build pretrained InceptionV3
+
+ Parameters
+ ----------
+ output_blocks : list of int
+ Indices of blocks to return features of. Possible values are:
+ - 0: corresponds to output of first max pooling
+ - 1: corresponds to output of second max pooling
+ - 2: corresponds to output which is fed to aux classifier
+ - 3: corresponds to output of final average pooling
+ resize_input : bool
+ If true, bilinearly resizes input to width and height 299 before
+ feeding input to model. As the network without fully connected
+ layers is fully convolutional, it should be able to handle inputs
+ of arbitrary size, so resizing might not be strictly needed
+ normalize_input : bool
+ If true, scales the input from range (0, 1) to the range the
+ pretrained Inception network expects, namely (-1, 1)
+ requires_grad : bool
+ If true, parameters of the model require gradients. Possibly useful
+ for finetuning the network
+ use_fid_inception : bool
+ If true, uses the pretrained Inception model used in Tensorflow's
+ FID implementation. If false, uses the pretrained Inception model
+ available in torchvision. The FID Inception model has different
+ weights and a slightly different structure from torchvision's
+ Inception model. If you want to compute FID scores, you are
+ strongly advised to set this parameter to true to get comparable
+ results.
+ """
+ super(InceptionV3, self).__init__()
+
+ self.resize_input = resize_input
+ self.normalize_input = normalize_input
+ self.output_blocks = sorted(output_blocks)
+ self.last_needed_block = max(output_blocks)
+
+ assert self.last_needed_block <= 3, \
+ 'Last possible output block index is 3'
+
+ self.blocks = nn.ModuleList()
+
+ if use_fid_inception:
+ inception = fid_inception_v3()
+ else:
+ inception = models.inception_v3(pretrained=True)
+
+ # Block 0: input to maxpool1
+ block0 = [
+ inception.Conv2d_1a_3x3,
+ inception.Conv2d_2a_3x3,
+ inception.Conv2d_2b_3x3,
+ nn.MaxPool2d(kernel_size=3, stride=2)
+ ]
+ self.blocks.append(nn.Sequential(*block0))
+
+ # Block 1: maxpool1 to maxpool2
+ if self.last_needed_block >= 1:
+ block1 = [
+ inception.Conv2d_3b_1x1,
+ inception.Conv2d_4a_3x3,
+ nn.MaxPool2d(kernel_size=3, stride=2)
+ ]
+ self.blocks.append(nn.Sequential(*block1))
+
+ # Block 2: maxpool2 to aux classifier
+ if self.last_needed_block >= 2:
+ block2 = [
+ inception.Mixed_5b,
+ inception.Mixed_5c,
+ inception.Mixed_5d,
+ inception.Mixed_6a,
+ inception.Mixed_6b,
+ inception.Mixed_6c,
+ inception.Mixed_6d,
+ inception.Mixed_6e,
+ ]
+ self.blocks.append(nn.Sequential(*block2))
+
+ # Block 3: aux classifier to final avgpool
+ if self.last_needed_block >= 3:
+ block3 = [
+ inception.Mixed_7a,
+ inception.Mixed_7b,
+ inception.Mixed_7c,
+ nn.AdaptiveAvgPool2d(output_size=(1, 1))
+ ]
+ self.blocks.append(nn.Sequential(*block3))
+
+ for param in self.parameters():
+ param.requires_grad = requires_grad
+
+ def forward(self, inp):
+ """Get Inception feature maps
+
+ Parameters
+ ----------
+ inp : torch.autograd.Variable
+ Input tensor of shape Bx3xHxW. Values are expected to be in
+ range (0, 1)
+
+ Returns
+ -------
+ List of torch.autograd.Variable, corresponding to the selected output
+ block, sorted ascending by index
+ """
+ outp = []
+ x = inp
+
+ if self.resize_input:
+ x = F.interpolate(x,
+ size=(299, 299),
+ mode='bilinear',
+ align_corners=False)
+
+ if self.normalize_input:
+ x = 2 * x - 1 # Scale from range (0, 1) to range (-1, 1)
+
+ for idx, block in enumerate(self.blocks):
+ x = block(x)
+ if idx in self.output_blocks:
+ outp.append(x)
+
+ if idx == self.last_needed_block:
+ break
+
+ return outp
+
+
+def fid_inception_v3():
+ """Build pretrained Inception model for FID computation
+
+ The Inception model for FID computation uses a different set of weights
+ and has a slightly different structure than torchvision's Inception.
+
+ This method first constructs torchvision's Inception and then patches the
+ necessary parts that are different in the FID Inception model.
+ """
+ inception = models.inception_v3(num_classes=1008,
+ aux_logits=False,
+ pretrained=False)
+ inception.Mixed_5b = FIDInceptionA(192, pool_features=32)
+ inception.Mixed_5c = FIDInceptionA(256, pool_features=64)
+ inception.Mixed_5d = FIDInceptionA(288, pool_features=64)
+ inception.Mixed_6b = FIDInceptionC(768, channels_7x7=128)
+ inception.Mixed_6c = FIDInceptionC(768, channels_7x7=160)
+ inception.Mixed_6d = FIDInceptionC(768, channels_7x7=160)
+ inception.Mixed_6e = FIDInceptionC(768, channels_7x7=192)
+ inception.Mixed_7b = FIDInceptionE_1(1280)
+ inception.Mixed_7c = FIDInceptionE_2(2048)
+
+ state_dict = load_state_dict_from_url(FID_WEIGHTS_URL, progress=True)
+ # state_dict = torch.load(FID_WEIGHTS_PATH)
+ inception.load_state_dict(state_dict)
+ return inception
+
+
+class FIDInceptionA(models.inception.InceptionA):
+ """InceptionA block patched for FID computation"""
+ def __init__(self, in_channels, pool_features):
+ super(FIDInceptionA, self).__init__(in_channels, pool_features)
+
+ def forward(self, x):
+ branch1x1 = self.branch1x1(x)
+
+ branch5x5 = self.branch5x5_1(x)
+ branch5x5 = self.branch5x5_2(branch5x5)
+
+ branch3x3dbl = self.branch3x3dbl_1(x)
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
+ branch3x3dbl = self.branch3x3dbl_3(branch3x3dbl)
+
+ # Patch: Tensorflow's average pool does not use the padded zero's in
+ # its average calculation
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
+ count_include_pad=False)
+ branch_pool = self.branch_pool(branch_pool)
+
+ outputs = [branch1x1, branch5x5, branch3x3dbl, branch_pool]
+ return torch.cat(outputs, 1)
+
+
+class FIDInceptionC(models.inception.InceptionC):
+ """InceptionC block patched for FID computation"""
+ def __init__(self, in_channels, channels_7x7):
+ super(FIDInceptionC, self).__init__(in_channels, channels_7x7)
+
+ def forward(self, x):
+ branch1x1 = self.branch1x1(x)
+
+ branch7x7 = self.branch7x7_1(x)
+ branch7x7 = self.branch7x7_2(branch7x7)
+ branch7x7 = self.branch7x7_3(branch7x7)
+
+ branch7x7dbl = self.branch7x7dbl_1(x)
+ branch7x7dbl = self.branch7x7dbl_2(branch7x7dbl)
+ branch7x7dbl = self.branch7x7dbl_3(branch7x7dbl)
+ branch7x7dbl = self.branch7x7dbl_4(branch7x7dbl)
+ branch7x7dbl = self.branch7x7dbl_5(branch7x7dbl)
+
+ # Patch: Tensorflow's average pool does not use the padded zero's in
+ # its average calculation
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
+ count_include_pad=False)
+ branch_pool = self.branch_pool(branch_pool)
+
+ outputs = [branch1x1, branch7x7, branch7x7dbl, branch_pool]
+ return torch.cat(outputs, 1)
+
+
+class FIDInceptionE_1(models.inception.InceptionE):
+ """First InceptionE block patched for FID computation"""
+ def __init__(self, in_channels):
+ super(FIDInceptionE_1, self).__init__(in_channels)
+
+ def forward(self, x):
+ branch1x1 = self.branch1x1(x)
+
+ branch3x3 = self.branch3x3_1(x)
+ branch3x3 = [
+ self.branch3x3_2a(branch3x3),
+ self.branch3x3_2b(branch3x3),
+ ]
+ branch3x3 = torch.cat(branch3x3, 1)
+
+ branch3x3dbl = self.branch3x3dbl_1(x)
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
+ branch3x3dbl = [
+ self.branch3x3dbl_3a(branch3x3dbl),
+ self.branch3x3dbl_3b(branch3x3dbl),
+ ]
+ branch3x3dbl = torch.cat(branch3x3dbl, 1)
+
+ # Patch: Tensorflow's average pool does not use the padded zero's in
+ # its average calculation
+ branch_pool = F.avg_pool2d(x, kernel_size=3, stride=1, padding=1,
+ count_include_pad=False)
+ branch_pool = self.branch_pool(branch_pool)
+
+ outputs = [branch1x1, branch3x3, branch3x3dbl, branch_pool]
+ return torch.cat(outputs, 1)
+
+
+class FIDInceptionE_2(models.inception.InceptionE):
+ """Second InceptionE block patched for FID computation"""
+ def __init__(self, in_channels):
+ super(FIDInceptionE_2, self).__init__(in_channels)
+
+ def forward(self, x):
+ branch1x1 = self.branch1x1(x)
+
+ branch3x3 = self.branch3x3_1(x)
+ branch3x3 = [
+ self.branch3x3_2a(branch3x3),
+ self.branch3x3_2b(branch3x3),
+ ]
+ branch3x3 = torch.cat(branch3x3, 1)
+
+ branch3x3dbl = self.branch3x3dbl_1(x)
+ branch3x3dbl = self.branch3x3dbl_2(branch3x3dbl)
+ branch3x3dbl = [
+ self.branch3x3dbl_3a(branch3x3dbl),
+ self.branch3x3dbl_3b(branch3x3dbl),
+ ]
+ branch3x3dbl = torch.cat(branch3x3dbl, 1)
+
+ # Patch: The FID Inception model uses max pooling instead of average
+ # pooling. This is likely an error in this specific Inception
+ # implementation, as other Inception models use average pooling here
+ # (which matches the description in the paper).
+ branch_pool = F.max_pool2d(x, kernel_size=3, stride=1, padding=1)
+ branch_pool = self.branch_pool(branch_pool)
+
+ outputs = [branch1x1, branch3x3, branch3x3dbl, branch_pool]
+ return torch.cat(outputs, 1)
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/evaluation/pytorch_i3d.py b/grn/tokenizer/videovae/evaluation/pytorch_i3d.py
new file mode 100644
index 0000000000000000000000000000000000000000..afb7209eb186e07637fe4d7717ebe60f3fb8dbd0
--- /dev/null
+++ b/grn/tokenizer/videovae/evaluation/pytorch_i3d.py
@@ -0,0 +1,426 @@
+# https://github.com/piergiaj/pytorch-i3d
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torch.autograd import Variable
+
+import numpy as np
+
+import os
+import sys
+from collections import OrderedDict
+from videovae.modules.lpips import normalize_tensor
+
+
+class NetLinLayer(nn.Module):
+ """ A single linear layer which does a 1x1 conv """
+ def __init__(self, chn_in, chn_out=1, use_dropout=False):
+ super(NetLinLayer, self).__init__()
+ layers = [nn.Dropout(), ] if (use_dropout) else []
+ layers += [nn.Conv3d(chn_in, chn_out, 1, stride=1, padding=0, bias=False), ]
+ self.model = nn.Sequential(*layers)
+
+
+class MaxPool3dSamePadding(nn.MaxPool3d):
+
+ def compute_pad(self, dim, s):
+ if s % self.stride[dim] == 0:
+ return max(self.kernel_size[dim] - self.stride[dim], 0)
+ else:
+ return max(self.kernel_size[dim] - (s % self.stride[dim]), 0)
+
+ def forward(self, x):
+ # compute 'same' padding
+ (batch, channel, t, h, w) = x.size()
+ #print t,h,w
+ out_t = np.ceil(float(t) / float(self.stride[0]))
+ out_h = np.ceil(float(h) / float(self.stride[1]))
+ out_w = np.ceil(float(w) / float(self.stride[2]))
+ #print out_t, out_h, out_w
+ pad_t = self.compute_pad(0, t)
+ pad_h = self.compute_pad(1, h)
+ pad_w = self.compute_pad(2, w)
+ #print pad_t, pad_h, pad_w
+
+ pad_t_f = pad_t // 2
+ pad_t_b = pad_t - pad_t_f
+ pad_h_f = pad_h // 2
+ pad_h_b = pad_h - pad_h_f
+ pad_w_f = pad_w // 2
+ pad_w_b = pad_w - pad_w_f
+
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
+ #print x.size()
+ #print pad
+ x = F.pad(x, pad)
+ return super(MaxPool3dSamePadding, self).forward(x)
+
+
+class Unit3D(nn.Module):
+
+ def __init__(self, in_channels,
+ output_channels,
+ kernel_shape=(1, 1, 1),
+ stride=(1, 1, 1),
+ padding=0,
+ activation_fn=F.relu,
+ use_batch_norm=True,
+ use_bias=False,
+ name='unit_3d'):
+
+ """Initializes Unit3D module."""
+ super(Unit3D, self).__init__()
+
+ self._output_channels = output_channels
+ self._kernel_shape = kernel_shape
+ self._stride = stride
+ self._use_batch_norm = use_batch_norm
+ self._activation_fn = activation_fn
+ self._use_bias = use_bias
+ self.name = name
+ self.padding = padding
+
+ self.conv3d = nn.Conv3d(in_channels=in_channels,
+ out_channels=self._output_channels,
+ kernel_size=self._kernel_shape,
+ stride=self._stride,
+ padding=0, # we always want padding to be 0 here. We will dynamically pad based on input size in forward function
+ bias=self._use_bias)
+
+ if self._use_batch_norm:
+ self.bn = nn.BatchNorm3d(self._output_channels, eps=1e-5, momentum=0.001)
+
+ def compute_pad(self, dim, s):
+ if s % self._stride[dim] == 0:
+ return max(self._kernel_shape[dim] - self._stride[dim], 0)
+ else:
+ return max(self._kernel_shape[dim] - (s % self._stride[dim]), 0)
+
+
+ def forward(self, x):
+ # compute 'same' padding
+ (batch, channel, t, h, w) = x.size()
+ #print t,h,w
+ out_t = np.ceil(float(t) / float(self._stride[0]))
+ out_h = np.ceil(float(h) / float(self._stride[1]))
+ out_w = np.ceil(float(w) / float(self._stride[2]))
+ #print out_t, out_h, out_w
+ pad_t = self.compute_pad(0, t)
+ pad_h = self.compute_pad(1, h)
+ pad_w = self.compute_pad(2, w)
+ #print pad_t, pad_h, pad_w
+
+ pad_t_f = pad_t // 2
+ pad_t_b = pad_t - pad_t_f
+ pad_h_f = pad_h // 2
+ pad_h_b = pad_h - pad_h_f
+ pad_w_f = pad_w // 2
+ pad_w_b = pad_w - pad_w_f
+
+ pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)
+ #print x.size()
+ #print pad
+ x = F.pad(x, pad)
+ #print x.size()
+
+ x = self.conv3d(x)
+ if self._use_batch_norm:
+ x = self.bn(x)
+ if self._activation_fn is not None:
+ x = self._activation_fn(x)
+ return x
+
+
+
+class InceptionModule(nn.Module):
+ def __init__(self, in_channels, out_channels, name):
+ super(InceptionModule, self).__init__()
+
+ self.b0 = Unit3D(in_channels=in_channels, output_channels=out_channels[0], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_0/Conv3d_0a_1x1')
+ self.b1a = Unit3D(in_channels=in_channels, output_channels=out_channels[1], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_1/Conv3d_0a_1x1')
+ self.b1b = Unit3D(in_channels=out_channels[1], output_channels=out_channels[2], kernel_shape=[3, 3, 3],
+ name=name+'/Branch_1/Conv3d_0b_3x3')
+ self.b2a = Unit3D(in_channels=in_channels, output_channels=out_channels[3], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_2/Conv3d_0a_1x1')
+ self.b2b = Unit3D(in_channels=out_channels[3], output_channels=out_channels[4], kernel_shape=[3, 3, 3],
+ name=name+'/Branch_2/Conv3d_0b_3x3')
+ self.b3a = MaxPool3dSamePadding(kernel_size=[3, 3, 3],
+ stride=(1, 1, 1), padding=0)
+ self.b3b = Unit3D(in_channels=in_channels, output_channels=out_channels[5], kernel_shape=[1, 1, 1], padding=0,
+ name=name+'/Branch_3/Conv3d_0b_1x1')
+ self.name = name
+
+ def forward(self, x):
+ b0 = self.b0(x)
+ b1 = self.b1b(self.b1a(x))
+ b2 = self.b2b(self.b2a(x))
+ b3 = self.b3b(self.b3a(x))
+ return torch.cat([b0,b1,b2,b3], dim=1)
+
+
+class InceptionI3d(nn.Module):
+ """Inception-v1 I3D architecture.
+ The model is introduced in:
+ Quo Vadis, Action Recognition? A New Model and the Kinetics Dataset
+ Joao Carreira, Andrew Zisserman
+ https://arxiv.org/pdf/1705.07750v1.pdf.
+ See also the Inception architecture, introduced in:
+ Going deeper with convolutions
+ Christian Szegedy, Wei Liu, Yangqing Jia, Pierre Sermanet, Scott Reed,
+ Dragomir Anguelov, Dumitru Erhan, Vincent Vanhoucke, Andrew Rabinovich.
+ http://arxiv.org/pdf/1409.4842v1.pdf.
+ """
+
+ # Endpoints of the model in order. During construction, all the endpoints up
+ # to a designated `final_endpoint` are returned in a dictionary as the
+ # second return value.
+ VALID_ENDPOINTS = (
+ 'Conv3d_1a_7x7',
+ 'MaxPool3d_2a_3x3',
+ 'Conv3d_2b_1x1',
+ 'Conv3d_2c_3x3',
+ 'MaxPool3d_3a_3x3',
+ 'Mixed_3b',
+ 'Mixed_3c',
+ 'MaxPool3d_4a_3x3',
+ 'Mixed_4b',
+ 'Mixed_4c',
+ 'Mixed_4d',
+ 'Mixed_4e',
+ 'Mixed_4f',
+ 'MaxPool3d_5a_2x2',
+ 'Mixed_5b',
+ 'Mixed_5c',
+ 'Logits',
+ 'Predictions',
+ )
+
+ FEAT_ENDPOINTS = (
+ 'Conv3d_1a_7x7',
+ 'Conv3d_2c_3x3',
+ 'Mixed_3c',
+ 'Mixed_4f',
+ 'Mixed_5c',
+ )
+ def __init__(self,
+ num_classes=400,
+ spatial_squeeze=True,
+ final_endpoint='Logits',
+ name='inception_i3d',
+ in_channels=3,
+ dropout_keep_prob=0.5,
+ is_coinrun=False,
+ ):
+ """Initializes I3D model instance.
+ Args:
+ num_classes: The number of outputs in the logit layer (default 400, which
+ matches the Kinetics dataset).
+ spatial_squeeze: Whether to squeeze the spatial dimensions for the logits
+ before returning (default True).
+ final_endpoint: The model contains many possible endpoints.
+ `final_endpoint` specifies the last endpoint for the model to be built
+ up to. In addition to the output at `final_endpoint`, all the outputs
+ at endpoints up to `final_endpoint` will also be returned, in a
+ dictionary. `final_endpoint` must be one of
+ InceptionI3d.VALID_ENDPOINTS (default 'Logits').
+ name: A string (optional). The name of this module.
+ Raises:
+ ValueError: if `final_endpoint` is not recognized.
+ """
+
+ if final_endpoint not in self.VALID_ENDPOINTS:
+ raise ValueError('Unknown final endpoint %s' % final_endpoint)
+
+ super(InceptionI3d, self).__init__()
+ self._num_classes = num_classes
+ self._spatial_squeeze = spatial_squeeze
+ self._final_endpoint = final_endpoint
+ self.logits = None
+ self.is_coinrun = is_coinrun
+
+ if self._final_endpoint not in self.VALID_ENDPOINTS:
+ raise ValueError('Unknown final endpoint %s' % self._final_endpoint)
+
+ self.end_points = {}
+ end_point = 'Conv3d_1a_7x7'
+ self.end_points[end_point] = Unit3D(in_channels=in_channels, output_channels=64, kernel_shape=[7, 7, 7],
+ stride=(1 if is_coinrun else 2, 2, 2), padding=(3,3,3), name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_2a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Conv3d_2b_1x1'
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=64, kernel_shape=[1, 1, 1], padding=0,
+ name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Conv3d_2c_3x3'
+ self.end_points[end_point] = Unit3D(in_channels=64, output_channels=192, kernel_shape=[3, 3, 3], padding=1,
+ name=name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_3a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_3b'
+ self.end_points[end_point] = InceptionModule(192, [64,96,128,16,32,32], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_3c'
+ self.end_points[end_point] = InceptionModule(256, [128,128,192,32,96,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_4a_3x3'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1 if is_coinrun else 3, 3, 3], stride=(1 if is_coinrun else 2, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4b'
+ self.end_points[end_point] = InceptionModule(128+192+96+64, [192,96,208,16,48,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4c'
+ self.end_points[end_point] = InceptionModule(192+208+48+64, [160,112,224,24,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4d'
+ self.end_points[end_point] = InceptionModule(160+224+64+64, [128,128,256,24,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4e'
+ self.end_points[end_point] = InceptionModule(128+256+64+64, [112,144,288,32,64,64], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_4f'
+ self.end_points[end_point] = InceptionModule(112+288+64+64, [256,160,320,32,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'MaxPool3d_5a_2x2'
+ self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[2, 2, 2], stride=(1 if is_coinrun else 2, 2, 2),
+ padding=0)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_5b'
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [256,160,320,32,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Mixed_5c'
+ self.end_points[end_point] = InceptionModule(256+320+128+128, [384,192,384,48,128,128], name+end_point)
+ if self._final_endpoint == end_point: return
+
+ end_point = 'Logits'
+ self.avg_pool = nn.AvgPool3d(kernel_size=[1, 8, 8] if is_coinrun else [2, 7, 7],
+ stride=(1, 1, 1))
+ self.dropout = nn.Dropout(dropout_keep_prob)
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
+ kernel_shape=[1, 1, 1],
+ padding=0,
+ activation_fn=None,
+ use_batch_norm=False,
+ use_bias=True,
+ name='logits')
+
+ self.build()
+ self.chns = [64, 192, 480, 832, 1024] # i3d features
+ """use_dropout= True
+ self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout)
+ self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout)
+ self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout)
+ self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout)
+ self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout)"""
+
+
+ def replace_logits(self, num_classes):
+ self._num_classes = num_classes
+ self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,
+ kernel_shape=[1, 1, 1],
+ padding=0,
+ activation_fn=None,
+ use_batch_norm=False,
+ use_bias=True,
+ name='logits')
+
+
+ def build(self):
+ for k in self.end_points.keys():
+ self.add_module(k, self.end_points[k])
+
+ def forward(self, x):
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x) # use _modules to work with dataparallel
+
+ x = self.logits(self.dropout(self.avg_pool(x)))
+ if self._spatial_squeeze:
+ logits = x.squeeze(3).squeeze(3)
+ logits = logits.mean(dim=2)
+ # logits is batch X time X classes, which is what we want to work with
+ return logits
+
+
+ def extract_features(self, x):
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x)
+ return self.avg_pool(x)
+
+
+ def extract_pre_pool_features(self, x):
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x)
+ return x
+
+
+ """
+ def forward(self, input, target):
+ in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target))
+ outs0, outs1 = self.net(in0_input), self.net(in1_input)
+ feats0, feats1, diffs = {}, {}, {}
+ lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
+ for kk in range(len(self.chns)):
+ feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(outs1[kk])
+ diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
+
+ res = [spatial_average(lins[kk].model(diffs[kk]), keepdim=True) for kk in range(len(self.chns))]
+ val = res[0]
+ for l in range(1, len(self.chns)):
+ val += res[l]
+ return val
+ """
+
+ def extract_features_multiscale(self, x):
+ xs = []
+ for end_point in self.VALID_ENDPOINTS:
+ if end_point in self.end_points:
+ x = self._modules[end_point](x)
+ if end_point in self.FEAT_ENDPOINTS:
+ xs.append(x)
+ return xs
+
+
+ def extract_perceptual(self, input, target):
+ outs0, outs1 = self.extract_features_multiscale(input), self.extract_features_multiscale(target)
+ feats0, feats1, diffs = {}, {}, {}
+ # lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
+ for kk in range(len(self.chns)):
+ feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(outs1[kk])
+ diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
+
+ assert len(self.chns) == len(diffs)
+ res = [spatial_temporal_average(diffs[kk], keepdim=True) for kk in range(len(self.chns))]
+ val = res[0]
+ for l in range(1, len(self.chns)):
+ val += res[l]
+ return val
+
+
+def spatial_temporal_average(x, keepdim=True):
+ return x.mean([1, 2, 3, 4], keepdim=keepdim)
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/models/discriminator.py b/grn/tokenizer/videovae/models/discriminator.py
new file mode 100644
index 0000000000000000000000000000000000000000..633186b349475f0534e7204973845d0d8938ccae
--- /dev/null
+++ b/grn/tokenizer/videovae/models/discriminator.py
@@ -0,0 +1,516 @@
+import numpy as np
+import random
+import functools
+import math
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.utils.checkpoint as checkpoint
+from einops import rearrange
+from videovae.utils.misc import set_tf32_flags
+from videovae.modules.normalization import Normalize
+
+
+class DiscriminatorPool:
+ def __init__(self, pool_size):
+ self.pool_size = int(pool_size)
+ self.num_imgs = 0
+ self.images = []
+
+ def query(self, images):
+ if self.pool_size == 0:
+ return images
+
+ return_images = []
+ for image in images:
+ if self.num_imgs < self.pool_size:
+ self.images.append(image)
+ self.num_imgs += 1
+ return_images.append(image)
+ else:
+ if random.uniform(0, 1) > 0.5:
+ i = random.randint(0, self.pool_size - 1)
+ tmp = self.images[i].clone()
+ self.images[i] = image
+ return_images.append(tmp)
+ else:
+ return_images.append(image)
+ return torch.stack(return_images)
+
+class ImageDiscriminator(nn.Module):
+ def __init__(self, args):
+ super().__init__()
+ if args.disc_type == "stylegan":
+ self.discriminator = StyleGANDiscriminator(use_blur=args.disc_use_blur, downsample_base=args.disc_stylegan_downsample_base)
+ elif args.disc_type == "spectralgan":
+ print("using NLayerSpectralDiscriminator")
+ self.discriminator = NLayerSpectralDiscriminator()
+ else:
+ self.discriminator = NLayerDiscriminator() # PatchGAN, default; args.disc_pool=no, default; temporal_compress=yes, default;
+ self.disc_pool = args.disc_pool
+ if args.disc_pool == "yes":
+ self.real_pool = DiscriminatorPool(pool_size=args.batch_size[0] * args.disc_pool_size)
+ self.fake_pool = DiscriminatorPool(pool_size=args.batch_size[0] * args.disc_pool_size)
+
+ def forward(self, x, pool_name=None):
+ if pool_name and self.disc_pool == "yes":
+ assert pool_name in ["real", "fake"]
+ if pool_name == "real":
+ x = self.real_pool.query(x)
+ elif pool_name == "fake":
+ x = self.fake_pool.query(x)
+ return self.discriminator(x)
+
+class VideoDiscriminator(nn.Module):
+ def __init__(self, args):
+ super().__init__()
+ if args.disc_type == "stylegan":
+ self.discriminator = StyleGANDiscriminator(conv_type="3d", use_blur=args.disc_use_blur, downsample_base=args.disc_stylegan_downsample_base)
+ # self.discriminator = MagvitDiscriminator(args.image_channels, apply_blur=args.apply_blur, apply_noise=args.apply_noise, model_type="3d", use_checkpoint=args.use_checkpoint, version=args.disc_version, norm_type=args.norm_type)
+ elif args.disc_type == "spectralgan":
+ print("using NLayerSpectralDiscriminator3D")
+ self.discriminator = NLayerSpectralDiscriminator3D(temporal_compress=args.disc_temporal_compress)
+ else: # PatchGAN, default; args.disc_pool=no, default; temporal_compress=yes, default;
+ self.discriminator = NLayerDiscriminator3D(temporal_compress=args.disc_temporal_compress)
+ # self.discriminator = NLayerDiscriminator3D(args.image_channels, args.disc_channels, args.disc_layers, args.norm_type, use_sigmoid=args.sigmoid_in_disc, activation=args.activation_in_disc, apply_blur=args.apply_blur, apply_noise=args.apply_noise, upcast_tf32=args.upcast_tf32)
+ self.disc_pool = args.disc_pool
+ if args.disc_pool == "yes":
+ self.real_pool = DiscriminatorPool(pool_size=args.batch_size[0] * args.disc_pool_size)
+ self.fake_pool = DiscriminatorPool(pool_size=args.batch_size[0] * args.disc_pool_size)
+
+ def forward(self, x, pool_name=None):
+ if pool_name and self.disc_pool == "yes":
+ assert pool_name in ["real", "fake"]
+ if pool_name == "real":
+ x = self.real_pool.query(x)
+ elif pool_name == "fake":
+ x = self.fake_pool.query(x)
+ return self.discriminator(x)
+
+
+class NLayerSpectralDiscriminator(nn.Module):
+ """Defines a PatchGAN discriminator as in Pix2Pix
+ --> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
+ """
+ def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False):
+ """Construct a PatchGAN discriminator
+ Parameters:
+ input_nc (int) -- the number of channels in input images
+ ndf (int) -- the number of filters in the last conv layer
+ n_layers (int) -- the number of conv layers in the discriminator
+ norm_layer -- normalization layer
+ """
+ super(NLayerSpectralDiscriminator, self).__init__()
+ kw = 4
+ padw = 1
+ sequence = [nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]
+ nf_mult = 1
+ nf_mult_prev = 1
+ for n in range(1, n_layers): # gradually increase the number of filters
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n, 8)
+ sequence += [
+ torch.nn.utils.spectral_norm(nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=2, padding=padw, bias=True)),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n_layers, 8)
+ sequence += [
+ nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=True),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ sequence += [
+ nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)] # output 1 channel prediction map
+ self.main = nn.Sequential(*sequence)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, module):
+ if isinstance(module, nn.Conv2d):
+ nn.init.normal_(module.weight.data, 0.0, 0.02)
+ elif isinstance(module, nn.BatchNorm2d):
+ nn.init.normal_(module.weight.data, 1.0, 0.02)
+ nn.init.constant_(module.bias.data, 0)
+
+ def forward(self, input):
+ """Standard forward."""
+ return self.main(input)
+
+
+class NLayerSpectralDiscriminator3D(nn.Module):
+ """Defines a PatchGAN discriminator as in Pix2Pix
+ --> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
+ """
+ def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False, temporal_compress="yes"):
+ """Construct a PatchGAN discriminator
+ Parameters:
+ input_nc (int) -- the number of channels in input images
+ ndf (int) -- the number of filters in the last conv layer
+ n_layers (int) -- the number of conv layers in the discriminator
+ norm_layer -- normalization layer
+ """
+ super(NLayerSpectralDiscriminator3D, self).__init__()
+ kw = 4
+ padw = 1
+ sequence = [nn.Conv3d(input_nc, ndf, kernel_size=kw, stride=(1,2,2), padding=padw), nn.LeakyReLU(0.2, True)]
+ nf_mult = 1
+ nf_mult_prev = 1
+ _stride = 2 if temporal_compress == "yes" else (1,2,2)
+ for n in range(1, n_layers): # gradually increase the number of filters
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n, 8)
+ sequence += [
+ torch.nn.utils.spectral_norm(nn.Conv3d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=_stride, padding=padw, bias=True)),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n_layers, 8)
+ sequence += [
+ torch.nn.utils.spectral_norm(nn.Conv3d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=True)),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ sequence += [
+ nn.Conv3d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)] # output 1 channel prediction map
+ self.main = nn.Sequential(*sequence)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, module):
+ if isinstance(module, nn.Conv3d):
+ nn.init.normal_(module.weight.data, 0.0, 0.02)
+ elif isinstance(module, nn.BatchNorm2d):
+ nn.init.normal_(module.weight.data, 1.0, 0.02)
+ nn.init.constant_(module.bias.data, 0)
+
+ def forward(self, input):
+ """Standard forward."""
+ return self.main(input)
+
+
+class NLayerDiscriminator(nn.Module):
+ """Defines a PatchGAN discriminator as in Pix2Pix
+ --> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
+ """
+ def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False):
+ """Construct a PatchGAN discriminator
+ Parameters:
+ input_nc (int) -- the number of channels in input images
+ ndf (int) -- the number of filters in the last conv layer
+ n_layers (int) -- the number of conv layers in the discriminator
+ norm_layer -- normalization layer
+ """
+ super(NLayerDiscriminator, self).__init__()
+ norm_type = "batch"
+ use_bias = norm_type != "batch"
+
+ kw = 4
+ padw = 1
+ sequence = [nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]
+ nf_mult = 1
+ nf_mult_prev = 1
+ for n in range(1, n_layers): # gradually increase the number of filters
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n, 8)
+ sequence += [
+ nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=2, padding=padw, bias=use_bias),
+ Normalize(ndf * nf_mult, norm_type=norm_type),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n_layers, 8)
+ sequence += [
+ nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=use_bias),
+ Normalize(ndf * nf_mult, norm_type=norm_type),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ sequence += [
+ nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)] # output 1 channel prediction map
+ self.main = nn.Sequential(*sequence)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, module):
+ if isinstance(module, nn.Conv2d):
+ nn.init.normal_(module.weight.data, 0.0, 0.02)
+ elif isinstance(module, nn.BatchNorm2d):
+ nn.init.normal_(module.weight.data, 1.0, 0.02)
+ nn.init.constant_(module.bias.data, 0)
+
+ def forward(self, input):
+ """Standard forward."""
+ return self.main(input)
+
+class NLayerDiscriminator3D(nn.Module):
+ """Defines a PatchGAN discriminator as in Pix2Pix
+ --> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
+ """
+ def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False, temporal_compress="yes"):
+ """Construct a PatchGAN discriminator
+ Parameters:
+ input_nc (int) -- the number of channels in input images
+ ndf (int) -- the number of filters in the last conv layer
+ n_layers (int) -- the number of conv layers in the discriminator
+ norm_layer -- normalization layer
+ """
+ super(NLayerDiscriminator3D, self).__init__()
+ norm_type = "batch"
+ use_bias = norm_type != "batch"
+
+ kw = 4
+ padw = 1
+ sequence = [nn.Conv3d(input_nc, ndf, kernel_size=kw, stride=(1,2,2), padding=padw), nn.LeakyReLU(0.2, True)]
+ nf_mult = 1
+ nf_mult_prev = 1
+ _stride = 2 if temporal_compress == "yes" else (1,2,2)
+ for n in range(1, n_layers): # gradually increase the number of filters
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n, 8)
+ sequence += [
+ nn.Conv3d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=_stride, padding=padw, bias=use_bias),
+ Normalize(ndf * nf_mult, norm_type=norm_type),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ nf_mult_prev = nf_mult
+ nf_mult = min(2 ** n_layers, 8)
+ sequence += [
+ nn.Conv3d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=use_bias),
+ Normalize(ndf * nf_mult, norm_type=norm_type),
+ nn.LeakyReLU(0.2, True)
+ ]
+
+ sequence += [
+ nn.Conv3d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)] # output 1 channel prediction map
+ self.main = nn.Sequential(*sequence)
+
+ self.apply(self._init_weights)
+
+ def _init_weights(self, module):
+ if isinstance(module, nn.Conv3d):
+ nn.init.normal_(module.weight.data, 0.0, 0.02)
+ elif isinstance(module, nn.BatchNorm2d):
+ nn.init.normal_(module.weight.data, 1.0, 0.02)
+ nn.init.constant_(module.bias.data, 0)
+
+ def forward(self, input):
+ """Standard forward."""
+ return self.main(input)
+
+class ActNorm(nn.Module):
+ def __init__(self, num_features, logdet=False, affine=True,
+ allow_reverse_init=False):
+ assert affine
+ super().__init__()
+ self.logdet = logdet
+ self.loc = nn.Parameter(torch.zeros(1, num_features, 1, 1))
+ self.scale = nn.Parameter(torch.ones(1, num_features, 1, 1))
+ self.allow_reverse_init = allow_reverse_init
+
+ self.register_buffer('initialized', torch.tensor(0, dtype=torch.uint8))
+
+ def initialize(self, input):
+ with torch.no_grad():
+ flatten = input.permute(1, 0, 2, 3).contiguous().view(input.shape[1], -1)
+ mean = (
+ flatten.mean(1)
+ .unsqueeze(1)
+ .unsqueeze(2)
+ .unsqueeze(3)
+ .permute(1, 0, 2, 3)
+ )
+ std = (
+ flatten.std(1)
+ .unsqueeze(1)
+ .unsqueeze(2)
+ .unsqueeze(3)
+ .permute(1, 0, 2, 3)
+ )
+
+ self.loc.data.copy_(-mean)
+ self.scale.data.copy_(1 / (std + 1e-6))
+
+ def forward(self, input, reverse=False):
+ if reverse:
+ return self.reverse(input)
+ if len(input.shape) == 2:
+ input = input[:,:,None,None]
+ squeeze = True
+ else:
+ squeeze = False
+
+ _, _, height, width = input.shape
+
+ if self.training and self.initialized.item() == 0:
+ self.initialize(input)
+ self.initialized.fill_(1)
+
+ h = self.scale * (input + self.loc)
+
+ if squeeze:
+ h = h.squeeze(-1).squeeze(-1)
+
+ if self.logdet:
+ log_abs = torch.log(torch.abs(self.scale))
+ logdet = height*width*torch.sum(log_abs)
+ logdet = logdet * torch.ones(input.shape[0]).to(input)
+ return h, logdet
+
+ return h
+
+ def reverse(self, output):
+ if self.training and self.initialized.item() == 0:
+ if not self.allow_reverse_init:
+ raise RuntimeError(
+ "Initializing ActNorm in reverse direction is "
+ "disabled by default. Use allow_reverse_init=True to enable."
+ )
+ else:
+ self.initialize(output)
+ self.initialized.fill_(1)
+
+ if len(output.shape) == 2:
+ output = output[:,:,None,None]
+ squeeze = True
+ else:
+ squeeze = False
+
+ h = output / self.scale - self.loc
+
+ if squeeze:
+ h = h.squeeze(-1).squeeze(-1)
+ return h
+
+class ResBlockDown(nn.Module):
+ def __init__(self, ic, oc, model_type="2d"):
+ super(ResBlockDown, self).__init__()
+ assert model_type in ["2d", "3d"]
+ activation_func = nn.LeakyReLU(0.2, True)
+
+ if model_type == "2d":
+ self.branch1 = nn.Sequential(
+ nn.Conv2d(ic, oc, kernel_size=3, stride=1, padding=1),
+ activation_func,
+ nn.AvgPool2d(kernel_size=2, stride=2),
+ nn.Conv2d(oc, oc, kernel_size=3, stride=1, padding=1),
+ activation_func,
+ )
+ self.branch2 = nn.Sequential(
+ nn.AvgPool2d(kernel_size=2, stride=2),
+ nn.Conv2d(ic, oc, kernel_size=1, stride=1, padding=0)
+ )
+ else:
+ raise NotImplementedError
+
+ def forward(self, x):
+ return self.branch1(x) + self.branch2(x)
+
+########################################
+# StyleGAN #
+########################################
+class StyleGANDiscriminator(nn.Module):
+ def __init__(self, input_nc=3, ndf=64, n_layers=3, channel_multiplier=1, image_size=256, conv_type="2d", use_blur=True, downsample_base=2):
+ super().__init__()
+ channels = {
+ 4: 512,
+ 8: 512,
+ 16: 512,
+ 32: 512,
+ 64: 256 * channel_multiplier,
+ 128: 128 * channel_multiplier,
+ 256: 64 * channel_multiplier,
+ 512: 32 * channel_multiplier,
+ 1024: 16 * channel_multiplier,
+ }
+ assert conv_type in ["2d", "3d"]
+ conv = nn.Conv2d if conv_type == "2d" else nn.Conv3d
+
+ log_size = int(math.log(image_size, 2))
+ in_channel = channels[image_size]
+
+ blocks = [conv(input_nc, in_channel, 3, padding=1), leaky_relu()]
+ for i in range(log_size, downsample_base, -1):
+ out_channel = channels[2 ** (i - 1)]
+ blocks.append(DiscriminatorBlock(in_channel, out_channel, conv_type=conv_type, use_blur=use_blur))
+ in_channel = out_channel
+ self.blocks = nn.ModuleList(blocks)
+
+ self.final_conv = nn.Sequential(
+ conv(in_channel, channels[4], 3, padding=1),
+ leaky_relu(),
+ )
+ self.final_linear = nn.Sequential(
+ conv(channels[4], 1, 1, padding=0)
+ )
+
+ def forward(self, x):
+ for block in self.blocks:
+ x = block(x)
+ x = self.final_conv(x)
+ x = self.final_linear(x)
+ return x
+
+
+class DiscriminatorBlock(nn.Module):
+ def __init__(self, input_channels, filters, downsample=True, conv_type="2d", use_blur=True):
+ super().__init__()
+ assert conv_type in ["2d", "3d"]
+ conv = nn.Conv2d if conv_type == "2d" else nn.Conv3d
+ if conv_type == "2d":
+ conv = nn.Conv2d
+ stride = 2 if downsample else 1
+ elif conv_type == "3d":
+ conv = nn.Conv3d
+ stride = (1,2,2) if downsample else (1,1,1)
+
+ self.conv_res = conv(input_channels, filters, 1, stride = stride)
+
+ self.net = nn.Sequential(
+ conv(input_channels, filters, 3, padding=1),
+ leaky_relu(),
+ conv(filters, filters, 3, padding=1),
+ leaky_relu()
+ )
+
+ self.downsample = nn.Sequential(
+ Blur() if use_blur else nn.Identity(),
+ conv(filters, filters, 3, padding = 1, stride = stride)
+ ) if downsample else None
+
+ def forward(self, x):
+ res = self.conv_res(x)
+ x = self.net(x)
+ if exists(self.downsample):
+ x = self.downsample(x)
+ x = (x + res) * (1 / math.sqrt(2))
+ return x
+
+class Blur(nn.Module):
+ def __init__(self):
+ super().__init__()
+ f = torch.Tensor([1, 2, 1])
+ self.register_buffer('f', f)
+
+ def forward(self, x):
+ is_image = x.ndim == 4
+ if not is_image:
+ b = x.shape[0]
+ x = rearrange(x, "b c t h w -> (b t) c h w")
+ f = self.f
+ f = f[None, None, :] * f [None, :, None]
+ from kornia.filters import filter2d
+ x = filter2d(x, f, normalized=True)
+ if not is_image:
+ x = rearrange(x, "(b t) c h w -> b c t h w", b=b)
+ return x
+
+def leaky_relu(p=0.2):
+ return nn.LeakyReLU(p, inplace=True)
+
+def exists(val):
+ return val is not None
diff --git a/grn/tokenizer/videovae/models/hbq_tokenizer.py b/grn/tokenizer/videovae/models/hbq_tokenizer.py
new file mode 100644
index 0000000000000000000000000000000000000000..bd4e3c2d7d38c97ca02470ee220f42f2204e5185
--- /dev/null
+++ b/grn/tokenizer/videovae/models/hbq_tokenizer.py
@@ -0,0 +1,1233 @@
+# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
+import logging
+import os
+import os.path as osp
+
+import torch
+import torch.cuda.amp as amp
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+from einops import rearrange
+
+from videovae.utils.dynamic_resolution_two_pyramid import dynamic_resolution_thw, h_div_w_templates
+
+
+CACHE_T = 2
+
+
+class CausalConv3d(nn.Conv3d):
+ """
+ Causal 3d convolusion.
+ """
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self._padding = (
+ self.padding[2],
+ self.padding[2],
+ self.padding[1],
+ self.padding[1],
+ 2 * self.padding[0],
+ 0,
+ )
+ self.padding = (0, 0, 0)
+
+ def forward(self, x, cache_x=None):
+ padding = list(self._padding)
+ if cache_x is not None and self._padding[4] > 0:
+ cache_x = cache_x.to(x.device)
+ x = torch.cat([cache_x, x], dim=2)
+ padding[4] -= cache_x.shape[2]
+ x = F.pad(x, padding)
+
+ return super().forward(x)
+
+
+class RMS_norm(nn.Module):
+
+ def __init__(self, dim, channel_first=True, images=True, bias=False):
+ super().__init__()
+ broadcastable_dims = (1, 1, 1) if not images else (1, 1)
+ shape = (dim, *broadcastable_dims) if channel_first else (dim,)
+
+ self.channel_first = channel_first
+ self.scale = dim**0.5
+ self.gamma = nn.Parameter(torch.ones(shape))
+ self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
+
+ def forward(self, x):
+ return (F.normalize(x, dim=(1 if self.channel_first else -1)) *
+ self.scale * self.gamma + self.bias)
+
+
+class Upsample(nn.Upsample):
+
+ def forward(self, x):
+ """
+ Fix bfloat16 support for nearest neighbor interpolation.
+ """
+ return super().forward(x.float()).type_as(x)
+
+
+class Resample(nn.Module):
+
+ def __init__(self, dim, mode):
+ assert mode in (
+ "none",
+ "upsample2d",
+ "upsample3d",
+ "downsample2d",
+ "downsample3d",
+ )
+ super().__init__()
+ self.dim = dim
+ self.mode = mode
+
+ # layers
+ if mode == "upsample2d":
+ self.resample = nn.Sequential(
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
+ nn.Conv2d(dim, dim, 3, padding=1),
+ )
+ elif mode == "upsample3d":
+ self.resample = nn.Sequential(
+ Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
+ nn.Conv2d(dim, dim, 3, padding=1),
+ # nn.Conv2d(dim, dim//2, 3, padding=1)
+ )
+ self.time_conv = CausalConv3d(
+ dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
+ elif mode == "downsample2d":
+ self.resample = nn.Sequential(
+ nn.ZeroPad2d((0, 1, 0, 1)),
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
+ elif mode == "downsample3d":
+ self.resample = nn.Sequential(
+ nn.ZeroPad2d((0, 1, 0, 1)),
+ nn.Conv2d(dim, dim, 3, stride=(2, 2)))
+ self.time_conv = CausalConv3d(
+ dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
+ else:
+ self.resample = nn.Identity()
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ b, c, t, h, w = x.size()
+ if self.mode == "upsample3d":
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ if feat_cache[idx] is None:
+ feat_cache[idx] = "Rep"
+ feat_idx[0] += 1
+ else:
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
+ feat_cache[idx] != "Rep"):
+ # cache last frame of last two chunk
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
+ feat_cache[idx] == "Rep"):
+ cache_x = torch.cat(
+ [
+ torch.zeros_like(cache_x).to(cache_x.device),
+ cache_x
+ ],
+ dim=2,
+ )
+ if feat_cache[idx] == "Rep":
+ x = self.time_conv(x)
+ else:
+ x = self.time_conv(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ x = x.reshape(b, 2, c, t, h, w)
+ x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
+ 3)
+ x = x.reshape(b, c, t * 2, h, w)
+ t = x.shape[2]
+ x = rearrange(x, "b c t h w -> (b t) c h w")
+ x = self.resample(x)
+ x = rearrange(x, "(b t) c h w -> b c t h w", t=t) # this 4 lines do spatial down / up sample
+
+ if self.mode == "downsample3d":
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ if feat_cache[idx] is None:
+ feat_cache[idx] = x.clone()
+ feat_idx[0] += 1
+ else:
+ cache_x = x[:, :, -1:, :, :].clone()
+ x = self.time_conv(
+ torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ return x
+
+ def init_weight(self, conv):
+ conv_weight = conv.weight.detach().clone()
+ nn.init.zeros_(conv_weight)
+ c1, c2, t, h, w = conv_weight.size()
+ one_matrix = torch.eye(c1, c2)
+ init_matrix = one_matrix
+ nn.init.zeros_(conv_weight)
+ conv_weight.data[:, :, 1, 0, 0] = init_matrix # * 0.5
+ conv.weight = nn.Parameter(conv_weight)
+ nn.init.zeros_(conv.bias.data)
+
+ def init_weight2(self, conv):
+ conv_weight = conv.weight.data.detach().clone()
+ nn.init.zeros_(conv_weight)
+ c1, c2, t, h, w = conv_weight.size()
+ init_matrix = torch.eye(c1 // 2, c2)
+ conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
+ conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
+ conv.weight = nn.Parameter(conv_weight)
+ nn.init.zeros_(conv.bias.data)
+
+
+class ResidualBlock(nn.Module):
+
+ def __init__(self, in_dim, out_dim, dropout=0.0):
+ super().__init__()
+ self.in_dim = in_dim
+ self.out_dim = out_dim
+
+ # layers
+ self.residual = nn.Sequential(
+ RMS_norm(in_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(in_dim, out_dim, 3, padding=1),
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ nn.Dropout(dropout),
+ CausalConv3d(out_dim, out_dim, 3, padding=1),
+ )
+ self.shortcut = (
+ CausalConv3d(in_dim, out_dim, 1)
+ if in_dim != out_dim else nn.Identity())
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ h = self.shortcut(x)
+ for layer in self.residual:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ # cache last frame of last two chunk
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+ return x + h
+
+
+class AttentionBlock(nn.Module):
+ """
+ Causal self-attention with a single head.
+ """
+
+ def __init__(self, dim):
+ super().__init__()
+ self.dim = dim
+
+ # layers
+ self.norm = RMS_norm(dim)
+ self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
+ self.proj = nn.Conv2d(dim, dim, 1)
+
+ # zero out the last layer params
+ nn.init.zeros_(self.proj.weight)
+
+ def forward(self, x):
+ identity = x
+ b, c, t, h, w = x.size()
+ x = rearrange(x, "b c t h w -> (b t) c h w")
+ x = self.norm(x)
+ # compute query, key, value
+ q, k, v = (
+ self.to_qkv(x).reshape(b * t, 1, c * 3,
+ -1).permute(0, 1, 3,
+ 2).contiguous().chunk(3, dim=-1))
+
+ # apply attention
+ x = F.scaled_dot_product_attention(
+ q,
+ k,
+ v,
+ )
+ x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
+
+ # output
+ x = self.proj(x)
+ x = rearrange(x, "(b t) c h w-> b c t h w", t=t)
+ return x + identity
+
+
+def patchify(x, patch_size):
+ if patch_size == 1:
+ return x
+ if x.dim() == 4:
+ x = rearrange(
+ x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
+ elif x.dim() == 5:
+ x = rearrange(
+ x,
+ "b c f (h q) (w r) -> b (c r q) f h w",
+ q=patch_size,
+ r=patch_size,
+ )
+ else:
+ raise ValueError(f"Invalid input shape: {x.shape}")
+
+ return x
+
+
+def unpatchify(x, patch_size):
+ if patch_size == 1:
+ return x
+
+ if x.dim() == 4:
+ x = rearrange(
+ x, "b (c r q) h w -> b c (h q) (w r)", q=patch_size, r=patch_size)
+ elif x.dim() == 5:
+ x = rearrange(
+ x,
+ "b (c r q) f h w -> b c f (h q) (w r)",
+ q=patch_size,
+ r=patch_size,
+ )
+ return x
+
+
+class AvgDown3D(nn.Module):
+
+ def __init__(
+ self,
+ in_channels,
+ out_channels,
+ factor_t,
+ factor_s=1,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+ self.factor_t = factor_t
+ self.factor_s = factor_s
+ self.factor = self.factor_t * self.factor_s * self.factor_s
+
+ assert in_channels * self.factor % out_channels == 0
+ self.group_size = in_channels * self.factor // out_channels
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
+ pad = (0, 0, 0, 0, pad_t, 0)
+ x = F.pad(x, pad)
+ B, C, T, H, W = x.shape
+ x = x.view(
+ B,
+ C,
+ T // self.factor_t,
+ self.factor_t,
+ H // self.factor_s,
+ self.factor_s,
+ W // self.factor_s,
+ self.factor_s,
+ )
+ x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
+ x = x.view(
+ B,
+ C * self.factor,
+ T // self.factor_t,
+ H // self.factor_s,
+ W // self.factor_s,
+ )
+ x = x.view(
+ B,
+ self.out_channels,
+ self.group_size,
+ T // self.factor_t,
+ H // self.factor_s,
+ W // self.factor_s,
+ )
+ x = x.mean(dim=2)
+ return x
+
+
+class DupUp3D(nn.Module):
+
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ factor_t,
+ factor_s=1,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+
+ self.factor_t = factor_t
+ self.factor_s = factor_s
+ self.factor = self.factor_t * self.factor_s * self.factor_s
+
+ assert out_channels * self.factor % in_channels == 0
+ self.repeats = out_channels * self.factor // in_channels
+
+ def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
+ x = x.repeat_interleave(self.repeats, dim=1)
+ x = x.view(
+ x.size(0),
+ self.out_channels,
+ self.factor_t,
+ self.factor_s,
+ self.factor_s,
+ x.size(2),
+ x.size(3),
+ x.size(4),
+ )
+ x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
+ x = x.view(
+ x.size(0),
+ self.out_channels,
+ x.size(2) * self.factor_t,
+ x.size(4) * self.factor_s,
+ x.size(6) * self.factor_s,
+ )
+ if first_chunk:
+ x = x[:, :, self.factor_t - 1:, :, :]
+ return x
+
+
+class Down_ResidualBlock(nn.Module):
+
+ def __init__(self,
+ in_dim,
+ out_dim,
+ dropout,
+ mult,
+ temperal_downsample=False,
+ down_flag=False):
+ super().__init__()
+
+ # Shortcut path with downsample
+ self.avg_shortcut = AvgDown3D(
+ in_dim,
+ out_dim,
+ factor_t=2 if temperal_downsample else 1,
+ factor_s=2 if down_flag else 1,
+ )
+
+ # Main path with residual blocks and downsample
+ downsamples = []
+ for _ in range(mult): # mult=2, two block
+ downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
+ in_dim = out_dim
+
+ # Add the final downsample block
+ if down_flag:
+ mode = "downsample3d" if temperal_downsample else "downsample2d"
+ downsamples.append(Resample(out_dim, mode=mode))
+
+ self.downsamples = nn.Sequential(*downsamples)
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+ x_copy = x.clone()
+ for module in self.downsamples:
+ x = module(x, feat_cache, feat_idx)
+
+ return x + self.avg_shortcut(x_copy)
+
+
+class Up_ResidualBlock(nn.Module):
+
+ def __init__(self,
+ in_dim,
+ out_dim,
+ dropout,
+ mult,
+ temperal_upsample=False,
+ up_flag=False):
+ super().__init__()
+ # Shortcut path with upsample
+ if up_flag:
+ self.avg_shortcut = DupUp3D(
+ in_dim,
+ out_dim,
+ factor_t=2 if temperal_upsample else 1,
+ factor_s=2 if up_flag else 1,
+ )
+ else:
+ self.avg_shortcut = None
+
+ # Main path with residual blocks and upsample
+ upsamples = []
+ for _ in range(mult):
+ upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
+ in_dim = out_dim
+
+ # Add the final upsample block
+ if up_flag:
+ mode = "upsample3d" if temperal_upsample else "upsample2d"
+ upsamples.append(Resample(out_dim, mode=mode))
+
+ self.upsamples = nn.Sequential(*upsamples)
+
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
+ x_main = x.clone()
+ for module in self.upsamples:
+ x_main = module(x_main, feat_cache, feat_idx)
+ if self.avg_shortcut is not None:
+ x_shortcut = self.avg_shortcut(x, first_chunk)
+ return x_main + x_shortcut
+ else:
+ return x_main
+
+
+class Encoder3d(nn.Module):
+
+ def __init__(
+ self,
+ dim=128,
+ z_dim=4,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, False],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_downsample = temperal_downsample
+
+ # dimensions
+ dims = [dim * u for u in [1] + dim_mult] # [1,2,4,4] -> [1,1,2,4,4] -> [128,128,256,512,512]
+
+ # init block
+ self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
+
+ # downsample blocks
+ downsamples = []
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
+ t_down_flag = (
+ temperal_downsample[i]
+ if i < len(temperal_downsample) else False)
+ downsamples.append(
+ Down_ResidualBlock(
+ in_dim=in_dim,
+ out_dim=out_dim,
+ dropout=dropout,
+ mult=num_res_blocks,
+ temperal_downsample=t_down_flag,
+ down_flag=i != len(dim_mult) - 1,
+ ))
+ self.downsamples = nn.Sequential(*downsamples)
+
+ # middle blocks
+ self.middle = nn.Sequential(
+ ResidualBlock(out_dim, out_dim, dropout),
+ AttentionBlock(out_dim),
+ ResidualBlock(out_dim, out_dim, dropout),
+ )
+
+ # # output blocks
+ self.head = nn.Sequential(
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(out_dim, z_dim, 3, padding=1),
+ )
+
+ def forward(self, x, feat_cache=None, feat_idx=[0]):
+
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = self.conv1(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = self.conv1(x)
+
+ ## downsamples
+ for layer in self.downsamples:
+ if feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## middle
+ for layer in self.middle:
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## head
+ for layer in self.head:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+
+ return x
+
+
+class Decoder3d(nn.Module):
+
+ def __init__(
+ self,
+ dim=128,
+ z_dim=4,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_upsample=[False, True, True],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_upsample = temperal_upsample
+
+ # dimensions
+ dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
+ # scale = 1.0 / 2**(len(dim_mult) - 2)
+ # init block
+ self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
+
+ # middle blocks
+ self.middle = nn.Sequential(
+ ResidualBlock(dims[0], dims[0], dropout),
+ AttentionBlock(dims[0]),
+ ResidualBlock(dims[0], dims[0], dropout),
+ )
+
+ # upsample blocks
+ upsamples = []
+ for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
+ t_up_flag = temperal_upsample[i] if i < len(
+ temperal_upsample) else False
+ upsamples.append(
+ Up_ResidualBlock(
+ in_dim=in_dim,
+ out_dim=out_dim,
+ dropout=dropout,
+ mult=num_res_blocks + 1,
+ temperal_upsample=t_up_flag,
+ up_flag=i != len(dim_mult) - 1,
+ ))
+ self.upsamples = nn.Sequential(*upsamples)
+
+ # output blocks
+ self.head = nn.Sequential(
+ RMS_norm(out_dim, images=False),
+ nn.SiLU(),
+ CausalConv3d(out_dim, 12, 3, padding=1),
+ )
+
+ def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
+ if feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = self.conv1(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = self.conv1(x)
+
+ for layer in self.middle:
+ if isinstance(layer, ResidualBlock) and feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx)
+ else:
+ x = layer(x)
+
+ ## upsamples
+ for layer in self.upsamples:
+ if feat_cache is not None:
+ x = layer(x, feat_cache, feat_idx, first_chunk)
+ else:
+ x = layer(x)
+
+ ## head
+ for layer in self.head:
+ if isinstance(layer, CausalConv3d) and feat_cache is not None:
+ idx = feat_idx[0]
+ cache_x = x[:, :, -CACHE_T:, :, :].clone()
+ if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
+ cache_x = torch.cat(
+ [
+ feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
+ cache_x.device),
+ cache_x,
+ ],
+ dim=2,
+ )
+ x = layer(x, feat_cache[idx])
+ feat_cache[idx] = cache_x
+ feat_idx[0] += 1
+ else:
+ x = layer(x)
+ return x
+
+
+def count_conv3d(model):
+ count = 0
+ for m in model.modules():
+ if isinstance(m, CausalConv3d):
+ count += 1
+ return count
+
+
+class WanVAE_(nn.Module):
+
+ def __init__(
+ self,
+ dim=160,
+ dec_dim=256,
+ z_dim=16,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, False],
+ dropout=0.0,
+ ):
+ super().__init__()
+ self.dim = dim
+ self.z_dim = z_dim
+ self.dim_mult = dim_mult
+ self.num_res_blocks = num_res_blocks
+ self.attn_scales = attn_scales
+ self.temperal_downsample = temperal_downsample
+ self.temperal_upsample = temperal_downsample[::-1]
+
+ # modules
+ self.encoder = Encoder3d(
+ dim,
+ z_dim * 2,
+ dim_mult,
+ num_res_blocks,
+ attn_scales,
+ self.temperal_downsample,
+ dropout,
+ )
+ self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
+ self.conv2 = CausalConv3d(z_dim, z_dim, 1)
+ self.decoder = Decoder3d(
+ dec_dim,
+ z_dim,
+ dim_mult,
+ num_res_blocks,
+ attn_scales,
+ self.temperal_upsample,
+ dropout,
+ )
+
+ def forward(self, x, scale=[0, 1]):
+ mu = self.encode(x, scale)
+ x_recon = self.decode(mu, scale)
+ return x_recon, mu
+
+ def encode(self, x, scale=None):
+ self.clear_cache()
+ x = patchify(x, patch_size=2)
+ # import pdb; pdb.set_trace()
+ # if self.training:
+ # out = self.encoder(x)
+ # else:
+ t = x.shape[2]
+ iter_ = 1 + (t - 1) // 4
+ for i in range(iter_):
+ self._enc_conv_idx = [0]
+ if i == 0:
+ out = self.encoder(
+ x[:, :, :1, :, :],
+ feat_cache=self._enc_feat_map,
+ feat_idx=self._enc_conv_idx,
+ )
+ else:
+ out_ = self.encoder(
+ x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
+ feat_cache=self._enc_feat_map,
+ feat_idx=self._enc_conv_idx,
+ )
+ out = torch.cat([out, out_], 2)
+ mu, log_var = self.conv1(out).chunk(2, dim=1)
+ if scale is not None:
+ if isinstance(scale[0], torch.Tensor):
+ mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
+ 1, self.z_dim, 1, 1, 1)
+ else:
+ mu = (mu - scale[0]) * scale[1]
+ self.clear_cache()
+ if self.other_args.encoder_out_type == 'feature_tanh':
+ mu = torch.tanh(mu)
+ elif self.other_args.encoder_out_type == 'feature':
+ pass
+ else:
+ raise ValueError(f'{self.other_args.encoder_out_type=} is not supported!')
+ return mu
+
+ def decode(self, z, scale=None):
+ self.clear_cache()
+ if scale is not None:
+ if isinstance(scale[0], torch.Tensor):
+ z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
+ 1, self.z_dim, 1, 1, 1)
+ else:
+ z = z / scale[1] + scale[0]
+ x = self.conv2(z)
+ # if self.training:
+ # out = self.decoder(x)
+ # else:
+ iter_ = z.shape[2]
+ for i in range(iter_):
+ self._conv_idx = [0]
+ if i == 0:
+ out = self.decoder(
+ x[:, :, i:i + 1, :, :],
+ feat_cache=self._feat_map,
+ feat_idx=self._conv_idx,
+ first_chunk=True,
+ )
+ else:
+ out_ = self.decoder(
+ x[:, :, i:i + 1, :, :],
+ feat_cache=self._feat_map,
+ feat_idx=self._conv_idx,
+ )
+ out = torch.cat([out, out_], 2)
+ out = unpatchify(out, patch_size=2)
+ self.clear_cache()
+ return out
+
+ def reparameterize(self, mu, log_var):
+ std = torch.exp(0.5 * log_var)
+ eps = torch.randn_like(std)
+ return eps * std + mu
+
+ def sample(self, imgs, deterministic=False):
+ import pdb; pdb.set_trace()
+ mu, log_var = self.encode(imgs)
+ if deterministic:
+ return mu
+ std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
+ return mu + std * torch.randn_like(std)
+
+ def clear_cache(self):
+ self._conv_num = count_conv3d(self.decoder)
+ self._conv_idx = [0]
+ self._feat_map = [None] * self._conv_num
+ # cache encode
+ self._enc_conv_num = count_conv3d(self.encoder)
+ self._enc_conv_idx = [0]
+ self._enc_feat_map = [None] * self._enc_conv_num
+
+
+
+class HBQ_Tokenizer(WanVAE_):
+ def __init__(self, args):
+ self.args = args
+ self.other_args = args
+ super().__init__(
+ dim=args.dim,
+ dec_dim=args.dec_dim,
+ z_dim=args.latent_channels,
+ dim_mult=args.dim_mult,
+ num_res_blocks=args.num_res_blocks,
+ temperal_downsample=args.temperal_downsample,
+ dropout=args.dropout_z
+ )
+ self.h_div_w_templates = h_div_w_templates
+ h_div_w2hw = {}
+ for (h, w) in dynamic_resolution_thw:
+ h_div_w = dynamic_resolution_thw[(h,w)]['h_div_w']
+ pn = dynamic_resolution_thw[(h,w)]['pn']
+ if h_div_w not in h_div_w2hw:
+ h_div_w2hw[h_div_w] = {}
+ h_div_w2hw[h_div_w][pn] = (h,w)
+ self.h_div_w2hw = h_div_w2hw
+
+ @staticmethod
+ def add_model_specific_args(parent_parser):
+ from videovae.utils import str2bool
+ import argparse
+
+ # dim=160,
+ # dec_dim=256,
+ # z_dim=16,
+ # dim_mult=[1, 2, 4, 4],
+ # num_res_blocks=2,
+ # attn_scales=[],
+ # temperal_downsample=[True, True, False],
+ # dropout=0.0,
+
+ parser = argparse.ArgumentParser(parents=[parent_parser], add_help=False)
+ parser.add_argument("--dim", type=int, default=160)
+ parser.add_argument("--dec_dim", type=int, default=256)
+ parser.add_argument("--latent_channels", type=int, default=16)
+ parser.add_argument("--dim_mult", type=int, nargs='+', default=[1, 2, 4, 4],)
+ parser.add_argument("--num_res_blocks", type=int, default=2)
+ parser.add_argument("--temperal_downsample", type=int, nargs='+', default=[0,1,1],)
+ parser.add_argument("--dropout_z", type=float, default=0.)
+ parser.add_argument("--encoder_out_type", type=str, default='')
+ return parser
+
+ def forward(self, x, disc_factor, image_disc=None, video_disc=None, image_perceptual_model=None, video_perceptual_model=None, is_train=True):
+ # disc_factor never used
+ device = x.device
+ is_image = x.ndim == 4
+ if is_image:
+ x = x.unsqueeze(2)
+ B, C, T, H, W = x.shape
+
+ h_div_w = H / W
+ spatial_compress_rate = int(os.environ['spatial_compress_rate'])
+ h_div_w_template = self.h_div_w_templates[np.argmin(np.abs(self.h_div_w_templates - h_div_w))]
+ for pn in self.h_div_w2hw[h_div_w_template]:
+ hh, ww = self.h_div_w2hw[h_div_w_template][pn]
+ if H * W >= hh * ww * spatial_compress_rate * spatial_compress_rate:
+ real_pn = pn
+
+ preserve_type = 'detail'
+ with torch.amp.autocast("cuda", dtype=torch.float):
+ z = self.encode(x)
+
+ loss_dict, log_dict = {}, {}
+
+ log_dict['z_mean'] = z.mean([0,2,3,4]).mean()
+ log_dict['z_std'] = z.std([0,2,3,4]).mean()
+
+ if 'hierarchical_binary_quant_round_' in self.other_args.quant_method:
+ assert self.other_args.encoder_out_type == 'feature_tanh'
+ magic_df_round = int(self.other_args.quant_method.split('hierarchical_binary_quant_round_')[1])
+ if magic_df_round >= 100: # contiguous
+ z_list = [z]
+ all_loss = torch.tensor([0.], dtype=z.dtype, device=z.device)
+ else: # discrete
+ approx_signal = 0.
+ for round_ind in range(magic_df_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ approx_signal = approx_signal + torch.where(z > approx_signal, interval, -interval)
+ approx_signal = (approx_signal - z).detach() + z
+ z_list = [approx_signal]
+ all_loss = F.mse_loss(approx_signal.detach(), z)
+
+ if is_train == False:
+ z = z_list[-1]
+ with torch.amp.autocast("cuda", dtype=torch.float):
+ x_recon = self.decode(z).to(torch.float32)
+ return x, x_recon, z
+
+ commitment_loss = torch.mean(all_loss)
+ # kl_loss = posterior.kl(reduction="mean")
+ # kl_loss = torch.mean(kl_loss) * self.args.kl_weight
+ # loss_dict['train/kl_loss'] = kl_loss
+ # log_dict['kl/mean'] = torch.mean(posterior.mean.detach())
+ # log_dict['kl/mean_square'] = torch.mean(torch.pow(posterior.mean.detach(), 2))
+ # log_dict['kl/var'] = torch.mean(posterior.var.detach())
+
+ for ind, z in enumerate(z_list):
+ pn_prefix = f'real_pn_{real_pn}/{preserve_type}/'
+ log_dict[f'train/{pn_prefix}commitment_loss_wo_w'] = commitment_loss.detach()
+ with torch.amp.autocast("cuda", dtype=torch.float):
+ x_recon = self.decode(z).to(torch.float32)
+
+ if x.ndim == 4:
+ x = x.unsqueeze(2)
+ assert x_recon.shape == x.shape, f'x_recon.shape {x_recon.shape}!= x.shape {x.shape}'
+ if self.args.recon_loss_type == 'l1':
+ recon_loss = F.l1_loss(x_recon, x) * self.args.l1_weight
+ else:
+ recon_loss = F.mse_loss(x_recon, x) * self.args.l1_weight
+ if 'train/recon_loss' not in loss_dict:
+ loss_dict[f'train/{pn_prefix}recon_loss'] = recon_loss
+ else:
+ loss_dict[f'train/{pn_prefix}recon_loss'] += recon_loss
+
+ if is_image: # handle the cases with 4 dims
+ flat_frames = x = x.squeeze(2)
+ flat_frames_recon = x_recon = x_recon.squeeze(2)
+ else:
+ flat_frames = rearrange(x, "B C T H W -> (B T) C H W")
+ flat_frames_recon = rearrange(x_recon, "B C T H W -> (B T) C H W")
+
+ # Perceptual loss
+ if is_image:
+ image_perceptual_loss = image_perceptual_model(flat_frames, flat_frames_recon).mean() * self.args.perceptual_weight
+ if "train/image_perceptual_loss" not in loss_dict:
+ loss_dict[f'train/{pn_prefix}image_perceptual_loss'] = image_perceptual_loss
+ else:
+ loss_dict[f'train/{pn_prefix}image_perceptual_loss'] += image_perceptual_loss
+ else:
+ if self.args.lpips_model == "swin3d_t":
+ video_perceptual_loss = video_perceptual_model(x, x_recon).mean() * self.args.video_perceptual_weight
+ else:
+ video_perceptual_loss = video_perceptual_model(flat_frames, flat_frames_recon).mean() * self.args.video_perceptual_weight
+ if "train/video_perceptual_loss" not in loss_dict:
+ loss_dict[f'train/{pn_prefix}video_perceptual_loss'] = video_perceptual_loss
+ else:
+ loss_dict[f'train/{pn_prefix}video_perceptual_loss'] += video_perceptual_loss
+
+ ### GAN loss
+ if self.args.image_gan_weight > 0 and (self.args.gan_image4video == "yes" or is_image):
+ logits_image_fake = image_disc(flat_frames_recon)
+ # g_image_loss = torch.mean(F.relu(1. - logits_image_fake)) * self.args.image_gan_weight
+ g_image_loss = -torch.mean(logits_image_fake) * self.args.image_gan_weight
+ if 'train/g_image_loss' not in loss_dict:
+ loss_dict[f'train/{pn_prefix}g_image_loss'] = g_image_loss
+ else:
+ loss_dict[f'train/{pn_prefix}g_image_loss'] += g_image_loss
+ if T > 1 and self.args.video_gan_weight > 0:
+ logits_video_fake = video_disc(x_recon)
+ g_video_loss = -torch.mean(logits_video_fake) * self.args.video_gan_weight
+ if 'train/g_video_loss' not in loss_dict:
+ loss_dict[f'train/{pn_prefix}g_video_loss'] = g_video_loss
+ else:
+ loss_dict[f'train/{pn_prefix}g_video_loss'] += g_video_loss
+
+ x_recon1, flat_frames1, flat_frames_recon1 = x_recon.detach(), flat_frames.detach(), flat_frames_recon.detach()
+ return (x, x_recon1, flat_frames1, flat_frames_recon1, loss_dict, log_dict)
+
+
+def _video_vae(pretrained_path=None, z_dim=16, dim=160, device="cpu", **kwargs):
+ # params
+ cfg = dict(
+ dim=dim,
+ z_dim=z_dim,
+ dim_mult=[1, 2, 4, 4],
+ num_res_blocks=2,
+ attn_scales=[],
+ temperal_downsample=[True, True, True],
+ dropout=0.0,
+ )
+ cfg.update(**kwargs)
+
+ # init model
+ with torch.device("meta"):
+ model = WanVAE_(**cfg)
+
+ # load checkpoint
+ logging.info(f"loading {pretrained_path}")
+ model.load_state_dict(
+ torch.load(pretrained_path, map_location=device), assign=True)
+
+ return model
+
+
+class Wan2_2_VAE:
+
+ def __init__(
+ self,
+ z_dim=48,
+ c_dim=160,
+ vae_pth=None,
+ dim_mult=[1, 2, 4, 4],
+ temperal_downsample=[False, True, True],
+ dtype=torch.float,
+ device="cuda",
+ ):
+
+ self.dtype = dtype
+ self.device = device
+
+ mean = torch.tensor(
+ [
+ -0.2289,
+ -0.0052,
+ -0.1323,
+ -0.2339,
+ -0.2799,
+ 0.0174,
+ 0.1838,
+ 0.1557,
+ -0.1382,
+ 0.0542,
+ 0.2813,
+ 0.0891,
+ 0.1570,
+ -0.0098,
+ 0.0375,
+ -0.1825,
+ -0.2246,
+ -0.1207,
+ -0.0698,
+ 0.5109,
+ 0.2665,
+ -0.2108,
+ -0.2158,
+ 0.2502,
+ -0.2055,
+ -0.0322,
+ 0.1109,
+ 0.1567,
+ -0.0729,
+ 0.0899,
+ -0.2799,
+ -0.1230,
+ -0.0313,
+ -0.1649,
+ 0.0117,
+ 0.0723,
+ -0.2839,
+ -0.2083,
+ -0.0520,
+ 0.3748,
+ 0.0152,
+ 0.1957,
+ 0.1433,
+ -0.2944,
+ 0.3573,
+ -0.0548,
+ -0.1681,
+ -0.0667,
+ ],
+ dtype=dtype,
+ device=device,
+ )
+ std = torch.tensor(
+ [
+ 0.4765,
+ 1.0364,
+ 0.4514,
+ 1.1677,
+ 0.5313,
+ 0.4990,
+ 0.4818,
+ 0.5013,
+ 0.8158,
+ 1.0344,
+ 0.5894,
+ 1.0901,
+ 0.6885,
+ 0.6165,
+ 0.8454,
+ 0.4978,
+ 0.5759,
+ 0.3523,
+ 0.7135,
+ 0.6804,
+ 0.5833,
+ 1.4146,
+ 0.8986,
+ 0.5659,
+ 0.7069,
+ 0.5338,
+ 0.4889,
+ 0.4917,
+ 0.4069,
+ 0.4999,
+ 0.6866,
+ 0.4093,
+ 0.5709,
+ 0.6065,
+ 0.6415,
+ 0.4944,
+ 0.5726,
+ 1.2042,
+ 0.5458,
+ 1.6887,
+ 0.3971,
+ 1.0600,
+ 0.3943,
+ 0.5537,
+ 0.5444,
+ 0.4089,
+ 0.7468,
+ 0.7744,
+ ],
+ dtype=dtype,
+ device=device,
+ )
+ self.scale = [mean, 1.0 / std]
+
+ # init model
+ self.model = (
+ _video_vae(
+ pretrained_path=vae_pth,
+ z_dim=z_dim,
+ dim=c_dim,
+ dim_mult=dim_mult,
+ temperal_downsample=temperal_downsample,
+ ).eval().requires_grad_(False).to(device))
+
+ def encode(self, videos):
+ try:
+ if not isinstance(videos, list):
+ raise TypeError("videos should be a list")
+ with amp.autocast(dtype=self.dtype):
+ return [
+ self.model.encode(u.unsqueeze(0),
+ self.scale).float().squeeze(0)
+ for u in videos
+ ]
+ except TypeError as e:
+ logging.info(e)
+ return None
+
+ def decode(self, zs):
+ try:
+ if not isinstance(zs, list):
+ raise TypeError("zs should be a list")
+ with amp.autocast(dtype=self.dtype):
+ return [
+ self.model.decode(u.unsqueeze(0),
+ self.scale).float().clamp_(-1,
+ 1).squeeze(0)
+ for u in zs
+ ]
+ except TypeError as e:
+ logging.info(e)
+ return None
diff --git a/grn/tokenizer/videovae/modules/__init__.py b/grn/tokenizer/videovae/modules/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..487806892a94a8b8b29d1e02057d1a9160dd30f1
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/__init__.py
@@ -0,0 +1,7 @@
+from .lpips import LPIPS, ResNet50LPIPS, build_lpips_model
+from .codebook import Codebook, MultiScaleCodebook
+from .normalization import Normalize, SpatialGroupNorm, RMS_norm
+from .conv import FluxConv, DCDownBlock2d, DCUpBlock2d, DCDownBlock3d, DCUpBlock3d, CogVideoXCausalConv3d, CogVideoXSafeConv3d
+from .commitments import DiagonalGaussianDistribution
+from .loss import adopt_weight
+from .misc import swish
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/attention.py b/grn/tokenizer/videovae/modules/attention.py
new file mode 100644
index 0000000000000000000000000000000000000000..a6f68d673d00e6293da527a76a97bb4bee1b6c23
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/attention.py
@@ -0,0 +1,699 @@
+import math
+from operator import truediv
+import torch
+import torch.nn.functional as F
+from torch import nn, einsum
+from beartype import beartype
+from typing import Tuple
+
+from einops import rearrange, repeat
+from einops.layers.torch import Rearrange
+from timm.models.layers import to_2tuple, trunc_normal_
+from videovae.modules.drop_path import DropPath
+
+from fairscale.nn import checkpoint_wrapper
+from torch.nn.attention import SDPBackend, sdpa_kernel
+from videovae.utils.misc import is_dtype_16
+from videovae.modules.normalization import l2norm, LayerNorm, RMSNorm
+
+
+def do_pool(x: torch.Tensor, stride: int) -> torch.Tensor:
+ # Refer to `Unroll` to see how this performs a maxpool-Nd
+ # B, N, C
+ return x.view(x.shape[0], stride, -1, x.shape[-1]).max(dim=1).values
+
+
+def exists(val):
+ return val is not None
+
+
+def default(val, d):
+ return val if exists(val) else d
+
+
+def leaky_relu(p=0.1):
+ # return nn.LeakyReLU(p)
+ return nn.Identity()
+
+
+def precompute_freqs_cis_2d(dim: int, end: int, H, W, theta: float = 10000.0, scale=1.0, use_cls=False):
+ # H = int( end**0.5 )
+ assert H * W == end
+ flat_patch_pos = torch.arange(0 if not use_cls else -1, end) # N = end
+ x_pos = flat_patch_pos % H # N
+ y_pos = flat_patch_pos // H # N
+ freqs = 1.0 / (theta ** (torch.arange(0, dim, 4)[: (dim // 4)].float() / dim)) # Hc/4
+ x_freqs = torch.outer(x_pos, freqs).float() # N Hc/4
+ y_freqs = torch.outer(y_pos, freqs).float() # N Hc/4
+ x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs)
+ y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs)
+ freqs_cis = torch.cat([x_cis.unsqueeze(dim=-1), y_cis.unsqueeze(dim=-1)], dim=-1) # N,Hc/4,2
+ freqs_cis = freqs_cis.reshape(end if not use_cls else end + 1, -1)
+ # we need to think how to implement this for multi heads.
+ # freqs_cis = torch.cat([x_cis, y_cis], dim=-1) # N, Hc/2
+ return freqs_cis
+
+
+def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
+ # x: B N H Hc/2
+ # freqs_cis: N, H*Hc/2 or N Hc/2
+ ndim = x.ndim
+ assert 0 <= 1 < ndim
+
+ if freqs_cis.shape[-1] == x.shape[-1]:
+ shape = [1 if i == 2 or i == 0 else d for i, d in enumerate(x.shape)] # 1, N, 1, Hc/2
+ else:
+ shape = [d if i != 0 else 1 for i, d in enumerate(x.shape)] # 1, N, H, Hc/2
+ # B, N, Hc/2
+ return freqs_cis.view(*shape)
+
+def apply_rotary_emb(
+ xq: torch.Tensor,
+ xk: torch.Tensor,
+ freqs_cis: torch.Tensor,
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ # xq : B N H Hc
+ xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) # B N H Hc/2
+ xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
+ freqs_cis = reshape_for_broadcast(freqs_cis, xq_)
+ xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3) # B, N, H, Hc
+ xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
+ return xq_out.type_as(xq), xk_out.type_as(xk)
+
+
+class Pooling(nn.Module):
+ def __init__(self, pool_type, dim):
+ super().__init__()
+ if pool_type == "a":
+ self.pool = nn.AvgPool2d(kernel_size=2)
+
+ elif pool_type == "m":
+ self.pool = nn.MaxPool2d(kernel_size=2)
+
+ elif pool_type == "l":
+ self.pool = nn.Linear(4 * dim, dim)
+
+ else:
+ raise NotImplementedError
+
+ self.pool_type = pool_type
+
+ def forward(self, x):
+ # B N C
+ B, N, C= x.shape
+ if self.pool_type in ["a", "m"]:
+ H, W = int(math.sqrt(N)), int(math.sqrt(N))
+ x = x.view(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ x = self.pool(x)
+ x = x.view(B, C, -1).transpose(1, 2).contiguous()
+
+ else:
+ x = x.view(B, N//4, -1)
+ x = self.pool(x)
+
+ return x
+
+
+class Up(nn.Module):
+ def __init__(self, up_type, dim):
+ super().__init__()
+ if up_type == "n":
+ self.up = nn.Upsample(scale_factor=2, mode='nearest')
+
+ elif up_type == "r":
+ self.up = nn.Sequential(
+ nn.Upsample(scale_factor=2, mode='nearest'),
+ Rearrange('b c h w -> b (h w) c'),
+ nn.Linear(dim, dim)
+ )
+
+ else:
+ raise NotImplementedError
+
+ self.up_type = up_type
+
+ def forward(self, x):
+ # B N C
+ B, N, C= x.shape
+ if self.up_type == "n":
+ H, W = int(math.sqrt(N)), int(math.sqrt(N))
+ x = x.view(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
+ x = self.up(x)
+ x = x.view(B, C, -1).transpose(1, 2).contiguous()
+
+ else:
+ #x = self.up(x) # B, N, 4c
+ #x = x.view(B, N * 4, -1)
+ H, W = int(math.sqrt(N)), int(math.sqrt(N))
+ x = x.view(B, H, W, -1).permute(0, 3, 1, 2).contiguous() # B, C, H, W
+ x = self.up(x) # B, (2H 2W), C
+
+ return x
+
+
+class GEGLU(nn.Module):
+ def forward(self, x):
+ x, gate = x.chunk(2, dim=-1)
+ return F.gelu(gate) * x
+
+
+def FeedForward(dim, mult=4, dropout=0.):
+ """ Check this paper to understand the computation: https://arxiv.org/pdf/2002.05202.pdf"""
+ inner_dim = int(mult * (2 / 3) * dim)
+ return nn.Sequential(
+ nn.LayerNorm(dim),
+ nn.Linear(dim, inner_dim * 2, bias=False),
+ GEGLU(),
+ nn.Dropout(dropout),
+ nn.Linear(inner_dim, dim, bias=False)
+ )
+
+# PEG - position generating module
+
+
+
+
+def window_partition(x, window_size):
+ """
+ Args:
+ x: (B, H, W, C)
+ window_size (int): window size
+
+ Returns:
+ windows: (num_windows*B, window_size, window_size, C)
+ """
+ B, H, W, C = x.shape
+ x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
+ windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
+ return windows
+
+
+def window_reverse(windows, window_size, H, W):
+ """
+ Args:
+ windows: (num_windows*B, window_size, window_size, C)
+ window_size (int): Window size
+ H (int): Height of image
+ W (int): Width of image
+
+ Returns:
+ x: (B, H, W, C)
+ """
+ B = int(windows.shape[0] / (H * W / window_size / window_size))
+ x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
+ return x
+
+
+class WindowAttention(nn.Module):
+ r""" Window based multi-head self attention (W-MSA) module with relative position bias.
+ It supports both of shifted and non-shifted window.
+
+ Args:
+ dim (int): Number of input channels.
+ window_size (tuple[int]): The height and width of the window.
+ num_heads (int): Number of attention heads.
+ qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
+ qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
+ attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
+ proj_drop (float, optional): Dropout ratio of output. Default: 0.0
+ """
+
+ def __init__(self, dim, window_size, num_heads, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.):
+
+ super().__init__()
+ self.dim = dim
+ if isinstance(window_size, int):
+ window_size = (window_size, window_size)
+
+ self.norm = LayerNorm(dim)
+ self.window_size = window_size # Wh, Ww
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.scale = qk_scale or head_dim ** -0.5
+
+ # define a parameter table of relative position bias
+ self.relative_position_bias_table = nn.Parameter(
+ torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH
+
+ # get pair-wise relative position index for each token inside the window
+ coords_h = torch.arange(self.window_size[0])
+ coords_w = torch.arange(self.window_size[1])
+ coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww
+ coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
+ relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
+ relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
+ relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
+ relative_coords[:, :, 1] += self.window_size[1] - 1
+ relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
+ relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
+ self.register_buffer("relative_position_index", relative_position_index)
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.attn_drop = nn.Dropout(attn_drop)
+ self.proj = nn.Linear(dim, dim)
+ self.proj_drop = nn.Dropout(proj_drop)
+
+ trunc_normal_(self.relative_position_bias_table, std=.02)
+ self.softmax = nn.Softmax(dim=-1)
+
+ def forward(self, x):
+ """
+ Args:
+ x: input features with shape of (num_windows*B, N, C)
+ mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
+ """
+ B_, N, C = x.shape
+ H, W = int(math.sqrt(N)), int(math.sqrt(N))
+ x = self.norm(x)
+
+ x = x.view(B_, H, W, -1)
+ # partition windows
+ x_windows = window_partition(x, self.window_size[0]) # nW*B, window_size, window_size, C
+ x_windows = x_windows.view(-1, self.window_size[0] * self.window_size[1], C) # nW*B, window_size*window_size, C
+
+ BW, NW = x_windows.shape[:2]
+
+ qkv = self.qkv(x_windows).reshape(BW, NW, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
+ q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
+
+ q = q * self.scale
+ attn = (q @ k.transpose(-2, -1))
+
+ relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
+ self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH
+ relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
+
+ attn = attn + relative_position_bias.unsqueeze(0)
+ attn = self.softmax(attn)
+
+ attn = self.attn_drop(attn)
+
+ x_windows = (attn @ v).transpose(1, 2).reshape(BW, NW, C)
+ x_windows = self.proj(x_windows)
+ x_windows = self.proj_drop(x_windows)
+
+ x = window_reverse(x_windows, self.window_size[0], H, W) # B H' W' C
+ x = x.view(B_, H * W, C)
+
+ return x
+
+
+
+
+class PEG(nn.Module):
+ def __init__(self, dim, causal=False):
+ super().__init__()
+ self.causal = causal
+ self.dsconv = nn.Conv3d(dim, dim, 3, groups=dim)
+
+ @beartype
+ def forward(self, x, shape: Tuple[int, int, int, int] = None):
+ needs_shape = x.ndim == 3
+ assert not (needs_shape and not exists(shape))
+
+ orig_shape = x.shape
+ if needs_shape:
+ x = x.reshape(*shape, -1)
+
+ x = rearrange(x, 'b ... d -> b d ...')
+
+ frame_padding = (2, 0) if self.causal else (1, 1)
+
+ x = F.pad(x, (1, 1, 1, 1, *frame_padding), value=0.)
+ x = self.dsconv(x)
+
+ x = rearrange(x, 'b d ... -> b ... d')
+
+ if needs_shape:
+ x = rearrange(x, 'b ... d -> b (...) d')
+
+ return x.reshape(orig_shape)
+
+# attention
+
+
+class Attention(nn.Module):
+ def __init__(
+ self,
+ dim,
+ dim_context=None,
+ dim_head=64,
+ heads=8,
+ causal=False,
+ norm_context=False,
+ dropout=0.,
+ spatial_pos="rel",
+ mlp_block=False,
+ qk_norm=None
+ ):
+ super().__init__()
+ self.heads = heads
+ self.causal = causal
+ inner_dim = dim_head * heads
+ dim_context = default(dim_context, dim)
+
+ # if spatial_pos == "rel":
+ # self.spatial_rel_pos_bias = ContinuousPositionBias(dim=dim, heads=heads) # HACK this: whether shared pos encoding is better or on the contrary
+
+ self.spatial_pos = spatial_pos
+ self.freqs_cis = None
+
+ # if causal:
+ # self.rel_pos_bias = AlibiPositionalBias(heads=heads)
+
+ self.p_dropout = dropout
+ self.attn_dropout = nn.Dropout(dropout)
+
+ self.norm = LayerNorm(dim)
+ self.context_norm = LayerNorm(
+ dim_context) if norm_context else nn.Identity()
+
+ self.qk_norm = qk_norm
+ if qk_norm == "l2norm":
+ self.q_scale = nn.Parameter(torch.ones(dim_head))
+ self.k_scale = nn.Parameter(torch.ones(dim_head))
+ elif qk_norm == "rmsnorm":
+ self.q_norm = RMSNorm(dim_head)
+ self.k_norm = RMSNorm(dim_head)
+
+ self.to_q = nn.Linear(dim, inner_dim, bias=False)
+ self.to_kv = nn.Linear(dim_context, inner_dim * 2, bias=False)
+ self.dim = inner_dim
+
+ # mlp branch
+ self.mlp_block = mlp_block
+ if mlp_block:
+ self.mlp_in = nn.Linear(dim, inner_dim)
+ self.mlp_gelu = nn.GELU()
+ self.mlp_out = nn.Linear(dim, inner_dim)
+
+ self.to_out = nn.Linear(inner_dim, dim)
+
+ def forward(
+ self,
+ x,
+ mask=None,
+ context=None,
+ is_spatial=True,
+ q_stride=1,
+ rope_cache=None,
+ upcast_attention=None
+ ):
+ batch, device, dtype = x.shape[0], x.device, x.dtype
+
+ if exists(context):
+ context = self.context_norm(context)
+
+ kv_input = default(context, x)
+
+ x = self.norm(x)
+ N = x.shape[1]
+
+ q, k, v = self.to_q(x), *self.to_kv(kv_input).chunk(2, dim=-1)
+ q, k, v = map(lambda t: rearrange(
+ t, 'b n (h d) -> b n h d', h=self.heads), (q, k, v))
+
+ if self.spatial_pos == "rope" and is_spatial and rope_cache != None:
+ q, k = apply_rotary_emb(q, k, freqs_cis=rope_cache)
+
+ q, k, v = map(lambda t: rearrange(
+ t, 'b n h d -> b h n d', h=self.heads), (q, k, v))
+
+ B, H, _, D = q.shape
+ if q_stride > 1:
+ # Refer to Unroll to see how this performs a maxpool-Nd
+ q = (
+ q.view(B, H, q_stride, -1, D)
+ .max(dim=2)
+ .values
+ )
+
+ if self.qk_norm == "l2norm":
+ q, k = map(l2norm, (q, k))
+ q = q * self.q_scale
+ k = k * self.k_scale
+ elif self.qk_norm == "rmsnorm":
+ q = self.q_norm(q)
+ k = self.k_norm(k)
+
+ if exists(mask):
+ mask = rearrange(mask, 'b j -> b 1 1 j')
+
+ if q.shape[-2] == 1 and k.shape[-2] == 1 and v.shape[-2] == 1:
+ dummy_op = torch.sum(q) * 0 + torch.sum(k) * 0 # # incorporate a dummy operation to ensure q and k are used
+ out = v + dummy_op
+ else:
+ # print(q.dtype, k.dtype, v.dtype)
+ q = q.to(torch.float32) if "q" in upcast_attention else q
+ k = k.to(torch.float32) if "k" in upcast_attention else k
+ v = v.to(torch.float32) if "v" in upcast_attention else v
+ if is_dtype_16(q) or is_dtype_16(k) or is_dtype_16(v):
+ with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=self.p_dropout, is_causal=self.causal)
+ else:
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=self.p_dropout, is_causal=self.causal)
+
+ out = rearrange(out, 'b h n d -> b n (h d)')
+
+ # mlp_block branch
+ if self.mlp_block:
+ mlp_x = self.mlp_in(x)
+ mlp_x = self.mlp_gelu(mlp_x)
+ mlp_out = self.mlp_out(mlp_x)
+ out = out + mlp_out
+
+ return self.to_out(out)
+
+
+# alibi positional bias for extrapolation
+class AlibiPositionalBias(nn.Module):
+ def __init__(self, heads):
+ super().__init__()
+ self.heads = heads
+ slopes = torch.Tensor(self._get_slopes(heads))
+ slopes = rearrange(slopes, 'h -> h 1 1')
+ self.register_buffer('slopes', slopes, persistent=False)
+ self.register_buffer('bias', None, persistent=False)
+
+ def get_bias(self, i, j, device):
+ i_arange = torch.arange(j - i, j, device=device)
+ j_arange = torch.arange(j, device=device)
+ bias = -torch.abs(rearrange(j_arange, 'j -> 1 1 j') -
+ rearrange(i_arange, 'i -> 1 i 1'))
+ return bias
+
+ @staticmethod
+ def _get_slopes(heads):
+ def get_slopes_power_of_2(n):
+ start = (2**(-2**-(math.log2(n)-3)))
+ ratio = start
+ return [start*ratio**i for i in range(n)]
+
+ if math.log2(heads).is_integer():
+ return get_slopes_power_of_2(heads)
+
+ closest_power_of_2 = 2 ** math.floor(math.log2(heads))
+ return get_slopes_power_of_2(closest_power_of_2) + get_slopes_power_of_2(2 * closest_power_of_2)[0::2][:heads-closest_power_of_2]
+
+ def forward(self, sim):
+ h, i, j, device = *sim.shape[-3:], sim.device
+
+ if exists(self.bias) and self.bias.shape[-1] >= j:
+ return self.bias[..., :i, :j]
+
+ bias = self.get_bias(i, j, device)
+ bias = bias * self.slopes
+
+ num_heads_unalibied = h - bias.shape[0]
+ bias = F.pad(bias, (0, 0, 0, 0, 0, num_heads_unalibied))
+ self.register_buffer('bias', bias, persistent=False)
+
+ return self.bias
+
+
+class ContinuousPositionBias(nn.Module):
+ """ from https://arxiv.org/abs/2111.09883 """
+
+ def __init__(
+ self,
+ *,
+ dim,
+ heads,
+ num_dims=2, # 2 for images, 3 for video
+ layers=2,
+ log_dist=True,
+ cache_rel_pos=False
+ ):
+ super().__init__()
+ self.num_dims = num_dims
+ self.log_dist = log_dist
+
+ self.net = nn.ModuleList([])
+ self.net.append(nn.Sequential(
+ nn.Linear(self.num_dims, dim), leaky_relu()))
+
+ for _ in range(layers - 1):
+ self.net.append(nn.Sequential(nn.Linear(dim, dim), leaky_relu()))
+
+ self.net.append(nn.Linear(dim, heads))
+
+ self.cache_rel_pos = cache_rel_pos
+ self.register_buffer('rel_pos', None, persistent=False)
+
+ def forward(self, *dimensions, device=torch.device('cpu')):
+
+ if not exists(self.rel_pos) or not self.cache_rel_pos:
+ positions = [torch.arange(d, device=device) for d in dimensions]
+ grid = torch.stack(torch.meshgrid(*positions, indexing='ij'))
+ grid = rearrange(grid, 'c ... -> (...) c')
+ rel_pos = rearrange(grid, 'i c -> i 1 c') - \
+ rearrange(grid, 'j c -> 1 j c')
+
+ if self.log_dist:
+ rel_pos = torch.sign(rel_pos) * torch.log(rel_pos.abs() + 1)
+
+ self.register_buffer('rel_pos', rel_pos, persistent=False)
+
+ rel_pos = self.rel_pos.float()
+
+ for layer in self.net:
+ rel_pos = layer(rel_pos)
+
+ return rearrange(rel_pos, 'i j h -> h i j')
+
+# transformer
+
+
+class Transformer(nn.Module):
+ def __init__(
+ self,
+ dim,
+ *,
+ depth,
+ block,
+ dim_context=None,
+ causal=False,
+ dim_head=64,
+ heads=8,
+ ff_mult=4,
+ peg=False,
+ peg_causal=False,
+ has_cross_attn=False,
+ attn_dropout=0.,
+ ff_dropout=0.,
+ window_size=4,
+ spatial_pos="rel",
+ mlp_block=False,
+ upcast_attention=None,
+ qk_norm=None,
+ drop_path=0.
+ ):
+ super().__init__()
+ self.dim = dim
+ self.dim_head = dim_head
+ self.heads = heads
+ self.upcast_attention = upcast_attention
+ assert len(block) == depth
+ self.layers = nn.ModuleList([])
+ dpr = [x.item() for x in torch.linspace(0, drop_path, depth)]
+ for i in range(depth):
+ if block[i] == 't':
+ self.layers.append(nn.ModuleList([
+ PEG(dim=dim, causal=peg_causal) if peg else None,
+ Attention(dim=dim, dim_head=dim_head, heads=heads,
+ causal=causal, dropout=attn_dropout, spatial_pos=spatial_pos, mlp_block=mlp_block, qk_norm=qk_norm),
+ Attention(dim=dim, dim_head=dim_head, dim_context=dim_context, heads=heads, causal=False,
+ dropout=attn_dropout, mlp_block=mlp_block, qk_norm=qk_norm) if has_cross_attn else None,
+ FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout),
+ DropPath(dpr[i]) if dpr[i] > 0. else nn.Identity()
+ ]))
+
+ # elif block[i] == 'w':
+ # self.layers.append(nn.ModuleList([
+ # None,
+ # WindowAttention(dim=dim, window_size=window_size, num_heads=heads, attn_drop=attn_dropout),
+ # None,
+ # FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)
+ # ]))
+
+ # # various pooling methods: B, N, C
+ # elif block[i] in ['a', 'm', 'l']:
+ # self.layers.append(nn.ModuleList([
+ # None,
+ # Pooling(block[i], dim),
+ # None,
+ # FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)
+ # ]))
+
+ # elif block[i] in ['n', 'r']:
+ # self.layers.append(nn.ModuleList([
+ # None,
+ # Up(block[i], dim),
+ # None,
+ # FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)
+ # ]))
+
+ else:
+ raise NotImplementedError
+
+ self.block = block
+ self.norm_out = nn.LayerNorm(dim)
+
+ @beartype
+ def forward(
+ self,
+ x,
+ video_shape: Tuple[int, int, int, int] = None,
+ context=None,
+ self_attn_mask=None,
+ cross_attn_context_mask=None,
+ q_strides=None,
+ is_spatial=True
+ ):
+ if q_strides is None:
+ q_strides = '1' * len(self.layers)
+
+ for blk, q_stride, (peg, self_attn, cross_attn, ff, drop_path) in zip(self.block, q_strides, self.layers):
+ if exists(peg):
+ with torch.amp.autocast("cuda", enabled=False):
+ x = peg(x, shape=video_shape) + x
+
+ if isinstance(self_attn, Attention):
+ H, W = video_shape[2], video_shape[3]
+ if x.shape[-2] == H * W:
+ rope_cache = precompute_freqs_cis_2d(self.dim_head, x.shape[1], H, W).to(x.device)
+ elif x.shape[-2] == 1 or is_spatial == False:
+ rope_cache = None
+ else:
+ raise NotImplementedError
+ x = drop_path(self_attn(
+ x, mask=self_attn_mask,
+ q_stride=int(q_stride), is_spatial=is_spatial,
+ rope_cache=rope_cache, upcast_attention=self.upcast_attention
+ )) + do_pool(x, int(q_stride))
+
+ elif isinstance(self_attn, WindowAttention):
+ x = drop_path(self_attn(x)) + x
+ else:
+ x = self_attn(x)
+
+ if exists(cross_attn) and exists(context):
+ x = cross_attn(x, context=context,
+ mask=cross_attn_context_mask) + x
+
+ x = ff(x) + x
+
+ # deal with downsampling:
+ if blk in ['a', 'm', 'l']:
+ video_shape = (video_shape[0], video_shape[1], video_shape[2]//2, video_shape[3]//2) # video_shape: B, T, H, W
+
+ elif blk in ['n', 'r']:
+ video_shape = (video_shape[0], video_shape[1], int(video_shape[2]*2), int(video_shape[3]*2))
+
+
+ if q_stride != '1':
+ down_ratio = int(math.sqrt(int(q_stride)))
+ video_shape = (video_shape[0], video_shape[1], video_shape[2]//down_ratio, video_shape[3]//down_ratio)
+
+ return self.norm_out(x)
diff --git a/grn/tokenizer/videovae/modules/cache/vgg.pth b/grn/tokenizer/videovae/modules/cache/vgg.pth
new file mode 100644
index 0000000000000000000000000000000000000000..47e943cfacabf7040b4af8cf4084ab91177f1b88
Binary files /dev/null and b/grn/tokenizer/videovae/modules/cache/vgg.pth differ
diff --git a/grn/tokenizer/videovae/modules/codebook.py b/grn/tokenizer/videovae/modules/codebook.py
new file mode 100644
index 0000000000000000000000000000000000000000..af7ff344a9d21839bef6989d9ffbbf39983f4dfe
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/codebook.py
@@ -0,0 +1,417 @@
+# Copyright (c) Meta Platforms, Inc. All Rights Reserved
+
+from enum import unique
+import numpy as np
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.distributed as dist
+
+from videovae.utils.misc import shift_dim
+
+class Codebook(nn.Module):
+ def __init__(self, n_codes, embedding_dim, no_random_restart=False, restart_thres=1.0, usage_sigma=0.99, fp32_quant=False):
+ super().__init__()
+ self.register_buffer('embeddings', torch.randn(n_codes, embedding_dim))
+ self.register_buffer('N', torch.zeros(n_codes))
+ self.register_buffer('z_avg', self.embeddings.data.clone())
+ self.register_buffer('codebook_usage', torch.zeros(n_codes))
+
+ self.call_cnt = 0
+ self.usage_sigma = usage_sigma
+
+ self.n_codes = n_codes
+ self.embedding_dim = embedding_dim
+ self._need_init = True
+ self.no_random_restart = no_random_restart
+ self.restart_thres = restart_thres
+
+ self.fp32_quant = fp32_quant
+
+ def _tile(self, x):
+ d, ew = x.shape
+ if d < self.n_codes:
+ n_repeats = (self.n_codes + d - 1) // d
+ std = 0.01 / np.sqrt(ew)
+ x = x.repeat(n_repeats, 1)
+ x = x + torch.randn_like(x) * std
+ return x
+
+ def _init_embeddings(self, z):
+ # z: [b, c, t, h, w]
+ self._need_init = False
+ flat_inputs = shift_dim(z, 1, -1).flatten(end_dim=-2)
+ y = self._tile(flat_inputs)
+
+ d = y.shape[0]
+ _k_rand = y[torch.randperm(y.shape[0])][:self.n_codes]
+ if dist.is_initialized():
+ dist.broadcast(_k_rand, 0)
+ self.embeddings.data.copy_(_k_rand)
+ self.z_avg.data.copy_(_k_rand)
+ self.N.data.copy_(torch.ones(self.n_codes))
+
+
+ def calculate_batch_codebook_usage_percentage(self, batch_encoding_indices):
+ # Flatten the batch of encoding indices into a single 1D tensor
+ all_indices = batch_encoding_indices.flatten()
+
+ # Obtain the total number of encoding indices in the batch to calculate percentages
+ total_indices = all_indices.numel()
+
+ # Initialize a tensor to store the percentage usage of each code
+ codebook_usage_percentage = torch.zeros(self.n_codes, device=all_indices.device)
+
+ # Count the number of occurrences of each index and get their frequency as percentages
+ unique_indices, counts = torch.unique(all_indices, return_counts=True)
+ # Calculate the percentage
+ percentages = (counts.float() / total_indices)
+
+ # Populate the corresponding percentages in the codebook_usage_percentage tensor
+ codebook_usage_percentage[unique_indices.long()] = percentages
+
+ return codebook_usage_percentage
+
+
+
+ def forward(self, z):
+ # z: [b, c, t, h, w]
+ if self._need_init and self.training:
+ self._init_embeddings(z)
+ flat_inputs = shift_dim(z, 1, -1).flatten(end_dim=-2) # [bthw, c]
+
+ distances = (flat_inputs ** 2).sum(dim=1, keepdim=True) \
+ - 2 * flat_inputs @ self.embeddings.t() \
+ + (self.embeddings.t() ** 2).sum(dim=0, keepdim=True) # [bthw, c]
+
+ encoding_indices = torch.argmin(distances, dim=1)
+ encode_onehot = F.one_hot(encoding_indices, self.n_codes).type_as(flat_inputs) # [bthw, ncode]
+ encoding_indices = encoding_indices.view(z.shape[0], *z.shape[2:]) # [b, t, h, w, ncode]
+
+ embeddings = F.embedding(encoding_indices, self.embeddings) # [b, t, h, w, c]
+ embeddings = shift_dim(embeddings, -1, 1) # [b, c, t, h, w]
+
+ commitment_loss = 0.25 * F.mse_loss(z, embeddings.detach())
+
+ # EMA codebook update
+ if self.training:
+ n_total = encode_onehot.sum(dim=0)
+ encode_sum = flat_inputs.t() @ encode_onehot
+ if dist.is_initialized():
+ dist.all_reduce(n_total)
+ dist.all_reduce(encode_sum)
+
+ self.N.data.mul_(0.99).add_(n_total, alpha=0.01)
+ self.z_avg.data.mul_(0.99).add_(encode_sum.t(), alpha=0.01)
+
+ n = self.N.sum()
+ weights = (self.N + 1e-7) / (n + self.n_codes * 1e-7) * n
+ encode_normalized = self.z_avg / weights.unsqueeze(1)
+ self.embeddings.data.copy_(encode_normalized)
+
+ y = self._tile(flat_inputs)
+ _k_rand = y[torch.randperm(y.shape[0])][:self.n_codes]
+ if dist.is_initialized():
+ dist.broadcast(_k_rand, 0)
+
+ if not self.no_random_restart:
+ usage = (self.N.view(self.n_codes, 1) >= self.restart_thres).float()
+ self.embeddings.data.mul_(usage).add_(_k_rand * (1 - usage))
+
+ embeddings_st = (embeddings - z).detach() + z
+
+ avg_probs = torch.mean(encode_onehot, dim=0)
+ perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10)))
+
+ try:
+ usage = self.calculate_batch_codebook_usage_percentage(encoding_indices)
+ except:
+ usage = torch.zeros(self.n_codes, device=encoding_indices.device)
+
+
+ # print(usage.shape, torch.zeros(self.n_codes).shape)
+
+ if self.call_cnt == 0:
+ self.codebook_usage.data = usage
+ else:
+ self.codebook_usage.data = self.usage_sigma * self.codebook_usage.data + (1 - self.usage_sigma) * usage
+
+ self.call_cnt += 1
+ # avg_distribution = self.codebook_usage.data.sum() / self.n_codes
+ avg_usage = (self.codebook_usage.data > (1/self.n_codes)).sum() / self.n_codes
+
+ return dict(embeddings=embeddings_st, encodings=encoding_indices,
+ commitment_loss=commitment_loss, perplexity=perplexity, avg_usage=avg_usage, batch_usage=usage)
+
+ def dictionary_lookup(self, encodings):
+ embeddings = F.embedding(encodings, self.embeddings)
+ return embeddings
+
+
+# Multi-scale Codebook
+from typing import List, Optional, Tuple, Sequence, Union
+
+
+class ResConvAfterUpsample(nn.Conv3d):
+ def __init__(self, embed_dim, quant_resi):
+ ks = 3 if quant_resi < 0 else 1
+ super().__init__(in_channels=embed_dim, out_channels=embed_dim, kernel_size=ks, stride=1, padding=ks//2)
+ self.resi_ratio = abs(quant_resi)
+
+ def forward(self, h_BCthw):
+ return h_BCthw.mul(1-self.resi_ratio) + super().forward(h_BCthw).mul_(self.resi_ratio)
+
+
+class SharedResConvAfterUpsample(nn.Module):
+ def __init__(self, qresi: ResConvAfterUpsample):
+ super().__init__()
+ self.qresi: ResConvAfterUpsample = qresi
+
+ def __getitem__(self, _) -> ResConvAfterUpsample:
+ return self.qresi
+
+
+class ResConvAfterUpsampleList(nn.Module):
+ def __init__(self, qresi_ls: nn.ModuleList):
+ super().__init__()
+ self.qresi_ls = qresi_ls
+ K = len(qresi_ls)
+ self.ticks = np.linspace(1/3/K, 1-1/3/K, K) if K == 4 else np.linspace(1/2/K, 1-1/2/K, K)
+
+ def __getitem__(self, at_from_0_to_1: float) -> ResConvAfterUpsample:
+ return self.qresi_ls[np.argmin(np.abs(self.ticks - at_from_0_to_1)).item()]
+
+ def extra_repr(self) -> str:
+ return f'ticks={self.ticks}'
+
+
+class ResConvAfterUpsampleModuleList(nn.ModuleList):
+ def __init__(self, qresi: List):
+ super().__init__(qresi)
+ # self.qresi = qresi
+ K = len(qresi)
+ self.ticks = np.linspace(1/3/K, 1-1/3/K, K) if K == 4 else np.linspace(1/2/K, 1-1/2/K, K)
+
+ def __getitem__(self, at_from_0_to_1: float) -> ResConvAfterUpsample:
+ return super().__getitem__(np.argmin(np.abs(self.ticks - at_from_0_to_1)).item())
+
+ def extra_repr(self) -> str:
+ return f'ticks={self.ticks}'
+
+class MultiScaleCodebook(nn.Module):
+ def __init__(self, n_codes,
+ embedding_dim, no_random_restart=False,
+ restart_thres=1.0, usage_sigma=0.99, fp32_quant=False,
+ quant_resi = -0.5, share_quant_resi = 4, default_qresi_counts = 10,
+ t_patch_nums = (1, 1, 2, 2, 2, 4, 4, 4, 4, 4),
+ v_patch_nums = (1, 2, 3, 4, 5, 6, 8, 10, 13, 16),
+ ):
+ super().__init__()
+ self.register_buffer('embeddings', torch.randn(n_codes, embedding_dim))
+ self.register_buffer('N', torch.zeros(n_codes))
+ self.register_buffer('z_avg', self.embeddings.data.clone())
+ self.register_buffer('codebook_usage', torch.zeros(n_codes))
+
+ self.call_cnt = 0
+ self.usage_sigma = usage_sigma
+
+ self.n_codes = n_codes
+ self.embedding_dim = embedding_dim
+ self._need_init = True
+ self.no_random_restart = no_random_restart
+ self.restart_thres = restart_thres
+
+ self.fp32_quant = fp32_quant
+
+ # quant resi
+
+ self.t_patch_nums = t_patch_nums
+ self.v_patch_nums = v_patch_nums
+ self.quant_resi_ratio = quant_resi
+
+ if share_quant_resi == 1: # args.qsr
+ self.quant_resi = SharedResConvAfterUpsample(ResConvAfterUpsample(embedding_dim, quant_resi) if abs(quant_resi) > 1e-6 else nn.Identity())
+ elif share_quant_resi == 0:
+ self.quant_resi = ResConvAfterUpsampleModuleList([(ResConvAfterUpsample(embedding_dim, quant_resi) if abs(quant_resi) > 1e-6 else nn.Identity()) for _ in range(default_qresi_counts or len(self.v_patch_nums))])
+ else:
+ self.quant_resi = ResConvAfterUpsampleList(nn.ModuleList([(ResConvAfterUpsample(embedding_dim, quant_resi) if abs(quant_resi) > 1e-6 else nn.Identity()) for _ in range(share_quant_resi)]))
+
+ self.z_interplote_down = 'area'
+ self.z_interplote_up = 'trilinear'
+
+
+
+ def _tile(self, x):
+ d, ew = x.shape
+ if d < self.n_codes:
+ n_repeats = (self.n_codes + d - 1) // d
+ std = 0.01 / np.sqrt(ew)
+ x = x.repeat(n_repeats, 1)
+ x = x + torch.randn_like(x) * std
+ return x
+
+ def _init_embeddings(self, z):
+ # z: [b, c, t, h, w]
+ self._need_init = False
+ flat_inputs = shift_dim(z, 1, -1).flatten(end_dim=-2)
+ y = self._tile(flat_inputs)
+
+ d = y.shape[0]
+ _k_rand = y[torch.randperm(y.shape[0])][:self.n_codes]
+ if dist.is_initialized():
+ dist.broadcast(_k_rand, 0)
+ self.embeddings.data.copy_(_k_rand)
+ self.z_avg.data.copy_(_k_rand)
+ self.N.data.copy_(torch.ones(self.n_codes))
+
+
+ def calculate_batch_codebook_usage_percentage(self, batch_encoding_indices):
+ # Flatten the batch of encoding indices into a single 1D tensor
+ all_indices = batch_encoding_indices.flatten()
+
+ # Obtain the total number of encoding indices in the batch to calculate percentages
+ total_indices = all_indices.numel()
+
+ # Initialize a tensor to store the percentage usage of each code
+ codebook_usage_percentage = torch.zeros(self.n_codes, device=all_indices.device)
+
+ # Count the number of occurrences of each index and get their frequency as percentages
+ unique_indices, counts = torch.unique(all_indices, return_counts=True)
+ # Calculate the percentage
+ percentages = (counts.float() / total_indices)
+
+ # Populate the corresponding percentages in the codebook_usage_percentage tensor
+ codebook_usage_percentage[unique_indices.long()] = percentages
+
+ return codebook_usage_percentage
+
+
+
+ def forward(self, z):
+ # z: [b, c, t, h, w]
+ if self._need_init and self.training:
+ self._init_embeddings(z)
+
+ # 永远维持THW的结构,差最近邻时候flat,然后会进行quant_res
+ B, C, T, H, W = z.shape
+
+ z_no_grad = z.detach()
+ accu_h = torch.zeros_like(z_no_grad)
+
+
+ if self.training:
+ all_flat_inputs, all_encode_onehot = [], []
+
+ commitment_loss = 0.0
+ scale_num = len(self.v_patch_nums)
+ ms_encoding_indices = []
+
+
+ with torch.cuda.amp.autocast(enabled=False):
+
+ for si, (tpn, pn) in enumerate(zip(self.t_patch_nums, self.v_patch_nums)):
+ tpn = min(tpn, T)
+
+ # latents
+ rest_z = z_no_grad - accu_h.data
+
+ if si != scale_num - 1: # z进行下采样
+ rest_z = F.interpolate(rest_z, size=(tpn, pn, pn), mode=self.z_interplote_down)
+
+ z_NC = rest_z.permute(0, 2, 3, 4, 1).reshape(-1, C)
+
+ # 这个尺度的 rest_z 与 codebook的 distances
+ d_no_grad = torch.sum(z_NC.square(), dim=1, keepdim=True) + torch.sum(self.embeddings.square(), dim=1, keepdim=False)
+ d_no_grad.addmm_(z_NC, self.embeddings.t(), alpha=-2, beta=1)
+
+ # 转成离散ids
+ encoding_indices = torch.argmin(d_no_grad, dim=1)
+ encode_onehot = F.one_hot(encoding_indices, self.n_codes).type_as(z_NC) # [bthw, ncode]
+ encoding_indices = encoding_indices.view(rest_z.shape[0], *rest_z.shape[2:]) # [b, t, h, w, ncode]
+
+ ms_encoding_indices.append(encoding_indices)
+
+ # id转回连续,用h_表述
+ h_BTHWC = F.embedding(encoding_indices, self.embeddings) # [b, t, h, w, c]
+ h_BCTHW = h_BTHWC.permute(0, 4, 1, 2, 3).contiguous() # [b, c, t, h, w]
+
+ # up & quant resi
+
+ h_BCTHW = F.interpolate(h_BCTHW, size=(T, H, W), mode=self.z_interplote_up).contiguous()
+
+ # 加一个quant resi做卷积运算
+ quant_head = si / max(1, (scale_num - 1))
+ h_BCTHW = self.quant_resi[quant_head](h_BCTHW)
+
+ # h累加
+ accu_h = accu_h + h_BCTHW
+
+ commitment_loss += 0.25 * F.mse_loss(accu_h, z.detach()) # 0.25是一个beta
+
+ if self.training:
+ all_flat_inputs.append(z_NC)
+ all_encode_onehot.append(encode_onehot)
+
+ if self.training:
+
+ encode_onehot = torch.cat(all_encode_onehot, dim=0)
+ flat_inputs = torch.cat(all_flat_inputs, dim=0)
+
+ n_total = encode_onehot.sum(dim=0)
+ encode_sum = flat_inputs.t() @ encode_onehot
+ if dist.is_initialized():
+ dist.all_reduce(n_total)
+ dist.all_reduce(encode_sum)
+
+ self.N.data.mul_(0.99).add_(n_total, alpha=0.01)
+ self.z_avg.data.mul_(0.99).add_(encode_sum.t(), alpha=0.01)
+
+ n = self.N.sum()
+ weights = (self.N + 1e-7) / (n + self.n_codes * 1e-7) * n
+ encode_normalized = self.z_avg / weights.unsqueeze(1)
+ self.embeddings.data.copy_(encode_normalized)
+
+ y = self._tile(flat_inputs)
+ _k_rand = y[torch.randperm(y.shape[0])][:self.n_codes]
+ if dist.is_initialized():
+ dist.broadcast(_k_rand, 0)
+
+ if not self.no_random_restart:
+ usage = (self.N.view(self.n_codes, 1) >= self.restart_thres).float()
+ self.embeddings.data.mul_(usage).add_(_k_rand * (1 - usage))
+
+ commitment_loss *= 1.0 / scale_num
+ embeddings_st = (accu_h - z_no_grad).detach() + z
+
+ avg_probs = torch.mean(encode_onehot, dim=0)
+ perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10)))
+
+ try:
+ usage = self.calculate_batch_codebook_usage_percentage(encoding_indices)
+ except:
+ usage = torch.zeros(self.n_codes, device=encoding_indices.device)
+
+
+ # print(usage.shape, torch.zeros(self.n_codes).shape)
+
+ if self.call_cnt == 0:
+ self.codebook_usage.data = usage
+ else:
+ self.codebook_usage.data = self.usage_sigma * self.codebook_usage.data + (1 - self.usage_sigma) * usage
+
+ self.call_cnt += 1
+ # avg_distribution = self.codebook_usage.data.sum() / self.n_codes
+ avg_usage = (self.codebook_usage.data > (1/self.n_codes)).sum() / self.n_codes
+
+ # print(f"training: {embeddings_st.size()=}, {encoding_indices.size()=}")
+ # for idx, en_idx in enumerate(ms_encoding_indices):
+ # print(f"{idx=}, {en_idx.size()=}", flush=True)
+
+ return dict(embeddings=embeddings_st, encodings=ms_encoding_indices,
+ commitment_loss=commitment_loss, perplexity=perplexity, avg_usage=avg_usage, batch_usage=usage)
+
+ def dictionary_lookup(self, encodings):
+ embeddings = F.embedding(encodings, self.embeddings)
+ return embeddings
+
diff --git a/grn/tokenizer/videovae/modules/commitments.py b/grn/tokenizer/videovae/modules/commitments.py
new file mode 100644
index 0000000000000000000000000000000000000000..4438fe41ad4ba0f72bd3c704f6f28e1965ef1c5d
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/commitments.py
@@ -0,0 +1,181 @@
+import torch
+import torch.nn as nn
+import numpy as np
+import torch.nn.functional as F
+
+
+class DiagonalGaussianDistribution(object):
+ def __init__(self, parameters, deterministic=False):
+ self.parameters = parameters
+ self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
+ self.logvar = torch.clamp(self.logvar, -30.0, 0.4) # exp(0.4)=1.5
+ self.deterministic = deterministic
+ self.std = torch.exp(0.5 * self.logvar)
+ self.var = torch.exp(self.logvar)
+ if self.deterministic:
+ self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
+
+ def sample(self):
+ # x = self.mean + self.std * torch.randn(self.mean.shape).to(device)
+ x = self.mean + self.std * torch.randn(self.mean.shape, device=self.parameters.device)
+ return x
+
+ def kl(self, other=None, reduction="sum"):
+ if reduction == "sum":
+ reduction_op = torch.sum
+ elif reduction == "mean":
+ reduction_op = torch.mean
+ if self.mean.ndim == 4:
+ dims = [1,2,3]
+ else:
+ dims = [1,2,3,4]
+ if self.deterministic:
+ return torch.Tensor([0.])
+ else:
+ if other is None:
+ return 0.5 * reduction_op(torch.pow(self.mean, 2)
+ + self.var - 1.0 - self.logvar,
+ dim=dims)
+ else:
+ return 0.5 * reduction_op(
+ torch.pow(self.mean - other.mean, 2) / other.var
+ + self.var / other.var - 1.0 - self.logvar + other.logvar,
+ dim=dims)
+
+ def nll(self, sample, dims=[1,2,3]):
+ if self.deterministic:
+ return torch.Tensor([0.])
+ logtwopi = np.log(2.0 * np.pi)
+ return 0.5 * torch.sum(
+ logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
+ dim=dims)
+
+ def mode(self):
+ return self.mean
+
+
+
+def normal_kl(mean1, logvar1, mean2, logvar2):
+ """
+ source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
+ Compute the KL divergence between two gaussians.
+ Shapes are automatically broadcasted, so batches can be compared to
+ scalars, among other use cases.
+ """
+ tensor = None
+ for obj in (mean1, logvar1, mean2, logvar2):
+ if isinstance(obj, torch.Tensor):
+ tensor = obj
+ break
+ assert tensor is not None, "at least one argument must be a Tensor"
+
+ # Force variances to be Tensors. Broadcasting helps convert scalars to
+ # Tensors, but it does not work for torch.exp().
+ logvar1, logvar2 = [
+ x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
+ for x in (logvar1, logvar2)
+ ]
+
+ return 0.5 * (
+ -1.0
+ + logvar2
+ - logvar1
+ + torch.exp(logvar1 - logvar2)
+ + ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
+ )
+
+class VectorQuantizer(nn.Module):
+ def __init__(self, n_e, e_dim, beta, entropy_loss_ratio, l2_norm, show_usage):
+ super().__init__()
+ self.n_e = n_e
+ self.e_dim = e_dim
+ self.beta = beta
+ self.entropy_loss_ratio = entropy_loss_ratio
+ self.l2_norm = l2_norm
+ self.show_usage = show_usage
+
+ self.embedding = nn.Embedding(self.n_e, self.e_dim)
+ self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
+ if self.l2_norm:
+ self.embedding.weight.data = F.normalize(self.embedding.weight.data, p=2, dim=-1)
+ if self.show_usage:
+ self.register_buffer("codebook_used", nn.Parameter(torch.zeros(65536)))
+
+
+ def forward(self, z):
+ # reshape z -> (batch, height, width, channel) and flatten
+ z = torch.einsum('b c h w -> b h w c', z).contiguous()
+ z_flattened = z.view(-1, self.e_dim)
+ # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z
+
+ if self.l2_norm:
+ z = F.normalize(z, p=2, dim=-1)
+ z_flattened = F.normalize(z_flattened, p=2, dim=-1)
+ embedding = F.normalize(self.embedding.weight, p=2, dim=-1)
+ else:
+ embedding = self.embedding.weight
+
+ d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \
+ torch.sum(embedding**2, dim=1) - 2 * \
+ torch.einsum('bd,dn->bn', z_flattened, torch.einsum('n d -> d n', embedding))
+
+ min_encoding_indices = torch.argmin(d, dim=1)
+ z_q = embedding[min_encoding_indices].view(z.shape)
+ perplexity = None
+ min_encodings = None
+ vq_loss = None
+ commit_loss = None
+ entropy_loss = None
+ codebook_usage = 0
+
+ if self.show_usage and self.training:
+ cur_len = min_encoding_indices.shape[0]
+ self.codebook_used[:-cur_len] = self.codebook_used[cur_len:].clone()
+ self.codebook_used[-cur_len:] = min_encoding_indices
+ codebook_usage = len(torch.unique(self.codebook_used)) / self.n_e
+
+ # compute loss for embedding
+ if self.training:
+ vq_loss = torch.mean((z_q - z.detach()) ** 2)
+ commit_loss = self.beta * torch.mean((z_q.detach() - z) ** 2)
+ entropy_loss = self.entropy_loss_ratio * compute_entropy_loss(-d)
+
+ # preserve gradients
+ z_q = z + (z_q - z).detach()
+
+ # reshape back to match original input shape
+ z_q = torch.einsum('b h w c -> b c h w', z_q)
+
+ return z_q, (vq_loss, commit_loss, entropy_loss, codebook_usage), (perplexity, min_encodings, min_encoding_indices)
+
+ def get_codebook_entry(self, indices, shape=None, channel_first=True):
+ # shape = (batch, channel, height, width) if channel_first else (batch, height, width, channel)
+ if self.l2_norm:
+ embedding = F.normalize(self.embedding.weight, p=2, dim=-1)
+ else:
+ embedding = self.embedding.weight
+ z_q = embedding[indices] # (b*h*w, c)
+
+ if shape is not None:
+ if channel_first:
+ z_q = z_q.reshape(shape[0], shape[2], shape[3], shape[1])
+ # reshape back to match original input shape
+ z_q = z_q.permute(0, 3, 1, 2).contiguous()
+ else:
+ z_q = z_q.view(shape)
+ return z_q
+
+def compute_entropy_loss(affinity, loss_type="softmax", temperature=0.01):
+ flat_affinity = affinity.reshape(-1, affinity.shape[-1])
+ flat_affinity /= temperature
+ probs = F.softmax(flat_affinity, dim=-1)
+ log_probs = F.log_softmax(flat_affinity + 1e-5, dim=-1)
+ if loss_type == "softmax":
+ target_probs = probs
+ else:
+ raise ValueError("Entropy loss {} not supported".format(loss_type))
+ avg_probs = torch.mean(target_probs, dim=0)
+ avg_entropy = - torch.sum(avg_probs * torch.log(avg_probs + 1e-5))
+ sample_entropy = - torch.mean(torch.sum(target_probs * log_probs, dim=-1))
+ loss = sample_entropy - avg_entropy
+ return loss
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/conv.py b/grn/tokenizer/videovae/modules/conv.py
new file mode 100644
index 0000000000000000000000000000000000000000..4217ebaa13fce7b4b753b80dfd513fa79b753d04
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/conv.py
@@ -0,0 +1,501 @@
+from typing import Dict, Optional, Tuple, Union
+import torch
+import torch.nn as nn
+from einops import rearrange, repeat
+import torch.nn.functional as F
+from .misc import swish
+from videovae.modules.normalization import get_norm
+from videovae.utils.context_parallel import ContextParallelUtils as cp
+from videovae.utils.context_parallel import dist_conv_cache_send, dist_conv_cache_recv
+
+
+class DCDownBlock3d(nn.Module):
+ def __init__(self,
+ in_channels: int,
+ out_channels: int,
+ shortcut: bool = True,
+ group_norm=False,
+ compress_time=False,
+ norm_type=None,
+ pad_mode="constant",
+ ) -> None:
+ super().__init__()
+ self.shortcut = shortcut
+ self.compress_time = compress_time
+ if group_norm:
+ norm_layer = get_norm(norm_type)
+ self.norm = norm_layer(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True)
+ self.nonlinearity = swish
+ else:
+ self.norm = nn.Identity()
+ self.nonlinearity = nn.Identity()
+ self.spatial_factor = 2
+ self.temporal_factor = int(compress_time) if compress_time else 1
+ out_ratio = self.spatial_factor**2
+ assert out_channels % out_ratio == 0
+ out_channels = out_channels // out_ratio
+
+ # self.conv = nn.Conv3d(
+ # in_channels,
+ # out_channels,
+ # kernel_size=3,
+ # stride=(1, 1, 1),
+ # padding=0,
+ # )
+ self.conv = CogVideoXCausalConv3d(in_channels, out_channels, kernel_size=3, pad_mode=pad_mode)
+
+ def forward(self, hidden_states: torch.Tensor, conv_cache: Optional[Dict[str, torch.Tensor]] = None, temporal_compress = True) -> torch.Tensor:
+ new_conv_cache = {}
+ conv_cache = conv_cache or {}
+
+ x = hidden_states
+ x = self.nonlinearity(self.norm(x))
+ assert x.ndim == 5, f"x.ndim must be (B C T H W)"
+
+ ### use nn.Conv3d
+ # x = F.pad(x, (1, 1, 1, 1, 2, 0)) # causal pad (left, right, top, bottom, front, back)
+ # x[:, :, :2, 1:-1, 1:-1] = x[:, :, 2:3, 1:-1, 1:-1].clone() # broadcast the first value
+ # x = self.conv(x)
+
+ ### use CogVideoXCausalConv3d
+ x, new_conv_cache["conv"] = self.conv(x, conv_cache=conv_cache.get("conv"))
+
+ if x.shape[2] > 1:
+ if x.shape[2] % 2 == 1:
+ x_first, x_rest = x[:, :, 0, ...], x[:, :, 1:, ...]
+ y_first, y_rest = hidden_states[:, :, 0, ...], hidden_states[:, :, 1:, ...]
+ else:
+ x_first, x_rest = None, x
+ y_first, y_rest = None, hidden_states
+ elif x.shape[2] == 1:
+ x_first, x_rest = x[:, :, 0, ...], None
+ y_first, y_rest = hidden_states[:, :, 0, ...], None
+ else:
+ raise NotImplementedError
+ if x_first is not None:
+ x_first = rearrange(x_first, "b c (h ph) (w pw) -> b (ph pw c) h w", ph=self.spatial_factor, pw=self.spatial_factor)
+ y_first = rearrange(y_first, "b c (h ph) (w pw) -> b (ph pw c) h w", ph=self.spatial_factor, pw=self.spatial_factor)
+ if x_rest is not None:
+ if temporal_compress:
+ x_rest = rearrange(x_rest, "b c (t pt) (h ph) (w pw) -> b (ph pw c) t pt h w", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+ x_rest = x_rest.mean(dim=3)
+ y_rest = rearrange(y_rest, "b c (t pt) (h ph) (w pw) -> b (ph pw c) t pt h w", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+ y_rest = y_rest.mean(dim=3)
+ else:
+ x_rest = rearrange(x_rest, "b c (t pt) (h ph) (w pw) -> b (ph pw c) (t pt) h w", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+ y_rest = rearrange(y_rest, "b c (t pt) (h ph) (w pw) -> b (ph pw c) (t pt) h w", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+ if x_first is not None and x_rest is not None:
+ x = torch.cat([x_first[:,:, None,...], x_rest], dim=2)
+ y = torch.cat([y_first[:,:, None,...], y_rest], dim=2)
+ else:
+ x = x_first[:,:, None,...] if x_first is not None else x_rest
+ y = y_first[:,:, None,...] if y_first is not None else y_rest
+ if self.shortcut:
+ y = rearrange(y, "b (g c) t h w -> b g c t h w", c=x.shape[1]).mean(dim=1)
+ hidden_states = x + y
+ else:
+ hidden_states = x
+ return hidden_states, new_conv_cache
+
+class DCUpBlock3d(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ shortcut: bool = True,
+ interpolation_mode: str = "nearest",
+ group_norm=False,
+ compress_time=False,
+ norm_type=None,
+ pad_mode="constant",
+ ) -> None:
+ super().__init__()
+
+ self.compress_time = compress_time
+ if group_norm:
+ norm_layer = get_norm(norm_type)
+ self.norm = norm_layer(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True)
+ self.nonlinearity = swish
+ else:
+ self.norm = nn.Identity()
+ self.nonlinearity = nn.Identity()
+ self.interpolation_mode = interpolation_mode
+ self.shortcut = shortcut
+ self.spatial_factor = 2
+ self.temporal_factor = int(compress_time) if compress_time else 1
+ out_channels = out_channels * self.spatial_factor**2 * self.temporal_factor
+ # self.conv = nn.Conv3d(in_channels, out_channels, 3, (1, 1, 1), 0)
+ self.conv = CogVideoXCausalConv3d(in_channels, out_channels, kernel_size=3, pad_mode=pad_mode)
+ assert out_channels % in_channels == 0
+ self.repeats = out_channels // in_channels
+
+ def forward(self, hidden_states: torch.Tensor, conv_cache: Optional[Dict[str, torch.Tensor]] = None, split_first=False) -> torch.Tensor:
+ new_conv_cache = {}
+ conv_cache = conv_cache or {}
+
+ x = hidden_states
+ x = self.nonlinearity(self.norm(x))
+
+ compress_first = False
+ if x.shape[2] % 2 == 1 or split_first:
+ compress_first = True
+
+ ### use nn.Conv3d
+ # x = F.pad(x, (1, 1, 1, 1, 2, 0)) # causal pad (left, right, top, bottom, front, back)
+ # x[:, :, :2, 1:-1, 1:-1] = x[:, :, 2:3, 1:-1, 1:-1].clone() # broadcast the first value
+ # x = self.conv(x)
+
+ ### use CogVideoXCausalConv3d
+ x, new_conv_cache["conv"] = self.conv(x, conv_cache=conv_cache.get("conv"))
+
+ x = rearrange(x, "b (pt ph pw c) t h w -> b c (t pt) (h ph) (w pw)", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+ y = repeat(hidden_states, "b c t h w -> b (r c) t h w", r=self.repeats)
+ y = rearrange(y, "b (pt ph pw c) t h w -> b c (t pt) (h ph) (w pw)", pt=self.temporal_factor, ph=self.spatial_factor, pw=self.spatial_factor)
+
+ # convert pt+pt*n -> 1+pt*n
+ if self.temporal_factor > 1 and compress_first:
+ if x.shape[2] > 1:
+ x_first, x_rest = x[:, :, :self.temporal_factor, ...], x[:, :, self.temporal_factor:, ...]
+ y_first, y_rest = y[:, :, :self.temporal_factor, ...], y[:, :, self.temporal_factor:, ...]
+ elif x.shape[2] == 1:
+ assert x.shape[2] == y.shape[2] == self.temporal_factor
+ x_first, x_rest = x, None
+ y_first, y_rest = y, None
+ else:
+ raise NotImplementedError
+ x = torch.cat([x_first.mean(dim=2, keepdim=True), x_rest], dim=2)
+ y = torch.cat([y_first.mean(dim=2, keepdim=True), y_rest], dim=2)
+ if self.shortcut:
+ hidden_states = x + y
+ else:
+ hidden_states = x
+ return hidden_states, new_conv_cache
+
+class DCDownBlock2d(nn.Module):
+ def __init__(self,
+ in_channels: int,
+ out_channels: int,
+ downsample: bool = False,
+ shortcut: bool = True,
+ group_norm=False,
+ pad_mode="contant",
+ ) -> None:
+ super().__init__()
+ if group_norm:
+ self.norm = nn.GroupNorm(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True)
+ self.nonlinearity = swish
+ else:
+ self.norm = nn.Identity()
+ self.nonlinearity = nn.Identity()
+ self.downsample = downsample
+ self.factor = 2
+ self.stride = 1 if downsample else 2
+ self.group_size = in_channels * self.factor**2 // out_channels
+ self.shortcut = shortcut
+
+ out_ratio = self.factor**2
+ if downsample:
+ assert out_channels % out_ratio == 0
+ out_channels = out_channels // out_ratio
+
+ self.conv = nn.Conv2d(
+ in_channels,
+ out_channels,
+ kernel_size=3,
+ stride=self.stride,
+ padding=1,
+ )
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ x = self.nonlinearity(self.norm(hidden_states))
+ x = self.conv(x)
+ if self.downsample:
+ x = F.pixel_unshuffle(x, self.factor)
+
+ if self.shortcut:
+ y = F.pixel_unshuffle(hidden_states, self.factor)
+ y = y.unflatten(1, (-1, self.group_size))
+ y = y.mean(dim=2)
+ hidden_states = x + y
+ else:
+ hidden_states = x
+
+ return hidden_states
+
+class DCUpBlock2d(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ interpolate: bool = False,
+ shortcut: bool = True,
+ interpolation_mode: str = "nearest",
+ group_norm=False,
+ pad_mode="constant",
+ ) -> None:
+ super().__init__()
+
+ if group_norm:
+ self.norm = nn.GroupNorm(num_channels=in_channels, num_groups=32, eps=1e-6, affine=True)
+ self.nonlinearity = swish
+ else:
+ self.norm = nn.Identity()
+ self.nonlinearity = nn.Identity()
+ self.interpolate = interpolate
+ self.interpolation_mode = interpolation_mode
+ self.shortcut = shortcut
+ self.factor = 2
+ self.repeats = out_channels * self.factor**2 // in_channels
+
+ out_ratio = self.factor**2
+
+ if not interpolate:
+ out_channels = out_channels * out_ratio
+
+ self.conv = nn.Conv2d(in_channels, out_channels, 3, 1, 1, padding_mode=pad_mode)
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ x = self.nonlinearity(self.norm(hidden_states))
+ if self.interpolate:
+ x = F.interpolate(x, scale_factor=self.factor, mode=self.interpolation_mode)
+ x = self.conv(x)
+ else:
+ x = self.conv(x)
+ x = F.pixel_shuffle(x, self.factor)
+
+ if self.shortcut:
+ y = hidden_states.repeat_interleave(self.repeats, dim=1)
+ y = F.pixel_shuffle(y, self.factor)
+ hidden_states = x + y
+ else:
+ hidden_states = x
+
+ return hidden_states
+
+class CogVideoXSafeConv3d(nn.Conv3d):
+ r"""
+ A 3D convolution layer that splits the input tensor into smaller parts to avoid OOM in CogVideoX Model.
+ """
+
+ def forward(self, input: torch.Tensor) -> torch.Tensor:
+ memory_count = (
+ (input.shape[1] * input.shape[2] * input.shape[3] * input.shape[4]) * 2 / 1024**3
+ )
+
+ # Set to 2GB, suitable for CuDNN
+ if memory_count > 2:
+ kernel_size = self.kernel_size[0]
+ part_num = int(memory_count / 2) + 1
+ input_chunks = torch.chunk(input, part_num, dim=2)
+
+ if kernel_size > 1:
+ input_chunks = [input_chunks[0]] + [
+ torch.cat((input_chunks[i - 1][:, :, -kernel_size + 1 :], input_chunks[i]), dim=2)
+ for i in range(1, len(input_chunks))
+ ]
+
+ output_chunks = []
+ for input_chunk in input_chunks:
+ output_chunks.append(super().forward(input_chunk))
+ output = torch.cat(output_chunks, dim=2)
+ return output
+ else:
+ return super().forward(input)
+
+
+class CogVideoXCausalConv3d(nn.Module):
+ r"""A 3D causal convolution layer that pads the input tensor to ensure causality in CogVideoX Model.
+
+ Args:
+ in_channels (`int`): Number of channels in the input tensor.
+ out_channels (`int`): Number of output channels produced by the convolution.
+ kernel_size (`int` or `Tuple[int, int, int]`): Kernel size of the convolutional kernel.
+ stride (`int`, defaults to `1`): Stride of the convolution.
+ dilation (`int`, defaults to `1`): Dilation rate of the convolution.
+ pad_mode (`str`, defaults to `"constant"`): Padding mode.
+ """
+
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ kernel_size: Union[int, Tuple[int, int, int]],
+ stride: int = 1,
+ dilation: int = 1,
+ pad_mode: str = "constant",
+ ):
+ super().__init__()
+
+ if isinstance(kernel_size, int):
+ kernel_size = (kernel_size,) * 3
+
+ time_kernel_size, height_kernel_size, width_kernel_size = kernel_size
+
+ self.pad_mode = pad_mode
+ time_pad = dilation * (time_kernel_size - 1) + (1 - stride)
+ height_pad = height_kernel_size // 2
+ width_pad = width_kernel_size // 2
+
+ self.height_pad = height_pad
+ self.width_pad = width_pad
+ self.time_pad = time_pad
+ self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0)
+
+ self.temporal_dim = 2
+ self.time_kernel_size = time_kernel_size
+
+ stride = (stride, 1, 1)
+ dilation = (dilation, 1, 1)
+ self.conv = CogVideoXSafeConv3d(
+ in_channels=in_channels,
+ out_channels=out_channels,
+ kernel_size=kernel_size,
+ stride=stride,
+ dilation=dilation,
+ )
+
+ def fake_context_parallel_forward(
+ self, inputs: torch.Tensor, conv_cache: Optional[torch.Tensor] = None
+ ) -> torch.Tensor:
+ kernel_size = self.time_kernel_size
+
+ if cp.cp_on():
+ conv_cache = dist_conv_cache_recv()
+
+ if kernel_size > 1:
+ cached_inputs = [conv_cache.to(inputs.device)] if conv_cache is not None else [inputs[:, :, :1]] * (kernel_size - 1)
+ inputs = torch.cat(cached_inputs + [inputs], dim=2)
+ return inputs
+
+ def forward(self, inputs: torch.Tensor, conv_cache: Optional[torch.Tensor] = None) -> torch.Tensor:
+ inputs = self.fake_context_parallel_forward(inputs, conv_cache)
+
+ if cp.cp_on():
+ dist_conv_cache_send(inputs[:, :, -self.time_kernel_size + 1 :])
+ else:
+ conv_cache = inputs[:, :, -self.time_kernel_size + 1 :].clone()
+
+
+ padding_2d = (self.width_pad, self.width_pad, self.height_pad, self.height_pad)
+ if self.pad_mode == "constant":
+ inputs = F.pad(inputs, padding_2d, mode="constant", value=0)
+ else:
+ _shape = inputs.shape
+ inputs = F.pad(inputs.view(-1, *inputs.shape[-2:]), padding_2d, mode="replicate")
+ inputs = inputs.view(*_shape[:-2], *inputs.shape[-2:])
+
+ output = self.conv(inputs)
+ return output, conv_cache
+
+class FluxConv(nn.Module):
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, cnn_type="2d", cnn_slice_seq_len=17, causal_offset=0, temporal_down=False):
+ super().__init__()
+ self.cnn_type = cnn_type
+ self.slice_seq_len = cnn_slice_seq_len
+
+ if cnn_type == "2d":
+ self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride=stride, padding=padding)
+ if cnn_type == "3d":
+ if temporal_down == False:
+ stride = (1, stride, stride)
+ else:
+ stride = (stride, stride, stride)
+ self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride=stride, padding=0)
+ if isinstance(kernel_size, int):
+ kernel_size = (kernel_size, kernel_size, kernel_size)
+ self.padding = (
+ kernel_size[0] - 1 + causal_offset, # Temporal causal padding
+ padding, # Height padding
+ padding # Width padding
+ )
+ self.causal_offset = causal_offset
+ self.stride = stride
+ self.kernel_size = kernel_size
+
+ def forward(self, x):
+ if self.cnn_type == "2d":
+ if type(x) == list:
+ for i in range(len(x)):
+ x[i] = self.forward(x[i])
+ return x
+ if x.ndim == 5:
+ B, C, T, H, W = x.shape
+ x = rearrange(x, "B C T H W -> (B T) C H W")
+ x = self.conv(x)
+ x = rearrange(x, "(B T) C H W -> B C T H W", T=T)
+ return x
+ else:
+ return self.conv(x)
+ if self.cnn_type == "3d":
+ if x.ndim == 5:
+ assert self.stride[0] == 1 or self.stride[0] == 2, f"only temporal stride = 1 or 2 are supported"
+ if self.stride[0] == 1:
+ for i in reversed(range(0, x.shape[2], self.slice_seq_len+self.stride[0]-1)):
+ st = i
+ en = min(i+self.slice_seq_len, x.shape[2])
+ _x = x[:,:,st:en,:,:]
+ if i == 0:
+ _x = F.pad(_x, (self.padding[2], self.padding[2], # Width
+ self.padding[1], self.padding[1], # Height
+ self.padding[0], 0)) # Temporal
+ _x[:,:,:self.padding[0],
+ self.padding[1]:_x.shape[-2]-self.padding[1],
+ self.padding[2]:_x.shape[-1]-self.padding[2]] = x[:,:,0:1,:,:].clone() # broadcast the first value
+ else:
+ padding_0 = self.kernel_size[0] - 1
+ _x = F.pad(_x, (self.padding[2], self.padding[2], # Width
+ self.padding[1], self.padding[1], # Height
+ padding_0, 0)) # Temporal
+ _x[:,:,:padding_0,
+ self.padding[1]:_x.shape[-2]-self.padding[1],
+ self.padding[2]:_x.shape[-1]-self.padding[2]] = x[:,:,i-padding_0:i,:,:].clone()
+ try:
+ _x = self.conv(_x)
+ except:
+ xs = [_x[:,:,:,:,i-1:i+2] for i in range(1,_x.shape[-1]-1)]
+ for i in range(len(xs)):
+ xs[i] = self.conv(xs[i])
+ _x = torch.cat(xs, dim=-1)
+ if i == 0:
+ x[:,:,st-self.causal_offset:en,:,:] = _x
+ x = x[:,:,1:,:,:]
+ else:
+ x[:,:,st:en,:,:] = _x
+ else:
+ xs = []
+ for i in range(0, x.shape[2], self.slice_seq_len+self.stride[0]-1):
+ st = i
+ en = min(i+self.slice_seq_len, x.shape[2])
+ _x = x[:,:,st:en,:,:]
+ if i == 0:
+ _x = F.pad(_x, (self.padding[2], self.padding[2], # Width
+ self.padding[1], self.padding[1], # Height
+ self.padding[0], 0)) # Temporal
+ _x[:,:,:self.padding[0],
+ self.padding[1]:_x.shape[-2]-self.padding[1],
+ self.padding[2]:_x.shape[-1]-self.padding[2]] = x[:,:,0:1,:,:].clone() # broadcast the first value
+ else:
+ padding_0 = self.kernel_size[0] - 1
+ _x = F.pad(_x, (self.padding[2], self.padding[2], # Width
+ self.padding[1], self.padding[1], # Height
+ padding_0, 0)) # Temporal
+ _x[:,:,:padding_0,
+ self.padding[1]:_x.shape[-2]-self.padding[1],
+ self.padding[2]:_x.shape[-1]-self.padding[2]] = x[:,:,i-padding_0:i,:,:].clone()
+ _x = self.conv(_x)
+ xs.append(_x)
+ try:
+ x = torch.cat(xs, dim=2)
+ except:
+ device = x.device
+ del x
+ xs = [_x.cpu().pin_memory() for _x in xs]
+ torch.cuda.empty_cache()
+ x = torch.cat([_x for _x in xs], dim=2).to(device=device)
+ else:
+ x = F.pad(x, (self.padding[2], self.padding[2], # Width
+ self.padding[1], self.padding[1])) # Height
+ weight = torch.sum(self.conv.weight, dim=2)
+ bias = self.conv.bias
+ x = F.conv2d(x, weight=weight, bias=bias,stride=self.conv.stride[1:])
+ return x
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/diffaug.py b/grn/tokenizer/videovae/modules/diffaug.py
new file mode 100644
index 0000000000000000000000000000000000000000..9310d12c601767f66291190b9cf5d5fa6305d7be
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/diffaug.py
@@ -0,0 +1,121 @@
+# BSD 2-Clause "Simplified" License
+# Copyright (c) 2020, Shengyu Zhao, Zhijian Liu, Ji Lin, Jun-Yan Zhu, and Song Han
+# All rights reserved.
+#
+# Redistribution and use in source and binary forms, with or without
+# modification, are permitted provided that the following conditions are met:
+#
+# * Redistributions of source code must retain the above copyright notice, this
+# list of conditions and the following disclaimer.
+#
+# * Redistributions in binary form must reproduce the above copyright notice,
+# this list of conditions and the following disclaimer in the documentation
+# and/or other materials provided with the distribution.
+#
+# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
+# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
+# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
+# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
+# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
+# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
+# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+#
+# Code from https://github.com/mit-han-lab/data-efficient-gans
+
+"""Training GANs with DiffAugment."""
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+
+def DiffAugment(x: torch.Tensor, policy: str = '', channels_first: bool = True) -> torch.Tensor:
+ if policy:
+ if not channels_first:
+ x = x.permute(0, 3, 1, 2)
+ for p in policy.split(','):
+ for f in AUGMENT_FNS[p]:
+ x = f(x)
+ if not channels_first:
+ x = x.permute(0, 2, 3, 1)
+ x = x.contiguous()
+ return x
+
+
+def rand_brightness(x: torch.Tensor) -> torch.Tensor:
+ x = x + (torch.rand(x.size(0), 1, 1, 1, dtype=x.dtype, device=x.device) - 0.5)
+ return x
+
+
+def rand_saturation(x: torch.Tensor) -> torch.Tensor:
+ x_mean = x.mean(dim=1, keepdim=True)
+ x = (x - x_mean) * (torch.rand(x.size(0), 1, 1, 1, dtype=x.dtype, device=x.device) * 2) + x_mean
+ return x
+
+
+def rand_contrast(x: torch.Tensor) -> torch.Tensor:
+ x_mean = x.mean(dim=[1, 2, 3], keepdim=True)
+ x = (x - x_mean) * (torch.rand(x.size(0), 1, 1, 1, dtype=x.dtype, device=x.device) + 0.5) + x_mean
+ return x
+
+
+def rand_translation(x: torch.Tensor, ratio: float = 0.125) -> torch.Tensor:
+ shift_x, shift_y = int(x.size(2) * ratio + 0.5), int(x.size(3) * ratio + 0.5)
+ translation_x = torch.randint(-shift_x, shift_x + 1, size=[x.size(0), 1, 1], device=x.device)
+ translation_y = torch.randint(-shift_y, shift_y + 1, size=[x.size(0), 1, 1], device=x.device)
+ grid_batch, grid_x, grid_y = torch.meshgrid(
+ torch.arange(x.size(0), dtype=torch.long, device=x.device),
+ torch.arange(x.size(2), dtype=torch.long, device=x.device),
+ torch.arange(x.size(3), dtype=torch.long, device=x.device),
+ )
+ grid_x = torch.clamp(grid_x + translation_x + 1, 0, x.size(2) + 1)
+ grid_y = torch.clamp(grid_y + translation_y + 1, 0, x.size(3) + 1)
+ x_pad = F.pad(x, [1, 1, 1, 1, 0, 0, 0, 0])
+ x = x_pad.permute(0, 2, 3, 1).contiguous()[grid_batch, grid_x, grid_y].permute(0, 3, 1, 2)
+ return x
+
+
+def rand_cutout(x: torch.Tensor, ratio: float = 0.2) -> torch.Tensor:
+ cutout_size = int(x.size(2) * ratio + 0.5), int(x.size(3) * ratio + 0.5)
+ offset_x = torch.randint(0, x.size(2) + (1 - cutout_size[0] % 2), size=[x.size(0), 1, 1], device=x.device)
+ offset_y = torch.randint(0, x.size(3) + (1 - cutout_size[1] % 2), size=[x.size(0), 1, 1], device=x.device)
+ grid_batch, grid_x, grid_y = torch.meshgrid(
+ torch.arange(x.size(0), dtype=torch.long, device=x.device),
+ torch.arange(cutout_size[0], dtype=torch.long, device=x.device),
+ torch.arange(cutout_size[1], dtype=torch.long, device=x.device),
+ )
+ grid_x = torch.clamp(grid_x + offset_x - cutout_size[0] // 2, min=0, max=x.size(2) - 1)
+ grid_y = torch.clamp(grid_y + offset_y - cutout_size[1] // 2, min=0, max=x.size(3) - 1)
+ mask = torch.ones(x.size(0), x.size(2), x.size(3), dtype=x.dtype, device=x.device)
+ mask[grid_batch, grid_x, grid_y] = 0
+ x = x * mask.unsqueeze(1)
+ return x
+
+
+def rand_resize(x: torch.Tensor, min_ratio: float = 0.8, max_ratio: float = 1.2) -> torch.Tensor:
+ resize_ratio = np.random.rand()*(max_ratio-min_ratio) + min_ratio
+ resized_img = F.interpolate(x, size=int(resize_ratio*x.shape[3]), mode='bilinear')
+ org_size = x.shape[3]
+ if int(resize_ratio*x.shape[3]) < x.shape[3]:
+ left_pad = (x.shape[3]-int(resize_ratio*x.shape[3]))/2.
+ left_pad = int(left_pad)
+ right_pad = x.shape[3] - left_pad - resized_img.shape[3]
+ x = F.pad(resized_img, (left_pad, right_pad, left_pad, right_pad), "constant", 0.)
+ else:
+ left = (int(resize_ratio*x.shape[3])-x.shape[3])/2.
+ left = int(left)
+ x = resized_img[:, :, left:(left+x.shape[3]), left:(left+x.shape[3])]
+ assert x.shape[2] == org_size
+ assert x.shape[3] == org_size
+ return x
+
+
+AUGMENT_FNS = {
+ 'color': [rand_brightness, rand_saturation, rand_contrast],
+ 'translation': [rand_translation],
+ 'resize': [rand_resize],
+ 'cutout': [rand_cutout],
+}
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/drop_path.py b/grn/tokenizer/videovae/modules/drop_path.py
new file mode 100644
index 0000000000000000000000000000000000000000..991a323e33304a2fe71c9757c13ace5c89fd8a5a
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/drop_path.py
@@ -0,0 +1,36 @@
+# from timm.models.layers import DropPath
+import torch
+
+def drop_path(x, drop_prob: float = 0., training: bool = False, scale_by_keep: bool = True):
+ """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
+
+ This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
+ the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
+ See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
+ changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
+ 'survival rate' as the argument.
+
+ """
+ if drop_prob == 0. or not training:
+ return x
+ keep_prob = 1 - drop_prob
+ shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
+ random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
+ if keep_prob > 0.0 and scale_by_keep:
+ random_tensor.div_(keep_prob)
+ return x * random_tensor
+
+
+class DropPath(torch.nn.Module):
+ """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
+ """
+ def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):
+ super(DropPath, self).__init__()
+ self.drop_prob = drop_prob
+ self.scale_by_keep = scale_by_keep
+
+ def forward(self, x):
+ return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
+
+ def extra_repr(self):
+ return f'drop_prob={round(self.drop_prob,3):0.3f}'
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/loss.py b/grn/tokenizer/videovae/modules/loss.py
new file mode 100644
index 0000000000000000000000000000000000000000..2dcd8254ac9d0e35da41bcc4ee861ce7a63354f4
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/loss.py
@@ -0,0 +1,86 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torch.autograd import grad
+
+def hinge_d_loss(logits_real, logits_fake):
+ loss_real = torch.mean(F.relu(1. - logits_real))
+ loss_fake = torch.mean(F.relu(1. + logits_fake))
+ d_loss = 0.5 * (loss_real + loss_fake)
+ return d_loss
+
+def vanilla_d_loss(logits_real, logits_fake):
+ d_loss = 0.5 * (
+ torch.mean(torch.nn.functional.softplus(-logits_real)) +
+ torch.mean(torch.nn.functional.softplus(logits_fake)))
+ return d_loss
+
+def get_disc_loss(disc_loss_type):
+ if disc_loss_type == 'vanilla':
+ disc_loss = vanilla_d_loss
+ elif disc_loss_type == 'hinge': # this way
+ disc_loss = hinge_d_loss
+ return disc_loss
+
+def adopt_weight(global_step, threshold=0, value=0., warmup=0):
+ if global_step < threshold or threshold < 0:
+ weight = value
+ else:
+ weight = 1
+ if global_step - threshold < warmup:
+ weight = min((global_step - threshold) / warmup, 1)
+ return weight
+
+def gradient_penalty(discriminator, real_data, fake_data, device):
+ alpha = torch.rand(real_data.size(0), 1, device=device)
+ alpha = alpha.expand_as(real_data)
+ interpolates = alpha * real_data + ((1 - alpha) * fake_data)
+ interpolates = torch.autograd.Variable(interpolates, requires_grad=True)
+
+ d_interpolates = discriminator(interpolates)
+ gradients = grad(
+ outputs=d_interpolates,
+ inputs=interpolates,
+ grad_outputs=torch.ones_like(d_interpolates, device=device),
+ create_graph=True,
+ retain_graph=True,
+ only_inputs=True
+ )[0]
+
+ gradients = gradients.view(gradients.size(0), -1)
+ gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
+ return gradient_penalty
+
+class InfoNCELoss(nn.Module):
+ def __init__(self, temperature: float = 0.07):
+ super(InfoNCELoss, self).__init__()
+ self.temperature = temperature
+
+ def forward(self, features: torch.Tensor, features_prime: torch.Tensor) -> torch.Tensor:
+ batch_size = features.shape[0]
+
+ # Normalize feature vectors
+ features = F.normalize(features, dim=1)
+ features_prime = F.normalize(features_prime, dim=1)
+
+ # Concatenate features and features_prime
+ combined_features = torch.cat([features, features_prime], dim=0)
+
+ # Compute similarity matrix
+ similarity_matrix = torch.matmul(combined_features, combined_features.T) / self.temperature
+
+ # Mask to exclude self-similarity
+ mask = torch.eye(2 * batch_size, dtype=torch.bool).to(features.device)
+ similarity_matrix.masked_fill_(mask, float('-inf'))
+
+ # Create labels for contrastive loss
+ labels = torch.arange(batch_size).to(features.device)
+ labels = torch.cat([labels, labels], dim=0)
+
+ # Compute logits (separate positive and negative pairs)
+ positives_logits = torch.cat([similarity_matrix[:batch_size, batch_size:], similarity_matrix[batch_size:, :batch_size]], dim=0)
+
+ # The labels are like: [0, 1, 2, ..., batch_size-1, 0, 1, 2, ..., batch_size-1]
+ loss = F.cross_entropy(positives_logits, labels)
+
+ return loss
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/lpips.py b/grn/tokenizer/videovae/modules/lpips.py
new file mode 100644
index 0000000000000000000000000000000000000000..9eaf045ff4461f644bca443c1f8368b29fa06032
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/lpips.py
@@ -0,0 +1,244 @@
+"""Stripped version of https://github.com/richzhang/PerceptualSimilarity/tree/master/models"""
+
+import os, hashlib
+import requests
+from tqdm import tqdm
+
+import torch
+import torch.nn as nn
+from torchvision import models
+from collections import namedtuple
+
+from videovae.utils.misc import data_prefix_manager, set_tf32_flags
+
+URL_MAP = {
+ "vgg_lpips": "https://heibox.uni-heidelberg.de/f/607503859c864bc1b30b/?dl=1"
+}
+
+CKPT_MAP = {
+ "vgg_lpips": "vgg.pth"
+}
+
+MD5_MAP = {
+ "vgg_lpips": "d507d7349b931f0638a25a48a722f98a"
+}
+
+def download(url, local_path, chunk_size=1024):
+ os.makedirs(os.path.split(local_path)[0], exist_ok=True)
+ with requests.get(url, stream=True) as r:
+ total_size = int(r.headers.get("content-length", 0))
+ with tqdm(total=total_size, unit="B", unit_scale=True) as pbar:
+ with open(local_path, "wb") as f:
+ for data in r.iter_content(chunk_size=chunk_size):
+ if data:
+ f.write(data)
+ pbar.update(chunk_size)
+
+
+def md5_hash(path):
+ with open(path, "rb") as f:
+ content = f.read()
+ return hashlib.md5(content).hexdigest()
+
+
+def get_ckpt_path(name, root, check=False):
+ assert name in URL_MAP
+ path = os.path.join(root, CKPT_MAP[name])
+ if not os.path.exists(path) or (check and not md5_hash(path) == MD5_MAP[name]):
+ print("Downloading {} model from {} to {}".format(name, URL_MAP[name], path))
+ download(URL_MAP[name], path)
+ md5 = md5_hash(path)
+ assert md5 == MD5_MAP[name], md5
+ return path
+
+def build_lpips_model(args):
+ if hasattr(args, 'lpips_model'):
+ if args.lpips_model == "vgg":
+ image_perceptual_model = LPIPS(upcast_tf32=args.upcast_tf32).eval()
+ video_perceptual_model = LPIPS(upcast_tf32=args.upcast_tf32).eval()
+ elif args.lpips_model == "resnet50":
+ image_perceptual_model = ResNet50LPIPS().eval()
+ video_perceptual_model = ResNet50LPIPS().eval()
+ elif args.lpips_model == "swin3d_t":
+ image_perceptual_model = LPIPS(upcast_tf32=args.upcast_tf32).eval()
+ video_perceptual_model = Swin3DTLPIPS(perceptual_layers=args.video_perceptual_layers).eval()
+ else:
+ raise NotImplementedError
+ for model in [image_perceptual_model, video_perceptual_model]:
+ for param in model.parameters():
+ param.requires_grad = False
+ else:
+ image_perceptual_model = video_perceptual_model = None
+ return image_perceptual_model, video_perceptual_model
+
+class ResNet50LPIPS(nn.Module):
+ def __init__(self):
+ super().__init__()
+ resnet50 = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)
+ self.lpips_net = nn.Sequential(*(list(resnet50.children())[:-2]))
+ self.lpips_loss = nn.MSELoss()
+
+ def forward(self, input, target):
+ return self.lpips_loss(self.lpips_net(input), self.lpips_net(target),)
+
+# Define a function to hook into each SwinTransformerBlockV2
+def get_features(outputs, hooks, x):
+ """Get features at each hook"""
+ hooks.append(x)
+
+# Custom feature extractor
+class FeatureExtractor(nn.Module):
+ def __init__(self, model, perceptual_layers=[]):
+ super(FeatureExtractor, self).__init__()
+ self.model = model
+ self.hooks = []
+
+ # Register hooks at each SwinTransformerBlockV2 layer
+ for name, module in self.model.named_modules():
+ if name in perceptual_layers:
+ module.register_forward_hook(lambda module, input, output: get_features(output, self.hooks, input))
+
+ def forward(self, x):
+ self.hooks = []
+ _ = self.model(x)
+ return self.hooks
+
+class Swin3DTLPIPS(nn.Module):
+ def __init__(self, perceptual_layers=[]):
+ super().__init__()
+ swin_v2_s_model = models.video.swin3d_t(weights=models.video.Swin3D_T_Weights.DEFAULT)
+ self.feature_extractor = FeatureExtractor(swin_v2_s_model, perceptual_layers=perceptual_layers)
+ self.lpips_loss = nn.MSELoss()
+
+ def forward(self, rec_video, gt_video):
+ rec_features = self.feature_extractor(rec_video)
+ gt_features = self.feature_extractor(gt_video)
+ loss = 0.
+ for rec_f, gt_f in zip(rec_features, gt_features):
+ mse_loss = self.lpips_loss(rec_f[0], gt_f[0])
+ loss += mse_loss
+ loss = loss / len(rec_features)
+ return loss
+
+class LPIPS(nn.Module):
+ # Learned perceptual metric
+ def __init__(self, use_dropout=True, upcast_tf32=False):
+ super().__init__()
+ self.upcast_tf32 = upcast_tf32
+ self.scaling_layer = ScalingLayer()
+ self.chns = [64, 128, 256, 512, 512] # vg16 features
+ self.net = vgg16(pretrained=True, requires_grad=False)
+ self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout)
+ self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout)
+ self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout)
+ self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout)
+ self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout)
+ self.load_from_pretrained()
+ for param in self.parameters():
+ param.requires_grad = False
+
+ def load_from_pretrained(self, name="vgg_lpips"):
+ ckpt = get_ckpt_path(name, os.path.join(os.path.dirname(os.path.abspath(__file__)), "cache"))
+ self.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu"), weights_only=True), strict=False)
+ print("loaded pretrained LPIPS loss from {}".format(ckpt))
+
+ @classmethod
+ def from_pretrained(cls, name="vgg_lpips"):
+ if name is not "vgg_lpips":
+ raise NotImplementedError
+ model = cls()
+ ckpt = get_ckpt_path(name, os.path.join(os.path.dirname(os.path.abspath(__file__)), "cache"))
+ model.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu"), weights_only=True), strict=False)
+ return model
+
+ def forward(self, input, target):
+ with set_tf32_flags(not self.upcast_tf32):
+ in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target))
+ outs0, outs1 = self.net(in0_input), self.net(in1_input)
+ feats0, feats1, diffs = {}, {}, {}
+ lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
+ for kk in range(len(self.chns)):
+ feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(outs1[kk])
+ diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
+
+ res = [spatial_average(lins[kk].model(diffs[kk]), keepdim=True) for kk in range(len(self.chns))]
+ val = res[0]
+ for l in range(1, len(self.chns)):
+ # print(res[l].shape)
+ val += res[l]
+
+ return val
+
+
+class ScalingLayer(nn.Module):
+ def __init__(self):
+ super(ScalingLayer, self).__init__()
+ self.register_buffer('shift', torch.Tensor([-.030, -.088, -.188])[None, :, None, None])
+ self.register_buffer('scale', torch.Tensor([.458, .448, .450])[None, :, None, None])
+
+ def forward(self, inp):
+ return (inp - self.shift) / self.scale
+
+
+class NetLinLayer(nn.Module):
+ """ A single linear layer which does a 1x1 conv """
+ def __init__(self, chn_in, chn_out=1, use_dropout=False):
+ super(NetLinLayer, self).__init__()
+ layers = [nn.Dropout(), ] if (use_dropout) else []
+ layers += [nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False), ]
+ self.model = nn.Sequential(*layers)
+
+
+class vgg16(torch.nn.Module):
+ def __init__(self, requires_grad=False, pretrained=True):
+ super(vgg16, self).__init__()
+ # load locally
+ assert pretrained == True
+ vgg_model = models.vgg16()
+ vgg_model.load_state_dict(torch.load(data_prefix_manager("checkpoints/vgg16-397923af.pth"), weights_only=True))
+ vgg_pretrained_features = vgg_model.features
+
+ self.slice1 = torch.nn.Sequential()
+ self.slice2 = torch.nn.Sequential()
+ self.slice3 = torch.nn.Sequential()
+ self.slice4 = torch.nn.Sequential()
+ self.slice5 = torch.nn.Sequential()
+ self.N_slices = 5
+ for x in range(4):
+ self.slice1.add_module(str(x), vgg_pretrained_features[x])
+ for x in range(4, 9):
+ self.slice2.add_module(str(x), vgg_pretrained_features[x])
+ for x in range(9, 16):
+ self.slice3.add_module(str(x), vgg_pretrained_features[x])
+ for x in range(16, 23):
+ self.slice4.add_module(str(x), vgg_pretrained_features[x])
+ for x in range(23, 30):
+ self.slice5.add_module(str(x), vgg_pretrained_features[x])
+ if not requires_grad:
+ for param in self.parameters():
+ param.requires_grad = False
+
+ def forward(self, X):
+ h = self.slice1(X)
+ h_relu1_2 = h
+ h = self.slice2(h)
+ h_relu2_2 = h
+ h = self.slice3(h)
+ h_relu3_3 = h
+ h = self.slice4(h)
+ h_relu4_3 = h
+ h = self.slice5(h)
+ h_relu5_3 = h
+ vgg_outputs = namedtuple("VggOutputs", ['relu1_2', 'relu2_2', 'relu3_3', 'relu4_3', 'relu5_3'])
+ out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3)
+ return out
+
+
+def normalize_tensor(x,eps=1e-10):
+ # norm_factor = torch.sqrt(torch.sum(x**2,dim=1,keepdim=True))
+ norm_factor = x.norm(p=2, dim=1, keepdim=True)
+ return x/(norm_factor+eps)
+
+
+def spatial_average(x, keepdim=True):
+ return x.mean([2,3],keepdim=keepdim)
diff --git a/grn/tokenizer/videovae/modules/misc.py b/grn/tokenizer/videovae/modules/misc.py
new file mode 100644
index 0000000000000000000000000000000000000000..8cf03f4b7885c4773fcaafde9aedc66ed205b0f9
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/misc.py
@@ -0,0 +1,13 @@
+import torch
+
+def swish(x):
+ if type(x) == list:
+ for i in range(len(x)):
+ x[i] = swish(x[i])
+ return x
+ try:
+ return x*torch.sigmoid(x)
+ except:
+ for _i in range(x.shape[2]):
+ x[:,:,_i:_i+1,:,:] = x[:,:,_i:_i+1,:,:]*torch.sigmoid(x[:,:,_i:_i+1,:,:])
+ return x
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/modules/normalization.py b/grn/tokenizer/videovae/modules/normalization.py
new file mode 100644
index 0000000000000000000000000000000000000000..b88e8164ebd25caee8fef3cae2fd02c2b6cde969
--- /dev/null
+++ b/grn/tokenizer/videovae/modules/normalization.py
@@ -0,0 +1,152 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from einops import rearrange
+
+
+def get_norm(norm_type):
+ if norm_type == "spatial-group":
+ return SpatialGroupNorm
+ elif norm_type == "rms":
+ return RMS_norm
+ elif norm_type == "group":
+ return nn.GroupNorm
+ else:
+ raise NotImplementedError
+
+class RMS_norm(nn.Module):
+
+ def __init__(self, num_channels, channel_first=True, bias=False, **kwargs):
+ super().__init__()
+ broadcastable_dims = (1, 1, 1)
+ shape = (num_channels, *broadcastable_dims)
+
+ self.channel_first = channel_first
+ self.scale = num_channels**0.5
+ self.gamma = nn.Parameter(torch.ones(shape))
+ self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
+
+ def forward(self, x):
+ return F.normalize(
+ x, dim=(1 if self.channel_first else
+ -1)) * self.scale * self.gamma + self.bias
+
+class SpatialGroupNorm(nn.GroupNorm):
+ def __init__(self, *args, **kwargs):
+ super(SpatialGroupNorm, self).__init__(*args, **kwargs)
+
+ def shard_norm(self, x):
+ dtype = x.dtype
+ x = x.to(torch.float32)
+ with torch.amp.autocast("cuda", torch.float32):
+ for _i in range(x.shape[0]):
+ x[_i:_i+1,...] = super(SpatialGroupNorm, self).forward(x[_i:_i+1,...])
+ x = x.to(dtype=dtype)
+ return x
+
+ def forward(self, x):
+ dtype = x.dtype
+ x = x.to(torch.float32)
+ assert x.ndim == 5
+ T = x.shape[2]
+ x = rearrange(x, "B C T H W -> (B T) C H W")
+ try:
+ x = super(SpatialGroupNorm, self).forward(x)
+ except:
+ x = self.shard_norm(x) # shard norm if OOM fallback
+ x = rearrange(x, "(B T) C H W -> B C T H W", T=T)
+ x = x.to(dtype=dtype)
+ return x
+
+class Normalize(nn.Module):
+ def __init__(self, in_channels, norm_type, norm_axis="spatial"):
+ super().__init__()
+ self.norm_axis = norm_axis
+ assert norm_type in ['group', 'batch', "no"]
+ if norm_type == 'group':
+ if in_channels % 32 == 0:
+ self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
+ elif in_channels % 24 == 0:
+ self.norm = nn.GroupNorm(num_groups=24, num_channels=in_channels, eps=1e-6, affine=True)
+ else:
+ raise NotImplementedError
+ elif norm_type == 'batch':
+ self.norm = nn.SyncBatchNorm(in_channels, track_running_stats=False) # Runtime Error: grad inplace if set track_running_stats to True
+ elif norm_type == 'no':
+ self.norm = nn.Identity()
+
+ def _norm(self, x):
+ try:
+ x = self.norm(x)
+ except:
+ device = x.device
+ self.norm_cpu = self.norm.cpu()
+ x = self.norm_cpu(x.cpu().pin_memory()).to(device=device)
+ return x
+
+ def shard_norm(self, x):
+ dtype = x.dtype
+ x = x.to(torch.float32)
+ with torch.amp.autocast("cuda", torch.float32):
+ for _i in range(x.shape[0]):
+ x[_i:_i+1,...] = self.norm(x[_i:_i+1,...])
+ x = x.to(dtype=dtype)
+ return x
+
+ def forward(self, x):
+ if self.norm_axis == "spatial":
+ if type(x) == list:
+ for i in range(len(x)):
+ x[i] = self.norm(x[i])
+ return x
+ if x.ndim == 4:
+ try:
+ x = self.norm(x)
+ except:
+ x = self.shard_norm(x)
+ else:
+ B, C, T, H, W = x.shape
+ x = rearrange(x, "B C T H W -> (B T) C H W")
+ # x = self.shard_norm(x)
+ try:
+ x = self.norm(x)
+ except:
+ x = self.shard_norm(x)
+ x = rearrange(x, "(B T) C H W -> B C T H W", T=T)
+ elif self.norm_axis == "spatial-temporal":
+ x = self._norm(x)
+ else:
+ raise NotImplementedError
+ return x
+
+def l2norm(t):
+ return F.normalize(t, dim=-1)
+
+class LayerNorm(nn.Module):
+ def __init__(self, dim):
+ super().__init__()
+ self.gamma = nn.Parameter(torch.ones(dim))
+ self.register_buffer("beta", torch.zeros(dim))
+
+ def forward(self, x):
+ return F.layer_norm(x, x.shape[-1:], self.gamma, self.beta)
+
+# https://github.com/huggingface/transformers/blob/2f12e408225b1ebceb0d2f701ce419d46678dc31/src/transformers/models/llama/modeling_llama.py#L76
+class RMSNorm(nn.Module):
+ def __init__(self, hidden_size, eps=1e-6):
+ """
+ LlamaRMSNorm is equivalent to T5LayerNorm
+ """
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(hidden_size))
+ self.variance_epsilon = eps
+
+ def forward(self, hidden_states, sp_slice=None):
+ input_dtype = hidden_states.dtype
+ hidden_states = hidden_states.to(torch.float32)
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
+ if sp_slice is None:
+ return (self.weight * hidden_states).to(input_dtype)
+ else:
+ return (self.weight[sp_slice] * hidden_states).to(input_dtype) # torch.float32 * torchbfloat16 in DDP will cast to torch.float32
diff --git a/grn/tokenizer/videovae/utils/__init__.py b/grn/tokenizer/videovae/utils/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..1d4e78e63ad402a475010b863bb9288a76489c92
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/__init__.py
@@ -0,0 +1 @@
+from .arguments import str2bool
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/utils/arguments.py b/grn/tokenizer/videovae/utils/arguments.py
new file mode 100644
index 0000000000000000000000000000000000000000..a6579353e1cfeb744cae7f396703b7f0e6da1489
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/arguments.py
@@ -0,0 +1,281 @@
+import argparse
+
+
+def str2bool(v):
+ if isinstance(v, bool):
+ return v
+ if v.lower() in ('true'):
+ return True
+ elif v.lower() in ('false'):
+ return False
+ else:
+ raise argparse.ArgumentTypeError('Boolean value expected.')
+
+def add_model_specific_args(args, parser):
+ from videovae.models.hbq_tokenizer import HBQ_Tokenizer
+ if args.tokenizer in ["hbq_tokenizer"]:
+ parser = HBQ_Tokenizer.add_model_specific_args(parser)
+ vae_model = HBQ_Tokenizer
+ else:
+ raise NotImplementedError
+ return args, parser, vae_model
+
+class MainArgs:
+ @staticmethod
+ def add_main_args(parser):
+ # training
+ parser.add_argument('--max_steps', type=int, default=1e6)
+ parser.add_argument('--log_every', type=int, default=1)
+ parser.add_argument('--ckpt_every', type=int, default=1000)
+ parser.add_argument('--default_root_dir', type=str, required=True)
+ parser.add_argument('--compile', type=str, default="no", choices=["no", "yes"])
+ parser.add_argument('--ema', type=str, default="no", choices=["no", "yes"])
+ parser.add_argument('--mfu_logging', type=str, default="no", choices=["no", "yes"])
+ parser.add_argument('--dataloader_init_epoch', type=int, default=-1)
+ parser.add_argument('--context_parallel_size', type=int, default=0)
+ parser.add_argument('--video_ranks_ratio', type=float, default=-1.0)
+
+ # optimization
+ parser.add_argument('--lr', type=float, default=1e-4)
+ parser.add_argument('--beta1', type=float, default=0.9)
+ parser.add_argument('--beta2', type=float, default=0.95)
+ parser.add_argument('--optim_type', type=str, default="Adam", choices=["Adam", "AdamW"])
+ parser.add_argument('--disc_optim_type', type=str, default=None, choices=[None, "rmsprop"])
+ parser.add_argument('--max_grad_norm', type=float, default=1.0)
+ parser.add_argument('--max_grad_norm_disc', type=float, default=1.0)
+ parser.add_argument('--disable_sch', action="store_true") # deprecated option
+ parser.add_argument('--scheduler', type=str, default="no", choices=["no", "linear"])
+ parser.add_argument('--warmup_steps', type=int, default=0)
+ parser.add_argument('--lr_min', type=float, default=0.)
+ parser.add_argument('--warmup_lr_init', type=float, default=0.)
+
+ # basic vae config
+ parser.add_argument('--patch_size', type=int, default=8)
+ parser.add_argument('--temporal_patch_size', type=int, default=4)
+ parser.add_argument('--embedding_dim', type=int, default=256)
+ parser.add_argument('--codebook_dim', type=int, default=16)
+ parser.add_argument('--codebook_dim_low', type=int, default=-1) # low-dim latents
+ parser.add_argument('--use_vae', action="store_true")
+
+ # EQ-VAE
+ parser.add_argument('--eq_scale_prior', type=float, default=0.)
+ parser.add_argument('--eq_angle_prior', type=float, default=0.)
+
+ # discrete vae config
+ parser.add_argument('--use_stochastic_depth', action="store_true")
+ parser.add_argument("--drop_rate", type=float, default=0.0)
+ parser.add_argument('--schedule_mode', type=str, default="original")
+ parser.add_argument('--lr_drop', nargs='*', type=int, default=None, help="A list of numeric values. Example: --values 270 300")
+ parser.add_argument('--lr_drop_rate', type=float, default=0.1)
+ parser.add_argument('--keep_first_quant', action="store_true")
+ parser.add_argument('--keep_last_quant', action="store_true")
+ parser.add_argument('--remove_residual_detach', action="store_true")
+ parser.add_argument('--use_out_phi', action="store_true")
+ parser.add_argument('--use_out_phi_res', action="store_true")
+ parser.add_argument('--use_lecam_reg', action="store_true")
+ parser.add_argument('--lecam_weight', type=float, default=0.05)
+ parser.add_argument('--perceptual_model', type=str, default="vgg16", choices=["vgg16", "resnet50", "resnet50_v2"])
+ parser.add_argument('--base_ch_disc', type=int, default=64)
+ parser.add_argument('--random_flip', action="store_true")
+ parser.add_argument('--flip_prob', type=float, default=0.5)
+ parser.add_argument('--flip_mode', type=str, default="stochastic", choices=["stochastic", "deterministic", "stochastic_dynamic"])
+ parser.add_argument('--max_flip_lvl', type=int, default=1)
+ parser.add_argument('--not_load_optimizer', action="store_true")
+ parser.add_argument('--use_lecam_reg_zero', action="store_true")
+ parser.add_argument('--freeze_encoder', action="store_true")
+ parser.add_argument('--freeze_decoder', action="store_true")
+ parser.add_argument('--rm_downsample', action="store_true")
+ parser.add_argument('--random_flip_1lvl', action="store_true")
+ parser.add_argument('--flip_lvl_idx', type=int, default=0)
+ parser.add_argument('--drop_when_test', action="store_true")
+ parser.add_argument('--drop_lvl_idx', type=int, default=None)
+ parser.add_argument('--drop_lvl_num', type=int, default=0)
+ parser.add_argument('--compute_all_commitment', action="store_true")
+ parser.add_argument('--disable_codebook_usage', action="store_true")
+ parser.add_argument('--freeze_enc_main', action="store_true")
+ parser.add_argument('--freeze_dec_main', action="store_true")
+ parser.add_argument('--random_short_schedule', action="store_true")
+ parser.add_argument('--short_schedule_prob', type=float, default=0.5)
+ parser.add_argument('--use_bernoulli', action="store_true")
+ parser.add_argument('--use_rot_trick', action="store_true")
+ parser.add_argument('--disable_flip_prob', type=float, default=0.0)
+ parser.add_argument('--dino_disc', action="store_true")
+ parser.add_argument('--quantizer_type', type=str, default='MultiScaleBSQ')
+ parser.add_argument('--lfq_weight', type=float, default=0.)
+ parser.add_argument('--entropy_loss_weight', type=float, default=0.1)
+ parser.add_argument('--visu_every', type=int, default=1000)
+ parser.add_argument('--commitment_loss_weight', type=float, default=0.25)
+ parser.add_argument('--bsq_version', type=str, default="v1", choices=["v1", "v2"])
+ parser.add_argument('--diversity_gamma', type=float, default=1)
+ parser.add_argument('--bs1_for1024', action="store_true")
+ parser.add_argument('--casual_multi_scale', action="store_true")
+ parser.add_argument('--double_compress_t', action="store_true")
+ parser.add_argument('--temporal_slicing', action="store_true")
+ parser.add_argument('--latent_adjust_type', type=str, default=None)
+ parser.add_argument('--compute_latent_loss', action="store_true")
+ parser.add_argument('--latent_loss_weight', type=float, default=0.0)
+ parser.add_argument('--use_raw_latentz', action="store_true")
+ parser.add_argument('--last_scale_repeat_n', type=int, default=0)
+ parser.add_argument('--num_lvl_fsq', type=int, default=5)
+ parser.add_argument('--use_midscale_sup', action="store_true")
+ parser.add_argument('--midscale_list', nargs='*', default=[0.5, 0.75, 1.0])
+ parser.add_argument('--use_eq', action="store_true") # Equivariance Regularization
+ parser.add_argument('--eq_prob', type=float, default=0.5)
+ parser.add_argument('--disc_prob', type=float, default=0.5)
+ parser.add_argument('--test_mode', type=str, default="discrete", choices=["discrete", "continuous"])
+ parser.add_argument('--remove_disc', type=int, default=0)
+
+ # discriminator config
+ parser.add_argument('--disc_version', type=str, default="v1")
+ parser.add_argument('--magvit_disc', action="store_true") # deprecated
+ parser.add_argument('--disc_type', type=str, default="patchgan", choices=["patchgan", "stylegan", "spectralgan"])
+ parser.add_argument('--sigmoid_in_disc', action="store_true")
+ parser.add_argument('--activation_in_disc', type=str, default="leaky_relu")
+ parser.add_argument('--apply_blur', action="store_true")
+ parser.add_argument('--apply_noise', action="store_true")
+ parser.add_argument('--dis_warmup_steps', type=int, default=0)
+ parser.add_argument('--dis_lr_multiplier', type=float, default=1.)
+ parser.add_argument('--dis_minlr_multiplier', action="store_true")
+ parser.add_argument('--disc_channels', type=int, default=64)
+ parser.add_argument('--disc_layers', type=int, default=3)
+ parser.add_argument('--discriminator_iter_start', type=int, default=0)
+ parser.add_argument('--disc_pretrain_iter', type=int, default=0)
+ parser.add_argument('--disc_optim_steps', type=int, default=1)
+ parser.add_argument('--disc_warmup', type=int, default=0)
+ parser.add_argument('--disc_pool', type=str, default="no", choices=["no", "yes"])
+ parser.add_argument('--disc_pool_size', type=int, default=100)
+ parser.add_argument('--disc_temporal_compress', type=str, default="yes", choices=["no", "yes"])
+ parser.add_argument('--disc_use_blur', type=str, default="yes", choices=["no", "yes"])
+ parser.add_argument('--disc_stylegan_downsample_base', type=int, default=2)
+
+ parser = MainArgs.add_loss_args(parser)
+ parser = MainArgs.add_accelerate_args(parser)
+
+ # initialization
+ parser.add_argument('--tokenizer', type=str, required=True)
+ parser.add_argument('--pretrained', type=str, default=None)
+ parser.add_argument('--pretrained_mode', type=str, default="full")
+ parser.add_argument('--pretrained_ema', type=str, default="no")
+ parser.add_argument('--inflation_pe', action="store_true")
+ parser.add_argument('--init_vgen', type=str, default='no', choices=['no', 'keep', 'average'])
+ parser.add_argument('--no_init_idis', action="store_true") # deprecated option
+ parser.add_argument('--init_idis', type=str, default='keep', choices=['no', 'keep']) # use keep by default following previous settings
+ parser.add_argument('--init_vdis', type=str, default="no")
+
+ # misc
+ parser.add_argument('--enable_nan_detector', action='store_true')
+ parser.add_argument('--turn_on_profiler', action='store_true')
+ parser.add_argument('--profiler_scheduler_wait_steps', type=int, default=10)
+ parser.add_argument('--debug', action='store_true')
+ parser.add_argument('--video_logger', action='store_true') # deprecated option
+ parser.add_argument('--data_root', type=str, default="sg")
+ parser.add_argument('--username', type=str, default="bin.yan")
+ parser.add_argument('--seed', type=int, default=1234)
+ parser.add_argument('--vq_to_vae', action='store_true')
+ parser.add_argument('--load_not_strict', action='store_true')
+ parser.add_argument('--zero', type=int, default=0, choices=[0, 1, 2, 3]) # 1 hybrid shard, 2 shard grad_op, 3 full shard
+ parser.add_argument('--bucket_cap_mb', type=int, default=40) # DDP
+ parser.add_argument('--manual_gc_interval', type=int, default=10000) # DDP
+
+ # hj arguments
+ parser.add_argument('--semantic_scale_dim', type=int, default=8)
+ parser.add_argument('--detail_scale_dim', type=int, default=64)
+ parser.add_argument('--use_feat_proj', type=int, default=1)
+ parser.add_argument('--skip_detail_scales_prob', type=float, default=0.5)
+ parser.add_argument('--semantic_scales', type=int, default=8)
+ parser.add_argument('--return_res', type=str, default='')
+ parser.add_argument('--use_multi_scale', type=int, default=1)
+ parser.add_argument('--quant_not_rely_256', type=int, default=1)
+ parser.add_argument('--image_scale_repetition', type=str, default='[5,5,5,5,5,5,5,4,3,2,1,1,1,1,1,1,1,1,1,1,1]')
+ parser.add_argument('--video_scale_repetition_times', type=int, nargs='+', default=[1, 2, 3, 4, 6, 8, 10, 20])
+ parser.add_argument('--video_scale_repetition_prob', type=float, nargs='+', default=[0.2, 0.2, 0.2, 0.15, 0.1, 0.05, 0.05, 0.05])
+ parser.add_argument('--use_fsq', type=int, default=0)
+ parser.add_argument('--semantic_num_lvl', type=int, default=2)
+ parser.add_argument('--detail_num_lvl', type=int, default=2)
+ parser.add_argument('--middle_scale_dim', type=int, default=16)
+ parser.add_argument('--div_delta_t', type=int, default=0)
+ parser.add_argument('--multi_scale_freq', type=str, default='[1.0,3.0,9.0,27.0]')
+ parser.add_argument('--train_continuous', type=int, default=0)
+ parser.add_argument('--test_type', type=str, default='detail', choices=['detail', 'semantic'])
+ parser.add_argument('--quant_method', type=str, default='default')
+ parser.add_argument('--remove_enlarge_factors', type=int, default=0)
+ parser.add_argument('--elementwise_enlarge_factor', type=int, default=0, choices=[0,1])
+ parser.add_argument('--quant_unit', type=int, default=-1)
+ parser.add_argument('--use_channelwise_std', type=int, default=0)
+ parser.add_argument('--enlarge_factors_type', type=str, default='')
+ parser.add_argument('--minus_mean_each_scale', type=int, default=0, choices=[0,1])
+ parser.add_argument('--normed_tgt_std', type=float, default=0.2)
+ parser.add_argument('--same_noise_among_scales', type=int, default=0, choices=[0,1])
+ parser.add_argument('--enable_online_download', type=int, default=0, choices=[0,1])
+ return parser
+
+ @staticmethod
+ def add_loss_args(parser):
+ parser.add_argument('--fix_model', type=str, nargs="+", default=['no'])
+ parser.add_argument("--recon_loss_type", type=str, default='l1', choices=['l1', 'l2'])
+ parser.add_argument('--image_gan_weight', type=float, default=1.0)
+ parser.add_argument('--video_gan_weight', type=float, default=1.0)
+ parser.add_argument('--image_disc_weight', type=float, default=0.)
+ parser.add_argument('--video_disc_weight', type=float, default=0.)
+ parser.add_argument('--vf_weight', type=float, default=0.)
+ parser.add_argument('--vf_weight_approx', type=float, default=-1)
+ parser.add_argument('--vf_distmat_margin', type=float, default=0.25)
+ parser.add_argument('--vf_cos_margin', type=float, default=0.5)
+ parser.add_argument('--temporal_alignment', type=str, default=None, choices=["aux_mean", "z_upsample"])
+ parser.add_argument('--l1_weight', type=float, default=4.0)
+ parser.add_argument('--gan_feat_weight', type=float, default=0.0)
+ parser.add_argument('--lpips_model', type=str, default='vgg', choices=['vgg', 'resnet50', 'swin3d_t'])
+ parser.add_argument('--perceptual_weight', type=float, default=0.0)
+ parser.add_argument('--video_perceptual_weight', type=float, default=None)
+ parser.add_argument('--video_perceptual_layers', type=str, nargs="+", default=[])
+ parser.add_argument('--kl_weight', type=float, default=0.)
+ parser.add_argument('--norm_type', type=str, default='group', choices=['batch', 'group', "spatial-group", "rms"])
+ parser.add_argument('--disc_loss_type', type=str, default='hinge', choices=['hinge', 'vanilla'])
+ parser.add_argument('--gan_image4video', type=str, default='yes', choices=['no', 'yes'])
+ return parser
+
+ @staticmethod
+ def add_accelerate_args(parser):
+ parser.add_argument('--use_checkpoint', action="store_true")
+ parser.add_argument('--precision', type=str, default="fp32", choices=['fp32', 'bf16']) # disable fp16
+ parser.add_argument('--encoder_dtype', type=str, default="fp32", choices=['fp32', 'bf16']) # disable fp16
+ parser.add_argument('--decoder_dtype', type=str, default="fp32", choices=['fp32', 'bf16']) # disable fp16
+ parser.add_argument('--upcast_attention', type=str, default="", choices=["qk", "qkv"])
+ parser.add_argument('--upcast_tf32', action="store_true")
+ return parser
+
+def format_args(args):
+ # Start building the script string
+ script_content = "#!/bin/bash\n\n"
+ script_content += "torchrun \\\n"
+ script_content += " --nproc_per_node=$ARNOLD_WORKER_GPU \\\n"
+ script_content += " --nnodes=$ARNOLD_WORKER_NUM --master_addr=$ARNOLD_WORKER_0_HOST \\\n"
+ script_content += " --node_rank=$ARNOLD_ID --master_port=$port \\\n"
+ script_content += " train.py \\\n"
+
+ # Iterate over each key-value pair and append it to the command
+ for k, v in args.__dict__.items():
+ script_content += f" --{k} {v} \\\n"
+
+ # Remove the last backslash and newline
+ script_content = script_content.rstrip(" \\\n") + "\n"
+ return script_content
+
+def init_resolution(resolution, num_datasets):
+ if len(resolution) == 1:
+ resolution = [(resolution[0], resolution[0])] * num_datasets
+ elif len(resolution) == num_datasets:
+ resolution = [(resolution[i], resolution[i]) for i in range(len(resolution))]
+ elif len(resolution) == num_datasets * 2:
+ resolution = [(resolution[i], resolution[i+1]) for i in range(0, len(resolution), 2)]
+ else:
+ raise NotImplementedError
+ return resolution
+
+def init_args(args):
+ args.resolution = init_resolution(args.resolution, len(args.dataset_list))
+ if args.pretrained == "None":
+ args.pretrained = None
+ if args.video_perceptual_weight is None:
+ args.video_perceptual_weight = args.perceptual_weight
+ return args
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/utils/context_parallel.py b/grn/tokenizer/videovae/utils/context_parallel.py
new file mode 100644
index 0000000000000000000000000000000000000000..80a35ec35f4056dad30a11665c62e71b769a882d
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/context_parallel.py
@@ -0,0 +1,170 @@
+import math
+import torch
+import torch.nn as nn
+import torch.distributed as dist
+
+import videovae.utils.diffdist.functional as distops
+
+class ContextParallelUtils:
+ _CONTEXT_PARALLEL_GROUP = None
+ _CONTEXT_PARALLEL_SIZE = 0
+ _CONTEXT_PARALLEL_ON = False
+
+ """
+ {
+ "cp_size": 2,
+ }
+ """
+ CP_CONFIG = None
+
+ @staticmethod
+ def set_cp_on(on=True):
+ ContextParallelUtils._CONTEXT_PARALLEL_ON = on
+
+ @staticmethod
+ def cp_on():
+ return ContextParallelUtils._CONTEXT_PARALLEL_ON
+
+ @staticmethod
+ def get_cp_cfg():
+ return ContextParallelUtils.CP_CONFIG
+
+ @staticmethod
+ def is_cp_initialized():
+ if ContextParallelUtils._CONTEXT_PARALLEL_GROUP is None:
+ return False
+ else:
+ return True
+
+ @staticmethod
+ def initialize_context_parallel(cp_config:dict):
+ assert ContextParallelUtils._CONTEXT_PARALLEL_GROUP is None, "context parallel group is already initialized"
+
+ context_parallel_size = cp_config["cp_size"]
+ if context_parallel_size > 1:
+ ContextParallelUtils.CP_CONFIG = cp_config
+ else:
+ print(f"WARN: context parallel size must > 1 but got {context_parallel_size}")
+ return
+
+ ContextParallelUtils._CONTEXT_PARALLEL_SIZE = context_parallel_size
+
+ rank = torch.distributed.get_rank()
+ world_size = torch.distributed.get_world_size()
+
+ for i in range(0, world_size, context_parallel_size):
+ ranks = range(i, i + context_parallel_size)
+ group = torch.distributed.new_group(ranks)
+ if rank in ranks:
+ ContextParallelUtils._CONTEXT_PARALLEL_GROUP = group
+ break
+
+ @staticmethod
+ def get_cp_group():
+ return ContextParallelUtils._CONTEXT_PARALLEL_GROUP
+
+ @staticmethod
+ def get_cp_size():
+ return ContextParallelUtils._CONTEXT_PARALLEL_SIZE
+
+ @staticmethod
+ def get_cp_world_size():
+ if ContextParallelUtils.is_cp_initialized():
+ world_size = torch.distributed.get_world_size()
+ return world_size // ContextParallelUtils._CONTEXT_PARALLEL_SIZE
+ else:
+ return 0
+
+ @staticmethod
+ def get_cp_rank():
+ if ContextParallelUtils.is_cp_initialized():
+ global_rank = torch.distributed.get_rank()
+ cp_rank = global_rank % ContextParallelUtils._CONTEXT_PARALLEL_SIZE
+ return cp_rank
+ else:
+ return 0
+
+ def get_cp_group_rank():
+ if ContextParallelUtils.is_cp_initialized():
+ rank = torch.distributed.get_rank()
+ cp_group_rank = rank // ContextParallelUtils._CONTEXT_PARALLEL_SIZE
+ return cp_group_rank
+ else:
+ return 0
+
+
+def _gather_tensor_shape(local_ts):
+ cp_size = ContextParallelUtils.get_cp_size()
+ local_shape = torch.tensor(local_ts.shape, dtype=torch.int64, device=local_ts.device)
+ gathered_shapes = [torch.zeros(len(local_shape), dtype=torch.int64, device=local_ts.device) for _ in range(cp_size)]
+ dist.all_gather(gathered_shapes, local_shape, group=ContextParallelUtils._CONTEXT_PARALLEL_GROUP)
+ return [shape.tolist() for shape in gathered_shapes]
+
+@torch.compiler.disable()
+def dist_encoder_gather_result(res)->list:
+ cp_size = ContextParallelUtils.get_cp_size()
+ if cp_size < 2:
+ return res
+
+ shape_list = _gather_tensor_shape(res) # [[1,2,3,4],[x,x,x,x]] list of shapes on different rank
+ encs=[torch.zeros(s, device=res.device, dtype=res.dtype) for s in shape_list]
+
+ dist.barrier()
+ encs = distops.all_gather(encs, res, group=ContextParallelUtils._CONTEXT_PARALLEL_GROUP)
+ return encs
+
+@torch.compiler.disable()
+def dist_decoder_gather_result(res)->list:
+ cp_size = ContextParallelUtils.get_cp_size()
+ if cp_size < 2:
+ return res
+
+ shape_list = _gather_tensor_shape(res) # [[1,2,3,4],[x,x,x,x]] list of shapes on different rank
+ decs = [torch.zeros(s, device=res.device, dtype=res.dtype) for s in shape_list]
+
+ dist.barrier()
+ decs = distops.all_gather(decs, res, group=ContextParallelUtils._CONTEXT_PARALLEL_GROUP)
+ return decs
+
+
+def _send_with_shape(local_ts, next_rank):
+ local_shape = torch.tensor(local_ts.shape, dtype=torch.int64, device=local_ts.device)
+ torch.distributed.send(local_shape.contiguous(), next_rank)
+ torch.distributed.send(local_ts.contiguous(), next_rank)
+
+def _recv_with_shape(pre_rank):
+ device = torch.cuda.current_device() if torch.cuda.is_available() else torch.device('cpu')
+
+ shape = torch.zeros(5, dtype=torch.int64, device=device)
+ torch.distributed.recv(shape, pre_rank)
+ ts = torch.zeros(shape.tolist(), device=device)
+ torch.distributed.recv(ts, pre_rank)
+ return ts
+
+
+@torch.compiler.disable()
+def dist_conv_cache_send(conv_cache):
+
+ cp_rank = ContextParallelUtils.get_cp_rank()
+ global_rank = torch.distributed.get_rank()
+ cp_size = ContextParallelUtils.get_cp_size()
+
+ if cp_rank == cp_size - 1:
+ return
+ if conv_cache is None:
+ return
+
+ next_rank = global_rank + 1
+ _send_with_shape(conv_cache, next_rank)
+
+@torch.compiler.disable()
+def dist_conv_cache_recv():
+ cp_rank = ContextParallelUtils.get_cp_rank()
+ global_rank = torch.distributed.get_rank()
+
+ if cp_rank == 0:
+ return None
+
+ pre_rank = global_rank - 1
+ return _recv_with_shape(pre_rank)
+
diff --git a/grn/tokenizer/videovae/utils/diffdist/README.md b/grn/tokenizer/videovae/utils/diffdist/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..a16b530b25182895b3d8050afe8353c9cd74569c
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/README.md
@@ -0,0 +1,9 @@
+# diffdist, for differentiable communication
+
+borrowed from https://github.com/ag14774/diffdist and fix code:
+```
+ # tmp = dist.reduce(tensor_list[i], i, op, group, async_op=True)
+ # to
+ # tmp = dist.reduce(tensor_list[i].contiguous(), i, op, group, async_op=True)
+
+```
diff --git a/grn/tokenizer/videovae/utils/diffdist/__init__.py b/grn/tokenizer/videovae/utils/diffdist/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..3519969938f5ce874c6bd81963c98c485ce1fc8c
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/__init__.py
@@ -0,0 +1,3 @@
+from videovae.utils.diffdist import extra_collectives
+from videovae.utils.diffdist import functional
+from videovae.utils.diffdist import modules
diff --git a/grn/tokenizer/videovae/utils/diffdist/extra_collectives.py b/grn/tokenizer/videovae/utils/diffdist/extra_collectives.py
new file mode 100644
index 0000000000000000000000000000000000000000..50728a05a81888da86ecbd493b83cdf342e2f8cf
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/extra_collectives.py
@@ -0,0 +1,44 @@
+import torch.distributed as dist
+from torch.distributed import ReduceOp
+
+
+class AsyncOpList(object):
+ def __init__(self, ops):
+ self.ops = ops
+
+ def wait(self):
+ for op in self.ops:
+ op.wait()
+
+ def is_completed(self):
+ for op in self.ops:
+ if not op.is_completed():
+ return False
+ return True
+
+
+def reduce_scatter(tensor,
+ tensor_list,
+ op=ReduceOp.SUM,
+ group=dist.group.WORLD,
+ async_op=False):
+ ranks = dist.get_process_group_ranks(group)
+ rank = dist.get_rank(group)
+ if tensor is None:
+ tensor = tensor_list[rank]
+ if tensor.dim() == 0:
+ tensor = tensor.view(-1)
+ tensor[:] = tensor_list[rank]
+ ops = []
+ for i in range(dist.get_world_size(group)):
+ if i == rank:
+ tmp = dist.reduce(tensor.contiguous(), ranks[i], op, group, async_op=True)
+ else:
+ tmp = dist.reduce(tensor_list[i].contiguous(), ranks[i], op, group, async_op=True)
+ ops.append(tmp)
+
+ oplist = AsyncOpList(ops)
+ if async_op:
+ return oplist
+ else:
+ oplist.wait()
diff --git a/grn/tokenizer/videovae/utils/diffdist/functional.py b/grn/tokenizer/videovae/utils/diffdist/functional.py
new file mode 100644
index 0000000000000000000000000000000000000000..2c8dd75fdea7cebe28e169341a75faf51baa0fec
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/functional.py
@@ -0,0 +1,55 @@
+import videovae.utils.diffdist.modules as mods
+import torch.distributed as dist
+
+
+def consume_variable(tensor_to_consume, tensors_to_return, set_ones_grad=True):
+ return mods.ConsumeVariable(set_ones_grad)(tensor_to_consume,
+ *tensors_to_return)
+
+
+def send(tensor, dst, group=dist.group.WORLD, tag=0):
+ return mods.Send(dst, group, tag)(tensor)
+
+
+def recv(tensor,
+ src=None,
+ group=dist.group.WORLD,
+ tag=0,
+ next_backprop=None,
+ inplace=True):
+ return mods.Recv(src, group, tag, next_backprop, inplace)(tensor)
+
+
+def broadcast(tensor,
+ src,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ return mods.Broadcast(src, group, next_backprop, inplace)(tensor)
+
+
+def gather(tensor,
+ gather_list=None,
+ dst=None,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ return mods.Gather(dst, group, next_backprop, inplace)(tensor, gather_list)
+
+
+def scatter(tensor,
+ scatter_list=None,
+ src=None,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ return mods.Scatter(src, group, next_backprop, inplace)(tensor,
+ scatter_list)
+
+
+def all_gather(gather_list,
+ tensor,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ return mods.AllGather(group, next_backprop, inplace)(gather_list, tensor)
diff --git a/grn/tokenizer/videovae/utils/diffdist/functions.py b/grn/tokenizer/videovae/utils/diffdist/functions.py
new file mode 100644
index 0000000000000000000000000000000000000000..f65a6fac72d72f0261f6afcb293cfb6546f3e67a
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/functions.py
@@ -0,0 +1,196 @@
+from torch.autograd import Function
+import videovae.utils.diffdist.extra_collectives as dist_extra
+import torch.distributed as dist
+import torch
+
+
+class ConsumeVariableFunc(Function):
+ @staticmethod
+ def forward(ctx, tensor_to_consume, set_ones_grad, *tensors_to_return):
+ ctx.save_for_backward(tensor_to_consume)
+ ctx.set_ones_grad = set_ones_grad
+ return tensors_to_return
+
+ @staticmethod
+ def backward(ctx, *grad_outputs):
+ tensor_to_consume, = ctx.saved_tensors
+ if ctx.set_ones_grad:
+ fake_grad = torch.ones_like(tensor_to_consume)
+ else:
+ fake_grad = torch.zeros_like(tensor_to_consume)
+
+ return (fake_grad, None) + grad_outputs
+
+
+class SendFunc(Function):
+ @staticmethod
+ def forward(ctx, tensor, dst, group=dist.group.WORLD, tag=0):
+ ctx.save_for_backward(tensor)
+ ctx.dst = dst
+ ctx.group = group
+ ctx.tag = tag
+ dist.send(tensor, dst, group, tag)
+ return tensor.new_tensor([])
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ tensor, = ctx.saved_tensors
+ # TODO: Add ctx.needs_input_grad check
+ grad_tensor = torch.zeros_like(tensor)
+ dist.recv(grad_tensor, ctx.dst, ctx.group, ctx.tag)
+
+ return grad_tensor, None, None, None
+
+
+class RecvFunc(Function):
+ @staticmethod
+ def forward(ctx,
+ tensor,
+ src=None,
+ group=dist.group.WORLD,
+ tag=0,
+ inplace=True):
+ if not inplace:
+ tensor = torch.zeros_like(tensor).requires_grad_(False)
+ ctx.src = src
+ ctx.group = group
+ ctx.tag = tag
+ sender = dist.recv(tensor, src, group, tag)
+ if src:
+ assert sender == src
+ else:
+ ctx.src = sender
+ sender = torch.tensor(sender)
+ ctx.mark_non_differentiable(sender)
+ return tensor, sender
+
+ @staticmethod
+ def backward(ctx, grad_tensor, grad_sender):
+ dist.send(grad_tensor, ctx.src, ctx.group, ctx.tag)
+ return grad_tensor, None, None, None, None
+
+
+class BroadcastFunc(Function):
+ @staticmethod
+ def forward(ctx, tensor, src, group=dist.group.WORLD, inplace=True):
+ ctx.src = src
+ ctx.group = group
+ if dist.get_rank(group) == src:
+ if not inplace:
+ with torch.no_grad():
+ tensor = tensor.clone().requires_grad_(False)
+ else:
+ if not inplace:
+ tensor = torch.zeros_like(tensor).requires_grad_(False)
+ dist.broadcast(tensor, src, group)
+ return tensor
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ dist.reduce(grad_output,
+ ctx.src,
+ op=dist.ReduceOp.SUM,
+ group=ctx.group)
+ return grad_output, None, None, None
+
+
+class AllReduceFunc(Function):
+ @staticmethod
+ def forward(ctx, i):
+ raise NotImplementedError
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ raise NotImplementedError
+
+
+class ReduceFunc(Function):
+ @staticmethod
+ def forward(ctx, i):
+ raise NotImplementedError
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ raise NotImplementedError
+
+
+class AllGatherFunc(Function):
+ @staticmethod
+ def forward(ctx, tensor, group, inplace, *gather_list):
+ ctx.save_for_backward(tensor)
+ ctx.group = group
+ gather_list = list(gather_list)
+ if not inplace:
+ gather_list = [torch.zeros_like(g) for g in gather_list]
+ dist.all_gather(gather_list, tensor, group)
+ return tuple(gather_list)
+
+ @staticmethod
+ def backward(ctx, *grads):
+ input, = ctx.saved_tensors
+ grad_out = torch.zeros_like(input)
+ dist_extra.reduce_scatter(grad_out, list(grads), group=ctx.group)
+ return (grad_out, None, None) + grads
+
+
+class GatherFunc(Function):
+ @staticmethod
+ def forward(ctx, input, dst, group, inplace, *gather_list):
+ ctx.dst = dst
+ ctx.group = group
+ ctx.save_for_backward(input)
+ if dist.get_rank(group) == dst:
+ gather_list = list(gather_list)
+ if not inplace:
+ gather_list = [torch.zeros_like(g) for g in gather_list]
+ dist.gather(input, gather_list=gather_list, dst=dst, group=group)
+ return tuple(gather_list)
+ else:
+ dist.gather(input, [], dst=dst, group=group)
+ return input.new_tensor([])
+
+ @staticmethod
+ def backward(ctx, *grads):
+ input, = ctx.saved_tensors
+ grad_input = torch.zeros_like(input)
+ if dist.get_rank(ctx.group) == ctx.dst:
+ grad_outputs = list(grads)
+ dist.scatter(grad_input,
+ grad_outputs,
+ src=ctx.dst,
+ group=ctx.group)
+ return (grad_input, None, None, None) + grads
+ else:
+ dist.scatter(grad_input, [], src=ctx.dst, group=ctx.group)
+ return grad_input, None, None, None, None
+
+
+class ScatterFunc(Function):
+ @staticmethod
+ def forward(ctx,
+ tensor,
+ src,
+ group=dist.group.WORLD,
+ inplace=True,
+ *scatter_list):
+ ctx.src = src
+ ctx.group = group
+ if not inplace:
+ tensor = torch.zeros_like(tensor)
+ if dist.get_rank(group) == src:
+ ctx.save_for_backward(*scatter_list)
+ scatter_list = list(scatter_list)
+ dist.scatter(tensor, scatter_list, src=src, group=group)
+ else:
+ dist.scatter(tensor, [], src=src, group=group)
+ return tensor
+
+ @staticmethod
+ def backward(ctx, grad_tensor):
+ if dist.get_rank(ctx.group) == ctx.src:
+ grad_outputs = [torch.zeros_like(g) for g in ctx.saved_tensors]
+ dist.gather(grad_tensor, grad_outputs, ctx.src, group=ctx.group)
+ return (grad_tensor, None, None, None) + tuple(grad_outputs)
+ else:
+ dist.gather(grad_tensor, [], ctx.src, group=ctx.group)
+ return grad_tensor, None, None, None, None
diff --git a/grn/tokenizer/videovae/utils/diffdist/modules.py b/grn/tokenizer/videovae/utils/diffdist/modules.py
new file mode 100644
index 0000000000000000000000000000000000000000..3c810a0d1b22f0c18404f99b50fac9e9516a626a
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/modules.py
@@ -0,0 +1,155 @@
+import torch.nn as nn
+import torch.distributed as dist
+import videovae.utils.diffdist.functions as funcs
+
+
+class ConsumeVariable(nn.Module):
+ def __init__(self, set_ones_grad=True):
+ """
+ If set_ones_grad=True then the gradient w.r.t tensor_to_consume
+ is set to 1 during backprop. Otherwise, it is set to 0.
+ """
+ super(ConsumeVariable, self).__init__()
+ self.set_ones_grad = set_ones_grad
+
+ def forward(self, tensor_to_consume, *tensors_to_return):
+ tensors_to_return = funcs.ConsumeVariableFunc.apply(
+ tensor_to_consume, self.set_ones_grad, *tensors_to_return)
+ return tensors_to_return
+
+
+class Send(nn.Module):
+ def __init__(self, dst, group=dist.group.WORLD, tag=0):
+ super(Send, self).__init__()
+ self.dst = dst
+ self.group = group
+ self.tag = tag
+
+ def forward(self, tensor):
+ return funcs.SendFunc.apply(tensor, self.dst, self.group, self.tag)
+
+
+class Recv(nn.Module):
+ def __init__(self,
+ src=None,
+ group=dist.group.WORLD,
+ tag=0,
+ next_backprop=None,
+ inplace=True):
+ super(Recv, self).__init__()
+ self.next_backprop = next_backprop
+ self.src = src
+ self.group = group
+ self.tag = tag
+ self.inplace = inplace
+
+ self.consume = None
+ if self.next_backprop is not None:
+ self.consume = ConsumeVariable()
+
+ def forward(self, tensor):
+ if self.consume:
+ tensor, = self.consume(self.next_backprop, tensor)
+ tensor, sender = funcs.RecvFunc.apply(tensor, self.src, self.group,
+ self.tag, self.inplace)
+ return tensor, sender.item()
+
+
+class Broadcast(nn.Module):
+ def __init__(self,
+ src,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ super(Broadcast, self).__init__()
+ self.src = src
+ self.group = group
+ self.next_backprop = next_backprop
+ self.inplace = inplace
+
+ self.consume = None
+ if self.next_backprop is not None:
+ self.consume = ConsumeVariable()
+
+ def forward(self, tensor):
+ if self.consume:
+ tensor, = self.consume(self.next_backprop, tensor)
+ return funcs.BroadcastFunc.apply(tensor, self.src, self.group,
+ self.inplace)
+
+
+class Gather(nn.Module):
+ def __init__(self,
+ dst=None,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ super(Gather, self).__init__()
+ self.dst = dst
+ self.group = group
+ self.next_backprop = next_backprop
+ self.inplace = inplace
+
+ self.consume = None
+ if self.next_backprop is not None:
+ self.consume = ConsumeVariable()
+
+ def forward(self, tensor, gather_list=None):
+ if self.consume:
+ tensor, = self.consume(self.next_backprop, tensor)
+ if dist.get_rank(self.group) == self.dst:
+ return list(
+ funcs.GatherFunc.apply(tensor, self.dst, self.group,
+ self.inplace, *gather_list))
+ else:
+ return funcs.GatherFunc.apply(tensor, self.dst, self.group,
+ self.inplace, None)
+
+
+class Scatter(nn.Module):
+ def __init__(self,
+ src=None,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ super(Scatter, self).__init__()
+ self.src = src
+ self.group = group
+ self.next_backprop = next_backprop
+ self.inplace = inplace
+
+ self.consume = None
+ if self.next_backprop is not None:
+ self.consume = ConsumeVariable()
+
+ def forward(self, tensor, scatter_list=None):
+ if self.consume:
+ tensor, = self.consume(self.next_backprop, tensor)
+ if dist.get_rank(self.group) == self.src:
+ return funcs.ScatterFunc.apply(tensor, self.src, self.group,
+ self.inplace, *scatter_list)
+ else:
+ return funcs.ScatterFunc.apply(tensor, self.src, self.group,
+ self.inplace, None)
+
+
+class AllGather(nn.Module):
+ def __init__(self,
+ group=dist.group.WORLD,
+ next_backprop=None,
+ inplace=True):
+ super(AllGather, self).__init__()
+ self.group = group
+ self.next_backprop = next_backprop
+ self.inplace = inplace
+
+ self.consume = None
+ if self.next_backprop is not None:
+ self.consume = ConsumeVariable()
+
+ def forward(self, gather_list, tensor):
+ if self.consume:
+ tensor, = self.consume(self.next_backprop, tensor)
+ return list(
+ funcs.AllGatherFunc.apply(tensor, self.group, self.inplace,
+ *gather_list))
diff --git a/grn/tokenizer/videovae/utils/diffdist/testing.py b/grn/tokenizer/videovae/utils/diffdist/testing.py
new file mode 100644
index 0000000000000000000000000000000000000000..9151ff92b584420bcf842feb40cc0a219dcf9f77
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/diffdist/testing.py
@@ -0,0 +1,273 @@
+import videovae.utils.diffdist.functional as distops
+import torch.distributed as dist
+import torch
+import videovae.utils.diffdist.extra_collectives as extra_comm
+
+
+def test_reduce_scatter():
+ if dist.get_rank() == 0:
+ print("REDUCE_SCATTER TEST\n")
+ x = torch.arange(dist.get_world_size()).float().split(1)
+ buff = torch.tensor(0.)
+ extra_comm.reduce_scatter(buff, x)
+ print(dist.get_rank(), x)
+ print(dist.get_rank(), buff)
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print('-' * 50)
+
+
+def test_all_gather():
+ if dist.get_rank() == 0:
+ print("ALL GATHER TEST\n")
+ dist.barrier()
+ x = torch.tensor(3., requires_grad=True)
+ y = (dist.get_rank() + 1) * x
+
+ print(dist.get_rank(), "Sending y:", y)
+ z = distops.all_gather(list(torch.zeros(dist.get_world_size())),
+ y,
+ next_backprop=None,
+ inplace=True)
+ print(dist.get_rank(), "Received tensor:", z)
+ l = torch.sum(torch.stack(z))
+ l = l * (dist.get_rank() + 1)
+ l.backward()
+
+ print(dist.get_rank(), "Gradient with MPI:", x.grad)
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print()
+ x = [
+ torch.tensor(3., requires_grad=True)
+ for i in range(dist.get_world_size())
+ ]
+ res = []
+ for i in range(1, dist.get_world_size() + 1):
+ res.append(i * x[i - 1])
+
+ res2 = []
+ for i in range(dist.get_world_size()):
+ temp = []
+ for j in range(dist.get_world_size()):
+ temp.append(torch.clone(res[j]))
+ res2.append(temp)
+ l_s = [torch.sum(torch.stack(i)) for i in res2]
+ final = [(i + 1) * k for i, k in enumerate(l_s)]
+ for i in range(dist.get_world_size() - 1):
+ final[i].backward(retain_graph=True)
+ final[-1].backward()
+ for i, x_i in enumerate(x):
+ print(i, "Gradient in single process:", x_i.grad)
+ print('-' * 50)
+
+
+def test_scatter():
+ if dist.get_rank() == 0:
+ print("SCATTER TEST\n")
+ x = [
+ torch.tensor(3., requires_grad=True)
+ for i in range(dist.get_world_size())
+ ]
+ y = [2 * x_i for x_i in x]
+
+ print("Sending y:", y)
+ buffer = torch.tensor(0.)
+ z = distops.scatter(buffer, y, src=0, inplace=False)
+ else:
+ buffer = torch.tensor(0., requires_grad=True)
+ z = distops.scatter(buffer, src=0, inplace=False)
+
+ print(dist.get_rank(), "Received tensor:", z)
+ # Computation
+ k = (dist.get_rank() + 1) * z
+ k.backward()
+
+ if dist.get_rank() == 0:
+ print("Gradient with MPI:", [x_i.grad for x_i in x])
+
+ if dist.get_rank() == 0:
+ print()
+ x = [
+ torch.tensor(3., requires_grad=True)
+ for i in range(dist.get_world_size())
+ ]
+ y = [2 * x_i for x_i in x]
+ res = []
+ for i in range(dist.get_world_size()):
+ res.append((i + 1) * y[i])
+
+ for i, k in enumerate(res):
+ k.backward()
+ print("Gradient in single process:", [x_i.grad for x_i in x])
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print('-' * 50)
+
+
+def test_gather():
+ if dist.get_rank() == 0:
+ print("GATHER TEST\n")
+ dist.barrier()
+ x = torch.tensor(3., requires_grad=True)
+ y = (dist.get_rank() + 1) * x
+
+ print(dist.get_rank(), "Sending y:", y)
+ if dist.get_rank() == 0:
+ z = distops.gather(y,
+ torch.zeros(dist.get_world_size()).split(1),
+ dst=0,
+ next_backprop=None,
+ inplace=True)
+ print(dist.get_rank(), "Received tensor:", z)
+ l = torch.sum(torch.stack(z))
+ l.backward()
+ else:
+ dummy = distops.gather(y, dst=0, next_backprop=None, inplace=True)
+ dummy.backward(torch.tensor([]))
+ print(dist.get_rank(), "Gradient with MPI:", x.grad)
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print()
+ x = [
+ torch.tensor(3., requires_grad=True)
+ for i in range(dist.get_world_size())
+ ]
+ res = []
+ for i in range(1, dist.get_world_size() + 1):
+ res.append(i * x[i - 1])
+
+ z = torch.stack(res)
+ l = torch.sum(z)
+ l.backward()
+ for i, x_i in enumerate(x):
+ print(i, "Gradient in single process:", x_i.grad)
+ print('-' * 50)
+
+
+def test_broadcast():
+ if dist.get_rank() == 0:
+ print("BROADCAST TEST\n")
+ x = torch.tensor(3., requires_grad=True)
+ y = 2 * x
+
+ print(dist.get_rank(), "Sending y:", y)
+ z = distops.broadcast(y, src=0, inplace=False)
+ print(dist.get_rank(), "Received tensor:", z)
+
+ # Computation
+ k = 3 * z
+ k.backward()
+ print("Gradient with MPI:", x.grad)
+
+ print()
+ x = torch.tensor(3., requires_grad=True)
+ y = 2 * x
+ res = [3 * y]
+ for i in range(1, dist.get_world_size()):
+ res.append(9 * y)
+
+ for i, k in enumerate(res):
+ if i == (len(res) - 1):
+ k.backward()
+ else:
+ k.backward(retain_graph=True)
+ print("Gradient in single process:", x.grad)
+ else:
+ x = torch.tensor(5., requires_grad=True)
+ y = 7 * x
+
+ buffer = torch.tensor(0.)
+ z = distops.broadcast(buffer, src=0, next_backprop=y)
+ print(dist.get_rank(), "Received tensor:", z)
+ k = 9 * z
+ k.backward()
+ print(dist.get_rank(), "Grad of disconnected part:", x.grad)
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print('-' * 50)
+
+
+def test_consume_variable():
+ x = torch.tensor(5., requires_grad=True)
+ y = 2 * x
+
+ z = 3 * y
+ j = 4 * y
+
+ z = distops.consume_variable(j, [z], set_ones_grad=True)[0]
+ print(z)
+ z.backward()
+ print(x.grad)
+ print()
+ x = torch.tensor(5., requires_grad=True)
+ y = 2 * x
+
+ z = 3 * y
+ j = 4 * y
+
+ z.backward(retain_graph=True)
+ j.backward()
+ print(x.grad)
+
+
+def test_send_recv():
+ if dist.get_rank() == 0:
+ print("SEND/RECV TEST\n")
+ x = torch.tensor(3., requires_grad=True)
+ y = 2 * x
+
+ print("Before sending y:", y)
+ connector = distops.send(y, dst=1)
+ # Computation happens in process 1
+ buffer = torch.tensor(0.)
+ z, _ = distops.recv(buffer, src=1, next_backprop=connector)
+ print("After receiving:", z)
+
+ k = 3 * z
+ k.backward()
+ print("Gradient with MPI:", x.grad)
+
+ print()
+ x = torch.tensor(3., requires_grad=True)
+ y = 2 * x
+ l = y * 10
+ k = 3 * l
+ k.backward()
+ print("Gradient in single process:", x.grad)
+ print('-' * 50)
+ elif dist.get_rank() == 1:
+ buffer = torch.tensor(0., requires_grad=True)
+ y, _ = distops.recv(buffer, src=0)
+
+ l = y * 10
+
+ connector = distops.send(l, dst=0)
+ connector.backward(torch.tensor([]))
+
+
+if __name__ == '__main__':
+ dist.init_process_group('mpi')
+
+ print(f'I am {dist.get_rank()}')
+ dist.barrier()
+ if dist.get_rank() == 0:
+ print('-' * 50)
+
+ if dist.get_rank() == 0:
+ print("EXTRA COLLECTIVES")
+
+ test_reduce_scatter()
+
+ if dist.get_rank() == 0:
+ print('-' * 50)
+
+ test_send_recv()
+
+ test_broadcast()
+
+ test_gather()
+
+ test_scatter()
+
+ test_all_gather()
diff --git a/grn/tokenizer/videovae/utils/distributed.py b/grn/tokenizer/videovae/utils/distributed.py
new file mode 100644
index 0000000000000000000000000000000000000000..4f64666e83745ae1ecd2dc7c4f913a115fa61078
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/distributed.py
@@ -0,0 +1,153 @@
+# from https://github.com/FoundationVision/LlamaGen/blob/main/utils/distributed.py
+import os
+import sys
+import glob
+import torch
+import subprocess
+import torch.distributed as dist
+import datetime
+import logging
+import builtins
+
+from videovae.utils.misc import rank_zero_only, COLOR_BLUE, COLOR_RESET
+
+from torch.distributed.fsdp.wrap import ModuleWrapPolicy
+from torch.distributed.fsdp import (
+ FullyShardedDataParallel as FSDP,
+ ShardingStrategy,
+ MixedPrecision,
+)
+
+
+def setup_for_distributed(is_master, logging_dir=""):
+ builtin_print = builtins.print
+ def print(*args, **kwargs):
+ if is_master:
+ builtin_print(*args, **kwargs)
+ builtins.print = print
+
+ if is_master:
+ os.makedirs(logging_dir, exist_ok=True)
+ existing_logs = glob.glob(os.path.join(logging_dir, 'log_out_*.txt'))
+ log_numbers = [int(log.split('.txt')[0].split('_')[-1]) for log in existing_logs]
+ next_log_number = max(log_numbers) + 1 if log_numbers else 1
+
+ log_out_path = os.path.join(logging_dir, f'log_out_{next_log_number}.txt')
+ log_err_path = os.path.join(logging_dir, f'log_err_{next_log_number}.txt')
+
+ print(f"{COLOR_BLUE}stdout will be written to {log_out_path}{COLOR_RESET}")
+ print(f"{COLOR_BLUE}stderr will be written to {log_err_path}{COLOR_RESET}")
+
+ # sys.stdout = Tee(sys.stdout, open(log_out_path, 'w'))
+ # sys.stderr = Tee(sys.stderr, open(log_err_path, 'w'))
+
+class Tee(object):
+ def __init__(self, *files):
+ self.files = files
+
+ def write(self, obj):
+ for f in self.files:
+ f.write(obj)
+ f.flush()
+
+ def flush(self):
+ for f in self.files:
+ f.flush()
+
+def init_distributed_mode(args, timeout_minutes=15):
+ if 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
+ args.rank = int(os.environ["RANK"])
+ args.world_size = int(os.environ['WORLD_SIZE'])
+ args.gpu = int(os.environ['LOCAL_RANK'])
+ args.dist_url = 'env://'
+ os.environ['LOCAL_SIZE'] = str(torch.cuda.device_count())
+ elif 'SLURM_PROCID' in os.environ:
+ proc_id = int(os.environ['SLURM_PROCID'])
+ ntasks = int(os.environ['SLURM_NTASKS'])
+ node_list = os.environ['SLURM_NODELIST']
+ num_gpus = torch.cuda.device_count()
+ addr = subprocess.getoutput(
+ 'scontrol show hostname {} | head -n1'.format(node_list))
+ os.environ['MASTER_PORT'] = os.environ.get('MASTER_PORT', '29500')
+ os.environ['MASTER_ADDR'] = addr
+ os.environ['WORLD_SIZE'] = str(ntasks)
+ os.environ['RANK'] = str(proc_id)
+ os.environ['LOCAL_RANK'] = str(proc_id % num_gpus)
+ os.environ['LOCAL_SIZE'] = str(num_gpus)
+ args.dist_url = 'env://'
+ args.world_size = ntasks
+ args.rank = proc_id
+ args.gpu = proc_id % num_gpus
+ else:
+ print('Not using distributed mode')
+ args.distributed = False
+ return
+
+ args.distributed = True
+
+ torch.cuda.set_device(args.gpu)
+ args.dist_backend = 'nccl'
+ print('| distributed init (rank {}): {}'.format(
+ args.rank, args.dist_url), flush=True)
+ torch.distributed.init_process_group(backend=args.dist_backend, init_method=args.dist_url,
+ world_size=args.world_size, rank=args.rank,
+ timeout=datetime.timedelta(seconds=timeout_minutes * 60)
+ )
+ torch.distributed.barrier()
+ setup_for_distributed(args.rank == 0, args.default_root_dir)
+
+def _FSDP(model: torch.nn.Module, device, zero) -> FSDP:
+ def my_policy(
+ module: torch.nn.Module,
+ recurse: bool,
+ **kwargs,
+ ) -> bool:
+ return True
+ auto_wrap_policy = my_policy
+ model = FSDP(
+ model,
+ auto_wrap_policy=auto_wrap_policy,
+ device_id=device,
+ sharding_strategy={1:ShardingStrategy.HYBRID_SHARD, 2:ShardingStrategy.SHARD_GRAD_OP, 3:ShardingStrategy.FULL_SHARD}.get(zero),
+ mixed_precision=MixedPrecision(
+ param_dtype=torch.float,
+ reduce_dtype=torch.float,
+ buffer_dtype=torch.float,
+ ),
+ sync_module_states=True,
+ limit_all_gathers=True,
+ use_orig_params=True,
+ )
+ torch.cuda.synchronize()
+ return model
+
+
+def reduce_losses(loss_dict, dst=0):
+ loss_names = list(loss_dict.keys())
+ loss_tensor = torch.stack([loss_dict[name] for name in loss_names])
+
+ dist.reduce(loss_tensor, dst=dst, op=dist.ReduceOp.SUM)
+ # Only average the loss values on the destination rank
+ if dist.get_rank() == dst:
+ loss_tensor /= dist.get_world_size()
+ averaged_losses = {name: loss_tensor[i].item() for i, name in enumerate(loss_names)}
+ else:
+ averaged_losses = {name: None for name in loss_names}
+
+ return averaged_losses
+
+@rank_zero_only
+def average_losses(loss_dict_list):
+ sum_dict = {}
+ count_dict = {}
+ for loss_dict in loss_dict_list:
+ for key, value in loss_dict.items():
+ if key in sum_dict:
+ sum_dict[key] += value
+ count_dict[key] += 1
+ else:
+ sum_dict[key] = value
+ count_dict[key] = 1
+
+ avg_dict = {key: sum_dict[key] / count_dict[key] for key in sum_dict}
+ return avg_dict
diff --git a/grn/tokenizer/videovae/utils/dynamic_resolution.py b/grn/tokenizer/videovae/utils/dynamic_resolution.py
new file mode 100644
index 0000000000000000000000000000000000000000..ec8c82486974d3af0c57df140d1555c8beb5823f
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/dynamic_resolution.py
@@ -0,0 +1,37 @@
+import json
+import numpy as np
+import tqdm
+
+vae_stride = 16
+ratio2hws = {
+ 1.000: [(1,1),(2,2),(4,4),(6,6),(8,8),(12,12),(16,16),(20,20),(24,24),(32,32),(40,40),(48,48),(64,64),(80,80),(96,96),(128,128)],
+ 1.250: [(1,1),(2,2),(3,3),(5,4),(10,8),(15,12),(20,16),(25,20),(30,24),(35,28),(45,36),(55,44),(70,56),(90,72),(110,88),(140,112)],
+ 1.333: [(1,1),(2,2),(4,3),(8,6),(12,9),(16,12),(20,15),(24,18),(28,21),(36,27),(48,36),(60,45),(72,54),(96,72),(120,90),(144,108)],
+ 1.500: [(1,1),(2,2),(3,2),(6,4),(9,6),(15,10),(21,14),(27,18),(33,22),(39,26),(48,32),(63,42),(78,52),(96,64),(126,84),(156,104)],
+ 1.750: [(1,1),(2,2),(3,3),(7,4),(11,6),(14,8),(21,12),(28,16),(35,20),(42,24),(56,32),(70,40),(84,48),(112,64),(140,80),(168,96)],
+ 2.000: [(1,1),(2,2),(4,2),(6,3),(10,5),(16,8),(22,11),(30,15),(38,19),(46,23),(60,30),(74,37),(90,45),(120,60),(148,74),(180,90)],
+ 2.500: [(1,1),(2,2),(5,2),(10,4),(15,6),(20,8),(25,10),(30,12),(40,16),(50,20),(65,26),(80,32),(100,40),(130,52),(160,64),(200,80)],
+ 3.000: [(1,1),(2,2),(6,2),(9,3),(15,5),(21,7),(27,9),(36,12),(45,15),(54,18),(72,24),(90,30),(111,37),(144,48),(180,60),(222,74)],
+}
+full_ratio2hws = {}
+for ratio, hws in ratio2hws.items():
+ full_ratio2hws[ratio] = hws
+ full_ratio2hws[int(1/ratio*1000)/1000] = [(item[1], item[0]) for item in hws]
+
+dynamic_resolution_h_w = {}
+predefined_HW_Scales_dynamic = {}
+aspect_ratio_scale_list = []
+bs_dict = {7: 8, 10: 4, 13: 1, 16: 1} # 256x256: batch=8, 512x512: batch=4, 1024x1024: batch=1 (bs=1 avoid OOM)
+for ratio in full_ratio2hws:
+ dynamic_resolution_h_w[ratio] ={}
+ for ind, leng in enumerate([7, 10, 13, 16]):
+ h, w = full_ratio2hws[ratio][leng-1][0], full_ratio2hws[ratio][leng-1][1] # feature map size
+ pixel = (h * vae_stride, w * vae_stride) # The original image (H, W)
+ dynamic_resolution_h_w[ratio][pixel[1]] = {
+ 'pixel': pixel,
+ 'scales': full_ratio2hws[ratio][:leng]
+ } # W as key
+ predefined_HW_Scales_dynamic[(h, w)] = full_ratio2hws[ratio][:leng]
+ # deal with aspect_ratio_scale_list
+ info_dict = {"ratio": ratio, "h": h * vae_stride, "w": w * vae_stride, "bs": bs_dict[leng]}
+ aspect_ratio_scale_list.append(info_dict)
diff --git a/grn/tokenizer/videovae/utils/dynamic_resolution_two_pyramid.py b/grn/tokenizer/videovae/utils/dynamic_resolution_two_pyramid.py
new file mode 100644
index 0000000000000000000000000000000000000000..c9f8a69bd725bbe3697e2fa758a549fd1d59c752
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/dynamic_resolution_two_pyramid.py
@@ -0,0 +1,176 @@
+import os
+import math
+
+import numpy as np
+
+video_frames = 97
+vae_stride = 16
+compressed_frames = video_frames // 4 + 1
+spatial_compress_rate = os.environ['spatial_compress_rate']
+
+def append_dummy_t(ratio2hws):
+ for key in ratio2hws:
+ for i in range(len(ratio2hws[key])):
+ h, w = ratio2hws[key][i]
+ ratio2hws[key][i] = (1, h, w)
+ return ratio2hws
+
+def get_full_ratio2hws(ratio2hws, total_pixels2scales):
+ full_ratio2hws = {}
+ for ratio, hws in ratio2hws.items():
+ real_ratio = hws[-1][1] / hws[-1][2]
+ full_ratio2hws[int(real_ratio*1000)/1000] = hws
+ if ratio != 1.000:
+ full_ratio2hws[int(1/real_ratio*1000)/1000] = [(item[0], item[2], item[1]) for item in hws]
+ if 'heritage_' in spatial_compress_rate:
+ dynamic_resolution_h_w = {}
+ for ratio in full_ratio2hws:
+ dynamic_resolution_h_w[ratio] = {}
+ for pn, scale_times in total_pixels2scales.items():
+ h, w = full_ratio2hws[ratio][-1][1] * scale_times, full_ratio2hws[ratio][-1][2] * scale_times
+ # pixel = (h * vae_stride, w * vae_stride)
+ scales = [(pt,scale_times*ph, scale_times*pw) for (pt, ph, pw) in full_ratio2hws[ratio]]
+ dynamic_resolution_h_w[ratio][(h, w)] = {
+ 'scales': scales,
+ 'pn': pn,
+ 'h_div_w': ratio,
+ }
+ else:
+ dynamic_resolution_h_w = {}
+ for ratio in full_ratio2hws:
+ dynamic_resolution_h_w[ratio] = {}
+ for pn, scales_num in total_pixels2scales.items():
+ h, w = full_ratio2hws[ratio][scales_num-1][1], full_ratio2hws[ratio][scales_num-1][2]
+ # pixel = (h * vae_stride, w * vae_stride)
+ scales = full_ratio2hws[ratio][:scales_num]
+ dynamic_resolution_h_w[ratio][(h, w)] = {
+ 'scales': scales,
+ 'pn': pn,
+ 'h_div_w': ratio,
+ }
+ return dynamic_resolution_h_w
+
+# ratio2hws = {
+# 1.000: [(1,1),(2,2),(3,3),(4,4),(5,5),(6,6),(7,7),(8,8),(10,10),(12,12),(16,16),(24,24),(32,32),(48,48),(60,60),(64,64)],
+# 1.250: [(1,1),(2,2),(3,3),(4,3),(5,4),(6,5),(7,5),(8,6),(10,8),(15,12),(20,16),(30,24),(35,28),(45,36),(66,52),(70,56)],
+# 1.333: [(1,1),(2,2),(3,2),(4,3),(5,4),(6,5),(7,5),(8,6),(12,9),(16,12),(20,15),(28,21),(36,27),(48,36),(68,50),(72,54)],
+# 1.500: [(1,1),(2,2),(3,2),(4,3),(5,3),(6,4),(7,4),(8,6),(12,8),(15,10),(21,14),(33,22),(39,26),(48,32),(72,48),(78,52)],
+# 1.750: [(1,1),(2,2),(3,3),(4,3),(5,3),(6,4),(7,4),(8,5),(12,7),(14,8),(21,12),(32,18),(42,24),(54,30),(80,45),(84,48)],
+# 2.000: [(1,1),(2,2),(3,2),(4,2),(5,3),(6,3),(7,4),(8,4),(12,6),(16,8),(22,11),(38,19),(46,23),(60,30),(82,41),(90,45)],
+# 2.500: [(1,1),(2,2),(3,2),(4,2),(5,2),(7,3),(8,3),(10,4),(15,6),(20,8),(25,10),(40,16),(50,20),(65,26),(90,36),(100,40)],
+# 3.000: [(1,1),(2,2),(3,2),(4,2),(5,2),(6,2),(8,3),(9,3),(15,5),(21,7),(27,9),(45,15),(54,18),(72,24),(96,32),(111,37)],
+# }
+# total_pixels2scales = {
+# '0.06M': 11,
+# '0.15M': 13,
+# '0.25M': 13,
+# '0.40M': 14,
+# '0.90M': 15,
+# '1M': 16,
+# }
+
+def get_ratio2hws_video_v2():
+ import os
+ if spatial_compress_rate == '16':
+ ratio2hws_video_common_v2 = {}
+ for h_div_w in [1, 100/116, 3/4, 2/3, 9/16, 1/2, 2/5, 1/3]:
+ scale_schedule = []
+ # 48*48 is 480p, 60*60 is 720p
+ for scale in [1,2,3,4,5,6,7,8,10,12,16] + [24, 32, 40, 48, 60, 64]:
+ # for scale in [4,8,12,20]:
+ area = scale * scale
+ pw_float = math.sqrt(area / h_div_w)
+ ph_float = pw_float * h_div_w
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
+ scale_schedule.append((ph, pw))
+ ratio2hws_video_common_v2[h_div_w] = scale_schedule
+ total_pixels2scales = {
+ # '16K': 8,
+ '0.06M': 11,
+ '0.25M': 13,
+ '0.40M': 14,
+ # '0.60M': 15,
+ '0.92M': 16,
+ '1M': 17,
+ }
+ elif spatial_compress_rate == '32':
+ ratio2hws_video_common_v2 = {}
+ for h_div_w in [1, 100/116, 3/4, 2/3, 9/16, 1/2, 2/5, 1/3]:
+ scale_schedule = []
+ # 48*48 is 480p, 60*60 is 720p
+ # for scale in list(range(1,1+16)) + [20, 24, 30, 40]:
+ for scale in [1,2,3,4,5,6,7,8] + [10,12,16,20,25,30]:
+ # for scale in [4,8,12,20]:
+ area = scale * scale
+ pw_float = math.sqrt(area / h_div_w)
+ ph_float = pw_float * h_div_w
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
+ scale_schedule.append((ph, pw))
+ ratio2hws_video_common_v2[h_div_w] = scale_schedule
+ total_pixels2scales = {
+ '16K': 4,
+ '0.06M': 8,
+ '0.40M': 12,
+ }
+ elif spatial_compress_rate == 'heritage_schedule_stride_32':
+ ratio2hws_video_common_v2 = {}
+ for h_div_w in [1, 100/116, 3/4, 2/3, 9/16, 1/2, 2/5, 1/3]:
+ scale_schedule = []
+ # 48*48 is 480p, 60*60 is 720p
+ # for scale in list(range(1,1+16)) + [20, 24, 30, 40]:
+ for scale in [1,2,3,4,5,6,7,8]:
+ # for scale in [4,8,12,20]:
+ area = scale * scale
+ pw_float = math.sqrt(area / h_div_w)
+ ph_float = pw_float * h_div_w
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
+ scale_schedule.append((ph, pw))
+ ratio2hws_video_common_v2[h_div_w] = scale_schedule
+ total_pixels2scales = {
+ '0.06M': 1,
+ '0.25M': 2,
+ '1M': 4,
+ }
+ elif spatial_compress_rate == 'heritage_schedule_stride_8':
+ ratio2hws_video_common_v2 = {}
+ for h_div_w in [1, 100/116, 3/4, 2/3, 9/16, 1/2, 2/5, 1/3]:
+ scale_schedule = []
+ # 48*48 is 480p, 60*60 is 720p
+ # for scale in list(range(1,1+16)) + [20, 24, 30, 40]:
+ for scale in [1,2,3,4,5,6,7,8]:
+ # for scale in [4,8,12,20]:
+ area = scale * scale
+ pw_float = math.sqrt(area / h_div_w)
+ ph_float = pw_float * h_div_w
+ ph, pw = int(np.round(ph_float)), int(np.round(pw_float))
+ scale_schedule.append((ph, pw))
+ ratio2hws_video_common_v2[h_div_w] = scale_schedule
+ total_pixels2scales = {
+ '0.06M': 4,
+ '0.25M': 8,
+ '1M': 16,
+ }
+ return ratio2hws_video_common_v2, total_pixels2scales
+
+ratio2hws, total_pixels2scales = get_ratio2hws_video_v2()
+ratio2hws = append_dummy_t(ratio2hws)
+dynamic_resolution_h_w = get_full_ratio2hws(ratio2hws, total_pixels2scales)
+dynamic_resolution_thw = {}
+for ratio in dynamic_resolution_h_w:
+ for (h, w) in dynamic_resolution_h_w[ratio]:
+ # spatial_time_schedule = []
+ # spatial_time_schedule.extend(image_scale_schedule)
+ # firstframe_scalecnt = len(image_scale_schedule)
+ # if compressed_frames > 1:
+ # scale_schedule = dynamic_resolution_h_w[ratio][pn]['scales']
+ # # predefined_t = np.linspace(1, compressed_frames - 1, len(scale_schedule))
+ # predefined_t = np.linspace(1, compressed_frames - 1, total_pixels2scales['0.06M']-1).tolist() + [compressed_frames - 1] * (len(scale_schedule)-total_pixels2scales['0.06M']+1)
+ # spatial_time_schedule.extend([(min(int(np.round(predefined_t[i])), compressed_frames - 1), h, w) for i, (_, h, w) in enumerate(scale_schedule)])
+ dynamic_resolution_thw[(h, w)] = dynamic_resolution_h_w[ratio][(h, w)]
+ # dynamic_resolution_thw[(h, w)]['tower_split_index'] = firstframe_scalecnt
+h_div_w_templates = np.array(list(dynamic_resolution_h_w.keys()))
+# print(dynamic_resolution_thw)
+
+if __name__ == '__main__':
+ print(dynamic_resolution_thw)
+
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/utils/ema.py b/grn/tokenizer/videovae/utils/ema.py
new file mode 100644
index 0000000000000000000000000000000000000000..e1a7e40237bd8759603e52c272eee0b36375c18e
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/ema.py
@@ -0,0 +1,27 @@
+import torch
+from collections import OrderedDict
+
+@torch.no_grad()
+def update_ema(ema_model, model, decay=0.9999):
+ """
+ Step the EMA model towards the current model.
+ """
+ ema_params = OrderedDict(ema_model.named_parameters())
+ model_params = OrderedDict(model.named_parameters())
+
+ for name, param in model_params.items():
+ # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
+ ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
+
+ ema_params = OrderedDict(ema_model.named_buffers())
+ model_params = OrderedDict(model.named_buffers())
+ for name, param in model_params.items():
+ if ('signal_' in name) or ('scale_learnable_parameters' in name):
+ ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
+
+def requires_grad(model, flag=True):
+ """
+ Set requires_grad flag for all parameters in a model.
+ """
+ for p in model.parameters():
+ p.requires_grad = flag
\ No newline at end of file
diff --git a/grn/tokenizer/videovae/utils/init_models.py b/grn/tokenizer/videovae/utils/init_models.py
new file mode 100644
index 0000000000000000000000000000000000000000..4c863cb949d220ec1fabf3b88e8502258f8e3812
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/init_models.py
@@ -0,0 +1,403 @@
+import os.path as osp
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import math
+from videovae.utils.misc import is_torch_optim_sch
+
+
+def inflate_gen(state_dict, temporal_patch_size, spatial_patch_size, strategy="average", inflation_pe=False):
+ new_state_dict = state_dict.copy()
+
+ pe_image0_w = state_dict["encoder.to_patch_emb_first_frame.1.weight"] # image_channel * patch_width * patch_height
+ pe_image0_b = state_dict["encoder.to_patch_emb_first_frame.1.bias"] # image_channel * patch_width * patch_height
+ pe_image1_w = state_dict["encoder.to_patch_emb_first_frame.2.weight"] # image_channel * patch_width * patch_height, dim
+ pe_image1_b = state_dict["encoder.to_patch_emb_first_frame.2.bias"] # image_channel * patch_width * patch_height
+ pe_image2_w = state_dict["encoder.to_patch_emb_first_frame.3.weight"] # image_channel * patch_width * patch_height
+ pe_image2_b = state_dict["encoder.to_patch_emb_first_frame.3.bias"] # image_channel * patch_width * patch_height
+
+ pd_image0_w = state_dict["decoder.to_pixels_first_frame.0.weight"] # dim, image_channel * patch_width * patch_height
+ pd_image0_b = state_dict["decoder.to_pixels_first_frame.0.bias"] # image_channel * patch_width * patch_height
+
+ pe_video0_w = state_dict["encoder.to_patch_emb.1.weight"]
+
+ old_patch_size = int(math.sqrt(pe_image0_w.shape[0] // 3))
+ old_patch_size_temporal = pe_video0_w.shape[0] // (3 * old_patch_size * old_patch_size)
+
+ if old_patch_size != spatial_patch_size or old_patch_size_temporal != temporal_patch_size:
+ if not inflation_pe:
+ del new_state_dict["encoder.to_patch_emb_first_frame.1.weight"]
+ del new_state_dict["encoder.to_patch_emb_first_frame.1.bias"]
+ del new_state_dict["encoder.to_patch_emb_first_frame.2.weight"]
+
+ del new_state_dict["decoder.to_pixels_first_frame.0.weight"]
+ del new_state_dict["decoder.to_pixels_first_frame.0.bias"]
+
+ del new_state_dict["encoder.to_patch_emb.1.weight"]
+ del new_state_dict["encoder.to_patch_emb.1.bias"]
+ del new_state_dict["encoder.to_patch_emb.2.weight"]
+
+ del new_state_dict["decoder.to_pixels.0.weight"]
+ del new_state_dict["decoder.to_pixels.0.bias"]
+
+ return new_state_dict
+
+
+ print(f"Inflate the patch embedding size from {old_patch_size_temporal}x{old_patch_size}x{old_patch_size} to {temporal_patch_size}x{spatial_patch_size}x{spatial_patch_size}.")
+ pe_image0_w = F.interpolate(pe_image0_w.unsqueeze(0).unsqueeze(0), size=(3 * spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+ pe_image0_b = F.interpolate(pe_image0_b.unsqueeze(0).unsqueeze(0), size=(3 * spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+ pe_image1_w = F.interpolate(pe_image1_w.unsqueeze(0), size=(3 * spatial_patch_size * spatial_patch_size)).squeeze(0)
+
+ new_state_dict["encoder.to_patch_emb_first_frame.1.weight"] = pe_image0_w
+ new_state_dict["encoder.to_patch_emb_first_frame.1.bias"] = pe_image0_b
+ new_state_dict["encoder.to_patch_emb_first_frame.2.weight"] = pe_image1_w
+
+ pd_image0_w = F.interpolate(pd_image0_w.permute(1, 0).unsqueeze(0), size=(3 * spatial_patch_size * spatial_patch_size)).squeeze(0).permute(1, 0)
+ pd_image0_b = F.interpolate(pd_image0_b.unsqueeze(0).unsqueeze(0), size=(3 * spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+
+ new_state_dict["decoder.to_pixels_first_frame.0.weight"] = pd_image0_w
+ new_state_dict["decoder.to_pixels_first_frame.0.bias"] = pd_image0_b
+
+ pe_video0_w = state_dict["encoder.to_patch_emb.1.weight"]
+ pe_video0_b = state_dict["encoder.to_patch_emb.1.bias"]
+ pe_video1_w = state_dict["encoder.to_patch_emb.2.weight"]
+
+ pe_video0_w = F.interpolate(pe_video0_w.unsqueeze(0).unsqueeze(0), size=(3 * temporal_patch_size * spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+ pe_video0_b = F.interpolate(pe_video0_b.unsqueeze(0).unsqueeze(0), size=(3 * temporal_patch_size* spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+ pe_video1_w = F.interpolate(pe_video1_w.unsqueeze(0), size=(3 * temporal_patch_size * spatial_patch_size * spatial_patch_size)).squeeze(0)
+
+ pd_video0_w = state_dict["decoder.to_pixels.0.weight"]
+ pd_video0_b = state_dict["decoder.to_pixels.0.bias"]
+
+ pd_video0_w = F.interpolate(pd_image0_w.permute(1, 0).unsqueeze(0), size=(3 * temporal_patch_size * spatial_patch_size * spatial_patch_size)).squeeze(0).permute(1, 0)
+ pd_video0_b = F.interpolate(pd_image0_b.unsqueeze(0).unsqueeze(0), size=(3 * temporal_patch_size * spatial_patch_size * spatial_patch_size)).squeeze(0).squeeze(0)
+
+ new_state_dict["encoder.to_patch_emb.1.weight"] = pe_video0_w
+ new_state_dict["encoder.to_patch_emb.1.bias"] = pe_video0_b
+ new_state_dict["encoder.to_patch_emb.2.weight"] = pe_video1_w
+
+ new_state_dict["decoder.to_pixels.0.weight"] = pd_video0_w
+ new_state_dict["decoder.to_pixels.0.bias"] = pd_video0_b
+
+ return new_state_dict
+
+
+ if strategy == "average":
+ pe_video0_w = torch.cat([pe_image0_w/temporal_patch_size] * temporal_patch_size)
+ pe_video0_b = torch.cat([pe_image0_b/temporal_patch_size] * temporal_patch_size)
+
+ pe_video1_w = torch.cat([pe_image1_w/temporal_patch_size] * temporal_patch_size, dim=-1)
+ pe_video1_b = pe_image1_b # torch.cat([pe_image1_b/temporal_patch_size] * temporal_patch_size)
+
+ pe_video2_w = pe_image2_w # torch.cat([pe_image2_w/temporal_patch_size] * temporal_patch_size)
+ pe_video2_b = pe_image2_b # torch.cat([pe_image2_b/temporal_patch_size] * temporal_patch_size)
+
+ elif strategy == "first":
+ pe_video0_w = torch.cat([pe_image0_w] + [torch.zeros_like(pe_image0_w, dtype=pe_image0_w.dtype)] * (temporal_patch_size - 1))
+ pe_video0_b = torch.cat([pe_image0_b] + [torch.zeros_like(pe_image0_b, dtype=pe_image0_b.dtype)] * (temporal_patch_size - 1))
+
+ pe_video1_w = torch.cat([pe_image1_w] + [torch.zeros_like(pe_image1_w, dtype=pe_image1_w.dtype)] * (temporal_patch_size - 1), dim=-1)
+ pe_video1_b = pe_image1_b # torch.cat([pe_image1_b] + [torch.zeros_like(pe_image1_b, dtype=pe_image1_b.dtype)] * (temporal_patch_size - 1))
+
+ pe_video2_w = pe_image2_w # torch.cat([pe_image2_w] + [torch.zeros_like(pe_image2_w, dtype=pe_image2_w.dtype)] * (temporal_patch_size - 1))
+ pe_video2_b = pe_image2_b # torch.cat([pe_image2_b] + [torch.zeros_like(pe_image2_b, dtype=pe_image2_b.dtype)] * (temporal_patch_size - 1))
+
+
+ else:
+ raise NotImplementedError
+
+
+ new_state_dict["encoder.to_patch_emb.1.weight"] = pe_video0_w
+ new_state_dict["encoder.to_patch_emb.1.bias"] = pe_video0_b
+
+ new_state_dict["encoder.to_patch_emb.2.weight"] = pe_video1_w
+ new_state_dict["encoder.to_patch_emb.2.bias"] = pe_video1_b
+
+ new_state_dict["encoder.to_patch_emb.3.weight"] = pe_video2_w
+ new_state_dict["encoder.to_patch_emb.3.bias"] = pe_video2_b
+
+
+ if strategy == "average":
+ pd_video0_w = torch.cat([pd_image0_w/temporal_patch_size] * temporal_patch_size)
+ pd_video0_b = torch.cat([pd_image0_b/temporal_patch_size] * temporal_patch_size)
+
+ elif strategy == "first":
+ pd_video0_w = torch.cat([pd_image0_w] + [torch.zeros_like(pd_image0_w, dtype=pd_image0_w.dtype)] * (temporal_patch_size - 1))
+ pd_video0_b = torch.cat([pd_image0_b] + [torch.zeros_like(pd_image0_b, dtype=pd_image0_b.dtype)] * (temporal_patch_size - 1))
+
+ else:
+ raise NotImplementedError
+
+
+ new_state_dict["decoder.to_pixels.0.weight"] = pd_video0_w
+ new_state_dict["decoder.to_pixels.0.bias"] = pd_video0_b
+
+ return new_state_dict
+
+
+def inflate_dis(state_dict, strategy="center"):
+ print("#" * 50)
+ print(f"Initialize the video discriminator with {strategy}.")
+ print("#" * 50)
+ idis_weights = {k: v for k, v in state_dict.items() if "image_discriminator" in k}
+ vids_weights = {k: v for k, v in state_dict.items() if "video_discriminator" in k}
+
+ new_state_dict = state_dict.copy()
+ for k in vids_weights.keys():
+ del new_state_dict[k]
+
+
+ for k in idis_weights.keys():
+ new_k = "video_discriminator" + k[len("image_discriminator"):]
+ if "weight" in k and new_state_dict[k].ndim == 4:
+ old_weight = state_dict[k]
+ if strategy == "average":
+ new_weight = old_weight.unsqueeze(2).repeat(1, 1, 4, 1, 1) / 4
+ elif strategy == "center":
+ new_weight_ = old_weight# .unsqueeze(2) # O I 1 K K
+ new_weight = torch.zeros((new_weight_.size(0), new_weight_.size(1), 4, new_weight_.size(2), new_weight_.size(3)), dtype=new_weight_.dtype)
+ new_weight[:, :, 1] = new_weight_
+
+ elif strategy == "first":
+ new_weight_ = old_weight# .unsqueeze(2)
+ new_weight = torch.zeros((new_weight_.size(0), new_weight_.size(1), 4, new_weight_.size(2), new_weight_.size(3)), dtype=new_weight_.dtype)
+ new_weight[:, :, 0] = new_weight_
+
+ elif strategy == "last":
+ new_weight_ = old_weight# .unsqueeze(2)
+ new_weight = torch.zeros((new_weight_.size(0), new_weight_.size(1), 4, new_weight_.size(2), new_weight_.size(3)), dtype=new_weight_.dtype)
+ new_weight[:, :, -1] = new_weight_
+ else:
+ raise NotImplementedError
+
+ new_state_dict[new_k] = new_weight
+
+ elif "bias" in k:
+ new_state_dict[new_k] = state_dict[k]
+ else:
+ new_state_dict[new_k] = state_dict[k]
+
+
+ return new_state_dict
+
+def load_unstrictly(state_dict, model, loaded_keys=[]):
+ missing_keys = []
+ for name, param in model.named_parameters():
+ if name in state_dict:
+ try:
+ param.data.copy_(state_dict[name])
+ except:
+ print(f"{name} mismatch: param {name}, shape {param.data.shape}, state_dict shape {state_dict[name].shape}")
+ missing_keys.append(name)
+ elif name not in loaded_keys:
+ missing_keys.append(name)
+ for name, param in model.named_buffers():
+ if name in state_dict:
+ try:
+ param.data.copy_(state_dict[name])
+ except:
+ print(f"{name} mismatch: param {name}, shape {param.data.shape}, state_dict shape {state_dict[name].shape}")
+ missing_keys.append(name)
+ elif name not in loaded_keys:
+ missing_keys.append(name)
+ return model, missing_keys
+
+def init_vae_only(state_dict, vae):
+ vae, missing_keys = load_unstrictly(state_dict, vae)
+ print(f"missing keys in loading vae: {[key for key in missing_keys if not key.startswith('flux')]}")
+ return vae
+
+def init_image_disc(state_dict, image_disc, args):
+ if args.no_init_idis or args.init_idis == "no":
+ state_dict = {}
+ else:
+ state_dict = state_dict["image_disc"]
+ # load nn.GroupNorm to Normalize class
+ delete_keys = []
+ loaded_keys = []
+ model = image_disc
+ for key in state_dict:
+ if key.endswith(".weight"):
+ norm_key = key.replace(".weight", ".norm.weight")
+ if norm_key and norm_key in model.state_dict():
+ model.state_dict()[norm_key].copy_(state_dict[key])
+ delete_keys.append(key)
+ loaded_keys.append(norm_key)
+ if key.endswith(".bias"):
+ norm_key = key.replace(".bias", ".norm.bias")
+ if norm_key and norm_key in model.state_dict():
+ model.state_dict()[norm_key].copy_(state_dict[key])
+ delete_keys.append(key)
+ loaded_keys.append(norm_key)
+ for key in delete_keys:
+ del state_dict[key]
+ msg = image_disc.load_state_dict(state_dict, strict=False)
+ print(f"image disc missing: {[key for key in msg.missing_keys if key not in loaded_keys]}")
+ print(f"image disc unexpected: {msg.unexpected_keys}")
+ return image_disc
+
+def init_video_disc(state_dict, video_disc, args):
+ # init video disc
+ if args.init_vdis == "no":
+ video_disc_state_dict = {}
+ elif args.init_vdis == "keep":
+ video_disc_state_dict = state_dict["video_disc"]
+ else:
+ video_disc_state_dict = inflate_dis(state_dict["video_disc"], strategy=args.init_vdis)
+ msg = video_disc.load_state_dict(video_disc_state_dict, strict=False)
+ print(f"video disc missing: {msg.missing_keys}")
+ print(f"video disc unexpected: {msg.unexpected_keys}")
+ return video_disc
+
+def init_vit_from_image(state_dict, vae, image_disc, video_disc, args):
+ if args.init_vgen == "no":
+ vae_state_dict = state_dict["vae"]
+ del vae_state_dict["encoder.to_patch_emb.1.weight"]
+ del vae_state_dict["encoder.to_patch_emb.1.bias"]
+ del vae_state_dict["encoder.to_patch_emb.2.weight"]
+ del vae_state_dict["encoder.to_patch_emb.2.bias"]
+ del vae_state_dict["encoder.to_patch_emb.3.weight"]
+ del vae_state_dict["encoder.to_patch_emb.3.bias"]
+
+ del vae_state_dict["decoder.to_pixels.0.weight"]
+ del vae_state_dict["decoder.to_pixels.0.bias"]
+ vae_state_dict = state_dict["vae"]
+
+ elif args.init_vgen == "keep":
+ vae_state_dict = state_dict["vae"]
+ else:
+ vae_state_dict = inflate_gen(state_dict["vae"], temporal_patch_size=args.temporal_patch_size, spatial_patch_size=args.patch_size, strategy=args.init_vgen, inflation_pe=args.inflation_pe)
+
+ if args.vq_to_vae:
+ del vae_state_dict["pre_vq_conv.1.weight"]
+ del vae_state_dict["pre_vq_conv.1.bias"]
+
+ msg = vae.load_state_dict(vae_state_dict, strict=False)
+ print(f"vae missing: {msg.missing_keys}")
+ print(f"vae unexpected: {msg.unexpected_keys}")
+
+ image_disc = init_image_disc(state_dict, image_disc, args)
+ # video_disc = init_video_disc(state_dict, image_disc, args) # random init video discriminator
+
+ return vae, image_disc, video_disc
+
+def load_cnn(model, state_dict, prefix, expand=False, use_linear=False):
+ delete_keys = []
+ loaded_keys = []
+ for key in state_dict:
+ if key.startswith(prefix):
+ _key = key[len(prefix):]
+ if _key in model.state_dict():
+ # load nn.Conv2d or nn.Linear to nn.Linear
+ if use_linear and (".q.weight" in key or ".k.weight" in key or ".v.weight" in key or ".proj_out.weight" in key):
+ load_weights = state_dict[key].squeeze()
+ elif _key.endswith(".conv.weight") and expand:
+ if model.state_dict()[_key].shape == state_dict[key].shape:
+ # 2D cnn to 2D cnn
+ load_weights = state_dict[key]
+ else:
+ # 2D cnn to 3D cnn
+ _expand_dim = model.state_dict()[_key].shape[2]
+ load_weights = state_dict[key].unsqueeze(2).repeat(1, 1, _expand_dim, 1, 1)
+ load_weights = load_weights / _expand_dim # normalize across expand dim
+ else:
+ load_weights = state_dict[key]
+ model.state_dict()[_key].copy_(load_weights)
+ delete_keys.append(key)
+ loaded_keys.append(prefix+_key)
+ # load nn.Conv2d to Conv class
+ conv_list = ["conv"] if use_linear else ["conv", ".q.", ".k.", ".v.", ".proj_out.", ".nin_shortcut."]
+ if any(k in _key for k in conv_list):
+ if _key.endswith(".weight"):
+ conv_key = _key.replace(".weight", ".conv.weight")
+ if conv_key and conv_key in model.state_dict():
+ if model.state_dict()[conv_key].shape == state_dict[key].shape:
+ # 2D cnn to 2D cnn
+ load_weights = state_dict[key]
+ else:
+ # 2D cnn to 3D cnn
+ _expand_dim = model.state_dict()[conv_key].shape[2]
+ load_weights = state_dict[key].unsqueeze(2).repeat(1, 1, _expand_dim, 1, 1)
+ load_weights = load_weights / _expand_dim # normalize across expand dim
+ model.state_dict()[conv_key].copy_(load_weights)
+ delete_keys.append(key)
+ loaded_keys.append(prefix+conv_key)
+ if _key.endswith(".bias"):
+ conv_key = _key.replace(".bias", ".conv.bias")
+ if conv_key and conv_key in model.state_dict():
+ model.state_dict()[conv_key].copy_(state_dict[key])
+ delete_keys.append(key)
+ loaded_keys.append(prefix+conv_key)
+ # load nn.GroupNorm to Normalize class
+ if "norm" in _key:
+ if _key.endswith(".weight"):
+ norm_key = _key.replace(".weight", ".norm.weight")
+ if norm_key and norm_key in model.state_dict():
+ model.state_dict()[norm_key].copy_(state_dict[key])
+ delete_keys.append(key)
+ loaded_keys.append(prefix+norm_key)
+ if _key.endswith(".bias"):
+ norm_key = _key.replace(".bias", ".norm.bias")
+ if norm_key and norm_key in model.state_dict():
+ model.state_dict()[norm_key].copy_(state_dict[key])
+ delete_keys.append(key)
+ loaded_keys.append(prefix+norm_key)
+
+ for key in delete_keys:
+ del state_dict[key]
+
+ return model, state_dict, loaded_keys
+
+def init_cnn_from_image(state_dict, vae, image_disc, video_disc, args, expand=False):
+ vae.encoder, state_dict["vae"], loaded_keys1 = load_cnn(vae.encoder, state_dict["vae"], prefix="encoder.", expand=expand)
+ vae.decoder, state_dict["vae"], loaded_keys2 = load_cnn(vae.decoder, state_dict["vae"], prefix="decoder.", expand=expand)
+ loaded_keys = loaded_keys1 + loaded_keys2
+ # msg = vae.load_state_dict(state_dict["vae"], strict=False)
+ # print(f"vae missing: {[key for key in msg.missing_keys if key not in loaded_keys]}")
+ # print(f"vae unexpected: {msg.unexpected_keys}")
+ vae, missing_keys = load_unstrictly(state_dict["vae"], vae, loaded_keys)
+
+ if image_disc:
+ image_disc = init_image_disc(state_dict, image_disc, args)
+ ### random init video discriminator
+ # if video_disc:
+ # video_disc = init_video_disc(state_dict, image_disc, args)
+ return vae, image_disc, video_disc
+
+def resume_from_ckpt(state_dict, model_optims, load_optims=True, remove_disc=False, remove_enlarge_factors=False, ckpt_path=None, args=None):
+ all_missing_keys = []
+ # load weights first
+
+ if remove_enlarge_factors:
+ for key in state_dict['vae']:
+ if 'scale_learnable_parameters' in key:
+ state_dict['vae'][key] = None
+ print(f'del {key}')
+ for key in state_dict['ema']:
+ if 'scale_learnable_parameters' in key:
+ state_dict['vae'][key] = None
+ print(f'del {key}')
+
+ for key1 in ['vae', 'ema']:
+ if (key1 not in state_dict) or (not state_dict[key1]):
+ continue
+ for key2 in state_dict[key1]:
+ if ('scale_learnable_parameters' in key2):
+ if len(state_dict[key1][key2].shape) == 1 and ('channelwise_' in args.enlarge_factors_type):
+ state_dict[key1][key2] = state_dict[key1][key2].reshape(len(state_dict[key1][key2]), 1, 1, 1, 1, 1).repeat(1, 1, args.detail_scale_dim, 1, 1, 1)
+
+ for k in model_optims:
+ if remove_disc:
+ if "image_disc" in k or "video_disc" in k:
+ continue
+ if k in state_dict and model_optims[k] and state_dict[k] and (not is_torch_optim_sch(model_optims[k])):
+ model_optims[k], missing_keys = load_unstrictly(state_dict[k], model_optims[k])
+ all_missing_keys += missing_keys
+
+ print(f"missing weights: {all_missing_keys}, do not load optimzer states")
+ if "step" not in state_dict:
+ if not ckpt_path:
+ state_dict["step"] = 0
+ else:
+ state_dict["step"] = int(osp.basename(ckpt_path).split('_')[-1].split('.')[0])
+ return model_optims, state_dict["step"]
diff --git a/grn/tokenizer/videovae/utils/mfu.py b/grn/tokenizer/videovae/utils/mfu.py
new file mode 100644
index 0000000000000000000000000000000000000000..0458333e0b2cfbfbf7742309f6c72a99b9742272
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/mfu.py
@@ -0,0 +1,179 @@
+import os
+import math
+from datetime import datetime
+from abc import ABC, abstractmethod
+
+import torch
+from torch import nn
+import torch.distributed as dist
+from torch.nn.modules.conv import _ConvNd
+from torch.utils.checkpoint import TorchDispatchMode
+
+
+def get_device_tflops():
+ peak_tflops = -1
+ arch = torch.cuda.get_device_capability()
+ if arch[0] == 8 and arch[1] == 0: # A100/A800
+ peak_tflops = 312
+ elif arch[0] == 9 and arch[1] == 0: # H100/H800
+ peak_tflops = 989
+ else:
+ print(f"unknown default tflops of device capability {arch[0]}.{arch[1]}")
+ return peak_tflops
+
+
+class NullCtx(TorchDispatchMode):
+
+ def __torch_dispatch__(self, func, types, args=(), kwargs=None):
+ if kwargs is None:
+ kwargs = {}
+ return func(*args, **kwargs)
+
+
+class DisableMfu(NullCtx):
+ def __enter__(self):
+ super().__enter__()
+ self.old_flop_enable = Flops.enable
+ Flops.enable = False
+
+ def __exit__(self, *args, **kwargs):
+ Flops.enable = self.old_flop_enable
+ super().__exit__(*args, **kwargs)
+
+
+def context_fn():
+ return NullCtx(), DisableMfu()
+
+
+class CustomFlops(ABC):
+ """
+ For functions,
+ 1. run the func within CustomFlops
+ 2. implement the hook `flops`
+ to support register_forward_hook
+ """
+ @abstractmethod
+ def flops(self, args, kwargs, output) -> dict:
+ pass
+
+
+def conv_flops_func(module, args, kwargs, output):
+ return 2 * math.prod(module.kernel_size) * module.in_channels * output.numel()
+
+
+def linear_flops_func(module, args, kwargs, output):
+ return 2 * module.in_features * output.numel()
+
+
+def layernorm_flops_func(module, args, kwargs, output):
+ return 4 * output.numel()
+
+
+def groupnorm_flops_func(module, args, kwargs, output):
+ return 2 * output.numel()
+
+
+def syncbatchnorm_flops_func(module, args, kwargs, output):
+ return 2 * output.numel()
+
+
+basic_flops_func = {
+ _ConvNd: conv_flops_func,
+ nn.Linear: linear_flops_func,
+ nn.LayerNorm: layernorm_flops_func,
+ nn.GroupNorm: groupnorm_flops_func,
+ nn.SyncBatchNorm: syncbatchnorm_flops_func,
+}
+
+@torch._dynamo.disable()
+def calculate_flops(module, args, kwargs, output):
+ flops = 0
+ flops_dict = {}
+ if isinstance(module, CustomFlops):
+ flops_dict = module.flops(args, kwargs, output)
+ else:
+ flops_func = basic_flops_func[module._base_m]
+ flops_dict = {module.__class__.__name__: flops_func(module, args, kwargs, output)}
+
+ for module_class, module_flops in flops_dict.items():
+ if module_class not in Flops.module_flops_dict:
+ Flops.module_flops_dict[module_class] = module_flops * (3 if module.training else 1)
+ else:
+ Flops.module_flops_dict[module_class] += module_flops * (3 if module.training else 1)
+
+ flops = sum(list(flops_dict.values()))
+ Flops.flops += flops * (3 if module.training else 1)
+
+
+class Flops:
+ handlers = []
+ flops = 0
+ enable = True
+ module_flops_dict = {}
+
+ @staticmethod
+ def reset():
+ tmp = Flops.flops
+ Flops.flops = 0
+ Flops.module_flops_dict = {}
+ return tmp
+
+ @staticmethod
+ def _hook(module, args, kwargs, output):
+ if not Flops.enable:
+ return
+
+ if module.training and not torch.is_grad_enabled():
+ # activation checkpoint mode
+ return
+ calculate_flops(module, args, kwargs, output)
+
+ @staticmethod
+ def _dfs_register_hooks(parent_name: str, cur_m: nn.Module):
+ for name, m in cur_m.named_children():
+ # custom hooks
+ if isinstance(m, CustomFlops):
+ assert isinstance(m, nn.Module)
+ Flops.handlers.append(
+ m.register_forward_hook(Flops._hook, with_kwargs=True)
+ )
+ continue
+ # built-in hooks
+ is_registered = False
+ for base_m, flops_func in basic_flops_func.items():
+ if isinstance(m, base_m):
+ m._base_m = base_m
+ Flops.handlers.append(
+ m.register_forward_hook(Flops._hook, with_kwargs=True)
+ )
+ is_registered = True
+ break
+ if not is_registered:
+ Flops._dfs_register_hooks(parent_name + "." + name, m)
+
+ @staticmethod
+ def unwrap(self):
+ for hdl in Flops.handlers:
+ hdl.remove()
+
+
+
+def register_mfu_hook(model):
+ Flops._dfs_register_hooks("root", model)
+
+
+def get_tflops():
+ return Flops.flops / 1e12
+
+
+def get_tflops_dict(record_iters=1):
+ tflops_dict = {module: round(flops / record_iters/ 1e12, 3) for module, flops in Flops.module_flops_dict.items()}
+ return tflops_dict
+
+
+def get_mfu(iter_time):
+ # compute MFU
+ ideal_TFLOPS = get_device_tflops()
+ achieve_TFLOPs = Flops.reset() / 1e12
+ mfu = achieve_TFLOPs / iter_time / ideal_TFLOPS
+ return mfu
diff --git a/grn/tokenizer/videovae/utils/misc.py b/grn/tokenizer/videovae/utils/misc.py
new file mode 100644
index 0000000000000000000000000000000000000000..57a24e13a43b58e487f29c2d533020f05dcdb0c2
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/misc.py
@@ -0,0 +1,280 @@
+# Copyright (c) Meta Platforms, Inc. All Rights Reserved
+
+import torch
+import torch.distributed as dist
+import imageio
+import os
+import random
+
+import math
+import numpy as np
+from einops import rearrange
+import torch.optim as optim
+import torch.optim.lr_scheduler as lr_scheduler
+
+import sys
+import pdb as pdb_original
+from contextlib import contextmanager
+
+COLOR_BLUE = "\033[94m"
+COLOR_RESET = "\033[0m"
+ptdtype = {None: torch.float32, 'fp32': torch.float32, 'bf16': torch.bfloat16}
+
+def rank_zero_only(fn):
+ def wrapped_fn(*args, **kwargs):
+ if not dist.is_initialized() or dist.get_rank() == 0:
+ return fn(*args, **kwargs)
+ return wrapped_fn
+
+@rank_zero_only
+def print_gpu_usage(model_name) -> None:
+ allocated_memory = torch.cuda.memory_allocated()
+ reserved_memory = torch.cuda.memory_reserved()
+ print(f"after {model_name} backward Allocated Memory: {allocated_memory}, Reserved Memory: {reserved_memory}")
+ torch.cuda.empty_cache()
+
+def seed_everything(seed=0, allow_tf32=True, benchmark=True, deterministic=False):
+ random.seed(seed)
+ np.random.seed(seed)
+ os.environ['PYTHONHASHSEED'] = str(seed)
+ torch.manual_seed(seed)
+ torch.cuda.manual_seed_all(seed)
+
+ torch.backends.cudnn.deterministic = deterministic
+ torch.backends.cudnn.benchmark = benchmark # default False in torch 2.3.1
+
+ # See https://pytorch.org/docs/stable/generated/torch.use_deterministic_algorithms.html
+ os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'
+ # See https://pytorch.org/docs/stable/notes/randomness.html
+ torch.use_deterministic_algorithms(deterministic)
+
+ torch.backends.cudnn.allow_tf32 = allow_tf32 # default True in torch 2.3.1
+ torch.backends.cuda.matmul.allow_tf32 = allow_tf32 # default True in torch 2.3.1
+
+# Function to print model summary in table format
+@rank_zero_only
+def print_model_summary(models):
+ # Table headers
+ print(f"{'Layer Name':<20} {'Param #':<20}")
+ print("="*40)
+ total_params = 0
+ learnable_params = 0
+ for model in models:
+ for name, module in model.named_children():
+ params = sum(p.numel() for p in module.parameters())
+ learnable_params += sum(p.numel() for p in module.parameters() if p.requires_grad)
+ total_params += params
+ params_str = f"{params/1e6:.2f}M"
+ print(f"{name:<20} {params_str:<20}")
+ print("="*40)
+ print(f"Total number of parameters: {total_params/1e6:.2f}M\nLearnable_params: {learnable_params/1e6:.2f}M")
+
+def version_checker(base_version, high_version):
+ try:
+ from bytedance.ndtimeline import __version__
+ from packaging.version import Version
+ if Version(__version__) < Version(base_version) or Version(__version__) >= Version(high_version):
+ raise RuntimeError(f"bytedance.ndtimeline's version should be >={base_version} <{high_version}, but {__version__} found")
+ except ImportError:
+ raise RuntimeError(f"bytedance.ndtimeline's version should be >={base_version} <{high_version}")
+
+def is_torch_optim_sch(obj):
+ return isinstance(obj, (optim.Optimizer, optim.lr_scheduler.LambdaLR))
+
+def rearranged_forward(x, func):
+ if x.ndim == 4:
+ x = rearrange(x, "B C H W -> B H W C")
+ elif x.ndim == 5:
+ x = rearrange(x, "B C T H W -> B T H W C")
+ x = func(x)
+ if x.ndim == 4:
+ x = rearrange(x, "B H W C -> B C H W")
+ elif x.ndim == 5:
+ x = rearrange(x, "B T H W C -> B C T H W")
+ return x
+
+def is_dtype_16(data):
+ return data.dtype == torch.float16 or data.dtype == torch.bfloat16
+
+@contextmanager
+def set_tf32_flags(flag):
+ old_matmul_flag = torch.backends.cuda.matmul.allow_tf32
+ old_cudnn_flag = torch.backends.cudnn.allow_tf32
+ torch.backends.cuda.matmul.allow_tf32 = flag
+ torch.backends.cudnn.allow_tf32 = flag
+ try:
+ yield
+ finally:
+ # Restore the original flags
+ torch.backends.cuda.matmul.allow_tf32 = old_matmul_flag
+ torch.backends.cudnn.allow_tf32 = old_cudnn_flag
+
+class DataPrefixManager:
+ @classmethod
+ def set_data_root(cls, data_root, username=""):
+ cls._current_data_root = data_root
+ cls._username = username
+
+ @classmethod
+ def get_work_dir(cls, use_username=True):
+ return os.path.join(cls._current_data_root, cls._username)
+
+ @classmethod
+ def __call__(cls, rel_path, use_username=True, prefix=""):
+ return os.path.join(cls.get_work_dir(use_username=use_username), prefix, rel_path)
+
+data_prefix_manager = DataPrefixManager()
+
+def get_last_ckpt(root_dir):
+ if not os.path.exists(root_dir): return None
+ ckpt_files = {}
+ for dirpath, dirnames, filenames in os.walk(root_dir):
+ for filename in filenames:
+ if filename.endswith('.ckpt') and ('slim_' not in dirpath):
+ num_iter = int(filename.split('.ckpt')[0].split('_')[-1])
+ ckpt_files[num_iter]=os.path.join(dirpath, filename)
+ iter_list = list(ckpt_files.keys())
+ if len(iter_list) == 0: return None
+ max_iter = max(iter_list)
+ return ckpt_files[max_iter]
+
+
+# Shifts src_tf dim to dest dim
+# i.e. shift_dim(x, 1, -1) would be (b, c, t, h, w) -> (b, t, h, w, c)
+def shift_dim(x, src_dim=-1, dest_dim=-1, make_contiguous=True):
+ n_dims = len(x.shape)
+ if src_dim < 0:
+ src_dim = n_dims + src_dim
+ if dest_dim < 0:
+ dest_dim = n_dims + dest_dim
+
+ assert 0 <= src_dim < n_dims and 0 <= dest_dim < n_dims
+
+ dims = list(range(n_dims))
+ del dims[src_dim]
+
+ permutation = []
+ ctr = 0
+ for i in range(n_dims):
+ if i == dest_dim:
+ permutation.append(src_dim)
+ else:
+ permutation.append(dims[ctr])
+ ctr += 1
+ x = x.permute(permutation)
+ if make_contiguous:
+ x = x.contiguous()
+ return x
+
+
+# reshapes tensor start from dim i (inclusive)
+# to dim j (exclusive) to the desired shape
+# e.g. if x.shape = (b, thw, c) then
+# view_range(x, 1, 2, (t, h, w)) returns
+# x of shape (b, t, h, w, c)
+def view_range(x, i, j, shape):
+ shape = tuple(shape)
+
+ n_dims = len(x.shape)
+ if i < 0:
+ i = n_dims + i
+
+ if j is None:
+ j = n_dims
+ elif j < 0:
+ j = n_dims + j
+
+ assert 0 <= i < j <= n_dims
+
+ x_shape = x.shape
+ target_shape = x_shape[:i] + shape + x_shape[j:]
+ return x.view(target_shape)
+
+
+def accuracy(output, target, topk=(1,)):
+ """Computes the accuracy over the k top predictions for the specified values of k"""
+ with torch.no_grad():
+ maxk = max(topk)
+ batch_size = target.size(0)
+
+ _, pred = output.topk(maxk, 1, True, True)
+ pred = pred.t()
+ correct = pred.eq(target.reshape(1, -1).expand_as(pred))
+
+ res = []
+ for k in topk:
+ correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)
+ res.append(correct_k.mul_(100.0 / batch_size))
+ return res
+
+
+def tensor_slice(x, begin, size):
+ assert all([b >= 0 for b in begin])
+ size = [l - b if s == -1 else s
+ for s, b, l in zip(size, begin, x.shape)]
+ assert all([s >= 0 for s in size])
+
+ slices = [slice(b, b + s) for b, s in zip(begin, size)]
+ return x[slices]
+
+
+def save_video_grid(video, fname, nrow=None, fps=16):
+ b, c, t, h, w = video.shape
+ video = video.permute(0, 2, 3, 4, 1).contiguous()
+
+ video = (video.detach().cpu().numpy() * 255).astype('uint8')
+ if nrow is None:
+ nrow = math.ceil(math.sqrt(b))
+ ncol = math.ceil(b / nrow)
+ padding = 1
+ video_grid = np.zeros((t, (padding + h) * nrow + padding,
+ (padding + w) * ncol + padding, c), dtype='uint8')
+ # print(video_grid.shape)
+ for i in range(b):
+ r = i // ncol
+ c = i % ncol
+ start_r = (padding + h) * r
+ start_c = (padding + w) * c
+ video_grid[:, start_r:start_r + h, start_c:start_c + w] = video[i]
+ video = []
+ for i in range(t):
+ video.append(video_grid[i])
+ imageio.mimsave(fname, video, fps=fps)
+ # skvideo.io.vwrite(fname, video_grid, inputdict={'-r': '5'})
+ # print('saved videos to', fname)
+
+
+def comp_getattr(args, attr_name, default=None):
+ if hasattr(args, attr_name):
+ return getattr(args, attr_name)
+ else:
+ return default
+
+
+def visualize_tensors(t, name=None, nest=0):
+ if name is not None:
+ print(name, "current nest: ", nest)
+ print("type: ", type(t))
+ if 'dict' in str(type(t)):
+ print(t.keys())
+ for k in t.keys():
+ if t[k] is None:
+ print(k, "None")
+ else:
+ if 'Tensor' in str(type(t[k])):
+ print(k, t[k].shape)
+ elif 'dict' in str(type(t[k])):
+ print(k, 'dict')
+ visualize_tensors(t[k], name, nest + 1)
+ elif 'list' in str(type(t[k])):
+ print(k, len(t[k]))
+ visualize_tensors(t[k], name, nest + 1)
+ elif 'list' in str(type(t)):
+ print("list length: ", len(t))
+ for t2 in t:
+ visualize_tensors(t2, name, nest + 1)
+ elif 'Tensor' in str(type(t)):
+ print(t.shape)
+ else:
+ print(t)
+ return ""
diff --git a/grn/tokenizer/videovae/utils/nan_detector.py b/grn/tokenizer/videovae/utils/nan_detector.py
new file mode 100644
index 0000000000000000000000000000000000000000..dca02d1bf9ca083573bc4721ab5361001e1b1db9
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/nan_detector.py
@@ -0,0 +1,107 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates. All Rights Reserved.
+#
+# This source code is licensed under the MIT license found in the
+# LICENSE file in the root directory of this source tree.
+
+import os
+import logging
+
+import torch
+
+logger = logging.getLogger(__name__)
+RANK = int(os.environ["RANK"]) if "RANK" in os.environ else 0
+
+class NanDetector:
+ """
+ Detects the first NaN or Inf in forward and/or backward pass and logs, together with the module name
+ """
+
+ def __init__(self, model, forward=True, backward=True):
+ self.bhooks = []
+ self.fhooks = []
+ self.forward = forward
+ self.backward = backward
+ self.named_parameters = list(model.named_parameters())
+ self.reset()
+
+ for name, mod in model.named_modules():
+ mod.__module_name = name
+ self.add_hooks(mod)
+
+ def __enter__(self):
+ return self
+
+ def __exit__(self, exc_type, exc_value, exc_traceback):
+ # Dump out all model gnorms to enable better debugging
+ norm = {}
+ gradients = {}
+ for name, param in self.named_parameters:
+ if param.grad is not None:
+ grad_norm = torch.norm(param.grad.data, p=2, dtype=torch.float32)
+ norm[name] = grad_norm.item()
+ if torch.isnan(grad_norm).any() or torch.isinf(grad_norm).any():
+ gradients[name] = param.grad.data
+ if len(gradients) > 0:
+ logger.info("Detected nan/inf grad norm, dumping norms...")
+ logger.info(f"norms: {norm}")
+ logger.info(f"gradients: {gradients}")
+
+ self.close()
+
+ def add_hooks(self, module):
+ if self.forward:
+ self.fhooks.append(module.register_forward_hook(self.fhook_fn))
+ if self.backward:
+ self.bhooks.append(module.register_backward_hook(self.bhook_fn))
+
+ def reset(self):
+ self.has_printed_f = False
+ self.has_printed_b = False
+
+ def _detect(self, tensor, name, backward):
+ err = None
+ if (
+ torch.is_floating_point(tensor)
+ # single value tensors (like the loss) will not provide much info
+ and tensor.numel() >= 2
+ ):
+ with torch.no_grad():
+ if torch.isnan(tensor).any():
+ err = "NaN"
+ elif torch.isinf(tensor).any():
+ err = "Inf"
+ if err is not None:
+ err = f"{err} detected in output of {name}, shape: {tensor.shape}, {'backward' if backward else 'forward'}"
+ return err
+
+ def _apply(self, module, inp, x, backward):
+ if torch.is_tensor(x):
+ if isinstance(inp, tuple) and len(inp) > 0:
+ inp = inp[0]
+ err = self._detect(x, module.__module_name, backward)
+ if err is not None:
+ if torch.is_tensor(inp) and not backward:
+ err += (
+ f" input max: {inp.max().item()}, input min: {inp.min().item()}"
+ )
+ has_printed_attr = "has_printed_b" if backward else "has_printed_f"
+ logger.warning(f"rank-{RANK}, err_info : {err}")
+ setattr(self, has_printed_attr, True)
+ elif isinstance(x, dict):
+ for v in x.values():
+ self._apply(module, inp, v, backward)
+ elif isinstance(x, list) or isinstance(x, tuple):
+ for v in x:
+ self._apply(module, inp, v, backward)
+
+ def fhook_fn(self, module, inp, output):
+ if not self.has_printed_f:
+ self._apply(module, inp, output, backward=False)
+
+ def bhook_fn(self, module, inp, output):
+ if not self.has_printed_b:
+ self._apply(module, inp, output, backward=True)
+
+ def close(self):
+ for hook in self.fhooks + self.bhooks:
+ hook.remove()
diff --git a/grn/tokenizer/videovae/utils/scheduler.py b/grn/tokenizer/videovae/utils/scheduler.py
new file mode 100644
index 0000000000000000000000000000000000000000..37247fc05a58efa4faee42c9e387c7f7658dbb53
--- /dev/null
+++ b/grn/tokenizer/videovae/utils/scheduler.py
@@ -0,0 +1,12 @@
+
+def get_lambda(args):
+ if args.scheduler == "linear":
+ def lr_lambda(step):
+ warmup_steps = args.warmup_steps
+ if step < warmup_steps:
+ return step / warmup_steps
+ else:
+ return 1.
+ return lr_lambda
+ else:
+ raise NotImplementedError
diff --git a/grn/trainer/__init__.py b/grn/trainer/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..3daef7b693c94338dd0a74f3870859b32e604afa
--- /dev/null
+++ b/grn/trainer/__init__.py
@@ -0,0 +1,3 @@
+def get_trainer(args):
+ from grn.trainer.sft_trainer import Trainer
+ return Trainer
\ No newline at end of file
diff --git a/grn/trainer/sft_trainer.py b/grn/trainer/sft_trainer.py
new file mode 100644
index 0000000000000000000000000000000000000000..51869f7bbe61f1c22b1f43108b29eddf7c3e6573
--- /dev/null
+++ b/grn/trainer/sft_trainer.py
@@ -0,0 +1,327 @@
+import random
+import time
+import gc
+from functools import partial
+from pprint import pformat
+from typing import List, Optional, Tuple, Union
+import os
+import os.path as osp
+import copy
+import json
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from matplotlib.colors import ListedColormap
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.fsdp.api import FullOptimStateDictConfig, FullStateDictConfig, StateDictType
+from torch.nn.parallel import DistributedDataParallel as DDP
+import numpy as np
+import torch.distributed as tdist
+from torch.amp import autocast
+import cv2
+
+import grn.utils_t2iv.dist as dist
+from grn.models.ema import update_ema
+from grn.utils_t2iv import arg_util, misc
+from grn.utils import wandb_utils
+from grn.schedules import get_encode_decode_func
+from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta
+from grn.utils.compress_tokens import save_packed_tensor
+from grn.models.grn import GRN
+
+Ten = torch.Tensor
+FTen = torch.Tensor
+ITen = torch.LongTensor
+BTen = torch.BoolTensor
+fullstate_save_policy = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
+fulloptstate_save_policy = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
+
+def sum_dict(acc_pt2scale_acc):
+ for full_pt in acc_pt2scale_acc:
+ for si in range(len(acc_pt2scale_acc[full_pt])):
+ acc_pt2scale_acc[full_pt][si] = torch.tensor(acc_pt2scale_acc[full_pt][si]).sum()
+ return acc_pt2scale_acc
+
+def dict2list(acc_pt2scale_acc):
+ flatten_acc_pt2scale_acc = []
+ for key, val in acc_pt2scale_acc.items():
+ flatten_acc_pt2scale_acc.extend(val)
+ return flatten_acc_pt2scale_acc
+
+def list2dict(acc_pt2scale_acc, flatten_acc_pt2scale_acc):
+ ptr = 0
+ for key in acc_pt2scale_acc:
+ for ind in range(len(acc_pt2scale_acc[key])):
+ acc_pt2scale_acc[key][ind] = flatten_acc_pt2scale_acc[ptr]
+ ptr += 1
+ return acc_pt2scale_acc
+
+import queue
+import threading
+
+def save_token():
+ while True:
+ try:
+ raw_features, feature_cache_files4images = save_token_queue.get()
+ for i in range(len(feature_cache_files4images)):
+ if not osp.exists(feature_cache_files4images[i]):
+ os.makedirs(osp.dirname(feature_cache_files4images[i]), exist_ok=True)
+ # torch.save(raw_features[i], feature_cache_files4images[i])
+ save_packed_tensor(feature_cache_files4images[i], raw_features[i])
+ print(f'Save to {feature_cache_files4images[i]}')
+ else:
+ print(f'{feature_cache_files4images[i]} exists, skip')
+ except Exception as e:
+ print(f"Error saving token: {e}")
+ finally:
+ save_token_queue.task_done()
+
+save_token_queue = queue.Queue()
+saver = threading.Thread(target=save_token, daemon=True)
+saver.start()
+
+class Trainer(object):
+ def __init__(
+ self, is_visualizer: bool, device,
+ vae_local, gpt_wo_ddp: GRN, gpt: DDP, gpt_opt: torch.optim.Optimizer,
+ dbg_unused=False,zero=0, vae_latent_dim=True, reweight_loss_by_scale=0,
+ gpt_wo_ddp_ema=None, gpt_ema=None, use_fsdp_model_ema=False, other_args=None,
+ ):
+ super(Trainer, self).__init__()
+ self.zero = zero
+ self.vae_latent_dim = vae_latent_dim
+ self.gpt: Union[DDP, FSDP, nn.Module]
+ self.gpt, self.vae_local = gpt, vae_local
+ self.dynamic_scale_schedule = other_args.dynamic_scale_schedule
+ 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)
+ self.gpt_opt = gpt_opt
+ self.gpt_wo_ddp = gpt_wo_ddp
+ self.gpt_wo_ddp_ema = gpt_wo_ddp_ema
+ self.gpt_ema = gpt_ema
+ self.use_fsdp_model_ema = use_fsdp_model_ema
+ self.batch_size, self.seq_len = 0, 0
+ self.reweight_loss_by_scale = reweight_loss_by_scale
+ video_encode, _, _, _ = get_encode_decode_func(other_args.dynamic_scale_schedule)
+ self.video_encode = video_encode
+ self.is_visualizer = is_visualizer
+ numpy_generator = np.random.default_rng(other_args.seed + tdist.get_rank())
+ torch_cuda_generator = torch.Generator(device=other_args.device)
+ torch_cuda_generator.manual_seed(other_args.seed + tdist.get_rank())
+ self.rank_vary_generator = {
+ 'numpy_generator': numpy_generator,
+ 'torch_cuda_generator': torch_cuda_generator,
+ }
+ gpt_uncompiled = self.gpt_wo_ddp._orig_mod if hasattr(self.gpt_wo_ddp, '_orig_mod') else self.gpt_wo_ddp
+ del gpt_uncompiled.rng
+ gpt_uncompiled.rng = torch.Generator(device=device)
+ del gpt_uncompiled
+
+ def train_step(
+ self, ep: int, it: int, g_it: int, stepping: bool, clip_decay_ratio: float, metric_lg: misc.MetricLogger, logging_params: bool,
+ raw_features_bcthw: FTen, feature_cache_files4images: list, media: str, meta_list: list,
+ inp_B3HW: FTen, text_cond_tuple: Union[ITen, FTen], args: arg_util.Args,
+ ) -> Tuple[torch.Tensor, Optional[float]]:
+ device = args.device
+ B = len(inp_B3HW) + len(raw_features_bcthw)
+
+ if media == 'images':
+ is_image_batch = 1
+ else:
+ is_image_batch = 0
+ # [forward]
+ with torch.autocast('cuda', enabled=True, dtype=torch.bfloat16, cache_enabled=args.zero == 0):
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ raw_features_list = []
+ if len(inp_B3HW):
+ with torch.no_grad():
+ for inp_ind, inp in enumerate(inp_B3HW):
+ raw_features_, _, _ = self.vae_local.encode_for_raw_features(inp.unsqueeze(0), scale_schedule=None, slice=args.use_slice)
+ raw_features = raw_features_[0]
+ save_tokens = args.use_vae_token_cache and args.save_vae_token_cache and (not osp.exists(feature_cache_files4images[inp_ind]))
+ if save_tokens:
+ from grn.utils_t2iv.hbq_util_t2iv import raw_feature2bit_label
+ saved_visual_features = raw_feature2bit_label(raw_features, args.hbq_round).to(torch.bool)
+ saved_visual_features = saved_visual_features.cpu().data
+ save_token_queue.put((saved_visual_features, [feature_cache_files4images[inp_ind]]))
+ raw_features_list.append([raw_features])
+
+ if len(raw_features_bcthw):
+ raw_features_bcthw = [[item.unsqueeze(0)] for item in raw_features_bcthw]
+ raw_features_list = raw_features_list + raw_features_bcthw
+
+ assert isinstance(raw_features_list[0], list)
+ full_pts_this_batch = [item[0].shape[-3] for item in raw_features_list]
+ kv_compact, lens, cu_seqlens_k, max_seqlen_k, caption_nums = text_cond_tuple
+
+ with torch.no_grad():
+ x_BLC, x_BLC_mask, scale_or_time_ids, gt_BLC, _, visual_rope_cache, sequece_packing_scales, super_scale_lengths, other_info_by_scale = self.video_encode(
+ vae=self.vae_local,
+ inp_B3HW=None,
+ vae_features=raw_features_list,
+ args=args,
+ device=device,
+ rope2d_freqs_grid=self.gpt.rope2d_freqs_grid,
+ dynamic_resolution_h_w=self.dynamic_resolution_h_w,
+ text_lens=lens,
+ caption_nums=caption_nums,
+ tokens_remain=args.train_max_token_len,
+ rank_vary_generator=self.rank_vary_generator,
+ vis_verbose=False, # g_it%20==0,
+ meta_list=meta_list,
+ )
+
+ # x_BLC_wo_prefix: torch.Size([bs, 2*2+3*3+...+64*64, d or 4d])
+
+ # import pdb; pdb.set_trace()
+ # from torchinfo import summary
+ # res = summary(self.gpt, input_data=(text_cond_tuple, x_BLC_wo_prefix, scale_schedule))
+
+ logits_norm, loss, acc_bit, valid_sequence_ratio = self.gpt(
+ text_cond_tuple,
+ x_BLC,
+ x_BLC_mask=x_BLC_mask,
+ gt_BL=gt_BLC,
+ is_image_batch=is_image_batch,
+ visual_rope_cache=visual_rope_cache,
+ sequece_packing_scales=sequece_packing_scales,
+ super_scale_lengths=super_scale_lengths,
+ other_info_by_scale=other_info_by_scale,
+ scale_or_time_ids=scale_or_time_ids,
+ ) # loss & acc_bit: [seq_len]
+
+ # [loss reweight]
+ # import pdb; pdb.set_trace()
+ example_global_scales = 10
+ acc_pt2scale_acc = {}
+ acc_pt2scale_acc_counter = {}
+ for full_pt in self.dynamic_resolution_h_w[self.h_div_w_templates[0]]['0.06M']['pt2scale_schedule']:
+ full_pt = int(np.round((full_pt-1) / 4)) * 4 + 1
+ if full_pt not in acc_pt2scale_acc:
+ acc_pt2scale_acc[full_pt] = [[] for _ in range(example_global_scales)]
+ acc_pt2scale_acc_counter[full_pt] = [0 for _ in range(example_global_scales)]
+
+ flatten_L_list, flatten_acc_bit_list, flatten_weight_list = [], [], []
+ ptr = 0
+ global_scale_ind = 0
+ for sample_ind, item in enumerate(sequece_packing_scales):
+ full_pt = full_pts_this_batch[sample_ind]
+ full_pt = int(np.round((full_pt-1) / 4)) * 4 + 1
+ for si, (pt, ph, pw) in enumerate(item):
+ mul_pt_ph_pw = pt * ph * pw
+ start, end = ptr, ptr+mul_pt_ph_pw
+ ptr = end
+ loss_this_scale = loss[start:end].mean()
+ acc_this_scale = acc_bit[start:end].mean()
+ wandb_plot_index = other_info_by_scale[global_scale_ind]['wandb_plot_index']
+ acc_pt2scale_acc[full_pt][wandb_plot_index].append(acc_this_scale)
+ acc_pt2scale_acc_counter[full_pt][wandb_plot_index] += 1
+ flatten_weight_list.append(mul_pt_ph_pw ** args.reweight_loss_by_scale)
+ flatten_L_list.append(loss_this_scale)
+ flatten_acc_bit_list.append(acc_this_scale)
+ global_scale_ind += 1
+ flatten_weight_list = torch.tensor(flatten_weight_list, dtype=loss.dtype, device=loss.device)
+ flatten_weight_list = flatten_weight_list / flatten_weight_list.sum()
+ final_loss = (torch.stack(flatten_L_list) * flatten_weight_list).sum()
+ final_acc_bit = (torch.stack(flatten_acc_bit_list) * flatten_weight_list).sum()
+
+ # [backward]
+ final_loss.backward(retain_graph=False, create_graph=False)
+ if self.zero:
+ grad_norm_t = self.gpt.clip_grad_norm_(args.tclip)
+ else:
+ grad_norm_t = torch.nn.utils.clip_grad_norm_(self.gpt.parameters(), args.tclip) # non zero mode
+ self.gpt_opt.step()
+
+ # update ema
+ if args.use_fsdp_model_ema and (args.model_ema_decay < 1):
+ update_ema(self.gpt_ema, self.gpt, args.model_ema_decay)
+
+ # [zero_grad]
+ if stepping:
+ self.gpt_opt.zero_grad(set_to_none=True)
+
+ # [metric logging]
+ if metric_lg.log_every_iter or it == 0 or it in metric_lg.log_iters:
+ acc_pt2scale_acc = sum_dict(acc_pt2scale_acc)
+ flatten_acc_pt2scale_acc = dict2list(acc_pt2scale_acc)
+ flatten_acc_pt2scale_acc_counter = dict2list(acc_pt2scale_acc_counter)
+
+ train_loss = final_loss.item()
+ train_acc = final_acc_bit.item()
+ metrics = torch.tensor(flatten_acc_pt2scale_acc + flatten_acc_pt2scale_acc_counter + [grad_norm_t.item(), train_loss, train_acc, is_image_batch, valid_sequence_ratio], device=loss.device)
+ tdist.all_reduce(metrics, op=tdist.ReduceOp.SUM)
+ flatten_acc_pt2scale_acc, flatten_acc_pt2scale_acc_counter = metrics[:len(flatten_acc_pt2scale_acc)], metrics[len(flatten_acc_pt2scale_acc):2*len(flatten_acc_pt2scale_acc)]
+ flatten_acc_pt2scale_acc = flatten_acc_pt2scale_acc / (flatten_acc_pt2scale_acc_counter + 1e-16)
+ acc_pt2scale_acc = list2dict(acc_pt2scale_acc, flatten_acc_pt2scale_acc)
+ acc_pt2scale_acc_counter = list2dict(acc_pt2scale_acc_counter, flatten_acc_pt2scale_acc_counter)
+ grad_norm_t, train_loss, train_acc, is_image_batch, valid_sequence_ratio = metrics[2*len(flatten_acc_pt2scale_acc):] / (dist.get_world_size() + 1e-16)
+ if args.num_of_label_value == 1:
+ key, base = 'Loss', 1
+ else:
+ key, base = 'Acc', 100
+ metric_lg.update(L=train_loss, Acc=train_acc*base, L_i=0., Acc_i=0., L_v=0., Acc_v=0., tnm=grad_norm_t, seq_usage=valid_sequence_ratio*100.) # todo: Accm, Acct
+ wandb_log_dict = {
+ 'Overall/train_loss': train_loss,
+ 'Overall/train_acc': train_acc*base,
+ 'Overall/grad_norm_t': grad_norm_t,
+ 'Overall/logits_abs_mean': logits_norm.item(),
+ 'Overall/video_batch_ratio': (1-is_image_batch)*100.,
+ 'Overall/valid_sequence_ratio': valid_sequence_ratio*100.,
+ }
+ for full_pt in acc_pt2scale_acc:
+ for si in range(len(acc_pt2scale_acc[full_pt])):
+ if acc_pt2scale_acc_counter[full_pt][si] > 0:
+ duration = (full_pt-1) / 4
+ prefix = f't{duration:04.1f}s/signal_{si*0.1:.01f}_{(si+1)*0.1:.01f}'
+ wandb_log_dict[f'Details/{key}/{prefix}'] = acc_pt2scale_acc[full_pt][si].item() * base
+ wandb_log_dict[f'Details/Num/{prefix}'] = acc_pt2scale_acc_counter[full_pt][si]
+ wandb_utils.log(wandb_log_dict, step=g_it)
+
+ def __repr__(self):
+ return (
+ f'\n'
+ f'[VGPTTr.config]: {pformat(self.get_config(), indent=2, width=250)}\n'
+ f'[VGPTTr.structure]: {super(Trainer, self).__repr__().replace(Trainer.__name__, "")}'
+ )
+
+ def get_config(self):
+ return {}
+
+ def state_dict(self):
+ m = self.vae_local
+ if hasattr(m, '_orig_mod'):
+ m = m._orig_mod
+ state = {'config': self.get_config()}
+
+ if self.zero:
+ state['gpt_fsdp'] = None
+ with FSDP.state_dict_type(self.gpt, StateDictType.FULL_STATE_DICT, fullstate_save_policy, fulloptstate_save_policy):
+ state['gpt_fsdp'] = self.gpt.state_dict()
+ if self.use_fsdp_model_ema:
+ state['gpt_ema_fsdp'] = self.gpt_ema.state_dict()
+ state['gpt_fsdp_opt'] = None
+ else:
+ if self.using_ema:
+ self.ema_load()
+ state['gpt_ema_for_vis'] = {k: v.cpu() for k, v in self.gpt_wo_ddp.state_dict().items()}
+ self.ema_recover()
+
+ for k in ('gpt_wo_ddp', 'gpt_opt'):
+ m = getattr(self, k)
+ if m is not None:
+ if hasattr(m, '_orig_mod'):
+ m = m._orig_mod
+ state[k] = m.state_dict()
+ return state
+
+ def load_state_dict(self, state, strict=True, skip_vae=False):
+ if self.zero:
+ with FSDP.state_dict_type(self.gpt, StateDictType.FULL_STATE_DICT, fullstate_save_policy, fulloptstate_save_policy):
+ self.gpt.load_state_dict(state['gpt_fsdp'])
+ if self.use_fsdp_model_ema:
+ self.gpt_ema.load_state_dict(state['gpt_ema_fsdp'])
+ one_group_opt_state = state['gpt_fsdp_opt']
+ optim_state_dict = FSDP.optim_state_dict_to_load(model=self.gpt, optim=self.gpt_opt.optimizer, optim_state_dict=one_group_opt_state)
+ else:
+ raise NotImplementedError
diff --git a/grn/utils/compress_tokens.py b/grn/utils/compress_tokens.py
new file mode 100644
index 0000000000000000000000000000000000000000..6bc5fc36da1e6a7c4f10abe50a3134f1f82b70dc
--- /dev/null
+++ b/grn/utils/compress_tokens.py
@@ -0,0 +1,20 @@
+import torch
+import numpy as np
+
+def save_packed_tensor(filename, tensor):
+ """use np.savez_compressed to save compressed tensor and its shape"""
+ if tensor.dtype != torch.bool:
+ raise TypeError("Input tensor must be of dtype torch.bool")
+ shape_array = np.array(tensor.shape)
+ packed_data = np.packbits(tensor.numpy())
+ np.savez_compressed(filename, shape=shape_array, data=packed_data)
+
+def load_packed_tensor(filename):
+ """read .npz file and decompress tensor"""
+ with np.load(filename) as loader:
+ shape = loader['shape']
+ packed_data = loader['data']
+ numel = np.prod(shape)
+ unpacked_data = np.unpackbits(packed_data, count=numel)
+ restored_tensor = torch.from_numpy(unpacked_data.reshape(shape))
+ return restored_tensor
diff --git a/grn/utils/safe_rm.py b/grn/utils/safe_rm.py
new file mode 100644
index 0000000000000000000000000000000000000000..d14bbc13ab3ac52e4e369503adc6d86ea811f694
--- /dev/null
+++ b/grn/utils/safe_rm.py
@@ -0,0 +1,58 @@
+import os
+import shutil
+import glob
+from pathlib import Path
+import sys
+
+def safe_remove(target_path, expected_working_dir=None):
+ """
+ Safely remove a file or directory.
+ - Resolves glob patterns
+ - Checks for '..' in path
+ - Prevents deleting root '/'
+ - Ensures target is within expected_working_dir
+ """
+ if expected_working_dir is None:
+ expected_working_dir = os.getcwd()
+
+ expected_working_dir = os.path.abspath(expected_working_dir)
+ target_path_str = str(target_path)
+
+ # Check if original path contains '..'
+ if '..' in target_path_str:
+ print(f"Warning: path {target_path_str} contains '..'. Deletion aborted.")
+ return
+
+ # Handle globs
+ matched_paths = glob.glob(target_path_str)
+ if not matched_paths and not '*' in target_path_str and not '?' in target_path_str:
+ matched_paths = [target_path_str]
+
+ for p in matched_paths:
+ abs_p = os.path.abspath(p)
+
+ # Security checks
+ if abs_p == '/':
+ print("Warning: attempting to delete root directory. Deletion aborted.")
+ continue
+
+ if not abs_p.startswith(expected_working_dir):
+ print(f"Warning: path {abs_p} is not within expected working directory {expected_working_dir}. Deletion aborted.")
+ continue
+
+ path_obj = Path(abs_p)
+ if path_obj.exists() or path_obj.is_symlink():
+ if path_obj.is_dir() and not path_obj.is_symlink():
+ shutil.rmtree(abs_p)
+ else:
+ path_obj.unlink()
+
+if __name__ == '__main__':
+ if len(sys.argv) < 2:
+ print("Usage: python safe_rm.py [expected_working_dir]")
+ sys.exit(1)
+
+ target = sys.argv[1]
+ expected_dir = sys.argv[2] if len(sys.argv) > 2 else os.path.dirname(os.path.abspath(target))
+
+ safe_remove(target, expected_dir)
diff --git a/grn/utils/video_decoder.py b/grn/utils/video_decoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..810356547f5ac9e4474ea6c48a4047443b77800a
--- /dev/null
+++ b/grn/utils/video_decoder.py
@@ -0,0 +1,288 @@
+from abc import ABC, abstractmethod
+import io
+import math
+import numpy as np
+from typing import Optional, TypeVar, Union
+import collections
+
+try:
+ import decord
+except ImportError:
+ _HAS_DECORD = False
+else:
+ _HAS_DECORD = True
+
+if _HAS_DECORD:
+ decord.bridge.set_bridge('native')
+
+DecordDevice = TypeVar("DecordDevice")
+
+# https://github.com/dmlc/decord/issues/208#issuecomment-1157632702
+class VideoReaderWrapper(decord.VideoReader):
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.seek(0)
+
+ def __getitem__(self, key):
+ frames = super().__getitem__(key)
+ self.seek(0)
+ return frames
+
+
+class Video(ABC):
+ """
+ Video provides an interface to access clips from a video container.
+ """
+
+ @abstractmethod
+ def __init__(
+ self,
+ file: Union[str, io.IOBase],
+ video_name: Optional[str] = None,
+ decode_audio: bool = True,
+ ) -> None:
+ """
+ Args:
+ file (BinaryIO): a file-like object (e.g. io.BytesIO or io.StringIO) that
+ contains the encoded video.
+ """
+ pass
+
+ @property
+ @abstractmethod
+ def duration(self) -> float:
+ """
+ Returns:
+ duration of the video in seconds
+ """
+ pass
+
+ @abstractmethod
+ def get_clip(
+ self, start_sec: float, end_sec: float, num_samples: int
+ ):
+ """
+ Retrieves frames from the internal video at the specified start and end times
+ in seconds (the video always starts at 0 seconds).
+
+ Args:
+ start_sec (float): the clip start time in seconds
+ end_sec (float): the clip end time in seconds
+ Returns:
+ video_data_dictonary: A dictionary mapping strings to tensor of the clip's
+ underlying data.
+
+ """
+ pass
+
+ def close(self):
+ pass
+
+
+class EncodedVideoDecord(Video):
+ """
+
+ Accessing clips from an encoded video using Decord video reading API
+ as the decoding backend. For more details, please refer to -
+ `Decord `
+ """
+
+ def __init__(
+ self,
+ file: Union[str, io.IOBase],
+ video_name: Optional[str] = None,
+ width: int = -1,
+ height: int = -1,
+ num_threads: int = 0,
+ fault_tol: int = -1,
+ ) -> None:
+ """
+ Args:
+ file str: file path.
+ video_name (str): An optional name assigned to the video.
+ decode_audio (bool): If disabled, audio is not decoded.
+ sample_rate: int, default is -1
+ Desired output sample rate of the audio, unchanged if `-1` is specified.
+ mono: bool, default is True
+ Desired output channel layout of the audio. `True` is mono layout. `False`
+ is unchanged.
+ width : int, default is -1
+ Desired output width of the video, unchanged if `-1` is specified.
+ height : int, default is -1
+ Desired output height of the video, unchanged if `-1` is specified.
+ num_threads : int, default is 0
+ Number of decoding thread, auto if `0` is specified.
+ fault_tol : int, default is -1
+ The threshold of corupted and recovered frames. This is to prevent silent fault
+ tolerance when for example 50% frames of a video cannot be decoded and duplicate
+ frames are returned. You may find the fault tolerant feature sweet in many
+ cases, but not for training models. Say `N = # recovered frames`
+ If `fault_tol` < 0, nothing will happen.
+ If 0 < `fault_tol` < 1.0, if N > `fault_tol * len(video)`,
+ raise `DECORDLimitReachedError`.
+ If 1 < `fault_tol`, if N > `fault_tol`, raise `DECORDLimitReachedError`.
+ """
+ self._video_name = video_name
+ if not _HAS_DECORD:
+ raise ImportError(
+ "decord is required to use EncodedVideoDecord decoder. Please "
+ "install with 'pip install decord' for CPU-only version and refer to"
+ "'https://github.com/dmlc/decord' for GPU-supported version"
+ )
+ try:
+ self._av_reader = VideoReaderWrapper(
+ uri=file,
+ ctx=decord.cpu(0),
+ width=width,
+ height=height,
+ num_threads=num_threads,
+ fault_tol=fault_tol,
+ )
+ except Exception as e:
+ raise RuntimeError(f"Failed to open video {video_name} with Decord. {e}")
+
+ self._fps = self._av_reader.get_avg_fps()
+ self._duration = float(len(self._av_reader)) / float(self._fps)
+
+ @property
+ def name(self) -> Optional[str]:
+ """
+ Returns:
+ name: the name of the stored video if set.
+ """
+ return self._video_name
+
+ @property
+ def duration(self) -> float:
+ """
+ Returns:
+ duration: the video's duration/end-time in seconds.
+ """
+ return self._duration
+
+ def close(self):
+ if self._av_reader is not None:
+ del self._av_reader
+ self._av_reader = None
+
+ def get_clip(
+ self, start_sec: float, end_sec: float, num_samples: int
+ ):
+ """
+ Retrieves frames from the encoded video at the specified start and end times
+ in seconds (the video always starts at 0 seconds).
+
+ Args:
+ start_sec (float): the clip start time in seconds
+ end_sec (float): the clip end time in seconds
+ Returns:
+ clip_data:
+ A dictionary mapping the entries at "video" and "audio" to a tensors.
+
+ "video": A tensor of the clip's RGB frames with shape:
+ (channel, time, height, width). The frames are of type torch.float32 and
+ in the range [0 - 255].
+
+ Returns None if no video or audio found within time range.
+
+ """
+ if start_sec > end_sec or start_sec > self._duration:
+ raise RuntimeError(
+ f"Incorrect time window for Decord decoding for video: {self._video_name}."
+ )
+
+ start_idx = math.ceil(self._fps * start_sec)
+ end_idx = math.ceil(self._fps * end_sec)
+ end_idx = min(end_idx, len(self._av_reader))
+ # frame_idxs = list(range(start_idx, end_idx))
+
+ frame_idxs = np.linspace(start_idx, end_idx - 1, num_samples, dtype=int)
+
+ try:
+ outputs = self._av_reader.get_batch(frame_idxs)
+ return outputs.asnumpy(), frame_idxs - frame_idxs[0]
+ except Exception as e:
+ print(f"Failed to decode video with Decord: {self._video_name}. {e}")
+ raise e
+
+try:
+ import cv2
+except ImportError:
+ print(f"ERR: import cv2 failed, install cv2 by 'pip install opencv-python'")
+
+class EncodedVideoOpencv():
+ def __init__(
+ self,
+ file: Union[str, io.IOBase],
+ video_name: Optional[str] = None,
+ width: int = -1,
+ height: int = -1,
+ num_threads: int = 0,
+ fault_tol: int = -1,
+ ) -> None:
+ """
+ Args:
+ file str: file path.
+ video_name (str): An optional name assigned to the video.
+ width : Not support yet.
+ height : Not support yet.
+ num_threads : Not support yet.
+ fault_tol : Not support yet.
+ """
+
+ self._video_name = video_name
+ self.cap = cv2.VideoCapture(file)
+ self._fps = self.cap.get(cv2.CAP_PROP_FPS)
+ self._vlen = int(self.cap.get(cv2.CAP_PROP_FRAME_COUNT))
+ self._duration = float(self._vlen) / float(self._fps)
+
+ @property
+ def name(self) -> Optional[str]:
+ """
+ Returns:
+ name: the name of the stored video if set.
+ """
+ return self._video_name
+
+ @property
+ def duration(self) -> float:
+ """
+ Returns:
+ duration: the video's duration/end-time in seconds.
+ """
+ return self._duration
+
+ def __del__(self):
+ self.close()
+
+ def close(self):
+ self.cap.release()
+
+ def get_clip(
+ self, start_sec: float, end_sec: float, num_samples: int
+ ):
+ if start_sec > end_sec or start_sec > self._duration:
+ raise RuntimeError(
+ f"Incorrect time window for Decord decoding for video: {self._video_name}."
+ )
+ start_idx = math.ceil(self._fps * start_sec)
+ end_idx = math.ceil(self._fps * end_sec)
+ end_idx = min(end_idx, self._vlen)
+ frame_idxs = np.linspace(start_idx, end_idx - 1, num_samples, dtype=int)
+ frame_idx2freq = collections.defaultdict(int)
+ for frame_idx in frame_idxs:
+ frame_idx2freq[frame_idx] += 1
+ try:
+ frames = []
+ for i in range(self._vlen):
+ if i > frame_idxs[-1]:
+ break
+ ret, frame = self.cap.read()
+ if i in frame_idx2freq:
+ frames.extend([frame] * frame_idx2freq[i])
+ frames = np.array(frames).astype(np.uint8) # BGR type
+ assert len(frames) == num_samples
+ return frames, frame_idxs - frame_idxs[0]
+ except Exception as e:
+ print(f"Failed to decode video with opencv: {self._video_name}. {e}")
+ raise e
diff --git a/grn/utils/wandb_utils.py b/grn/utils/wandb_utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..fabbe635e3910a7608ffbb9227f3f1b84d47c1f5
--- /dev/null
+++ b/grn/utils/wandb_utils.py
@@ -0,0 +1,55 @@
+import wandb
+import torch
+from torchvision.utils import make_grid
+import torch.distributed as dist
+from PIL import Image
+import os
+import argparse
+import hashlib
+import math
+
+
+def is_main_process():
+ return dist.get_rank() == 0
+
+def namespace_to_dict(namespace):
+ return {
+ k: namespace_to_dict(v) if isinstance(v, argparse.Namespace) else v
+ for k, v in vars(namespace).items()
+ }
+
+
+def generate_run_id(exp_name):
+ # https://stackoverflow.com/questions/16008670/how-to-hash-a-string-into-8-digits
+ return str(int(hashlib.sha256(exp_name.encode('utf-8')).hexdigest(), 16) % 10 ** 8)
+
+
+def initialize(args, entity, exp_name, project_name):
+ config_dict = namespace_to_dict(args)
+ wandb.login(key=os.environ["WANDB_KEY"])
+ wandb.init(
+ entity=entity,
+ project=project_name,
+ name=exp_name,
+ config=config_dict,
+ id=generate_run_id(exp_name),
+ resume="allow",
+ )
+
+
+def log(stats, step=None):
+ if is_main_process():
+ wandb.log({k: v for k, v in stats.items()}, step=step)
+
+
+def log_image(name, sample, step=None):
+ if is_main_process():
+ sample = array2grid(sample)
+ wandb.log({f"{name}": wandb.Image(sample), "train_step": step})
+
+
+def array2grid(x):
+ nrow = round(math.sqrt(x.size(0)))
+ x = make_grid(x, nrow=nrow, normalize=True, value_range=(-1,1))
+ x = x.mul(255).add_(0.5).clamp_(0,255).permute(1,2,0).to('cpu', torch.uint8).numpy()
+ return x
\ No newline at end of file
diff --git a/grn/utils_c2i/crop.py b/grn/utils_c2i/crop.py
new file mode 100644
index 0000000000000000000000000000000000000000..7582690e6ebbbc5ce397906d5e3724f21b8755be
--- /dev/null
+++ b/grn/utils_c2i/crop.py
@@ -0,0 +1,23 @@
+import numpy as np
+from PIL import Image
+
+
+def center_crop_arr(pil_image, image_size):
+ """
+ Center cropping implementation from ADM.
+ https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
+ """
+ while min(*pil_image.size) >= 2 * image_size:
+ pil_image = pil_image.resize(
+ tuple(x // 2 for x in pil_image.size), resample=Image.BOX
+ )
+
+ scale = image_size / min(*pil_image.size)
+ pil_image = pil_image.resize(
+ tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC
+ )
+
+ arr = np.array(pil_image)
+ crop_y = (arr.shape[0] - image_size) // 2
+ crop_x = (arr.shape[1] - image_size) // 2
+ return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size])
diff --git a/grn/utils_c2i/denoiser.py b/grn/utils_c2i/denoiser.py
new file mode 100644
index 0000000000000000000000000000000000000000..3b8dc2e05fe9d9572f3d68ceac5e9b4c913d93e3
--- /dev/null
+++ b/grn/utils_c2i/denoiser.py
@@ -0,0 +1,168 @@
+import torch
+import torch.nn as nn
+from grn.models.grn_c2i import GRN_models
+
+class Denoiser(nn.Module):
+ def __init__(
+ self,
+ args
+ ):
+ super().__init__()
+ vae_down_sample = 16
+ self.net = GRN_models[args.model](
+ input_size=args.img_size // vae_down_sample,
+ in_channels=args.in_channels,
+ num_classes=args.class_num,
+ attn_drop=args.attn_dropout,
+ proj_drop=args.proj_dropout,
+ args=args,
+ )
+ self.img_size = args.img_size // vae_down_sample
+ self.num_classes = args.class_num
+
+ self.label_drop_prob = args.label_drop_prob
+ self.P_mean = args.P_mean
+ self.P_std = args.P_std
+ self.t_eps = args.t_eps
+ self.noise_scale = args.noise_scale
+
+ # ema
+ self.ema_decay1 = args.ema_decay1
+ self.ema_decay2 = args.ema_decay2
+ self.ema_params1 = None
+ self.ema_params2 = None
+
+ # generation hyper params
+ self.method = args.sampling_method
+ self.steps = args.num_sampling_steps
+ self.cfg_scale = args.cfg
+ self.cfg_interval = (args.interval_min, args.interval_max)
+ self.args = args
+
+ def drop_labels(self, labels):
+ drop = torch.rand(labels.shape[0], device=labels.device) < self.label_drop_prob
+ out = torch.where(drop, torch.full_like(labels, self.num_classes), labels)
+ return out
+
+ def sample_t(self, n: int, device=None):
+ z = torch.randn(n, device=device) * self.P_std + self.P_mean
+ return torch.sigmoid(z)
+
+ def forward(self, x, labels):
+ labels_dropped = self.drop_labels(labels) if self.training else labels
+
+ t = self.sample_t(x.size(0), device=x.device).view(-1, *([1] * (x.ndim - 1)))
+
+ # x shape: [B,d,h,w]
+ if self.args.method == 'GRN_ind':
+ classes = 2**self.args.hbq_round
+ elif self.args.method == 'GRN_bit':
+ classes = 2
+ random_labels = torch.randint(0, classes, size=x.shape, device=x.device)
+ x_mask = torch.rand(size=x.shape, device=x.device) < t
+ z = torch.where(x_mask, x, random_labels)
+ x_pred = self.net(z, t.flatten(), labels_dropped) # x_pred shape: [B, classes, d, h, w]
+ # ce loss
+ gt_labels = x # [B,d,h,w]
+ pred_labels = torch.argmax(x_pred, 1)# [B,d,h,w]
+ pred_acc = pred_labels == gt_labels
+ t_bin2acc, t_bin2freq = [0. for _ in range(10)], [0. for _ in range(10)]
+ pred_acc_1d = pred_acc.float().mean([1,2,3]).reshape(-1)
+ t_1d = t.reshape(-1)
+ for ind in range(len(pred_acc_1d)):
+ t_round = int(t_1d[ind] / 0.1)
+ t_round = min(t_round, 9)
+ t_bin2acc[t_round] += pred_acc_1d[ind]
+ t_bin2freq[t_round] += 1
+ t_bin2acc = torch.tensor(t_bin2acc, device=x_pred.device)
+ t_bin2freq = torch.tensor(t_bin2freq, device=x_pred.device)
+ with torch.amp.autocast('cuda', dtype=torch.float32):
+ loss = torch.nn.functional.cross_entropy(x_pred, gt_labels)
+ return loss, t_bin2acc, t_bin2freq
+
+ @torch.no_grad()
+ def generate(self, labels):
+ device = labels.device
+ bsz = labels.size(0)
+ if self.args.method in ['GRN_ind']:
+ classes = 2 ** self.args.hbq_round
+ rand_labels = torch.randint(0, classes, (bsz, self.net.in_channels//classes, self.img_size, self.img_size), device=device)
+ z = rand_labels
+ elif self.args.method in ['GRN_bit']:
+ classes = 2
+ rand_labels = torch.randint(0, classes, (bsz, self.net.in_channels//classes, self.img_size, self.img_size), device=device)
+ z = rand_labels
+ else:
+ raise ValueError(f'{self.args.method=} is not supported')
+ timesteps = torch.linspace(0.0, 1.0, self.steps+1, device=device).view(-1, *([1] * z.ndim)).expand(-1, bsz, -1, -1, -1)
+
+ for i in range(self.steps): # self.steps=50
+ t = timesteps[i] # len(timesteps)=51
+ t_next = timesteps[i + 1]
+
+ # conditional
+ x_cond = self.net(z, t.flatten(), labels)
+
+ # unconditional
+ x_uncond = self.net(z, t.flatten(), torch.full_like(labels, self.num_classes))
+
+ # cfg interval
+ low, high = self.cfg_interval
+ if low < 0: # power-cos cfg
+ rescale_cfg_weight = (1 - torch.cos((t ** (-low)) * torch.pi)) * 1/2
+ x_pred = x_cond + self.cfg_scale * rescale_cfg_weight.unsqueeze(1) * (x_cond - x_uncond)
+ else:
+ interval_mask = (t < high) & ((low == 0) | (t > low))
+ cfg_scale_interval = torch.where(interval_mask, self.cfg_scale, 1.0)
+ x_pred = x_uncond + cfg_scale_interval.unsqueeze(1) * (x_cond - x_uncond)
+ x_pred = x_pred / self.args.tau
+ x_pred = x_pred.softmax(dim=1)
+ B, classes, d, h, w = x_pred.shape
+ x_pred = x_pred.permute(0,2,3,4,1) # [B,d,h,w,classes]
+ pred_labels = torch.multinomial(x_pred.reshape(-1, classes), num_samples=1, replacement=True, generator=None).reshape(B,d,h,w)
+ use_pred_mask = torch.rand(size=pred_labels.shape, device=pred_labels.device) < t_next
+ z = torch.where(use_pred_mask, pred_labels, rand_labels)
+ return z
+
+
+ @torch.no_grad()
+ def _forward_sample(self, z, t, labels):
+ # conditional
+ x_cond = self.net(z, t.flatten(), labels)
+ v_cond = (x_cond - z) / (1.0 - t).clamp_min(self.t_eps)
+
+ # unconditional
+ x_uncond = self.net(z, t.flatten(), torch.full_like(labels, self.num_classes))
+ v_uncond = (x_uncond - z) / (1.0 - t).clamp_min(self.t_eps)
+
+ # cfg interval
+ low, high = self.cfg_interval
+ interval_mask = (t < high) & ((low == 0) | (t > low))
+ cfg_scale_interval = torch.where(interval_mask, self.cfg_scale, 1.0)
+
+ return v_uncond + cfg_scale_interval * (v_cond - v_uncond)
+
+ @torch.no_grad()
+ def _euler_step(self, z, t, t_next, labels):
+ v_pred = self._forward_sample(z, t, labels)
+ z_next = z + (t_next - t) * v_pred
+ return z_next
+
+ @torch.no_grad()
+ def _heun_step(self, z, t, t_next, labels):
+ v_pred_t = self._forward_sample(z, t, labels)
+
+ z_next_euler = z + (t_next - t) * v_pred_t
+ v_pred_t_next = self._forward_sample(z_next_euler, t_next, labels)
+
+ v_pred = 0.5 * (v_pred_t + v_pred_t_next)
+ z_next = z + (t_next - t) * v_pred
+ return z_next
+
+ @torch.no_grad()
+ def update_ema(self):
+ source_params = list(self.parameters())
+ for targ, src in zip(self.ema_params1, source_params):
+ targ.detach().mul_(self.ema_decay1).add_(src, alpha=1 - self.ema_decay1)
+ for targ, src in zip(self.ema_params2, source_params):
+ targ.detach().mul_(self.ema_decay2).add_(src, alpha=1 - self.ema_decay2)
diff --git a/grn/utils_c2i/engine.py b/grn/utils_c2i/engine.py
new file mode 100644
index 0000000000000000000000000000000000000000..a8cd10ccb0459ddcad4d8ac355a782187864997e
--- /dev/null
+++ b/grn/utils_c2i/engine.py
@@ -0,0 +1,229 @@
+import math
+import sys
+import os
+import shutil
+import json
+import copy
+
+import torch
+import numpy as np
+import cv2
+import torch_fidelity
+
+import grn.utils_c2i.misc as misc
+import grn.utils_c2i.lr_sched as lr_sched
+import grn.utils.wandb_utils as wandb_utils
+from grn.utils_c2i.hbq_util_c2i import raw_feature2label, raw_feature2bit_label
+
+def train_one_epoch(model, model_without_ddp, data_loader, optimizer, device, epoch, log_writer=None, args=None, vae=None):
+ model.train(True)
+ metric_logger = misc.MetricLogger(delimiter=" ")
+ metric_logger.add_meter('lr', misc.SmoothedValue(window_size=1, fmt='{value:.6f}'))
+ metric_logger.add_meter('grad_norm', misc.SmoothedValue(window_size=1, fmt='{value:.6f}'))
+ header = 'Epoch: [{}]'.format(epoch)
+ print_freq = 20
+
+ optimizer.zero_grad()
+
+ if log_writer is not None:
+ print('log_dir: {}'.format(log_writer.log_dir))
+
+ for data_iter_step, (x, labels) in enumerate(metric_logger.log_every(data_loader, print_freq, header)):
+ # per iteration (instead of per epoch) lr scheduler
+ lr_sched.adjust_learning_rate(optimizer, data_iter_step / len(data_loader) + epoch, args)
+
+ # normalize image to [-1, 1]
+ x = x.to(device, non_blocking=True).to(torch.float32).div_(255)
+ x = x * 2.0 - 1.0
+ labels = labels.to(device, non_blocking=True)
+
+ with torch.no_grad():
+ if args.method == 'GRN_ind':
+ raw_features_, _, _ = vae.encode_for_raw_features(x.unsqueeze(2), scale_schedule=None, slice=True)
+ raw_features = raw_features_[0].squeeze(2)
+ x = raw_feature2label(raw_features, hbq_round=args.hbq_round)
+ elif args.method == 'GRN_bit':
+ raw_features_, _, _ = vae.encode_for_raw_features(x.unsqueeze(2), scale_schedule=None, slice=True)
+ raw_features = raw_features_[0].squeeze(2)
+ x = raw_feature2bit_label(raw_features, hbq_round=args.hbq_round)
+
+ with torch.amp.autocast('cuda', dtype=torch.bfloat16):
+ loss, t_bin2acc, t_bin2freq = model(x, labels)
+
+ import torch.distributed as dist
+ dist.all_reduce(t_bin2acc)
+ dist.all_reduce(t_bin2freq)
+ t_bin2acc = t_bin2acc / (t_bin2freq + 1e-8) * 100.
+
+ loss_value = loss.item()
+ if not math.isfinite(loss_value):
+ print("Loss is {}, stopping training".format(loss_value))
+ sys.exit(1)
+
+ optimizer.zero_grad()
+ loss.backward()
+ grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.clip_grad_norm)
+ optimizer.step()
+
+ torch.cuda.synchronize()
+
+ model_without_ddp.update_ema()
+
+ loss_value_reduce = misc.all_reduce_mean(loss_value)
+ grad_norm_reduce = misc.all_reduce_mean(grad_norm.item())
+
+ metric_logger.update(loss=loss_value_reduce)
+ lr = optimizer.param_groups[0]["lr"]
+ metric_logger.update(lr=lr)
+ metric_logger.update(grad_norm=grad_norm_reduce)
+
+ if log_writer is not None:
+ # Use epoch_1000x as the x-axis in TensorBoard to calibrate curves.
+ epoch_1000x = int((data_iter_step / len(data_loader) + epoch) * 1000)
+ if data_iter_step % args.log_freq == 0:
+ log_writer.add_scalar('train_loss', loss_value_reduce, epoch_1000x)
+ log_writer.add_scalar('lr', lr, epoch_1000x)
+ if args.wandb:
+ wandb_utils.log(
+ { "train loss": loss_value_reduce, "lr": lr, "grad_norm_t": grad_norm_reduce},
+ step=epoch_1000x
+ )
+ visual_dict = {}
+ for t_round in range(10):
+ if t_bin2freq[t_round] > 0:
+ visual_dict.update({f"accuracy/signal_{t_round*0.1:.1f}_{(t_round+1)*0.1:.1f}": t_bin2acc[t_round].item()})
+ visual_dict.update({f"frequency/signal_{t_round*0.1:.1f}_{(t_round+1)*0.1:.1f}": t_bin2freq[t_round].item()})
+ wandb_utils.log(visual_dict, step=epoch_1000x)
+
+
+def evaluate(model_without_ddp, args, epoch, batch_size=64, log_writer=None, vae=None):
+
+ model_without_ddp.eval()
+ world_size = misc.get_world_size()
+ local_rank = misc.get_rank()
+ num_steps = args.num_images // (batch_size * world_size) + 1
+
+ # Construct the folder name for saving generated images.
+ save_folder = os.path.join(
+ args.generation_dir,
+ f'Epoch{epoch:04d}',
+ "{}-steps{}-cfg{}-tau{}-interval{}-{}-image{}-res{}".format(
+ model_without_ddp.method, model_without_ddp.steps, model_without_ddp.cfg_scale, args.tau,
+ model_without_ddp.cfg_interval[0], model_without_ddp.cfg_interval[1], args.num_images, args.img_size
+ ),
+ 'images',
+ )
+ print("Save to:", save_folder)
+ if misc.get_rank() == 0 and not os.path.exists(save_folder):
+ os.makedirs(save_folder)
+
+ json_file = os.path.join(
+ os.path.dirname(os.environ.get('CKPT_FILE', '/tmp/res')),
+ f'testing/Epoch{epoch:04d}/images_{args.num_images}',
+ "{}-steps{}-cfg{}-tau{}-interval{}-{}-image{}-res{}".format(
+ model_without_ddp.method, model_without_ddp.steps, model_without_ddp.cfg_scale, args.tau,
+ model_without_ddp.cfg_interval[0], model_without_ddp.cfg_interval[1], args.num_images, args.img_size
+ ),
+ 'metrics.json',
+ )
+
+ # switch to ema params, hard-coded to be the first one
+ print("Switch to ema")
+ model_params_backup = [p.detach().clone() for p in model_without_ddp.parameters()]
+ if hasattr(model_without_ddp, 'module') and hasattr(model_without_ddp.module, 'ema_params1'):
+ ema_params = model_without_ddp.module.ema_params1
+ else:
+ ema_params = model_without_ddp.ema_params1
+
+ for param, ema_param in zip(model_without_ddp.parameters(), ema_params):
+ param.data.copy_(ema_param.data)
+
+ # ensure that the number of images per class is equal.
+ class_num = args.class_num
+ assert args.num_images % class_num == 0, "Number of images per class must be the same"
+ class_label_gen_world = np.arange(0, class_num).repeat(args.num_images // class_num)
+ class_label_gen_world = np.hstack([class_label_gen_world, np.zeros(50000)])
+
+ for i in range(num_steps):
+ print("Generation step {}/{}".format(i, num_steps))
+
+ start_idx = world_size * batch_size * i + local_rank * batch_size
+ end_idx = start_idx + batch_size
+ labels_gen = class_label_gen_world[start_idx:end_idx]
+ labels_gen = torch.Tensor(labels_gen).long().cuda()
+
+ with torch.amp.autocast('cuda', dtype=torch.bfloat16):
+ sampled_images = model_without_ddp.generate(labels_gen)
+
+ if args.method == 'GRN_ind':
+ from grn.utils_c2i.hbq_util_c2i import label2quant_features
+ sampled_images = label2quant_features(sampled_images, hbq_round=args.hbq_round)
+ sampled_images = vae.decode(sampled_images.unsqueeze(2), slice=True).squeeze(2)
+ elif args.method == 'GRN_bit':
+ from grn.utils_c2i.hbq_util_c2i import bit_label2raw_feature
+ sampled_images = bit_label2raw_feature(sampled_images, hbq_round=args.hbq_round)
+ sampled_images = vae.decode(sampled_images.unsqueeze(2), slice=True).squeeze(2)
+
+ torch.distributed.barrier()
+
+ # denormalize images
+ sampled_images = (sampled_images + 1) / 2
+ sampled_images = sampled_images.detach().cpu()
+
+ # distributed save images
+ for b_id in range(sampled_images.size(0)):
+ img_id = i * sampled_images.size(0) * world_size + local_rank * sampled_images.size(0) + b_id
+ if img_id >= args.num_images:
+ break
+ gen_img = np.round(np.clip(sampled_images[b_id].numpy().transpose([1, 2, 0]) * 255, 0, 255))
+ gen_img = gen_img.astype(np.uint8)[:, :, ::-1]
+ cv2.imwrite(os.path.join(save_folder, '{}.png'.format(str(img_id).zfill(5))), gen_img)
+
+ torch.distributed.barrier()
+
+ # back to no ema
+ print("Switch back from ema")
+ for param, backup in zip(model_without_ddp.parameters(), model_params_backup):
+ param.data.copy_(backup.data)
+ del model_params_backup
+ torch.cuda.empty_cache()
+
+ # compute FID and IS
+ if log_writer is not None:
+ if args.img_size == 256:
+ fid_statistics_file = 'fid_stats/jit_in256_stats.npz'
+ elif args.img_size == 512:
+ fid_statistics_file = 'fid_stats/jit_in512_stats.npz'
+ else:
+ raise NotImplementedError
+ metrics_dict = torch_fidelity.calculate_metrics(
+ input1=save_folder,
+ input2=None,
+ fid_statistics_file=fid_statistics_file,
+ cuda=True,
+ isc=True,
+ fid=True,
+ kid=False,
+ prc=False,
+ verbose=False,
+ restrict_data_size=-1,
+ shuffle=False,
+ )
+ fid = metrics_dict['frechet_inception_distance']
+ inception_score = metrics_dict['inception_score_mean']
+ postfix = "_cfg{}_res{}".format(model_without_ddp.cfg_scale, args.img_size)
+ log_writer.add_scalar('fid{}'.format(postfix), fid, epoch)
+ log_writer.add_scalar('is{}'.format(postfix), inception_score, epoch)
+ if args.wandb:
+ wandb_utils.log(
+ {"fid": fid, "is": inception_score},
+ step=epoch
+ )
+ print("FID: {:.4f}, Inception Score: {:.4f}".format(fid, inception_score))
+ os.makedirs(os.path.dirname(json_file), exist_ok=True)
+ with open(json_file, 'w') as f:
+ json.dump(metrics_dict, f)
+ if args.delete_images:
+ shutil.rmtree(save_folder)
+
+ torch.distributed.barrier()
diff --git a/grn/utils_c2i/hbq_util_c2i.py b/grn/utils_c2i/hbq_util_c2i.py
new file mode 100644
index 0000000000000000000000000000000000000000..48c4a41b4c849c7b8d756741da30f08fca976897
--- /dev/null
+++ b/grn/utils_c2i/hbq_util_c2i.py
@@ -0,0 +1,49 @@
+import torch
+
+
+def label2quant_features(pred_sample_labels, hbq_round):
+ approx_signal = 0.
+ pred_sample_labels = pred_sample_labels.to(torch.long)
+ for round_ind in range(hbq_round):
+ interval = (1/2)**(round_ind+1)
+ base = 2**(hbq_round-1-round_ind)
+ approx_signal = approx_signal + interval * torch.where(pred_sample_labels>=base, +1, -1)
+ pred_sample_labels = pred_sample_labels % base
+ return approx_signal
+
+def raw_feature2label(feature, hbq_round):
+ quant_features = 0.
+ labels = 0
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ labels = labels * 2 + torch.where(feature > quant_features, 1, 0)
+ quant_features = quant_features + torch.where(feature > quant_features, interval, -interval)
+ return labels
+
+def raw_feature2bit_label(feature, hbq_round):
+ quant_features = 0.
+ labels = []
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ labels.append(torch.where(feature > quant_features, 1, 0))
+ quant_features = quant_features + torch.where(feature > quant_features, interval, -interval)
+ labels = torch.stack(labels, dim=1) # [B,hbq_round,d,h,w]
+ B, _, d, h, w = labels.shape
+ labels = labels.reshape(B, hbq_round * d, h, w)
+ return labels
+
+def bit_label2raw_feature(bit_labels, hbq_round):
+ B, hbq_round_mul_d, h, w = bit_labels.shape
+ d = hbq_round_mul_d // hbq_round
+ bit_labels = bit_labels.reshape(B, hbq_round, d, h, w).to(torch.long)
+ raw_features = 0.
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ raw_features = raw_features + interval * torch.where(bit_labels[:,round_ind] == 1, +1, -1)
+ return raw_features
+
+def multiclass_labels2onehot_input(labels, num_classes):
+ B,d,h,w = labels.shape
+ onehot_input = torch.nn.functional.one_hot(labels.to(torch.long), num_classes) # [B,d,h,w] -> [B,d,h,w,num_classes]
+ onehot_input = onehot_input.permute(0,4,1,2,3).reshape(B,num_classes*d,h,w).float() # [B,d,h,w,num_classes] -> [B,num_classes,d,h,w] -> [B,num_classes*d,h,w]
+ return onehot_input
diff --git a/grn/utils_c2i/lr_sched.py b/grn/utils_c2i/lr_sched.py
new file mode 100644
index 0000000000000000000000000000000000000000..f4c4e610a48f0828100123b07f34e3a8a967dd48
--- /dev/null
+++ b/grn/utils_c2i/lr_sched.py
@@ -0,0 +1,27 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+
+import math
+
+
+def adjust_learning_rate(optimizer, epoch, args):
+ """Decay the learning rate with half-cycle cosine after warmup"""
+ if epoch < args.warmup_epochs:
+ lr = args.lr * epoch / args.warmup_epochs
+ else:
+ if args.lr_schedule == "constant":
+ lr = args.lr
+ elif args.lr_schedule == "cosine":
+ lr = args.min_lr + (args.lr - args.min_lr) * 0.5 * \
+ (1. + math.cos(math.pi * (epoch - args.warmup_epochs) / (args.epochs - args.warmup_epochs)))
+ else:
+ raise NotImplementedError
+ for param_group in optimizer.param_groups:
+ if "lr_scale" in param_group:
+ param_group["lr"] = lr * param_group["lr_scale"]
+ else:
+ param_group["lr"] = lr
+ return lr
diff --git a/grn/utils_c2i/misc.py b/grn/utils_c2i/misc.py
new file mode 100644
index 0000000000000000000000000000000000000000..d2ca491887ee35d8259bc78026f5324dd4f29be5
--- /dev/null
+++ b/grn/utils_c2i/misc.py
@@ -0,0 +1,337 @@
+# Copyright (c) Meta Platforms, Inc. and affiliates.
+# All rights reserved.
+
+# This source code is licensed under the license found in the
+# LICENSE file in the root directory of this source tree.
+# --------------------------------------------------------
+# References:
+# DeiT: https://github.com/facebookresearch/deit
+# BEiT: https://github.com/microsoft/unilm/tree/master/beit
+# --------------------------------------------------------
+
+import builtins
+import datetime
+import os
+import time
+from collections import defaultdict, deque
+from pathlib import Path
+import copy
+
+import torch
+import torch.distributed as dist
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.fsdp import StateDictType, FullStateDictConfig
+
+
+class SmoothedValue(object):
+ """Track a series of values and provide access to smoothed values over a
+ window or the global series average.
+ """
+
+ def __init__(self, window_size=20, fmt=None):
+ if fmt is None:
+ fmt = "{median:.4f} ({global_avg:.4f})"
+ self.deque = deque(maxlen=window_size)
+ self.total = 0.0
+ self.count = 0
+ self.fmt = fmt
+
+ def update(self, value, n=1):
+ self.deque.append(value)
+ self.count += n
+ self.total += value * n
+
+ def synchronize_between_processes(self):
+ """
+ Warning: does not synchronize the deque!
+ """
+ if not is_dist_avail_and_initialized():
+ return
+ t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda')
+ dist.barrier()
+ dist.all_reduce(t)
+ t = t.tolist()
+ self.count = int(t[0])
+ self.total = t[1]
+
+ @property
+ def median(self):
+ d = torch.tensor(list(self.deque))
+ return d.median().item()
+
+ @property
+ def avg(self):
+ d = torch.tensor(list(self.deque), dtype=torch.float32)
+ return d.mean().item()
+
+ @property
+ def global_avg(self):
+ return self.total / self.count
+
+ @property
+ def max(self):
+ return max(self.deque)
+
+ @property
+ def value(self):
+ return self.deque[-1]
+
+ def __str__(self):
+ return self.fmt.format(
+ median=self.median,
+ avg=self.avg,
+ global_avg=self.global_avg,
+ max=self.max,
+ value=self.value)
+
+
+class MetricLogger(object):
+ def __init__(self, delimiter="\t"):
+ self.meters = defaultdict(SmoothedValue)
+ self.delimiter = delimiter
+
+ def update(self, **kwargs):
+ for k, v in kwargs.items():
+ if v is None:
+ continue
+ if isinstance(v, torch.Tensor):
+ v = v.item()
+ assert isinstance(v, (float, int))
+ self.meters[k].update(v)
+
+ def __getattr__(self, attr):
+ if attr in self.meters:
+ return self.meters[attr]
+ if attr in self.__dict__:
+ return self.__dict__[attr]
+ raise AttributeError("'{}' object has no attribute '{}'".format(
+ type(self).__name__, attr))
+
+ def __str__(self):
+ loss_str = []
+ for name, meter in self.meters.items():
+ loss_str.append(
+ "{}: {}".format(name, str(meter))
+ )
+ return self.delimiter.join(loss_str)
+
+ def synchronize_between_processes(self):
+ for meter in self.meters.values():
+ meter.synchronize_between_processes()
+
+ def add_meter(self, name, meter):
+ self.meters[name] = meter
+
+ def log_every(self, iterable, print_freq, header=None):
+ i = 0
+ if not header:
+ header = ''
+ start_time = time.time()
+ end = time.time()
+ iter_time = SmoothedValue(fmt='{avg:.4f}')
+ data_time = SmoothedValue(fmt='{avg:.4f}')
+ space_fmt = ':' + str(len(str(len(iterable)))) + 'd'
+ log_msg = [
+ header,
+ '[{0' + space_fmt + '}/{1}]',
+ 'eta: {eta}',
+ '{meters}',
+ 'time: {time}',
+ 'data: {data}'
+ ]
+ if torch.cuda.is_available():
+ log_msg.append('max mem: {memory:.0f}')
+ log_msg = self.delimiter.join(log_msg)
+ MB = 1024.0 * 1024.0
+ for obj in iterable:
+ data_time.update(time.time() - end)
+ yield obj
+ iter_time.update(time.time() - end)
+ if i % print_freq == 0 or i == len(iterable) - 1:
+ eta_seconds = iter_time.global_avg * (len(iterable) - i)
+ eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))
+ if torch.cuda.is_available():
+ print(log_msg.format(
+ i, len(iterable), eta=eta_string,
+ meters=str(self),
+ time=str(iter_time), data=str(data_time),
+ memory=torch.cuda.max_memory_allocated() / MB))
+ else:
+ print(log_msg.format(
+ i, len(iterable), eta=eta_string,
+ meters=str(self),
+ time=str(iter_time), data=str(data_time)))
+ i += 1
+ end = time.time()
+ total_time = time.time() - start_time
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
+ print('{} Total time: {} ({:.4f} s / it)'.format(
+ header, total_time_str, total_time / len(iterable)))
+
+
+def setup_for_distributed(is_master):
+ """
+ This function disables printing when not in master process
+ """
+ builtin_print = builtins.print
+
+ def print(*args, **kwargs):
+ force = kwargs.pop('force', False)
+ force = force or (get_world_size() > 8)
+ if is_master or force:
+ now = datetime.datetime.now().time()
+ builtin_print('[{}] '.format(now), end='') # print with time stamp
+ builtin_print(*args, **kwargs)
+
+ builtins.print = print
+
+
+def is_dist_avail_and_initialized():
+ if not dist.is_available():
+ return False
+ if not dist.is_initialized():
+ return False
+ return True
+
+
+def get_world_size():
+ if not is_dist_avail_and_initialized():
+ return 1
+ return dist.get_world_size()
+
+
+def get_rank():
+ if not is_dist_avail_and_initialized():
+ return 0
+ return dist.get_rank()
+
+
+def is_main_process():
+ return get_rank() == 0
+
+
+def save_on_master(*args, **kwargs):
+ if is_main_process():
+ torch.save(*args, **kwargs)
+
+
+def init_distributed_mode(args):
+ if args.dist_on_itp:
+ args.rank = int(os.environ['OMPI_COMM_WORLD_RANK'])
+ args.world_size = int(os.environ['OMPI_COMM_WORLD_SIZE'])
+ args.gpu = int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK'])
+ args.dist_url = "tcp://%s:%s" % (os.environ['MASTER_ADDR'], os.environ['MASTER_PORT'])
+ os.environ['LOCAL_RANK'] = str(args.gpu)
+ os.environ['RANK'] = str(args.rank)
+ os.environ['WORLD_SIZE'] = str(args.world_size)
+ # ["RANK", "WORLD_SIZE", "MASTER_ADDR", "MASTER_PORT", "LOCAL_RANK"]
+ elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
+ args.rank = int(os.environ["RANK"])
+ args.world_size = int(os.environ['WORLD_SIZE'])
+ args.gpu = int(os.environ['LOCAL_RANK'])
+ elif 'SLURM_PROCID' in os.environ:
+ args.rank = int(os.environ['SLURM_PROCID'])
+ args.gpu = args.rank % torch.cuda.device_count()
+ else:
+ print('Not using distributed mode')
+ setup_for_distributed(is_master=True) # hack
+ args.distributed = False
+ return
+
+ args.distributed = True
+
+ torch.cuda.set_device(args.gpu)
+ args.dist_backend = 'nccl'
+ print('| distributed init (rank {}): {}, gpu {}'.format(
+ args.rank, args.dist_url, args.gpu), flush=True)
+ torch.distributed.init_process_group(backend=args.dist_backend, init_method=args.dist_url,
+ world_size=args.world_size, rank=args.rank)
+ torch.distributed.barrier()
+ setup_for_distributed(args.rank == 0)
+
+
+def add_weight_decay(model, weight_decay=0, skip_list=()):
+ decay = []
+ no_decay = []
+ for name, param in model.named_parameters():
+ if not param.requires_grad:
+ continue # frozen weights
+ if len(param.shape) == 1 or name.endswith(".bias") or name in skip_list or 'diffloss' in name:
+ no_decay.append(param) # no weight decay on bias, norm and diffloss
+ else:
+ decay.append(param)
+ return [
+ {'params': no_decay, 'weight_decay': 0.},
+ {'params': decay, 'weight_decay': weight_decay}]
+
+
+def save_model(args, model_without_ddp, optimizer, epoch, epoch_name=None):
+ if epoch_name is None:
+ epoch_name = str(epoch)
+ output_dir = Path(args.output_dir)
+ checkpoint_path = output_dir / ('checkpoint-%s.pth' % epoch_name)
+
+ if isinstance(model_without_ddp, FSDP):
+ save_policy = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
+ with FSDP.state_dict_type(model_without_ddp, StateDictType.FULL_STATE_DICT, save_policy):
+ model_state = model_without_ddp.state_dict()
+
+ # Backup current params
+ current_params = [p.detach().clone() for p in model_without_ddp.parameters()]
+
+ # EMA 1
+ for p, ema_p in zip(model_without_ddp.parameters(), model_without_ddp.module.ema_params1):
+ p.data.copy_(ema_p.data)
+ ema_state1 = model_without_ddp.state_dict()
+
+ # EMA 2
+ for p, ema_p in zip(model_without_ddp.parameters(), model_without_ddp.module.ema_params2):
+ p.data.copy_(ema_p.data)
+ ema_state2 = model_without_ddp.state_dict()
+
+ # Restore
+ for p, backup in zip(model_without_ddp.parameters(), current_params):
+ p.data.copy_(backup.data)
+ del current_params
+
+ # Optimizer state dict
+ opt_state = FSDP.full_optim_state_dict(model_without_ddp, optimizer)
+
+ to_save = {
+ 'model': model_state,
+ 'optimizer': opt_state,
+ 'epoch': epoch,
+ 'args': args,
+ 'model_ema1': ema_state1,
+ 'model_ema2': ema_state2,
+ }
+ else:
+ to_save = {
+ 'model': model_without_ddp.state_dict(),
+ 'optimizer': optimizer.state_dict(),
+ 'epoch': epoch,
+ 'args': args,
+ }
+
+ # ema
+ ema_state_dict1 = copy.deepcopy(model_without_ddp.state_dict())
+ ema_state_dict2 = copy.deepcopy(model_without_ddp.state_dict())
+ for i, (name, _value) in enumerate(model_without_ddp.named_parameters()):
+ assert name in ema_state_dict1 and name in ema_state_dict2
+ ema_state_dict1[name] = model_without_ddp.ema_params1[i]
+ ema_state_dict2[name] = model_without_ddp.ema_params2[i]
+ to_save['model_ema1'] = ema_state_dict1
+ to_save['model_ema2'] = ema_state_dict2
+
+ save_on_master(to_save, checkpoint_path)
+
+
+def all_reduce_mean(x):
+ world_size = get_world_size()
+ if world_size > 1:
+ x_reduce = torch.tensor(x).cuda()
+ dist.all_reduce(x_reduce)
+ x_reduce /= world_size
+ return x_reduce.item()
+ else:
+ return x
\ No newline at end of file
diff --git a/grn/utils_c2i/model_util.py b/grn/utils_c2i/model_util.py
new file mode 100644
index 0000000000000000000000000000000000000000..f772c134d9dd6400d9c99553e093a1db9ce09fdf
--- /dev/null
+++ b/grn/utils_c2i/model_util.py
@@ -0,0 +1,201 @@
+# --------------------------------------------------------
+# References:
+# Lightning-DiT: https://github.com/hustvl/LightningDiT
+# --------------------------------------------------------
+
+from math import pi
+
+import torch
+from torch import nn
+import numpy as np
+
+from einops import rearrange, repeat
+
+
+def broadcat(tensors, dim = -1):
+ num_tensors = len(tensors)
+ shape_lens = set(list(map(lambda t: len(t.shape), tensors)))
+ assert len(shape_lens) == 1, 'tensors must all have the same number of dimensions'
+ shape_len = list(shape_lens)[0]
+ dim = (dim + shape_len) if dim < 0 else dim
+ dims = list(zip(*map(lambda t: list(t.shape), tensors)))
+ expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
+ assert all([*map(lambda t: len(set(t[1])) <= 2, expandable_dims)]), 'invalid dimensions for broadcastable concatentation'
+ max_dims = list(map(lambda t: (t[0], max(t[1])), expandable_dims))
+ expanded_dims = list(map(lambda t: (t[0], (t[1],) * num_tensors), max_dims))
+ expanded_dims.insert(dim, (dim, dims[dim]))
+ expandable_shapes = list(zip(*map(lambda t: t[1], expanded_dims)))
+ tensors = list(map(lambda t: t[0].expand(*t[1]), zip(tensors, expandable_shapes)))
+ return torch.cat(tensors, dim = dim)
+
+
+def rotate_half(x):
+ x = rearrange(x, '... (d r) -> ... d r', r = 2)
+ x1, x2 = x.unbind(dim = -1)
+ x = torch.stack((-x2, x1), dim = -1)
+ return rearrange(x, '... d r -> ... (d r)')
+
+
+class VisionRotaryEmbedding(nn.Module):
+ def __init__(
+ self,
+ dim,
+ pt_seq_len,
+ ft_seq_len=None,
+ custom_freqs = None,
+ freqs_for = 'lang',
+ theta = 10000,
+ max_freq = 10,
+ num_freqs = 1,
+ ):
+ super().__init__()
+ if custom_freqs:
+ freqs = custom_freqs
+ elif freqs_for == 'lang':
+ freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim))
+ elif freqs_for == 'pixel':
+ freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi
+ elif freqs_for == 'constant':
+ freqs = torch.ones(num_freqs).float()
+ else:
+ raise ValueError(f'unknown modality {freqs_for}')
+
+ if ft_seq_len is None: ft_seq_len = pt_seq_len
+ t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len
+
+ freqs_h = torch.einsum('..., f -> ... f', t, freqs)
+ freqs_h = repeat(freqs_h, '... n -> ... (n r)', r = 2)
+
+ freqs_w = torch.einsum('..., f -> ... f', t, freqs)
+ freqs_w = repeat(freqs_w, '... n -> ... (n r)', r = 2)
+
+ freqs = broadcat((freqs_h[:, None, :], freqs_w[None, :, :]), dim = -1)
+
+ self.register_buffer("freqs_cos", freqs.cos())
+ self.register_buffer("freqs_sin", freqs.sin())
+
+ def forward(self, t, start_index = 0):
+ rot_dim = self.freqs_cos.shape[-1]
+ end_index = start_index + rot_dim
+ assert rot_dim <= t.shape[-1], f'feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}'
+ t_left, t, t_right = t[..., :start_index], t[..., start_index:end_index], t[..., end_index:]
+ t = (t * self.freqs_cos) + (rotate_half(t) * self.freqs_sin)
+ return torch.cat((t_left, t, t_right), dim = -1)
+
+
+class VisionRotaryEmbeddingFast(nn.Module):
+ def __init__(
+ self,
+ dim,
+ pt_seq_len=16,
+ ft_seq_len=None,
+ custom_freqs = None,
+ freqs_for = 'lang',
+ theta = 10000,
+ max_freq = 10,
+ num_freqs = 1,
+ num_cls_token = 0
+ ):
+ super().__init__()
+ if custom_freqs:
+ freqs = custom_freqs
+ elif freqs_for == 'lang':
+ freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim))
+ elif freqs_for == 'pixel':
+ freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi
+ elif freqs_for == 'constant':
+ freqs = torch.ones(num_freqs).float()
+ else:
+ raise ValueError(f'unknown modality {freqs_for}')
+
+ if ft_seq_len is None: ft_seq_len = pt_seq_len
+ t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len
+
+ freqs = torch.einsum('..., f -> ... f', t, freqs)
+ freqs = repeat(freqs, '... n -> ... (n r)', r = 2)
+ freqs = broadcat((freqs[:, None, :], freqs[None, :, :]), dim = -1)
+
+ if num_cls_token > 0:
+ freqs_flat = freqs.view(-1, freqs.shape[-1]) # [N_img, D]
+ cos_img = freqs_flat.cos()
+ sin_img = freqs_flat.sin()
+
+ # prepend in-context cls token
+ N_img, D = cos_img.shape
+ cos_pad = torch.ones(num_cls_token, D, dtype=cos_img.dtype, device=cos_img.device)
+ sin_pad = torch.zeros(num_cls_token, D, dtype=sin_img.dtype, device=sin_img.device)
+
+ self.freqs_cos = torch.cat([cos_pad, cos_img], dim=0).cuda() # [N_cls+N_img, D]
+ self.freqs_sin = torch.cat([sin_pad, sin_img], dim=0).cuda()
+ else:
+ self.freqs_cos = freqs.cos().view(-1, freqs.shape[-1]).cuda()
+ self.freqs_sin = freqs.sin().view(-1, freqs.shape[-1]).cuda()
+
+ def forward(self, t): return t * self.freqs_cos + rotate_half(t) * self.freqs_sin
+
+
+class RMSNorm(nn.Module):
+ def __init__(self, hidden_size, eps=1e-6):
+ """
+ LlamaRMSNorm is equivalent to T5LayerNorm
+ """
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(hidden_size))
+ self.variance_epsilon = eps
+
+ def forward(self, hidden_states):
+ input_dtype = hidden_states.dtype
+ hidden_states = hidden_states.to(torch.float32)
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
+ return (self.weight * hidden_states).to(input_dtype)
+
+
+def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0):
+ """
+ grid_size: int of the grid height and width
+ return:
+ pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
+ """
+ grid_h = np.arange(grid_size, dtype=np.float32)
+ grid_w = np.arange(grid_size, dtype=np.float32)
+ grid = np.meshgrid(grid_w, grid_h) # here w goes first
+ grid = np.stack(grid, axis=0)
+
+ grid = grid.reshape([2, 1, grid_size, grid_size])
+ pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
+ if cls_token and extra_tokens > 0:
+ pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
+ return pos_embed
+
+
+def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
+ assert embed_dim % 2 == 0
+
+ # use half of dimensions to encode grid_h
+ emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
+ emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
+
+ emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
+ return emb
+
+
+def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
+ """
+ embed_dim: output dimension for each position
+ pos: a list of positions to be encoded: size (M,)
+ out: (M, D)
+ """
+ assert embed_dim % 2 == 0
+ omega = np.arange(embed_dim // 2, dtype=np.float64)
+ omega /= embed_dim / 2.
+ omega = 1. / 10000**omega # (D/2,)
+
+ pos = pos.reshape(-1) # (M,)
+ out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
+
+ emb_sin = np.sin(out) # (M, D/2)
+ emb_cos = np.cos(out) # (M, D/2)
+
+ emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
+ return emb
\ No newline at end of file
diff --git a/grn/utils_t2iv/amp_opt.py b/grn/utils_t2iv/amp_opt.py
new file mode 100644
index 0000000000000000000000000000000000000000..debb004243e164a247e47d2eb6af47c8f416e244
--- /dev/null
+++ b/grn/utils_t2iv/amp_opt.py
@@ -0,0 +1,188 @@
+import math
+import os
+import signal
+import sys
+import time
+from typing import List, Optional, Tuple, Union
+
+import torch
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+# from memory_profiler import profile
+
+import grn.utils_t2iv.dist as dist
+from grn.utils import misc
+
+class NullCtx:
+ def __enter__(self):
+ pass
+
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ pass
+
+
+def handle_timeout(signum, frame):
+ raise TimeoutError('took too long')
+
+
+def per_param_clip_grad_norm_(parameters, thresh: float, stable=False, fp=None) -> (float, float):
+ skipped, max_grad = [], 0
+ for pi, p in enumerate(parameters):
+ if p.grad is not None:
+ g = p.grad.data.norm(2).item() + 1e-7
+ max_grad = max(max_grad, g)
+ clip_coef = thresh / g
+ if clip_coef < 1:
+ if stable and clip_coef < 0.2:
+ skipped.append(clip_coef)
+ p.grad.data.mul_(0) # todo NOTE: inf.mul_(0)==nan will shrink the scale ratio, but inf.zero_()==0 won't
+ else:
+ p.grad.data.mul_(clip_coef)
+
+ # if fp is not None: fp.write(f'[per_param_clip_grad_norm_:47] finished.\n'); fp.flush()
+ return 0 if len(skipped) == 0 else math.log10(max(min(skipped), 1e-7)), max_grad
+
+
+class AmpOptimizer:
+ def __init__(
+ self,
+ model_name_3letters: str, mixed_precision: int,
+ optimizer: torch.optim.Optimizer, model_maybe_fsdp: Union[torch.nn.Module, FSDP],
+ r_accu: float, grad_clip: float, zero: int,
+ ):
+ self.enable_amp = mixed_precision > 0
+ self.zero = zero
+ if self.enable_amp:
+ self.using_fp16_rather_bf16 = mixed_precision != 2
+ self.max_sc = float(mixed_precision if mixed_precision > 128 else 32768)
+
+ # todo: on both V100 and A100, torch.get_autocast_gpu_dtype() returns fp16, not bf16.
+ self.amp_ctx = torch.autocast('cuda', enabled=True, dtype=torch.float16 if self.using_fp16_rather_bf16 else torch.bfloat16, cache_enabled=self.zero == 0) # todo: cache_enabled=False
+ if self.using_fp16_rather_bf16:
+ self.scaler = torch.cuda.amp.GradScaler(init_scale=2. ** 11, growth_interval=1000)
+ else:
+ self.scaler = None
+ else:
+ self.using_fp16_rather_bf16 = True
+ self.amp_ctx = NullCtx()
+ self.scaler = None
+
+ t = torch.zeros(dist.get_world_size())
+ t[dist.get_rank()] = float(self.enable_amp)
+ dist.allreduce(t)
+ assert round(t.sum().item()) in {0, dist.get_world_size()}, f'enable_amp: {t}'
+
+ t = torch.zeros(dist.get_world_size())
+ t[dist.get_rank()] = float(self.using_fp16_rather_bf16)
+ dist.allreduce(t)
+ assert round(t.sum().item()) in {0, dist.get_world_size()}, f'using_fp16_rather_bf16: {t}'
+
+ self.model_name_3letters = model_name_3letters
+ self.optimizer, self.model_maybe_fsdp = optimizer, model_maybe_fsdp
+ self.r_accu = r_accu
+
+ self.paras = self.names = ... # todo: solve EMA-related codes
+
+ self.grad_clip, self.grad_clip_we = grad_clip, 0 # todo: disable wclip
+ if self.grad_clip > 100:
+ self.grad_clip %= 100
+ self.per_param = True
+ else:
+ self.per_param = False
+ self.per_param = False # todo: disable wclip
+
+ self.early_clipping = grad_clip > 0 and not hasattr(optimizer, 'global_grad_norm')
+ self.late_clipping = grad_clip > 0 and hasattr(optimizer, 'global_grad_norm') # deepspeed's optimizer
+
+ self.fp = None
+ self.last_orig_norm: torch.Tensor = torch.tensor(0.1)
+
+ @torch.no_grad()
+ def log_param(self, ep: int):
+ if self.zero == 0:
+ for name, values in get_param_for_log(self.model_name_3letters, self.model_maybe_fsdp.named_parameters()).items():
+ values: List[float]
+ if len(values) == 1: # e.g., cls token will only have one value
+ values.append(values[0])
+ else:
+ ...
+ # todo: log params
+
+ # @profile(precision=4, stream=open('amp_sc.log', 'w+'))
+ def backward_clip_step(
+ self, ep: int, it: int, g_it: int, stepping: bool, logging_params: bool, loss: torch.Tensor, clip_decay_ratio=1, stable=False,
+ ) -> Tuple[torch.Tensor, Optional[float]]:
+ # hj note: self.r_accu=1. self.scaler=None stepping=True self.fp=None self.early_clipping=True self.late_clipping=False
+ # backward
+ loss = loss.mul(self.r_accu) # r_accu == 1.0 / n_gradient_accumulation
+ orig_norm = scaler_sc = None
+ # if self.fp is not None:
+ # if g_it % 20 == 0: self.fp.seek(0); self.fp.truncate(0)
+ if self.scaler is not None:
+ self.scaler.scale(loss).backward(retain_graph=False, create_graph=False) # retain_graph=retain_graph, create_graph=create_graph
+ else:
+ loss.backward(retain_graph=False, create_graph=False)
+ # if self.fp is not None: self.fp.write(f'[backward_clip_step:131] [it{it}, g_it{g_it}] after backward\n'); self.fp.flush()
+
+ # clip gradients then step optimizer
+ if stepping:
+ if self.scaler is not None: self.scaler.unscale_(self.optimizer) # now the gradient can be correctly got
+ # if self.fp is not None: self.fp.write(f'[backward_clip_step:137] [it{it}, g_it{g_it}] after scaler.unscale_\n'); self.fp.flush()
+
+ skipped, orig_norm = 0, self.last_orig_norm
+ # try:
+ if self.fp is not None:
+ if g_it % 10 == 0: self.fp.seek(0); self.fp.truncate(0)
+ self.fp.write(f'\n'); self.fp.flush()
+ if self.early_clipping:
+ c = self.grad_clip
+ if self.zero:
+ orig_norm: Optional[torch.Tensor] = self.model_maybe_fsdp.clip_grad_norm_(c)
+ else:
+ orig_norm: Optional[torch.Tensor] = torch.nn.utils.clip_grad_norm_(self.model_maybe_fsdp.parameters(), c)
+
+ # if self.fp is not None: self.fp.write(f'[backward_clip_step:175] [it{it}, g_it{g_it}] before opt step\n'); self.fp.flush()
+ if self.scaler is not None:
+ self.scaler: torch.cuda.amp.GradScaler
+ if self.zero:
+ # synchronize found_inf_per_device before calling step, so that even if only some ranks found inf on their sharded params, all other ranks will know
+ # otherwise, when saving FSDP optimizer state, it will cause AssertionError saying "Different ranks have different values for step."
+ for optimizer_state in self.scaler._per_optimizer_states.values():
+ for t in optimizer_state['found_inf_per_device'].values():
+ dist.allreduce(t) # ideally, each rank only has one single t; so no need to use async allreduce
+
+ self.scaler.step(self.optimizer)
+ scaler_sc: Optional[float] = self.scaler.get_scale()
+ if scaler_sc > self.max_sc: # fp16 will overflow when >65536, so multiply 32768 could be dangerous
+ # print(f'[fp16 scaling] too large loss scale {scaler_sc}! (clip to {self.max_sc:g})')
+ self.scaler.update(new_scale=self.max_sc)
+ else:
+ self.scaler.update()
+ try:
+ scaler_sc = float(math.log2(scaler_sc))
+ except Exception as e:
+ print(f'[scaler_sc = {scaler_sc}]\n' * 15, flush=True)
+ time.sleep(1)
+ print(f'[scaler_sc = {scaler_sc}]\n' * 15, flush=True)
+ raise e
+ else:
+ self.optimizer.step()
+
+ if self.late_clipping:
+ orig_norm: Optional[torch.Tensor] = self.optimizer.global_grad_norm
+ self.last_orig_norm = orig_norm
+ # no zero_grad calling here, gonna log those gradients!
+ return orig_norm, scaler_sc
+
+ def state_dict(self):
+ return {
+ 'optimizer': self.optimizer.state_dict()
+ } if self.scaler is None else {
+ 'scaler': self.scaler.state_dict(),
+ 'optimizer': self.optimizer.state_dict()
+ }
+
+ def load_state_dict(self, state, strict=True):
+ if self.scaler is not None:
+ try: self.scaler.load_state_dict(state['scaler'])
+ except Exception as e: print(f'[fp16 load_state_dict err] {e}')
+ self.optimizer.load_state_dict(state['optimizer'])
diff --git a/grn/utils_t2iv/arg_util.py b/grn/utils_t2iv/arg_util.py
new file mode 100644
index 0000000000000000000000000000000000000000..c1f0eba787982c0a2869e14704edfb9c614cbf7c
--- /dev/null
+++ b/grn/utils_t2iv/arg_util.py
@@ -0,0 +1,289 @@
+import json
+import math
+import os
+import random
+import subprocess
+import sys
+import time
+from collections import OrderedDict, deque
+from typing import Optional, Union, Literal
+
+import numpy as np
+import torch
+from tap import Tap
+
+import grn.utils_t2iv.dist as dist
+from grn.utils_t2iv.sequence_parallel import SequenceParallelManager as sp_manager
+
+
+class Args(Tap):
+ local_out_path: str = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'local_output') # directory for save checkpoints
+ data_path: str = '' # dataset path
+ video_fps: int = 24 # video fps
+ video_frames: int = 1 # video frames
+ hdfs_mode: str = 'read' # hdfs_mode
+ bed: str = '' # bed directory for copy checkpoints apart from local_out_path
+ vae_path: str = '' # VAE ckpt
+ exp_name: str = '' # experiment name
+ model: str = '' # for VAE training, 'b' or any other for GPT training
+ short_cap_prob: float = 0.2 # prob for training with short captions
+ project_name: str = '' # name of wandb project
+ tf32: bool = True # whether to use TensorFloat32
+ auto_resume: bool = True # whether to automatically resume from the last checkpoint found in args.bed
+ rush_resume: str = '' # pretrained checkpoint
+ enable_hybrid_shard: int = 0 # whether to use hybrid FSDP
+ inner_shard_degree: int = 1 # inner degree for FSDP
+ zero: int = 0 # ds zero
+ enable_checkpointing: str = None # checkpointing strategy: full-block, self-attn
+ pad_to_multiplier: int = 1 # >1 for padding the seq len to a multiplier of this
+ log_every_iter: bool = False
+ checkpoint_type: str = 'torch' # checkpoint_type: torch, onmistore
+ device: str = 'cpu'
+ is_master_node: bool = None
+ # dir
+ log_txt_path: str = ''
+ t5_path: str = ''
+ online_t5: bool = True # whether to use online t5 or load local features
+ # GPT
+ sdpa_mem: bool = True # whether to use with torch.backends.cuda.sdp_kernel(enable_flash=False, enable_math=False, enable_mem_efficient=True)
+ tfast: int = 0 # compile GPT
+ tau: float = 1 # tau of self attention in GPT
+ tp: float = 0.0 # top-p
+ tk: float = 0.0 # top-k
+ drop_condition_prob: float = 0.1 # >0: classifier-free guidance, drop cond with prob cfg
+ fp16: int = 0 # 1: fp16, 2: bf16, >2: fp16's max scaling multiplier todo: 记得让quantize相关的feature都强制fp32!另外residueal最好也是fp32(根据flash-attention)nn.Conv2d有一个参数是use_float16?
+ use_flex_attn: bool = False # whether to use flex_attn to speedup training
+ tlr: float = 2e-5 # learning rate
+ twd: float = 0.005 # vqgan: 0.01
+ twde: float = 0
+ ep: int = 100
+ wp: float = 0
+ wp0: float = 0.005
+ wpe: float = 0.3 # 0.001, final cosine lr = wpe * peak lr
+ sche: str = '' # cos, exp, lin
+ log_freq: int = 50 # log frequency in the stdout
+ tclip: float = 2. # <=0 for not grad clip GPT; >100 for per-param clip (%= 100 automatically)
+ # data
+ workers: int = 0 # num workers; 0: auto, -1: don't use multiprocessing in DataLoader
+ norm_eps: float = 1e-6 # norm eps
+ tlen: int = 512 # truncate text embedding to this length
+ Ct5: int = 2048 # feature dimension of text encoder
+ num_of_label_value: int = 2 # num_of_label_value, =2 means bitwise label, =0 means index-wise label, others means fsq, never set to 1
+ enable_dynamic_length_prompt: int = 0 # enable dynamic length prompt during training
+ save_model_iters_freq: int = 1000 # save model iter freq
+ reweight_loss_by_scale: float = 0 # reweight loss by scale
+ vae_latent_dim: int = 1 # here 16/32/64 is bsq vae of different quant bits
+ model_init_device: str = 'cuda' # model_init_device
+ fsdp_init_device: str = 'cuda' # model_init_device
+ apply_spatial_patchify: int = 0 # apply apply_spatial_patchify or not
+ dynamic_scale_schedule: str = '' # dynamic scale schedule
+ use_slice: int = 0 # whether use slice for vae encoding
+ use_vae_token_cache: int = 1 # whether use token cache for speedup
+ save_vae_token_cache: int = 0 # whether save_vae_token_cache
+ allow_online_vae_feature_extraction: int = 1 # whether allow_online_vae_feature_extraction, if False, only load cached features
+ use_text_token_cache: int = 1 # whether use text token cache for speedup
+ token_cache_dir: str = '' # token_cache_dir
+ down_size_limit: int = 10000 # down_size_limit in MB, larger video won't download
+ addition_pn_list: str = '[]'
+ video_caption_type: str = 'merged_caption'
+ video_caption_type: str = '' # video caption type, we use tarsier2_caption
+ only_images4extract_feats: int = 0 # only extract feats for images, set true when extract features, set false for training
+ temporal_compress_rate: int = 4 # temporal_compress_rate, set to 4 by default
+ cached_video_frames: int = 81 # load cache files' video_frames, set to 81 by default
+ rope_type: str = '3d' # rope type, choose from ['2d', '3d', '4d'], default set to '3d'
+ loop_data_per_epoch: int = 0
+ meta_folders: str = ''
+ meta_folder_repeats: str = ''
+ meta_folder_identifiers: str = ''
+ pn_list: str = ''
+ pn_probs: str = ''
+ i2v_ratio: float = 0.
+ dense_ratio4seqpack: float = 1.
+
+ # RL Arguments
+ pair_input: int = 0 # dpo needs pair_input=1
+ rl_with_ref_model: int = 1
+ dpo_training: int = 0 # whether enable dpo training
+ dpo_loss_type: str = 'dpo'
+ dpo_beta: float = 1.0 # 0.1~0.5, larger dpo_beta will be more close to reference model
+ dpo_func: str = 'sigmoid'
+ dpo_label_smoothing: float = 0.0
+ dpo_sft_weight: float = 1.0 # use sft_weight when doing dpo
+ scale_wise_dpo: int = 0
+
+ # ema model or rl reference model
+ use_fsdp_model_ema: int = 0
+ model_ema_decay: float = 0.9999 # model_ema_decay < 1 will update the ema model, >=1 will fix the model and is used as rl reference model
+
+ # seq parallel
+ sp_size: int = 0
+
+ train_max_token_len: int = -1
+ duration_resolution: float = 1
+ cache_check_mode: int = 0 # 0 means not check chche file, 1 means check at the begining, 2 means check at each iteration, -1 means include no cache meta only, used for token cache
+ wp_it: int = 100
+ drop_long_video: int = 1
+
+ image_scale_repetition: str = '[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]'
+ video_scale_repetition: str = '[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]'
+ video_scale_probs: str = '[1, 1, 1, 1, 1, 1, 1, 1, 1, 0.2, 0.1, 0.05]'
+ hbq_round: int = 4
+ train_h_div_w_list: str = '[]'
+ simple_text_proj: int = 0
+ min_video_frames: int = -1
+ fsdp_save_flatten_model: int = 0
+ semantic_scale_dim: int = 16
+ detail_scale_dim: int = 64
+ restrict_data_size: int = -1
+ skip_count_text_token: int = 0
+ semantic_num_lvl: int = 2
+ detail_num_lvl: int = 2
+ use_fsq_cls_head: Literal[0, 1] = 0
+ use_clipwise_caption: int = 0
+ fsdp_warp_mode: str = 'trans_block' # trans_block or full
+ save_start_model: int = 0
+ use_ada_layer_norm: int = 0
+ add_scale_token: Literal[0, 1] = 0 # must >= 0
+ add_class_token: int = 0 # must >= 0, > 0 means class2image
+ vae_encoder_out_type: str = 'mu_sigma'
+ alpha: float = 0.0
+ refine_mode: str = ''
+ log_norm_mean: float = 0
+ log_norm_sigma: float = -1. # < 0 mean disable log-norm sampling
+ gradient_accumulation: int = 1
+ # would be automatically set in runtime
+ cmd: str = ' '.join(a.replace('--exp_name=', '').replace('--exp_name ', '') for a in sys.argv[7:]) # [automatically set; don't specify this]
+
+ @property
+ def gpt_training(self):
+ return len(self.model) > 0
+
+ def set_initial_seed(self, benchmark: bool):
+ torch.backends.cudnn.enabled = True
+ torch.backends.cudnn.benchmark = benchmark
+ assert self.seed
+ seed = self.seed
+ torch.backends.cudnn.deterministic = True
+ os.environ['PYTHONHASHSEED'] = str(seed)
+ random.seed(seed)
+ np.random.seed(seed)
+ torch.manual_seed(seed)
+ if torch.cuda.is_available():
+ torch.cuda.manual_seed(seed)
+ torch.cuda.manual_seed_all(seed)
+
+ def compile_model(self, m, fast):
+ if fast == 0:
+ return m
+ return torch.compile(m, mode={
+ 1: 'reduce-overhead',
+ 2: 'max-autotune',
+ 3: 'default',
+ }[fast]) if hasattr(torch, 'compile') else m
+
+ def state_dict(self, key_ordered=True) -> Union[OrderedDict, dict]:
+ d = (OrderedDict if key_ordered else dict)()
+ # self.as_dict() would contain methods, but we only need variables
+ for k in self.class_variables.keys():
+ if k not in {'device', 'dbg_ks_fp'}: # these are not serializable
+ d[k] = getattr(self, k)
+ return d
+
+ @staticmethod
+ def set_tf32(tf32: bool):
+ if torch.cuda.is_available():
+ torch.backends.cudnn.allow_tf32 = bool(tf32)
+ torch.backends.cuda.matmul.allow_tf32 = bool(tf32)
+ if hasattr(torch, 'set_float32_matmul_precision'):
+ torch.set_float32_matmul_precision('high' if tf32 else 'highest')
+ print(f'[tf32] [precis] torch.get_float32_matmul_precision(): {torch.get_float32_matmul_precision()}')
+ print(f'[tf32] [ conv ] torch.backends.cudnn.allow_tf32: {torch.backends.cudnn.allow_tf32}')
+ print(f'[tf32] [matmul] torch.backends.cuda.matmul.allow_tf32: {torch.backends.cuda.matmul.allow_tf32}')
+
+ def __str__(self):
+ s = []
+ for k in self.class_variables.keys():
+ if k not in {'device', 'dbg_ks_fp'}: # these are not serializable
+ s.append(f' {k:20s}: {getattr(self, k)}')
+ s = '\n'.join(s)
+ return f'{{\n{s}\n}}\n'
+
+
+def init_dist_and_get_args():
+ for i in range(len(sys.argv)):
+ if sys.argv[i].startswith('--local-rank=') or sys.argv[i].startswith('--local_rank='):
+ del sys.argv[i]
+ break
+ args = Args(explicit_bool=True).parse_args(known_only=True)
+ args.chunk_nodes = int(os.environ.get('CK', '') or '0')
+
+ if len(args.extra_args) > 0 and args.is_master_node == 0:
+ print(f'======================================================================================')
+ print(f'=========================== WARNING: UNEXPECTED EXTRA ARGS ===========================\n{args.extra_args}')
+ print(f'=========================== WARNING: UNEXPECTED EXTRA ARGS ===========================')
+ print(f'======================================================================================\n\n')
+
+ args.set_tf32(args.tf32)
+
+ try: os.makedirs(args.bed, exist_ok=True)
+ except: pass
+ try: os.makedirs(args.local_out_path, exist_ok=True)
+ except: pass
+
+ day3 = 60*24*3
+ dist.init_distributed_mode(local_out_path=args.local_out_path, fork=False, timeout_minutes=day3 if int(os.environ.get('LONG_DBG', '0') or '0') > 0 else 30)
+ args.device = dist.get_device()
+
+ # sync seed
+ args.seed = int(time.time())
+ seed = torch.tensor([args.seed], device=args.device)
+ if torch.distributed.is_initialized():
+ torch.distributed.all_reduce(seed, op=torch.distributed.ReduceOp.MIN)
+ args.seed = seed.item()
+
+ if args.sp_size > 1:
+ print(f"INFO: sp_size={args.sp_size}")
+ sp_manager.init_sp(args.sp_size)
+
+ args.sche = args.sche or ('lin0' if args.gpt_training else 'cos')
+ if args.wp == 0:
+ args.wp = args.ep * 1/100
+
+ di = {
+ 'b': 'bilinear', 'c': 'bicubic', 'n': 'nearest', 'a': 'area', 'aa': 'area+area',
+ 'at': 'auto', 'auto': 'auto',
+ 'v': 'vae',
+ 'x': 'pix', 'xg': 'pix_glu', 'gx': 'pix_glu', 'g': 'pix_glu'
+ }
+
+ args.log_txt_path = os.path.join(args.local_out_path, 'log.txt')
+ args.video_scale_probs = json.loads(args.video_scale_probs)
+
+ ls = '[]'
+ if 'AUTO_RESUME' in os.environ:
+ ls.append(int(os.environ['AUTO_RESUME']))
+ ls = sorted(ls, reverse=True)
+ ls = [str(i) for i in ls]
+ args.ckpt_trials = ls
+ args.real_trial_id = args.trial_id if len(ls) == 0 else str(ls[-1])
+
+ args.enable_checkpointing = None if args.enable_checkpointing in [False, 0, "0"] else args.enable_checkpointing
+ args.enable_checkpointing = "full-block" if args.enable_checkpointing in [True, 1, "1"] else args.enable_checkpointing
+ assert args.enable_checkpointing in [None, "full-block", "full-attn", "self-attn"], \
+ f"only support no-checkpointing or full-block/full-attn checkpointing, but got {args.enable_checkpointing}."
+
+ if len(args.exp_name) == 0:
+ args.exp_name = os.path.basename(args.bed) or 'test_exp'
+
+ if dist.is_master():
+ from grn.utils.safe_rm import safe_remove
+ safe_remove(os.path.join(args.bed, "ready-node*"), args.bed)
+ safe_remove(os.path.join(args.local_out_path, "ready-node*"), args.local_out_path)
+
+ if args.sdpa_mem:
+ from torch.backends.cuda import enable_flash_sdp, enable_math_sdp, enable_mem_efficient_sdp
+ enable_flash_sdp(True)
+ enable_mem_efficient_sdp(True)
+ enable_math_sdp(False)
+ print(args)
+ return args
diff --git a/grn/utils_t2iv/comm/comm.py b/grn/utils_t2iv/comm/comm.py
new file mode 100644
index 0000000000000000000000000000000000000000..6eaa6678eeec231b2f31ac1daa176ac9be7c812b
--- /dev/null
+++ b/grn/utils_t2iv/comm/comm.py
@@ -0,0 +1,422 @@
+from typing import Any, Optional, Tuple
+
+import torch
+import torch.distributed as dist
+import torch.nn.functional as F
+from einops import rearrange
+from torch import Tensor
+from torch.distributed import ProcessGroup
+
+if torch.__version__ >= "2.4.0":
+ _torch_custom_op_wrapper = torch.library.custom_op
+ _torch_register_fake_wrapper = torch.library.register_fake
+else:
+ def noop_custom_op_wrapper(name, fn=None, /, *, mutates_args, device_types=None, schema=None):
+ def wrap(func):
+ return func
+ if fn is None:
+ return wrap
+ return fn
+ def noop_register_fake_wrapper(op, fn=None, /, *, lib=None, _stacklevel=1):
+ def wrap(func):
+ return func
+ if fn is None:
+ return wrap
+ return fn
+ _torch_custom_op_wrapper = noop_custom_op_wrapper
+ _torch_register_fake_wrapper = noop_register_fake_wrapper
+
+
+__sp_comm_group__ = None
+
+def set_sp_comm_group(group=None):
+ global __sp_comm_group__
+ assert __sp_comm_group__ is None and group is not None
+ __sp_comm_group__ = group
+
+def get_sp_comm_group():
+ global __sp_comm_group__
+ assert __sp_comm_group__ is not None
+ return __sp_comm_group__
+
+
+# ======================================================
+# Model
+# ======================================================
+
+
+def model_sharding(model: torch.nn.Module):
+ global_rank = dist.get_rank()
+ world_size = dist.get_world_size()
+ for _, param in model.named_parameters():
+ padding_size = (world_size - param.numel() % world_size) % world_size
+ if padding_size > 0:
+ padding_param = torch.nn.functional.pad(param.data.view(-1), [0, padding_size])
+ else:
+ padding_param = param.data.view(-1)
+ splited_params = padding_param.split(padding_param.numel() // world_size)
+ splited_params = splited_params[global_rank]
+ param.data = splited_params
+
+
+# ======================================================
+# AllGather & ReduceScatter
+# ======================================================
+
+
+class AsyncAllGatherForTwo(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ weight: Tensor,
+ bias: Tensor,
+ sp_rank: int,
+ sp_size: int,
+ group: Optional[ProcessGroup] = None,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ from torch.distributed._functional_collectives import all_gather_tensor
+
+ ctx.group = group
+ ctx.sp_rank = sp_rank
+ ctx.sp_size = sp_size
+
+ # all gather inputs
+ all_inputs = all_gather_tensor(inputs.unsqueeze(0), 0, group)
+ # compute local qkv
+ local_qkv = F.linear(inputs, weight, bias).unsqueeze(0)
+
+ # remote compute
+ remote_inputs = all_inputs[1 - sp_rank].view(list(local_qkv.shape[:-1]) + [-1])
+ # compute remote qkv
+ remote_qkv = F.linear(remote_inputs, weight, bias)
+
+ # concat local and remote qkv
+ if sp_rank == 0:
+ qkv = torch.cat([local_qkv, remote_qkv], dim=0)
+ else:
+ qkv = torch.cat([remote_qkv, local_qkv], dim=0)
+ qkv = rearrange(qkv, "sp b n c -> b (sp n) c")
+
+ ctx.save_for_backward(inputs, weight, remote_inputs)
+ return qkv
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ from torch.distributed._functional_collectives import reduce_scatter_tensor
+
+ group = ctx.group
+ sp_rank = ctx.sp_rank
+ sp_size = ctx.sp_size
+ inputs, weight, remote_inputs = ctx.saved_tensors
+
+ # split qkv_grad
+ qkv_grad = grad_outputs[0]
+ qkv_grad = rearrange(qkv_grad, "b (sp n) c -> sp b n c", sp=sp_size)
+ qkv_grad = torch.chunk(qkv_grad, 2, dim=0)
+ if sp_rank == 0:
+ local_qkv_grad, remote_qkv_grad = qkv_grad
+ else:
+ remote_qkv_grad, local_qkv_grad = qkv_grad
+
+ # compute remote grad
+ remote_inputs_grad = torch.matmul(remote_qkv_grad, weight).squeeze(0)
+ weight_grad = torch.matmul(remote_qkv_grad.transpose(-1, -2), remote_inputs).squeeze(0).sum(0)
+ bias_grad = remote_qkv_grad.squeeze(0).sum(0).sum(0)
+
+ # launch async reduce scatter
+ remote_inputs_grad_zero = torch.zeros_like(remote_inputs_grad)
+ if sp_rank == 0:
+ remote_inputs_grad = torch.cat([remote_inputs_grad_zero, remote_inputs_grad], dim=0)
+ else:
+ remote_inputs_grad = torch.cat([remote_inputs_grad, remote_inputs_grad_zero], dim=0)
+ remote_inputs_grad = reduce_scatter_tensor(remote_inputs_grad, "sum", 0, group)
+
+ # compute local grad and wait for reduce scatter
+ local_input_grad = torch.matmul(local_qkv_grad, weight).squeeze(0)
+ weight_grad += torch.matmul(local_qkv_grad.transpose(-1, -2), inputs).squeeze(0).sum(0)
+ bias_grad += local_qkv_grad.squeeze(0).sum(0).sum(0)
+
+ # sum remote and local grad
+ inputs_grad = remote_inputs_grad + local_input_grad
+ return inputs_grad, weight_grad, bias_grad, None, None, None
+
+
+class AllGather(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ group: Optional[ProcessGroup] = None,
+ overlap: bool = False,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ assert ctx is not None or not overlap
+
+ if ctx is not None:
+ ctx.comm_grp = group
+
+ comm_size = dist.get_world_size(group)
+ if comm_size == 1:
+ return inputs.unsqueeze(0), None
+
+ buffer_shape = (comm_size,) + inputs.shape
+ outputs = torch.empty(buffer_shape, dtype=inputs.dtype, device=inputs.device)
+ buffer_list = list(torch.chunk(outputs, comm_size, dim=0))
+ if not overlap:
+ dist.all_gather(buffer_list, inputs, group=group)
+ return outputs, None
+ else:
+ handle = dist.all_gather(buffer_list, inputs, group=group, async_op=True)
+ return outputs, handle
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ return (
+ ReduceScatter.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
+ None,
+ None,
+ )
+
+
+class ReduceScatter(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ group: ProcessGroup,
+ overlap: bool = False,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ assert ctx is not None or not overlap
+
+ if ctx is not None:
+ ctx.comm_grp = group
+
+ comm_size = dist.get_world_size(group)
+ if comm_size == 1:
+ return inputs.squeeze(0), None
+
+ if not inputs.is_contiguous():
+ inputs = inputs.contiguous()
+
+ output_shape = inputs.shape[1:]
+ outputs = torch.empty(output_shape, dtype=inputs.dtype, device=inputs.device)
+ buffer_list = list(torch.chunk(inputs, comm_size, dim=0))
+ if not overlap:
+ dist.reduce_scatter(outputs, buffer_list, group=group)
+ return outputs, None
+ else:
+ handle = dist.reduce_scatter(outputs, buffer_list, group=group, async_op=True)
+ return outputs, handle
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ # TODO: support async backward
+ return (
+ AllGather.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
+ None,
+ None,
+ )
+
+
+# ======================================================
+# AlltoAll
+# ======================================================
+
+
+@_torch_custom_op_wrapper("distributed::_all_to_all_func", mutates_args=(), device_types="cuda")
+def _all_to_all_func(input_: torch.Tensor, world_size: int = 1, scatter_dim: int = 0, gather_dim: int = 0) -> torch.Tensor:
+ input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
+ output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
+ group = get_sp_comm_group()
+ dist.all_to_all(output_list, input_list, group=group)
+ return torch.cat(output_list, dim=gather_dim).contiguous()
+
+
+@_torch_register_fake_wrapper("distributed::_all_to_all_func")
+def _all_to_all_func_fake(input_: torch.Tensor, world_size: int = 1, scatter_dim: int = 0, gather_dim: int = 0) -> torch.Tensor:
+ inp_shape = list(input_.shape)
+ group = get_sp_comm_group()
+ world_size = dist.get_world_size(group)
+ if world_size == 1:
+ return input_
+
+ inp_shape[gather_dim] = inp_shape[gather_dim] * world_size
+ inp_shape[scatter_dim] = inp_shape[scatter_dim] // world_size
+ outputs = torch.empty(torch.Size(inp_shape), dtype=input_.dtype, device=input_.device, layout=input_.layout)
+ return outputs
+
+
+class _AllToAll(torch.autograd.Function):
+ """All-to-all communication.
+
+ Args:
+ input_: input matrix
+ process_group: communication group
+ scatter_dim: scatter dimension
+ gather_dim: gather dimension
+ """
+
+ @staticmethod
+ def forward(ctx, input_, process_group, scatter_dim, gather_dim):
+ ctx.process_group = process_group
+ ctx.scatter_dim = scatter_dim
+ ctx.gather_dim = gather_dim
+ world_size = dist.get_world_size(process_group)
+
+ return _wrapper_all_to_all_func(input_, world_size, scatter_dim, gather_dim)
+
+ @staticmethod
+ def backward(ctx, *grad_output):
+ process_group = ctx.process_group
+ scatter_dim = ctx.gather_dim
+ gather_dim = ctx.scatter_dim
+ return_grad = _AllToAll.apply(*grad_output, process_group, scatter_dim, gather_dim)
+ return (return_grad, None, None, None)
+
+
+def all_to_all_comm(input_, process_group=None, scatter_dim=2, gather_dim=1):
+ return _AllToAll.apply(input_, process_group, scatter_dim, gather_dim)
+
+
+# ======================================================
+# Sequence Gather & Split
+# ======================================================
+
+
+def _split_sequence_func(inputs, pg: dist.ProcessGroup, dim=-1):
+ world_size = dist.get_world_size(pg)
+ if world_size == 1:
+ return inputs
+
+ # Split along last dimension.
+ rank = dist.get_rank(pg)
+ dim_size = inputs.size(dim)
+ assert dim_size % world_size == 0, (
+ f"The dimension to split ({dim_size}) is not a multiple of world size ({world_size}), "
+ f"cannot split tensor evenly"
+ )
+
+ outputs = torch.split(inputs, dim_size // world_size, dim=dim)[rank]
+ return outputs
+
+
+@_torch_custom_op_wrapper("distributed::_gather_sequence_func", mutates_args=(), device_types="cuda")
+def _gather_sequence_func(inputs: torch.Tensor, dim: int = -1) -> torch.Tensor:
+ pg = get_sp_comm_group()
+ world_size = dist.get_world_size(pg)
+ if world_size == 1:
+ return inputs
+
+ # all gather
+ inputs = inputs.contiguous()
+ outputs = [torch.empty_like(inputs) for _ in range(world_size)]
+ dist.all_gather(outputs, inputs, group=pg)
+
+ # concat
+ outputs = torch.cat(outputs, dim=dim)
+ return outputs
+
+
+@_torch_register_fake_wrapper("distributed::_gather_sequence_func")
+def _gather_sequence_func_fake(inputs: torch.Tensor, dim: int = -1) -> torch.Tensor:
+ inp_shape = list(inputs.shape)
+ pg = get_sp_comm_group()
+ world_size = dist.get_world_size(pg)
+ if world_size == 1:
+ return inputs
+
+ inp_shape[dim] = inp_shape[dim] * world_size
+ outputs = torch.empty(torch.Size(inp_shape), dtype=inputs.dtype, device=inputs.device, layout=inputs.layout)
+ return outputs
+
+
+if torch.__version__ >= "2.4.0":
+ _wrapper_all_to_all_func = torch.ops.distributed._all_to_all_func
+ _wrapper_gather_sequence_func = torch.ops.distributed._gather_sequence_func
+else:
+ _wrapper_all_to_all_func = _all_to_all_func
+ _wrapper_gather_sequence_func = _gather_sequence_func
+
+
+class _GatherForwardSplitBackward(torch.autograd.Function):
+ """
+ Gather the input sequence.
+
+ Args:
+ input_: input matrix.
+ process_group: process group.
+ dim: dimension
+ """
+
+ @staticmethod
+ def symbolic(graph, input_):
+ return _wrapper_gather_sequence_func(input_)
+
+ @staticmethod
+ def forward(ctx, input_, process_group, dim, grad_scale):
+ ctx.process_group = process_group
+ ctx.dim = dim
+ ctx.grad_scale = grad_scale
+ return _wrapper_gather_sequence_func(input_, dim)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == "up":
+ grad_output = grad_output * dist.get_world_size(ctx.process_group)
+ elif ctx.grad_scale == "down":
+ grad_output = grad_output / dist.get_world_size(ctx.process_group)
+
+ return _split_sequence_func(grad_output, ctx.process_group, ctx.dim), None, None, None
+
+
+class _SplitForwardGatherBackward(torch.autograd.Function):
+ """
+ Split sequence.
+
+ Args:
+ input_: input matrix.
+ process_group: parallel mode.
+ dim: dimension
+ """
+
+ @staticmethod
+ def symbolic(graph, input_):
+ return _split_sequence_func(input_)
+
+ @staticmethod
+ def forward(ctx, input_, process_group, dim, grad_scale):
+ ctx.process_group = process_group
+ ctx.dim = dim
+ ctx.grad_scale = grad_scale
+ return _split_sequence_func(input_, process_group, dim)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == "up":
+ grad_output = grad_output * dist.get_world_size(ctx.process_group)
+ elif ctx.grad_scale == "down":
+ grad_output = grad_output / dist.get_world_size(ctx.process_group)
+ return _wrapper_gather_sequence_func(grad_output, ctx.dim), None, None, None
+
+
+def split_sequence(input_, process_group, dim, grad_scale=1.0):
+ return _SplitForwardGatherBackward.apply(input_, process_group, dim, grad_scale)
+
+
+def gather_sequence(input_, process_group, dim, grad_scale=None):
+ return _GatherForwardSplitBackward.apply(input_, process_group, dim, grad_scale)
diff --git a/grn/utils_t2iv/comm/dist.py b/grn/utils_t2iv/comm/dist.py
new file mode 100644
index 0000000000000000000000000000000000000000..7b60ade7851cc28a50047a1c7560d500352de496
--- /dev/null
+++ b/grn/utils_t2iv/comm/dist.py
@@ -0,0 +1,187 @@
+import torch
+import torch.distributed as dist
+
+
+# ====================
+# All-To-All
+# ====================
+def _all_to_all(
+ input_: torch.Tensor,
+ world_size: int,
+ group: dist.ProcessGroup,
+ scatter_dim: int,
+ gather_dim: int,
+):
+ input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
+ output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
+ dist.all_to_all(output_list, input_list, group=group)
+ return torch.cat(output_list, dim=gather_dim).contiguous()
+
+
+class _AllToAll(torch.autograd.Function):
+ """All-to-all communication.
+
+ Args:
+ input_: input matrix
+ process_group: communication group
+ scatter_dim: scatter dimension
+ gather_dim: gather dimension
+ """
+
+ @staticmethod
+ def forward(ctx, input_, process_group, scatter_dim, gather_dim):
+ ctx.process_group = process_group
+ ctx.scatter_dim = scatter_dim
+ ctx.gather_dim = gather_dim
+ ctx.world_size = dist.get_world_size(process_group)
+ output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
+ return output
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ grad_output = _all_to_all(
+ grad_output,
+ ctx.world_size,
+ ctx.process_group,
+ ctx.gather_dim,
+ ctx.scatter_dim,
+ )
+ return (
+ grad_output,
+ None,
+ None,
+ None,
+ )
+
+
+def all_to_all(
+ input_: torch.Tensor,
+ process_group: dist.ProcessGroup,
+ scatter_dim: int = 2,
+ gather_dim: int = 1,
+):
+ return _AllToAll.apply(input_, process_group, scatter_dim, gather_dim)
+
+
+def _gather(
+ input_: torch.Tensor,
+ world_size: int,
+ group: dist.ProcessGroup,
+ gather_dim: int,
+):
+ if gather_list is None:
+ gather_list = [torch.empty_like(input_) for _ in range(world_size)]
+ dist.gather(input_, gather_list, group=group, gather_dim=gather_dim)
+ return gather_list
+
+
+# ====================
+# Gather-Split
+# ====================
+
+
+def _split(input_, pg: dist.ProcessGroup, dim=-1):
+ # skip if only one rank involved
+ world_size = dist.get_world_size(pg)
+ rank = dist.get_rank(pg)
+ if world_size == 1:
+ return input_
+
+ # split last dim
+ dim_size = input_.size(dim)
+ assert dim_size % world_size == 0, (
+ f"Dim to split ({dim_size}) is not divisible by world_size ({world_size}), "
+ f"uneven split"
+ )
+
+ tensor_list = torch.split(input_, dim_size // world_size, dim=dim)
+ output = tensor_list[rank].contiguous()
+
+ return output
+
+
+def _gather(input_, pg: dist.ProcessGroup, dim=-1):
+ input_ = input_.contiguous()
+ world_size = dist.get_world_size(pg)
+ dist.get_rank(pg)
+
+ if world_size == 1:
+ return input_
+
+ # all gather
+ tensor_list = [torch.empty_like(input_) for _ in range(world_size)]
+ assert input_.device.type == "cuda"
+ torch.distributed.all_gather(tensor_list, input_, group=pg)
+
+ # concat
+ output = torch.cat(tensor_list, dim=dim).contiguous()
+
+ return output
+
+
+class _GatherForwardSplitBackward(torch.autograd.Function):
+ """Gather the input from model parallel region and concatenate.
+
+ Args:
+ input_: input matrix.
+ process_group: parallel mode.
+ dim: dimension
+ """
+
+ @staticmethod
+ def symbolic(graph, input_):
+ return _gather(input_)
+
+ @staticmethod
+ def forward(ctx, input_, process_group, dim, grad_scale):
+ ctx.mode = process_group
+ ctx.dim = dim
+ ctx.grad_scale = grad_scale
+ return _gather(input_, process_group, dim)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == "up":
+ grad_output = grad_output * dist.get_world_size(ctx.mode)
+ elif ctx.grad_scale == "down":
+ grad_output = grad_output / dist.get_world_size(ctx.mode)
+
+ return _split(grad_output, ctx.mode, ctx.dim), None, None, None
+
+
+class _SplitForwardGatherBackward(torch.autograd.Function):
+ """
+ Split the input and keep only the corresponding chuck to the rank.
+
+ Args:
+ input_: input matrix.
+ process_group: parallel mode.
+ dim: dimension
+ """
+
+ @staticmethod
+ def symbolic(graph, input_):
+ return _split(input_)
+
+ @staticmethod
+ def forward(ctx, input_, process_group, dim, grad_scale):
+ ctx.mode = process_group
+ ctx.dim = dim
+ ctx.grad_scale = grad_scale
+ return _split(input_, process_group, dim)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == "up":
+ grad_output = grad_output * dist.get_world_size(ctx.mode)
+ elif ctx.grad_scale == "down":
+ grad_output = grad_output / dist.get_world_size(ctx.mode)
+ return _gather(grad_output, ctx.mode, ctx.dim), None, None, None
+
+
+def split_forward_gather_backward(input_, process_group, dim, grad_scale=1.0):
+ return _SplitForwardGatherBackward.apply(input_, process_group, dim, grad_scale)
+
+
+def gather_forward_split_backward(input_, process_group, dim, grad_scale=None):
+ return _GatherForwardSplitBackward.apply(input_, process_group, dim, grad_scale)
\ No newline at end of file
diff --git a/grn/utils_t2iv/comm/operation.py b/grn/utils_t2iv/comm/operation.py
new file mode 100644
index 0000000000000000000000000000000000000000..1204278d17513d0d4145dcb0e180c47ca4608bcb
--- /dev/null
+++ b/grn/utils_t2iv/comm/operation.py
@@ -0,0 +1,378 @@
+from typing import Any, Optional, Tuple
+
+import torch
+import torch.distributed as dist
+import torch.nn.functional as F
+from einops import rearrange
+from torch import Tensor
+from torch.distributed import ProcessGroup
+
+
+class AllToAll(torch.autograd.Function):
+ """Dispatches input tensor [e, c, h] to all experts by all_to_all_single
+ operation in torch.distributed.
+ """
+
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ group: ProcessGroup,
+ overlap: bool = False,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ assert ctx is not None or not overlap
+
+ if ctx is not None:
+ ctx.comm_grp = group
+ if not inputs.is_contiguous():
+ inputs = inputs.contiguous()
+ if dist.get_world_size(group) == 1:
+ return inputs, None
+ output = torch.empty_like(inputs)
+ if not overlap:
+ dist.all_to_all_single(output, inputs, group=group)
+ return output, None
+ else:
+ handle = dist.all_to_all_single(output, inputs, group=group, async_op=True)
+ return output, handle
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ return (
+ AllToAll.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
+ None,
+ None,
+ )
+
+
+class AsyncAllGatherForTwo(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ weight: Tensor,
+ bias: Tensor,
+ sp_rank: int,
+ sp_size: int,
+ group: Optional[ProcessGroup] = None,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ from torch.distributed._functional_collectives import all_gather_tensor
+
+ ctx.group = group
+ ctx.sp_rank = sp_rank
+ ctx.sp_size = sp_size
+
+ # all gather inputs
+ all_inputs = all_gather_tensor(inputs.unsqueeze(0), 0, group)
+ # compute local qkv
+ local_qkv = F.linear(inputs, weight, bias).unsqueeze(0)
+
+ # remote compute
+ remote_inputs = all_inputs[1 - sp_rank].view(list(local_qkv.shape[:-1]) + [-1])
+ # compute remote qkv
+ remote_qkv = F.linear(remote_inputs, weight, bias)
+
+ # concat local and remote qkv
+ if sp_rank == 0:
+ qkv = torch.cat([local_qkv, remote_qkv], dim=0)
+ else:
+ qkv = torch.cat([remote_qkv, local_qkv], dim=0)
+ qkv = rearrange(qkv, "sp b n c -> b (sp n) c")
+
+ ctx.save_for_backward(inputs, weight, remote_inputs)
+ return qkv
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ from torch.distributed._functional_collectives import reduce_scatter_tensor
+
+ group = ctx.group
+ sp_rank = ctx.sp_rank
+ sp_size = ctx.sp_size
+ inputs, weight, remote_inputs = ctx.saved_tensors
+
+ # split qkv_grad
+ qkv_grad = grad_outputs[0]
+ qkv_grad = rearrange(qkv_grad, "b (sp n) c -> sp b n c", sp=sp_size)
+ qkv_grad = torch.chunk(qkv_grad, 2, dim=0)
+ if sp_rank == 0:
+ local_qkv_grad, remote_qkv_grad = qkv_grad
+ else:
+ remote_qkv_grad, local_qkv_grad = qkv_grad
+
+ # compute remote grad
+ remote_inputs_grad = torch.matmul(remote_qkv_grad, weight).squeeze(0)
+ weight_grad = torch.matmul(remote_qkv_grad.transpose(-1, -2), remote_inputs).squeeze(0).sum(0)
+ bias_grad = remote_qkv_grad.squeeze(0).sum(0).sum(0)
+
+ # launch async reduce scatter
+ remote_inputs_grad_zero = torch.zeros_like(remote_inputs_grad)
+ if sp_rank == 0:
+ remote_inputs_grad = torch.cat([remote_inputs_grad_zero, remote_inputs_grad], dim=0)
+ else:
+ remote_inputs_grad = torch.cat([remote_inputs_grad, remote_inputs_grad_zero], dim=0)
+ remote_inputs_grad = reduce_scatter_tensor(remote_inputs_grad, "sum", 0, group)
+
+ # compute local grad and wait for reduce scatter
+ local_input_grad = torch.matmul(local_qkv_grad, weight).squeeze(0)
+ weight_grad += torch.matmul(local_qkv_grad.transpose(-1, -2), inputs).squeeze(0).sum(0)
+ bias_grad += local_qkv_grad.squeeze(0).sum(0).sum(0)
+
+ # sum remote and local grad
+ inputs_grad = remote_inputs_grad + local_input_grad
+ return inputs_grad, weight_grad, bias_grad, None, None, None
+
+
+class AllGather(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ group: Optional[ProcessGroup] = None,
+ overlap: bool = False,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ assert ctx is not None or not overlap
+
+ if ctx is not None:
+ ctx.comm_grp = group
+
+ comm_size = dist.get_world_size(group)
+ # print(f"XW debug, All Gather Dist world size {comm_size}")
+ if comm_size == 1:
+ return inputs.unsqueeze(0), None
+
+ buffer_shape = (comm_size,) + inputs.shape
+ outputs = torch.empty(buffer_shape, dtype=inputs.dtype, device=inputs.device)
+ buffer_list = list(torch.chunk(outputs, comm_size, dim=0))
+ # buffer_list = list([
+ # t.squeeze(0) for t in torch.chunk(outputs, comm_size, dim=0)
+ # ])
+
+ if not overlap:
+ # print("buffer list", len(buffer_list), [t.shape for t in buffer_list])
+ # print("inputs", inputs.shape, inputs.is_contiguous())
+ # print(group)
+
+ dist.all_gather(buffer_list, inputs, group=group)
+ return outputs, None
+ else:
+ handle = dist.all_gather(buffer_list, inputs, group=group, async_op=True)
+ return outputs, handle
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ return (
+ ReduceScatter.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
+ None,
+ None,
+ )
+
+
+class ReduceScatter(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: Any,
+ inputs: Tensor,
+ group: ProcessGroup,
+ overlap: bool = False,
+ ) -> Tuple[Tensor, Any]:
+ """
+ Returns:
+ outputs: Tensor
+ handle: Optional[Work], if overlap is True
+ """
+ assert ctx is not None or not overlap
+
+ if ctx is not None:
+ ctx.comm_grp = group
+
+ comm_size = dist.get_world_size(group)
+ if comm_size == 1:
+ return inputs.squeeze(0), None
+
+ if not inputs.is_contiguous():
+ inputs = inputs.contiguous()
+
+ output_shape = inputs.shape[1:]
+ outputs = torch.empty(output_shape, dtype=inputs.dtype, device=inputs.device)
+ buffer_list = list(torch.chunk(inputs, comm_size, dim=0))
+ if not overlap:
+ dist.reduce_scatter(outputs, buffer_list, group=group)
+ return outputs, None
+ else:
+ handle = dist.reduce_scatter(outputs, buffer_list, group=group, async_op=True)
+ return outputs, handle
+
+ @staticmethod
+ def backward(ctx: Any, *grad_outputs) -> Tuple[Tensor, None, None]:
+ # TODO: support async backward
+ return (
+ AllGather.forward(None, grad_outputs[0], ctx.comm_grp, False)[0],
+ None,
+ None,
+ )
+
+
+# using all_to_all_single api to perform all to all communication
+def _all_to_all_single(input_, seq_world_size, group, scatter_dim, gather_dim):
+ inp_shape = list(input_.shape)
+ inp_shape[scatter_dim] = inp_shape[scatter_dim] // seq_world_size
+ if scatter_dim < 2:
+ input_t = input_.reshape([seq_world_size, inp_shape[scatter_dim]] + inp_shape[scatter_dim + 1 :]).contiguous()
+ else:
+ input_t = (
+ input_.reshape([-1, seq_world_size, inp_shape[scatter_dim]] + inp_shape[scatter_dim + 1 :])
+ .transpose(0, 1)
+ .contiguous()
+ )
+
+ output = torch.empty_like(input_t)
+ dist.all_to_all_single(output, input_t, group=group)
+
+ if scatter_dim < 2:
+ output = output.transpose(0, 1).contiguous()
+
+ return output.reshape(
+ inp_shape[:gather_dim]
+ + [
+ inp_shape[gather_dim] * seq_world_size,
+ ]
+ + inp_shape[gather_dim + 1 :]
+ ).contiguous()
+
+
+# using all_to_all api to perform all to all communication
+def _all_to_all(input_, world_size, group, scatter_dim, gather_dim):
+ input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
+ output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
+ dist.all_to_all(output_list, input_list, group=group)
+ return torch.cat(output_list, dim=gather_dim).contiguous()
+
+
+class _AllToAll(torch.autograd.Function):
+ """All-to-all communication.
+
+ Args:
+ input_: input matrix
+ process_group: communication group
+ scatter_dim: scatter dimension
+ gather_dim: gather dimension
+ """
+
+ @staticmethod
+ def forward(ctx, input_, process_group, scatter_dim, gather_dim):
+ ctx.process_group = process_group
+ ctx.scatter_dim = scatter_dim
+ ctx.gather_dim = gather_dim
+ world_size = dist.get_world_size(process_group)
+ bsz, _, _ = input_.shape
+
+ # Todo: Try to make all_to_all_single compatible with a large batch size
+ if bsz == 1:
+ return _all_to_all_single(input_, world_size, process_group, scatter_dim, gather_dim)
+ else:
+ return _all_to_all(input_, world_size, process_group, scatter_dim, gather_dim)
+
+ @staticmethod
+ def backward(ctx, *grad_output):
+ process_group = ctx.process_group
+ scatter_dim = ctx.gather_dim
+ gather_dim = ctx.scatter_dim
+ return_grad = _AllToAll.apply(*grad_output, process_group, scatter_dim, gather_dim)
+ return (return_grad, None, None, None)
+
+
+def model_sharding(model: torch.nn.Module):
+ global_rank = dist.get_rank()
+ world_size = dist.get_world_size()
+ for _, param in model.named_parameters():
+ padding_size = (world_size - param.numel() % world_size) % world_size
+ if padding_size > 0:
+ padding_param = torch.nn.functional.pad(param.data.view(-1), [0, padding_size])
+ else:
+ padding_param = param.data.view(-1)
+ splited_params = padding_param.split(padding_param.numel() // world_size)
+ splited_params = splited_params[global_rank]
+ param.data = splited_params
+
+
+def all_to_all_comm(input_, process_group=None, scatter_dim=2, gather_dim=1):
+ return _AllToAll.apply(input_, process_group, scatter_dim, gather_dim)
+
+
+def _gather(input_, dim=-1, process_group=None):
+ # skip if only one rank involved
+ world_size = dist.get_world_size(process_group)
+ if world_size == 1:
+ return input_
+
+ # all gather
+ input_ = input_.contiguous()
+ tensor_list = [torch.empty_like(input_) for _ in range(world_size)]
+ torch.distributed.all_gather(tensor_list, input_, group=process_group)
+
+ # concat
+ output = torch.cat(tensor_list, dim=dim).contiguous()
+
+ return output
+
+
+def _split(input_, dim=-1, process_group=None):
+ # skip if only one rank involved
+ world_size = dist.get_world_size(process_group)
+ if world_size == 1:
+ return input_
+
+ # Split along last dimension.
+ dim_size = input_.size(dim)
+ assert dim_size % world_size == 0, (
+ f"The dimension to split ({dim_size}) is not a multiple of world size ({world_size}), "
+ f"cannot split tensor evenly"
+ )
+
+ tensor_list = torch.split(input_, dim_size // world_size, dim=dim)
+ rank = dist.get_rank(process_group)
+ output = tensor_list[rank].clone().contiguous()
+
+ return output
+
+
+class _GatherForwardSplitBackward(torch.autograd.Function):
+ """Gather the input from model parallel region and concatenate.
+
+ Args:
+ input_: input matrix.
+ parallel_mode: parallel mode.
+ dim: dimension
+ """
+
+ @staticmethod
+ def forward(ctx, input_, dim, process_group):
+ ctx.process_group = process_group
+ ctx.dim = dim
+ return _gather(input_, dim, process_group)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ return _split(grad_output, ctx.dim, ctx.process_group), None, None
+
+
+def gather_forward_split_backward(input_, dim, process_group):
+ return _GatherForwardSplitBackward.apply(input_, dim, process_group)
+
+
diff --git a/grn/utils_t2iv/comm/pg_utils.py b/grn/utils_t2iv/comm/pg_utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..d29591f793e91caf01b1d59100bad04cefd2a41c
--- /dev/null
+++ b/grn/utils_t2iv/comm/pg_utils.py
@@ -0,0 +1,230 @@
+# copy from colossalai and opendit
+#
+import itertools
+from functools import reduce
+from operator import mul
+from typing import Dict, List, Optional, Tuple, Union
+
+import numpy as np
+import torch.distributed as dist
+from torch.distributed import ProcessGroup
+
+
+def prod(nums: List[int]) -> int:
+ """Product of a list of numbers.
+
+ Args:
+ nums (List[int]): A list of numbers.
+
+ Returns:
+ int: The product of the numbers.
+ """
+ return reduce(mul, nums)
+
+
+class ProcessGroupMesh:
+ """A helper class to manage the process group mesh. It only describes how to organize process groups, and it's decoupled with parallel method.
+ It just initialize process groups and cache them. The parallel method should manage them and use them to do the parallel computation.
+
+ We use a ND-tuple to represent the process group mesh. And a ND-coordinate is to represent each process.
+ For example, ``(0, 1, 0)`` represents the process whose rank is 2 in a 3D process group mesh with size ``(2, 2, 2)``.
+
+ Args:
+ *size (int): The size of each dimension of the process group mesh. The product of the size must be equal to the world size.
+
+ Attributes:
+ shape (Tuple[int, ...]): The shape of the process group mesh.
+ rank (int): The rank of the current process.
+ """
+
+ def __init__(self, *size: int) -> None:
+ assert dist.is_initialized(), "Please initialize torch.distributed first."
+ assert prod(size) == dist.get_world_size(), f"The product of the size must be equal to the world size. However, got {prod(size)} and {dist.get_world_size()}."
+ self._shape = size
+ self._rank = dist.get_rank()
+ self._coord = ProcessGroupMesh.unravel(self._rank, self._shape)
+ self._ranks_to_group: Dict[Tuple[int, ...], ProcessGroup] = {}
+ self._group_to_ranks: Dict[ProcessGroup, Tuple[int, ...]] = {}
+
+ @property
+ def shape(self) -> Tuple[int, ...]:
+ return self._shape
+
+ @property
+ def rank(self) -> int:
+ return self._rank
+
+ def size(self, dim: Optional[int] = None) -> Union[int, Tuple[int, ...]]:
+ """Get the size of the process group mesh.
+
+ Args:
+ dim (Optional[int], optional): Dimension of the process group mesh. `None` means all dimensions. Defaults to None.
+
+ Returns:
+ Union[int, Tuple[int, ...]]: Size of the target dimension or the whole process group mesh.
+ """
+ if dim is None:
+ return self._shape
+ else:
+ return self._shape[dim]
+
+ def coordinate(self, dim: Optional[int] = None) -> Union[int, Tuple[int, ...]]:
+ """Get the coordinate of the process group mesh.
+
+ Args:
+ dim (Optional[int], optional): Dimension of the process group mesh. `None` means all dimensions. Defaults to None.
+
+ Returns:
+ Union[int, Tuple[int, ...]]: Coordinate of the target dimension or the whole process group mesh.
+ """
+ if dim is None:
+ return self._coord
+ else:
+ return self._coord[dim]
+
+ @staticmethod
+ def unravel(rank: int, shape: Tuple[int, ...]) -> Tuple[int, ...]:
+ """Convert a rank to a coordinate.
+
+ Args:
+ rank (int): Rank to be converted.
+ shape (Tuple[int, ...]): Shape of the process group mesh.
+
+ Returns:
+ Tuple[int, ...]: Coordinate of the rank.
+ """
+ res = np.unravel_index(rank, shape)
+ return tuple(int(i) for i in res)
+
+ @staticmethod
+ def ravel(coord: Tuple[int, ...], shape: Tuple[int, ...], mode: str = "raise") -> int:
+ """Convert a coordinate to a rank.
+ mode: ['raise', 'wrap', 'clip'], see https://numpy.org/doc/stable/reference/generated/numpy.ravel_multi_index.html.
+ with wrap, index out of range would be wrapped around.
+ For instance, ravel((0, i, 0), (1, 2, 1), 'wrap') returns (i % 2)
+
+ Args:
+ coords (Tuple[int, ...]): Coordinate to be converted.
+ shape (Tuple[int, ...]): Shape of the process group mesh.
+ mode (Optional[str]): The mode for numpy.ravel_multi_index.
+
+ Returns:
+ int: Rank of the coordinate.
+ """
+
+ assert mode in ["raise", "wrap", "clip"]
+ return int(np.ravel_multi_index(coord, shape, mode))
+
+ def get_group(self, ranks_in_group: List[int], backend: Optional[str] = None) -> ProcessGroup:
+ """Get the process group with the given ranks. It the process group doesn't exist, it will be created.
+
+ Args:
+ ranks_in_group (List[int]): Ranks in the process group.
+ backend (Optional[str], optional): Backend of the process group. Defaults to None.
+
+ Returns:
+ ProcessGroup: The process group with the given ranks.
+ """
+ ranks_in_group = sorted(ranks_in_group)
+ if tuple(ranks_in_group) not in self._group_to_ranks:
+ group = dist.new_group(ranks_in_group, backend=backend)
+ self._ranks_to_group[tuple(ranks_in_group)] = group
+ self._group_to_ranks[group] = tuple(ranks_in_group)
+ return self._ranks_to_group[tuple(ranks_in_group)]
+
+ def get_ranks_in_group(self, group: ProcessGroup) -> List[int]:
+ """Get the ranks in the given process group. The process group must be created by this class.
+
+ Args:
+ group (ProcessGroup): The process group.
+
+ Returns:
+ List[int]: Ranks in the process group.
+ """
+ return list(self._group_to_ranks[group])
+
+ @staticmethod
+ def get_coords_along_axis(
+ base_coord: Tuple[int, ...], axis: int, indices_at_axis: List[int]
+ ) -> List[Tuple[int, ...]]:
+ """Get coordinates along the given axis.
+
+ Args:
+ base_coord (Tuple[int, ...]): Base coordinate which the coordinates along the axis are based on.
+ axis (int): Axis along which the coordinates are generated.
+ indices_at_axis (List[int]): Indices at the axis.
+
+ Returns:
+ List[Tuple[int, ...]]: Coordinates along the axis.
+ """
+ coords_in_group = []
+ for idx in indices_at_axis:
+ coords_in_group.append(base_coord[:axis] + (idx,) + base_coord[axis + 1 :])
+ return coords_in_group
+
+ def create_group_along_axis(
+ self, axis: int, indices_at_axis: Optional[List[int]] = None, backend: Optional[str] = None
+ ) -> ProcessGroup:
+ """Create all process groups along the given axis, and return the one which the current process belongs to.
+
+ Args:
+ axis (int): Axis along which the process groups are created.
+ indices_at_axis (Optional[List[int]], optional): Indices at the axis. Defaults to None.
+ backend (Optional[str], optional): Backend of the process group. Defaults to None.
+
+ Returns:
+ ProcessGroup: The process group along the given axis which the current process belongs to.
+ """
+ indices_at_axis = indices_at_axis or list(range(self._shape[axis]))
+ reduced_shape = list(self._shape)
+ # the choices on the axis are reduced to 1, since it's determined by `indices_at_axis`
+ reduced_shape[axis] = 1
+ target_group = None
+ # use Cartesian product to generate all combinations of coordinates
+ for base_coord in itertools.product(*[range(s) for s in reduced_shape]):
+ coords_in_group = ProcessGroupMesh.get_coords_along_axis(base_coord, axis, indices_at_axis)
+ ranks_in_group = tuple([ProcessGroupMesh.ravel(coord, self._shape) for coord in coords_in_group])
+ group = self.get_group(ranks_in_group, backend=backend)
+ if self._rank in ranks_in_group:
+ target_group = group
+ return target_group
+
+ def get_group_along_axis(
+ self, axis: int, indices_at_axis: Optional[List[int]] = None, backend: Optional[str] = None
+ ) -> ProcessGroup:
+ """Get the process group along the given axis which the current process belongs to. If the process group doesn't exist, it will be created.
+
+ Args:
+ axis (int): Axis along which the process groups are created.
+ indices_at_axis (Optional[List[int]], optional): Indices at the axis. Defaults to None.
+ backend (Optional[str], optional): Backend of the process group. Defaults to None.
+
+ Returns:
+ ProcessGroup: The process group along the given axis which the current process belongs to.
+ """
+ indices_at_axis = indices_at_axis or list(range(self._shape[axis]))
+ coords_in_group = ProcessGroupMesh.get_coords_along_axis(self._coord, axis, indices_at_axis)
+ ranks_in_group = tuple([ProcessGroupMesh.ravel(coord, self._shape) for coord in coords_in_group])
+ if ranks_in_group not in self._ranks_to_group:
+ # no need to cache it explicitly, since it will be cached in `create_group_along_axis`
+ return self.create_group_along_axis(axis, indices_at_axis, backend=backend)
+ return self._ranks_to_group[ranks_in_group]
+
+from torch.distributed import ProcessGroup
+
+
+class ProcessGroupManager(ProcessGroupMesh):
+ def __init__(self, *size: int, dp_axis, sp_axis):
+ super().__init__(*size)
+ self.dp_axis = dp_axis
+ self.sp_axis = sp_axis
+ self._dp_group: ProcessGroup = self.get_group_along_axis(self.dp_axis)
+ self._sp_group: ProcessGroup = self.get_group_along_axis(self.sp_axis)
+
+ @property
+ def dp_group(self) -> ProcessGroup:
+ return self._dp_group
+
+ @property
+ def sp_group(self) -> ProcessGroup:
+ return self._sp_group
diff --git a/grn/utils_t2iv/dist.py b/grn/utils_t2iv/dist.py
new file mode 100644
index 0000000000000000000000000000000000000000..675b6f03d0be774f713a4c34778d550daef3ef08
--- /dev/null
+++ b/grn/utils_t2iv/dist.py
@@ -0,0 +1,325 @@
+import datetime
+import functools
+import os
+import sys
+from typing import List
+from typing import Union
+
+import pytz
+import torch
+import torch.distributed as tdist
+import torch.multiprocessing as mp
+
+
+__rank, __local_rank, __world_size, __device = 0, 0, 1, 'cpu'
+__rank_str_zfill = '0'
+__initialized = False
+
+
+def initialized():
+ return __initialized
+
+
+def __initialize(fork=False, backend='nccl', gpu_id_if_not_distibuted=0, timeout_minutes=30):
+ global __device
+ if not torch.cuda.is_available():
+ print(f'[dist initialize] cuda is not available, use cpu instead', file=sys.stderr)
+ return
+ elif 'RANK' not in os.environ:
+ torch.cuda.set_device(gpu_id_if_not_distibuted)
+ __device = torch.empty(1).cuda().device
+ print(f'[dist initialize] env variable "RANK" is not set, use {__device} as the device', file=sys.stderr)
+ return
+ # then 'RANK' must exist
+ global_rank, num_gpus = int(os.environ['RANK']), torch.cuda.device_count()
+ local_rank = global_rank % num_gpus
+ torch.cuda.set_device(local_rank)
+
+ # ref: https://github.com/open-mmlab/mmcv/blob/master/mmcv/runner/dist_utils.py#L29
+ """
+ if mp.get_start_method(allow_none=True) is None:
+ method = 'fork' if fork else 'spawn'
+ print(f'[dist initialize] mp method={method}')
+ mp.set_start_method(method)
+ """
+ tdist.init_process_group(backend=backend, timeout=datetime.timedelta(seconds=timeout_minutes * 60))
+
+ global __rank, __local_rank, __world_size, __initialized, __rank_str_zfill
+ __local_rank = local_rank
+ __rank, __world_size = tdist.get_rank(), tdist.get_world_size()
+ __rank_str_zfill = str(__rank).zfill(len(str(__world_size)))
+ __device = torch.device(local_rank)
+ __initialized = True
+
+ assert tdist.is_initialized(), 'torch.distributed is not initialized!'
+ print(f'[lrk={get_local_rank()}, rk={get_rank()}]')
+
+
+def get_rank():
+ return __rank
+
+
+def get_rank_given_group(group: tdist.ProcessGroup):
+ return tdist.get_rank(group=group)
+
+
+def get_rank_str_zfill():
+ return __rank_str_zfill
+
+
+def get_local_rank():
+ return __local_rank
+
+
+def get_world_size():
+ return __world_size
+
+
+def get_device():
+ return __device
+
+
+def set_gpu_id(gpu_id: int):
+ if gpu_id is None: return
+ global __device
+ if isinstance(gpu_id, (str, int)):
+ torch.cuda.set_device(int(gpu_id))
+ __device = torch.empty(1).cuda().device
+ else:
+ raise NotImplementedError
+
+
+def is_master():
+ return __rank == 0
+
+
+def is_local_master():
+ return __local_rank == 0
+
+
+def is_visualizer():
+ return __rank == 0
+ # return __rank == max(__world_size - 8, 0)
+
+
+def parallelize(net, syncbn=False):
+ if syncbn:
+ net = torch.nn.SyncBatchNorm.convert_sync_batchnorm(net)
+ net = net.cuda()
+ net = torch.nn.parallel.DistributedDataParallel(net, device_ids=[get_local_rank()], find_unused_parameters=False, broadcast_buffers=False)
+ return net
+
+
+def new_group(ranks: List[int]):
+ if __initialized:
+ return tdist.new_group(ranks=ranks)
+ return None
+
+
+def new_local_machine_group():
+ if __initialized:
+ cur_subgroup, subgroups = tdist.new_subgroups()
+ return cur_subgroup
+ return None
+
+
+def barrier():
+ if __initialized:
+ tdist.barrier()
+
+
+def allreduce(t: torch.Tensor, async_op=False):
+ if __initialized:
+ if not t.is_cuda:
+ cu = t.detach().cuda()
+ ret = tdist.all_reduce(cu, async_op=async_op)
+ t.copy_(cu.cpu())
+ else:
+ ret = tdist.all_reduce(t, async_op=async_op)
+ return ret
+ return None
+
+
+def allgather(t: torch.Tensor, cat=True) -> Union[List[torch.Tensor], torch.Tensor]:
+ if __initialized:
+ if not t.is_cuda:
+ t = t.cuda()
+ ls = [torch.empty_like(t) for _ in range(__world_size)]
+ tdist.all_gather(ls, t)
+ else:
+ ls = [t]
+ if cat:
+ ls = torch.cat(ls, dim=0)
+ return ls
+
+
+def allgather_diff_shape(t: torch.Tensor, cat=True) -> Union[List[torch.Tensor], torch.Tensor]:
+ if __initialized:
+ if not t.is_cuda:
+ t = t.cuda()
+
+ t_size = torch.tensor(t.size(), device=t.device)
+ ls_size = [torch.empty_like(t_size) for _ in range(__world_size)]
+ tdist.all_gather(ls_size, t_size)
+
+ max_B = max(size[0].item() for size in ls_size)
+ pad = max_B - t_size[0].item()
+ if pad:
+ pad_size = (pad, *t.size()[1:])
+ t = torch.cat((t, t.new_empty(pad_size)), dim=0)
+
+ ls_padded = [torch.empty_like(t) for _ in range(__world_size)]
+ tdist.all_gather(ls_padded, t)
+ ls = []
+ for t, size in zip(ls_padded, ls_size):
+ ls.append(t[:size[0].item()])
+ else:
+ ls = [t]
+ if cat:
+ ls = torch.cat(ls, dim=0)
+ return ls
+
+
+def broadcast(t: torch.Tensor, src_rank) -> None:
+ if __initialized:
+ if not t.is_cuda:
+ cu = t.detach().cuda()
+ tdist.broadcast(cu, src=src_rank)
+ t.copy_(cu.cpu())
+ else:
+ tdist.broadcast(t, src=src_rank)
+
+
+def dist_fmt_vals(val: float, fmt: Union[str, None] = '%.2f') -> Union[torch.Tensor, List]:
+ if not initialized():
+ return torch.tensor([val]) if fmt is None else [fmt % val]
+
+ ts = torch.zeros(__world_size)
+ ts[__rank] = val
+ allreduce(ts)
+ if fmt is None:
+ return ts
+ return [fmt % v for v in ts.cpu().numpy().tolist()]
+
+
+def master_only(func):
+ @functools.wraps(func)
+ def wrapper(*args, **kwargs):
+ force = kwargs.pop('force', False)
+ if force or is_master():
+ ret = func(*args, **kwargs)
+ else:
+ ret = None
+ barrier()
+ return ret
+ return wrapper
+
+
+def local_master_only(func):
+ @functools.wraps(func)
+ def wrapper(*args, **kwargs):
+ force = kwargs.pop('force', False)
+ if force or is_local_master():
+ ret = func(*args, **kwargs)
+ else:
+ ret = None
+ barrier()
+ return ret
+ return wrapper
+
+
+def for_visualize(func):
+ @functools.wraps(func)
+ def wrapper(*args, **kwargs):
+ if is_visualizer():
+ # with torch.no_grad():
+ ret = func(*args, **kwargs)
+ else:
+ ret = None
+ return ret
+ return wrapper
+
+
+def finalize():
+ if __initialized:
+ tdist.destroy_process_group()
+
+
+def init_distributed_mode(local_out_path, fork=False, only_sync_master=False, timeout_minutes=30):
+ try:
+ __initialize(fork=fork, timeout_minutes=timeout_minutes)
+ barrier()
+ except RuntimeError as e:
+ print(f'{"!"*80} dist init error (NCCL Error?), stopping training! {"!"*80}', flush=True)
+ raise e
+
+ if local_out_path is not None: os.makedirs(local_out_path, exist_ok=True)
+ _change_builtin_print(is_local_master())
+ if (is_master() if only_sync_master else is_local_master()) and local_out_path is not None and len(local_out_path):
+ sys.stdout, sys.stderr = BackupStreamToFile(local_out_path, for_stdout=True), BackupStreamToFile(local_out_path, for_stdout=False)
+
+
+def _change_builtin_print(is_master):
+ import builtins as __builtin__
+
+ builtin_print = __builtin__.print
+ if type(builtin_print) != type(open):
+ return
+
+ def prt(*args, **kwargs):
+ force = kwargs.pop('force', False)
+ clean = kwargs.pop('clean', False)
+ deeper = kwargs.pop('deeper', False)
+ if is_master or force:
+ if not clean:
+ f_back = sys._getframe().f_back
+ if deeper and f_back.f_back is not None:
+ f_back = f_back.f_back
+ file_desc = f'{f_back.f_code.co_filename:24s}'[-24:]
+ time_str = datetime.datetime.now(tz=pytz.timezone('Asia/Shanghai')).strftime('[%m-%d %H:%M:%S]')
+ builtin_print(f'{time_str} ({file_desc}, line{f_back.f_lineno:-4d})=>', *args, **kwargs)
+ else:
+ builtin_print(*args, **kwargs)
+
+ __builtin__.print = prt
+
+
+class BackupStreamToFile(object):
+ def __init__(self, local_output_dir, for_stdout=True):
+ self.for_stdout = for_stdout
+ self.terminal_stream = sys.stdout if for_stdout else sys.stderr
+ fname = os.path.join(local_output_dir, 'b1_stdout.txt' if for_stdout else 'b2_stderr.txt')
+ existing = os.path.exists(fname)
+ self.file_stream = open(fname, 'a')
+ if existing:
+ time_str = datetime.datetime.now(tz=pytz.timezone('Asia/Shanghai')).strftime('[%m-%d %H:%M:%S]')
+ self.file_stream.write('\n'*7 + '='*55 + f' RESTART {time_str} ' + '='*55 + '\n')
+ self.file_stream.flush()
+ self.enabled = True
+
+ def write(self, message):
+ self.terminal_stream.write(message)
+ self.file_stream.write(message)
+
+ def flush(self):
+ self.terminal_stream.flush()
+ self.file_stream.flush()
+
+ def isatty(self):
+ return True
+
+ def close(self):
+ if not self.enabled:
+ return
+ self.enabled = False
+ self.file_stream.flush()
+ self.file_stream.close()
+ if self.for_stdout:
+ sys.stdout = self.terminal_stream
+ sys.stdout.flush()
+ else:
+ sys.stderr = self.terminal_stream
+ sys.stderr.flush()
+
+ def __del__(self):
+ self.close()
diff --git a/grn/utils_t2iv/hbq_util_t2iv.py b/grn/utils_t2iv/hbq_util_t2iv.py
new file mode 100644
index 0000000000000000000000000000000000000000..132eff3060ce099cc689351444b7c5e0dd59882c
--- /dev/null
+++ b/grn/utils_t2iv/hbq_util_t2iv.py
@@ -0,0 +1,50 @@
+import torch
+
+def index_label2quant_features(pred_sample_labels, hbq_round):
+ approx_signal = 0.
+ pred_sample_labels = pred_sample_labels.to(torch.long)
+ for round_ind in range(hbq_round):
+ interval = (1/2)**(round_ind+1)
+ base = 2**(hbq_round-1-round_ind)
+ approx_signal = approx_signal + interval * torch.where(pred_sample_labels>=base, +1, -1)
+ pred_sample_labels = pred_sample_labels % base
+ return approx_signal
+
+def raw_feature2index_label(feature, hbq_round):
+ quant_features = 0.
+ labels = 0
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ labels = labels * 2 + torch.where(feature > quant_features, 1, 0)
+ quant_features = quant_features + torch.where(feature > quant_features, interval, -interval)
+ return labels
+
+def raw_feature2bit_label(feature, hbq_round):
+ # [B,d,t,h,w] -> [B,hbq_round_mul_d,t,h,w]
+ quant_features = 0.
+ labels = []
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ labels.append(torch.where(feature > quant_features, 1, 0))
+ quant_features = quant_features + torch.where(feature > quant_features, interval, -interval)
+ labels = torch.stack(labels, dim=1) # [B,hbq_round,d,t,h,w]
+ B, _, d, t, h, w = labels.shape
+ labels = labels.reshape(B, hbq_round * d, t, h, w)
+ return labels
+
+def bit_label2raw_feature(bit_labels, hbq_round):
+ # [B,hbq_round_mul_d,t,h,w] -> [B,d,t,h,w]
+ B, hbq_round_mul_d, t, h, w = bit_labels.shape
+ d = hbq_round_mul_d // hbq_round
+ bit_labels = bit_labels.reshape(B, hbq_round, d, t, h, w).to(torch.long)
+ raw_features = 0.
+ for round_ind in range(hbq_round):
+ interval = (1/2) ** (round_ind + 1) # 0.5, 0.25, 0.125, ...
+ raw_features = raw_features + interval * torch.where(bit_labels[:,round_ind] == 1, +1, -1)
+ return raw_features # [B,d,t,h,w]
+
+def multiclass_labels2onehot_input(labels, num_classes):
+ B,d,t,h,w = labels.shape
+ onehot_input = torch.nn.functional.one_hot(labels.to(torch.long), num_classes) # [B,d,t,h,w] -> [B,d,t,h,w,num_classes]
+ onehot_input = onehot_input.permute(0,1,5,2,3,4).reshape(B,d*num_classes,t,h,w).float() # [B,d,t,h,w,num_classes] -> [B,d,num_classes,t,h,w] -> [B,d*num_classes,t,h,w]
+ return onehot_input
diff --git a/grn/utils_t2iv/infer.py b/grn/utils_t2iv/infer.py
new file mode 100644
index 0000000000000000000000000000000000000000..895edb5c0f5df111c73148562364854543542148
--- /dev/null
+++ b/grn/utils_t2iv/infer.py
@@ -0,0 +1,259 @@
+import hashlib
+import os
+import os.path as osp
+import re
+import time
+
+os.environ["TOKENIZERS_PARALLELISM"] = "false"
+
+import cv2
+import imageio
+import numpy as np
+import torch
+from PIL import Image
+from timm.models import create_model
+from torchvision.transforms.functional import to_tensor
+
+torch._dynamo.config.cache_size_limit = 64
+
+from grn.models.basic import *
+from grn.models.grn import GRN
+from grn.models.umt5.t5 import T5EncoderModel
+
+def extract_key_val(text):
+ return {k: v.lstrip() for k, v in re.findall(r'<(.+?):(.+?)>', text)}
+
+def encode_prompt(t5_path, text_tokenizer, text_encoder, prompt, args=None):
+ print(f't5 encode prompt: {prompt}')
+ text_encoder.model.to(args.other_device)
+ text_features = text_encoder([prompt], args.other_device)
+ lens = [len(item) for item in text_features]
+ cu_seqlens_k = [0]
+ for len_i in lens:
+ cu_seqlens_k.append(cu_seqlens_k[-1] + len_i)
+ cu_seqlens_k = torch.tensor(cu_seqlens_k, dtype=torch.int32)
+ Ltext = max(lens)
+ kv_compact = torch.cat(text_features, dim=0).float()
+ kv_compact = kv_compact.to(args.other_device)
+ text_cond_tuple = (kv_compact, lens, cu_seqlens_k, Ltext)
+ return text_cond_tuple
+
+def gen_one_example(
+ model,
+ vae,
+ text_tokenizer,
+ text_encoder,
+ prompt,
+ cfg_list=[],
+ tau_list=[],
+ negative_prompt="",
+ scale_schedule=None,
+ top_k=900,
+ top_p=0.97,
+ cfg_sc=3,
+ cfg_exp_k=0.0,
+ cfg_insertion_layer=-5,
+ vae_latent_dim=0,
+ gumbel=0,
+ softmax_merge_topk=-1,
+ gt_leak=-1,
+ gt_ls_Bl=None,
+ g_seed=None,
+ input_use_interplote_up=False,
+ args=None,
+ get_visual_rope_embeds=None,
+ noise_list=None,
+ return_summed_code_only=False,
+ class_token_id=0,
+ first_frame_condition=False,
+):
+ sstt = time.time()
+ if not isinstance(cfg_list, list):
+ cfg_list = [cfg_list] * len(scale_schedule)
+ if not isinstance(tau_list, list):
+ tau_list = [tau_list] * len(scale_schedule)
+
+ text_cond_tuple = []
+ for prompt_str in prompt:
+ text_cond_tuple.append(encode_prompt(args.text_encoder_ckpt, text_tokenizer, text_encoder, prompt_str, args=args))
+ negative_label_B_or_BLT = encode_prompt(args.text_encoder_ckpt, text_tokenizer, text_encoder, negative_prompt, args=args)
+ print(f'cfg: {cfg_list}, tau: {tau_list}')
+ with torch.cuda.amp.autocast(enabled=True, dtype=torch.bfloat16, cache_enabled=True):
+ stt = time.time()
+ out = model.autoregressive_infer(
+ vae=vae,
+ scale_schedule=scale_schedule,
+ label_B_or_BLT=text_cond_tuple, g_seed=g_seed,
+ B=1, negative_label_B_or_BLT=negative_label_B_or_BLT, force_gt_Bhw=None,
+ cfg_sc=cfg_sc, cfg_list=cfg_list, tau_list=tau_list, top_k=top_k, top_p=top_p,
+ returns_vemb=1, ratio_Bl1=None, gumbel=gumbel, norm_cfg=False,
+ cfg_exp_k=cfg_exp_k, cfg_insertion_layer=cfg_insertion_layer,
+ vae_latent_dim=vae_latent_dim, softmax_merge_topk=softmax_merge_topk,
+ ret_img=True, trunk_scale=1000,
+ gt_leak=gt_leak, gt_ls_Bl=gt_ls_Bl, inference_mode=True,
+ input_use_interplote_up=input_use_interplote_up,
+ args=args,
+ get_visual_rope_embeds=get_visual_rope_embeds,
+ noise_list=noise_list,
+ return_summed_code_only=return_summed_code_only,
+ class_token_id=class_token_id,
+ first_frame_condition=first_frame_condition,
+ )
+ _, pred_multi_scale_bit_labels, img_list = out
+
+ print(f"cost: {time.time() - sstt}, model cost={time.time() - stt}")
+ img = img_list[0]
+ return img
+
+def get_prompt_id(prompt):
+ return hash_string(prompt)
+
+def load_tokenizer(t5_path ='', device='cuda'):
+ print(f'[Loading tokenizer and text encoder]')
+ if isinstance(device, str):
+ device = torch.device(device)
+ text_encoder = T5EncoderModel(
+ text_len=512,
+ dtype=torch.bfloat16,
+ device=device,
+ checkpoint_path=osp.join(t5_path, 'models_t5_umt5-xxl-enc-bf16.pth'),
+ tokenizer_path=osp.join(t5_path, 'umt5-xxl'),
+ enable_fsdp=False)
+ text_tokenizer = text_encoder.tokenizer
+ return text_tokenizer, text_encoder
+
+def transform(pil_img, tgt_h, tgt_w):
+ width, height = pil_img.size
+ if width / height <= tgt_w / tgt_h:
+ resized_width = tgt_w
+ resized_height = int(np.round(tgt_w / (width / height)))
+ else:
+ resized_height = tgt_h
+ resized_width = int(np.round((width / height) * tgt_h))
+ pil_img = pil_img.resize((resized_width, resized_height), resample=Image.LANCZOS)
+ # crop the center out
+ arr = np.array(pil_img)
+ crop_y = (arr.shape[0] - tgt_h) // 2
+ crop_x = (arr.shape[1] - tgt_w) // 2
+ im = to_tensor(arr[crop_y: crop_y + tgt_h, crop_x: crop_x + tgt_w])
+ return im * 2 - 1
+
+def hash_string(input_string):
+ md5 = hashlib.md5()
+ md5.update(input_string.encode('utf-8'))
+ return md5.hexdigest()
+
+def joint_vi_vae_encode_decode(vae, image_path, scale_schedule, device, tgt_h, tgt_w):
+ pil_image = Image.open(image_path).convert('RGB')
+ inp = transform(pil_image, tgt_h, tgt_w)
+ inp = inp.unsqueeze(0).to(device)
+ scale_schedule = [(item[0], item[1], item[2]) for item in scale_schedule]
+ t1 = time.time()
+ h, z, _, all_bit_indices, _, _ = vae.encode(inp, scale_schedule=scale_schedule)
+ t2 = time.time()
+ recons_img = vae.decode(z)[0]
+ if len(recons_img.shape) == 4:
+ recons_img = recons_img.squeeze(1)
+ print(f'recons: z.shape: {z.shape}, recons_img shape: {recons_img.shape}')
+ t3 = time.time()
+ print(f'vae encode takes {t2-t1:.2f}s, decode takes {t3-t2:.2f}s')
+ recons_img = ((recons_img + 1) / 2).clamp(0, 1)
+ recons_img = recons_img.permute(1, 2, 0).mul(255).cpu().numpy().astype(np.uint8)
+ gt_img = ((inp[0] + 1) / 2).clamp(0, 1)
+ gt_img = gt_img.permute(1, 2, 0).mul(255).cpu().numpy().astype(np.uint8)
+ print(recons_img.shape, gt_img.shape)
+ return gt_img, recons_img, all_bit_indices
+
+
+def load_transformer(vae, args):
+ device = torch.device(args.other_device)
+ model_path = args.model_path
+ if not model_path:
+ state_dict = None
+ elif args.checkpoint_type == 'torch':
+ # copy large model to local; save slim to local; and copy slim to nas; load local slim model
+ slim_model_path = model_path
+ print(f'load checkpoint from {slim_model_path}')
+ state_dict = torch.load(slim_model_path, map_location=device)
+
+ print(f'[Loading Model]')
+ # Check if device is CUDA before enabling autocast
+ if device.type == 'cuda':
+ with torch.cuda.amp.autocast(enabled=True, dtype=torch.bfloat16, cache_enabled=True), torch.no_grad():
+ model = create_model(
+ args.model,
+ vae_local=vae, text_channels=args.text_channels, text_maxlen=512,
+ shared_aln=True, raw_scale_schedule=None,
+ checkpointing='full-block',
+ customized_flash_attn=False,
+ fused_norm=True,
+ pad_to_multiplier=128,
+ use_flex_attn=False,
+ num_of_label_value=args.num_of_label_value,
+ rope2d_normalized_by_hw=args.rope2d_normalized_by_hw,
+ pn=args.pn,
+ apply_spatial_patchify=args.apply_spatial_patchify,
+ inference_mode=True,
+ train_h_div_w_list=args.train_h_div_w_list,
+ dynamic_scale_schedule=args.dynamic_scale_schedule,
+ video_frames=args.video_frames,
+ other_args=args,
+ ).to(device=device)
+ else:
+ with torch.no_grad():
+ model = create_model(
+ args.model,
+ vae_local=vae, text_channels=args.text_channels, text_maxlen=512,
+ shared_aln=True, raw_scale_schedule=None,
+ checkpointing='full-block',
+ customized_flash_attn=False,
+ fused_norm=True,
+ pad_to_multiplier=128,
+ use_flex_attn=False,
+ num_of_label_value=args.num_of_label_value,
+ rope2d_normalized_by_hw=args.rope2d_normalized_by_hw,
+ pn=args.pn,
+ apply_spatial_patchify=args.apply_spatial_patchify,
+ inference_mode=True,
+ train_h_div_w_list=args.train_h_div_w_list,
+ dynamic_scale_schedule=args.dynamic_scale_schedule,
+ video_frames=args.video_frames,
+ other_args=args,
+ ).to(device=device)
+ print(f'[you selected model with {args.model}] model size: {sum(p.numel() for p in model.parameters())/1e9:.2f}B, bf16={args.bf16}')
+ if args.bf16:
+ for block in model.unregistered_blocks:
+ block.bfloat16()
+ model.eval()
+ model.requires_grad_(False)
+ # Only call cuda() if device is CUDA
+ if device.type == 'cuda':
+ model.cuda()
+ torch.cuda.empty_cache()
+ print(f'[Load model weights]')
+ if state_dict:
+ if 'trainer' in state_dict:
+ print(model.load_state_dict(state_dict['trainer']['gpt_fsdp'], strict=True))
+ else:
+ print(model.load_state_dict(state_dict, strict=True))
+ return model
+
+def images2video(ndarray_image_list, fps=24, save_filepath='tmp.mp4'):
+ # ndarray_image_list: bgr sequence
+ save_dir = osp.dirname(save_filepath)
+ if save_dir:
+ os.makedirs(save_dir, exist_ok=True)
+
+ if len(ndarray_image_list) == 1:
+ save_filepath = osp.splitext(save_filepath)[0] + '.jpg'
+ cv2.imwrite(save_filepath, ndarray_image_list[0]) # bgr
+ print(f"Image saved as {osp.abspath(save_filepath)}")
+ else:
+ # imageio takes rgb, so convert bgr to rgb
+ imageio.mimsave(save_filepath, ndarray_image_list[..., ::-1], fps=fps)
+ print(f"Video saved as {osp.abspath(save_filepath)}")
+
+def imgs_tensor2uint8_imgs(imgs_tensor):
+ imgs_tensor = imgs_tensor.permute(1, 2, 3, 0) # [c,t,h,w] -> [t,h,w,c]
+ imgs_tensor = ((imgs_tensor + 1) / 2).clamp(0, 1)
+ return imgs_tensor.mul(255).to(torch.uint8).flip(dims=(3,)).cpu().numpy()
\ No newline at end of file
diff --git a/grn/utils_t2iv/load.py b/grn/utils_t2iv/load.py
new file mode 100644
index 0000000000000000000000000000000000000000..8d3a63a77ad57cd0294ed7cc8265fab6d19b18ed
--- /dev/null
+++ b/grn/utils_t2iv/load.py
@@ -0,0 +1,55 @@
+import gc
+import os
+import os.path as osp
+import random
+import sys
+from copy import deepcopy
+from typing import Tuple, Union
+
+import torch
+import yaml
+
+from grn.models.grn import *
+from grn.models.hbq_tokenizer import HBQ_Tokenizer
+from timm.models import create_model
+
+def load_visual_tokenizer(args, device=None):
+ if not device:
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
+ vae = HBQ_Tokenizer(args=args, latent_channels=args.detail_scale_dim, encoder_out_type='feature_tanh')
+ vae.eval()
+ vae = vae.to('cuda')
+ for param in vae.parameters():
+ param.requires_grad = False
+ state_dict = torch.load(args.vae_path, map_location='cuda')
+ if 'ema' in state_dict:
+ print(f'Load ema vae weights')
+ state_dict = state_dict['ema']
+ else:
+ print(f'Load non ema vae weights')
+ state_dict = state_dict['vae']
+ print('Load vae: ', vae.load_state_dict(state_dict, assign=True))
+ return vae
+
+def build_vae_gpt(args, device='cuda'):
+ vae_local = load_visual_tokenizer(args, device)
+ gpt_kw = dict(
+ pretrained=False, global_pool='',
+ text_channels=args.Ct5, text_maxlen=args.tlen,
+ norm_eps=args.norm_eps,
+ top_p=args.tp, top_k=args.tk, tau=args.tau,
+ checkpointing=args.enable_checkpointing,
+ pad_to_multiplier=args.pad_to_multiplier,
+ use_flex_attn=args.use_flex_attn,
+ num_of_label_value=args.num_of_label_value,
+ train_h_div_w_list=None,
+ apply_spatial_patchify=args.apply_spatial_patchify,
+ dynamic_scale_schedule=args.dynamic_scale_schedule,
+ video_frames=args.video_frames,
+ other_args=args,
+ )
+ print(f'[create gpt_wo_ddp] constructor kw={gpt_kw}\n')
+ gpt_kw['vae_local'] = vae_local
+ gpt_wo_ddp = create_model(args.model, **gpt_kw)
+ assert all(p.requires_grad for n, p in gpt_wo_ddp.named_parameters())
+ return vae_local, gpt_wo_ddp
diff --git a/grn/utils_t2iv/lr_control.py b/grn/utils_t2iv/lr_control.py
new file mode 100644
index 0000000000000000000000000000000000000000..9c27e680f84d200dad242129572a149e3b6616e5
--- /dev/null
+++ b/grn/utils_t2iv/lr_control.py
@@ -0,0 +1,61 @@
+import math
+from pprint import pformat
+from typing import Tuple, List, Dict, Union
+
+import torch.nn
+import grn.utils_t2iv.dist as dist
+
+
+def filter_params(model, ndim_dict, nowd_keys=(), lr_scale=0.0) -> Tuple[
+ List[str], List[torch.nn.Parameter], List[Dict[str, Union[torch.nn.Parameter, float]]]
+]:
+ with_lr_scale = hasattr(model, 'get_layer_id_and_scale_exp') and 0 < lr_scale <= 1
+ print(f'[get_param_groups][lr decay] with_lr_scale={with_lr_scale}, lr_scale={lr_scale}')
+ para_groups, para_groups_dbg = {}, {}
+ names, paras = [], []
+ names_no_grad = []
+ count, numel = 0, 0
+ for name, para in model.named_parameters():
+ name = name.replace('_fsdp_wrapped_module.', '')
+ if not para.requires_grad:
+ names_no_grad.append(name)
+ continue # frozen weights
+ count += 1
+ numel += para.numel()
+ names.append(name)
+ paras.append(para)
+
+ if ndim_dict.get(name, 2) == 1 or name.endswith('bias') or any(k in name for k in nowd_keys):
+ cur_wd_sc, group_name = 0., 'ND'
+ else:
+ cur_wd_sc, group_name = 1., 'D'
+
+ if with_lr_scale:
+ layer_id, scale_exp = model.get_layer_id_and_scale_exp(name)
+ group_name = f'layer{layer_id}_' + group_name
+ cur_lr_sc = lr_scale ** scale_exp
+ dbg = f'[layer {layer_id}][sc = {lr_scale} ** {scale_exp}]'
+ else:
+ cur_lr_sc = 1.
+ dbg = f'[no scale]'
+
+ if group_name not in para_groups:
+ para_groups[group_name] = {'params': [], 'wd_sc': cur_wd_sc, 'lr_sc': cur_lr_sc}
+ para_groups_dbg[group_name] = {'params': [], 'wd_sc': cur_wd_sc, 'lr_sc': dbg}
+ para_groups[group_name]['params'].append(para)
+ para_groups_dbg[group_name]['params'].append(name)
+
+ for g in para_groups_dbg.values():
+ g['params'] = pformat(', '.join(g['params']), width=200)
+
+ print(f'[get_param_groups] param_groups = \n{pformat(para_groups_dbg, indent=2, width=240)}\n')
+
+ for rk in range(dist.get_world_size()):
+ dist.barrier()
+ if dist.get_rank() == rk:
+ print(f'[get_param_groups][rank{dist.get_rank()}] {type(model).__name__=} {count=}, {numel=}', flush=True, force=True)
+ print('')
+
+ assert len(names_no_grad) == 0, f'[get_param_groups] names_no_grad = \n{pformat(names_no_grad, indent=2, width=240)}\n'
+ del ndim_dict
+ return names, paras, list(para_groups.values())
diff --git a/grn/utils_t2iv/misc.py b/grn/utils_t2iv/misc.py
new file mode 100644
index 0000000000000000000000000000000000000000..50eb1cc07b30f6c92ad8a4f79e6f3cc7749ce776
--- /dev/null
+++ b/grn/utils_t2iv/misc.py
@@ -0,0 +1,363 @@
+import datetime
+import functools
+import math
+import os
+import random
+import subprocess
+import sys
+import shlex
+import threading
+import time
+from collections import defaultdict, deque
+from typing import Iterator, List, Tuple
+
+import numpy as np
+import pytz
+import torch
+import torch.distributed as tdist
+import torch.nn.functional as F
+
+import grn.utils_t2iv.dist as dist
+
+
+def os_system(cmd):
+ if isinstance(cmd, str):
+ cmd = shlex.split(cmd)
+ return subprocess.call(cmd)
+
+def echo(info):
+ time_str = datetime.datetime.now().strftime("%m-%d-%H:%M:%S")
+ caller = sys._getframe().f_back
+ print(f'[{time_str}] ({os.path.basename(caller.f_code.co_filename)}, line{caller.f_lineno})=> {info}')
+
+def os_system_get_stdout(cmd):
+ if isinstance(cmd, str):
+ cmd = shlex.split(cmd)
+ return subprocess.run(cmd, stdout=subprocess.PIPE).stdout.decode('utf-8')
+
+def os_system_get_stdout_stderr(cmd):
+ if isinstance(cmd, str):
+ cmd = shlex.split(cmd)
+ cnt = 0
+ while True:
+ try:
+ sp = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, timeout=30)
+ except subprocess.TimeoutExpired:
+ cnt += 1
+ print(f'[fetch free_port file] timeout cnt={cnt}')
+ else:
+ return sp.stdout.decode('utf-8'), sp.stderr.decode('utf-8')
+
+
+def is_pow2n(x):
+ return x > 0 and (x & (x - 1) == 0)
+
+
+def time_str(fmt='[%m-%d %H:%M:%S]'):
+ return datetime.datetime.now(tz=pytz.timezone('Asia/Shanghai')).strftime(fmt)
+
+
+class DistLogger(object):
+ def __init__(self, lg):
+ self._lg = lg
+
+ @staticmethod
+ def do_nothing(*args, **kwargs):
+ pass
+
+ def __getattr__(self, attr: str):
+ return getattr(self._lg, attr) if self._lg is not None else DistLogger.do_nothing
+
+class TensorboardLogger(object):
+ def __init__(self, log_dir, filename_suffix):
+ try: import tensorflow_io as tfio
+ except: pass
+ from torch.utils.tensorboard import SummaryWriter
+ self.writer = SummaryWriter(log_dir=log_dir, filename_suffix=filename_suffix)
+ self.step = 0
+
+ def set_step(self, step=None):
+ if step is not None:
+ self.step = step
+ else:
+ self.step += 1
+
+ def loggable(self):
+ return self.step == 0 or (self.step + 1) % 500 == 0
+
+ def update(self, head='scalar', step=None, **kwargs):
+ if step is None:
+ step = self.step
+ if not self.loggable(): return
+ for k, v in kwargs.items():
+ if v is None: continue
+ if hasattr(v, 'item'): v = v.item()
+ self.writer.add_scalar(f'{head}/{k}', v, step)
+
+ def log_tensor_as_distri(self, tag, tensor1d, step=None):
+ if step is None:
+ step = self.step
+ if not self.loggable(): return
+ try:
+ self.writer.add_histogram(tag=tag, values=tensor1d, global_step=step)
+ except Exception as e:
+ print(f'[log_tensor_as_distri writer.add_histogram failed]: {e}')
+
+ def log_image(self, tag, img_chw, step=None):
+ if step is None:
+ step = self.step
+ if not self.loggable(): return
+ self.writer.add_image(tag, img_chw, step, dataformats='CHW')
+
+ def flush(self):
+ self.writer.flush()
+
+ def close(self):
+ self.writer.close()
+
+class SmoothedValue(object):
+ """Track a series of values and provide access to smoothed values over a
+ window or the global series average.
+ """
+
+ def __init__(self, window_size=30, fmt=None):
+ if fmt is None:
+ fmt = "{median:.4f} ({global_avg:.4f})"
+ self.deque = deque(maxlen=window_size)
+ self.total = 0.0
+ self.count = 0
+ self.fmt = fmt
+
+ def update(self, value, n=1):
+ self.deque.append(value)
+ self.count += n
+ self.total += value * n
+
+ def synchronize_between_processes(self):
+ """
+ Warning: does not synchronize the deque!
+ """
+ t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda')
+ tdist.barrier()
+ tdist.all_reduce(t)
+ t = t.tolist()
+ self.count = int(t[0])
+ self.total = t[1]
+
+ @property
+ def median(self):
+ return np.median(self.deque) if len(self.deque) else 0
+
+ @property
+ def avg(self):
+ return sum(self.deque) / (len(self.deque) or 1)
+
+ @property
+ def global_avg(self):
+ return self.total / (self.count or 1)
+
+ @property
+ def max(self):
+ return max(self.deque) if len(self.deque) else 0
+
+ @property
+ def value(self):
+ return self.deque[-1] if len(self.deque) else 0
+
+ def time_preds(self, counts) -> Tuple[float, str, str]:
+ remain_secs = counts * self.median
+ return remain_secs, str(datetime.timedelta(seconds=round(remain_secs))), time.strftime("%Y-%m-%d %H:%M", time.localtime(time.time() + remain_secs))
+
+ def __str__(self):
+ return self.fmt.format(median=self.median, avg=self.avg, global_avg=self.global_avg, max=self.max, value=self.value)
+
+
+class MetricLogger(object):
+ def __init__(self):
+ self.meters = defaultdict(SmoothedValue)
+ self.iter_end_t = time.time()
+ self.log_iters = set()
+ self.log_every_iter = False
+
+ def update(self, **kwargs):
+ # if it != 0 and it not in self.log_iters: return
+ for k, v in kwargs.items():
+ if v is None: continue
+ if hasattr(v, 'item'): v = v.item()
+ # assert isinstance(v, (float, int)), type(v)
+ self.meters[k].update(v)
+
+ def __getattr__(self, attr):
+ if attr in self.meters:
+ return self.meters[attr]
+ if attr in self.__dict__:
+ return self.__dict__[attr]
+ raise AttributeError("'{}' object has no attribute '{}'".format(
+ type(self).__name__, attr))
+
+ def __str__(self):
+ loss_str = []
+ for name, meter in self.meters.items():
+ if len(meter.deque):
+ loss_str.append(
+ "{}: {}".format(name, str(meter))
+ )
+ return ' '.join(loss_str)
+
+ def synchronize_between_processes(self):
+ for meter in self.meters.values():
+ meter.synchronize_between_processes()
+
+ def add_meter(self, name, meter):
+ self.meters[name] = meter
+
+ def log_every(self, start_it, max_iters, itrt, log_freq, log_every_iter=False, header=''): # also solve logging & skipping iterations before start_it
+ start_it = start_it % max_iters
+ self.log_iters = set(range(start_it, max_iters, log_freq))
+ self.log_iters.add(start_it)
+ self.log_iters.add(max_iters-1)
+ self.log_iters.add(max_iters)
+ self.log_every_iter = log_every_iter
+ self.iter_end_t = time.time()
+ self.iter_time = SmoothedValue(fmt='{value:.4f}')
+ self.data_time = SmoothedValue(fmt='{value:.3f}')
+ header_fmt = header + ': [{0:' + str(len(str(max_iters))) + 'd}/{1}]'
+
+ start_time = time.time()
+ if isinstance(itrt, Iterator) and not hasattr(itrt, 'preload') and not hasattr(itrt, 'set_epoch'): # this
+ for it in range(start_it, max_iters):
+ obj = next(itrt)
+ if it < start_it: continue
+ self.data_time.update(time.time() - self.iter_end_t)
+ yield it, obj
+ self.iter_time.update(time.time() - self.iter_end_t)
+ if self.log_every_iter or it in self.log_iters:
+ eta_seconds = self.iter_time.avg * (max_iters - it)
+ print(f'{header_fmt.format(it, max_iters)} eta: {str(datetime.timedelta(seconds=int(eta_seconds)))} {str(self)} T: {self.iter_time.value:.3f}s dataT: {self.data_time.value*1e3:.1f}ms', flush=True)
+ self.iter_end_t = time.time()
+ else:
+ if isinstance(itrt, int): itrt = range(itrt)
+ for it, obj in enumerate(itrt):
+ if it < start_it:
+ self.iter_end_t = time.time()
+ continue
+ self.data_time.update(time.time() - self.iter_end_t)
+ yield it, obj
+ self.iter_time.update(time.time() - self.iter_end_t)
+ if self.log_every_iter or it in self.log_iters:
+ eta_seconds = self.iter_time.avg * (max_iters - it)
+ print(f'{header_fmt.format(it, max_iters)} eta: {str(datetime.timedelta(seconds=int(eta_seconds)))} {str(self)} T: {self.iter_time.value:.3f}s dataT: {self.data_time.value*1e3:.1f}ms', flush=True)
+ self.iter_end_t = time.time()
+ cost = time.time() - start_time
+ cost_str = str(datetime.timedelta(seconds=int(cost)))
+ print(f'{header} Cost of this ep: {cost_str} ({cost / (max_iters-start_it):.3f} s / it)', flush=True)
+
+
+class NullDDP(torch.nn.Module):
+ def __init__(self, module, *args, **kwargs):
+ super(NullDDP, self).__init__()
+ self.module = module
+ self.require_backward_grad_sync = False
+
+ def forward(self, *args, **kwargs):
+ return self.module(*args, **kwargs)
+
+
+def build_2d_sincos_position_embedding(h, w, embed_dim, temperature=10000., sc=0, verbose=True): # (1, hw**2, embed_dim)
+ # DiT: sc=0
+ # DETR: sc=2?
+ grid_w = torch.arange(w, dtype=torch.float32)
+ grid_h = torch.arange(h, dtype=torch.float32)
+ grid_w, grid_h = torch.meshgrid([grid_w, grid_h], indexing='ij')
+ if sc == 0:
+ scale = 1
+ elif sc == 1:
+ scale = math.pi * 2 / w
+ else:
+ scale = 1 / w
+ grid_w = scale * grid_w.reshape(h*w, 1) # scale * [0, 0, 0, 1, 1, 1, 2, 2, 2]
+ grid_h = scale * grid_h.reshape(h*w, 1) # scale * [0, 1, 2, 0, 1, 2, 0, 1, 2]
+
+ assert embed_dim % 4 == 0, f'Embed dimension ({embed_dim}) must be divisible by 4 for 2D sin-cos position embedding!'
+ pos_dim = embed_dim // 4
+ omega = torch.arange(pos_dim, dtype=torch.float32) / pos_dim
+ omega = (-math.log(temperature) * omega).exp()
+ # omega == (1/T) ** (arange(pos_dim) / pos_dim), a vector only dependent on C
+ out_w = grid_w * omega.view(1, pos_dim) # out_w: scale * [0*ome, 0*ome, 0*ome, 1*ome, 1*ome, 1*ome, 2*ome, 2*ome, 2*ome]
+ out_h = grid_h * omega.view(1, pos_dim) # out_h: scale * [0*ome, 1*ome, 2*ome, 0*ome, 1*ome, 2*ome, 0*ome, 1*ome, 2*ome]
+ pos_emb = torch.cat([torch.sin(out_w), torch.cos(out_w), torch.sin(out_h), torch.cos(out_h)], dim=1)[None, :, :]
+ if verbose: print(f'[build_2d_sincos_position_embedding @ {hw} x {hw}] scale_type={sc}, temperature={temperature:g}, shape={pos_emb.shape}')
+ return pos_emb # (1, hw**2, embed_dim)
+
+
+if __name__ == '__main__':
+ import seaborn as sns
+ import matplotlib.pyplot as plt
+ cmap_div = sns.color_palette('icefire', as_cmap=True)
+
+ scs = [0, 1, 2]
+ temps = [20, 50, 100, 1000]
+ reso = 3.0
+ RR, CC = len(scs), len(temps)
+ plt.figure(figsize=(CC * reso, RR * reso)) # figsize=(16, 16)
+ for row, sc in enumerate(scs):
+ for col, temp in enumerate(temps):
+ name = f'sc={sc}, T={temp}'
+ hw, C = 16, 512
+ N = hw*hw
+ pe = build_2d_sincos_position_embedding(hw, C, temperature=temp, sc=sc, verbose=False)[0] # N, C = 64, 16
+
+ hw2 = 16
+ N2 = hw2*hw2
+ pe2 = build_2d_sincos_position_embedding(hw2, C, temperature=temp, sc=sc, verbose=False)[0] # N, C = 64, 16
+ # pe2 = pe2.flip(dims=(0,))
+ bchw, bchw2 = F.normalize(pe.view(hw, hw, C).permute(2, 0, 1).unsqueeze(0), dim=1), F.normalize(pe2.view(hw2, hw2, C).permute(2, 0, 1).unsqueeze(0), dim=1)
+ dis = [
+ f'{F.mse_loss(bchw, F.interpolate(bchw2, size=bchw.shape[-2], mode=inter)).item():.3f}'
+ for inter in ('bilinear', 'bicubic', 'nearest')
+ ]
+ dis += [
+ f'{F.mse_loss(F.interpolate(bchw, size=bchw2.shape[-2], mode=inter), bchw2).item():.3f}'
+ for inter in ('area', 'nearest')
+ ]
+ print(f'[{name:^20s}] dis: {dis}')
+ """
+ [ sc=0, T=20 ] dis: ['0.010', '0.011', '0.011', '0.009', '0.010']
+ [ sc=0, T=100 ] dis: ['0.007', '0.007', '0.007', '0.006', '0.007']
+ [ sc=0, T=1000 ] dis: ['0.005', '0.005', '0.005', '0.004', '0.005']
+ [ sc=0, T=10000 ] dis: ['0.004', '0.004', '0.004', '0.003', '0.004']
+ [ sc=1, T=20 ] dis: ['0.007', '0.008', '0.008', '0.007', '0.008']
+ [ sc=1, T=100 ] dis: ['0.005', '0.005', '0.005', '0.005', '0.005']
+ [ sc=1, T=1000 ] dis: ['0.003', '0.003', '0.003', '0.003', '0.003']
+ [ sc=1, T=10000 ] dis: ['0.003', '0.003', '0.003', '0.003', '0.003']
+ [ sc=2, T=20 ] dis: ['0.000', '0.000', '0.000', '0.000', '0.000']
+ [ sc=2, T=100 ] dis: ['0.000', '0.000', '0.000', '0.000', '0.000']
+ [ sc=2, T=1000 ] dis: ['0.000', '0.000', '0.000', '0.000', '0.000']
+ [ sc=2, T=10000 ] dis: ['0.000', '0.000', '0.000', '0.000', '0.000']
+ Process finished with exit code 0
+ """
+
+ pe = torch.from_numpy(cmap_div(pe.T.numpy())[:, :, :3]) # C, N, 3
+ tar_h, tar_w = 1024, 1024
+ pe = pe.repeat_interleave(tar_w//pe.shape[0], dim=0).repeat_interleave(tar_h//pe.shape[1], dim=1)
+ plt.subplot(RR, CC, 1+row*CC+col)
+ plt.title(name)
+ plt.xlabel('hxw'), plt.ylabel('C')
+ plt.xticks([]), plt.yticks([])
+ plt.imshow(pe.mul(255).round().clamp(0, 255).byte().numpy())
+ plt.tight_layout(h_pad=0.02)
+ plt.show()
+
+
+def check_randomness(args):
+ U = 16384
+ t = torch.zeros(dist.get_world_size(), 4, dtype=torch.float32, device=args.device)
+ t0 = torch.zeros(1, dtype=torch.float32, device=args.device).random_(U)
+ t[dist.get_rank(), 0] = float(random.randrange(U))
+ t[dist.get_rank(), 1] = float(np.random.randint(U))
+ t[dist.get_rank(), 2] = float(torch.randint(0, U, (1,))[0])
+ t[dist.get_rank(), 3] = float(t0[0])
+ dist.allreduce(t)
+ for rk in range(1, dist.get_world_size()):
+ assert torch.allclose(t[rk - 1], t[rk]), f't={t}'
+ del t0, t, U
diff --git a/grn/utils_t2iv/save_and_load.py b/grn/utils_t2iv/save_and_load.py
new file mode 100644
index 0000000000000000000000000000000000000000..66443006b066d766e5693ad496746a3cb0e65d91
--- /dev/null
+++ b/grn/utils_t2iv/save_and_load.py
@@ -0,0 +1,157 @@
+import gc
+import os
+import os.path as osp
+import subprocess
+import time
+import re
+from typing import List, Optional, Tuple
+
+import torch
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+
+import glob
+import shutil
+from grn.utils_t2iv import arg_util
+import grn.utils_t2iv.dist as dist
+
+def glob_with_global_step(pattern, recursive=False):
+ def extract_ep_iter(filename):
+ match = re.search(r'global_step_(\d+)', filename)
+ if match:
+ iter_idx = int(match.group(1))
+ return iter_idx
+ return 0
+ return sorted(glob.glob(pattern, recursive=recursive), key=lambda x: extract_ep_iter(os.path.basename(x)), reverse=True)
+
+
+class CKPTSaver(object):
+ def __init__(self, is_master: bool, eval_milestone: List[Tuple[float, float]]):
+ self.is_master = is_master
+ self.time_stamp = torch.tensor([time.time() - 1e5, time.time()], device=dist.get_device())
+ self.sp_also: subprocess.Popen = None
+ self.sp_best: subprocess.Popen = None
+ self.sp_backup: subprocess.Popen = None
+ self.acc_str, self.eval_milestone = '[no acc str]', eval_milestone
+
+ def sav(
+ self, args: arg_util.Args, g_it: int, next_ep: int, next_it: int, trainer,
+ acc_str: Optional[str] = None, eval_milestone: Optional[List[Tuple[float, float]]] = None,
+ also_save_to: str = None, best_save_to: str = None,
+ ):
+
+ if acc_str is not None: self.acc_str = acc_str
+ if eval_milestone is not None: self.eval_milestone = eval_milestone
+
+ fname = f'global_step_{g_it}.pth'
+ local_out_ckpt = os.path.join(args.local_out_path, fname)
+
+ # NOTE: all rank should call this state_dict(), not master only!
+ trainer_state = trainer.state_dict()
+
+ if self.is_master:
+ stt = time.time()
+ torch.save({
+ 'args': args.state_dict(),
+ 'gpt_training': args.gpt_training,
+ 'arch': args.model,
+ 'epoch': next_ep,
+ 'iter': next_it,
+ 'trainer': trainer_state,
+ 'g_it': g_it,
+ }, local_out_ckpt)
+
+ print(f'[CKPTSaver][rank00] dbg: {args.bed=}', flush=True)
+ def auto_sync(source_filename, target_filename):
+ print(f'[CKPTSaver] auto_save {source_filename} -> {target_filename}', flush=True)
+ def _sync_worker():
+ try:
+ import shutil
+ from grn.utils.safe_rm import safe_remove
+ if os.path.isdir(source_filename):
+ shutil.copytree(source_filename, target_filename, dirs_exist_ok=True)
+ else:
+ shutil.copy2(source_filename, target_filename)
+ if source_filename.endswith('.pth') and (osp.abspath(source_filename) != osp.abspath(target_filename)):
+ safe_remove(source_filename, osp.dirname(source_filename))
+ except Exception as e:
+ print(f'[CKPTSaver] auto_save failed: {e}', flush=True)
+
+ import threading
+ self.sp_backup = threading.Thread(target=_sync_worker)
+ self.sp_backup.start()
+
+ local_files = glob.glob(f"{args.local_out_path}/*")
+ for filename in local_files:
+ basename = os.path.basename(filename)
+ target_filename = f'{args.bed}/{basename}'
+ if basename.endswith('.pth'):
+ if not os.path.isfile(target_filename):
+ auto_sync(filename, target_filename)
+ else:
+ auto_sync(filename, target_filename)
+ cost = time.time() - stt
+ print(f'[CKPTSaver][rank00] cost: {cost:.2f}s', flush=True)
+
+ del trainer_state
+ time.sleep(3), gc.collect(), torch.cuda.empty_cache(), time.sleep(3)
+ dist.barrier()
+
+def auto_resume(args: arg_util.Args, pattern='ckpt*.pth') -> Tuple[List[str], int, int, str, List[Tuple[float, float]], dict, dict]:
+ info = []
+ resume = ''
+ if args.auto_resume:
+ for dd in (args.local_out_path, args.bed):
+ all_ckpt = glob_with_global_step(os.path.join(dd, pattern))
+ if len(all_ckpt): break
+ if len(all_ckpt) == 0:
+ info.append(f'[auto_resume] no ckpt found @ {pattern}')
+ info.append(f'[auto_resume quit]')
+ else:
+ resume = all_ckpt[0]
+ info.append(f'[auto_resume] auto load from @ {resume} ...')
+ else:
+ info.append(f'[auto_resume] disabled')
+ info.append(f'[auto_resume quit]')
+
+ if len(resume) == 0:
+ return info, 0, 0, '[no acc str]', [], {}, {}
+
+ print(f'auto resume from {resume}')
+
+ try:
+ import os.path as osp
+ tgt_file = os.path.join(args.local_out_path, osp.basename(tgt_file))
+ os.makedirs(osp.dirname(tgt_file), exist_ok=True)
+ print(f'[load model] copy {resume} to {tgt_file}')
+ shutil.copyfile(resume, tgt_file)
+ ckpt = torch.load(tgt_file, map_location='cpu')
+ except Exception as e:
+ info.append(f'[auto_resume] failed, {e} @ {resume}')
+ if len(all_ckpt) < 2:
+ return info, 0, 0, '[no acc str]', [], {}, {}
+ try: # another chance to load from bytenas
+ ckpt = torch.load(all_ckpt[1], map_location='cpu')
+ except Exception as e:
+ info.append(f'[auto_resume] failed, {e} @ {all_ckpt[1]}')
+ return info, 0, 0, '[no acc str]', [], {}, {}
+
+ dist.barrier()
+ ep, it, g_it = ckpt['epoch'], ckpt['iter'], ckpt.get('g_it', 0)
+ eval_milestone = ckpt.get('milestones', [])
+ info.append(f'[auto_resume success] resume from ep{ep}, it{it}, eval_milestone: {eval_milestone}')
+ return info, ep, it, ckpt.get('acc_str', '[no acc str]'), eval_milestone, ckpt['trainer'], ckpt['args']
+
+def _is_trace_training_duration():
+ if (
+ metrics_client is not None
+ and os.environ.get('ENABLE_TRAINING_DURATION_METRICS_COLLECTION', 'false') == 'true'
+ and os.environ.get('MERLIN_JOB_ID', None) is not None
+ and os.environ.get("ARNOLD_ROBUST_TRAINING", None) == '1'
+ ):
+ # collect metrics from any of the following executor
+ executor_list = ["executor-0-0", "executor-0-1", "executor-0-2",
+ "executor-0", "executor-1", "executor-2"]
+ for ending in executor_list:
+ if os.environ.get("MY_POD_NAME", "").endswith(ending):
+ return True
+ return False
diff --git a/grn/utils_t2iv/sequence_parallel.py b/grn/utils_t2iv/sequence_parallel.py
new file mode 100644
index 0000000000000000000000000000000000000000..7f4222d43f81594eb1148f775ceb3d842dd7b55a
--- /dev/null
+++ b/grn/utils_t2iv/sequence_parallel.py
@@ -0,0 +1,96 @@
+import torch
+import torch.nn as nn
+import torch.distributed as dist
+from .comm.pg_utils import ProcessGroupManager
+from .comm.comm import set_sp_comm_group, split_sequence, gather_sequence, all_to_all_comm
+from .comm.operation import gather_forward_split_backward
+
+class SequenceParallelManager:
+ _SP_GROUP = None
+ _SP_SIZE = 0
+
+ @staticmethod
+ def sp_on():
+ return SequenceParallelManager._SP_GROUP is not None
+
+ @staticmethod
+ def init_sp(sp_size):
+ if SequenceParallelManager._SP_GROUP is not None:
+ print("WARN: sequence parallel group is already initialized")
+ return
+
+ if sp_size <= 1:
+ print(f"WARN: sequence parallel size must > 1 but got {sp_size}")
+ return
+
+ world_size = dist.get_world_size()
+ assert world_size % sp_size == 0, f"world_size {world_size} must be divisible by sp_size({sp_size})"
+ SequenceParallelManager._SP_SIZE = sp_size
+
+ pm = ProcessGroupManager(
+ world_size // sp_size,
+ sp_size,
+ dp_axis=0,
+ sp_axis=1,
+ )
+ pm_group = pm.sp_group
+ set_sp_comm_group(pm_group)
+ SequenceParallelManager._SP_GROUP = pm_group
+ return
+
+ @staticmethod
+ def get_sp_group():
+ return SequenceParallelManager._SP_GROUP
+
+ @staticmethod
+ def get_sp_size():
+ return SequenceParallelManager._SP_SIZE
+
+ @staticmethod
+ def get_sp_group_nums():
+ # if 2 sp_size, 8 ranks, group nums is 4
+ if SequenceParallelManager.sp_on():
+ world_size = torch.distributed.get_world_size()
+ return world_size // SequenceParallelManager._SP_SIZE
+ else:
+ return 0
+
+ @staticmethod
+ def get_sp_rank():
+ if SequenceParallelManager.sp_on():
+ global_rank = torch.distributed.get_rank()
+ sp_rank = global_rank % SequenceParallelManager._SP_SIZE
+ return sp_rank
+ else:
+ return 0
+
+ def get_sp_group_rank():
+ if SequenceParallelManager.sp_on():
+ global_rank = torch.distributed.get_rank()
+ sp_group_rank = global_rank // SequenceParallelManager._SP_SIZE
+ return sp_group_rank
+ else:
+ return 0
+
+def sp_split_sequence_by_dim(seq, seqlen_dim=1) -> torch.Tensor:
+ """
+ split the raw sequence by seqlen_dim
+ """
+ return split_sequence(seq, SequenceParallelManager.get_sp_group(), seqlen_dim, 'down')
+
+def sp_gather_sequence_by_dim(seq, seqlen_dim=1) -> torch.Tensor:
+ """
+ gather seqlen_dim to recover raw sequence
+ """
+ return gather_sequence(seq, SequenceParallelManager.get_sp_group(), seqlen_dim, 'up')
+
+def sp_all_to_all(ts, scatter_dim, gather_dim):
+ """
+ reorder the tensor's dimension, like [raw_seq_len/sp_size, hidden_dim] to [raw_seq_len, hidden_dim/sp_size]
+
+ scatter_dim: the dimension to split the tensor
+ gather_dim: the dimension to concatenate
+ """
+
+ return all_to_all_comm(ts, SequenceParallelManager.get_sp_group(), scatter_dim, gather_dim)
+
diff --git a/grn_pipeline.py b/grn_pipeline.py
new file mode 100644
index 0000000000000000000000000000000000000000..edf530bb850fe4f572b1ac588a609e342e6cc2dc
--- /dev/null
+++ b/grn_pipeline.py
@@ -0,0 +1,238 @@
+import os
+import json
+import numpy as np
+import torch
+from PIL import Image
+import tempfile
+
+from grn.utils_t2iv.infer import (
+ load_tokenizer,
+ load_transformer,
+ gen_one_example,
+ images2video
+)
+from grn.utils_t2iv.load import load_visual_tokenizer
+from grn.schedules.dynamic_resolution import get_dynamic_resolution_meta, get_first_full_spatial_size_scale_index
+from grn.schedules import get_encode_decode_func
+
+
+class GRNPipeline:
+ def __init__(self, model, vae, text_tokenizer, text_encoder, args, device='cuda'):
+ self.model = model
+ self.vae = vae
+ self.text_tokenizer = text_tokenizer
+ self.text_encoder = text_encoder
+ self.args = args
+ self.device = device
+ self.video_encode, self.video_decode, self.get_visual_rope_embeds, self.get_scale_pack_info = \
+ get_encode_decode_func(args.dynamic_scale_schedule)
+
+ @classmethod
+ def from_pretrained(
+ cls,
+ model_path='./weights/model.pth',
+ vae_path='./weights/hbq_tokenizer.ckpt',
+ text_encoder_ckpt='./weights/umt5-xxl',
+ device='cuda',
+ torch_dtype=torch.bfloat16,
+ hf_repo_id=None,
+ task='t2i',
+ pn='1M',
+ model='GRN2b',
+ use_slow_attn=False,
+ ):
+ # download weights from Hugging Face Hub
+ if hf_repo_id:
+ from huggingface_hub import hf_hub_download, snapshot_download
+ print(f"download weights from Hugging Face Hub: {hf_repo_id}")
+ if task == 'T2I':
+ model_path = hf_hub_download(repo_id=hf_repo_id, filename="GRN_T2I_2B_FSA_94600.pth")
+ elif task == 'T2V':
+ model_path = hf_hub_download(repo_id=hf_repo_id, filename="GRN_T2V_2B.pth")
+ else:
+ raise ValueError(f"Unknown task: {task}")
+ vae_path = hf_hub_download(repo_id=hf_repo_id, filename="HBQ_tokenizer_64dim_M4.ckpt")
+ snapshot_path = snapshot_download(repo_id=hf_repo_id, allow_patterns="umt5-xxl/**")
+ text_encoder_ckpt = os.path.join(snapshot_path, "umt5-xxl")
+ print(os.listdir(snapshot_path))
+
+ args = cls._get_default_args()
+ args.model = model
+ args.model_path = model_path
+ args.vae_path = vae_path
+ args.text_encoder_ckpt = text_encoder_ckpt
+ args.use_slow_attn = use_slow_attn
+ if isinstance(device, str):
+ device = torch.device(device)
+ args.other_device = device
+ args.task = task
+ args.pn = pn
+
+ # Derived parameters
+ args.max_duration = (args.video_frames - 1) / 4
+ args.num_of_label_value = args.num_lvl
+ args.semantic_num_lvl = args.num_lvl
+ args.detail_num_lvl = args.num_lvl
+ args.semantic_scale_dim = args.vae_latent_dim
+ args.detail_scale_dim = args.vae_latent_dim
+
+ # Load models
+ text_tokenizer, text_encoder = load_tokenizer(t5_path=args.text_encoder_ckpt, device=device)
+ vae = load_visual_tokenizer(args, device=device)
+ model = load_transformer(vae, args)
+
+ return cls(model, vae, text_tokenizer, text_encoder, args, device)
+
+ @staticmethod
+ def _get_default_args():
+ class Args:
+ def __init__(self):
+ self.video_frames = 81
+ self.model_path = './weights/model.pth'
+ self.vae_path = './weights/hbq_tokenizer.ckpt'
+ self.text_encoder_ckpt = './weights/umt5-xxl'
+ self.cfg = 1
+ self.fps = 16
+ self.cfg_insertion_layer = 0
+ self.vae_latent_dim = 64
+ self.hbq_round = 4
+ self.rope_type = '3d'
+ self.num_lvl = 2
+ self.rope2d_normalized_by_hw = 2
+ self.text_channels = 4096
+ self.apply_spatial_patchify = 0
+ self.h_div_w_template = 1.0
+ self.cache_dir = '/tmp'
+ self.checkpoint_type = 'torch'
+ self.seed = 42
+ self.bf16 = 0
+ self.dynamic_scale_schedule = 'GRN_vae_stride16'
+ self.train_h_div_w_list = '[]'
+ self.max_infer_steps = 50
+ self.min_infer_steps = 50
+ self.video_caption_type = 'tarsier2_caption'
+ self.temporal_compress_rate = 4
+ self.cached_video_frames = 81
+ self.duration_resolution = 0.25
+ self.video_fps = 16
+ self.simple_text_proj = 1
+ self.min_duration = -1
+ self.fsdp_save_flatten_model = 1
+ self.use_learnable_dim_proj = 0
+ self.use_fsq_cls_head = 1
+ self.use_feat_proj = 0
+ self.use_clipwise_caption = 0
+ self.use_ada_layer_norm = 0
+ self.cfg_type = 'cfg_interval_0.0'
+ self.add_scale_token = 1
+ self.vae_encoder_out_type = 'feature_tanh'
+ self.alpha = 1004
+ self.refine_mode = 'ar_discrete_GRN_bit'
+ self.add_class_token = 0
+ self.resample_rand_labels_per_step = 0
+ self.cfg_val = 3.0
+ self.scale_repetition = ''
+ self.gt_leak = -1
+ self.use_refined_prompt = None
+ self.use_prompt_engineering = 0
+ self.quality_prompt = ''
+ self.meta = ''
+ self.train_split_file = ''
+ self.n_sampes = 1
+ self.other_device = 'cuda' if torch.cuda.is_available() else 'cpu'
+ return Args()
+
+ def to(self, device):
+ if isinstance(device, str):
+ device = torch.device(device)
+ self.device = device
+ self.args.other_device = device
+ if self.model:
+ self.model = self.model.to(device)
+ if self.vae:
+ self.vae = self.vae.to(device)
+ return self
+
+ def __call__(
+ self,
+ prompt,
+ negative_prompt='',
+ guidance_scale=3.0,
+ temperature=1.0,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ first_frame_condition = False,
+ snr_shift = 1.,
+ h_div_w=1.,
+ duration=2.,
+ generator=None,
+ content_type='image',
+ seed=None,
+ **kwargs
+ ):
+ if seed is not None:
+ self.args.seed = seed
+ else:
+ self.args.seed = np.random.randint(0, 10000)
+
+ self.args.cfg_val = guidance_scale
+ self.args.tau = temperature
+ self.args.complexity_aware_Tmin = complexity_aware_Tmin
+ self.args.complexity_aware_Tmax = complexity_aware_Tmax
+ self.args.complexity_aware_k = complexity_aware_k
+ self.args.complexity_aware_b = complexity_aware_b
+ self.args.complexity_aware_wp = complexity_aware_wp
+ self.args.snr_shift = snr_shift
+
+ # Get dynamic resolution meta
+ dynamic_resolution_h_w, h_div_w_templates = get_dynamic_resolution_meta(
+ self.args.dynamic_scale_schedule,
+ self.args.train_h_div_w_list,
+ self.args.video_frames
+ )
+
+ # Get scale schedule based on aspect ratio
+ h_div_w_template_ = h_div_w_templates[np.argmin(np.abs(h_div_w_templates - h_div_w))]
+ self.args.mapped_h_div_w_template = h_div_w_template_
+
+ if content_type == "image":
+ num_frames = 1
+ duration = 0
+ else:
+ num_frames = min(self.args.video_frames, int(duration * self.args.video_fps + 1))
+
+ scale_schedule = dynamic_resolution_h_w[h_div_w_template_][self.args.pn]['pt2scale_schedule'][(num_frames - 1) // 4 + 1]
+
+ self.args.first_full_spatial_size_scale_index = get_first_full_spatial_size_scale_index(scale_schedule)
+ self.args.tower_split_index = self.args.first_full_spatial_size_scale_index + 1
+
+ # Generate content
+ generated_image = gen_one_example(
+ self.model, self.vae, self.text_tokenizer, self.text_encoder, [prompt],
+ negative_prompt=negative_prompt, g_seed=seed, gt_leak=self.args.gt_leak, gt_ls_Bl=None,
+ cfg_list=self.args.cfg_val, tau_list=self.args.tau, scale_schedule=scale_schedule,
+ cfg_insertion_layer=[self.args.cfg_insertion_layer], vae_latent_dim=self.args.vae_latent_dim,
+ args=self.args, get_visual_rope_embeds=self.get_visual_rope_embeds,
+ noise_list=None, first_frame_condition=first_frame_condition,
+ )
+
+ if len(generated_image.shape) == 3:
+ generated_image = generated_image.unsqueeze(0)
+
+ generated_image = generated_image.cpu().numpy()
+
+ ext = '.jpg' if num_frames == 1 else '.mp4'
+
+ with tempfile.NamedTemporaryFile(suffix=ext, delete=False) as tmp:
+ output_path = tmp.name
+
+ images2video(generated_image, fps=self.args.fps, save_filepath=output_path)
+
+ if ext == '.jpg':
+ img = Image.open(output_path)
+ return type('Result', (object,), {'images': [img]})
+ else:
+ return type('Result', (object,), {'videos': [output_path]})
\ No newline at end of file
diff --git a/t2i_infer.py b/t2i_infer.py
new file mode 100644
index 0000000000000000000000000000000000000000..5c5e4bfb982528cb6ae7625a284d7ebaa4f277ee
--- /dev/null
+++ b/t2i_infer.py
@@ -0,0 +1,30 @@
+from PIL import Image
+from grn_pipeline import GRNPipeline
+
+# Load pipeline
+pipeline = GRNPipeline.from_pretrained(
+ hf_repo_id='bytedance-research/GRN',
+ task='T2I',
+ pn='1M',
+ model='GRN2b',
+ use_slow_attn=True,
+ device='cpu',
+).to('cuda')
+
+# Generate one image
+result = pipeline(
+ prompt="" + "A cute cat playing in the garden",
+ guidance_scale=3.0,
+ temperature=1.1,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ snr_shift = 1.,
+ h_div_w=1.,
+ content_type='image',
+ seed=42,
+)
+image = result.images[0]
+image.save('./generated_image.jpg')
diff --git a/t2iv_train.py b/t2iv_train.py
new file mode 100644
index 0000000000000000000000000000000000000000..c8df0ef3ba0e1c31b3166517c22a43ba9f055e49
--- /dev/null
+++ b/t2iv_train.py
@@ -0,0 +1,410 @@
+import gc
+import json
+import math
+import os
+import os.path as osp
+import random
+import sys
+import time
+from functools import partial
+from typing import List, Optional, Tuple
+os.environ["TOKENIZERS_PARALLELISM"] = "false"
+os.environ['XFORMERS_FORCE_DISABLE_TRITON'] = '1'
+
+import numpy as np
+import torch
+torch._dynamo.config.cache_size_limit = 64
+from torch.nn import functional as F
+from torch.profiler import record_function
+from torch.utils.data import DataLoader
+from transformers import T5EncoderModel, T5TokenizerFast
+import torch.distributed as tdist
+
+import grn.utils_t2iv.dist as dist
+from grn.dataset.build import build_joint_dataset
+from grn.models.ema import get_ema_model
+from grn.utils_t2iv import arg_util, misc
+from grn.utils import wandb_utils
+from grn.trainer import get_trainer
+
+def build_everything_from_args(args: arg_util.Args, saver):
+ args.set_initial_seed(benchmark=True)
+ print(f'Loading T5 from {args.t5_path}...')
+ from grn.models.umt5.t5 import T5EncoderModel
+ text_encoder = T5EncoderModel(
+ text_len=args.tlen, # 512
+ dtype=torch.bfloat16, # torch.bfloat16
+ device=args.device,
+ checkpoint_path=osp.join(args.t5_path, 'models_t5_umt5-xxl-enc-bf16.pth'),
+ tokenizer_path=osp.join(args.t5_path, 'umt5-xxl'),
+ enable_fsdp=True) # False
+ # text_encoder.model.to(args.device)
+ text_tokenizer = text_encoder.tokenizer
+ args.text_tokenizer_type = 'umt5'
+ args.text_tokenizer = text_tokenizer
+
+ # build models. Note that here gpt is the causal VAR transformer which performs next scale prediciton with text guidance
+ vae_local, gpt_uncompiled, gpt_wo_ddp, gpt_ddp, gpt_wo_ddp_ema, gpt_ddp_ema, gpt_optim = build_model_optimizer(args)
+
+ Trainer = get_trainer(args)
+ # build trainer
+ trainer = Trainer(
+ is_visualizer=dist.is_visualizer(), device=args.device,
+ vae_local=vae_local, gpt_wo_ddp=gpt_wo_ddp, gpt=gpt_ddp,
+ zero=args.zero, vae_latent_dim=args.vae_latent_dim, gpt_opt=gpt_optim,
+ reweight_loss_by_scale=args.reweight_loss_by_scale, gpt_wo_ddp_ema=gpt_wo_ddp_ema,
+ gpt_ema=gpt_ddp_ema, use_fsdp_model_ema=args.use_fsdp_model_ema, other_args=args,
+ )
+
+ # auto resume from broken experiment
+ global_it = 0
+ if args.checkpoint_type == 'torch':
+ from grn.utils_t2iv.save_and_load import auto_resume
+ auto_resume_info, start_ep, start_it, acc_str, eval_milestone, trainer_state, args_state = auto_resume(args, 'ar-ckpt*.pth')
+ print(f'initial args:\n{str(args)}')
+ if start_ep == args.ep:
+ print(f'[vgpt] AR finished ({acc_str}), skipping ...\n\n')
+ return None
+ if trainer_state is not None and len(trainer_state):
+ trainer.load_state_dict(trainer_state, strict=False, skip_vae=True) # don't load vae again
+
+ del vae_local, gpt_uncompiled, gpt_wo_ddp, gpt_ddp, gpt_wo_ddp_ema, gpt_ddp_ema, gpt_optim
+ dist.barrier()
+ return text_tokenizer, text_encoder, trainer, global_it
+
+
+def build_model_optimizer(args):
+ from torch.nn.parallel import DistributedDataParallel as DDP
+ from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+ from grn.models.grn import MultipleLayers
+ from grn.models.init_param import init_weights
+ from grn.utils_t2iv.lr_control import filter_params
+ from grn.utils_t2iv.load import build_vae_gpt
+
+ # disable builtin initialization for speed
+ setattr(torch.nn.Linear, 'reset_parameters', lambda self: None)
+ setattr(torch.nn.LayerNorm, 'reset_parameters', lambda self: None)
+ vae_local, gpt_wo_ddp = build_vae_gpt(args, device=args.model_init_device)
+ count_p = lambda m: sum(p.numel() for p in m.parameters()) / 1e6
+ num_para = count_p(gpt_wo_ddp)
+ if num_para/1000 < 20: # < 20B
+ gpt_wo_ddp = gpt_wo_ddp.to('cuda')
+
+ init_weights(gpt_wo_ddp)
+ gpt_wo_ddp.special_init()
+ if args.use_fsdp_model_ema:
+ gpt_wo_ddp_ema = get_ema_model(gpt_wo_ddp)
+ else:
+ gpt_wo_ddp_ema = None
+
+ if args.rush_resume:
+ print(f"{args.rush_resume=}")
+ if '.pth' in args.rush_resume:
+ cpu_d = torch.load(args.rush_resume, 'cpu')
+ else:
+ from grn.utils_t2iv.save_and_load import merge_ckpt
+ cpu_d = merge_ckpt(args.rush_resume, osp.join(args.rush_resume, 'ouput'), save=False, use_ema_model=False, fsdp_save_flatten_model=args.fsdp_save_flatten_model)
+ if 'trainer' in cpu_d:
+ state_dict = cpu_d['trainer']['gpt_fsdp']
+ ema_state_dict = cpu_d['trainer'].get('gpt_ema_fsdp', state_dict)
+ else:
+ state_dict = cpu_d
+ ema_state_dict = state_dict
+ def drop_unfit_weights(state_dict):
+ try:
+ if 'word_embed.weight' in state_dict and (state_dict['word_embed.weight'].shape[1] != gpt_wo_ddp.word_embed.in_features):
+ print(f'[rush_resume] drop word_embed.weight')
+ del state_dict['word_embed.weight']
+ if 'head.proj.weight' in state_dict and (state_dict['head.proj.weight'].shape[0] != gpt_wo_ddp.head.proj.out_features):
+ print(f'[rush_resume] drop head')
+ del state_dict['head.proj.weight']
+ del state_dict['head.proj.bias']
+ except Exception as e:
+ print(e)
+ for key in ['word_embed.weight', 'head.proj.weight', 'head.proj.bias']:
+ if key in state_dict:
+ del state_dict[key]
+ print(f'[rush_resume] drop {key}')
+
+ if 'text_proj_for_sos.ca.mat_kv.weight' in state_dict and \
+ (state_dict['text_proj_for_sos.ca.mat_kv.weight'].shape != gpt_wo_ddp.text_proj_for_sos.ca.mat_kv.weight.shape):
+ print(f'[rush_resume] drop cfg_uncond')
+ del state_dict['cfg_uncond']
+ for key in list(state_dict.keys()):
+ if 'text' in key:
+ del state_dict[key]
+ return state_dict
+ print(gpt_wo_ddp.load_state_dict(drop_unfit_weights(state_dict), strict=False))
+ if args.use_fsdp_model_ema:
+ gpt_wo_ddp_ema.load_state_dict(drop_unfit_weights(ema_state_dict), strict=False)
+
+ ndim_dict = {name: para.ndim for name, para in gpt_wo_ddp.named_parameters() if para.requires_grad}
+
+ print(f'[PT] GPT model = {gpt_wo_ddp}\n\n')
+ print(f'[PT] GPT model details:')
+ for name, param in gpt_wo_ddp.named_parameters():
+ print(f"Name: {name}, Shape: {param.shape}")
+ print(f'[PT][#para], GPT={num_para:.2f}M parameters\n\n')
+
+ gpt_uncompiled = gpt_wo_ddp
+ gpt_wo_ddp = args.compile_model(gpt_wo_ddp, args.tfast)
+
+ gpt_ddp_ema = None
+ if args.zero:
+ from torch.distributed.fsdp import ShardingStrategy
+ from torch.distributed.fsdp.wrap import ModuleWrapPolicy
+ from torch.distributed.device_mesh import init_device_mesh
+ # use mix prec: https://github.com/pytorch/pytorch/issues/76607
+ if args.fsdp_warp_mode == 'full':
+ print(f'warp all modules for fsdp')
+ def my_policy(
+ module: torch.nn.Module,
+ recurse: bool,
+ **kwargs,
+ ) -> bool:
+ return True
+ auto_wrap_policy = my_policy
+ else:
+ print(f'warp transformer blocks for fsdp')
+ auto_wrap_policy = ModuleWrapPolicy([MultipleLayers, ])
+
+ if args.enable_hybrid_shard == 1:
+ sharding_strategy = ShardingStrategy.HYBRID_SHARD if args.zero == 3 else ShardingStrategy._HYBRID_SHARD_ZERO2
+ world_size = dist.get_world_size()
+ assert world_size % args.inner_shard_degree == 0
+ assert args.inner_shard_degree > 1 and args.inner_shard_degree <= world_size
+ device_mesh = init_device_mesh('cuda', (world_size // args.inner_shard_degree, args.inner_shard_degree))
+ elif args.enable_hybrid_shard == -1: # no shard
+ sharding_strategy = ShardingStrategy.NO_SHARD
+ device_mesh = None
+ else:
+ sharding_strategy = ShardingStrategy.FULL_SHARD if args.zero == 3 else ShardingStrategy.SHARD_GRAD_OP
+ device_mesh = None
+ print(f'{">" * 45 + " " * 5} FSDP INIT with {args.zero=} {sharding_strategy=} {auto_wrap_policy=} {" " * 5 + "<" * 45}', flush=True)
+
+ if args.fsdp_init_device == 'cpu':
+ gpt_wo_ddp = gpt_wo_ddp.cpu()
+
+ gpt_ddp: FSDP = FSDP(
+ gpt_wo_ddp,
+ device_id=dist.get_local_rank(),
+ sharding_strategy=sharding_strategy,
+ mixed_precision=None,
+ auto_wrap_policy=auto_wrap_policy,
+ use_orig_params=True,
+ sync_module_states=True,
+ limit_all_gathers=True,
+ device_mesh=device_mesh,
+ ).to(args.device)
+
+ if args.use_fsdp_model_ema:
+ gpt_wo_ddp_ema = gpt_wo_ddp_ema.to(args.device)
+ gpt_ddp_ema: FSDP = FSDP(
+ gpt_wo_ddp_ema,
+ device_id=dist.get_local_rank(),
+ sharding_strategy=sharding_strategy,
+ mixed_precision=None,
+ auto_wrap_policy=auto_wrap_policy,
+ use_orig_params=True,
+ sync_module_states=True,
+ limit_all_gathers=True,
+ device_mesh=device_mesh,
+ )
+ else:
+ ddp_class = DDP if dist.initialized() else misc.NullDDP
+ gpt_ddp: DDP = ddp_class(gpt_wo_ddp, device_ids=[dist.get_local_rank()], find_unused_parameters=args.dbg, broadcast_buffers=False)
+ torch.cuda.synchronize()
+
+ # =============== build optimizer ===============
+ nowd_keys = set()
+ nowd_keys |= {
+ 'cls_token', 'start_token', 'task_token', 'cfg_uncond',
+ 'pos_embed', 'pos_1LC', 'pos_start', 'start_pos', 'lvl_embed',
+ 'gamma', 'beta',
+ 'ada_gss', 'moe_bias',
+ 'scale_mul',
+ 'text_proj_for_sos.ca.mat_q',
+ 'scale_tokens', 'class_tokens'
+ }
+ names, paras, para_groups = filter_params(gpt_ddp if args.zero else gpt_wo_ddp, ndim_dict, nowd_keys=nowd_keys)
+ del ndim_dict
+ opt_clz = partial(torch.optim.AdamW, betas=(0.9, 0.999), fused=True)
+ opt_kw = dict(lr=args.tlr, weight_decay=args.twd)
+ print(f'[vgpt] optim={opt_clz}, opt_kw={opt_kw}\n')
+ gpt_optim = opt_clz(params=para_groups, **opt_kw)
+ del names, paras, para_groups
+ return vae_local, gpt_uncompiled, gpt_wo_ddp, gpt_ddp, gpt_wo_ddp_ema, gpt_ddp_ema, gpt_optim
+
+def build_dataset(args):
+ train_dataset = build_joint_dataset(
+ args,
+ args.meta_folders,
+ args.meta_folder_repeats,
+ max_caption_len=args.tlen,
+ short_prob=args.short_cap_prob,
+ load_vae_instead_of_image=False
+ )
+ return train_dataset
+
+def main_train(args: arg_util.Args):
+ if args.checkpoint_type == 'torch':
+ from grn.utils_t2iv.save_and_load import CKPTSaver, auto_resume
+ saver = CKPTSaver(dist.is_master(), eval_milestone=None)
+ else:
+ raise ValueError(f'{args.checkpoint_type=}')
+ text_tokenizer, text_encoder, trainer, start_global_it = build_everything_from_args(args, saver)
+ gc.collect(), torch.cuda.empty_cache()
+ logging_params_milestone: List[int] = np.linspace(1, args.ep, 10+1, dtype=int).tolist()
+ time.sleep(3), gc.collect(), torch.cuda.empty_cache(), time.sleep(3)
+
+ # ============================================= epoch loop begins =============================================
+ # build wandb logger
+ if dist.is_master():
+ wandb_utils.wandb.init(project=args.project_name, name=args.exp_name, config={})
+
+ dataloader_generator = torch.Generator()
+ dataloader_generator.manual_seed(args.seed + tdist.get_rank())
+ for ep in range(args.ep):
+ # build data at each epoch to ensure read meta take effects for each dataloader worker
+ args.epoch = ep
+ train_dataset = build_dataset(args)
+ iters_train = len(train_dataset)
+ print(f'[PT info] from {start_global_it=} {iters_train=}=======> bed: {args.bed} <=======\n')
+
+ # build dataloader
+ train_dataloader = DataLoader(dataset=train_dataset, num_workers=args.workers, pin_memory=True, batch_size=None, shuffle=True, generator=dataloader_generator)
+ train_dataloader_iter_obj = iter(train_dataloader)
+
+ # [train one epoch]
+ stats, (sec, remain_time, finish_time) = train_one_ep(
+ ep=ep,
+ start_global_it=start_global_it,
+ me=None,
+ saver=saver,
+ args=args,
+ ld_or_itrt=train_dataloader_iter_obj,
+ iters_train=iters_train,
+ text_tokenizer=text_tokenizer, text_encoder=text_encoder,
+ trainer=trainer,
+ logging_params_milestone=logging_params_milestone,
+ )
+ start_global_it += iters_train
+ del stats, train_dataset, train_dataloader
+ time.sleep(10), gc.collect(), time.sleep(10) # torch.cuda.empty_cache()
+ return
+
+
+def train_one_ep(
+ ep: int, start_global_it: int, me: misc.MetricLogger,
+ saver, args: arg_util.Args, ld_or_itrt, iters_train: int,
+ text_tokenizer: T5TokenizerFast, text_encoder: T5EncoderModel, trainer, logging_params_milestone,
+):
+ # IMPORTANT: import heavy packages after the Dataloader object creation/iteration to avoid OOM
+ step_cnt = 0
+ header = f'[Ep]: [{ep:4d}/{args.ep}]'
+
+ g_it, max_it = start_global_it, args.ep * iters_train
+
+ me = misc.MetricLogger()
+ [me.add_meter(x, misc.SmoothedValue(window_size=1, fmt='{value:.2g}')) for x in ['tlr']]
+ [me.add_meter(x, misc.SmoothedValue(window_size=1, fmt='{median:.2f} ({global_avg:.2f})')) for x in ['tnm']]
+ [me.add_meter(x, misc.SmoothedValue(window_size=1, fmt='{median:.3f} ({global_avg:.3f})')) for x in ['L', 'L_i', 'L_v']]
+ [me.add_meter(x, misc.SmoothedValue(window_size=1, fmt='{median:.2f} ({global_avg:.2f})')) for x in ['Acc', 'Acc_i', 'Acc_v']]
+ [me.add_meter(x, misc.SmoothedValue(window_size=1, fmt='{median:.2f} ({global_avg:.2f})')) for x in ['seq_usage']]
+ # ============================================= iteration loop begins =============================================
+ start_it = 0
+ for it, data in me.log_every(start_it, iters_train, ld_or_itrt, args.log_freq, args.log_every_iter, header):
+ # for dpo training, we will save the first iter model for comparison
+ g_it += 1
+ if (g_it > 0 and g_it % args.save_model_iters_freq == 0) or (args.save_start_model and g_it == 1):
+ if args.checkpoint_type == 'torch':
+ saver.sav(args=args, g_it=g_it, next_ep=ep, next_it=it+1, trainer=trainer, acc_str=f'[todo]', eval_milestone=None, also_save_to=None, best_save_to=None)
+
+ # [get data]
+ images, captions, raw_features_bcthw, feature_cache_files4images, media, meta_list = data['images'], data['captions'], data['raw_features_bcthw'], data['feature_cache_files4images'], data['media'], data['meta_list']
+
+ # # [prepare text features]
+ if args.add_class_token > 0: # c2i task
+ text_cond_tuple = [[] for _ in range(5)]
+ else:
+ caption_nums = [len(item) for item in captions]
+ flatten_captions = []
+ for item in captions:
+ flatten_captions.extend(item)
+ if args.text_tokenizer_type == 'flan_t5':
+ tokens = text_tokenizer(text=flatten_captions, max_length=text_tokenizer.model_max_length, padding='max_length', truncation=True, return_tensors='pt') # todo: put this into dataset
+ input_ids = tokens.input_ids.cuda(non_blocking=True)
+ mask = tokens.attention_mask.cuda(non_blocking=True)
+ text_features = text_encoder(input_ids=input_ids, attention_mask=mask)['last_hidden_state'].float()
+ lens: List[int] = mask.sum(dim=-1).tolist()
+ cu_seqlens_k = F.pad(mask.sum(dim=-1).to(dtype=torch.int32).cumsum_(0), (1, 0))
+ Ltext = max(lens)
+ kv_compact = []
+ for text_ind, (len_i, feat_i) in enumerate(zip(lens, text_features.unbind(0))):
+ kv_compact.append(feat_i[:len_i])
+ kv_compact = torch.cat(kv_compact, dim=0)
+ text_cond_tuple: Tuple[torch.FloatTensor, List[int], torch.LongTensor, int] = (kv_compact, lens, cu_seqlens_k, Ltext, caption_nums)
+ else:
+ text_features = text_encoder(flatten_captions, args.device)
+ lens = [len(item) for item in text_features]
+ cu_seqlens_k = [0]
+ for len_i in lens:
+ cu_seqlens_k.append(cu_seqlens_k[-1] + len_i)
+ cu_seqlens_k = torch.tensor(cu_seqlens_k, dtype=torch.int32)
+ Ltext = max(lens)
+ kv_compact = torch.cat(text_features, dim=0).float()
+ text_cond_tuple = (kv_compact, lens, cu_seqlens_k, Ltext, caption_nums)
+
+ if len(images):
+ images = [item.to(args.device, non_blocking=True) for item in images]
+ if len(raw_features_bcthw):
+ raw_features_bcthw = [item.to(args.device, non_blocking=True) for item in raw_features_bcthw]
+
+ # [schedule learning rate and weight decay]
+ if ep == 0 and (g_it-start_global_it) < args.wp_it:
+ cur_lr_ratio = args.wp0 + (1-args.wp0) * (g_it-start_global_it) / args.wp_it
+ else:
+ cur_lr_ratio = 1
+ cur_lr = args.tlr * cur_lr_ratio
+ if cur_lr_ratio < 1:
+ cur_wd = args.twd
+ for param_group in trainer.gpt_opt.param_groups:
+ param_group['lr'] = cur_lr * param_group.get('lr_sc', 1) # 'lr_sc' could be assigned
+ param_group['weight_decay'] = cur_wd * param_group.get('wd_sc', 1)
+
+ # [get scheduled hyperparameters]
+ stepping = (g_it + 1) % args.gradient_accumulation == 0
+ step_cnt += int(stepping)
+
+ trainer.train_step(
+ ep=ep, it=it, g_it=g_it, stepping=stepping, clip_decay_ratio=1,
+ metric_lg=me,
+ logging_params=stepping and step_cnt == 1 and (ep < 4 or ep in logging_params_milestone),
+ inp_B3HW=images,
+ raw_features_bcthw=raw_features_bcthw,
+ feature_cache_files4images=feature_cache_files4images,
+ text_cond_tuple=text_cond_tuple,
+ media=media,
+ meta_list=meta_list,
+ args=args,
+ )
+
+ me.update(tlr=cur_lr)
+ # ============================================= iteration loop ends =============================================
+
+ me.synchronize_between_processes()
+ return {k: meter.global_avg for k, meter in me.meters.items()}, me.iter_time.time_preds(max_it - (g_it + 1) + (args.ep - ep) * 15) # +15: other cost
+
+
+def main():
+ args: arg_util.Args = arg_util.init_dist_and_get_args()
+ main_train(args)
+ print(f'final args:\n\n{str(args)}')
+ if isinstance(sys.stdout, dist.BackupStreamToFile) and isinstance(sys.stderr, dist.BackupStreamToFile):
+ sys.stdout.close(), sys.stderr.close()
+ dist.barrier()
+ time.sleep(120)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/t2v_infer.py b/t2v_infer.py
new file mode 100644
index 0000000000000000000000000000000000000000..21e30980c29944a9c32322f639945075703ae214
--- /dev/null
+++ b/t2v_infer.py
@@ -0,0 +1,51 @@
+from grn_pipeline import GRNPipeline
+
+negative_prompt = (
+ # --- quality ---
+ "ugly, blurry, low-resolution, low-detail, low-quality, noisy, grainy, "
+ "overexposed, underexposed, oversaturated, undersaturated, soft focus, "
+ "artifacts, compression artifacts, jpeg artifacts, flickering, "
+ # --- style ---
+ "painting, oil painting, illustration, drawing, sketch, cartoon, anime, manga, "
+ "3d, cgi, render, digital art, "
+ "plastic, waxy, glossy, fake, unnatural, "
+ # --- skin, figure ---
+ "plastic skin, waxy skin, over-smoothed skin, doll-like, "
+ "deformed, mutated, disfigured, bad anatomy, bad hands, extra fingers, missing fingers, extra limbs, "
+ # --- motion ---
+ "still image, static, motionless, frozen, "
+ "unnatural motion, reversed motion, stuttering, choppy, "
+ # --- misc ---
+ "text, watermark, logo, signature, username, "
+ "crowded background, bad composition"
+)
+
+# Load pipeline
+pipeline = GRNPipeline.from_pretrained(
+ hf_repo_id='bytedance-research/GRN',
+ task='T2V',
+ pn='0.41M',
+ device='cpu'
+).to('cuda')
+
+prompt="The man, of medium build with short, dark, curly hair, stands centered in the frame, wearing a simple white t-shirt that contrasts with the greenery behind him. He holds a dark smartphone, likely a modern model with a triple-lens camera setup, in his right hand, angled slightly toward his body. His gaze is fixed on the screen, and his facial expression shifts subtly\u2014smiling, nodding, and occasionally pursing his lips\u2014as if reacting to content on the phone. The background features a mix of tall green trees and shrubs, with a light blue metal fence running horizontally across the mid-ground, suggesting a garden or rural boundary. The overcast sky diffuses the light, creating soft shadows and a calm, neutral atmosphere. The man\u2019s slight head movements and micro-expressions indicate engagement, possibly reading or responding to a message or video. The composition places him as the focal point, with the natural, slightly blurred background reinforcing his isolation in the moment. The relative stillness of the scene, apart from his subtle gestures, suggests a private, introspective interaction with technology in a serene outdoor setting"
+
+# Generate one video
+result = pipeline(
+ prompt=f"{prompt}. masterpiece, high quality.",
+ negative_prompt=negative_prompt,
+ guidance_scale=4.0,
+ temperature=1.0,
+ complexity_aware_Tmin=10,
+ complexity_aware_Tmax=50,
+ complexity_aware_k = 0,
+ complexity_aware_b = 50,
+ complexity_aware_wp = 5,
+ snr_shift = 1.,
+ h_div_w=9/16,
+ duration=2.,
+ first_frame_condition=False,
+ content_type='video',
+ seed=42,
+)
+video_file = result.videos[0]
diff --git a/tools/api_key.py b/tools/api_key.py
new file mode 100644
index 0000000000000000000000000000000000000000..8285305e51306727bd6bc50cab6b93ba6ee1b323
--- /dev/null
+++ b/tools/api_key.py
@@ -0,0 +1,3 @@
+HF_TOKEN = '[YOUR HF_TOKEN]'
+HF_HOME = '[YOUR HF_HOME]'
+GPT_AK = '[YOUR GPT_AK]'
diff --git a/tools/read_metrics.py b/tools/read_metrics.py
new file mode 100644
index 0000000000000000000000000000000000000000..a297de06c1cf10dc3dfa60927163badef2a99336
--- /dev/null
+++ b/tools/read_metrics.py
@@ -0,0 +1,20 @@
+import os
+import os.path as osp
+import json
+import glob
+import sys
+
+import numpy as np
+
+data_root = sys.argv[1]
+min_fid, best_dict, best_cfg = 999, None, None
+for metric_file in sorted(glob.glob(osp.join(data_root, '*/metrics.json'))):
+ with open(metric_file, 'r') as f:
+ metrics = json.load(f)
+ cur_fid = metrics['fid'] if 'fid' in metrics else metrics['frechet_inception_distance']
+ if cur_fid < min_fid:
+ min_fid = cur_fid
+ best_dict = metrics
+ best_cfg = [metric_file, metrics]
+ print(metric_file, metrics)
+print(f'\n {min_fid=} {best_dict=} {best_cfg=}')
diff --git a/tools/split_jsonl.py b/tools/split_jsonl.py
new file mode 100644
index 0000000000000000000000000000000000000000..d54475c77a663a57ba4428801f0cfb7210ffee81
--- /dev/null
+++ b/tools/split_jsonl.py
@@ -0,0 +1,121 @@
+import os
+import os.path as osp
+import time
+import itertools
+import shutil
+import glob
+import argparse
+import json
+from concurrent.futures import ThreadPoolExecutor, as_completed
+
+import tqdm
+import numpy as np
+
+def save_lines(lines, filename):
+ os.makedirs(osp.dirname(filename), exist_ok=True)
+ with open(filename, 'w') as f:
+ f.writelines(lines)
+ del lines
+
+def get_part_jsonls(save_dir, total_line_number, ext='.jsonl', chunk_size=1000, bucket_size=1000):
+ if osp.exists(save_dir):
+ shutil.rmtree(save_dir)
+ chunk_id2save_files = {}
+ missing = False
+ parts = int(np.ceil(total_line_number / chunk_size))
+ for chunk_id in range(1, parts+1):
+ if chunk_id == parts:
+ num_of_lines = total_line_number - chunk_size * (parts-1)
+ else:
+ num_of_lines = chunk_size
+ bucket = (chunk_id-1) // args.bucket_size + 1
+ chunk_id2save_files[chunk_id] = osp.join(save_dir, f'{bucket:06d}', f'{chunk_id:04d}_{parts:04d}_{num_of_lines:09d}{ext}')
+ if not osp.exists(chunk_id2save_files[chunk_id]):
+ missing = True
+ return missing, chunk_id2save_files
+
+def split_large_txt_files(all_lines, chunk_id2save_files):
+ chunk_id = 1
+ total = len(all_lines)
+ pbar = tqdm.tqdm(total=len(chunk_id2save_files))
+ chunk = []
+ futures = []
+
+ max_workers = 128
+ with ThreadPoolExecutor(max_workers=max_workers) as executor:
+ for line in all_lines:
+ chunk.append(line)
+ cur_chunk_size = int(osp.splitext(osp.basename(chunk_id2save_files[chunk_id]))[0].split('_')[-1])
+ if len(chunk) >= cur_chunk_size:
+ futures.append(
+ executor.submit(save_lines, chunk, chunk_id2save_files[chunk_id])
+ )
+ pbar.update(1)
+ chunk = []
+ chunk_id += 1
+
+ if len(chunk):
+ raise ValueError("last chunk not save, means misalign data!")
+
+ for future in as_completed(futures):
+ future.result()
+
+ pbar.close()
+
+
+from multiprocessing import Manager
+lock = Manager().Lock()
+def read_jsonl(jsonl_file):
+ with open(jsonl_file, 'r') as f:
+ lines = f.readlines()
+ global pbar
+ with lock:
+ pbar.update(1)
+ return lines
+
+def read_jsonls(jsonl_files, worker):
+ global pbar
+ from multiprocessing.pool import ThreadPool
+ pbar = tqdm.tqdm(total=len(jsonl_files))
+ print(f'[Data Loading] Reading {len(jsonl_files)} meta files...')
+ all_lines = []
+ if len(jsonl_files) == 1:
+ lines_num = int(osp.splitext(jsonl_files[0])[0].split('_')[-1])
+ pbar = tqdm.tqdm(total=lines_num)
+ with open(jsonl_files[0], 'r') as f:
+ for line in f:
+ pbar.update(1)
+ all_lines.append(line)
+ else:
+ with ThreadPool(worker) as pool:
+ for img_metas in pool.starmap(read_jsonl, [(bin_file,) for bin_file in jsonl_files]):
+ all_lines.extend(img_metas)
+ np.random.shuffle(all_lines)
+ return all_lines
+
+if __name__ == '__main__':
+ parser = argparse.ArgumentParser()
+ parser.add_argument('--jsonl_folder_list', type=str, default='', nargs='+', help='patha pathb pathc')
+ parser.add_argument('--save_dir', type=str, default='')
+ parser.add_argument('--chunk_size', type=int, default=100)
+ parser.add_argument('--bucket_size', type=int, default=10000)
+ parser.add_argument('--worker', type=int, default=128)
+ args = parser.parse_args()
+
+ global pbar
+ t1 = time.time()
+ jsonl_files = []
+ for item in args.jsonl_folder_list:
+ jsonl_files += glob.glob(osp.join(item, '*.jsonl'))
+ # jsonl_files += glob.glob(osp.join(item, '*/*.jsonl' ))
+ np.random.shuffle(jsonl_files)
+
+ pbar = tqdm.tqdm(total=len(jsonl_files))
+ lines = read_jsonls(jsonl_files, args.worker)
+ print(f'total {len(lines)} lines')
+ line_num = len(lines)
+ missing, chunk_id2save_files = get_part_jsonls(args.save_dir, line_num, chunk_size=args.chunk_size, bucket_size=args.bucket_size)
+
+ split_large_txt_files(lines, chunk_id2save_files)
+ t2 = time.time()
+ print(f'split takes {t2-t1}s')