Compare commits

..
Author SHA1 Message Date
Your Name 1856f01c06 update 2024-12-12 20:59:03 -08:00
Your Name 55442027f8 update 2024-12-12 20:50:16 -08:00
Your Name 710b9c549b update 2024-12-12 18:33:49 -08:00
rlsu9 40120579f2 update 2024-12-12 22:28:41 +00:00
rlsu9 f8409886a1 update 2024-12-12 21:35:29 +00:00
rlsu9 f37e1fa18e update 2024-12-12 21:32:43 +00:00
rlsu9 4917d58d8e update 2024-12-12 21:30:45 +00:00
rlsu9 bbeb227bc5 add scripts 2024-12-12 21:26:36 +00:00
rlsu9 163f9227ff update 2024-12-12 21:26:03 +00:00
rlsu9 98070d9bca add val gen 2024-12-12 19:54:22 +00:00
rlsu9 dfb15828a8 rename preprocess 2024-12-12 19:53:57 +00:00
rlsu9 ef667b9942 move data preprocess 2024-12-12 19:12:58 +00:00
rlsu9 01282b7622 debug ac; disitll; move fsdp_util 2024-12-12 19:07:43 +00:00
rlsu9 b14cc034e4 update 2024-12-12 18:11:04 +00:00
Peiyuan Zhang 6e5cb4b8e5 update 2024-12-12 06:34:22 +00:00
Peiyuan Zhang 87913376be add training 2024-12-12 05:59:11 +00:00
Peiyuan Zhang ce22b1279d vae done 2024-12-12 00:41:02 +00:00
Peiyuan Zhang f1dba33df7 update 2024-12-11 23:56:13 +00:00
Peiyuan Zhang e8ce4290ff update 2024-12-11 20:27:18 +00:00
Peiyuan Zhang fb1cc737e3 the same input output format for transformer 2024-12-11 19:56:57 +00:00
Peiyuan Zhang f5074a5c31 refactor single card attention to no pad 2024-12-11 19:00:10 +00:00
Peiyuan Zhang 0d4ae825c9 add no pad; support batch size > 1 2024-12-11 18:39:30 +00:00
Peiyuan Zhang da66e0fcea move rope freq inside transformer 2024-12-11 18:18:44 +00:00
foreverpiano 97f9db9b71 update 2024-12-11 13:36:38 +00:00
foreverpiano e297fe3949 make distill a framework 2024-12-11 13:33:24 +00:00
foreverpiano dd2cf8ad69 update 2024-12-11 13:04:48 +00:00
foreverpiano f5f5903aa4 update sp 2024-12-10 13:59:33 +00:00
foreverpiano 20b03a45d7 update 2024-12-09 10:25:53 +00:00
foreverpiano 2e4b523bee fps=24 2024-12-09 10:16:26 +00:00
foreverpiano ea0332f62a inference hunyuan ok 2024-12-09 10:10:59 +00:00
283 changed files with 17481 additions and 15794 deletions
+29 -112
View File
@@ -1,12 +1,17 @@
ucf101_stride4x4x4
__pycache__
*.mp4
.ipynb_checkpoints
*.pth
UCF-101/
results/
build/
fastvideo.egg-info/
wandb/
.idea
*.ipynb
*.jpg
!examples/dataset/lingbotworld2/image.jpg
*.mp3
*.safetensors
*.mp4
*.png
@@ -15,123 +20,35 @@ wandb/
*.pt
cache_dir/
wandb/
venv/
.venv/
test*
sample_video*
sample_image*
512*
720*
1024*
debug*
private*
caption*
*deepspeed*
revised*
129f*
all*
read*
YSH*
*pick*
*ysh*
hw*
257f*
513f*
taming*
221hw*
65x512x512
runs/
samples/
Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
outputs_video
checkpoints/
sbatch.sh
*.out
env
*.o
**/build/
**.pyc
**.txt
*.log
weights/
logs/
/Z-Image/
official_weights/
converted_weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
# Distribution / packaging
build/
dist/
*.egg-info/
*.egg
eggs/
.eggs/
# MkDocs documentation
site/
docs/getting_started/examples/
docs/examples/
docs/inference/examples/
docs/training/examples/
docs/distillation/examples/
!requirements-mkdocs.txt
# VSCode
.vscode/
# DS Store
.DS_Store
# vim swap files
*.swo
*.swp
# Python pickle files
*.pkl
# Reference videos (negations must come after the catch-all on line below)
# Static images
!docs/assets/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
!assets/images/**/*.png
!assets/images/**/*.jpg
!assets/images/**/*.jpeg
!assets/images/**/*.gif
!assets/videos/**/*.mp4
dmd_t2v_output/
preprocess_output_text/
# SvelteKit / Node artifacts under apps/fastvideo_studio/: see apps/fastvideo_studio/.gitignore
# Next.js / Node artifacts under apps/dreamverse/web/
apps/dreamverse/web/node_modules/
apps/dreamverse/web/.next/
apps/dreamverse/web/out/
apps/dreamverse/web/coverage/
apps/dreamverse/web/test-results/
apps/dreamverse/web/playwright-report/
apps/dreamverse/web/.env.local
apps/dreamverse/web/.env.development.local
apps/dreamverse/web/.env.test.local
# Generated by apps/dreamverse/scripts/install_native_ffmpeg.sh — host-specific
apps/dreamverse/scripts/ffmpeg-env.sh
apps/dreamverse/web/.env.production.local
# Unignore migrated Dreamverse product assets — root .gitignore globally
# ignores *.png/*.jpg/*.mp4/*.gif, but apps/dreamverse/web/public/ MUST
# be tracked (logo, icons, k2.png, etc.).
!apps/dreamverse/web/public/**/*.png
!apps/dreamverse/web/public/**/*.jpg
!apps/dreamverse/web/public/**/*.jpeg
!apps/dreamverse/web/public/**/*.mp4
!apps/dreamverse/web/public/**/*.gif
!apps/dreamverse/web/prompts/**/*.png
!apps/dreamverse/web/prompts/**/*.jpg
!apps/dreamverse/web/prompts/**/*.jpeg
!apps/dreamverse/web/prompts/**/*.mp4
!apps/dreamverse/web/prompts/**/*.gif
!apps/dreamverse/gpu-pool.svg
!apps/dreamverse/gpu-pool.drawio
.claude/
.codex/
.agents/tmp/
.sisyphus/
openspec/
fastvideo/tests/ssim/reference_videos/**
!fastvideo/tests/ssim/reference_videos/**/*.mp4
!fastvideo/tests/ssim/reference_videos/**/*.png
# Editor logs and local Python version pins (accidentally committed)
*.nvimlog
.nvimlog
.python-version
-10
View File
@@ -1,10 +0,0 @@
# fastvideo2/rl_rewards is vendored byte-identical from upstream (gated by
# sha256 in tests) — formatters must not touch it
exclude: ^fastvideo2/rl_rewards/
repos:
- repo: https://github.com/pre-commit/pre-commit-hooks
rev: v5.0.0
hooks:
- id: trailing-whitespace
- id: end-of-file-fixer
- id: check-merge-conflict
-89
View File
@@ -1,89 +0,0 @@
# Repository guidelines (branch `will/v2.1`)
This branch is the fastvideo2 MVP — one package, one model (Wan2.1), four
surfaces. Read `README.md` first.
## Layout
| Path | Role |
|---|---|
| `fastvideo2/card.py` | Frozen data cards, `derive()`, digest, validation (stdlib-only) |
| `fastvideo2/loop.py` | Driven-loop protocol + `LoopRunner` (stdlib; NVTX lazily) |
| `fastvideo2/pipeline.py` | Stage list with enforced `reads`/`writes` (stdlib-only) |
| `fastvideo2/loading.py` | Checkpoint → modules, standalone; component fingerprints |
| `fastvideo2/layers/` | Shared model layers (norms, MLP, rotary, attention) — torch-only, checkpoint-key compatible, cast semantics preserved (anchor-proven) |
| `fastvideo2/engine.py` | One-shot runner: request → outputs + identity-chained trace |
| `fastvideo2/verify.py` | Gates T0–T3 + evidence ledger |
| `fastvideo2/registry.py` | The only catalog: name → (card, pipeline builder) |
| `fastvideo2/wan21/` | The family's logic: card constant, loop, pipeline, vendored `model.py`, `reference.py` (the executable spec) |
| `fastvideo2/wan21/gates/` | The family's measurement side: goldens, anchor adapters, official capture shim, comparison CLI, diagnostics — never imported by logic code |
| `fastvideo2/evidence/` | Append-only ledger + blessed baselines (see its README) |
| `fastvideo2/tests/` | T0 contract tests — CPU, no torch, no weights |
## Invariants (enforced by review; violating them is the bug)
1. **Cards are pure data.** No callables, no live objects, no deploy-local
paths. If it can't round-trip through JSON, it doesn't belong on a card.
2. **Import direction is one-way:** `card` → `loop`/`pipeline` → `engine` →
`verify`. Family packages depend on core, never the reverse.
`wan21/reference.py` imports only the vendored official model file
(`wan21/model.py`, itself standalone) — never core/runtime modules — and
nothing outside `verify.py` may import the reference.
3. **Loop modules import torch-free** (torch inside methods) so contracts
validate anywhere; `import fastvideo2` must work without torch installed.
4. **Model-specific inputs are typed** (`WanForwardInputs`); never add an
untyped passthrough kwarg to a forward call.
5. **Evidence is append-only and human-owned.** Agents run `verify` and commit
the records; agents do not edit tolerances, re-bless baselines, or touch
`reference.py` to make a failing gate pass — say so instead.
6. **One catalog.** New servable ⇒ card constant + registry entry. No parallel
model lists.
7. **Official implementations are the numerics authority.** Where fidelity is
the requirement, run the authors' modeling code: the Wan DiT is vendored
from the pinned official commit (`wan21/model.py`, provenance in its
header). Restructuring vendored code (e.g. extracting `layers/`) is
allowed ONLY when the anchor stays bitwise 0.0 — the gate, not "verbatim",
is the equivalence guarantee. Two invariants for any extraction:
checkpoint keys unchanged (Sequential indices preserved) and cast/dtype
semantics unchanged (fp32 islands, promotion order, fp64 RoPE). Ports
(diffusers, etc.) may serve components only with anchor certification,
never on trust; the official repo never becomes a dependency or submodule;
when conventions conflict, official wins.
8. **One environment for goldens and gates.** The supported env is python 3.12
+ torch 2.12 (the fastvideo cluster venv). Goldens are captured with that
same env — official code rides `PYTHONPATH`, its extra deps go to a pip
`--target` dir — so anchor deltas measure implementation differences, never
torch/kernel version differences.
## Commands
```bash
pytest # T0, runs on a laptop
python -m fastvideo2 verify <model> --tier N # gates; appends evidence
python -m fastvideo2 describe <model> # card JSON + digest
python -m fastvideo2 generate <model> --prompt ... # one request
```
GPU work runs on dlcluster via the `run-fastvideo-dlcluster` skill from the
main FastVideo checkout (sync this branch with `git push origin HEAD`, then
run inside the branch clone at `/mnt/fv21` — do not disturb `/mnt/FastVideo`'s
checkout). One-time per environment, install the package editable with no
dependency changes (torch etc. already live in the venv) — after this,
scripts and `fastvideo2 <cmd>` work from any directory, no PYTHONPATH:
```bash
/mnt/FastVideo/.venv/bin/pip install -e /mnt/fv21 --no-deps -q
```
Cluster runs append to `fastvideo2/evidence/` and those files get
fetched and committed locally, so before every cluster `git pull`, reset that
tree or the pull conflicts:
```bash
git checkout -- fastvideo2/evidence; git clean -qfd fastvideo2/evidence; git pull
```
## Commit style
Short subject with a tag prefix (`[feat]: ...`, `[fix]: ...`, `[docs]: ...`).
Do not add AI co-author trailers.
+15 -1
View File
@@ -184,4 +184,18 @@
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
identification within third-party archives.
Copyright [2023] Lightning AI
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+146 -56
View File
@@ -1,72 +1,162 @@
# fastvideo2 — the v2.1 MVP (branch `will/v2.1`)
# FastVideo
A from-scratch, deliberately small substrate for the FastVideo big bet:
**post-training → inference-optimized serving**, designed so that both kinds of
agents — the ones that *build* the framework and the ones that will *operate*
video models inside products — get inspectable contracts, a ground-truth
oracle, machine-readable verification, and an identity-chained runtime.
<div align="center">
<a href=""><img src="https://img.shields.io/static/v1?label=API:H100&message=Replicate&color=pink"></a> &ensp;
<a href=""><img src="https://img.shields.io/static/v1?label=Discuss&message=Discord&color=purple&logo=discord"></a> &ensp;
</div>
<br>
<div align="center">
<img src=assets/logo.png width="50%"/>
</div>
This branch is a clean slate: everything except `LICENSE` was removed, and the
MVP supports exactly one model, **Wan2.1-T2V-1.3B**, end to end.
FastVideo is a scalable framework for post-training video diffusion models, addressing the growing challenges of fine-tuning, distillation, and inference as model sizes and sequence lengths increase. As a first step, it provides an efficient script for distilling and fine-tuning the 10B Mochi model, with plans to expand features and support for more models.
## The four surfaces
### Features
| Surface | Where | What it guarantees |
|---|---|---|
| **Contracts** | `fastvideo2/card.py`, `pipeline.py`, `loop.py` | Cards are frozen *data* (no callables): JSON round-trip, stable content digest — the identity used by deploy configs, trainers, and RL environment manifests alike. Pipeline stage edges (`reads`/`writes`) are enforced at run time. Loop classes carry a `semantics` id and provenance pins it: distilled weights cannot silently run under a base sampler. |
| **Reference** | `fastvideo2/wan21/reference.py` | The complete model in one standalone eager file — the textbook an agent copies, and the oracle the production path is measured against. Never imported by production code. |
| **Verifier** | `fastvideo2/verify.py`, `fastvideo2/evidence/` | Tiered gates: T0 contracts (CPU, seconds) → T1 component fingerprints vs a blessed baseline → T2 trajectory parity vs the reference, with tolerance calibrated by measured run-to-run self-noise (the determinism contract) → T3 decoded-output parity + anti-degeneracy anchors. Every run appends typed records (card digest + env fingerprint) to the evidence ledger. |
| **Trace** | `fastvideo2/engine.py`, `loop.py` | Every unit of work is named `request/stage/loop.step`; the same identity chain lands in the returned trace (typed timings) and in nested NVTX ranges, so Nsight correlates kernels to model-level identity for free. |
- FastMochi, a distilled Mochi model that can generate videos with merely 8 sampling steps.
- Finetuning with FSDP (both master weight and ema weight), sequence parallelism, and selective gradient checkpointing.
- LoRA coupled with pecomputed the latents and text embedding for minumum memory consumption.
- Finetuning with both image and videos.
## Quickstart
## Change Log
```bash
# machine-readable capability discovery
python -m fastvideo2 describe wan2.1-t2v-1.3b
# contracts only — CPU, no weights, no torch
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 0
pytest # the same contracts, as tests
- ```2024/12/06```: `FastMochi` v0.0.1 is released.
# GPU: bless the component baseline once, then gate against it
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 1 --bless
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 3
# generate — CLI or the SDK (the handle is the loaded card, modality-neutral)
python -m fastvideo2 generate wan2.1-t2v-1.3b --prompt "a cat surfing a wave" --out cat.mp4
python -c '
import fastvideo2 as fv2
model = fv2.load("wan2.1-t2v-1.3b") # -> Model (capabilities from the card)
model.generate("a cat surfing a wave", seed=7).save("cat.mp4")'
## Fast and High-Quality Text-to-video Generation
# the oracle, standalone (this file works copied out of the repo)
python -m fastvideo2.wan21.reference --prompt "a cat surfing a wave" --out ref.mp4
### 8-Step Results of FastMochi
<table class="center">
<td><img src=assets/8steps/1.gif width="320"></td></td>
<td><img src=assets/8steps/2.gif width="320"></td></td></td>
<tr>
<td style="text-align:center;" width="320">tmp</td>
<td style="text-align:center;" width="320">tmp</td>
<tr>
</table >
## Table of Contents
Jump to a specific section:
- [🔧 Installation](#-installation)
- [🚀 Inference](#-inference)
- [🎯 Distill](#-distill)
- [⚡ Finetune](#-lora-finetune)
## 🔧 Installation
```
conda create -n fastmochi python=3.10.0 -y && conda activate fastmochi
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
pip install "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo && pip install -e .
```
Weights resolve from the HF cache (`Wan-AI/Wan2.1-T2V-1.3B-Diffusers`) or an
explicit `--root`; components are stock diffusers/transformers modules, so
there is no conversion step and `load_component()` works standalone in a REPL.
## Design lineage (what this MVP encodes)
- **Cards as declared constants; variants as `derive()` diffs** — no builder
functions, no factory bags, no toy backends welded into production cards.
- **The card digest is the axle artifact**: the same identity a deploy config
points at, a trainer stamps provenance into, and an RL environment manifest
pins (`substitution: exact | bounded | quality-changing` is already on
`Provenance` for the post-training flywheel).
- **Typed conditioning** (`WanForwardInputs`): a new control channel is a new
field the forward must consume — never a silently dropped kwarg.
- **Verification is the product**: gates fail closed, evidence is append-only
data, baselines and tolerances are human-owned.
- **One loop, runtime-visible**: the driven-loop contract is what sessions,
interleaved serving, and RL rollout branching will consume next; the engine
stays a deliberately dumb one-shot runner until those consumers land.
## Scope and non-goals (MVP)
## 🚀 Inference
In: Wan2.1 T2V, bidirectional, single GPU, one-shot generation, tiers T0–T3.
Out (next, in order): causal/self-forcing students + sessions with forkable
state, the post-training flywheel emitting derived cards + evidence, the RL
environment server (`reset/step/branch`) over the same contracts, additional
model families via `derive()` and new recipe packages.
Use [scripts/download_hf.py](scripts/download_hf.py) to download the hugging-face style model to a local directory. Use it like this:
```bash
python scripts/download_hf.py --repo_id=FastVideo/FastMochi --local_dir=data/FastMochi --repo_type=model
```
Start the gradio UI with
```
python fastvideo/demo/gradio_web_demo.py --model_path data/FastMochi
```
We also provide CLI inference script featured with sequence parallelism.
```
export NUM_GPUS=4
torchrun --nnodes=1 --nproc_per_node=$NUM_GPUS \
fastvideo/sample/sample_t2v_mochi.py \
--model_path data/FastMochi \
--prompt_path assets/prompt.txt \
--num_frames 163 \
--height 480 \
--width 848 \
--num_inference_steps 8 \
--guidance_scale 1.5 \
--output_path outputs_video/demo_video \
--seed 12345 \
--scheduler_type "pcm_linear_quadratic" \
--linear_threshold 0.1 \
--linear_range 0.75
```
For the mochi style, simply following the scripts list in mochi repo.
```
git clone https://github.com/genmoai/mochi.git
cd mochi
# install env
...
python3 ./demos/cli.py --model_dir weights/ --cpu_offload
```
## 🎯 Distill
## 💰Hardware requirement
- VRAM is required for both distill 10B mochi model
To launch distillation, you will first need to prepare data in the following formats
```bash
asset/example_data
├── AAA.txt
├── AAA.png
├── BCC.txt
├── BCC.png
├── ......
├── CCC.txt
└── CCC.png
```
We provide a dataset example here. First download testing data. Use [scripts/download_hf.py](scripts/download_hf.py) to download the data to a local directory. Use it like this:
```bash
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-425-Data --local_dir=data/Merge-425-Data --repo_type=dataset
python scripts/download_hf.py --repo_id=Stealths-Video/validation_embeddings --local_dir=data/validation_embeddings --repo_type=dataset
```
Then the distillation can be launched by:
```
bash scripts/distill_t2v.sh
```
## ⚡ Lora Finetune
## 💰Hardware requirement
- VRAM is required for both distill 10B mochi model
To launch finetuning, you will first need to prepare data in the following formats.
Then the finetuning can be launched by:
```
bash scripts/lora_finetune.sh
```
## Acknowledgement
We learned from and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), and [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan).
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 380 KiB

+9
View File
@@ -0,0 +1,9 @@
A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.
A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.
A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.
En "The Matrix", Neo, interpretado por Keanu Reeves, personifica la lucha contra un sistema opresor a través de su icónica imagen, que incluye unos anteojos oscuros. Estos lentes no son solo un accesorio de moda; representan una barrera entre la realidad y la percepción. Al usar estos anteojos, Neo se sumerge en un mundo donde la verdad se oculta detrás de ilusiones y engaños. La oscuridad de los lentes simboliza la ignorancia y el control que las máquinas tienen sobre la humanidad, mientras que su propia búsqueda de la verdad lo lleva a descubrir sus auténticos poderes. La escena en que se los pone se convierte en un momento crucial, marcando su transformación de un simple programador a "El Elegido". Esta imagen se ha convertido en un ícono cultural, encapsulando el mensaje de que, al enfrentar la oscuridad, podemos encontrar la luz que nos guía hacia la libertad. Así, los anteojos de Neo se convierten en un símbolo de resistencia y autoconocimiento en un mundo manipulado.
Medium close up. Low-angle shot. A woman in a 1950s retro dress sits in a diner bathed in neon light, surrounded by classic decor and lively chatter. The camera starts with a medium shot of her sitting at the counter, then slowly zooms in as she blows a shiny pink bubblegum bubble. The bubble swells dramatically before popping with a soft, playful burst. The scene is vibrant and nostalgic, evoking the fun and carefree spirit of the 1950s.
Will Smith eats noodles.
A short clip of the blonde woman taking a sip from her whiskey glass, her eyes locking with the camera as she smirks playfully. The background shows a group of people laughing and enjoying the party, with vibrant neon signs illuminating the space. The shot is taken in a way that conveys the feeling of a tipsy, carefree night out. The camera then zooms in on her face as she winks, creating a cheeky, flirtatious vibe.
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robot's immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
A chimpanzee lead vocalist singing into a microphone on stage. The camera zooms in to show him singing. There is a spotlight on him.
-10
View File
@@ -1,10 +0,0 @@
# Examples
- `generate_t2v.py` — text-to-video via the SDK (`fv2.load(...)` →
`model.generate(...)` → `result.save(...)`). Card defaults, everything
overridable by flag. The standalone reference implementation (no SDK, no
runtime) lives at `fastvideo2/wan21/reference.py`.
The `--model` flag takes any catalog id — e.g. the 3-step FastWan students
(`fastwan-qad-fp8-1.3b`, `fastwan-t2v-1.3b`); their step count, sampler, and
sparsity/quant recipe come from the card, so no other flags change.
-47
View File
@@ -1,47 +0,0 @@
#!/usr/bin/env python3
"""Text-to-video with the fastvideo2 SDK — the canonical example.
python examples/generate_t2v.py --prompt "a cat surfing a wave" --out cat.mp4
Loads the card resident once, generates with card defaults (50 steps, 81
frames, 480x832 — override anything via flags), saves an mp4. Requires a CUDA
box; weights resolve from the HF cache on first use.
"""
from __future__ import annotations
import argparse
import fastvideo2 as fv2
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--prompt",
default="a golden retriever puppy running through a sprinkler "
"on a sunny lawn, water droplets sparkling, slow motion, cinematic")
p.add_argument("--model", default="wan2.1-t2v-1.3b",
help="a model id from the catalog (see `python -m fastvideo2 describe`)")
p.add_argument("--out", default="out.mp4")
p.add_argument("--seed", type=int, default=42)
p.add_argument("--num-steps", dest="num_steps", type=int, default=None,
help="unset -> the card's default")
p.add_argument("--num-frames", dest="num_frames", type=int, default=None)
p.add_argument("--guidance-scale", dest="guidance_scale", type=float, default=None)
args = p.parse_args()
model = fv2.load(args.model)
print(model)
overrides = {k: getattr(args, k) for k in ("seed", "num_steps", "num_frames", "guidance_scale")
if getattr(args, k) is not None}
result = model.generate(args.prompt, **overrides)
steps = [t for t in result.trace if "/denoise." in t["label"]]
print(f"video {result.video.shape} | {len(steps)} denoise steps, "
f"{sum(t['seconds'] for t in steps) / max(len(steps), 1):.2f}s/step, "
f"{result.seconds:.1f}s total")
print(f"saved -> {result.save(args.out)}")
if __name__ == "__main__":
main()
@@ -0,0 +1,156 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from diffusers.utils import export_to_video
import json
import os
import torch.distributed as dist
logger = get_logger(__name__)
from torch.utils.data import Dataset
from torch.utils.data.distributed import DistributedSampler
from torch.utils.data import DataLoader
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
class T5dataset(Dataset):
def __init__(
self,
json_path,
vae_debug,
):
self.json_path = json_path
self.vae_debug = vae_debug
with open(self.json_path, "r") as f:
train_dataset = json.load(f)
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
def __getitem__(self, idx):
caption = self.train_dataset[idx]["caption"]
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
length = self.train_dataset[idx]["length"]
if self.vae_debug:
latents = torch.load(
os.path.join(
args.output_dir, "latent", self.train_dataset[idx]["latent_path"]
),
map_location="cpu",
)
else:
latents = []
return dict(caption=caption, latents=latents, filename=filename, length=length)
def __len__(self):
return len(self.train_dataset)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
videoprocessor = VideoProcessor(vae_scale_factor=8)
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp_replace.json")
train_dataset = T5dataset(latents_json_path, args.vae_debug)
text_encoder = load_text_encoder(args.model_type,args.model_path, device=device)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
sampler = DistributedSampler(
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt=data["caption"],
)
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(
args.output_dir, "prompt_embed", video_name + ".pt"
)
video_path = os.path.join(
args.output_dir, "video", video_name + ".mp4"
)
prompt_attention_mask_path = os.path.join(
args.output_dir, "prompt_attention_mask", video_name + ".pt"
)
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
print(f"sample {video_name} saved")
if args.vae_debug:
export_to_video(video[idx], video_path, fps=fps)
item = {}
item["length"] = int(data["length"][idx])
item["latent_path"] = video_name + ".pt"
item["prompt_embed_path"] = video_name + ".pt"
item["prompt_attention_mask"] = video_name + ".pt"
item["caption"] = data["caption"][idx]
json_data.append(item)
dist.barrier()
local_data = json_data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
# os.remove(latents_json_path)
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
# text encoder & vae & diffusion model
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=1,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--vae_debug", action="store_true")
args = parser.parse_args()
main(args)
@@ -0,0 +1,128 @@
from fastvideo.dataset import getdataset
from torch.utils.data import DataLoader
from fastvideo.utils.dataset_utils import Collate
import argparse
import torch
from accelerate import Accelerator
from accelerate.logging import get_logger
from accelerate.utils import ProjectConfiguration
import json
import os
from diffusers import AutoencoderKLMochi
import torch.distributed as dist
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.load import load_vae
from tqdm import tqdm
logger = get_logger(__name__)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
train_dataset = getdataset(args)
sampler = DistributedSampler(
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
encoder_device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
vae, autocast_type = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
latents = vae.encode(data["pixel_values"].to(encoder_device))[
"latent_dist"
].sample()
for idx, video_path in enumerate(data["path"]):
video_name = os.path.basename(video_path).split(".")[0]
latent_path = os.path.join(
args.output_dir, "latent", video_name + ".pt"
)
torch.save(latents[idx].to(torch.bfloat16), latent_path)
item = {}
item["length"] = latents[idx].shape[1]
item["latent_path"] = video_name + ".pt"
item["caption"] = data["text"][idx]
json_data.append(item)
print(f"{video_name} processed")
dist.barrier()
local_data = json_data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
)
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=(
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
),
)
args = parser.parse_args()
main(args)
@@ -0,0 +1,69 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from diffusers.utils import export_to_video
import json
import os
import torch.distributed as dist
logger = get_logger(__name__)
from torch.utils.data import Dataset
from torch.utils.data.distributed import DistributedSampler
from torch.utils.data import DataLoader
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
text_encoder = load_text_encoder(args.model_type,args.model_path, device=device)
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
os.makedirs(os.path.join(args.output_dir,"validation"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir,"validation", "prompt_attention_mask"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir,"validation", "prompt_embed"), exist_ok=True)
json_data = []
with open(args.validation_prompt_txt, 'r', encoding='utf-8') as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt
)
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(args.output_dir,"validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(args.output_dir,"validation", "prompt_attention_mask", f"{file_name}.pt")
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
+111
View File
@@ -0,0 +1,111 @@
from transformers import AutoTokenizer
from torchvision import transforms
from torchvision.transforms import Lambda
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.latent_datasets import LatentDataset
from fastvideo.dataset.transform import (
Normalize255,
TemporalRandomCrop,
CenterCropResizeVideo,
)
def getdataset(args):
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
]
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose(
[
# Normalize255(),
*resize,
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
# norm_fun
]
)
transform_topcrop = transforms.Compose(
[
Normalize255(),
*resize_topcrop,
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
norm_fun,
]
)
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(
args.text_encoder_name, cache_dir=args.cache_dir
)
if args.dataset == "t2v":
return T2V_dataset(
args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
)
raise NotImplementedError(args.dataset)
if __name__ == "__main__":
from accelerate import Accelerator
from fastvideo.dataset.t2v_datasets import dataset_prog
import random
from tqdm import tqdm
args = type(
"args",
(),
{
"ae": "CausalVAEModel_4x8x8",
"dataset": "t2v",
"attention_mode": "xformers",
"use_rope": True,
"text_max_length": 300,
"max_height": 320,
"max_width": 240,
"num_frames": 1,
"use_image_num": 0,
"interpolation_scale_t": 1,
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
"use_img_from_vid": False,
"speed_factor": 1.0,
"cfg": 0.1,
"text_encoder_name": "google/mt5-xxl",
"dataloader_num_workers": 10,
},
)
accelerator = Accelerator()
dataset = getdataset(args)
num = len(dataset_prog.img_cap_list)
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [
i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data
]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
print(e)
# import ipdb;ipdb.set_trace()
print(image_data)
zero += 1
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
+126
View File
@@ -0,0 +1,126 @@
import torch
from torch.utils.data import Dataset
import json
import os
import random
class LatentDataset(Dataset):
def __init__(
self,
json_path,
num_latent_t,
cfg_rate,
):
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
self.cfg_rate = cfg_rate
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
self.prompt_attention_mask_dir = os.path.join(
self.datase_dir_path, "prompt_attention_mask"
)
with open(self.json_path, "r") as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
self.num_latent_t = num_latent_t
# just zero embeddings [256, 4096]
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, -self.num_latent_t :]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
else:
prompt_embed = torch.load(
os.path.join(self.prompt_embed_dir, prompt_embed_file),
map_location="cpu",
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(
self.prompt_attention_mask_dir, prompt_attention_mask_file
),
map_location="cpu",
weights_only=True,
)
return latent, prompt_embed, prompt_attention_mask
def __len__(self):
return len(self.data_anno)
def latent_collate_function(batch):
# return latent, prompt, latent_attn_mask, text_attn_mask
# latent_attn_mask: # b t h w
# text_attn_mask: b 1 l
# needs to check if the latent/prompt' size and apply padding & attn mask
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
# calculate max shape
max_t = max([latent.shape[1] for latent in latents])
max_h = max([latent.shape[2] for latent in latents])
max_w = max([latent.shape[3] for latent in latents])
# padding
latents = [
torch.nn.functional.pad(
latent,
(
0,
max_t - latent.shape[1],
0,
max_h - latent.shape[2],
0,
max_w - latent.shape[3],
),
)
for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latents):
latent_attn_mask[i, latent.shape[1] :, :, :] = 0
latent_attn_mask[i, :, latent.shape[2] :, :] = 0
latent_attn_mask[i, :, :, latent.shape[3] :] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latents, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
dataloader = torch.utils.data.DataLoader(
dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function
)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
prompt_embed.shape,
latent_attn_mask.shape,
prompt_attention_mask.shape,
)
import pdb
pdb.set_trace()
+348
View File
@@ -0,0 +1,348 @@
import json
import os, io, csv, math, random
import numpy as np
from einops import rearrange
from decord import VideoReader
from os.path import join as opj
from collections import Counter
import torch
from torch.utils.data.dataset import Dataset
from torch.utils.data import DataLoader, Dataset, get_worker_info
from tqdm import tqdm
from PIL import Image
from fastvideo.utils.dataset_utils import DecordInit
import torchvision
from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self):
self.cap_list = []
self.elements = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements = dict()
self.n_used_elements = dict()
def set_cap_list(self, num_workers, cap_list, n_elements):
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info):
if work_info is None:
worker_id = 0
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])
]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if not ("mt5" in args.text_encoder_name):
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx):
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW"
)
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(
pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
)
def get_image(self, idx):
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (
self.transform_topcrop(image)
if "human_images" in image_data["path"]
else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps = (
image_data["cap"]
if isinstance(image_data["cap"], list)
else [image_data["cap"]]
)
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
text = text if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list):
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (
resolution.get("height", None) is None
or resolution.get("width", None) is None
):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"]["width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# import ipdb;ipdb.set_trace()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if (
i["num_frames"] / fps
> self.video_length_tolerance_range
* (self.num_frames / self.train_fps * self.speed_factor)
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(
start_frame_idx, i["num_frames"], frame_interval
).astype(int)
# comment out it to enable dynamic frames training
if (
len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio
):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extention {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices):
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data):
cap_lists = []
with open(data, "r") as f:
folder_anno = [
i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno, "r") as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self):
cap_lists = self.read_jsons(self.data)
return cap_lists
+645
View File
@@ -0,0 +1,645 @@
import torch
import random
import numbers
from torchvision.transforms import RandomCrop, RandomResizedCrop
def _is_tensor_video_clip(clip):
if not torch.is_tensor(clip):
raise TypeError("clip should be Tensor. Got %s" % type(clip))
if not clip.ndimension() == 4:
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
return True
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]
)
def crop(clip, i, j, h, w):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i : i + h, j : j + w]
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
return torch.nn.functional.interpolate(
clip,
size=target_size,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
clip,
scale_factor=scale_,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
"""
Do spatial cropping and resizing to the video clip
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
i (int): i in (i,j) i.e coordinates of the upper left corner.
j (int): j in (i,j) i.e coordinates of the upper left corner.
h (int): Height of the cropped region.
w (int): Width of the cropped region.
size (tuple(int, int)): height and width of resized clip
Returns:
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
clip = crop(clip, i, j, h, w)
clip = resize(clip, size, interpolation_mode)
return clip
def center_crop(clip, crop_size):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
th, tw = crop_size
if h < th or w < tw:
raise ValueError("height and width must be no smaller than crop_size")
i = int(round((h - th) / 2.0))
j = int(round((w - tw) / 2.0))
return crop(clip, i, j, th, tw)
def center_crop_using_short_edge(clip):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h < w:
th, tw = h, h
i = 0
j = int(round((w - tw) / 2.0))
else:
th, tw = w, w
i = int(round((h - th) / 2.0))
j = 0
return crop(clip, i, j, th, tw)
def center_crop_th_tw(clip, th, tw, top_crop):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
# import ipdb;ipdb.set_trace()
h, w = clip.size(-2), clip.size(-1)
tr = th / tw
if h / w > tr:
new_h = int(w * tr)
new_w = w
else:
new_h = h
new_w = int(h / tr)
i = 0 if top_crop else int(round((h - new_h) / 2.0))
j = int(round((w - new_w) / 2.0))
return crop(clip, i, j, new_h, new_w)
def random_shift_crop(clip):
"""
Slide along the long edge, with the short edge as crop size
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h <= w:
long_edge = w
short_edge = h
else:
long_edge = h
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1,)).item()
j = torch.randint(0, w - tw + 1, size=(1,)).item()
return crop(clip, i, j, th, tw)
def normalize_video(clip):
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
permute the dimensions of clip tensor
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError(
"clip tensor should have data type uint8. Got %s" % str(clip.dtype)
)
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
def normalize(clip, mean, std, inplace=False):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
mean (tuple): pixel RGB mean. Size is (3)
std (tuple): pixel standard deviation. Size is (3)
Returns:
normalized clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
if not inplace:
clip = clip.clone()
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
# print(mean)
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
return clip
def hflip(clip):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
Returns:
flipped clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
return clip.flip(-1)
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
else:
self.size = size
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: randomly cropped video clip.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = self.size
if h < th or w < tw:
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1,)).item()
j = torch.randint(0, w - tw + 1, size=(1,)).item()
return i, j, th, tw
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: cropped video clip by stride.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
return 0, 0, th, tw # from top-left
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class LongSideResizeVideo:
"""
First use the long side,
then resize to the specified size
"""
def __init__(
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized video clip.
size is (T, C, 512, *) or (T, C, *, 512)
"""
_, _, h, w = clip.shape
if self.skip_low_resolution and max(h, w) <= self.size:
return clip
if h > w:
w = int(w * self.size / h)
h = self.size
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(
clip, target_size=(h, w), interpolation_mode=self.interpolation_mode
)
return resize_clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
center crop video, then resize to the specified size
"""
def __init__(
self,
size,
top_crop=False,
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(
clip, self.size[0], self.size[1], top_crop=self.top_crop
)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
interpolation_mode=self.interpolation_mode,
)
return clip_center_crop_resize
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class UCFCenterCropVideo:
"""
First scale to the specified size in equal proportion to the short edge,
then center cropping
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_resize = resize_scale(
clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode
)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class KineticsRandomCropResizeVideo:
"""
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
return clip_resize
class CenterCropVideo:
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop(clip, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize:
"""
Normalize the video clip by mean subtraction and division by standard deviation
Args:
mean (3-tuple): pixel RGB mean
std (3-tuple): pixel RGB standard deviation
inplace (boolean): whether do in-place normalization
"""
def __init__(self, mean, std, inplace=False):
self.mean = mean
self.std = std
self.inplace = inplace
def __call__(self, clip):
"""
Args:
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
"""
return normalize(clip, self.mean, self.std, self.inplace)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self):
pass
def __call__(self, clip):
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
return normalize_video(clip)
def __repr__(self) -> str:
return self.__class__.__name__
class RandomHorizontalFlipVideo:
"""
Flip the video clip along the horizontal direction with a given probability
Args:
p (float): probability of the clip being flipped. Default value is 0.5
"""
def __init__(self, p=0.5):
self.p = p
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Size is (T, C, H, W)
Return:
clip (torch.tensor): Size is (T, C, H, W)
"""
if random.random() < self.p:
clip = hflip(clip)
return clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(p={self.p})"
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop(object):
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size):
self.size = size
def __call__(self, total_frames):
rand_end = max(0, total_frames - self.size - 1)
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
class DynamicSampleDuration(object):
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, t_stride, extra_1):
self.t_stride = t_stride
self.extra_1 = extra_1
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(range(t + 1))[t // 2 :][
:: self.t_stride
] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
return 0, truncate_t
if __name__ == "__main__":
from torchvision import transforms
import torchvision.io as io
import numpy as np
from torchvision.utils import save_image
import os
vframes, aframes, info = io.read_video(
filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW"
)
trans = transforms.Compose(
[
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(
mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True
),
]
)
target_video_len = 32
frame_interval = 1
total_frames = len(vframes)
print(total_frames)
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
# Sampling video frames
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(
start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int
)
print(frame_indice)
select_vframes = vframes[frame_indice]
print(select_vframes.shape)
print(select_vframes.dtype)
select_vframes_trans = trans(select_vframes)
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(
dtype=torch.uint8
)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
for i in range(target_video_len):
save_image(
select_vframes_trans[i],
os.path.join("./test000", "%04d.png" % i),
normalize=True,
value_range=(-1, 1),
)
+142
View File
@@ -0,0 +1,142 @@
import gradio as gr
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
import tempfile
import os
import argparse
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs='+', default=[])
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="euler")
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument("--linear_threshold", type=float, default=0.025)
parser.add_argument("--linear_range", type=float, default=0.5)
parser.add_argument("--cpu_offload", action="store_true")
return parser.parse_args()
def load_model(args):
device = "cuda" if torch.cuda.is_available() else "cpu"
if args.scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, False, args.linear_threshold, args.linear_range)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder='transformer/')
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
pipe.to(device)
if args.cpu_offload:
pipe.enable_model_cpu_offload()
return pipe
def generate_video(prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, num_inference_steps, randomize_seed=False):
if randomize_seed:
seed = torch.randint(0, 1000000, (1,)).item()
pipe = load_model(args)
print("load model successfully")
generator = torch.Generator(device="cuda").manual_seed(seed)
if not use_negative_prompt:
negative_prompt = None
with torch.autocast("cuda", dtype=torch.bfloat16):
output = pipe(
prompt=[prompt],
negative_prompt=negative_prompt,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
).frames[0]
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
export_to_video(output, output_path, fps=30)
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze."
]
args = init_args()
with gr.Blocks() as demo:
gr.Markdown("# Mochi Video Generation Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(label="Height", minimum=256, maximum=1024, step=32, value=args.height)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(label="Number of Frames", minimum=8, maximum=256, value=args.num_frames)
guidance_scale = gr.Slider(label="Guidance Scale", minimum=1, maximum=20, value=args.guidance_scale)
num_inference_steps = gr.Slider(label="Inference Steps", minimum=10, maximum=100, value=args.num_inference_steps)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False
)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, num_inference_steps, randomize_seed],
outputs=[result, seed_output]
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+905
View File
@@ -0,0 +1,905 @@
import argparse
import math
import os
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
destroy_sequence_parallel_group,
get_sequence_parallel_state,
nccl_info,
)
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
import torch
from torch.distributed.fsdp import ShardingStrategy
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
import json
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.dataset_utils import LengthGroupedSampler
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from fastvideo.utils.load import get_no_split_modules, load_transformer
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from copy import deepcopy
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from safetensors.torch import save_file
from peft import LoraConfig
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
)
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
)
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
import time
from collections import deque
def main_print(content):
if int(os.environ["LOCAL_RANK"]) <= 0:
print(content)
def save_checkpoint(transformer, rank, output_dir, step):
main_print(f"--> saving checkpoint at step {step}")
with FSDP.state_dict_type(
transformer,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
cpu_state = transformer.state_dict()
# todo move to get_state_dict
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(transformer.config)
if 'dtype' in config_dict: del config_dict['dtype'] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
main_print(f"--> checkpoint saved at step {step}")
def reshard_fsdp(model):
for m in FSDP.fsdp_modules(model):
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / gradient_accumulation_steps
)
largest_singular_value = (
torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps
)
absolute_mean = torch.mean(torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item()
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_student_cfg_embed
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0,
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
model_input = normalize_dit_input(model_type, latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(
0, num_euler_timesteps, (bsz,), device=model_input.device
).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (
sigmas_prev * noise_scheduler.config.num_train_timesteps
).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
student_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if hunyuan_student_cfg_embed:
student_kwargs["guidance"] = torch.tensor([hunyuan_student_cfg_embed], device=noisy_model_input.device, dtype=torch.bfloat16) * 1000
model_pred = transformer(**student_kwargs)[0]
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase
)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = cond_teacher_output + w * (
cond_teacher_output - uncond_teacher_output
)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True
)
huber_c = 0.001
# loss = loss.mean()
loss = (
torch.mean(
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c**2)
- huber_c
)
/ gradient_accumulation_steps
)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float() ** 2))
* pred_decay_weight
/ gradient_accumulation_steps
)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (
torch.mean(model_pred.float() ** 2)
* pred_decay_weight
/ gradient_accumulation_steps
)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(
model_pred.detach().float(), model_pred_norm, gradient_accumulation_steps
)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(
ema_transformer.parameters(), transformer.parameters()
):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(), 1 - ema_decay)
)
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
dist.init_process_group("nccl")
torch.cuda.set_device(local_rank)
device = torch.cuda.current_device()
initialize_sequence_parallel_state(args.sp_size)
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weigths to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
transformer = load_transformer(args.model_type,args.dit_model_name_or_path, args.pretrained_model_name_or_path,torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16)
teacher_transformer = deepcopy(transformer)
if args.use_ema:
ema_transformer = deepcopy(transformer)
else:
ema_transformer = None
if args.use_lora:
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
transformer,
args.fsdp_sharding_startegy,
args.use_lora,
args.use_cpu_offload,
args.master_weight_type,
)
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
teacher_transformer = FSDP(
teacher_transformer,
**fsdp_kwargs,
)
if args.use_ema:
ema_transformer = FSDP(
ema_transformer,
**fsdp_kwargs,
)
main_print(f"--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
if args.use_ema:
apply_fsdp_checkpointing(ema_transformer, no_split_modules, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(
noise_scheduler.config.num_train_timesteps * args.linear_range
)
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps,
args.linear_quadratic_threshold,
linear_steps,
)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
sigmas = noise_scheduler.sigmas
solver = EulerSolver(
sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer
)
main_print(f"optimizer: {optimizer}")
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (
LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
)
if (args.group_frame or args.group_resolution)
else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
)
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader)
/ args.gradient_accumulation_steps
* args.sp_size
/ args.train_sp_batch_size
)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (
args.train_batch_size
* world_size
* args.gradient_accumulation_steps
/ args.sp_size
* args.train_sp_batch_size
)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError("resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_student_cfg_embed
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix(
{
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
}
)
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"pred_fro_norm": pred_norm["fro"],
"pred_largest_singular_value": pred_norm["largest singular value"],
"pred_absolute_mean": pred_norm["absolute mean"],
"pred_absolute_max": pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, step
)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir, step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, args.max_train_steps
)
else:
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_type",
type=str,
default="mochi",
help="The type of model to train."
)
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument(
"--seed", type=int, default=None, help="A seed for reproducible training."
)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=(
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."
),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=(
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
)
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument(
"--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA."
)
parser.add_argument(
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
)
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=(
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'
),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument(
"--distill_cfg", type=float, default=3.0, help="Distillation coefficient."
)
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument(
"--scheduler_type", type=str, default="pcm", help="The scheduler type to use."
)
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument(
"--weight_decay", type=float, default=0.001, help="Weight decay to apply."
)
parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_student_cfg_embed", type=float)
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+102
View File
@@ -0,0 +1,102 @@
from typing import Any, Dict, Optional, Union
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.models.attention import JointTransformerBlock
from diffusers.models.attention_processor import Attention, AttentionProcessor
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import AdaLayerNormContinuous
from diffusers.utils import (
USE_PEFT_BACKEND,
is_torch_version,
logging,
scale_lora_layers,
unscale_lora_layers,
)
from diffusers.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
from diffusers.models.transformers.transformer_2d import Transformer2DModelOutput
from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class DiscriminatorHead(nn.Module):
def __init__(self, input_channel, output_channel=1):
super().__init__()
inner_channel = 1024
self.conv1 = nn.Sequential(
nn.Conv2d(input_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv2 = nn.Sequential(
nn.Conv2d(inner_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
def forward(self, x):
b, twh, c = x.shape
t = twh // (30 * 53)
x = x.view(-1, 30 * 53, c)
x = x.permute(0, 2, 1)
x = x.view(b * t, c, 30, 53)
x = self.conv1(x)
x = self.conv2(x) + x
x = self.conv_out(x)
return x
class Discriminator(nn.Module):
def __init__(
self,
stride=8,
num_h_per_head=1,
adapter_channel_dims=[3072],
):
super().__init__()
adapter_channel_dims = adapter_channel_dims * (48 // stride)
self.stride = stride
self.num_h_per_head = num_h_per_head
self.head_num = len(adapter_channel_dims)
self.heads = nn.ModuleList(
[
nn.ModuleList(
[
DiscriminatorHead(adapter_channel)
for _ in range(self.num_h_per_head)
]
)
for adapter_channel in adapter_channel_dims
]
)
def forward(self, features):
outputs = []
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
assert len(features) // self.stride == len(self.heads)
for i in range(0, len(features), self.stride):
for h in self.heads[i // self.stride]:
# out = torch.utils.checkpoint.checkpoint(
# create_custom_forward(h),
# features[i],
# use_reentrant=False
# )
out = h(features[i])
outputs.append(out)
return outputs
+309
View File
@@ -0,0 +1,309 @@
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.utils.torch_utils import randn_tensor
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class PCMFMSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
pcm_timesteps: int = 50,
linear_quadratic=False,
linear_quadratic_threshold=0.025,
linear_range=0.5,
):
if linear_quadratic:
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(
num_train_timesteps, linear_quadratic_threshold, linear_steps
)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
timesteps = np.linspace(
1, num_train_timesteps, num_train_timesteps, dtype=np.float32
)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.euler_timesteps = (
np.arange(1, pcm_timesteps + 1) * (num_train_timesteps // pcm_timesteps)
).round().astype(np.int64) - 1
self.sigmas = sigmas.numpy()[::-1][self.euler_timesteps]
self.sigmas = torch.from_numpy((self.sigmas[::-1].copy()))
self.timesteps = self.sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
noise: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
if self.step_index is None:
self._init_step_index(timestep)
sigma = self.sigmas[self.step_index]
sample = sigma * noise + (1.0 - sigma) * sample
return sample
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(
self, num_inference_steps: int, device: Union[str, torch.device] = None
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
self.num_inference_steps = num_inference_steps
inference_indices = np.linspace(
0, self.config.pcm_timesteps, num=num_inference_steps, endpoint=False
)
inference_indices = np.floor(inference_indices).astype(np.int64)
inference_indices = torch.from_numpy(inference_indices).long()
self.sigmas_ = self.sigmas[inference_indices]
timesteps = self.sigmas_ * self.config.num_train_timesteps
self.timesteps = timesteps.to(device=device)
self.sigmas_ = torch.cat(
[self.sigmas_, torch.zeros(1, device=self.sigmas_.device)]
)
self._step_index = None
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
generator: Optional[torch.Generator] = None,
return_dict: bool = True,
) -> Union[PCMFMSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
s_churn (`float`):
s_tmin (`float`):
s_tmax (`float`):
s_noise (`float`, defaults to 1.0):
Scaling factor for noise added to the sample.
generator (`torch.Generator`, *optional*):
A random number generator.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if self.step_index is None:
self._init_step_index(timestep)
sample = sample.to(torch.float32)
sigma = self.sigmas_[self.step_index]
denoised = sample - model_output * sigma
derivative = (sample - denoised) / sigma
dt = self.sigmas_[self.step_index + 1] - sigma
prev_sample = sample + derivative * dt
prev_sample = prev_sample.to(model_output.dtype)
self._step_index += 1
if not return_dict:
return (prev_sample,)
return PCMFMSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
class EulerSolver:
def __init__(self, sigmas, timesteps=1000, euler_timesteps=50):
self.step_ratio = timesteps // euler_timesteps
self.euler_timesteps = (
np.arange(1, euler_timesteps + 1) * self.step_ratio
).round().astype(np.int64) - 1
self.euler_timesteps_prev = np.asarray([0] + self.euler_timesteps[:-1].tolist())
self.sigmas = sigmas[self.euler_timesteps]
self.sigmas_prev = np.asarray(
[sigmas[0]] + sigmas[self.euler_timesteps[:-1]].tolist()
) # either use sigma0 or 0
self.euler_timesteps = torch.from_numpy(self.euler_timesteps).long()
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
self.sigmas = torch.from_numpy(self.sigmas)
self.sigmas_prev = torch.from_numpy(self.sigmas_prev)
def to(self, device):
self.euler_timesteps = self.euler_timesteps.to(device)
self.euler_timesteps_prev = self.euler_timesteps_prev.to(device)
self.sigmas = self.sigmas.to(device)
self.sigmas_prev = self.sigmas_prev.to(device)
return self
def euler_step(self, sample, model_pred, timestep_index):
sigma = extract_into_tensor(self.sigmas, timestep_index, model_pred.shape)
sigma_prev = extract_into_tensor(
self.sigmas_prev, timestep_index, model_pred.shape
)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev
def euler_style_multiphase_pred(
self,
sample,
model_pred,
timestep_index,
multiphase,
is_target=False,
):
inference_indices = np.linspace(
0, len(self.euler_timesteps), num=multiphase, endpoint=False
)
inference_indices = np.floor(inference_indices).astype(np.int64)
inference_indices = (
torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device)
)
expanded_timestep_index = timestep_index.unsqueeze(1).expand(
-1, inference_indices.size(0)
)
valid_indices_mask = expanded_timestep_index >= inference_indices
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
timestep_index_end = inference_indices[last_valid_index]
if is_target:
sigma = extract_into_tensor(self.sigmas_prev, timestep_index, sample.shape)
else:
sigma = extract_into_tensor(self.sigmas, timestep_index, sample.shape)
sigma_prev = extract_into_tensor(
self.sigmas_prev, timestep_index_end, sample.shape
)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev, timestep_index_end
+936
View File
@@ -0,0 +1,936 @@
import argparse
from email.policy import strict
import logging
import math
import os
import shutil
from pathlib import Path
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
destroy_sequence_parallel_group,
get_sequence_parallel_state,
nccl_info,
)
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
import torch
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
import json
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.dataset_utils import LengthGroupedSampler
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import (
get_dit_fsdp_kwargs,
apply_fsdp_checkpointing,
get_discriminator_fsdp_kwargs,
)
import diffusers
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from copy import deepcopy
from diffusers.optimization import get_scheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from peft import LoraConfig
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
)
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
resume_training,
save_checkpoint_generator_discriminator,
resume_training_generator_discriminator,
)
from fastvideo.utils.logging_ import main_print
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
import time
from collections import deque
def gan_d_loss(
discriminator,
teacher_transformer,
sample_fake,
sample_real,
timestep,
encoder_hidden_states,
encoder_attention_mask,
weight,
):
loss = 0.0
# collate sample_fake and sample_real
with torch.no_grad():
fake_features = teacher_transformer(
sample_fake,
encoder_hidden_states,
timestep,
encoder_attention_mask,
output_attn=True,
return_dict=False,
)[1]
real_features = teacher_transformer(
sample_real,
encoder_hidden_states,
timestep,
encoder_attention_mask,
output_attn=True,
return_dict=False,
)[1]
fake_outputs = discriminator(fake_features)
real_outputs = discriminator(real_features)
for fake_output, real_output in zip(fake_outputs, real_outputs):
loss += (
torch.mean(weight * torch.relu(fake_output.float() + 1))
+ torch.mean(weight * torch.relu(1 - real_output.float()))
) / (discriminator.head_num * discriminator.num_h_per_head)
return loss
def gan_g_loss(
discriminator,
teacher_transformer,
sample_fake,
timestep,
encoder_hidden_states,
encoder_attention_mask,
weight,
):
loss = 0.0
features = teacher_transformer(
sample_fake,
encoder_hidden_states,
timestep,
encoder_attention_mask,
output_attn=True,
return_dict=False,
)[1]
fake_outputs = discriminator(
features,
)
for fake_output in fake_outputs:
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
discriminator.head_num * discriminator.num_h_per_head
)
return loss
def train_one_step_mochi(
transformer,
teacher_transformer,
optimizer,
discriminator,
discriminator_optimizer,
global_step,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
sp_size,
precondition_outputs,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
adv_weight,
):
optimizer.zero_grad()
discriminator_optimizer.zero_grad()
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
model_input = normalize_mochi_dit_input(latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(
0, num_euler_timesteps, (bsz,), device=model_input.device
).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev * noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
model_pred = transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase
)
weighting = 1.0
# # simplified flow matching aka 0-rectified flow matching loss
# # target = model_input - noise
# target = model_input
adv_index = torch.empty_like(end_index)
for i in range(end_index.size(0)):
adv_index[i] = torch.randint(
end_index[i].item(),
end_index[i].item() + num_euler_timesteps // multiphase,
(1,),
dtype=end_index.dtype,
device=end_index.device,
)
sigmas_end = extract_into_tensor(solver.sigmas_prev, end_index, model_input.shape)
sigmas_adv = extract_into_tensor(solver.sigmas_prev, adv_index, model_input.shape)
timesteps_end = (sigmas_end * noise_scheduler.config.num_train_timesteps).view(-1)
timesteps_adv = (sigmas_adv * noise_scheduler.config.num_train_timesteps).view(-1)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = cond_teacher_output + w * (
cond_teacher_output - uncond_teacher_output
)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True
)
real_adv = (
(1 - sigmas_adv) * target + (sigmas_adv - sigmas_end) * torch.randn_like(target)
) / (1 - sigmas_end)
fake_adv = (
(1 - sigmas_adv) * model_pred
+ (sigmas_adv - sigmas_end) * torch.randn_like(model_pred)
) / (1 - sigmas_end)
huber_c = 0.001
g_loss = torch.mean(
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c**2) - huber_c
)
discriminator.requires_grad_(False)
with torch.autocast("cuda", dtype=torch.bfloat16):
g_gan_loss = adv_weight * gan_g_loss(
discriminator,
teacher_transformer,
fake_adv.float(),
timesteps_adv,
encoder_hidden_states.float(),
encoder_attention_mask,
1.0,
)
g_loss += g_gan_loss
g_loss.backward()
g_loss = g_loss.detach().clone()
dist.all_reduce(g_loss, op=dist.ReduceOp.AVG)
g_grad_norm = transformer.clip_grad_norm_(max_grad_norm).item()
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
discriminator_optimizer.zero_grad()
discriminator.requires_grad_(True)
with torch.autocast("cuda", dtype=torch.bfloat16):
d_loss = gan_d_loss(
discriminator,
teacher_transformer,
fake_adv.detach(),
real_adv.detach(),
timesteps_adv,
encoder_hidden_states,
encoder_attention_mask,
1.0,
)
d_loss.backward()
d_grad_norm = discriminator.clip_grad_norm_(max_grad_norm).item()
discriminator_optimizer.step()
discriminator_optimizer.zero_grad()
return g_loss, g_grad_norm, d_loss, d_grad_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
dist.init_process_group("nccl")
torch.cuda.set_device(local_rank)
device = torch.cuda.current_device()
initialize_sequence_parallel_state(args.sp_size)
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weigths to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
# keep the master weight to float32
if args.dit_model_name_or_path:
transformer = transformer = MochiTransformer3DModel.from_pretrained(
args.dit_model_name_or_path,
torch_dtype=torch.float32,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=torch.float32,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
teacher_transformer = deepcopy(transformer)
discriminator = Discriminator(args.discriminator_head_stride)
if args.use_lora:
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
main_print(
f" Total transformer parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
# discriminator
main_print(
f" Total discriminator parameters = {sum(p.numel() for p in discriminator.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
fsdp_kwargs = get_dit_fsdp_kwargs(
args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload
)
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(args.master_weight_type)
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = ["MochiTransformerBlock"]
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
teacher_transformer = FSDP(
teacher_transformer,
**fsdp_kwargs,
)
discriminator = FSDP(
discriminator,
**discriminator_fsdp_kwargs,
)
main_print(f"--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
teacher_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold
)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
sigmas = noise_scheduler.sigmas
solver = EulerSolver(
sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=1e-3,
eps=1e-8,
)
discriminator_optimizer = torch.optim.AdamW(
discriminator.parameters(),
lr=args.discriminator_learning_rate,
betas=(0, 0.999),
weight_decay=1e-3,
eps=1e-8,
)
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer
)
elif args.resume_from_checkpoint:
(
transformer,
optimizer,
discriminator,
discriminator_optimizer,
init_steps,
) = resume_training_generator_discriminator(
transformer,
optimizer,
discriminator,
discriminator_optimizer,
args.resume_from_checkpoint,
rank,
)
main_print(f"optimizer: {optimizer}")
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (
LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
)
if (args.group_frame or args.group_resolution)
else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
)
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
assert args.gradient_accumulation_steps == 1
num_update_steps_per_epoch = math.ceil(
len(train_dataloader)
/ args.gradient_accumulation_steps
* args.sp_size
/ args.train_sp_batch_size
)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (
args.train_batch_size
* world_size
* args.gradient_accumulation_steps
/ args.sp_size
* args.train_sp_batch_size
)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# log_validation(args, transformer, device,
# torch.bfloat16, init_steps, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold, ema=False)
for i in range(init_steps):
_ = next(loader)
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
(
generator_loss,
generator_grad_norm,
discriminator_loss,
discriminator_grad_norm,
) = train_one_step_mochi(
transformer,
teacher_transformer,
optimizer,
discriminator,
discriminator_optimizer,
step,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
args.validation_sampling_steps,
args.not_apply_cfg_solver,
args.distill_cfg,
args.adv_weight,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix(
{
"g_loss": f"{generator_loss:.4f}",
"d_loss": f"{discriminator_loss:.4f}",
"g_grad_norm": generator_grad_norm,
"d_grad_norm": discriminator_grad_norm,
"step_time": f"{step_time:.2f}s",
}
)
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"generator_loss": generator_loss,
"discriminator_loss": discriminator_loss,
"generator_grad_norm": generator_grad_norm,
"discriminator_grad_norm": discriminator_grad_norm,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
},
step=step,
)
if step % args.checkpointing_steps == 0:
main_print(f"--> saving checkpoint at step {step}")
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, step
)
else:
# Your existing checkpoint saving code
save_checkpoint_generator_discriminator(
transformer,
optimizer,
discriminator,
discriminator_optimizer,
rank,
args.output_dir,
step,
)
main_print(f"--> checkpoint saved at step {step}")
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
ema=False,
)
if args.use_lora:
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, args.max_train_steps
)
else:
save_checkpoint(
transformer, optimizer, rank, args.output_dir, args.max_train_steps
)
save_checkpoint(
discriminator,
discriminator_optimizer,
rank,
args.output_dir,
step,
discriminator=True,
)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.999)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
parser.add_argument(
"--precondition_outputs",
action="store_true",
help="Whether to precondition the outputs of the model.",
)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=int, default=64)
parser.add_argument("--validation_guidance_scale", type=float, default=4.5)
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument(
"--seed", type=int, default=None, help="A seed for reproducible training."
)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=(
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."
),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=(
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--discriminator_learning_rate",
type=float,
default=1e-5,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
)
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument(
"--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA."
)
parser.add_argument(
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
)
parser.add_argument("--fsdp_sharding_startegy", default="full")
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=(
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'
),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument(
"--distill_cfg", type=float, default=3.0, help="Distillation coefficient."
)
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument(
"--scheduler_type", type=str, default="pcm", help="The scheduler type to use."
)
parser.add_argument(
"--adv_weight",
type=float,
default=0.1,
help="The weight of the adversarial loss.",
)
parser.add_argument(
"--discriminator_head_stride",
type=int,
default=2,
help="The stride of the discriminator head.",
)
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="The threshold of the linear quadratic scheduler.",
)
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+34
View File
@@ -0,0 +1,34 @@
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from einops import rearrange
def flash_attn_no_pad(
qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None
):
# adapted from https://github.com/Dao-AILab/flash-attention/blob/13403e81157ba37ca525890f2f0f2137edf75311/flash_attn/flash_attention.py#L27
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
x, key_padding_mask
)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
output_unpad = flash_attn_varlen_qkvpacked_func(
x_unpad,
cu_seqlens,
max_s,
dropout_p,
softmax_scale=softmax_scale,
causal=causal,
)
output = rearrange(
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size, seqlen
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
+90
View File
@@ -0,0 +1,90 @@
import os
import torch
__all__ = [
"C_SCALE",
"PROMPT_TEMPLATE",
"MODEL_BASE",
"PRECISIONS",
"NORMALIZATION_TYPE",
"ACTIVATION_TYPE",
"VAE_PATH",
"TEXT_ENCODER_PATH",
"TOKENIZER_PATH",
"TEXT_PROJECTION",
"DATA_TYPE",
"NEGATIVE_PROMPT",
]
PRECISION_TO_TYPE = {
'fp32': torch.float32,
'fp16': torch.float16,
'bf16': torch.bfloat16,
}
# =================== Constant Values =====================
# Computation scale factor, 1P = 1_000_000_000_000_000. Tensorboard will display the value in PetaFLOPS to avoid
# overflow error when tensorboard logging values.
C_SCALE = 1_000_000_000_000_000
# When using decoder-only models, we must provide a prompt template to instruct the text encoder
# on how to generate the text.
# --------------------------------------------------------------------
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"dit-llm-encode": {
"template": PROMPT_TEMPLATE_ENCODE,
"crop_start": 36,
},
"dit-llm-encode-video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
},
}
# ======================= Model ======================
PRECISIONS = {"fp32", "fp16", "bf16"}
NORMALIZATION_TYPE = {"layer", "rms"}
ACTIVATION_TYPE = {"relu", "silu", "gelu", "gelu_tanh"}
# =================== Model Path =====================
MODEL_BASE = os.getenv("MODEL_BASE", "./data/hunyuan")
# =================== Data =======================
DATA_TYPE = {"image", "video", "image_video"}
# 3D VAE
VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
# Text Encoder
TEXT_ENCODER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
# Tokenizer
TOKENIZER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
TEXT_PROJECTION = {
"linear", # Default, an nn.Linear() layer
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
}
@@ -0,0 +1,2 @@
from .pipelines import HunyuanVideoPipeline
from .schedulers import FlowMatchDiscreteScheduler
@@ -0,0 +1 @@
from .pipeline_hunyuan_video import HunyuanVideoPipeline
File diff suppressed because it is too large Load Diff
@@ -0,0 +1 @@
from .scheduling_flow_match_discrete import FlowMatchDiscreteScheduler
@@ -0,0 +1,257 @@
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.schedulers.scheduling_utils import SchedulerMixin
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
"""
Euler scheduler.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(
self,
num_inference_steps: int,
device: Union[str, torch.device] = None,
n_tokens: int = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
if not self.config.reverse:
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device
)
# Reset step index
self._step_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def scale_model_input(
self, sample: torch.Tensor, timestep: Optional[int] = None
) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return (prev_sample,)
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
+392
View File
@@ -0,0 +1,392 @@
import argparse
from .constants import *
import re
from .modules.models import HUNYUAN_VIDEO_CONFIG
def parse_args(namespace=None):
parser = argparse.ArgumentParser(description="HunyuanVideo inference script")
parser = add_network_args(parser)
parser = add_extra_models_args(parser)
parser = add_denoise_schedule_args(parser)
parser = add_inference_args(parser)
parser = add_parallel_args(parser)
args = parser.parse_args(namespace=namespace)
args = sanity_check_args(args)
return args
def add_network_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="HunyuanVideo network args")
# Main model
group.add_argument(
"--model",
type=str,
choices=list(HUNYUAN_VIDEO_CONFIG.keys()),
default="HYVideo-T/2-cfgdistill",
)
group.add_argument(
"--latent-channels",
type=str,
default=16,
help="Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
"it still needs to match the latent channels of the VAE model.",
)
group.add_argument(
"--precision",
type=str,
default="bf16",
choices=PRECISIONS,
help="Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
)
# RoPE
group.add_argument(
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
)
return parser
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(
title="Extra models args, including vae, text encoders and tokenizers)"
)
# - VAE
group.add_argument(
"--vae",
type=str,
default="884-16c-hy",
choices=list(VAE_PATH),
help="Name of the VAE model.",
)
group.add_argument(
"--vae-precision",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the VAE model.",
)
group.add_argument(
"--vae-tiling",
action="store_true",
help="Enable tiling for the VAE model to save GPU memory.",
)
group.set_defaults(vae_tiling=True)
group.add_argument(
"--text-encoder",
type=str,
default="llm",
choices=list(TEXT_ENCODER_PATH),
help="Name of the text encoder model.",
)
group.add_argument(
"--text-encoder-precision",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the text encoder model.",
)
group.add_argument(
"--text-states-dim",
type=int,
default=4096,
help="Dimension of the text encoder hidden states.",
)
group.add_argument(
"--text-len", type=int, default=256, help="Maximum length of the text input."
)
group.add_argument(
"--tokenizer",
type=str,
default="llm",
choices=list(TOKENIZER_PATH),
help="Name of the tokenizer model.",
)
group.add_argument(
"--prompt-template",
type=str,
default="dit-llm-encode",
choices=PROMPT_TEMPLATE,
help="Image prompt template for the decoder-only text encoder model.",
)
group.add_argument(
"--prompt-template-video",
type=str,
default="dit-llm-encode-video",
choices=PROMPT_TEMPLATE,
help="Video prompt template for the decoder-only text encoder model.",
)
group.add_argument(
"--hidden-state-skip-layer",
type=int,
default=2,
help="Skip layer for hidden states.",
)
group.add_argument(
"--apply-final-norm",
action="store_true",
help="Apply final normalization to the used text encoder hidden states.",
)
# - CLIP
group.add_argument(
"--text-encoder-2",
type=str,
default="clipL",
choices=list(TEXT_ENCODER_PATH),
help="Name of the second text encoder model.",
)
group.add_argument(
"--text-encoder-precision-2",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the second text encoder model.",
)
group.add_argument(
"--text-states-dim-2",
type=int,
default=768,
help="Dimension of the second text encoder hidden states.",
)
group.add_argument(
"--tokenizer-2",
type=str,
default="clipL",
choices=list(TOKENIZER_PATH),
help="Name of the second tokenizer model.",
)
group.add_argument(
"--text-len-2",
type=int,
default=77,
help="Maximum length of the second text input.",
)
return parser
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Denoise schedule args")
group.add_argument(
"--denoise-type",
type=str,
default="flow",
help="Denoise type for noised inputs.",
)
# Flow Matching
group.add_argument(
"--flow-shift",
type=float,
default=7.0,
help="Shift factor for flow matching schedulers.",
)
group.add_argument(
"--flow-reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow-solver",
type=str,
default="euler",
help="Solver for flow matching.",
)
group.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
help="Use linear quadratic schedule for flow matching."
"Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
)
group.add_argument(
"--linear-schedule-end",
type=int,
default=25,
help="End step for linear quadratic schedule for flow matching.",
)
return parser
def add_inference_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Inference args")
# ======================== Model loads ========================
group.add_argument(
"--model-base",
type=str,
default="ckpts",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--dit-weight",
type=str,
default="ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
help="Path to the HunyuanVideo model. If None, search the model in the args.model_root."
"1. If it is a file, load the model directly."
"2. If it is a directory, search the model in the directory. Support two types of models: "
"1) named `pytorch_model_*.pt`"
"2) named `*_model_states.pt`, where * can be `mp_rank_00`.",
)
group.add_argument(
"--model-resolution",
type=str,
default="540p",
choices=["540p", "720p"],
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
type=str,
default="module",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
group.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
# ======================== Inference general setting ========================
group.add_argument(
"--batch-size",
type=int,
default=1,
help="Batch size for inference and evaluation.",
)
group.add_argument(
"--infer-steps",
type=int,
default=50,
help="Number of denoising steps for inference.",
)
group.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
group.add_argument(
"--save-path",
type=str,
default="./results",
help="Path to save the generated samples.",
)
group.add_argument(
"--save-path-suffix",
type=str,
default="",
help="Suffix for the directory of saved samples.",
)
group.add_argument(
"--name-suffix",
type=str,
default="",
help="Suffix for the names of saved samples.",
)
group.add_argument(
"--num-videos",
type=int,
default=1,
help="Number of videos to generate for each prompt.",
)
# ---sample size---
group.add_argument(
"--video-size",
type=int,
nargs="+",
default=(720, 1280),
help="Video size for training. If a single value is provided, it will be used for both height "
"and width. If two values are provided, they will be used for height and width "
"respectively.",
)
group.add_argument(
"--video-length",
type=int,
default=129,
help="How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
)
# --- prompt ---
group.add_argument(
"--prompt",
type=str,
default=None,
help="Prompt for sampling during evaluation.",
)
group.add_argument(
"--seed-type",
type=str,
default="auto",
choices=["file", "random", "fixed", "auto"],
help="Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
"random seed. If fixed, use the fixed seed given by `--seed`. If auto, `csv` will use the "
"seed column if available, otherwise use the fixed `seed` value. `prompt` will use the "
"fixed `seed` value.",
)
group.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument(
"--neg-prompt", type=str, default=None, help="Negative prompt for sampling."
)
group.add_argument(
"--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale."
)
group.add_argument(
"--embedded-cfg-scale",
type=float,
default=6.0,
help="Embeded classifier free guidance scale.",
)
group.add_argument(
"--reproduce",
action="store_true",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
return parser
def add_parallel_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Parallel args")
# ======================== Model loads ========================
group.add_argument(
"--ulysses-degree",
type=int,
default=1,
help="Ulysses degree.",
)
group.add_argument(
"--ring-degree",
type=int,
default=1,
help="Ulysses degree.",
)
return parser
def sanity_check_args(args):
# VAE channels
vae_pattern = r"\d{2,3}-\d{1,2}c-\w+"
if not re.match(vae_pattern, args.vae):
raise ValueError(
f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'."
)
vae_channels = int(args.vae.split("-")[1][:-1])
if args.latent_channels is None:
args.latent_channels = vae_channels
if vae_channels != args.latent_channels:
raise ValueError(
f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels})."
)
return args
+519
View File
@@ -0,0 +1,519 @@
import os
import time
import random
import functools
from typing import List, Optional, Tuple, Union
from pathlib import Path
from loguru import logger
import torch
import torch.distributed as dist
from fastvideo.models.hunyuan.constants import PROMPT_TEMPLATE, NEGATIVE_PROMPT, PRECISION_TO_TYPE
from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.models.hunyuan.modules import load_model
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.models.hunyuan.utils.data_utils import align_to
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
class Inference(object):
def __init__(
self,
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=None,
pipeline=None,
use_cpu_offload=False,
device=None,
logger=None,
parallel_args=None,
):
self.vae = vae
self.vae_kwargs = vae_kwargs
self.text_encoder = text_encoder
self.text_encoder_2 = text_encoder_2
self.model = model
self.pipeline = pipeline
self.use_cpu_offload = use_cpu_offload
self.args = args
self.device = (
device
if device is not None
else "cuda"
if torch.cuda.is_available()
else "cpu"
)
self.logger = logger
self.parallel_args = parallel_args
@classmethod
def from_pretrained(cls, pretrained_model_path, args, device=None, **kwargs):
"""
Initialize the Inference pipeline.
Args:
pretrained_model_path (str or pathlib.Path): The model path, including t2v, text encoder and vae checkpoints.
args (argparse.Namespace): The arguments for the pipeline.
device (int): The device for inference. Default is 0.
"""
# ========================================================================
logger.info(f"Got text-to-video model root path: {pretrained_model_path}")
# ==================== Initialize Distributed Environment ================
if nccl_info.sp_size > 1:
device = torch.device(f"cuda:{os.environ['LOCAL_RANK']}")
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
parallel_args = None #{"ulysses_degree": args.ulysses_degree, "ring_degree": args.ring_degree}
# ======================== Get the args path =============================
# Disable gradient
torch.set_grad_enabled(False)
# =========================== Build main model ===========================
logger.info("Building model...")
factor_kwargs = {"device": device, "dtype": PRECISION_TO_TYPE[args.precision]}
in_channels = args.latent_channels
out_channels = args.latent_channels
model = load_model(
args,
in_channels=in_channels,
out_channels=out_channels,
factor_kwargs=factor_kwargs,
)
model = model.to(device)
model = Inference.load_state_dict(args, model, pretrained_model_path)
model.eval()
# ============================= Build extra models ========================
# VAE
vae, _, s_ratio, t_ratio = load_vae(
args.vae,
args.vae_precision,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
vae_kwargs = {"s_ratio": s_ratio, "t_ratio": t_ratio}
# Text encoder
if args.prompt_template_video is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get(
"crop_start", 0
)
elif args.prompt_template is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
else:
crop_start = 0
max_length = args.text_len + crop_start
# prompt_template
prompt_template = (
PROMPT_TEMPLATE[args.prompt_template]
if args.prompt_template is not None
else None
)
# prompt_template_video
prompt_template_video = (
PROMPT_TEMPLATE[args.prompt_template_video]
if args.prompt_template_video is not None
else None
)
text_encoder = TextEncoder(
text_encoder_type=args.text_encoder,
max_length=max_length,
text_encoder_precision=args.text_encoder_precision,
tokenizer_type=args.tokenizer,
prompt_template=prompt_template,
prompt_template_video=prompt_template_video,
hidden_state_skip_layer=args.hidden_state_skip_layer,
apply_final_norm=args.apply_final_norm,
reproduce=args.reproduce,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
text_encoder_2 = None
if args.text_encoder_2 is not None:
text_encoder_2 = TextEncoder(
text_encoder_type=args.text_encoder_2,
max_length=args.text_len_2,
text_encoder_precision=args.text_encoder_precision_2,
tokenizer_type=args.tokenizer_2,
reproduce=args.reproduce,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
return cls(
args=args,
vae=vae,
vae_kwargs=vae_kwargs,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
model=model,
use_cpu_offload=args.use_cpu_offload,
device=device,
logger=logger,
parallel_args=parallel_args
)
@staticmethod
def load_state_dict(args, model, pretrained_model_path):
load_key = args.load_key
dit_weight = Path(args.dit_weight)
if dit_weight is None:
model_dir = pretrained_model_path / f"t2v_{args.model_resolution}"
files = list(model_dir.glob("*.pt"))
if len(files) == 0:
raise ValueError(f"No model weights found in {model_dir}")
if str(files[0]).startswith("pytorch_model_"):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
else:
if dit_weight.is_dir():
files = list(dit_weight.glob("*.pt"))
if len(files) == 0:
raise ValueError(f"No model weights found in {dit_weight}")
if str(files[0]).startswith("pytorch_model_"):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
elif dit_weight.is_file():
model_path = dit_weight
bare_model = "unknown"
else:
raise ValueError(f"Invalid model path: {dit_weight}")
if not model_path.exists():
raise ValueError(f"model_path not exists: {model_path}")
logger.info(f"Loading torch model {model_path}...")
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage)
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}."
)
model.load_state_dict(state_dict, strict=True)
return model
@staticmethod
def parse_size(size):
if isinstance(size, int):
size = [size]
if not isinstance(size, (list, tuple)):
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
if len(size) == 1:
size = [size[0], size[0]]
if len(size) != 2:
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
return size
class HunyuanVideoSampler(Inference):
def __init__(
self,
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=None,
pipeline=None,
use_cpu_offload=False,
device=0,
logger=None,
parallel_args=None
):
super().__init__(
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=text_encoder_2,
pipeline=pipeline,
use_cpu_offload=use_cpu_offload,
device=device,
logger=logger,
parallel_args=parallel_args
)
self.pipeline = self.load_diffusion_pipeline(
args=args,
vae=self.vae,
text_encoder=self.text_encoder,
text_encoder_2=self.text_encoder_2,
model=self.model,
device=self.device,
)
self.default_negative_prompt = NEGATIVE_PROMPT
def load_diffusion_pipeline(
self,
args,
vae,
text_encoder,
text_encoder_2,
model,
scheduler=None,
device=None,
progress_bar_config=None,
data_type="video",
):
"""Load the denoising scheduler for inference."""
if scheduler is None:
if args.denoise_type == "flow":
scheduler = FlowMatchDiscreteScheduler(
shift=args.flow_shift,
reverse=args.flow_reverse,
solver=args.flow_solver,
)
else:
raise ValueError(f"Invalid denoise type {args.denoise_type}")
pipeline = HunyuanVideoPipeline(
vae=vae,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
transformer=model,
scheduler=scheduler,
progress_bar_config=progress_bar_config,
args=args,
)
if self.use_cpu_offload:
pipeline.enable_sequential_cpu_offload()
else:
pipeline = pipeline.to(device)
return pipeline
@torch.no_grad()
def predict(
self,
prompt,
height=192,
width=336,
video_length=129,
seed=None,
negative_prompt=None,
infer_steps=50,
guidance_scale=6,
flow_shift=5.0,
embedded_guidance_scale=None,
batch_size=1,
num_videos_per_prompt=1,
**kwargs,
):
"""
Predict the image/video from the given text.
Args:
prompt (str or List[str]): The input text.
kwargs:
height (int): The height of the output video. Default is 192.
width (int): The width of the output video. Default is 336.
video_length (int): The frame number of the output video. Default is 129.
seed (int or List[str]): The random seed for the generation. Default is a random integer.
negative_prompt (str or List[str]): The negative text prompt. Default is an empty string.
guidance_scale (float): The guidance scale for the generation. Default is 6.0.
num_images_per_prompt (int): The number of images per prompt. Default is 1.
infer_steps (int): The number of inference steps. Default is 100.
"""
out_dict = dict()
# ========================================================================
# Arguments: seed
# ========================================================================
if isinstance(seed, torch.Tensor):
seed = seed.tolist()
if seed is None:
seeds = [
random.randint(0, 1_000_000)
for _ in range(batch_size * num_videos_per_prompt)
]
elif isinstance(seed, int):
seeds = [
seed + i
for _ in range(batch_size)
for i in range(num_videos_per_prompt)
]
elif isinstance(seed, (list, tuple)):
if len(seed) == batch_size:
seeds = [
int(seed[i]) + j
for i in range(batch_size)
for j in range(num_videos_per_prompt)
]
elif len(seed) == batch_size * num_videos_per_prompt:
seeds = [int(s) for s in seed]
else:
raise ValueError(
f"Length of seed must be equal to number of prompt(batch_size) or "
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}."
)
else:
raise ValueError(
f"Seed must be an integer, a list of integers, or None, got {seed}."
)
generator = [torch.Generator(self.device).manual_seed(seed) for seed in seeds]
out_dict["seeds"] = seeds
# ========================================================================
# Arguments: target_width, target_height, target_video_length
# ========================================================================
if width <= 0 or height <= 0 or video_length <= 0:
raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length-1` must be a multiple of 4, got {video_length}"
)
logger.info(
f"Input (height, width, video_length) = ({height}, {width}, {video_length})"
)
target_height = align_to(height, 16)
target_width = align_to(width, 16)
target_video_length = video_length
out_dict["size"] = (target_height, target_width, target_video_length)
# ========================================================================
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
prompt = [prompt.strip()]
# negative prompt
if negative_prompt is None or negative_prompt == "":
negative_prompt = self.default_negative_prompt
if not isinstance(negative_prompt, str):
raise TypeError(
f"`negative_prompt` must be a string, but got {type(negative_prompt)}"
)
negative_prompt = [negative_prompt.strip()]
# ========================================================================
# Scheduler
# ========================================================================
scheduler = FlowMatchDiscreteScheduler(
shift=flow_shift,
reverse=self.args.flow_reverse,
solver=self.args.flow_solver
)
self.pipeline.scheduler = scheduler
if "884" in self.args.vae:
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
elif "888" in self.args.vae:
latents_size = [(video_length - 1) // 8 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
# Print infer args
# ========================================================================
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {target_video_length}
prompt: {prompt}
neg_prompt: {negative_prompt}
seed: {seed}
infer_steps: {infer_steps}
num_videos_per_prompt: {num_videos_per_prompt}
guidance_scale: {guidance_scale}
n_tokens: {n_tokens}
flow_shift: {flow_shift}
embedded_guidance_scale: {embedded_guidance_scale}"""
logger.debug(debug_str)
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.time()
samples = self.pipeline(
prompt=prompt,
height=target_height,
width=target_width,
video_length=target_video_length,
num_inference_steps=infer_steps,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
generator=generator,
output_type="pil",
n_tokens=n_tokens,
embedded_guidance_scale=embedded_guidance_scale,
data_type="video" if target_video_length > 1 else "image",
is_progress_bar=True,
vae_ver=self.args.vae,
enable_tiling=self.args.vae_tiling,
)[0]
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.time() - start_time
logger.info(f"Success, time: {gen_time}")
return out_dict
@@ -0,0 +1,26 @@
from .models import HYVideoDiffusionTransformer, HUNYUAN_VIDEO_CONFIG
def load_model(args, in_channels, out_channels, factor_kwargs):
"""load hunyuan video model
Args:
args (dict): model args
in_channels (int): input channels number
out_channels (int): output channels number
factor_kwargs (dict): factor kwargs
Returns:
model (nn.Module): The hunyuan video model
"""
if args.model in HUNYUAN_VIDEO_CONFIG.keys():
model = HYVideoDiffusionTransformer(
in_channels=in_channels,
out_channels=out_channels,
**HUNYUAN_VIDEO_CONFIG[args.model],
**factor_kwargs,
)
return model
else:
raise NotImplementedError()
@@ -0,0 +1,23 @@
import torch.nn as nn
def get_activation_layer(act_type):
"""get activation layer
Args:
act_type (str): the activation type
Returns:
torch.nn.functional: the activation layer
"""
if act_type == "gelu":
return lambda: nn.GELU()
elif act_type == "gelu_tanh":
# Approximate `tanh` requires torch >= 1.13
return lambda: nn.GELU(approximate="tanh")
elif act_type == "relu":
return nn.ReLU
elif act_type == "silu":
return nn.SiLU
else:
raise ValueError(f"Unknown activation type: {act_type}")
@@ -0,0 +1,91 @@
import importlib.metadata
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
def attention(
q,
k,
v,
drop_rate=0,
attn_mask=None,
causal=False,
):
qkv = torch.stack([q, k, v], dim=2)
if attn_mask is not None and attn_mask.dtype != torch.bool:
attn_mask = attn_mask.bool()
x = flash_attn_no_pad(qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None)
b, s, a, d = x.shape
out = x.reshape(b, s, -1)
return out
def parallel_attention(
q,
k,
v,
img_q_len,
img_kv_len,
text_mask
):
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
query, encoder_query = q
key, encoder_key = k
value, encoder_value = v
if get_sequence_parallel_state():
# batch_size, seq_len, attn_heads, head_dim
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
encoder_query = shrink_head(encoder_query, dim=2)
encoder_key = shrink_head(encoder_key, dim=2)
encoder_value = shrink_head(encoder_value, dim=2)
# [b, s, h, d]
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# Hint: please check encoder_query.shape
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
value = torch.cat([value, encoder_value], dim=1)
# B, S, 3, H, D
qkv = torch.stack([query, key, value], dim=2)
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1
)
if get_sequence_parallel_state():
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
hidden_states = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
attn = torch.cat([hidden_states, encoder_hidden_states], dim=1)
b, s, a, d = attn.shape
attn = attn.reshape(b, s, -1)
return attn
@@ -0,0 +1,157 @@
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from ..utils.helpers import to_2tuple
class PatchEmbed(nn.Module):
"""2D Image to Patch Embedding
Image to Patch Embedding using Conv2d
A convolution based approach to patchifying a 2D image w/ embedding projection.
Based on the impl in https://github.com/google-research/vision_transformer
Hacked together by / Copyright 2020 Ross Wightman
Remove the _assert function in forward function to be compatible with multi-resolution images.
"""
def __init__(
self,
patch_size=16,
in_chans=3,
embed_dim=768,
norm_layer=None,
flatten=True,
bias=True,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
patch_size = to_2tuple(patch_size)
self.patch_size = patch_size
self.flatten = flatten
self.proj = nn.Conv3d(
in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
**factory_kwargs
)
nn.init.xavier_uniform_(self.proj.weight.view(self.proj.weight.size(0), -1))
if bias:
nn.init.zeros_(self.proj.bias)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x):
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = self.norm(x)
return x
class TextProjection(nn.Module):
"""
Projects text embeddings. Also handles dropout for classifier-free guidance.
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
"""
def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.linear_1 = nn.Linear(
in_features=in_channels,
out_features=hidden_size,
bias=True,
**factory_kwargs
)
self.act_1 = act_layer()
self.linear_2 = nn.Linear(
in_features=hidden_size,
out_features=hidden_size,
bias=True,
**factory_kwargs
)
def forward(self, caption):
hidden_states = self.linear_1(caption)
hidden_states = self.act_1(hidden_states)
hidden_states = self.linear_2(hidden_states)
return hidden_states
def timestep_embedding(t, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
Args:
t (torch.Tensor): a 1-D Tensor of N indices, one per batch element. These may be fractional.
dim (int): the dimension of the output.
max_period (int): controls the minimum frequency of the embeddings.
Returns:
embedding (torch.Tensor): An (N, D) Tensor of positional embeddings.
.. ref_link: 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
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(
self,
hidden_size,
act_layer,
frequency_embedding_size=256,
max_period=10000,
out_size=None,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
self.max_period = max_period
if out_size is None:
out_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(
frequency_embedding_size, hidden_size, bias=True, **factory_kwargs
),
act_layer(),
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
)
nn.init.normal_(self.mlp[0].weight, std=0.02)
nn.init.normal_(self.mlp[2].weight, std=0.02)
def forward(self, t):
t_freq = timestep_embedding(
t, self.frequency_embedding_size, self.max_period
).type(self.mlp[0].weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
@@ -0,0 +1,118 @@
# Modified from timm library:
# https://github.com/huggingface/pytorch-image-models/blob/648aaa41233ba83eb38faf5ba9d415d574823241/timm/layers/mlp.py#L13
from functools import partial
import torch
import torch.nn as nn
from .modulate_layers import modulate
from ..utils.helpers import to_2tuple
class MLP(nn.Module):
"""MLP as used in Vision Transformer, MLP-Mixer and related networks"""
def __init__(
self,
in_channels,
hidden_channels=None,
out_features=None,
act_layer=nn.GELU,
norm_layer=None,
bias=True,
drop=0.0,
use_conv=False,
device=None,
dtype=None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
out_features = out_features or in_channels
hidden_channels = hidden_channels or in_channels
bias = to_2tuple(bias)
drop_probs = to_2tuple(drop)
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
self.fc1 = linear_layer(
in_channels, hidden_channels, bias=bias[0], **factory_kwargs
)
self.act = act_layer()
self.drop1 = nn.Dropout(drop_probs[0])
self.norm = (
norm_layer(hidden_channels, **factory_kwargs)
if norm_layer is not None
else nn.Identity()
)
self.fc2 = linear_layer(
hidden_channels, out_features, bias=bias[1], **factory_kwargs
)
self.drop2 = nn.Dropout(drop_probs[1])
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop1(x)
x = self.norm(x)
x = self.fc2(x)
x = self.drop2(x)
return x
#
class MLPEmbedder(nn.Module):
"""copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py"""
def __init__(self, in_dim: int, hidden_dim: int, device=None, dtype=None):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True, **factory_kwargs)
self.silu = nn.SiLU()
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True, **factory_kwargs)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.out_layer(self.silu(self.in_layer(x)))
class FinalLayer(nn.Module):
"""The final layer of DiT."""
def __init__(
self, hidden_size, patch_size, out_channels, act_layer, device=None, dtype=None
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
# Just use LayerNorm for the final layer
self.norm_final = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
if isinstance(patch_size, int):
self.linear = nn.Linear(
hidden_size,
patch_size * patch_size * out_channels,
bias=True,
**factory_kwargs
)
else:
self.linear = nn.Linear(
hidden_size,
patch_size[0] * patch_size[1] * patch_size[2] * out_channels,
bias=True,
)
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
# Here we don't distinguish between the modulate types. Just use the simple one.
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
nn.init.zeros_(self.adaLN_modulation[1].bias)
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift=shift, scale=scale)
x = self.linear(x)
return x
+759
View File
@@ -0,0 +1,759 @@
from typing import Any, List, Tuple, Optional, Union, Dict
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from .activation_layers import get_activation_layer
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, PatchEmbed, TextProjection
from .attenion import parallel_attention
from .posemb_layers import apply_rotary_emb
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
from .modulate_layers import ModulateDiT, modulate, apply_gate
from .token_refiner import SingleTokenRefiner
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
from fastvideo.utils.parallel_states import (
nccl_info,
)
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal dit block with seperate modulation for
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
def __init__(
self,
hidden_size: int,
heads_num: int,
mlp_width_ratio: float,
mlp_act_type: str = "gelu_tanh",
qk_norm: bool = True,
qk_norm_type: str = "rms",
qkv_bias: bool = False,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.deterministic = False
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.img_mod = ModulateDiT(
hidden_size,
factor=6,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.img_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.img_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.img_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.img_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.img_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_mlp = MLP(
hidden_size,
mlp_hidden_dim,
act_layer=get_activation_layer(mlp_act_type),
bias=True,
**factory_kwargs,
)
self.txt_mod = ModulateDiT(
hidden_size,
factor=6,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.txt_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
self.txt_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.txt_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.txt_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.txt_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_mlp = MLP(
hidden_size,
mlp_hidden_dim,
act_layer=get_activation_layer(mlp_act_type),
bias=True,
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
def enable_deterministic(self):
self.deterministic = True
def disable_deterministic(self):
self.deterministic = False
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple = None,
text_mask: torch.Tensor = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(vec).chunk(6, dim=-1)
(
txt_mod1_shift,
txt_mod1_scale,
txt_mod1_gate,
txt_mod2_shift,
txt_mod2_scale,
txt_mod2_gate,
) = self.txt_mod(vec).chunk(6, dim=-1)
# Prepare image for attention.
img_modulated = self.img_norm1(img)
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply RoPE if needed.
if freqs_cis is not None:
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
freqs_cis = (
shrink_head(freqs_cis[0], dim=0),
shrink_head(freqs_cis[1], dim=0)
)
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask
)
# attention computation end
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
# Calculate the img bloks.
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
),
gate=txt_mod2_gate,
)
return img, txt
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers as described in
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
Also refer to (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
def __init__(
self,
hidden_size: int,
heads_num: int,
mlp_width_ratio: float = 4.0,
mlp_act_type: str = "gelu_tanh",
qk_norm: bool = True,
qk_norm_type: str = "rms",
qk_scale: float = None,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.deterministic = False
self.hidden_size = hidden_size
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
self.scale = qk_scale or head_dim ** -0.5
# qkv and mlp_in
self.linear1 = nn.Linear(
hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs
)
# proj and mlp_out
self.linear2 = nn.Linear(
hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.pre_norm = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.mlp_act = get_activation_layer(mlp_act_type)()
self.modulation = ModulateDiT(
hidden_size,
factor=3,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
def enable_deterministic(self):
self.deterministic = True
def disable_deterministic(self):
self.deterministic = False
def forward(
self,
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
text_mask: torch.Tensor = None,
) -> torch.Tensor:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
freqs_cis = (
shrink_head(freqs_cis[0], dim=0),
shrink_head(freqs_cis[1], dim=0)
)
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_v, txt_v = v[:, :-txt_len, :, :], v[:, -txt_len:, :, :]
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask
)
# attention computation end
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
"""
HunyuanVideo Transformer backbone
Inherited from ModelMixin and ConfigMixin for compatibility with diffusers' sampler StableDiffusionPipeline.
Reference:
[1] Flux.1: https://github.com/black-forest-labs/flux
[2] MMDiT: http://arxiv.org/abs/2403.03206
Parameters
----------
args: argparse.Namespace
The arguments parsed by argparse.
patch_size: list
The size of the patch.
in_channels: int
The number of input channels.
out_channels: int
The number of output channels.
hidden_size: int
The hidden size of the transformer backbone.
heads_num: int
The number of attention heads.
mlp_width_ratio: float
The ratio of the hidden size of the MLP in the transformer block.
mlp_act_type: str
The activation function of the MLP in the transformer block.
depth_double_blocks: int
The number of transformer blocks in the double blocks.
depth_single_blocks: int
The number of transformer blocks in the single blocks.
rope_dim_list: list
The dimension of the rotary embedding for t, h, w.
qkv_bias: bool
Whether to use bias in the qkv linear layer.
qk_norm: bool
Whether to use qk norm.
qk_norm_type: str
The type of qk norm.
guidance_embed: bool
Whether to use guidance embedding for distillation.
text_projection: str
The type of the text projection, default is single_refiner.
use_attention_mask: bool
Whether to use attention mask for text encoder.
dtype: torch.dtype
The dtype of the model.
device: torch.device
The device of the model.
"""
@register_to_config
def __init__(
self,
patch_size: list = [1, 2, 2],
in_channels: int = 4, # Should be VAE.config.latent_channels.
out_channels: int = None,
hidden_size: int = 3072,
heads_num: int = 24,
mlp_width_ratio: float = 4.0,
mlp_act_type: str = "gelu_tanh",
mm_double_blocks_depth: int = 20,
mm_single_blocks_depth: int = 40,
rope_dim_list: List[int] = [16, 56, 56],
qkv_bias: bool = True,
qk_norm: bool = True,
qk_norm_type: str = "rms",
guidance_embed: bool = False, # For modulation.
text_projection: str = "single_refiner",
use_attention_mask: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
text_states_dim: int = 4096,
text_states_dim_2: int = 768,
rope_theta:int = 256,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.patch_size = patch_size
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.unpatchify_channels = self.out_channels
self.guidance_embed = guidance_embed
self.rope_dim_list = rope_dim_list
self.rope_theta = rope_theta
# Text projection. Default to linear projection.
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
self.use_attention_mask = use_attention_mask
self.text_projection = text_projection
if hidden_size % heads_num != 0:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
)
pe_dim = hidden_size // heads_num
if sum(rope_dim_list) != pe_dim:
raise ValueError(
f"Got {rope_dim_list} but expected positional dim {pe_dim}"
)
self.hidden_size = hidden_size
self.heads_num = heads_num
# image projection
self.img_in = PatchEmbed(
self.patch_size, self.in_channels, self.hidden_size, **factory_kwargs
)
# text projection
if self.text_projection == "linear":
self.txt_in = TextProjection(
self.config.text_states_dim,
self.hidden_size,
get_activation_layer("silu"),
**factory_kwargs,
)
elif self.text_projection == "single_refiner":
self.txt_in = SingleTokenRefiner(
self.config.text_states_dim, hidden_size, heads_num, depth=2, **factory_kwargs
)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
# time modulation
self.time_in = TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
# text modulation
self.vector_in = MLPEmbedder(
self.config.text_states_dim_2, self.hidden_size, **factory_kwargs
)
# guidance modulation
self.guidance_in = (
TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
if guidance_embed
else None
)
# double blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(mm_double_blocks_depth)
]
)
# single blocks
self.single_blocks = nn.ModuleList(
[
MMSingleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
**factory_kwargs,
)
for _ in range(mm_single_blocks_depth)
]
)
self.final_layer = FinalLayer(
self.hidden_size,
self.patch_size,
self.out_channels,
get_activation_layer("silu"),
**factory_kwargs,
)
def enable_deterministic(self):
for block in self.double_blocks:
block.enable_deterministic()
for block in self.single_blocks:
block.enable_deterministic()
def disable_deterministic(self):
for block in self.double_blocks:
block.disable_deterministic()
for block in self.single_blocks:
block.disable_deterministic()
def get_rotary_pos_embed(self, rope_sizes):
target_ndim = 3
ndim = 5 - 2
head_dim = self.hidden_size // self.heads_num
rope_dim_list = self.rope_dim_list
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
theta=self.rope_theta,
use_real=True,
theta_rescale_factor=1,
)
return freqs_cos, freqs_sin
# x: torch.Tensor,
# t: torch.Tensor, # Should be in range(0, 1000).
# text_states: torch.Tensor = None,
# text_mask: torch.Tensor = None, # Now we don't use it.
# text_states_2: Optional[torch.Tensor] = None, # Text embedding for modulation.
# guidance: torch.Tensor = None, # Guidance for modulation, should be cfg_scale x 1000.
# return_dict: bool = True,
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
output_attn=False,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
guidance = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance == None:
guidance = torch.tensor([6016.], device=hidden_states.device, dtype=torch.bfloat16)
out = {}
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
_, _, ot, oh, ow = x.shape
tt, th, tw = (
ot // self.patch_size[0],
oh // self.patch_size[1],
ow // self.patch_size[2],
)
original_tt = nccl_info.sp_size * tt
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
# Prepare modulation vectors.
vec = self.time_in(t)
# text modulation
vec = vec + self.vector_in(text_states_2)
# guidance modulation
if self.guidance_embed:
if guidance is None:
raise ValueError(
"Didn't get guidance strength for guidance distilled model."
)
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
vec = vec + self.guidance_in(guidance)
# Embed image and text.
img = self.img_in(img)
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
for _, block in enumerate(self.double_blocks):
double_block_args = [
img,
txt,
vec,
freqs_cis,
text_mask
]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if len(self.single_blocks) > 0:
for _, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask
]
x = block(*single_block_args)
img = x[:, :img_seq_len, ...]
# ---------------------------- Final layer ------------------------------
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
img = self.unpatchify(img, tt, th, tw)
if return_dict:
out["x"] = img
return out
return (img, )
def unpatchify(self, x, t, h, w):
"""
x: (N, T, patch_size**2 * C)
imgs: (N, H, W, C)
"""
c = self.unpatchify_channels
pt, ph, pw = self.patch_size
assert t * h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
def params_count(self):
counts = {
"double": sum(
[
sum(p.numel() for p in block.img_attn_qkv.parameters())
+ sum(p.numel() for p in block.img_attn_proj.parameters())
+ sum(p.numel() for p in block.img_mlp.parameters())
+ sum(p.numel() for p in block.txt_attn_qkv.parameters())
+ sum(p.numel() for p in block.txt_attn_proj.parameters())
+ sum(p.numel() for p in block.txt_mlp.parameters())
for block in self.double_blocks
]
),
"single": sum(
[
sum(p.numel() for p in block.linear1.parameters())
+ sum(p.numel() for p in block.linear2.parameters())
for block in self.single_blocks
]
),
"total": sum(p.numel() for p in self.parameters()),
}
counts["attn+mlp"] = counts["double"] + counts["single"]
return counts
#################################################################################
# HunyuanVideo Configs #
#################################################################################
HUNYUAN_VIDEO_CONFIG = {
"HYVideo-T/2": {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
},
"HYVideo-T/2-cfgdistill": {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
"guidance_embed": True,
},
}
@@ -0,0 +1,155 @@
from typing import Callable
import torch
import torch.nn as nn
class ModulateDiT(nn.Module):
"""Modulation layer for DiT."""
def __init__(
self,
hidden_size: int,
factor: int,
act_layer: Callable,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.act = act_layer()
self.linear = nn.Linear(
hidden_size, factor * hidden_size, bias=True, **factory_kwargs
)
# Zero-initialize the modulation
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(self.act(x))
def modulate(x, shift=None, scale=None):
"""modulate by shift and scale
Args:
x (torch.Tensor): input tensor.
shift (torch.Tensor, optional): shift tensor. Defaults to None.
scale (torch.Tensor, optional): scale tensor. Defaults to None.
Returns:
torch.Tensor: the output tensor after modulate.
"""
if scale is None and shift is None:
return x
elif shift is None:
return x * (1 + scale.unsqueeze(1))
elif scale is None:
return x + shift.unsqueeze(1)
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def apply_gate(x, gate=None, tanh=False):
"""AI is creating summary for apply_gate
Args:
x (torch.Tensor): input tensor.
gate (torch.Tensor, optional): gate tensor. Defaults to None.
tanh (bool, optional): whether to use tanh function. Defaults to False.
Returns:
torch.Tensor: the output tensor after apply gate.
"""
if gate is None:
return x
if tanh:
return x * gate.unsqueeze(1).tanh()
else:
return x * gate.unsqueeze(1)
def ckpt_wrapper(module):
def ckpt_forward(*inputs):
outputs = module(*inputs)
return outputs
return ckpt_forward
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
def get_norm_layer(norm_layer):
"""
Get the normalization layer.
Args:
norm_layer (str): The type of normalization layer.
Returns:
norm_layer (nn.Module): The normalization layer.
"""
if norm_layer == "layer":
return nn.LayerNorm
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -0,0 +1,77 @@
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
def get_norm_layer(norm_layer):
"""
Get the normalization layer.
Args:
norm_layer (str): The type of normalization layer.
Returns:
norm_layer (nn.Module): The normalization layer.
"""
if norm_layer == "layer":
return nn.LayerNorm
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -0,0 +1,310 @@
import torch
from typing import Union, Tuple, List
def _to_tuple(x, dim=2):
if isinstance(x, int):
return (x,) * dim
elif len(x) == dim:
return x
else:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start, *args, dim=2):
"""
Get n-D meshgrid with start, stop and num.
Args:
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
n-tuples.
*args: See above.
dim (int): Dimension of the meshgrid. Defaults to 2.
Returns:
grid (np.ndarray): [dim, ...]
"""
if len(args) == 0:
# start is grid_size
num = _to_tuple(start, dim=dim)
start = (0,) * dim
stop = num
elif len(args) == 1:
# start is start, args[0] is stop, step is 1
start = _to_tuple(start, dim=dim)
stop = _to_tuple(args[0], dim=dim)
num = [stop[i] - start[i] for i in range(dim)]
elif len(args) == 2:
# start is start, args[0] is stop, args[1] is num
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
else:
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
axis_grid = []
for i in range(dim):
a, b, n = start[i], stop[i], num[i]
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
axis_grid.append(g)
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
return grid
#################################################################################
# Rotary Positional Embedding Functions #
#################################################################################
# https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L80
def reshape_for_broadcast(
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
x: torch.Tensor,
head_first=False,
):
"""
Reshape frequency tensor for broadcasting it with another tensor.
This function reshapes the frequency tensor to have the same shape as the target tensor 'x'
for the purpose of broadcasting the frequency tensor during element-wise operations.
Notes:
When using FlashMHAModified, head_first should be False.
When using Attention, head_first should be True.
Args:
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.
x (torch.Tensor): Target tensor for broadcasting compatibility.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
torch.Tensor: Reshaped frequency tensor.
Raises:
AssertionError: If the frequency tensor doesn't match the expected shape.
AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.
"""
ndim = x.ndim
assert 0 <= 1 < ndim
if isinstance(freqs_cis, tuple):
# freqs_cis: (cos, sin) in real space
if head_first:
assert freqs_cis[0].shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis[0].shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
else:
# freqs_cis: values in complex space
if head_first:
assert freqs_cis.shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis.shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis.view(*shape)
def rotate_half(x):
x_real, x_imag = (
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
) # [B, S, H, D//2]
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
def apply_rotary_emb(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
head_first: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor.
This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided
frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor
is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are
returned as real tensors.
Args:
xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]
xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D]
freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
xk_out = None
if isinstance(freqs_cis, tuple):
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
cos, sin = cos.to(xq.device), sin.to(xq.device)
# real * cos - imag * sin
# imag * cos + real * sin
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
else:
# view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)
xq_ = torch.view_as_complex(
xq.float().reshape(*xq.shape[:-1], -1, 2)
) # [B, S, H, D//2]
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(
xq.device
) # [S, D//2] --> [1, S, 1, D//2]
# (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)
# view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
xk_ = torch.view_as_complex(
xk.float().reshape(*xk.shape[:-1], -1, 2)
) # [B, S, H, D//2]
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
return xq_out, xk_out
def get_nd_rotary_pos_embed(
rope_dim_list,
start,
*args,
theta=10000.0,
use_real=False,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
):
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
Args:
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
sum(rope_dim_list) should equal to head_dim of attention layer.
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
*args: See above.
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.
Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real
part and an imaginary part separately.
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
Returns:
pos_embed (torch.Tensor): [HW, D/2]
"""
grid = get_meshgrid_nd(
start, *args, dim=len(rope_dim_list)
) # [3, W, H, D] / [2, W, H]
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list
), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(
rope_dim_list[i],
grid[i].reshape(-1),
theta,
use_real=use_real,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
if use_real:
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
else:
emb = torch.cat(embs, dim=1) # (WHD, D/2)
return emb
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
theta: float = 10000.0,
use_real: bool = False,
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
and the end index 'end'. The 'theta' parameter scales the frequencies.
The returned tensor contains complex values in complex64 data type.
Args:
dim (int): Dimension of the frequency tensor.
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool, optional): If True, return real part and imaginary part separately.
Otherwise, return complex numbers.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
Returns:
freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor ** (dim / (dim - 2))
freqs = 1.0 / (
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
) # [D/2]
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
if use_real:
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(
torch.ones_like(freqs), freqs
) # complex64 # [S, D/2]
return freqs_cis
@@ -0,0 +1,223 @@
from typing import Optional
from einops import rearrange
import torch
import torch.nn as nn
from .activation_layers import get_activation_layer
from .attenion import attention
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, TextProjection
from .attenion import attention
from .mlp_layers import MLP
from .modulate_layers import modulate, apply_gate
class IndividualTokenRefinerBlock(nn.Module):
def __init__(
self,
hidden_size,
heads_num,
mlp_width_ratio: str = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
)
self.self_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.self_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.self_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.self_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
)
act_layer = get_activation_layer(act_type)
self.mlp = MLP(
in_channels=hidden_size,
hidden_channels=mlp_hidden_dim,
act_layer=act_layer,
drop=mlp_drop_rate,
**factory_kwargs,
)
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
nn.init.zeros_(self.adaLN_modulation[1].bias)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
attn_mask: torch.Tensor = None,
):
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
norm_x = self.norm1(x)
qkv = self.self_attn_qkv(norm_x)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed
q = self.self_attn_q_norm(q).to(v)
k = self.self_attn_k_norm(k).to(v)
# Self-Attention
attn = attention(q, k, v, attn_mask=attn_mask)
x = x + apply_gate(self.self_attn_proj(attn), gate_msa)
# FFN Layer
x = x + apply_gate(self.mlp(self.norm2(x)), gate_mlp)
return x
class IndividualTokenRefiner(nn.Module):
def __init__(
self,
hidden_size,
heads_num,
depth,
mlp_width_ratio: float = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.blocks = nn.ModuleList(
[
IndividualTokenRefinerBlock(
hidden_size=hidden_size,
heads_num=heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(depth)
]
)
def forward(
self,
x: torch.Tensor,
c: torch.LongTensor,
mask: Optional[torch.Tensor] = None,
):
mask = mask.clone().bool()
# avoid attention weight become NaN
mask[:, 0] = True
for block in self.blocks:
x = block(x, c, mask)
return x
class SingleTokenRefiner(nn.Module):
"""
A single token refiner block for llm text embedding refine.
"""
def __init__(
self,
in_channels,
hidden_size,
heads_num,
depth,
mlp_width_ratio: float = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
attn_mode: str = "torch",
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.attn_mode = attn_mode
assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."
self.input_embedder = nn.Linear(
in_channels, hidden_size, bias=True, **factory_kwargs
)
act_layer = get_activation_layer(act_type)
# Build timestep embedding layer
self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)
# Build context embedding layer
self.c_embedder = TextProjection(
in_channels, hidden_size, act_layer, **factory_kwargs
)
self.individual_token_refiner = IndividualTokenRefiner(
hidden_size=hidden_size,
heads_num=heads_num,
depth=depth,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
def forward(
self,
x: torch.Tensor,
t: torch.LongTensor,
mask: Optional[torch.LongTensor] = None,
):
timestep_aware_representations = self.t_embedder(t)
if mask is None:
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [b, s1, 1]
context_aware_representations = (x * mask_float).sum(
dim=1
) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(context_aware_representations)
c = timestep_aware_representations + context_aware_representations
x = self.input_embedder(x)
x = self.individual_token_refiner(x, c, mask)
return x
@@ -0,0 +1,51 @@
normal_mode_prompt = """Normal mode - Video Recaption Task:
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
0. Preserve ALL information, including style words and technical terms.
1. If the input is in Chinese, translate the entire description to English.
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
4. Output ALL must be in English.
Given Input:
input: "{input}"
"""
master_mode_prompt = """Master mode - Video Recaption Task:
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
0. Preserve ALL information, including style words and technical terms.
1. If the input is in Chinese, translate the entire description to English.
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
4. Output ALL must be in English.
Given Input:
input: "{input}"
"""
def get_rewrite_prompt(ori_prompt, mode="Normal"):
if mode == "Normal":
prompt = normal_mode_prompt.format(input=ori_prompt)
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Normal", mode)
return prompt
ori_prompt = "一只小狗在草地上奔跑。"
normal_prompt = get_rewrite_prompt(ori_prompt, mode="Normal")
master_prompt = get_rewrite_prompt(ori_prompt, mode="Master")
# Then you can use the normal_prompt or master_prompt to access the hunyuan-large rewrite model to get the final prompt.
@@ -0,0 +1,357 @@
from dataclasses import dataclass
from typing import Optional, Tuple
from copy import deepcopy
import torch
import torch.nn as nn
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel
from transformers.utils import ModelOutput
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
from ..constants import PRECISION_TO_TYPE
def use_default(value, default):
return value if value is not None else default
def load_text_encoder(
text_encoder_type,
text_encoder_precision=None,
text_encoder_path=None,
logger=None,
device=None,
):
if text_encoder_path is None:
text_encoder_path = TEXT_ENCODER_PATH[text_encoder_type]
if logger is not None:
logger.info(
f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}"
)
if text_encoder_type == "clipL":
text_encoder = CLIPTextModel.from_pretrained(text_encoder_path)
text_encoder.final_layer_norm = text_encoder.text_model.final_layer_norm
elif text_encoder_type == "llm":
text_encoder = AutoModel.from_pretrained(
text_encoder_path, low_cpu_mem_usage=True
)
text_encoder.final_layer_norm = text_encoder.norm
else:
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
# from_pretrained will ensure that the model is in eval mode.
if text_encoder_precision is not None:
text_encoder = text_encoder.to(dtype=PRECISION_TO_TYPE[text_encoder_precision])
text_encoder.requires_grad_(False)
if logger is not None:
logger.info(f"Text encoder to dtype: {text_encoder.dtype}")
if device is not None:
text_encoder = text_encoder.to(device)
return text_encoder, text_encoder_path
def load_tokenizer(
tokenizer_type, tokenizer_path=None, padding_side="right", logger=None
):
if tokenizer_path is None:
tokenizer_path = TOKENIZER_PATH[tokenizer_type]
if logger is not None:
logger.info(f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
if tokenizer_type == "clipL":
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path, max_length=77)
elif tokenizer_type == "llm":
tokenizer = AutoTokenizer.from_pretrained(
tokenizer_path, padding_side=padding_side
)
else:
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
return tokenizer, tokenizer_path
@dataclass
class TextEncoderModelOutput(ModelOutput):
"""
Base class for model's outputs that also contains a pooling of the last hidden states.
Args:
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
List of decoded texts.
"""
hidden_state: torch.FloatTensor = None
attention_mask: Optional[torch.LongTensor] = None
hidden_states_list: Optional[Tuple[torch.FloatTensor, ...]] = None
text_outputs: Optional[list] = None
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder_type: str,
max_length: int,
text_encoder_precision: Optional[str] = None,
text_encoder_path: Optional[str] = None,
tokenizer_type: Optional[str] = None,
tokenizer_path: Optional[str] = None,
output_key: Optional[str] = None,
use_attention_mask: bool = True,
input_max_length: Optional[int] = None,
prompt_template: Optional[dict] = None,
prompt_template_video: Optional[dict] = None,
hidden_state_skip_layer: Optional[int] = None,
apply_final_norm: bool = False,
reproduce: bool = False,
logger=None,
device=None,
):
super().__init__()
self.text_encoder_type = text_encoder_type
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.tokenizer_type = (
tokenizer_type if tokenizer_type is not None else text_encoder_type
)
self.tokenizer_path = (
tokenizer_path if tokenizer_path is not None else text_encoder_path
)
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (
use_attention_mask is True
), "Attention mask is True required when training videos."
self.input_max_length = (
input_max_length if input_max_length is not None else max_length
)
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
self.apply_final_norm = apply_final_norm
self.reproduce = reproduce
self.logger = logger
self.use_template = self.prompt_template is not None
if self.use_template:
assert (
isinstance(self.prompt_template, dict)
and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert "{}" in str(self.prompt_template["template"]), (
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template['template']}"
)
self.use_video_template = self.prompt_template_video is not None
if self.use_video_template:
if self.prompt_template_video is not None:
assert (
isinstance(self.prompt_template_video, dict)
and "template" in self.prompt_template_video
), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}"
assert "{}" in str(self.prompt_template_video["template"]), (
"`prompt_template_video['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template_video['template']}"
)
if "t5" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
elif "clip" in text_encoder_type:
self.output_key = output_key or "pooler_output"
elif "llm" in text_encoder_type or "glm" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
self.model, self.model_path = load_text_encoder(
text_encoder_type=self.text_encoder_type,
text_encoder_precision=self.precision,
text_encoder_path=self.model_path,
logger=self.logger,
device=device,
)
self.dtype = self.model.dtype
self.device = self.model.device
self.tokenizer, self.tokenizer_path = load_tokenizer(
tokenizer_type=self.tokenizer_type,
tokenizer_path=self.tokenizer_path,
padding_side="right",
logger=self.logger,
)
def __repr__(self):
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
@staticmethod
def apply_text_to_template(text, template, prevent_empty_text=True):
"""
Apply text to template.
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If Ture, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
# Will send string to tokenizer. Used for llm
return template.format(text)
else:
raise TypeError(f"Unsupported template type: {type(template)}")
def text2tokens(self, text, data_type="image"):
"""
Tokenize the input text.
Args:
text (str or list): Input text.
"""
tokenize_input_type = "str"
if self.use_template:
if data_type == "image":
prompt_template = self.prompt_template["template"]
elif data_type == "video":
prompt_template = self.prompt_template_video["template"]
else:
raise ValueError(f"Unsupported data type: {data_type}")
if isinstance(text, (list, tuple)):
text = [
self.apply_text_to_template(one_text, prompt_template)
for one_text in text
]
if isinstance(text[0], list):
tokenize_input_type = "list"
elif isinstance(text, str):
text = self.apply_text_to_template(text, prompt_template)
if isinstance(text, list):
tokenize_input_type = "list"
else:
raise TypeError(f"Unsupported text type: {type(text)}")
kwargs = dict(
truncation=True,
max_length=self.max_length,
padding="max_length",
return_tensors="pt",
)
if tokenize_input_type == "str":
return self.tokenizer(
text,
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
**kwargs,
)
elif tokenize_input_type == "list":
return self.tokenizer.apply_chat_template(
text,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
**kwargs,
)
else:
raise ValueError(f"Unsupported tokenize_input_type: {tokenize_input_type}")
def encode(
self,
batch_encoding,
use_attention_mask=None,
output_hidden_states=False,
do_sample=None,
hidden_state_skip_layer=None,
return_texts=False,
data_type="image",
device=None,
):
"""
Args:
batch_encoding (dict): Batch encoding from tokenizer.
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
Defaults to None.
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
output_hidden_states will be set True. Defaults to False.
do_sample (bool): Whether to sample from the model. Used for Decoder-Only LLMs. Defaults to None.
When self.produce is False, do_sample is set to True by default.
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
If None, self.output_key will be used. Defaults to None.
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
hidden_state_skip_layer = use_default(
hidden_state_skip_layer, self.hidden_state_skip_layer
)
do_sample = use_default(do_sample, not self.reproduce)
attention_mask = (
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
)
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
attention_mask=attention_mask,
output_hidden_states=output_hidden_states
or hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
# Remove hidden states of instruction tokens, only keep prompt tokens.
if self.use_template:
if data_type == "image":
crop_start = self.prompt_template.get("crop_start", -1)
elif data_type == "video":
crop_start = self.prompt_template_video.get("crop_start", -1)
else:
raise ValueError(f"Unsupported data type: {data_type}")
if crop_start > 0:
last_hidden_state = last_hidden_state[:, crop_start:]
attention_mask = (
attention_mask[:, crop_start:] if use_attention_mask else None
)
if output_hidden_states:
return TextEncoderModelOutput(
last_hidden_state, attention_mask, outputs.hidden_states
)
return TextEncoderModelOutput(last_hidden_state, attention_mask)
def forward(
self,
text,
use_attention_mask=None,
output_hidden_states=False,
do_sample=False,
hidden_state_skip_layer=None,
return_texts=False,
):
batch_encoding = self.text2tokens(text)
return self.encode(
batch_encoding,
use_attention_mask=use_attention_mask,
output_hidden_states=output_hidden_states,
do_sample=do_sample,
hidden_state_skip_layer=hidden_state_skip_layer,
return_texts=return_texts,
)
@@ -0,0 +1,15 @@
import numpy as np
import math
def align_to(value, alignment):
"""align hight, width according to alignment
Args:
value (int): height or width
alignment (int): target alignment factor
Returns:
int: the aligned value
"""
return int(math.ceil(value / alignment) * alignment)
@@ -0,0 +1,70 @@
import os
from pathlib import Path
from einops import rearrange
import torch
import torchvision
import numpy as np
import imageio
CODE_SUFFIXES = {
".py", # Python codes
".sh", # Shell scripts
".yaml",
".yml", # Configuration files
}
def safe_dir(path):
"""
Create a directory (or the parent directory of a file) if it does not exist.
Args:
path (str or Path): Path to the directory.
Returns:
path (Path): Path object of the directory.
"""
path = Path(path)
path.mkdir(exist_ok=True, parents=True)
return path
def safe_file(path):
"""
Create the parent directory of a file if it does not exist.
Args:
path (str or Path): Path to the file.
Returns:
path (Path): Path object of the file.
"""
path = Path(path)
path.parent.mkdir(exist_ok=True, parents=True)
return path
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24):
"""save videos by video tensor
copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61
Args:
videos (torch.Tensor): video tensor predicted by the model
path (str): path to save video
rescale (bool, optional): rescale the video tensor from [-1, 1] to . Defaults to False.
n_rows (int, optional): Defaults to 1.
fps (int, optional): video save fps. Defaults to 8.
"""
videos = rearrange(videos, "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=n_rows)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
if rescale:
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
x = torch.clamp(x, 0, 1)
x = (x * 255).numpy().astype(np.uint8)
outputs.append(x)
os.makedirs(os.path.dirname(path), exist_ok=True)
imageio.mimsave(path, outputs, fps=fps)
+40
View File
@@ -0,0 +1,40 @@
import collections.abc
from itertools import repeat
def _ntuple(n):
def parse(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
x = tuple(x)
if len(x) == 1:
x = tuple(repeat(x[0], n))
return x
return tuple(repeat(x, n))
return parse
to_1tuple = _ntuple(1)
to_2tuple = _ntuple(2)
to_3tuple = _ntuple(3)
to_4tuple = _ntuple(4)
def as_tuple(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
return tuple(x)
if x is None or isinstance(x, (int, float, str)):
return (x,)
else:
raise ValueError(f"Unknown type {type(x)}")
def as_list_of_2tuple(x):
x = as_tuple(x)
if len(x) == 1:
x = (x[0], x[0])
assert len(x) % 2 == 0, f"Expect even length, got {len(x)}."
lst = []
for i in range(0, len(x), 2):
lst.append((x[i], x[i + 1]))
return lst
@@ -0,0 +1,46 @@
import argparse
import torch
from transformers import (
AutoProcessor,
LlavaForConditionalGeneration,
)
def preprocess_text_encoder_tokenizer(args):
processor = AutoProcessor.from_pretrained(args.input_dir)
model = LlavaForConditionalGeneration.from_pretrained(
args.input_dir,
torch_dtype=torch.float16,
low_cpu_mem_usage=True,
).to(0)
model.language_model.save_pretrained(
f"{args.output_dir}"
)
processor.tokenizer.save_pretrained(
f"{args.output_dir}"
)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--input_dir",
type=str,
required=True,
help="The path to the llava-llama-3-8b-v1_1-transformers.",
)
parser.add_argument(
"--output_dir",
type=str,
default="",
help="The output path of the llava-llama-3-8b-text-encoder-tokenizer."
"if '', the parent dir of output will be the same as input dir.",
)
args = parser.parse_args()
if len(args.output_dir) == 0:
args.output_dir = "/".join(args.input_dir.split("/")[:-1])
preprocess_text_encoder_tokenizer(args)
+62
View File
@@ -0,0 +1,62 @@
from pathlib import Path
import torch
from .autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from ..constants import VAE_PATH, PRECISION_TO_TYPE
def load_vae(vae_type: str="884-16c-hy",
vae_precision: str=None,
sample_size: tuple=None,
vae_path: str=None,
logger=None,
device=None
):
"""the fucntion to load the 3D VAE model
Args:
vae_type (str): the type of the 3D VAE model. Defaults to "884-16c-hy".
vae_precision (str, optional): the precision to load vae. Defaults to None.
sample_size (tuple, optional): the tiling size. Defaults to None.
vae_path (str, optional): the path to vae. Defaults to None.
logger (_type_, optional): logger. Defaults to None.
device (_type_, optional): device to load vae. Defaults to None.
"""
if vae_path is None:
vae_path = VAE_PATH[vae_type]
if logger is not None:
logger.info(f"Loading 3D VAE model ({vae_type}) from: {vae_path}")
config = AutoencoderKLCausal3D.load_config(vae_path)
if sample_size:
vae = AutoencoderKLCausal3D.from_config(config, sample_size=sample_size)
else:
vae = AutoencoderKLCausal3D.from_config(config)
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
ckpt = torch.load(vae_ckpt, map_location=vae.device)
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
vae.load_state_dict(ckpt)
spatial_compression_ratio = vae.config.spatial_compression_ratio
time_compression_ratio = vae.config.time_compression_ratio
if vae_precision is not None:
vae = vae.to(dtype=PRECISION_TO_TYPE[vae_precision])
vae.requires_grad_(False)
if logger is not None:
logger.info(f"VAE to dtype: {vae.dtype}")
if device is not None:
vae = vae.to(device)
vae.eval()
return vae, vae_path, spatial_compression_ratio, time_compression_ratio
@@ -0,0 +1,603 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Dict, Optional, Tuple, Union
from dataclasses import dataclass
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
try:
# This diffusers is modified and packed in the mirror.
from diffusers.loaders import FromOriginalVAEMixin
except ImportError:
# Use this to be compatible with the original diffusers.
from diffusers.loaders.single_file_model import FromOriginalModelMixin as FromOriginalVAEMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
from diffusers.models.attention_processor import (
ADDED_KV_ATTENTION_PROCESSORS,
CROSS_ATTENTION_PROCESSORS,
Attention,
AttentionProcessor,
AttnAddedKVProcessor,
AttnProcessor,
)
from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from .vae import DecoderCausal3D, BaseOutput, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D
@dataclass
class DecoderOutput2(BaseOutput):
sample: torch.FloatTensor
posterior: Optional[DiagonalGaussianDistribution] = None
class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
r"""
A VAE model with KL loss for encoding images/videos into latents and decoding latent representations into images/videos.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
"""
_supports_gradient_checkpointing = True
@register_to_config
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str] = ("DownEncoderBlockCausal3D",),
up_block_types: Tuple[str] = ("UpDecoderBlockCausal3D",),
block_out_channels: Tuple[int] = (64,),
layers_per_block: int = 1,
act_fn: str = "silu",
latent_channels: int = 4,
norm_num_groups: int = 32,
sample_size: int = 32,
sample_tsize: int = 64,
scaling_factor: float = 0.18215,
force_upcast: float = True,
spatial_compression_ratio: int = 8,
time_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
):
super().__init__()
self.time_compression_ratio = time_compression_ratio
self.encoder = EncoderCausal3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
double_z=True,
time_compression_ratio=time_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
)
self.decoder = DecoderCausal3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=time_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
)
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
self.use_slicing = False
self.use_spatial_tiling = False
self.use_temporal_tiling = False
# only relevant if vae tiling is enabled
self.tile_sample_min_tsize = sample_tsize
self.tile_latent_min_tsize = sample_tsize // time_compression_ratio
self.tile_sample_min_size = self.config.sample_size
sample_size = (
self.config.sample_size[0]
if isinstance(self.config.sample_size, (list, tuple))
else self.config.sample_size
)
self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1)))
self.tile_overlap_factor = 0.25
def _set_gradient_checkpointing(self, module, value=False):
if isinstance(module, (EncoderCausal3D, DecoderCausal3D)):
module.gradient_checkpointing = value
def enable_temporal_tiling(self, use_tiling: bool = True):
self.use_temporal_tiling = use_tiling
def disable_temporal_tiling(self):
self.enable_temporal_tiling(False)
def enable_spatial_tiling(self, use_tiling: bool = True):
self.use_spatial_tiling = use_tiling
def disable_spatial_tiling(self):
self.enable_spatial_tiling(False)
def enable_tiling(self, use_tiling: bool = True):
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger videos.
"""
self.enable_spatial_tiling(use_tiling)
self.enable_temporal_tiling(use_tiling)
def disable_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.disable_spatial_tiling()
self.disable_temporal_tiling()
def enable_slicing(self):
r"""
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
"""
self.use_slicing = True
def disable_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.use_slicing = False
@property
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.attn_processors
def attn_processors(self) -> Dict[str, AttentionProcessor]:
r"""
Returns:
`dict` of attention processors: A dictionary containing all attention processors used in the model with
indexed by its weight name.
"""
# set recursively
processors = {}
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
return processors
for name, module in self.named_children():
fn_recursive_add_processors(name, module, processors)
return processors
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(
self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]], _remove_lora=False
):
r"""
Sets the attention processor to use to compute attention.
Parameters:
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
The instantiated processor class or a dictionary of processor classes that will be set as the processor
for **all** `Attention` layers.
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
processor. This is strongly recommended when setting trainable attention processors.
"""
count = len(self.attn_processors.keys())
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor, _remove_lora=_remove_lora)
else:
module.set_processor(processor.pop(f"{name}.processor"), _remove_lora=_remove_lora)
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
def set_default_attn_processor(self):
"""
Disables custom attention processors and sets the default attention implementation.
"""
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
processor = AttnAddedKVProcessor()
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
processor = AttnProcessor()
else:
raise ValueError(
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
)
self.set_attn_processor(processor, _remove_lora=True)
@apply_forward_hook
def encode(
self, x: torch.FloatTensor, return_dict: bool = True
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
"""
Encode a batch of images/videos into latents.
Args:
x (`torch.FloatTensor`): Input batch of images/videos.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
The latent representations of the encoded images/videos. If `return_dict` is True, a
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
"""
assert len(x.shape) == 5, "The input tensor should have 5 dimensions."
if self.use_temporal_tiling and x.shape[2] > self.tile_sample_min_tsize:
return self.temporal_tiled_encode(x, return_dict=return_dict)
if self.use_spatial_tiling and (x.shape[-1] > self.tile_sample_min_size or x.shape[-2] > self.tile_sample_min_size):
return self.spatial_tiled_encode(x, return_dict=return_dict)
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self.encoder(x_slice) for x_slice in x.split(1)]
h = torch.cat(encoded_slices)
else:
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
assert len(z.shape) == 5, "The input tensor should have 5 dimensions."
if self.use_temporal_tiling and z.shape[2] > self.tile_latent_min_tsize:
return self.temporal_tiled_decode(z, return_dict=return_dict)
if self.use_spatial_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size):
return self.spatial_tiled_decode(z, return_dict=return_dict)
z = self.post_quant_conv(z)
dec = self.decoder(z)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
@apply_forward_hook
def decode(
self, z: torch.FloatTensor, return_dict: bool = True, generator=None
) -> Union[DecoderOutput, torch.FloatTensor]:
"""
Decode a batch of images/videos.
Args:
z (`torch.FloatTensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z).sample
if not return_dict:
return (decoded,)
return DecoderOutput(sample=decoded)
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
for y in range(blend_extent):
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent)
return b
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
for x in range(blend_extent):
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent)
return b
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
for x in range(blend_extent):
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * (x / blend_extent)
return b
def spatial_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True, return_moments: bool = False) -> AutoencoderKLOutput:
r"""Encode a batch of images/videos using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image/videos size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.FloatTensor`): Input batch of images/videos.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split video into tiles and encode them separately.
rows = []
for i in range(0, x.shape[-2], overlap_size):
row = []
for j in range(0, x.shape[-1], overlap_size):
tile = x[:, :, :, i: i + self.tile_sample_min_size, j: j + self.tile_sample_min_size]
tile = self.encoder(tile)
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=-1))
moments = torch.cat(result_rows, dim=-2)
if return_moments:
return moments
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def spatial_tiled_decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
r"""
Decode a batch of images/videos using a tiled decoder.
Args:
z (`torch.FloatTensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping tiles and decode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, z.shape[-2], overlap_size):
row = []
for j in range(0, z.shape[-1], overlap_size):
tile = z[:, :, :, i: i + self.tile_latent_min_size, j: j + self.tile_latent_min_size]
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
row.append(decoded)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=-1))
dec = torch.cat(result_rows, dim=-2)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
def temporal_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
B, C, T, H, W = x.shape
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_latent_min_tsize - blend_extent
# Split the video into tiles and encode them separately.
row = []
for i in range(0, T, overlap_size):
tile = x[:, :, i: i + self.tile_sample_min_tsize + 1, :, :]
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_sample_min_size or tile.shape[-2] > self.tile_sample_min_size):
tile = self.spatial_tiled_encode(tile, return_moments=True)
else:
tile = self.encoder(tile)
tile = self.quant_conv(tile)
if i > 0:
tile = tile[:, :, 1:, :, :]
row.append(tile)
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_extent)
result_row.append(tile[:, :, :t_limit, :, :])
else:
result_row.append(tile[:, :, :t_limit + 1, :, :])
moments = torch.cat(result_row, dim=2)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def temporal_tiled_decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
# Split z into overlapping tiles and decode them separately.
B, C, T, H, W = z.shape
overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_sample_min_tsize - blend_extent
row = []
for i in range(0, T, overlap_size):
tile = z[:, :, i: i + self.tile_latent_min_tsize + 1, :, :]
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_latent_min_size or tile.shape[-2] > self.tile_latent_min_size):
decoded = self.spatial_tiled_decode(tile, return_dict=True).sample
else:
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
if i > 0:
decoded = decoded[:, :, 1:, :, :]
row.append(decoded)
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_extent)
result_row.append(tile[:, :, :t_limit, :, :])
else:
result_row.append(tile[:, :, :t_limit + 1, :, :])
dec = torch.cat(result_row, dim=2)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
def forward(
self,
sample: torch.FloatTensor,
sample_posterior: bool = False,
return_dict: bool = True,
return_posterior: bool = False,
generator: Optional[torch.Generator] = None,
) -> Union[DecoderOutput2, torch.FloatTensor]:
r"""
Args:
sample (`torch.FloatTensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z).sample
if not return_dict:
if return_posterior:
return (dec, posterior)
else:
return (dec,)
if return_posterior:
return DecoderOutput2(sample=dec, posterior=posterior)
else:
return DecoderOutput2(sample=dec)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections
def fuse_qkv_projections(self):
"""
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query,
key, value) are fused. For cross-attention modules, key and value projection matrices are fused.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
self.original_attn_processors = None
for _, attn_processor in self.attn_processors.items():
if "Added" in str(attn_processor.__class__.__name__):
raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.")
self.original_attn_processors = self.attn_processors
for module in self.modules():
if isinstance(module, Attention):
module.fuse_projections(fuse=True)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
def unfuse_qkv_projections(self):
"""Disables the fused QKV projection if enabled.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
if self.original_attn_processors is not None:
self.set_attn_processor(self.original_attn_processors)
@@ -0,0 +1,764 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional, Tuple, Union
import torch
import torch.nn.functional as F
from torch import nn
from einops import rearrange
from diffusers.utils import logging
from diffusers.models.activations import get_activation
from diffusers.models.attention_processor import SpatialNorm
from diffusers.models.attention_processor import Attention
from diffusers.models.normalization import AdaGroupNorm
from diffusers.models.normalization import RMSNorm
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
"""
Implements a causal 3D convolution layer where each position only depends on previous timesteps and current spatial locations.
This maintains temporal causality in video generation tasks.
"""
def __init__(
self,
chan_in,
chan_out,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]] = 1,
dilation: Union[int, Tuple[int, int, int]] = 1,
pad_mode='replicate',
**kwargs
):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
self.time_causal_padding = padding
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class UpsampleCausal3D(nn.Module):
"""
A 3D upsampling layer with an optional convolution.
"""
def __init__(
self,
channels: int,
use_conv: bool = False,
use_conv_transpose: bool = False,
out_channels: Optional[int] = None,
name: str = "conv",
kernel_size: Optional[int] = None,
padding=1,
norm_type=None,
eps=None,
elementwise_affine=None,
bias=True,
interpolate=True,
upsample_factor=(2, 2, 2),
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.use_conv_transpose = use_conv_transpose
self.name = name
self.interpolate = interpolate
self.upsample_factor = upsample_factor
if norm_type == "ln_norm":
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
elif norm_type == "rms_norm":
self.norm = RMSNorm(channels, eps, elementwise_affine)
elif norm_type is None:
self.norm = None
else:
raise ValueError(f"unknown norm_type: {norm_type}")
conv = None
if use_conv_transpose:
raise NotImplementedError
elif use_conv:
if kernel_size is None:
kernel_size = 3
conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, bias=bias)
if name == "conv":
self.conv = conv
else:
self.Conv2d_0 = conv
def forward(
self,
hidden_states: torch.FloatTensor,
output_size: Optional[int] = None,
scale: float = 1.0,
) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
if self.norm is not None:
raise NotImplementedError
if self.use_conv_transpose:
return self.conv(hidden_states)
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
dtype = hidden_states.dtype
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(torch.float32)
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
if hidden_states.shape[0] >= 64:
hidden_states = hidden_states.contiguous()
# if `output_size` is passed we force the interpolation output
# size and do not make use of `scale_factor=2`
if self.interpolate:
B, C, T, H, W = hidden_states.shape
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
if output_size is None:
if T > 1:
other_h = F.interpolate(other_h, scale_factor=self.upsample_factor, mode="nearest")
first_h = first_h.squeeze(2)
first_h = F.interpolate(first_h, scale_factor=self.upsample_factor[1:], mode="nearest")
first_h = first_h.unsqueeze(2)
else:
raise NotImplementedError
if T > 1:
hidden_states = torch.cat((first_h, other_h), dim=2)
else:
hidden_states = first_h
# If the input is bfloat16, we cast back to bfloat16
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(dtype)
if self.use_conv:
if self.name == "conv":
hidden_states = self.conv(hidden_states)
else:
hidden_states = self.Conv2d_0(hidden_states)
return hidden_states
class DownsampleCausal3D(nn.Module):
"""
A 3D downsampling layer with an optional convolution.
"""
def __init__(
self,
channels: int,
use_conv: bool = False,
out_channels: Optional[int] = None,
padding: int = 1,
name: str = "conv",
kernel_size=3,
norm_type=None,
eps=None,
elementwise_affine=None,
bias=True,
stride=2,
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.padding = padding
stride = stride
self.name = name
if norm_type == "ln_norm":
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
elif norm_type == "rms_norm":
self.norm = RMSNorm(channels, eps, elementwise_affine)
elif norm_type is None:
self.norm = None
else:
raise ValueError(f"unknown norm_type: {norm_type}")
if use_conv:
conv = CausalConv3d(
self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias
)
else:
raise NotImplementedError
if name == "conv":
self.Conv2d_0 = conv
self.conv = conv
elif name == "Conv2d_0":
self.conv = conv
else:
self.conv = conv
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
if self.norm is not None:
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)
assert hidden_states.shape[1] == self.channels
hidden_states = self.conv(hidden_states)
return hidden_states
class ResnetBlockCausal3D(nn.Module):
r"""
A Resnet block.
"""
def __init__(
self,
*,
in_channels: int,
out_channels: Optional[int] = None,
conv_shortcut: bool = False,
dropout: float = 0.0,
temb_channels: int = 512,
groups: int = 32,
groups_out: Optional[int] = None,
pre_norm: bool = True,
eps: float = 1e-6,
non_linearity: str = "swish",
skip_time_act: bool = False,
# default, scale_shift, ada_group, spatial
time_embedding_norm: str = "default",
kernel: Optional[torch.FloatTensor] = None,
output_scale_factor: float = 1.0,
use_in_shortcut: Optional[bool] = None,
up: bool = False,
down: bool = False,
conv_shortcut_bias: bool = True,
conv_3d_out_channels: Optional[int] = None,
):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.up = up
self.down = down
self.output_scale_factor = output_scale_factor
self.time_embedding_norm = time_embedding_norm
self.skip_time_act = skip_time_act
linear_cls = nn.Linear
if groups_out is None:
groups_out = groups
if self.time_embedding_norm == "ada_group":
self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm1 = SpatialNorm(in_channels, temb_channels)
else:
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
if temb_channels is not None:
if self.time_embedding_norm == "default":
self.time_emb_proj = linear_cls(temb_channels, out_channels)
elif self.time_embedding_norm == "scale_shift":
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
elif self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
self.time_emb_proj = None
else:
raise ValueError(f"Unknown time_embedding_norm : {self.time_embedding_norm} ")
else:
self.time_emb_proj = None
if self.time_embedding_norm == "ada_group":
self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm2 = SpatialNorm(out_channels, temb_channels)
else:
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
self.nonlinearity = get_activation(non_linearity)
self.upsample = self.downsample = None
if self.up:
self.upsample = UpsampleCausal3D(in_channels, use_conv=False)
elif self.down:
self.downsample = DownsampleCausal3D(in_channels, use_conv=False, name="op")
self.use_in_shortcut = self.in_channels != conv_3d_out_channels if use_in_shortcut is None else use_in_shortcut
self.conv_shortcut = None
if self.use_in_shortcut:
self.conv_shortcut = CausalConv3d(
in_channels,
conv_3d_out_channels,
kernel_size=1,
stride=1,
bias=conv_shortcut_bias,
)
def forward(
self,
input_tensor: torch.FloatTensor,
temb: torch.FloatTensor,
scale: float = 1.0,
) -> torch.FloatTensor:
hidden_states = input_tensor
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
hidden_states = self.norm1(hidden_states, temb)
else:
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
if self.upsample is not None:
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
if hidden_states.shape[0] >= 64:
input_tensor = input_tensor.contiguous()
hidden_states = hidden_states.contiguous()
input_tensor = (
self.upsample(input_tensor, scale=scale)
)
hidden_states = (
self.upsample(hidden_states, scale=scale)
)
elif self.downsample is not None:
input_tensor = (
self.downsample(input_tensor, scale=scale)
)
hidden_states = (
self.downsample(hidden_states, scale=scale)
)
hidden_states = self.conv1(hidden_states)
if self.time_emb_proj is not None:
if not self.skip_time_act:
temb = self.nonlinearity(temb)
temb = (
self.time_emb_proj(temb, scale)[:, :, None, None]
)
if temb is not None and self.time_embedding_norm == "default":
hidden_states = hidden_states + temb
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
hidden_states = self.norm2(hidden_states, temb)
else:
hidden_states = self.norm2(hidden_states)
if temb is not None and self.time_embedding_norm == "scale_shift":
scale, shift = torch.chunk(temb, 2, dim=1)
hidden_states = hidden_states * (1 + scale) + shift
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.conv_shortcut is not None:
input_tensor = (
self.conv_shortcut(input_tensor)
)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
def get_down_block3d(
down_block_type: str,
num_layers: int,
in_channels: int,
out_channels: int,
temb_channels: int,
add_downsample: bool,
downsample_stride: int,
resnet_eps: float,
resnet_act_fn: str,
transformer_layers_per_block: int = 1,
num_attention_heads: Optional[int] = None,
resnet_groups: Optional[int] = None,
cross_attention_dim: Optional[int] = None,
downsample_padding: Optional[int] = None,
dual_cross_attention: bool = False,
use_linear_projection: bool = False,
only_cross_attention: bool = False,
upcast_attention: bool = False,
resnet_time_scale_shift: str = "default",
attention_type: str = "default",
resnet_skip_time_act: bool = False,
resnet_out_scale_factor: float = 1.0,
cross_attention_norm: Optional[str] = None,
attention_head_dim: Optional[int] = None,
downsample_type: Optional[str] = None,
dropout: float = 0.0,
):
# If attn head dim is not defined, we default it to the number of heads
if attention_head_dim is None:
logger.warn(
f"It is recommended to provide `attention_head_dim` when calling `get_down_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
)
attention_head_dim = num_attention_heads
down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
dropout=dropout,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
resnet_time_scale_shift=resnet_time_scale_shift,
)
raise ValueError(f"{down_block_type} does not exist.")
def get_up_block3d(
up_block_type: str,
num_layers: int,
in_channels: int,
out_channels: int,
prev_output_channel: int,
temb_channels: int,
add_upsample: bool,
upsample_scale_factor: Tuple,
resnet_eps: float,
resnet_act_fn: str,
resolution_idx: Optional[int] = None,
transformer_layers_per_block: int = 1,
num_attention_heads: Optional[int] = None,
resnet_groups: Optional[int] = None,
cross_attention_dim: Optional[int] = None,
dual_cross_attention: bool = False,
use_linear_projection: bool = False,
only_cross_attention: bool = False,
upcast_attention: bool = False,
resnet_time_scale_shift: str = "default",
attention_type: str = "default",
resnet_skip_time_act: bool = False,
resnet_out_scale_factor: float = 1.0,
cross_attention_norm: Optional[str] = None,
attention_head_dim: Optional[int] = None,
upsample_type: Optional[str] = None,
dropout: float = 0.0,
) -> nn.Module:
# If attn head dim is not defined, we default it to the number of heads
if attention_head_dim is None:
logger.warn(
f"It is recommended to provide `attention_head_dim` when calling `get_up_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
)
attention_head_dim = num_attention_heads
up_block_type = up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
if up_block_type == "UpDecoderBlockCausal3D":
return UpDecoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
resolution_idx=resolution_idx,
dropout=dropout,
add_upsample=add_upsample,
upsample_scale_factor=upsample_scale_factor,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
resnet_time_scale_shift=resnet_time_scale_shift,
temb_channels=temb_channels,
)
raise ValueError(f"{up_block_type} does not exist.")
class UNetMidBlockCausal3D(nn.Module):
"""
A 3D UNet mid-block [`UNetMidBlockCausal3D`] with multiple residual blocks and optional attention blocks.
"""
def __init__(
self,
in_channels: int,
temb_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
attn_groups: Optional[int] = None,
resnet_pre_norm: bool = True,
add_attention: bool = True,
attention_head_dim: int = 1,
output_scale_factor: float = 1.0,
):
super().__init__()
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
self.add_attention = add_attention
if attn_groups is None:
attn_groups = resnet_groups if resnet_time_scale_shift == "default" else None
# there is always at least one resnet
resnets = [
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
]
attentions = []
if attention_head_dim is None:
logger.warn(
f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {in_channels}."
)
attention_head_dim = in_channels
for _ in range(num_layers):
if self.add_attention:
attentions.append(
Attention(
in_channels,
heads=in_channels // attention_head_dim,
dim_head=attention_head_dim,
rescale_output_factor=output_scale_factor,
eps=resnet_eps,
norm_num_groups=attn_groups,
spatial_norm_dim=temb_channels if resnet_time_scale_shift == "spatial" else None,
residual_connection=True,
bias=True,
upcast_softmax=True,
_from_deprecated_attn_block=True,
)
)
else:
attentions.append(None)
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
hidden_states = self.resnets[0](hidden_states, temb)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
B, C, T, H, W = hidden_states.shape
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
attention_mask = prepare_causal_attention_mask(
T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B
)
hidden_states = attn(hidden_states, temb=temb, attention_mask=attention_mask)
hidden_states = rearrange(hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W)
hidden_states = resnet(hidden_states, temb)
return hidden_states
class DownEncoderBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default",
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_downsample: bool = True,
downsample_stride: int = 2,
downsample_padding: int = 1,
):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
temb_channels=None,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList(
[
DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)
]
)
else:
self.downsamplers = None
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
class UpDecoderBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
resolution_idx: Optional[int] = None,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_upsample: bool = True,
upsample_scale_factor=(2, 2, 2),
temb_channels: Optional[int] = None,
):
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=input_channels,
out_channels=out_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.resnets = nn.ModuleList(resnets)
if add_upsample:
self.upsamplers = nn.ModuleList(
[
UpsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
upsample_factor=upsample_scale_factor,
)
]
)
else:
self.upsamplers = None
self.resolution_idx = resolution_idx
def forward(
self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0
) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=temb, scale=scale)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states)
return hidden_states
+355
View File
@@ -0,0 +1,355 @@
from dataclasses import dataclass
from typing import Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from diffusers.utils import BaseOutput, is_torch_version
from diffusers.utils.torch_utils import randn_tensor
from diffusers.models.attention_processor import SpatialNorm
from .unet_causal_3d_blocks import (
CausalConv3d,
UNetMidBlockCausal3D,
get_down_block3d,
get_up_block3d,
)
@dataclass
class DecoderOutput(BaseOutput):
r"""
Output of decoding method.
Args:
sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
The decoded output sample from the last layer of the model.
"""
sample: torch.FloatTensor
class EncoderCausal3D(nn.Module):
r"""
The `EncoderCausal3D` layer of a variational autoencoder that encodes its input into a latent representation.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (64,),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
double_z: bool = True,
mid_block_add_attention=True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
add_time_downsample = bool(
i >= (len(block_out_channels) - 1 - num_time_downsample_layers)
and not is_final_block
)
else:
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2,) if add_time_downsample else (1,)
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
temb_channels=None,
)
self.down_blocks.append(down_block)
# mid
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
output_scale_factor=1,
resnet_time_scale_shift="default",
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
temb_channels=None,
add_attention=mid_block_add_attention,
)
# out
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = CausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
def forward(self, sample: torch.FloatTensor) -> torch.FloatTensor:
r"""The forward method of the `EncoderCausal3D` class."""
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
# middle
sample = self.mid_block(sample)
# post-process
sample = self.conv_norm_out(sample)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class DecoderCausal3D(nn.Module):
r"""
The `DecoderCausal3D` layer of a variational autoencoder that decodes its latent representation into an output sample.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (64,),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
norm_type: str = "group", # group, spatial
mid_block_add_attention=True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
self.mid_block = None
self.up_blocks = nn.ModuleList([])
temb_channels = in_channels if norm_type == "spatial" else None
# mid
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
output_scale_factor=1,
resnet_time_scale_shift="default" if norm_type == "group" else norm_type,
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
temb_channels=temb_channels,
add_attention=mid_block_add_attention,
)
# up
reversed_block_out_channels = list(reversed(block_out_channels))
output_channel = reversed_block_out_channels[0]
for i, up_block_type in enumerate(up_block_types):
prev_output_channel = output_channel
output_channel = reversed_block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
num_time_upsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
add_time_upsample = bool(
i >= len(block_out_channels) - 1 - num_time_upsample_layers
and not is_final_block
)
else:
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
up_block = get_up_block3d(
up_block_type,
num_layers=self.layers_per_block + 1,
in_channels=prev_output_channel,
out_channels=output_channel,
prev_output_channel=None,
add_upsample=bool(add_spatial_upsample or add_time_upsample),
upsample_scale_factor=upsample_scale_factor,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
temb_channels=temb_channels,
resnet_time_scale_shift=norm_type,
)
self.up_blocks.append(up_block)
prev_output_channel = output_channel
# out
if norm_type == "spatial":
self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
else:
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(
self,
sample: torch.FloatTensor,
latent_embeds: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
r"""The forward method of the `DecoderCausal3D` class."""
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions."
sample = self.conv_in(sample)
upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
if is_torch_version(">=", "1.11.0"):
# middle
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block),
sample,
latent_embeds,
use_reentrant=False,
)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block),
sample,
latent_embeds,
use_reentrant=False,
)
else:
# middle
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block), sample, latent_embeds
)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(up_block), sample, latent_embeds)
else:
# middle
sample = self.mid_block(sample, latent_embeds)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = up_block(sample, latent_embeds)
# post-process
if latent_embeds is None:
sample = self.conv_norm_out(sample)
else:
sample = self.conv_norm_out(sample, latent_embeds)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class DiagonalGaussianDistribution(object):
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
if parameters.ndim == 3:
dim = 2 # (B, L, C)
elif parameters.ndim == 5 or parameters.ndim == 4:
dim = 1 # (B, C, T, H ,W) / (B, C, H, W)
else:
raise NotImplementedError
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=dim)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
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, device=self.parameters.device, dtype=self.parameters.dtype
)
def sample(self, generator: Optional[torch.Generator] = None) -> torch.FloatTensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = randn_tensor(
self.mean.shape,
generator=generator,
device=self.parameters.device,
dtype=self.parameters.dtype,
)
x = self.mean + self.std * sample
return x
def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
else:
reduce_dim = list(range(1, self.mean.ndim))
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=reduce_dim,
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var
- 1.0
- self.logvar
+ other.logvar,
dim=reduce_dim,
)
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.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) -> torch.Tensor:
return self.mean
@@ -0,0 +1,431 @@
import torch
import argparse
from safetensors.torch import save_file
import os
parser = argparse.ArgumentParser()
parser.add_argument("--diffusers_path", required=True, type=str)
parser.add_argument("--transformer_path", type=str, default=None, help="Path to save transformer model")
parser.add_argument("--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model")
parser.add_argument("--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model")
args = parser.parse_args()
def reverse_scale_shift(weight, dim):
scale, shift = weight.chunk(2, dim=0)
new_weight = torch.cat([shift, scale], dim=0)
return new_weight
def reverse_proj_gate(weight):
gate, proj = weight.chunk(2, dim=0)
new_weight = torch.cat([proj, gate], dim=0)
return new_weight
def convert_diffusers_transformer_to_mochi(state_dict):
original_state_dict = state_dict.copy()
new_state_dict = {}
# Convert patch_embed
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop("patch_embed.proj.weight")
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop("patch_embed.proj.bias")
# Convert time_embed
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.weight")
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.bias")
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.weight")
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.bias")
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop("time_embed.pooler.to_kv.weight")
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop("time_embed.pooler.to_kv.bias")
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop("time_embed.pooler.to_q.weight")
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop("time_embed.pooler.to_q.bias")
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop("time_embed.pooler.to_out.weight")
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop("time_embed.pooler.to_out.bias")
new_state_dict["t5_yproj.weight"] = original_state_dict.pop("time_embed.caption_proj.weight")
new_state_dict["t5_yproj.bias"] = original_state_dict.pop("time_embed.caption_proj.bias")
# Convert transformer blocks
num_layers = 48
for i in range(num_layers):
block_prefix = f"transformer_blocks.{i}."
new_prefix = f"blocks.{i}."
# norm1
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(block_prefix + "norm1.linear.weight")
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(block_prefix + "norm1.linear.bias")
if i < num_layers - 1:
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.weight"
)
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.bias"
)
else:
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.weight"
)
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.bias"
)
# Visual attention
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
k = original_state_dict.pop(block_prefix + "attn1.to_k.weight")
v = original_state_dict.pop(block_prefix + "attn1.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_x.weight"] = qkv_weight
new_state_dict[new_prefix + "attn.q_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_q.weight"
)
new_state_dict[new_prefix + "attn.k_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_k.weight"
)
new_state_dict[new_prefix + "attn.proj_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.weight"
)
new_state_dict[new_prefix + "attn.proj_x.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.bias"
)
# Context attention
q = original_state_dict.pop(block_prefix + "attn1.add_q_proj.weight")
k = original_state_dict.pop(block_prefix + "attn1.add_k_proj.weight")
v = original_state_dict.pop(block_prefix + "attn1.add_v_proj.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_y.weight"] = qkv_weight
new_state_dict[new_prefix + "attn.q_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_q.weight"
)
new_state_dict[new_prefix + "attn.k_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_k.weight"
)
if i < num_layers - 1:
new_state_dict[new_prefix + "attn.proj_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.weight"
)
new_state_dict[new_prefix + "attn.proj_y.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.bias"
)
# MLP
new_state_dict[new_prefix + "mlp_x.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight")
)
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(block_prefix + "ff.net.2.weight")
if i < num_layers - 1:
new_state_dict[new_prefix + "mlp_y.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight")
)
new_state_dict[new_prefix + "mlp_y.w2.weight"] = original_state_dict.pop(
block_prefix + "ff_context.net.2.weight"
)
# Output layers
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.weight"), dim=0
)
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.bias"), dim=0
)
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop("proj_out.weight")
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop("proj_out.bias")
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
print("Remaining Keys:", original_state_dict.keys())
return new_state_dict
def convert_diffusers_vae_to_mochi(state_dict):
original_state_dict = state_dict.copy()
encoder_state_dict = {}
decoder_state_dict = {}
# Convert encoder
prefix = "encoder."
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(f"{prefix}proj_in.weight")
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(f"{prefix}proj_in.bias")
# Convert block_in
for i in range(3):
encoder_state_dict[f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
)
# Convert down_blocks
down_block_layers = [3, 4, 6]
for block in range(3):
encoder_state_dict[f"layers.{block+4}.layers.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.bias"
)
for i in range(down_block_layers[block]):
# Convert resnets
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
# Convert attentions
q = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight")
k = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight")
v = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight"
)
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias"
)
# Convert block_out
for i in range(3):
encoder_state_dict[f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
)
q = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_q.weight")
k = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_k.weight")
v = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.weight"
)
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.bias"
)
encoder_state_dict[f"layers.{i+7}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+7}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.bias"
)
# Convert output layers
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.weight")
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.bias")
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
# Convert decoder
prefix = "decoder."
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(f"{prefix}conv_in.weight")
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(f"{prefix}conv_in.bias")
# Convert block_in
for i in range(3):
decoder_state_dict[f"blocks.0.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
)
# Convert up_blocks
up_block_layers = [6, 4, 3]
for block in range(3):
for i in range(up_block_layers[block]):
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
decoder_state_dict[f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.weight"
)
decoder_state_dict[f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.bias"
)
# Convert block_out
for i in range(3):
decoder_state_dict[f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
)
# Convert output layers
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(f"{prefix}proj_out.bias")
return encoder_state_dict, decoder_state_dict
def ensure_safetensors_extension(path):
if not path.endswith('.safetensors'):
path = path + '.safetensors'
return path
def ensure_directory_exists(path):
directory = os.path.dirname(path)
if directory:
os.makedirs(directory, exist_ok=True)
def main(args):
from diffusers import MochiPipeline
pipe = MochiPipeline.from_pretrained(args.diffusers_path)
if args.transformer_path:
transformer_path = ensure_safetensors_extension(args.transformer_path)
ensure_directory_exists(transformer_path)
print(f"Converting transformer model...")
transformer_state_dict = convert_diffusers_transformer_to_mochi(pipe.transformer.state_dict())
save_file(transformer_state_dict, transformer_path)
print(f"Saved transformer to {transformer_path}")
if args.vae_encoder_path and args.vae_decoder_path:
encoder_path = ensure_safetensors_extension(args.vae_encoder_path)
decoder_path = ensure_safetensors_extension(args.vae_decoder_path)
ensure_directory_exists(encoder_path)
ensure_directory_exists(decoder_path)
print(f"Converting VAE models...")
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(pipe.vae.state_dict())
save_file(encoder_state_dict, encoder_path)
print(f"Saved VAE encoder to {encoder_path}")
save_file(decoder_state_dict, decoder_path)
print(f"Saved VAE decoder to {decoder_path}")
elif args.vae_encoder_path or args.vae_decoder_path:
print("Warning: Both VAE encoder and decoder paths must be specified to convert VAE models.")
if __name__ == "__main__":
main(args)
@@ -0,0 +1,48 @@
import torch
mochi_latents_mean = torch.tensor(
[
-0.06730895953510081,
-0.038011381506090416,
-0.07477820912866141,
-0.05565264470995561,
0.012767231469026969,
-0.04703542746246419,
0.043896967884726704,
-0.09346305707025976,
-0.09918314763016893,
-0.008729793427399178,
-0.011931556316503654,
-0.0321993391887285,
]
).view(1, 12, 1, 1, 1)
mochi_latents_std = torch.tensor(
[
0.9263795028493863,
0.9248894543193766,
0.9393059390890617,
0.959253732819592,
0.8244560132752793,
0.917259975397747,
0.9294154431013696,
1.3720942357788521,
0.881393668867029,
0.9168315692124348,
0.9185249279345552,
0.9274757570805041,
]
).view(1, 12, 1, 1, 1)
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
latents = (latents - latents_mean) / latents_std
return latents
elif model_type == "hunyuan":
return latents * 0.476986
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+735
View File
@@ -0,0 +1,735 @@
# Copyright 2024 The Genmo team and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Dict, Optional, Tuple
import torch
import torch.nn as nn
import diffusers
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import is_torch_version, logging
from diffusers.utils import (
USE_PEFT_BACKEND,
is_torch_version,
logging,
scale_lora_layers,
unscale_lora_layers,
)
from diffusers.utils.torch_utils import maybe_allow_in_graph
from diffusers.models.attention import FeedForward as HF_FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.embeddings import (
MochiCombinedTimestepCaptionEmbedding,
PatchEmbed,
)
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.loaders import PeftAdapterMixin
from fastvideo.models.mochi_hf.norm import (
MochiLayerNormContinuous,
MochiRMSNormZero,
MochiModulatedRMSNorm,
MochiRMSNorm,
)
from diffusers.models.normalization import AdaLayerNormContinuous
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather, all_to_all_4D
import torch.nn.functional as F
from diffusers.utils.torch_utils import is_torch_version, maybe_allow_in_graph
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class FeedForward(HF_FeedForward):
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
activation_fn: str = "geglu",
final_dropout: bool = False,
inner_dim=None,
bias: bool = True,
):
super().__init__(
dim, dim_out, mult, dropout, activation_fn, final_dropout, inner_dim, bias
)
assert activation_fn == "swiglu"
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.net[0].proj(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
return self.net[2](LigerSiLUMulFunction.apply(gate, hidden_states))
class MochiAttention(nn.Module):
def __init__(
self,
query_dim: int,
processor: "MochiAttnProcessor2_0",
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias: bool = False,
added_kv_proj_dim: Optional[int] = None,
added_proj_bias: Optional[bool] = True,
out_dim: int = None,
out_context_dim: int = None,
out_bias: bool = True,
context_pre_only: bool = False,
eps: float = 1e-5,
):
super().__init__()
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.out_dim = out_dim if out_dim is not None else query_dim
self.out_context_dim = out_context_dim if out_context_dim else query_dim
self.context_pre_only = context_pre_only
self.heads = out_dim // dim_head if out_dim is not None else heads
self.norm_q = MochiRMSNorm(dim_head, eps)
self.norm_k = MochiRMSNorm(dim_head, eps)
self.norm_added_q = MochiRMSNorm(dim_head, eps)
self.norm_added_k = MochiRMSNorm(dim_head, eps)
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.add_k_proj = nn.Linear(
added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
)
self.add_v_proj = nn.Linear(
added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
)
if self.context_pre_only is not None:
self.add_q_proj = nn.Linear(
added_kv_proj_dim, self.inner_dim, bias=added_proj_bias
)
self.to_out = nn.ModuleList([])
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
if not self.context_pre_only:
self.to_add_out = nn.Linear(
self.inner_dim, self.out_context_dim, bias=out_bias
)
self.processor = processor
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
**kwargs,
):
return self.processor(
self,
hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=attention_mask,
**kwargs,
)
class MochiAttnProcessor2_0:
"""Attention processor used in Mochi."""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0."
)
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# [b, s, h * d]
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
# [b, s, h=24, d=128]
query = query.unflatten(2, (attn.heads, -1))
key = key.unflatten(2, (attn.heads, -1))
value = value.unflatten(2, (attn.heads, -1))
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# [b, 256, h * d]
encoder_query = attn.add_q_proj(encoder_hidden_states)
encoder_key = attn.add_k_proj(encoder_hidden_states)
encoder_value = attn.add_v_proj(encoder_hidden_states)
# [b, 256, h=24, d=128]
encoder_query = encoder_query.unflatten(2, (attn.heads, -1))
encoder_key = encoder_key.unflatten(2, (attn.heads, -1))
encoder_value = encoder_value.unflatten(2, (attn.heads, -1))
if attn.norm_added_q is not None:
encoder_query = attn.norm_added_q(encoder_query)
if attn.norm_added_k is not None:
encoder_key = attn.norm_added_k(encoder_key)
if image_rotary_emb is not None:
freqs_cos, freqs_sin = image_rotary_emb[0], image_rotary_emb[1]
# shard the head dimension
if get_sequence_parallel_state():
# B, S, H, D to (S, B,) H, D
# batch_size, seq_len, attn_heads, head_dim
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(
dim, nccl_info.rank_within_group * local_heads, local_heads
)
encoder_query = shrink_head(encoder_query, dim=2)
encoder_key = shrink_head(encoder_key, dim=2)
encoder_value = shrink_head(encoder_value, dim=2)
if image_rotary_emb is not None:
freqs_cos = shrink_head(freqs_cos, dim=1)
freqs_sin = shrink_head(freqs_sin, dim=1)
if image_rotary_emb is not None:
def apply_rotary_emb(x, freqs_cos, freqs_sin):
x_even = x[..., 0::2].float()
x_odd = x[..., 1::2].float()
cos = (x_even * freqs_cos - x_odd * freqs_sin).to(x.dtype)
sin = (x_even * freqs_sin + x_odd * freqs_cos).to(x.dtype)
return torch.stack([cos, sin], dim=-1).flatten(-2)
query = apply_rotary_emb(query, freqs_cos, freqs_sin)
key = apply_rotary_emb(key, freqs_cos, freqs_sin)
# query, key, value = query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
# encoder_query, encoder_key, encoder_value = (
# encoder_query.transpose(1, 2),
# encoder_key.transpose(1, 2),
# encoder_value.transpose(1, 2),
# )
# [b, s, h, d]
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# H
query = torch.cat([query, encoder_query], dim=1).unsqueeze(2)
key = torch.cat([key, encoder_key], dim=1).unsqueeze(2)
value = torch.cat([value, encoder_value], dim=1).unsqueeze(2)
# B, S, 3, H, D
qkv = torch.cat([query, key, value], dim=2)
attn_mask = encoder_attention_mask[:, :].bool()
attn_mask = F.pad(attn_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(
qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None
)
# hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask = None, dropout_p=0.0, is_causal=False)
# valid_lengths = encoder_attention_mask.sum(dim=1) + sequence_length
# def no_padding_mask(score, b, h, q_idx, kv_idx):
# return torch.where(kv_idx < valid_lengths[b],score, -float("inf"))
# hidden_states = flex_attention(query, key, value, score_mod=no_padding_mask)
if get_sequence_parallel_state():
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1
)
# B, S, H, D
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
encoder_hidden_states = all_gather(
encoder_hidden_states, dim=2
).contiguous()
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
else:
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1
)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if hasattr(attn, "to_add_out"):
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
return hidden_states, encoder_hidden_states
@maybe_allow_in_graph
class MochiTransformerBlock(nn.Module):
r"""
Transformer block used in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
Args:
dim (`int`):
The number of channels in the input and output.
num_attention_heads (`int`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`):
The number of channels in each head.
qk_norm (`str`, defaults to `"rms_norm"`):
The normalization layer to use.
activation_fn (`str`, defaults to `"swiglu"`):
Activation function to use in feed-forward.
context_pre_only (`bool`, defaults to `False`):
Whether or not to process context-related conditions with additional layers.
eps (`float`, defaults to `1e-6`):
Epsilon value for normalization layers.
"""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
pooled_projection_dim: int,
qk_norm: str = "rms_norm",
activation_fn: str = "swiglu",
context_pre_only: bool = False,
eps: float = 1e-6,
) -> None:
super().__init__()
self.context_pre_only = context_pre_only
self.ff_inner_dim = (4 * dim * 2) // 3
self.ff_context_inner_dim = (4 * pooled_projection_dim * 2) // 3
self.norm1 = MochiRMSNormZero(dim, 4 * dim, eps=eps, elementwise_affine=False)
if not context_pre_only:
self.norm1_context = MochiRMSNormZero(
dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False
)
else:
self.norm1_context = MochiLayerNormContinuous(
embedding_dim=pooled_projection_dim,
conditioning_embedding_dim=dim,
eps=eps,
)
self.attn1 = MochiAttention(
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
bias=False,
added_kv_proj_dim=pooled_projection_dim,
added_proj_bias=False,
out_dim=dim,
out_context_dim=pooled_projection_dim,
context_pre_only=context_pre_only,
processor=MochiAttnProcessor2_0(),
eps=1e-5,
)
# TODO(aryan): norm_context layers are not needed when `context_pre_only` is True
self.norm2 = MochiModulatedRMSNorm(eps=eps)
self.norm2_context = (
MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None
)
self.norm3 = MochiModulatedRMSNorm(eps)
self.norm3_context = (
MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None
)
self.ff = FeedForward(
dim, inner_dim=self.ff_inner_dim, activation_fn=activation_fn, bias=False
)
self.ff_context = None
if not context_pre_only:
self.ff_context = FeedForward(
pooled_projection_dim,
inner_dim=self.ff_context_inner_dim,
activation_fn=activation_fn,
bias=False,
)
self.norm4 = MochiModulatedRMSNorm(eps=eps)
self.norm4_context = MochiModulatedRMSNorm(eps=eps)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[torch.Tensor] = None,
output_attn=False,
) -> Tuple[torch.Tensor, torch.Tensor]:
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(
hidden_states, temb
)
if not self.context_pre_only:
(
norm_encoder_hidden_states,
enc_gate_msa,
enc_scale_mlp,
enc_gate_mlp,
) = self.norm1_context(encoder_hidden_states, temb)
else:
norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb)
attn_hidden_states, context_attn_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
encoder_attention_mask=encoder_attention_mask,
)
hidden_states = hidden_states + self.norm2(
attn_hidden_states, torch.tanh(gate_msa).unsqueeze(1)
)
norm_hidden_states = self.norm3(
hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32))
)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + self.norm4(
ff_output, torch.tanh(gate_mlp).unsqueeze(1)
)
if not self.context_pre_only:
encoder_hidden_states = encoder_hidden_states + self.norm2_context(
context_attn_hidden_states, torch.tanh(enc_gate_msa).unsqueeze(1)
)
norm_encoder_hidden_states = self.norm3_context(
encoder_hidden_states,
(1 + enc_scale_mlp.unsqueeze(1).to(torch.float32)),
)
context_ff_output = self.ff_context(norm_encoder_hidden_states)
encoder_hidden_states = encoder_hidden_states + self.norm4_context(
context_ff_output, torch.tanh(enc_gate_mlp).unsqueeze(1)
)
if not output_attn:
attn_hidden_states = None
return hidden_states, encoder_hidden_states, attn_hidden_states
class MochiRoPE(nn.Module):
r"""
RoPE implementation used in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
Args:
base_height (`int`, defaults to `192`):
Base height used to compute interpolation scale for rotary positional embeddings.
base_width (`int`, defaults to `192`):
Base width used to compute interpolation scale for rotary positional embeddings.
"""
def __init__(self, base_height: int = 192, base_width: int = 192) -> None:
super().__init__()
self.target_area = base_height * base_width
def _centers(self, start, stop, num, device, dtype) -> torch.Tensor:
edges = torch.linspace(start, stop, num + 1, device=device, dtype=dtype)
return (edges[:-1] + edges[1:]) / 2
def _get_positions(
self,
num_frames: int,
height: int,
width: int,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
) -> torch.Tensor:
scale = (self.target_area / (height * width)) ** 0.5
t = torch.arange(num_frames * nccl_info.sp_size, device=device, dtype=dtype)
h = self._centers(
-height * scale / 2, height * scale / 2, height, device, dtype
)
w = self._centers(-width * scale / 2, width * scale / 2, width, device, dtype)
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
positions = torch.stack([grid_t, grid_h, grid_w], dim=-1).view(-1, 3)
return positions
def _create_rope(self, freqs: torch.Tensor, pos: torch.Tensor) -> torch.Tensor:
with torch.autocast(freqs.device.type, enabled=False):
# Always run ROPE freqs computation in FP32
freqs = torch.einsum(
"nd,dhf->nhf", pos.to(torch.float32), freqs.to(torch.float32)
)
freqs_cos = torch.cos(freqs)
freqs_sin = torch.sin(freqs)
return freqs_cos, freqs_sin
def forward(
self,
pos_frequencies: torch.Tensor,
num_frames: int,
height: int,
width: int,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
pos = self._get_positions(num_frames, height, width, device, dtype)
rope_cos, rope_sin = self._create_rope(pos_frequencies, pos)
return rope_cos, rope_sin
@maybe_allow_in_graph
class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
r"""
A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
Args:
patch_size (`int`, defaults to `2`):
The size of the patches to use in the patch embedding layer.
num_attention_heads (`int`, defaults to `24`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`, defaults to `128`):
The number of channels in each head.
num_layers (`int`, defaults to `48`):
The number of layers of Transformer blocks to use.
in_channels (`int`, defaults to `12`):
The number of channels in the input.
out_channels (`int`, *optional*, defaults to `None`):
The number of channels in the output.
qk_norm (`str`, defaults to `"rms_norm"`):
The normalization layer to use.
text_embed_dim (`int`, defaults to `4096`):
Input dimension of text embeddings from the text encoder.
time_embed_dim (`int`, defaults to `256`):
Output dimension of timestep embeddings.
activation_fn (`str`, defaults to `"swiglu"`):
Activation function to use in feed-forward.
max_sequence_length (`int`, defaults to `256`):
The maximum sequence length of text embeddings supported.
"""
_supports_gradient_checkpointing = True
@register_to_config
def __init__(
self,
patch_size: int = 2,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
num_layers: int = 48,
pooled_projection_dim: int = 1536,
in_channels: int = 12,
out_channels: Optional[int] = None,
qk_norm: str = "rms_norm",
text_embed_dim: int = 4096,
time_embed_dim: int = 256,
activation_fn: str = "swiglu",
max_sequence_length: int = 256,
) -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
out_channels = out_channels or in_channels
self.patch_embed = PatchEmbed(
patch_size=patch_size,
in_channels=in_channels,
embed_dim=inner_dim,
pos_embed_type=None,
)
self.time_embed = MochiCombinedTimestepCaptionEmbedding(
embedding_dim=inner_dim,
pooled_projection_dim=pooled_projection_dim,
text_embed_dim=text_embed_dim,
time_embed_dim=time_embed_dim,
num_attention_heads=8,
)
self.pos_frequencies = nn.Parameter(
torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0)
)
self.rope = MochiRoPE()
self.transformer_blocks = nn.ModuleList(
[
MochiTransformerBlock(
dim=inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
pooled_projection_dim=pooled_projection_dim,
qk_norm=qk_norm,
activation_fn=activation_fn,
context_pre_only=i == num_layers - 1,
)
for i in range(num_layers)
]
)
self.norm_out = AdaLayerNormContinuous(
inner_dim,
inner_dim,
elementwise_affine=False,
eps=1e-6,
norm_type="layer_norm",
)
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
self.gradient_checkpointing = False
def _set_gradient_checkpointing(self, module, value=False):
if hasattr(module, "gradient_checkpointing"):
module.gradient_checkpointing = value
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
output_attn=False,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
) -> torch.Tensor:
assert (
return_dict is False
), "return_dict is not supported in MochiTransformer3DModel"
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if (
attention_kwargs is not None
and attention_kwargs.get("scale", None) is not None
):
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p = self.config.patch_size
post_patch_height = height // p
post_patch_width = width // p
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
timestep = 1000 - timestep
temb, encoder_hidden_states = self.time_embed(
timestep,
encoder_hidden_states,
encoder_attention_mask,
hidden_dtype=hidden_states.dtype,
)
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
hidden_states = self.patch_embed(hidden_states)
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
image_rotary_emb = self.rope(
self.pos_frequencies,
num_frames,
post_patch_height,
post_patch_width,
device=hidden_states.device,
dtype=torch.float32,
)
attn_outputs_list = []
for i, block in enumerate(self.transformer_blocks):
if self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = (
{"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
)
(
hidden_states,
encoder_hidden_states,
attn_outputs,
) = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
image_rotary_emb,
output_attn,
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states, attn_outputs = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
temb=temb,
image_rotary_emb=image_rotary_emb,
output_attn=output_attn,
)
attn_outputs_list.append(attn_outputs)
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(
batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1
)
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5)
output = hidden_states.reshape(batch_size, -1, num_frames, height, width)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not output_attn:
attn_outputs_list = None
else:
attn_outputs_list = torch.stack(attn_outputs_list, dim=0)
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
return (-output, attn_outputs_list)
+130
View File
@@ -0,0 +1,130 @@
# Copyright 2024 The Genmo team and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import numbers
from typing import Dict, Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
class MochiModulatedRMSNorm(nn.Module):
def __init__(self, eps: float):
super().__init__()
self.eps = eps
def forward(self, hidden_states, scale=None):
hidden_states_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.eps)
if scale is not None:
hidden_states = hidden_states * scale
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states
class MochiRMSNorm(nn.Module):
def __init__(self, dim, eps: float, elementwise_affine=True):
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim))
else:
self.weight = None
def forward(self, hidden_states):
hidden_states_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.eps)
if self.weight is not None:
# convert into half-precision if necessary
if self.weight.dtype in [torch.float16, torch.bfloat16]:
hidden_states = hidden_states.to(self.weight.dtype)
hidden_states = hidden_states * self.weight
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states
class MochiLayerNormContinuous(nn.Module):
def __init__(
self,
embedding_dim: int,
conditioning_embedding_dim: int,
eps=1e-5,
bias=True,
):
super().__init__()
# AdaLN
self.silu = nn.SiLU()
self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias)
self.norm = MochiModulatedRMSNorm(eps=eps)
def forward(
self,
x: torch.Tensor,
conditioning_embedding: torch.Tensor,
) -> torch.Tensor:
input_dtype = x.dtype
# convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT)
scale = self.linear_1(self.silu(conditioning_embedding).to(x.dtype))
x = self.norm(x, (1 + scale.unsqueeze(1).to(torch.float32)))
return x.to(input_dtype)
class MochiRMSNormZero(nn.Module):
r"""
Adaptive RMS Norm used in Mochi.
Parameters:
embedding_dim (`int`): The size of each embedding vector.
"""
def __init__(
self,
embedding_dim: int,
hidden_dim: int,
eps: float = 1e-5,
elementwise_affine: bool = False,
) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, hidden_dim)
self.norm = MochiModulatedRMSNorm(eps=eps)
def forward(
self, hidden_states: torch.Tensor, emb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
hidden_states_dtype = hidden_states.dtype
emb = self.linear(self.silu(emb))
scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1)
hidden_states = self.norm(
hidden_states, (1 + scale_msa[:, None].to(torch.float32))
)
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states, gate_msa, scale_mlp, gate_mlp
+848
View File
@@ -0,0 +1,848 @@
# Copyright 2024 Black Forest Labs and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
from typing import Callable, Dict, List, Optional, Union, Any
import copy
import numpy as np
import torch
from transformers import T5EncoderModel, T5TokenizerFast
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.models.autoencoders import AutoencoderKL
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import (
is_torch_xla_available,
logging,
replace_example_docstring,
)
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
from einops import rearrange
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """
Examples:
```py
>>> import torch
>>> from diffusers import MochiPipeline
>>> from diffusers.utils import export_to_video
>>> pipe = MochiPipeline.from_pretrained("genmo/mochi-1-preview", torch_dtype=torch.bfloat16)
>>> pipe.to("cuda")
>>> prompt = "Close-up of a chameleon's eye, with its scaly skin changing color. Ultra high resolution 4k."
>>> frames = pipe(prompt, num_inference_steps=28, guidance_scale=3.5).frames[0]
>>> export_to_video(frames, "mochi.mp4")
```
"""
def calculate_shift(
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.16,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
# from: https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
if linear_steps is None:
linear_steps = num_steps // 2
linear_sigma_schedule = [
i * threshold_noise / linear_steps for i in range(linear_steps)
]
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
quadratic_steps = num_steps - linear_steps
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (
quadratic_steps**2
)
const = quadratic_coef * (linear_steps**2)
quadratic_sigma_schedule = [
quadratic_coef * (i**2) + linear_coef * i + const
for i in range(linear_steps, num_steps)
]
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule
sigma_schedule = [1.0 - x for x in sigma_schedule]
return sigma_schedule
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
def retrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] = None,
device: Optional[Union[str, torch.device]] = None,
timesteps: Optional[List[int]] = None,
sigmas: Optional[List[float]] = None,
**kwargs,
):
r"""
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
Args:
scheduler (`SchedulerMixin`):
The scheduler to get timesteps from.
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
must be `None`.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
timesteps (`List[int]`, *optional*):
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
`num_inference_steps` and `sigmas` must be `None`.
sigmas (`List[float]`, *optional*):
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
`num_inference_steps` and `timesteps` must be `None`.
Returns:
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError(
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
)
if timesteps is not None:
accepts_timesteps = "timesteps" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys()
)
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys()
)
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
class MochiPipeline(DiffusionPipeline):
r"""
The mochi pipeline for text-to-video generation.
Reference: https://github.com/genmoai/models
Args:
transformer ([`MochiTransformer3DModel`]):
Conditional Transformer architecture to denoise the encoded video latents.
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
vae ([`AutoencoderKL`]):
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
text_encoder ([`T5EncoderModel`]):
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
tokenizer (`CLIPTokenizer`):
Tokenizer of class
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
tokenizer (`T5TokenizerFast`):
Second Tokenizer of class
[T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
"""
model_cpu_offload_seq = "text_encoder->transformer->vae"
_optional_components = []
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(
self,
scheduler: FlowMatchEulerDiscreteScheduler,
vae: AutoencoderKL,
text_encoder: T5EncoderModel,
tokenizer: T5TokenizerFast,
transformer: MochiTransformer3DModel,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer,
scheduler=scheduler,
)
self.vae_spatial_scale_factor = 8
self.vae_temporal_scale_factor = 6
self.patch_size = 2
self.video_processor = VideoProcessor(
vae_scale_factor=self.vae_spatial_scale_factor
)
self.tokenizer_max_length = (
self.tokenizer.model_max_length
if hasattr(self, "tokenizer") and self.tokenizer is not None
else 77
)
self.default_height = 480
self.default_width = 848
# Adapted from diffusers.pipelines.cogvideo.pipeline_cogvideox.CogVideoXPipeline._get_t5_prompt_embeds
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_videos_per_prompt: int = 1,
max_sequence_length: int = 256,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(
prompt, padding="longest", return_tensors="pt"
).input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
removed_text = self.tokenizer.batch_decode(
untruncated_ids[:, max_sequence_length - 1 : -1]
)
logger.warning(
"The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length} tokens: {removed_text}"
)
prompt_embeds = self.text_encoder(
text_input_ids.to(device), attention_mask=prompt_attention_mask
)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
batch_size * num_videos_per_prompt, seq_len, -1
)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
prompt_attention_mask = prompt_attention_mask.repeat(num_videos_per_prompt, 1)
return prompt_embeds, prompt_attention_mask
# Adapted from diffusers.pipelines.cogvideo.pipeline_cogvideox.CogVideoXPipeline.encode_prompt
def encode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Optional[Union[str, List[str]]] = None,
do_classifier_free_guidance: bool = True,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
max_sequence_length: int = 256,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
Whether to use classifier free guidance or not.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
device: (`torch.device`, *optional*):
torch device
dtype: (`torch.dtype`, *optional*):
torch dtype
"""
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds, prompt_attention_mask = self._get_t5_prompt_embeds(
prompt=prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = (
batch_size * [negative_prompt]
if isinstance(negative_prompt, str)
else negative_prompt
)
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
(
negative_prompt_embeds,
negative_prompt_attention_mask,
) = self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return (
prompt_embeds,
prompt_attention_mask,
negative_prompt_embeds,
negative_prompt_attention_mask,
)
def check_inputs(
self,
prompt,
height,
width,
callback_on_step_end_tensor_inputs=None,
prompt_embeds=None,
negative_prompt_embeds=None,
prompt_attention_mask=None,
negative_prompt_attention_mask=None,
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(
f"`height` and `width` have to be divisible by 8 but are {height} and {width}."
)
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (
not isinstance(prompt, str) and not isinstance(prompt, list)
):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
if prompt_embeds is not None and prompt_attention_mask is None:
raise ValueError(
"Must provide `prompt_attention_mask` when specifying `prompt_embeds`."
)
if (
negative_prompt_embeds is not None
and negative_prompt_attention_mask is None
):
raise ValueError(
"Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
if prompt_attention_mask.shape != negative_prompt_attention_mask.shape:
raise ValueError(
"`prompt_attention_mask` and `negative_prompt_attention_mask` must have the same shape when passed directly, but"
f" got: `prompt_attention_mask` {prompt_attention_mask.shape} != `negative_prompt_attention_mask`"
f" {negative_prompt_attention_mask.shape}."
)
def enable_vae_slicing(self):
r"""
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
"""
self.vae.enable_slicing()
def disable_vae_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
computing decoding in one step.
"""
self.vae.disable_slicing()
def enable_vae_tiling(self):
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger images.
"""
self.vae.enable_tiling()
def disable_vae_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
computing decoding in one step.
"""
self.vae.disable_tiling()
def prepare_latents(
self,
batch_size,
num_channels_latents,
height,
width,
num_frames,
dtype,
device,
generator,
latents=None,
):
height = height // self.vae_spatial_scale_factor
width = width // self.vae_spatial_scale_factor
num_frames = (num_frames - 1) // self.vae_temporal_scale_factor + 1
shape = (batch_size, num_channels_latents, num_frames, height, width)
if latents is not None:
return latents.to(device=device, dtype=dtype)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@property
def guidance_scale(self):
return self._guidance_scale
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1.0
@property
def num_timesteps(self):
return self._num_timesteps
@property
def attention_kwargs(self):
return self._attention_kwargs
@property
def interrupt(self):
return self._interrupt
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
negative_prompt: Optional[Union[str, List[str]]] = None,
height: Optional[int] = None,
width: Optional[int] = None,
num_frames: int = 16,
num_inference_steps: int = 28,
timesteps: List[int] = None,
guidance_scale: float = 4.5,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
output_type: Optional[str] = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 256,
return_all_states=False,
):
r"""
Function invoked when calling the pipeline for generation.
Args:
prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
instead.
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
The height in pixels of the generated image. This is set to 1024 by default for the best results.
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
The width in pixels of the generated image. This is set to 1024 by default for the best results.
num_frames (`int`, defaults to 16):
The number of video frames to generate
num_inference_steps (`int`, *optional*, defaults to 50):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
timesteps (`List[int]`, *optional*):
Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument
in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is
passed will be used. Must be in descending order.
guidance_scale (`float`, defaults to `4.5`):
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
`guidance_scale` is defined as `w` of equation 2. of [Imagen
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
usually at the expense of lower image quality.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
The number of videos to generate per prompt.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
to make generation deterministic.
latents (`torch.Tensor`, *optional*):
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor will ge generated by sampling using the supplied random `generator`.
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
prompt_attention_mask (`torch.Tensor`, *optional*):
Pre-generated attention mask for text embeddings.
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated negative text embeddings. For PixArt-Sigma this negative prompt should be "". If not
provided, negative_prompt_embeds will be generated from `negative_prompt` input argument.
negative_prompt_attention_mask (`torch.FloatTensor`, *optional*):
Pre-generated attention mask for negative text embeddings.
output_type (`str`, *optional*, defaults to `"pil"`):
The output format of the generate image. Choose between
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~pipelines.mochi.MochiPipelineOutput`] instead of a plain tuple.
attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
`self.processor` in
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
callback_on_step_end (`Callable`, *optional*):
A function that calls at the end of each denoising steps during the inference. The function is called
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
`callback_on_step_end_tensor_inputs`.
callback_on_step_end_tensor_inputs (`List`, *optional*):
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
`._callback_tensor_inputs` attribute of your pipeline class.
max_sequence_length (`int` defaults to `256`):
Maximum sequence length to use with the `prompt`.
Examples:
Returns:
[`~pipelines.mochi.MochiPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`~pipelines.mochi.MochiPipelineOutput`] is returned, otherwise a `tuple`
is returned where the first element is a list with the generated images.
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
height = height or self.default_height
width = width or self.default_width
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt=prompt,
height=height,
width=width,
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._interrupt = False
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
# 3. Prepare text embeddings
(
prompt_embeds,
prompt_attention_mask,
negative_prompt_embeds,
negative_prompt_attention_mask,
) = self.encode_prompt(
prompt=prompt,
negative_prompt=negative_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
num_videos_per_prompt=num_videos_per_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
max_sequence_length=max_sequence_length,
device=device,
)
if self.do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
prompt_attention_mask = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask], dim=0
)
# 4. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
height,
width,
num_frames,
prompt_embeds.dtype,
device,
generator,
latents,
)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(
latents, "b t (n s) h w -> b t n s h w", n=world_size
).contiguous()
latents = latents[:, :, rank, :, :, :]
original_noise = copy.deepcopy(latents)
# 5. Prepare timestep
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
threshold_noise = 0.025
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
sigmas = np.array(sigmas)
# check if of type FlowMatchEulerDiscreteScheduler
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
device,
timesteps,
sigmas,
)
else:
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
device,
)
num_warmup_steps = max(
len(timesteps) - num_inference_steps * self.scheduler.order, 0
)
self._num_timesteps = len(timesteps)
# 6. Denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
latent_model_input = (
torch.cat([latents] * 2)
if self.do_classifier_free_guidance
else latents
)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
# Mochi CFG + Sampling runs in FP32
noise_pred = noise_pred.to(torch.float32)
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond
)
# compute the previous noisy sample x_t -> x_t-1
latents_dtype = latents.dtype
latents = self.scheduler.step(
noise_pred, t, latents.to(torch.float32), return_dict=False
)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
if torch.backends.mps.is_available():
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
latents = latents.to(latents_dtype)
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0
):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
# latents_shape = list(latents.shape)
# full_shape = [latents_shape[0] * world_size] + latents_shape[1:]
# all_latents = torch.zeros(full_shape, dtype=latents.dtype, device=latents.device)
# torch.distributed.all_gather_into_tensor(all_latents, latents)
# latents_list = list(all_latents.chunk(world_size, dim=0))
# latents = torch.cat(latents_list, dim=2)
if output_type == "latent":
video = latents
else:
# unscale/denormalize the latents
# denormalize with the mean and std if available and not None
has_latents_mean = (
hasattr(self.vae.config, "latents_mean")
and self.vae.config.latents_mean is not None
)
has_latents_std = (
hasattr(self.vae.config, "latents_std")
and self.vae.config.latents_std is not None
)
if has_latents_mean and has_latents_std:
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, 12, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = (
torch.tensor(self.vae.config.latents_std)
.view(1, 12, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents = (
latents * latents_std / self.vae.config.scaling_factor
+ latents_mean
)
else:
latents = latents / self.vae.config.scaling_factor
video = self.vae.decode(latents, return_dict=False)[0]
video = self.video_processor.postprocess_video(
video, output_type=output_type
)
# Offload all models
self.maybe_free_model_hooks()
if return_all_states:
# Pay extra attention here:
# prompt_embeds with shape torch.Size([2, 256]), where prompt_embeds[1] is the prompt_embeds for the actual prompt
# prompt_embeds[0] is for negative prompt
return original_noise, video, latents, prompt_embeds, prompt_attention_mask
if not return_dict:
return (video,)
return MochiPipelineOutput(frames=video)
+133
View File
@@ -0,0 +1,133 @@
import json
import torch.distributed as dist
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
import os
from diffusers.utils import export_to_video
import argparse
def generate_video_and_latent(
pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale
):
# Set the random seed for reproducibility
generator = torch.Generator("cuda").manual_seed(12345)
# Generate videos from the input prompt
noise, video, latent, prompt_embed, prompt_attention_mask = pipe(
prompt=prompt,
height=height,
width=width,
num_frames=num_frames,
generator=generator,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
output_type="latent_and_video",
)
# prompt_embed has negative prompt at index 0
return noise[0], video[0], latent[0], prompt_embed[1], prompt_attention_mask[1]
# return dummy tensor to debug first
# return torch.zeros(1, 3, 480, 848), torch.zeros(1, 256, 16, 16)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument(
"--prompt_path", type=str, default="data/dummyVid/videos2caption.json"
)
parser.add_argument("--dataset_output_dir", type=str, default="data/dummySynthetic")
args = parser.parse_args()
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
torch.cuda.set_device(local_rank)
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
if not isinstance(args.prompt_path, list):
args.prompt_path = [args.prompt_path]
if len(args.prompt_path) == 1 and args.prompt_path[0].endswith("txt"):
text_prompt = open(args.prompt_path[0], "r").readlines()
text_prompt = [i.strip() for i in text_prompt]
pipe = MochiPipeline.from_pretrained(args.model_path, torch_dtype=torch.bfloat16)
pipe.enable_vae_tiling()
pipe.enable_model_cpu_offload(gpu_id=local_rank)
# make dir if not exist
os.makedirs(args.dataset_output_dir, exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "noise"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(
os.path.join(args.dataset_output_dir, "prompt_attention_mask"), exist_ok=True
)
data = []
for i, prompt in enumerate(text_prompt):
if i % world_size != local_rank:
continue
(
noise,
video,
latent,
prompt_embed,
prompt_attention_mask,
) = generate_video_and_latent(
pipe,
prompt,
args.height,
args.width,
args.num_frames,
args.num_inference_steps,
args.guidance_scale,
)
# save latent
video_name = str(i)
noise_path = os.path.join(args.dataset_output_dir, "noise", video_name + ".pt")
latent_path = os.path.join(
args.dataset_output_dir, "latent", video_name + ".pt"
)
prompt_embed_path = os.path.join(
args.dataset_output_dir, "prompt_embed", video_name + ".pt"
)
video_path = os.path.join(args.dataset_output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(
args.dataset_output_dir, "prompt_attention_mask", video_name + ".pt"
)
# save latent
torch.save(noise, noise_path)
torch.save(latent, latent_path)
torch.save(prompt_embed, prompt_embed_path)
torch.save(prompt_attention_mask, prompt_attention_mask_path)
export_to_video(video, video_path, fps=30)
item = {}
item["cap"] = prompt
item["video"] = video_name + ".mp4"
item["noise"] = video_name + ".pt"
item["latent_path"] = video_name + ".pt"
item["prompt_embed_path"] = video_name + ".pt"
item["prompt_attention_mask"] = video_name + ".pt"
data.append(item)
dist.barrier()
local_data = data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
# save json
if local_rank == 0:
all_data = [item for sublist in gathered_data for item in sublist]
with open(
os.path.join(args.dataset_output_dir, "videos2caption.json"), "w"
) as f:
json.dump(all_data, f, indent=4)
+148
View File
@@ -0,0 +1,148 @@
import os
import imageio
import time
from einops import rearrange
import torch
import torchvision
import numpy as np
from pathlib import Path
from loguru import logger
from datetime import datetime
import argparse
from diffusers.utils import export_to_video
from fastvideo.models.hunyuan.utils.file_utils import save_videos_grid
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
import torch.distributed as dist
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
def initialize_distributed():
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
initialize_sequence_parallel_state(world_size)
def main(args):
initialize_distributed()
print(nccl_info.sp_size)
device = torch.cuda.current_device()
print(args)
models_root_path = Path(args.model_path)
if not models_root_path.exists():
raise ValueError(f"`models_root` not exists: {models_root_path}")
# Create save folder to save the samples
save_path = args.output_path
os.makedirs(os.path.dirname(save_path), exist_ok=True)
# Load models
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
# Get the updated args
args = hunyuan_video_sampler.args
# Start sampling
samples = []
for prompt in args.prompts:
outputs = hunyuan_video_sampler.predict(
prompt=prompt,
height=args.height,
width=args.width,
video_length=args.num_frames,
seed=args.seed,
negative_prompt=args.neg_prompt,
infer_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
num_videos_per_prompt=args.num_videos,
flow_shift=args.flow_shift,
batch_size=args.batch_size,
embedded_guidance_scale=args.embedded_cfg_scale
)
samples.append(outputs['samples'][0])
for prompt, video in zip(args.prompts, samples):
videos = rearrange(video.unsqueeze(0), "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
outputs.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
imageio.mimsave(args.output_path + f"{prompt[:100]}.mp4", outputs, fps=args.fps)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Basic parameters
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--model_path", type=str, default="data/hunyuan")
parser.add_argument("--output_path", type=str, default="./outputs/video")
parser.add_argument("--fps", type=int, default=24)
# Additional parameters
parser.add_argument("--denoise-type", type=str, default="flow", help="Denoise type for noised inputs.")
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
parser.add_argument("--guidance_scale", type=float, default=1.0, help="Classifier free guidance scale.")
parser.add_argument("--embedded_cfg_scale", type=float, default=6.0, help="Embedded classifier free guidance scale.")
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
parser.add_argument("--num_videos", type=int, default=1, help="Number of videos to generate per prompt.")
parser.add_argument("--load-key", type=str, default="module", help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.")
parser.add_argument("--use-cpu-offload", action="store_true", help="Use CPU offload for the model load.")
parser.add_argument("--dit-weight", type=str, default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt")
parser.add_argument("--reproduce", action="store_true", help="Enable reproducibility by setting random seeds and deterministic algorithms.")
parser.add_argument("--disable-autocast", action="store_true", help="Disable autocast for denoising loop and vae decoding in pipeline sampling.")
# Flow Matching
parser.add_argument("--flow-reverse", action="store_true", help="If reverse, learning/sampling from t=1 -> t=0.")
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
parser.add_argument("--use-linear-quadratic-schedule", action="store_true",
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)")
parser.add_argument("--linear-schedule-end", type=int, default=25,
help="End step for linear quadratic schedule for flow matching.")
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--text-encoder", type=str, default="llm")
parser.add_argument("--text-encoder-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
parser.add_argument("--text-encoder-2", type=str, default="clipL")
parser.add_argument("--text-encoder-precision-2", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--text-states-dim-2", type=int, default=768)
parser.add_argument("--tokenizer-2", type=str, default="clipL")
parser.add_argument("--text-len-2", type=int, default=77)
args = parser.parse_args()
main(args)
@@ -0,0 +1,126 @@
import os
import imageio
import time
from einops import rearrange
import torch
import torchvision
import numpy as np
from pathlib import Path
from loguru import logger
from datetime import datetime
import argparse
from diffusers.utils import export_to_video
from fastvideo.models.hunyuan.utils.file_utils import save_videos_grid
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
def main(args):
print(args)
models_root_path = Path(args.model_path)
if not models_root_path.exists():
raise ValueError(f"`models_root` not exists: {models_root_path}")
# Create save folder to save the samples
save_path = args.output_path
os.makedirs(os.path.dirname(save_path), exist_ok=True)
# Load models
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
# Get the updated args
args = hunyuan_video_sampler.args
# Start sampling
samples = []
for prompt in args.prompts:
outputs = hunyuan_video_sampler.predict(
prompt=prompt,
height=args.height,
width=args.width,
video_length=args.num_frames,
seed=args.seed,
negative_prompt=args.neg_prompt,
infer_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
num_videos_per_prompt=args.num_videos,
flow_shift=args.flow_shift,
batch_size=args.batch_size,
embedded_guidance_scale=args.embedded_cfg_scale
)
samples.append(outputs['samples'][0])
for prompt, video in zip(args.prompts, samples):
videos = rearrange(video.unsqueeze(0), "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
outputs.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
imageio.mimsave(args.output_path + f"{prompt[:100]}.mp4", outputs, fps=args.fps)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Basic parameters
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--model_path", type=str, default="data/hunyuan")
parser.add_argument("--output_path", type=str, default="./outputs/video")
parser.add_argument("--fps", type=int, default=24)
# Additional parameters
parser.add_argument("--denoise-type", type=str, default="flow", help="Denoise type for noised inputs.")
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
parser.add_argument("--guidance_scale", type=float, default=1.0, help="Classifier free guidance scale.")
parser.add_argument("--embedded_cfg_scale", type=float, default=6.0, help="Embedded classifier free guidance scale.")
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
parser.add_argument("--num_videos", type=int, default=1, help="Number of videos to generate per prompt.")
parser.add_argument("--load-key", type=str, default="module", help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.")
parser.add_argument("--use-cpu-offload", action="store_true", help="Use CPU offload for the model load.")
parser.add_argument("--dit-weight", type=str, default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt")
parser.add_argument("--reproduce", action="store_true", help="Enable reproducibility by setting random seeds and deterministic algorithms.")
parser.add_argument("--disable-autocast", action="store_true", help="Disable autocast for denoising loop and vae decoding in pipeline sampling.")
# Flow Matching
parser.add_argument("--flow-reverse", action="store_true", help="If reverse, learning/sampling from t=1 -> t=0.")
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
parser.add_argument("--use-linear-quadratic-schedule", action="store_true",
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)")
parser.add_argument("--linear-schedule-end", type=int, default=25,
help="End step for linear quadratic schedule for flow matching.")
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--text-encoder", type=str, default="llm")
parser.add_argument("--text-encoder-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
parser.add_argument("--text-encoder-2", type=str, default="clipL")
parser.add_argument("--text-encoder-precision-2", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--text-states-dim-2", type=int, default=768)
parser.add_argument("--tokenizer-2", type=str, default="clipL")
parser.add_argument("--text-len-2", type=int, default=77)
args = parser.parse_args()
main(args)
+177
View File
@@ -0,0 +1,177 @@
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
import torch.distributed as dist
from diffusers.utils import export_to_video
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
import argparse
import os
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
import json
from typing import Optional
from safetensors.torch import save_file, load_file
from peft import set_peft_model_state_dict, inject_adapter_in_model, load_peft_weights
from peft import LoraConfig
import sys
import pdb
import copy
from typing import Dict
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import convert_unet_state_dict_to_peft
from fastvideo.distill.solver import PCMFMScheduler
def initialize_distributed():
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
initialize_sequence_parallel_state(world_size)
def main(args):
initialize_distributed()
print(nccl_info.sp_size)
device = torch.cuda.current_device()
generator = torch.Generator(device).manual_seed(args.seed)
weight_dtype = torch.bfloat16
if args.scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
scheduler = PCMFMScheduler(
1000,
args.shift,
args.num_euler_timesteps,
linear_quadratic,
args.linear_threshold,
args.linear_range,
)
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/"
)
pipe = MochiPipeline.from_pretrained(
args.model_path, transformer=transformer, scheduler=scheduler
)
pipe.enable_vae_tiling()
if args.lora_checkpoint_dir is not None:
print(f"Loading LoRA weights from {args.lora_checkpoint_dir}")
config_path = os.path.join(args.lora_checkpoint_dir, "lora_config.json")
with open(config_path, "r") as f:
lora_config_dict = json.load(f)
rank = lora_config_dict["lora_params"]["lora_rank"]
lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
lora_scaling = lora_alpha / rank
pipe.load_lora_weights(args.lora_checkpoint_dir, adapter_name="default")
pipe.set_adapters(["default"], [lora_scaling])
print(f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}")
# pipe.to(device)
pipe.enable_model_cpu_offload(device)
# Generate videos from the input prompt
if args.prompt_embed_path is not None:
prompt_embeds = (
torch.load(args.prompt_embed_path, map_location="cpu", weights_only=True)
.to(device)
.unsqueeze(0)
)
encoder_attention_mask = (
torch.load(
args.encoder_attention_mask_path, map_location="cpu", weights_only=True
)
.to(device)
.unsqueeze(0)
)
prompts = None
elif args.prompt_path is not None:
prompts = [line.strip() for line in open(args.prompt_path, "r")]
prompt_embeds = None
encoder_attention_mask = None
else:
prompts = args.prompts
prompt_embeds = None
encoder_attention_mask = None
if prompts is not None:
videos = []
with torch.autocast("cuda", dtype=torch.bfloat16):
for prompt in prompts:
video = pipe(
prompt=[prompt],
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
generator=generator,
).frames
videos.append(video[0])
else:
with torch.autocast("cuda", dtype=torch.bfloat16):
videos = pipe(
prompt_embeds=prompt_embeds,
prompt_attention_mask=encoder_attention_mask,
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
generator=generator,
).frames
if nccl_info.global_rank <= 0:
if prompts is not None:
# mkdir
os.makedirs(args.output_path, exist_ok=True)
for video, prompt in zip(videos, prompts):
suffix = prompt.split(".")[0]
export_to_video(
video, os.path.join(args.output_path, f"{suffix}.mp4"), fps=30
)
else:
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
if __name__ == "__main__":
# arg parse
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--output_path", type=str, default="./outputs.mp4")
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--prompt_embed_path", type=str, default=None)
parser.add_argument("--prompt_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="euler")
parser.add_argument("--encoder_attention_mask_path", type=str, default=None)
parser.add_argument(
"--lora_checkpoint_dir",
type=str,
default=None,
help="Path to the directory containing LoRA checkpoints",
)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument("--linear_threshold", type=float, default=0.025)
parser.add_argument("--linear_range", type=float, default=0.5)
args = parser.parse_args()
main(args)
@@ -0,0 +1,57 @@
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.utils import export_to_video, load_image, load_video
import argparse
from diffusers import FlowMatchEulerDiscreteScheduler
def main(args):
# Set the random seed for reproducibility
generator = torch.Generator("cuda").manual_seed(args.seed)
# do not invert
scheduler = FlowMatchEulerDiscreteScheduler()
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/"
)
pipe = MochiPipeline.from_pretrained(
args.model_path, transformer=transformer, scheduler=scheduler
)
pipe.enable_vae_tiling()
# pipe.to("cuda:1")
pipe.enable_model_cpu_offload()
# Generate videos from the input prompt
with torch.autocast("cuda", dtype=torch.bfloat16):
videos = pipe(
prompt=args.prompts,
height=args.height,
width=args.width,
num_frames=args.num_frames,
generator=generator,
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
).frames
for prompt, video in zip(args.prompts, videos):
export_to_video(video, args.output_path + f"_{prompt}.mp4", fps=30)
if __name__ == "__main__":
# arg parse
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--output_path", type=str, default="./outputs.mp4")
args = parser.parse_args()
main(args)
+733
View File
@@ -0,0 +1,733 @@
import argparse
from email.policy import strict
import logging
import math
import os
import shutil
from pathlib import Path
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
destroy_sequence_parallel_group,
get_sequence_parallel_state,
nccl_info,
)
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
import torch
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
import json
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.dataset_utils import LengthGroupedSampler
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
from diffusers.utils import convert_unet_state_dict_to_peft
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from diffusers.optimization import get_scheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from safetensors.torch import save_file, load_file
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
)
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
)
from fastvideo.utils.logging_ import main_print
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
import time
from collections import deque
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
):
"""
Compute the density for sampling the timesteps when doing SD3 training.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if weighting_scheme == "logit_normal":
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
u = torch.normal(
mean=logit_mean,
std=logit_std,
size=(batch_size,),
device="cpu",
generator=generator,
)
u = torch.nn.functional.sigmoid(u)
elif weighting_scheme == "mode":
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
else:
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
return u
def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
def train_one_step_mochi(
transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
latents = normalize_mochi_dit_input(latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
broadcast(timesteps)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
# if rank<=0:
# print("2222222222222222222222222222222222222222222222")
# print(type(latents_attention_mask))
# print(latents_attention_mask)
with torch.autocast("cuda", torch.bfloat16):
model_pred = transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
# if rank<=0:
# print("333333333333333333333333333333333333333333333333")
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (
torch.mean((model_pred.float() - target.float()) ** 2)
/ gradient_accumulation_steps
)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item()
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
dist.init_process_group("nccl")
torch.cuda.set_device(local_rank)
device = torch.cuda.current_device()
initialize_sequence_parallel_state(args.sp_size)
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weigths to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
# keep the master weight to float32
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=torch.float32
if args.master_weight_type == "fp32"
else torch.bfloat16,
)
if args.use_lora:
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
if args.resume_from_lora_checkpoint:
lora_state_dict = MochiPipeline.lora_state_dict(
args.resume_from_lora_checkpoint
)
transformer_state_dict = {
f'{k.replace("transformer.", "")}': v
for k, v in lora_state_dict.items()
if k.startswith("transformer.")
}
transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict)
incompatible_keys = set_peft_model_state_dict(
transformer, transformer_state_dict, adapter_name="default"
)
if incompatible_keys is not None:
# check only for unexpected keys
unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
if unexpected_keys:
main_print(
f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. "
)
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
fsdp_kwargs = get_dit_fsdp_kwargs(
args.fsdp_sharding_startegy,
args.use_lora,
args.use_cpu_offload,
args.master_weight_type,
)
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = ["MochiTransformerBlock"]
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
main_print(f"--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
noise_scheduler = FlowMatchEulerDiscreteScheduler()
params_to_optimize = transformer.parameters()
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer
)
main_print(f"optimizer: {optimizer}")
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps,
num_training_steps=args.max_train_steps,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
sampler = (
LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
)
if (args.group_frame or args.group_resolution)
else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
)
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader)
/ args.gradient_accumulation_steps
* args.sp_size
/ args.train_sp_batch_size
)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (
args.train_batch_size
* world_size
* args.gradient_accumulation_steps
/ args.sp_size
* args.train_sp_batch_size
)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError("resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
loss, grad_norm = train_one_step_mochi(
transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix(
{
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
}
)
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, step
)
else:
# Your existing checkpoint saving code
save_checkpoint(transformer, optimizer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(args, transformer, device, torch.bfloat16, step)
if args.use_lora:
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, args.max_train_steps
)
else:
save_checkpoint(
transformer, optimizer, rank, args.output_dir, args.max_train_steps
)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.999)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
parser.add_argument(
"--precondition_outputs",
action="store_true",
help="Whether to precondition the outputs of the model.",
)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--uncond_prompt_dir", type=str)
parser.add_argument(
"--validation_sampling_steps",
type=str,
default="64",
help="use ',' to split multi sampling steps",
)
parser.add_argument(
"--validation_guidance_scale",
type=str,
default="4.5",
help="use ',' to split multi scale",
)
parser.add_argument("--validation_steps", type=int, default=50)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument(
"--seed", type=int, default=None, help="A seed for reproducible training."
)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=(
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."
),
)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=(
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=(
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
)
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument(
"--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA."
)
parser.add_argument(
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
)
parser.add_argument("--fsdp_sharding_startegy", default="full")
parser.add_argument(
"--weighting_scheme",
type=str,
default="uniform",
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
)
parser.add_argument(
"--logit_mean",
type=float,
default=0.0,
help="mean to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--logit_std",
type=float,
default=1.0,
help="std to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--mode_scale",
type=float,
default=1.29,
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=(
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'
),
)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--weight_decay", type=float, default=0.01, help="Weight decay to apply."
)
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+282
View File
@@ -0,0 +1,282 @@
# import
import os
import json
import torch
from fastvideo.utils.logging_ import main_print
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
from safetensors.torch import save_file, load_file
import torch.distributed.checkpoint as dist_cp
from torch.distributed.checkpoint.default_planner import (
DefaultSavePlanner,
DefaultLoadPlanner,
)
from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
from torch.distributed.fsdp import FullOptimStateDictConfig
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=False):
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
cpu_state = model.state_dict()
optim_state = FSDP.optim_state_dict(
model,
optimizer,
)
# todo move to get_state_dict
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
if rank <= 0 and not discriminator:
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(model.config)
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
optimizer_path = os.path.join(save_dir, "optimizer.pt")
torch.save(optim_state, optimizer_path)
else:
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
def save_checkpoint_generator_discriminator(
model,
optimizer,
discriminator,
discriminator_optimizer,
rank,
output_dir,
step,
):
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
cpu_state = model.state_dict()
# todo move to get_state_dict
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
hf_weight_dir = os.path.join(save_dir, "hf_weights")
os.makedirs(hf_weight_dir, exist_ok=True)
# save using safetensors
if rank <= 0:
config_dict = dict(model.config)
config_path = os.path.join(hf_weight_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
weight_path = os.path.join(hf_weight_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
main_print(f"--> saved HF weight checkpoint at path {hf_weight_dir}")
model_weight_dir = os.path.join(save_dir, "model_weights_state")
os.makedirs(model_weight_dir, exist_ok=True)
model_optimizer_dir = os.path.join(save_dir, "model_optimizer_state")
os.makedirs(model_optimizer_dir, exist_ok=True)
with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
optim_state = FSDP.optim_state_dict(model, optimizer)
model_state = model.state_dict()
weight_state_dict = {"model": model_state}
dist_cp.save_state_dict(
state_dict=weight_state_dict,
storage_writer=dist_cp.FileSystemWriter(model_weight_dir),
planner=DefaultSavePlanner(),
)
optimizer_state_dict = {"optimizer": optim_state}
dist_cp.save_state_dict(
state_dict=optimizer_state_dict,
storage_writer=dist_cp.FileSystemWriter(model_optimizer_dir),
planner=DefaultSavePlanner(),
)
discriminator_fsdp_state_dir = os.path.join(save_dir, "discriminator_fsdp_state")
os.makedirs(discriminator_fsdp_state_dir, exist_ok=True)
with FSDP.state_dict_type(
discriminator,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
model_state = discriminator.state_dict()
state_dict = {"optimizer": optim_state, "model": model_state}
if rank <= 0:
discriminator_fsdp_state_fil = os.path.join(
discriminator_fsdp_state_dir, "discriminator_state.pt"
)
torch.save(state_dict, discriminator_fsdp_state_fil)
main_print("--> saved FSDP state checkpoint")
def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
weight_state_dict = {"model": model.state_dict()}
optim_state = load_sharded_optimizer_state_dict(
model_state_dict=weight_state_dict["model"],
optimizer_key="optimizer",
storage_reader=dist_cp.FileSystemReader(optimizer_dir),
)
optim_state = optim_state["optimizer"]
flattened_osd = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optim_state
)
optimizer.load_state_dict(flattened_osd)
dist_cp.load_state_dict(
state_dict=weight_state_dict,
storage_reader=dist_cp.FileSystemReader(model_dir),
planner=DefaultLoadPlanner(),
)
model_state = weight_state_dict["model"]
model.load_state_dict(model_state)
main_print(f"--> loaded model and optimizer from path {model_dir}")
return model, optimizer
def load_full_state_model(model, optimizer, checkpoint_file, rank):
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
discriminator_state = torch.load(checkpoint_file)
model_state = discriminator_state["model"]
if rank <= 0:
optim_state = discriminator_state["optimizer"]
else:
optim_state = None
model.load_state_dict(model_state)
discriminator_optim_state = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optim_state
)
optimizer.load_state_dict(discriminator_optim_state)
main_print(
f"--> loaded discriminator and discriminator optimizer from path {checkpoint_file}"
)
return model, optimizer
def resume_training_generator_discriminator(
model, optimizer, discriminator, discriminator_optimizer, checkpoint_dir, rank
):
step = int(checkpoint_dir.split("-")[-1])
model_weight_dir = os.path.join(checkpoint_dir, "model_weights_state")
model_optimizer_dir = os.path.join(checkpoint_dir, "model_optimizer_state")
model, optimizer = load_sharded_model(
model, optimizer, model_weight_dir, model_optimizer_dir
)
discriminator_ckpt_file = os.path.join(
checkpoint_dir, "discriminator_fsdp_state", "discriminator_state.pt"
)
discriminator, discriminator_optimizer = load_full_state_model(
discriminator, discriminator_optimizer, discriminator_ckpt_file, rank
)
return model, optimizer, discriminator, discriminator_optimizer, step
def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
weight_path = os.path.join(checkpoint_dir, "diffusion_pytorch_model.safetensors")
if discriminator:
weight_path = os.path.join(
checkpoint_dir, "discriminator_pytorch_model.safetensors"
)
model_weights = load_file(weight_path)
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
current_state = model.state_dict()
current_state.update(model_weights)
model.load_state_dict(current_state, strict=False)
if discriminator:
optim_path = os.path.join(checkpoint_dir, "discriminator_optimizer.pt")
else:
optim_path = os.path.join(checkpoint_dir, "optimizer.pt")
optimizer_state_dict = torch.load(optim_path, weights_only=False)
optim_state = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optimizer_state_dict
)
optimizer.load_state_dict(optim_state)
step = int(checkpoint_dir.split("-")[-1])
return model, optimizer, step
def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step):
with FSDP.state_dict_type(
transformer,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
full_state_dict = transformer.state_dict()
lora_optim_state = FSDP.optim_state_dict(
transformer,
optimizer,
)
if rank <= 0:
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save optimizer
optim_path = os.path.join(save_dir, "lora_optimizer.pt")
torch.save(lora_optim_state, optim_path)
# save lora weight
main_print(f"--> saving LoRA checkpoint at step {step}")
transformer_lora_layers = get_peft_model_state_dict(
model=transformer, state_dict=full_state_dict
)
MochiPipeline.save_lora_weights(
save_directory=save_dir,
transformer_lora_layers=transformer_lora_layers,
is_main_process=True,
)
# save config
lora_config = {
"step": step,
"lora_params": {
"lora_rank": transformer.config.lora_rank,
"lora_alpha": transformer.config.lora_alpha,
"target_modules": transformer.config.lora_target_modules,
},
}
config_path = os.path.join(save_dir, "lora_config.json")
with open(config_path, "w") as f:
json.dump(lora_config, f, indent=4)
main_print(f"--> LoRA checkpoint saved at step {step}")
def resume_lora_optimizer(transformer, checkpoint_dir, optimizer):
config_path = os.path.join(checkpoint_dir, "lora_config.json")
with open(config_path, "r") as f:
config_dict = json.load(f)
optim_path = os.path.join(checkpoint_dir, "lora_optimizer.pt")
optimizer_state_dict = torch.load(optim_path, weights_only=False)
optim_state = FSDP.optim_state_dict_to_load(
model=transformer, optim=optimizer, optim_state_dict=optimizer_state_dict
)
optimizer.load_state_dict(optim_state)
step = config_dict["step"]
main_print(f"--> Successfully resuming LoRA optimizer from step {step}")
return transformer, optimizer, step
+334
View File
@@ -0,0 +1,334 @@
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team
import torch
import torch.distributed as dist
from fastvideo.utils.parallel_states import nccl_info
from typing import Any, Tuple
from torch import Tensor
from torch.nn import Module
def broadcast(input_: torch.Tensor):
src = nccl_info.group_id * nccl_info.sp_size
dist.broadcast(input_, src=src, group=nccl_info.group)
def _all_to_all_4D(
input: torch.tensor, scatter_idx: int = 2, gather_idx: int = 1, group=None
) -> torch.tensor:
"""
all-to-all for QKV
Args:
input (torch.tensor): a tensor sharded along dim scatter dim
scatter_idx (int): default 1
gather_idx (int): default 2
group : torch process group
Returns:
torch.tensor: resharded tensor (bs, seqlen/P, hc, hs)
"""
assert (
input.dim() == 4
), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"
seq_world_size = dist.get_world_size(group)
if scatter_idx == 2 and gather_idx == 1:
# input (torch.tensor): a tensor sharded along dim 1 (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
bs, shard_seqlen, hc, hs = input.shape
seqlen = shard_seqlen * seq_world_size
shard_hc = hc // seq_world_size
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
# (bs, seqlen/P, hc, hs) -reshape-> (bs, seq_len/P, P, hc/P, hs) -transpose(0,2)-> (P, seq_len/P, bs, hc/P, hs)
input_t = (
input.reshape(bs, shard_seqlen, seq_world_size, shard_hc, hs)
.transpose(0, 2)
.contiguous()
)
output = torch.empty_like(input_t)
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
# (P, seq_len/P, bs, hc/P, hs) scatter seqlen -all2all-> (P, seq_len/P, bs, hc/P, hs) scatter head
if seq_world_size > 1:
dist.all_to_all_single(output, input_t, group=group)
torch.cuda.synchronize()
else:
output = input_t
# if scattering the seq-dim, transpose the heads back to the original dimension
output = output.reshape(seqlen, bs, shard_hc, hs)
# (seq_len, bs, hc/P, hs) -reshape-> (bs, seq_len, hc/P, hs)
output = output.transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
return output
elif scatter_idx == 1 and gather_idx == 2:
# input (torch.tensor): a tensor sharded along dim 1 (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
bs, seqlen, shard_hc, hs = input.shape
hc = shard_hc * seq_world_size
shard_seqlen = seqlen // seq_world_size
seq_world_size = dist.get_world_size(group)
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
# (bs, seqlen, hc/P, hs) -reshape-> (bs, P, seq_len/P, hc/P, hs) -transpose(0, 3)-> (hc/P, P, seqlen/P, bs, hs) -transpose(0, 1) -> (P, hc/P, seqlen/P, bs, hs)
input_t = (
input.reshape(bs, seq_world_size, shard_seqlen, shard_hc, hs)
.transpose(0, 3)
.transpose(0, 1)
.contiguous()
.reshape(seq_world_size, shard_hc, shard_seqlen, bs, hs)
)
output = torch.empty_like(input_t)
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
# (P, bs x hc/P, seqlen/P, hs) scatter seqlen -all2all-> (P, bs x seq_len/P, hc/P, hs) scatter head
if seq_world_size > 1:
dist.all_to_all_single(output, input_t, group=group)
torch.cuda.synchronize()
else:
output = input_t
# if scattering the seq-dim, transpose the heads back to the original dimension
output = output.reshape(hc, shard_seqlen, bs, hs)
# (hc, seqlen/N, bs, hs) -tranpose(0,2)-> (bs, seqlen/N, hc, hs)
output = output.transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError("scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")
class SeqAllToAll4D(torch.autograd.Function):
@staticmethod
def forward(
ctx: Any,
group: dist.ProcessGroup,
input: Tensor,
scatter_idx: int,
gather_idx: int,
) -> Tensor:
ctx.group = group
ctx.scatter_idx = scatter_idx
ctx.gather_idx = gather_idx
return _all_to_all_4D(input, scatter_idx, gather_idx, group=group)
@staticmethod
def backward(ctx: Any, *grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
return (
None,
SeqAllToAll4D.apply(
ctx.group, *grad_output, ctx.gather_idx, ctx.scatter_idx
),
None,
None,
)
def all_to_all_4D(
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim, gather_dim)
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,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return _AllToAll.apply(input_, nccl_info.group, scatter_dim, gather_dim)
class _AllGather(torch.autograd.Function):
"""All-gather communication with autograd support.
Args:
input_: input tensor
dim: dimension along which to concatenate
"""
@staticmethod
def forward(ctx, input_, dim):
ctx.dim = dim
world_size = nccl_info.sp_size
group = nccl_info.group
input_size = list(input_.size())
ctx.input_size = input_size[dim]
tensor_list = [torch.empty_like(input_) for _ in range(world_size)]
input_ = input_.contiguous()
dist.all_gather(tensor_list, input_, group=group)
output = torch.cat(tensor_list, dim=dim)
return output
@staticmethod
def backward(ctx, grad_output):
world_size = nccl_info.sp_size
rank = nccl_info.rank_within_group
dim = ctx.dim
input_size = ctx.input_size
sizes = [input_size] * world_size
grad_input_list = torch.split(grad_output, sizes, dim=dim)
grad_input = grad_input_list[rank]
return grad_input, None
def all_gather(input_: torch.Tensor, dim: int = 1):
"""Performs an all-gather operation on the input tensor along the specified dimension.
Args:
input_ (torch.Tensor): Input tensor of shape [B, H, S, D].
dim (int, optional): Dimension along which to concatenate. Defaults to 1.
Returns:
torch.Tensor: Output tensor after all-gather operation, concatenated along 'dim'.
"""
return _AllGather.apply(input_, dim)
def prepare_sequence_parallel_data(
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
):
if nccl_info.sp_size == 1:
return (
hidden_states,
encoder_hidden_states,
attention_mask,
encoder_attention_mask,
)
def prepare(
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
):
hidden_states = all_to_all(hidden_states, scatter_dim=2, gather_dim=0)
encoder_hidden_states = all_to_all(
encoder_hidden_states, scatter_dim=1, gather_dim=0
)
attention_mask = all_to_all(attention_mask, scatter_dim=1, gather_dim=0)
encoder_attention_mask = all_to_all(
encoder_attention_mask, scatter_dim=1, gather_dim=0
)
return (
hidden_states,
encoder_hidden_states,
attention_mask,
encoder_attention_mask,
)
sp_size = nccl_info.sp_size
frame = hidden_states.shape[2]
assert frame % sp_size == 0, "frame should be a multiple of sp_size"
(
hidden_states,
encoder_hidden_states,
attention_mask,
encoder_attention_mask,
) = prepare(
hidden_states,
encoder_hidden_states.repeat(1, sp_size, 1),
attention_mask.repeat(1, sp_size, 1, 1),
encoder_attention_mask.repeat(1, sp_size),
)
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
def sp_parallel_dataloader_wrapper(
dataloader, device, train_batch_size, sp_size, train_sp_batch_size
):
while True:
for data_item in dataloader:
latents, cond, attn_mask, cond_mask = data_item
latents = latents.to(device)
cond = cond.to(device)
attn_mask = attn_mask.to(device)
cond_mask = cond_mask.to(device)
frame = latents.shape[2]
if frame == 1:
yield latents, cond, attn_mask, cond_mask
else:
latents, cond, attn_mask, cond_mask = prepare_sequence_parallel_data(
latents, cond, attn_mask, cond_mask
)
assert (
train_batch_size * sp_size >= train_sp_batch_size
), "train_batch_size * sp_size should be greater than train_sp_batch_size"
for iter in range(train_batch_size * sp_size // train_sp_batch_size):
st_idx = iter * train_sp_batch_size
ed_idx = (iter + 1) * train_sp_batch_size
encoder_hidden_states = cond[st_idx:ed_idx]
attention_mask = attn_mask[st_idx:ed_idx]
encoder_attention_mask = cond_mask[st_idx:ed_idx]
yield (
latents[st_idx:ed_idx],
encoder_hidden_states,
attention_mask,
encoder_attention_mask,
)
+377
View File
@@ -0,0 +1,377 @@
import math
from einops import rearrange
import decord
from torch.nn import functional as F
import torch
from typing import Optional
import torch.utils
import torch.utils.data
import torch
from torch.utils.data import Sampler
from typing import List
from collections import Counter
import random
IMG_EXTENSIONS = [".jpg", ".JPG", ".jpeg", ".JPEG", ".png", ".PNG"]
def is_image_file(filename):
return any(filename.endswith(extension) for extension in IMG_EXTENSIONS)
class DecordInit(object):
"""Using Decord(https://github.com/dmlc/decord) to initialize the video_reader."""
def __init__(self, num_threads=1):
self.num_threads = num_threads
self.ctx = decord.cpu(0)
def __call__(self, filename):
"""Perform the Decord initialization.
Args:
results (dict): The resulting dict to be modified and passed
to the next transform in pipeline.
"""
reader = decord.VideoReader(
filename, ctx=self.ctx, num_threads=self.num_threads
)
return reader
def __repr__(self):
repr_str = (
f"{self.__class__.__name__}("
f"sr={self.sr},"
f"num_threads={self.num_threads})"
)
return repr_str
def pad_to_multiple(number, ds_stride):
remainder = number % ds_stride
if remainder == 0:
return number
else:
padding = ds_stride - remainder
return number + padding
# TODO
class Collate:
def __init__(self, args):
self.batch_size = args.train_batch_size
self.group_frame = args.group_frame
self.group_resolution = args.group_resolution
self.max_height = args.max_height
self.max_width = args.max_width
self.ae_stride = args.ae_stride
self.ae_stride_t = args.ae_stride_t
self.ae_stride_thw = (self.ae_stride_t, self.ae_stride, self.ae_stride)
self.patch_size = args.patch_size
self.patch_size_t = args.patch_size_t
self.num_frames = args.num_frames
self.use_image_num = args.use_image_num
self.max_thw = (self.num_frames, self.max_height, self.max_width)
def package(self, batch):
batch_tubes = [i["pixel_values"] for i in batch] # b [c t h w]
input_ids = [i["input_ids"] for i in batch] # b [1 l]
cond_mask = [i["cond_mask"] for i in batch] # b [1 l]
return batch_tubes, input_ids, cond_mask
def __call__(self, batch):
batch_tubes, input_ids, cond_mask = self.package(batch)
ds_stride = self.ae_stride * self.patch_size
t_ds_stride = self.ae_stride_t * self.patch_size_t
pad_batch_tubes, attention_mask, input_ids, cond_mask = self.process(
batch_tubes,
input_ids,
cond_mask,
t_ds_stride,
ds_stride,
self.max_thw,
self.ae_stride_thw,
)
assert not torch.any(torch.isnan(pad_batch_tubes)), "after pad_batch_tubes"
return pad_batch_tubes, attention_mask, input_ids, cond_mask
def process(
self,
batch_tubes,
input_ids,
cond_mask,
t_ds_stride,
ds_stride,
max_thw,
ae_stride_thw,
):
# pad to max multiple of ds_stride
batch_input_size = [i.shape for i in batch_tubes] # [(c t h w), (c t h w)]
assert len(batch_input_size) == self.batch_size
if self.group_frame or self.group_resolution or self.batch_size == 1: #
len_each_batch = batch_input_size
idx_length_dict = dict([*zip(list(range(self.batch_size)), len_each_batch)])
count_dict = Counter(len_each_batch)
if len(count_dict) != 1:
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
pick_length = sorted_by_value[-1][0] # the highest frequency
candidate_batch = [
idx
for idx, length in idx_length_dict.items()
if length == pick_length
]
random_select_batch = [
random.choice(candidate_batch)
for _ in range(len(len_each_batch) - len(candidate_batch))
]
print(
batch_input_size,
idx_length_dict,
count_dict,
sorted_by_value,
pick_length,
candidate_batch,
random_select_batch,
)
pick_idx = candidate_batch + random_select_batch
batch_tubes = [batch_tubes[i] for i in pick_idx]
batch_input_size = [
i.shape for i in batch_tubes
] # [(c t h w), (c t h w)]
input_ids = [input_ids[i] for i in pick_idx] # b [1, l]
cond_mask = [cond_mask[i] for i in pick_idx] # b [1, l]
for i in range(1, self.batch_size):
assert batch_input_size[0] == batch_input_size[i]
max_t = max([i[1] for i in batch_input_size])
max_h = max([i[2] for i in batch_input_size])
max_w = max([i[3] for i in batch_input_size])
else:
max_t, max_h, max_w = max_thw
pad_max_t, pad_max_h, pad_max_w = (
pad_to_multiple(max_t - 1 + self.ae_stride_t, t_ds_stride),
pad_to_multiple(max_h, ds_stride),
pad_to_multiple(max_w, ds_stride),
)
pad_max_t = pad_max_t + 1 - self.ae_stride_t
each_pad_t_h_w = [
[pad_max_t - i.shape[1], pad_max_h - i.shape[2], pad_max_w - i.shape[3]]
for i in batch_tubes
]
pad_batch_tubes = [
F.pad(im, (0, pad_w, 0, pad_h, 0, pad_t), value=0)
for (pad_t, pad_h, pad_w), im in zip(each_pad_t_h_w, batch_tubes)
]
pad_batch_tubes = torch.stack(pad_batch_tubes, dim=0)
max_tube_size = [pad_max_t, pad_max_h, pad_max_w]
max_latent_size = [
((max_tube_size[0] - 1) // ae_stride_thw[0] + 1),
max_tube_size[1] // ae_stride_thw[1],
max_tube_size[2] // ae_stride_thw[2],
]
valid_latent_size = [
[
int(math.ceil((i[1] - 1) / ae_stride_thw[0])) + 1,
int(math.ceil(i[2] / ae_stride_thw[1])),
int(math.ceil(i[3] / ae_stride_thw[2])),
]
for i in batch_input_size
]
attention_mask = [
F.pad(
torch.ones(i, dtype=pad_batch_tubes.dtype),
(
0,
max_latent_size[2] - i[2],
0,
max_latent_size[1] - i[1],
0,
max_latent_size[0] - i[0],
),
value=0,
)
for i in valid_latent_size
]
attention_mask = torch.stack(attention_mask) # b t h w
if self.batch_size == 1 or self.group_frame or self.group_resolution:
assert torch.all(attention_mask.bool())
input_ids = torch.stack(input_ids) # b 1 l
cond_mask = torch.stack(cond_mask) # b 1 l
return pad_batch_tubes, attention_mask, input_ids, cond_mask
def split_to_even_chunks(indices, lengths, num_chunks, batch_size):
"""
Split a list of indices into `chunks` chunks of roughly equal lengths.
"""
if len(indices) % num_chunks != 0:
chunks = [indices[i::num_chunks] for i in range(num_chunks)]
else:
num_indices_per_chunk = len(indices) // num_chunks
chunks = [[] for _ in range(num_chunks)]
chunks_lengths = [0 for _ in range(num_chunks)]
for index in indices:
shortest_chunk = chunks_lengths.index(min(chunks_lengths))
chunks[shortest_chunk].append(index)
chunks_lengths[shortest_chunk] += lengths[index]
if len(chunks[shortest_chunk]) == num_indices_per_chunk:
chunks_lengths[shortest_chunk] = float("inf")
# return chunks
pad_chunks = []
for idx, chunk in enumerate(chunks):
if batch_size != len(chunk):
assert batch_size > len(chunk)
if len(chunk) != 0:
chunk = chunk + [
random.choice(chunk) for _ in range(batch_size - len(chunk))
]
else:
chunk = random.choice(pad_chunks)
print(chunks[idx], "->", chunk)
pad_chunks.append(chunk)
return pad_chunks
def group_frame_fun(indices, lengths):
# sort by num_frames
indices.sort(key=lambda i: lengths[i], reverse=True)
return indices
def megabatch_frame_alignment(megabatches, lengths):
aligned_magabatches = []
for _, megabatch in enumerate(megabatches):
assert len(megabatch) != 0
len_each_megabatch = [lengths[i] for i in megabatch]
idx_length_dict = dict([*zip(megabatch, len_each_megabatch)])
count_dict = Counter(len_each_megabatch)
# mixed frame length, align megabatch inside
if len(count_dict) != 1:
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
pick_length = sorted_by_value[-1][0] # the highest frequency
candidate_batch = [
idx for idx, length in idx_length_dict.items() if length == pick_length
]
random_select_batch = [
random.choice(candidate_batch)
for i in range(len(idx_length_dict) - len(candidate_batch))
]
aligned_magabatch = candidate_batch + random_select_batch
aligned_magabatches.append(aligned_magabatch)
# already aligned megabatches
else:
aligned_magabatches.append(megabatch)
return aligned_magabatches
def get_length_grouped_indices(
lengths,
batch_size,
world_size,
generator=None,
group_frame=False,
group_resolution=False,
seed=42,
):
# We need to use torch for the random part as a distributed sampler will set the random seed for torch.
if generator is None:
generator = torch.Generator().manual_seed(
seed
) # every rank will generate a fixed order but random index
indices = torch.randperm(len(lengths), generator=generator).tolist()
# sort dataset according to frame
indices = group_frame_fun(indices, lengths)
# chunk dataset to megabatches
megabatch_size = world_size * batch_size
megabatches = [
indices[i : i + megabatch_size] for i in range(0, len(lengths), megabatch_size)
]
# make sure the length in each magabatch is align with each other
megabatches = megabatch_frame_alignment(megabatches, lengths)
# aplit aligned megabatch into batches
megabatches = [
split_to_even_chunks(megabatch, lengths, world_size, batch_size)
for megabatch in megabatches
]
# random megabatches to do video-image mix training
indices = torch.randperm(len(megabatches), generator=generator).tolist()
shuffled_megabatches = [megabatches[i] for i in indices]
# expand indices and return
return [
i for megabatch in shuffled_megabatches for batch in megabatch for i in batch
]
class LengthGroupedSampler(Sampler):
r"""
Sampler that samples indices in a way that groups together features of the dataset of roughly the same length while
keeping a bit of randomness.
"""
def __init__(
self,
batch_size: int,
rank: int,
world_size: int,
lengths: Optional[List[int]] = None,
group_frame=False,
group_resolution=False,
generator=None,
):
if lengths is None:
raise ValueError("Lengths must be provided.")
self.batch_size = batch_size
self.rank = rank
self.world_size = world_size
self.lengths = lengths
self.group_frame = group_frame
self.group_resolution = group_resolution
self.generator = generator
def __len__(self):
return len(self.lengths)
def __iter__(self):
indices = get_length_grouped_indices(
self.lengths,
self.batch_size,
self.world_size,
group_frame=self.group_frame,
group_resolution=self.group_resolution,
generator=self.generator,
)
def distributed_sampler(lst, rank, batch_size, world_size):
result = []
index = rank * batch_size
while index < len(lst):
result.extend(lst[index : index + batch_size])
index += batch_size * world_size
return result
indices = distributed_sampler(
indices, self.rank, self.batch_size, self.world_size
)
return iter(indices)
+148
View File
@@ -0,0 +1,148 @@
from sympy import use
import torch
import os
import torch.distributed as dist
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper,
CheckpointImpl,
apply_activation_checkpointing,
)
from peft.utils.other import fsdp_auto_wrap_policy
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig, # general model non-sharded, non-flattened params
LocalStateDictConfig, # flattened params, usable only by FSDP
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
)
from fastvideo.utils.load import get_no_split_modules
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
from functools import partial
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
import functools
non_reentrant_wrapper = partial(
checkpoint_wrapper,
checkpoint_impl=CheckpointImpl.NO_REENTRANT,
)
check_fn = lambda submodule: isinstance(submodule, MochiTransformerBlock)
def apply_fsdp_checkpointing(model,no_split_modules, p=1):
# https://github.com/foundation-model-stack/fms-fsdp/blob/408c7516d69ea9b6bcd4c0f5efab26c0f64b3c2d/fms_fsdp/policies/ac_handler.py#L16
"""apply activation checkpointing to model
returns None as model is updated directly
"""
print(f"--> applying fdsp activation checkpointing...")
block_idx = 0
cut_off = 1 / 2
# when passing p as a fraction number (e.g. 1/3), it will be interpreted
# as a string in argv, thus we need eval("1/3") here for fractions.
p = eval(p) if isinstance(p, str) else p
def selective_checkpointing(submodule):
nonlocal block_idx
nonlocal cut_off
if isinstance(submodule, no_split_modules):
block_idx += 1
if block_idx * p >= cut_off:
cut_off += 1
return True
return False
apply_activation_checkpointing(
model,
checkpoint_wrapper_fn=non_reentrant_wrapper,
check_fn=selective_checkpointing,
)
def get_mixed_precision(master_weight_type="fp32"):
weight_type = torch.float32 if master_weight_type == "fp32" else torch.bfloat16
mixed_precision = MixedPrecision(
param_dtype=weight_type,
# Gradient communication precision.
reduce_dtype=weight_type,
# Buffer precision.
buffer_dtype=weight_type,
cast_forward_inputs=False,
)
return mixed_precision
def get_dit_fsdp_kwargs(
transformer, sharding_strategy, use_lora=False, cpu_offload=False, master_weight_type="fp32"
):
no_split_modules = get_no_split_modules(transformer)
if use_lora:
auto_wrap_policy = fsdp_auto_wrap_policy
else:
auto_wrap_policy = functools.partial(
transformer_auto_wrap_policy,
transformer_layer_cls=no_split_modules,
)
# we use float32 for fsdp but autocast during training
mixed_precision = get_mixed_precision(master_weight_type)
if sharding_strategy == "full":
sharding_strategy = ShardingStrategy.FULL_SHARD
elif sharding_strategy == "hybrid_full":
sharding_strategy = ShardingStrategy.HYBRID_SHARD
elif sharding_strategy == "none":
sharding_strategy = ShardingStrategy.NO_SHARD
auto_wrap_policy = None
elif sharding_strategy == "hybrid_zero2":
sharding_strategy = ShardingStrategy._HYBRID_SHARD_ZERO2
device_id = torch.cuda.current_device()
cpu_offload = (
torch.distributed.fsdp.CPUOffload(offload_params=True) if cpu_offload else None
)
fsdp_kwargs = {
"auto_wrap_policy": auto_wrap_policy,
"mixed_precision": mixed_precision,
"sharding_strategy": sharding_strategy,
"device_id": device_id,
"limit_all_gathers": True,
"cpu_offload": cpu_offload,
}
# Add LoRA-specific settings when LoRA is enabled
if use_lora:
fsdp_kwargs.update(
{
"use_orig_params": False, # Required for LoRA memory savings
"sync_module_states": True,
}
)
return fsdp_kwargs, no_split_modules
def get_discriminator_fsdp_kwargs(master_weight_type="fp32"):
auto_wrap_policy = None
# Use existing mixed precision settings
mixed_precision = get_mixed_precision(master_weight_type)
sharding_strategy = ShardingStrategy.NO_SHARD
device_id = torch.cuda.current_device()
fsdp_kwargs = {
"auto_wrap_policy": auto_wrap_policy,
"mixed_precision": mixed_precision,
"sharding_strategy": sharding_strategy,
"device_id": device_id,
"limit_all_gathers": True,
}
return fsdp_kwargs
+320
View File
@@ -0,0 +1,320 @@
import torch
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
from fastvideo.models.hunyuan.modules.models import HYVideoDiffusionTransformer, MMDoubleStreamBlock, MMSingleStreamBlock
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from diffusers import AutoencoderKLMochi
from transformers import T5EncoderModel, AutoTokenizer
import os
from torch import nn
# Path
from pathlib import Path
import torch.nn.functional as F
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.utils.logging_ import main_print
hunyuan_config = {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
"guidance_embed": True,
}
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"dit-llm-encode": {
"template": PROMPT_TEMPLATE_ENCODE,
"crop_start": 36,
},
"dit-llm-encode-video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
},
}
class HunyuanTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
text_len = 256
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get("crop_start", 0 )
max_length = text_len + crop_start
# prompt_template
prompt_template = PROMPT_TEMPLATE["dit-llm-encode"]
# prompt_template_video
prompt_template_video = PROMPT_TEMPLATE["dit-llm-encode-video"]
text_encoder_path = os.path.join(pretrained_model_name_or_path, "text_encoder")
self.text_encoder = TextEncoder(
text_encoder_type="llm",
text_encoder_path=text_encoder_path,
max_length=max_length,
text_encoder_precision="fp16",
tokenizer_type="llm",
prompt_template=prompt_template,
prompt_template_video=prompt_template_video,
hidden_state_skip_layer=2,
apply_final_norm=False,
reproduce=False,
logger=None,
device=device,
)
text_encoder_path_2 = os.path.join(pretrained_model_name_or_path, "text_encoder_2")
self.text_encoder_2 = TextEncoder(
text_encoder_type="clipL",
text_encoder_path=text_encoder_path_2,
max_length=77,
text_encoder_precision="fp16",
tokenizer_type="clipL",
reproduce=False,
logger=None,
device=device,
)
def encode_(self, prompt, text_encoder, clip_skip=None):
# TODO
device = self.text_encoder.device
data_type = "video"
num_videos_per_prompt = 1
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
if clip_skip is None:
prompt_outputs = text_encoder.encode(
text_inputs, data_type="video", device=device
)
prompt_embeds = prompt_outputs.hidden_state
else:
prompt_outputs = text_encoder.encode(
text_inputs,
output_hidden_states=True,
data_type=data_type,
device=device,
)
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
prompt_embeds
)
attention_mask = prompt_outputs.attention_mask
if attention_mask is not None:
attention_mask = attention_mask.to(device)
bs_embed, seq_len = attention_mask.shape
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
attention_mask = attention_mask.view(
bs_embed * num_videos_per_prompt, seq_len
)
if text_encoder is not None:
prompt_embeds_dtype = text_encoder.dtype
elif self.transformer is not None:
prompt_embeds_dtype = self.transformer.dtype
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
if prompt_embeds.ndim == 2:
bs_embed, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
else:
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, seq_len, -1
)
return (prompt_embeds, attention_mask)
def encode_prompt(self, prompt):
prompt_embeds, attention_mask = self.encode_(prompt, self.text_encoder)
prompt_embeds_2, attention_mask_2 = self.encode_(prompt, self.text_encoder_2)
prompt_embeds_2 = F.pad(prompt_embeds_2, (0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]), value=0).unsqueeze(1)
prompt_embeds = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
return prompt_embeds, attention_mask
class MochiTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
self.text_encoder = T5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path, "text_encoder")).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "text_encoder"))
self.max_sequence_length = 256
def encode_prompt(self, prompt):
device = self.text_encoder.device
dtype = self.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(
prompt, padding="longest", return_tensors="pt"
).input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
removed_text = self.tokenizer.batch_decode(
untruncated_ids[:, self.max_sequence_length - 1 : -1]
)
main_print(
f"Truncated text input: {prompt} to: {removed_text} for model input."
)
prompt_embeds = self.text_encoder(
text_input_ids.to(device), attention_mask=prompt_attention_mask
)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(
batch_size , seq_len, -1
)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
return prompt_embeds, prompt_attention_mask
def load_hunyuan_state_dict(model, dit_model_name_or_path):
load_key = "module"
model_path = dit_model_name_or_path
bare_model = "unknown"
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage, weights_only=True)
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}."
)
model.load_state_dict(state_dict, strict=True)
return model
def load_transformer(model_type,dit_model_name_or_path, pretrained_model_name_or_path, master_weight_type):
if model_type == "mochi":
if dit_model_name_or_path:
transformer = MochiTransformer3DModel.from_pretrained(
dit_model_name_or_path,
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
else:
transformer = MochiTransformer3DModel.from_pretrained(
pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "hunyuan":
transformer = HYVideoDiffusionTransformer(
in_channels=16,
out_channels=16,
**hunyuan_config,
dtype=master_weight_type,
)
transformer = load_hunyuan_state_dict(transformer, dit_model_name_or_path)
else:
raise ValueError(f"Unsupported model type: {model_type}")
return transformer
def load_vae(model_type, pretrained_model_name_or_path):
weight_dtype = torch.float32
if model_type == "mochi":
vae = AutoencoderKLMochi.from_pretrained(
pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype
).to("cuda")
autocast_type = torch.bfloat16
fps = 30
elif model_type == "hunyuan":
vae_precision = torch.float32
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
config = AutoencoderKLCausal3D.load_config(vae_path)
vae = AutoencoderKLCausal3D.from_config(config)
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
ckpt = torch.load(vae_ckpt, map_location=vae.device, weights_only=True)
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
vae.load_state_dict(ckpt)
vae = vae.to(dtype=vae_precision)
vae.requires_grad_(False)
vae = vae.to("cuda")
vae.eval()
autocast_type = torch.float32
fps = 24
return vae, autocast_type, fps
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
if model_type == "mochi":
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "hunyuan":
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
else:
raise ValueError(f"Unsupported model type: {model_type}")
return text_encoder
def get_no_split_modules(transformer):
# if of type MochiTransformer3DModel
if isinstance(transformer, MochiTransformer3DModel):
return (MochiTransformerBlock,)
elif isinstance(transformer, HYVideoDiffusionTransformer):
return (MMDoubleStreamBlock, MMSingleStreamBlock)
else:
raise ValueError(f"Unsupported transformer type: {type(transformer)}")
if __name__ == "__main__":
# test encode prompt
device = torch.cuda.current_device()
pretrained_model_name_or_path = "data/hunyuan"
text_encoder = load_text_encoder("hunyuan", pretrained_model_name_or_path, device)
prompt = "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."
prompt_embeds, attention_mask = text_encoder.encode_prompt(prompt)
+24
View File
@@ -0,0 +1,24 @@
import sys
import pdb
import os
def main_print(content):
if int(os.environ["LOCAL_RANK"]) <= 0:
print(content)
# ForkedPdb().set_trace()
class ForkedPdb(pdb.Pdb):
"""A Pdb subclass that may be used
from a forked multiprocessing child
"""
def interaction(self, *args, **kwargs):
_stdin = sys.stdin
try:
sys.stdin = open("/dev/stdin")
pdb.Pdb.interaction(self, *args, **kwargs)
finally:
sys.stdin = _stdin
+77
View File
@@ -0,0 +1,77 @@
from accelerate.logging import get_logger
import torch
logger = get_logger(__name__)
def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
# Optimizer creation
supported_optimizers = ["adam", "adamw", "prodigy"]
if args.optimizer not in supported_optimizers:
logger.warning(
f"Unsupported choice of optimizer: {args.optimizer}. Supported optimizers include {supported_optimizers}. Defaulting to AdamW"
)
args.optimizer = "adamw"
if args.use_8bit_adam and not (args.optimizer.lower() not in ["adam", "adamw"]):
logger.warning(
f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
f"set to {args.optimizer.lower()}"
)
if args.use_8bit_adam:
try:
import bitsandbytes as bnb
except ImportError:
raise ImportError(
"To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`."
)
if args.optimizer.lower() == "adamw":
optimizer_class = (
bnb.optim.AdamW8bit if args.use_8bit_adam else torch.optim.AdamW
)
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
)
elif args.optimizer.lower() == "adam":
optimizer_class = bnb.optim.Adam8bit if args.use_8bit_adam else torch.optim.Adam
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
)
elif args.optimizer.lower() == "prodigy":
try:
import prodigyopt
except ImportError:
raise ImportError(
"To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`"
)
optimizer_class = prodigyopt.Prodigy
if args.learning_rate <= 0.1:
logger.warning(
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
)
optimizer = optimizer_class(
params_to_optimize,
lr=args.learning_rate,
betas=(args.adam_beta1, args.adam_beta2),
beta3=args.prodigy_beta3,
weight_decay=args.adam_weight_decay,
eps=args.adam_epsilon,
decouple=args.prodigy_decouple,
use_bias_correction=args.prodigy_use_bias_correction,
safeguard_warmup=args.prodigy_safeguard_warmup,
)
return optimizer
+63
View File
@@ -0,0 +1,63 @@
import torch
import torch.distributed as dist
import os
class COMM_INFO:
def __init__(self):
self.group = None
self.sp_size = 1
self.global_rank = 0
self.rank_within_group = 0
self.group_id = 0
nccl_info = COMM_INFO()
_SEQUENCE_PARALLEL_STATE = False
def initialize_sequence_parallel_state(sequence_parallel_size):
global _SEQUENCE_PARALLEL_STATE
if sequence_parallel_size > 1:
_SEQUENCE_PARALLEL_STATE = True
initialize_sequence_parallel_group(sequence_parallel_size)
else:
nccl_info.sp_size = 1
nccl_info.global_rank = int(os.getenv("RANK", "0"))
nccl_info.rank_within_group = 0
nccl_info.group_id = int(os.getenv("RANK", "0"))
def set_sequence_parallel_state(state):
global _SEQUENCE_PARALLEL_STATE
_SEQUENCE_PARALLEL_STATE = state
def get_sequence_parallel_state():
return _SEQUENCE_PARALLEL_STATE
def initialize_sequence_parallel_group(sequence_parallel_size):
"""Initialize the sequence parallel group."""
rank = int(os.getenv("RANK", "0"))
world_size = int(os.getenv("WORLD_SIZE", "1"))
assert (
world_size % sequence_parallel_size == 0
), "world_size must be divisible by sequence_parallel_size, but got world_size: {}, sequence_parallel_size: {}".format(
world_size, sequence_parallel_size
)
nccl_info.sp_size = sequence_parallel_size
nccl_info.global_rank = rank
num_sequence_parallel_groups: int = world_size // sequence_parallel_size
for i in range(num_sequence_parallel_groups):
ranks = range(i * sequence_parallel_size, (i + 1) * sequence_parallel_size)
group = dist.new_group(ranks)
if rank in ranks:
nccl_info.group = group
nccl_info.rank_within_group = rank - i * sequence_parallel_size
nccl_info.group_id = i
def destroy_sequence_parallel_group():
"""Destroy the sequence parallel group."""
dist.destroy_process_group()
+354
View File
@@ -0,0 +1,354 @@
from typing import Optional, Union, List
import numpy as np
import torch
from einops import rearrange
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.models.mochi_hf.pipeline_mochi import (
linear_quadratic_schedule,
retrieve_timesteps,
)
from tqdm import tqdm
from diffusers.video_processor import VideoProcessor
from diffusers import (
FlowMatchEulerDiscreteScheduler,
AutoencoderKLMochi,
)
from fastvideo.utils.logging_ import main_print
from fastvideo.distill.solver import PCMFMScheduler
from diffusers.utils import export_to_video
import os
import wandb
import gc
from fastvideo.utils.load import load_vae
def prepare_latents(
batch_size,
num_channels_latents,
height,
width,
num_frames,
dtype,
device,
generator,
vae_spatial_scale_factor,
vae_temporal_scale_factor,
):
height = height // vae_spatial_scale_factor
width = width // vae_spatial_scale_factor
num_frames = (num_frames - 1) // vae_temporal_scale_factor + 1
shape = (batch_size, num_channels_latents, num_frames, height, width)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
def sample_validation_video(
transformer,
vae,
scheduler,
scheduler_type="euler",
height: Optional[int] = None,
width: Optional[int] = None,
num_frames: int = 16,
num_inference_steps: int = 28,
timesteps: List[int] = None,
guidance_scale: float = 4.5,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
output_type: Optional[str] = "pil",
vae_spatial_scale_factor=8,
vae_temporal_scale_factor=6,
num_channels_latents=12
):
device = vae.device
batch_size = prompt_embeds.shape[0]
do_classifier_free_guidance = guidance_scale > 1.0
if do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
prompt_attention_mask = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask], dim=0
)
# 4. Prepare latent variables
# TODO: Remove hardcore
latents = prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
height,
width,
num_frames,
prompt_embeds.dtype,
device,
generator,
vae_spatial_scale_factor,
vae_temporal_scale_factor,
)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(
latents, "b t (n s) h w -> b t n s h w", n=world_size
).contiguous()
latents = latents[:, :, rank, :, :, :]
# 5. Prepare timestep
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
threshold_noise = 0.025
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
sigmas = np.array(sigmas)
if scheduler_type == "euler":
timesteps, num_inference_steps = retrieve_timesteps(
scheduler,
num_inference_steps,
device,
timesteps,
sigmas,
)
else:
timesteps, num_inference_steps = retrieve_timesteps(
scheduler,
num_inference_steps,
device,
)
num_warmup_steps = max(len(timesteps) - num_inference_steps * scheduler.order, 0)
# 6. Denoising loop
# with self.progress_bar(total=num_inference_steps) as progress_bar:
# write with tqdm instead
# only enable if nccl_info.global_rank == 0
with tqdm(
total=num_inference_steps,
disable=nccl_info.rank_within_group != 0,
desc="Validation sampling...",
) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = (
torch.cat([latents] * 2) if do_classifier_free_guidance else latents
)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
with torch.autocast("cuda", dtype=torch.bfloat16):
noise_pred = transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask,
return_dict=False,
)[0]
# Mochi CFG + Sampling runs in FP32
noise_pred = noise_pred.to(torch.float32)
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (
noise_pred_text - noise_pred_uncond
)
# compute the previous noisy sample x_t -> x_t-1
latents_dtype = latents.dtype
latents = scheduler.step(
noise_pred, t, latents.to(torch.float32), return_dict=False
)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
if torch.backends.mps.is_available():
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
latents = latents.to(latents_dtype)
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and (i + 1) % scheduler.order == 0
):
progress_bar.update()
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
if output_type == "latent":
video = latents
else:
# unscale/denormalize the latents
# denormalize with the mean and std if available and not None
has_latents_mean = (
hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None
)
has_latents_std = (
hasattr(vae.config, "latents_std") and vae.config.latents_std is not None
)
if has_latents_mean and has_latents_std:
latents_mean = (
torch.tensor(vae.config.latents_mean)
.view(1, 12, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = (
torch.tensor(vae.config.latents_std)
.view(1, 12, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
else:
latents = latents / vae.config.scaling_factor
with torch.autocast("cuda", dtype=vae.dtype):
video = vae.decode(latents, return_dict=False)[0]
video_processor = VideoProcessor(vae_scale_factor=vae_spatial_scale_factor)
video = video_processor.postprocess_video(video, output_type=output_type)
return (video,)
@torch.no_grad()
@torch.autocast("cuda", dtype=torch.bfloat16)
def log_validation(
args,
transformer,
device,
weight_dtype, # TODO
global_step,
scheduler_type="euler",
shift=1.0,
num_euler_timesteps=100,
linear_quadratic_threshold=0.025,
linear_range=0.5,
ema=False,
):
# TODO
print(f"Running validation....\n")
if args.model_type == "mochi":
vae_spatial_scale_factor=8
vae_temporal_scale_factor=6
num_channels_latents=12
elif args.model_type == "hunyuan":
vae_spatial_scale_factor=8
vae_temporal_scale_factor=4
num_channels_latents=16
else:
raise ValueError(f"Model type {args.model_type} not supported")
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
vae.enable_tiling()
if scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
linear_quadraic = True if scheduler_type == "pcm_linear_quadratic" else False
scheduler = PCMFMScheduler(
1000,
shift,
num_euler_timesteps,
linear_quadraic,
linear_quadratic_threshold,
linear_range,
)
# args.validation_prompt_dir
validation_guidance_scale_ls = args.validation_guidance_scale.split(",")
validation_guidance_scale_ls = [
float(scale) for scale in validation_guidance_scale_ls
]
for validation_sampling_step in args.validation_sampling_steps.split(","):
validation_sampling_step = int(validation_sampling_step)
for validation_guidance_scale in validation_guidance_scale_ls:
videos = []
# prompt_embed are named embed0 to embedN
# check how many embeds are there
embe_dir = os.path.join(args.validation_prompt_dir, "prompt_embed")
mask_dir = os.path.join(args.validation_prompt_dir, "prompt_attention_mask")
embeds = sorted([f for f in os.listdir(embe_dir)])
masks = sorted([f for f in os.listdir(mask_dir)])
num_embeds = len(embeds)
validation_prompt_ids = list(range(num_embeds))
num_sp_groups = int(os.getenv("WORLD_SIZE", "1")) // nccl_info.sp_size
# pad to multiple of groups
if num_embeds % num_sp_groups != 0:
validation_prompt_ids += [0] * (num_sp_groups - num_embeds % num_sp_groups)
num_embeds_per_group = len(validation_prompt_ids) // num_sp_groups
local_prompt_ids = validation_prompt_ids[
nccl_info.group_id * num_embeds_per_group : (nccl_info.group_id + 1)
* num_embeds_per_group
]
for i in local_prompt_ids:
prompt_embed_path = os.path.join(
embe_dir, f"{embeds[i]}"
)
prompt_mask_path = os.path.join(
mask_dir, f"{masks[i]}"
)
prompt_embeds = (
torch.load(prompt_embed_path, map_location="cpu", weights_only=True)
.to(device)
.unsqueeze(0)
)
prompt_attention_mask = (
torch.load(prompt_mask_path, map_location="cpu", weights_only=True)
.to(device)
.unsqueeze(0)
)
negative_prompt_embeds = (
torch.zeros(256, 4096).to(device).unsqueeze(0)
)
negative_prompt_attention_mask = (
torch.zeros(256).bool().to(device).unsqueeze(0)
)
generator = torch.Generator(device="cuda").manual_seed(12345)
video = sample_validation_video(
transformer,
vae,
scheduler,
scheduler_type=scheduler_type,
num_frames=args.num_frames,
# Peiyuan TODO: remove hardcode
height=480,
width=848,
num_inference_steps=validation_sampling_step,
guidance_scale=validation_guidance_scale,
generator=generator,
prompt_embeds=prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_embeds=negative_prompt_embeds,
negative_prompt_attention_mask=negative_prompt_attention_mask,
vae_spatial_scale_factor=vae_spatial_scale_factor,
vae_temporal_scale_factor=vae_temporal_scale_factor,
num_channels_latents=num_channels_latents
)[0]
if nccl_info.rank_within_group == 0:
videos.append(video[0])
# collect videos from all process to process zero
gc.collect()
torch.cuda.empty_cache()
# log if main process
torch.distributed.barrier()
all_videos = [
None for i in range(int(os.getenv("WORLD_SIZE", "1")))
] # remove padded videos
torch.distributed.all_gather_object(all_videos, videos)
if nccl_info.global_rank == 0:
# remove padding
videos = [video for videos in all_videos for video in videos]
videos = videos[:num_embeds]
# linearize all videos
video_filenames = []
for i, video in enumerate(videos):
filename = os.path.join(
args.output_dir,
f"validation_step_{global_step}_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}_video_{i}.mp4",
)
export_to_video(video, filename, fps=fps)
video_filenames.append(filename)
logs = {
f"{'ema_' if ema else ''}validation_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}": [
wandb.Video(filename)
for i, filename in enumerate(video_filenames)
]
}
wandb.log(logs, step=global_step)
-31
View File
@@ -1,31 +0,0 @@
"""fastvideo2 — a post-training-to-serving substrate for video models, MVP.
Four surfaces (see README.md at the repo root):
contracts fastvideo2.card / pipeline / loop — frozen data cards, enforced
stage edges, the driven-loop protocol
reference fastvideo2.<family>.reference — the standalone eager oracle
verifier fastvideo2.verify — tiered gates + evidence ledger
trace engine identity chain — request/stage/loop.step -> NVTX
``import fastvideo2`` is dependency-light: torch / diffusers / transformers
load lazily, only when weights are actually touched.
"""
from fastvideo2.card import ModelCard, derive
from fastvideo2.engine import Instance, Output, Request, run
from fastvideo2.registry import resolve
from fastvideo2.sdk import Model, Result, load
__version__ = "0.1.0"
__all__ = ["ModelCard", "derive", "Model", "Result", "load", "resolve",
"Instance", "Output", "Request", "run", "generate", "__version__"]
def generate(model: str, prompt: str, *, root: str | None = None,
device: str | None = None, **request_kwargs) -> Result:
"""One-call convenience over the SDK: load then generate (loads per call —
hold a :class:`Model` via :func:`load` to amortize residency).
>>> result = fastvideo2.generate("wan2.1-t2v-1.3b", "a cat surfing", seed=7)
>>> result.video.shape # [T, H, W, C] uint8
"""
return load(model, root=root, device=device).generate(prompt, **request_kwargs)
-3
View File
@@ -1,3 +0,0 @@
from fastvideo2.cli import main
raise SystemExit(main())
-213
View File
@@ -1,213 +0,0 @@
"""Model cards — the contract surface.
A card is a frozen, pure-data description of one servable artifact: its
components, the loops its weights assume, and the sampling defaults that are
part of the trained artifact. Components and loops are declared as
``"module:attr"`` reference strings and loop params must be plain JSON values
— no callables — so a card is:
* **serializable** — ``to_json``/``from_dict`` are lossless (T0-gated), so a
card ships as ``card.json`` beside a checkpoint or inside a deploy config;
* **content-addressed** — ``digest()`` hashes the canonical JSON (think git
object id) and is the card's identity in evidence records, T1 baselines,
and environment manifests. It names the declaration only; weights, code,
and environment drift are checked separately (T1, T2, env fingerprint).
Variants are expressed as diffs against a base card via :func:`derive` — never
as builder functions with keyword arguments. Derivation is additive: a variant
that needs to *remove* something picked the wrong base and should be declared
fresh.
Import discipline: this module is stdlib-only. ``validate()`` imports the
declared loop modules to check their ``semantics`` ids, so loop modules must be
importable without torch.
"""
from __future__ import annotations
import hashlib
import importlib
import json
from dataclasses import asdict, dataclass, field, fields, is_dataclass, replace
from typing import Any
class CardError(ValueError):
pass
def resolve_ref(ref: str) -> Any:
"""Resolve a ``"module:attr"`` reference string to the live object."""
mod, _, attr = ref.partition(":")
if not mod or not attr:
raise CardError(f"bad reference {ref!r} (expected 'module:attr')")
return getattr(importlib.import_module(mod), attr)
@dataclass(frozen=True)
class ComponentSpec:
"""One weight-bearing (or processing) component of the artifact."""
component_id: str
kind: str # dit | vae | text_encoder | tokenizer
module: str # loader reference, e.g. "fastvideo2.wan21.model:WanModel"
subfolder: str # subfolder in the checkpoint layout ("" = repo root)
dtype: str = "bf16" # bf16 | fp32 | "" (dtype-less, e.g. tokenizer)
source: str = "" # weights repo override; "" = the card-level `weights`
@dataclass(frozen=True)
class LoopSpec:
"""One iterative computation the card can run.
``loop`` names the implementation class; ``params`` are plain JSON values
passed to its constructor. The class carries a ``semantics`` id that
provenance pins (see ``Provenance.assumes_loop``).
"""
loop_id: str
loop: str # "fastvideo2.wan21.loop:WanDenoiseLoop"
params: dict[str, Any] = field(default_factory=dict)
@dataclass(frozen=True)
class SamplingDefaults:
"""Per-model generation defaults that are part of the trained artifact."""
num_steps: int
guidance_scale: float
height: int
width: int
num_frames: int
fps: int
shift: float
negative_prompt: str = ""
@dataclass(frozen=True)
class Provenance:
"""Where the weights came from and what they assume.
``assumes_loop`` is a *semantics id* (e.g. ``"wan.flow_euler.cfg/v1"``),
not a loop_id: validation resolves every declared loop class and requires
one whose ``semantics`` matches. A distilled student that requires a
different sampler therefore cannot validate against a base card.
``substitution`` classifies this artifact relative to ``parents``:
``exact`` | ``bounded`` | ``quality-changing``.
"""
method: str = "base"
parents: tuple[str, ...] = ()
assumes_loop: str = ""
precision: str = "bf16"
substitution: str = "exact"
tolerances: dict[str, float] = field(default_factory=dict)
@dataclass(frozen=True)
class ModelCard:
model_id: str
family: str
weights: str # canonical source (HF repo id)
components: dict[str, ComponentSpec]
loops: dict[str, LoopSpec]
capabilities: tuple[str, ...]
provenance: Provenance
sampling_defaults: SamplingDefaults
determinism: str = "tolerance" # bitwise | tolerance
# --- identity ---------------------------------------------------------- #
def to_dict(self) -> dict:
return asdict(self)
def to_json(self) -> str:
return json.dumps(self.to_dict(), sort_keys=True, indent=2)
def digest(self) -> str:
"""Content digest over the canonical JSON — the card's identity in the
evidence ledger and every environment manifest."""
canon = json.dumps(self.to_dict(), sort_keys=True, separators=(",", ":"))
return hashlib.sha256(canon.encode()).hexdigest()[:16]
@classmethod
def from_dict(cls, d: dict) -> "ModelCard":
return cls(
model_id=d["model_id"],
family=d["family"],
weights=d["weights"],
components={k: ComponentSpec(**v) for k, v in d["components"].items()},
loops={k: LoopSpec(**v) for k, v in d["loops"].items()},
capabilities=tuple(d["capabilities"]),
provenance=Provenance(**{**d["provenance"], "parents": tuple(d["provenance"]["parents"])}),
sampling_defaults=SamplingDefaults(**d["sampling_defaults"]),
determinism=d.get("determinism", "tolerance"),
)
# --- validation -------------------------------------------------------- #
def validate(self) -> "ModelCard":
errs: list[str] = []
if not self.components:
errs.append("card declares no components")
if not self.loops:
errs.append("card declares no loops")
for cid, spec in self.components.items():
if cid != spec.component_id:
errs.append(f"component key {cid!r} != component_id {spec.component_id!r}")
semantics_seen: list[str] = []
for lid, spec in self.loops.items():
if lid != spec.loop_id:
errs.append(f"loop key {lid!r} != loop_id {spec.loop_id!r}")
try:
cls = resolve_ref(spec.loop)
except Exception as e: # unresolvable ref is a contract violation
errs.append(f"loop {lid!r}: cannot resolve {spec.loop!r} ({e})")
continue
sem = getattr(cls, "semantics", None)
if not sem:
errs.append(f"loop {lid!r}: class {spec.loop!r} declares no `semantics` id")
else:
semantics_seen.append(sem)
try:
json.dumps(spec.params)
except TypeError:
errs.append(f"loop {lid!r}: params are not plain JSON values")
# the teeth: weights may only be served under a loop whose semantics
# they were trained for.
if self.provenance.assumes_loop and self.provenance.assumes_loop not in semantics_seen:
errs.append(
f"provenance.assumes_loop={self.provenance.assumes_loop!r} matches no declared "
f"loop semantics (have {semantics_seen}) — these weights cannot run on this card")
if self.determinism not in ("bitwise", "tolerance"):
errs.append(f"unknown determinism class {self.determinism!r}")
if errs:
raise CardError(f"ModelCard {self.model_id!r} failed validation:\n - " + "\n - ".join(errs))
return self
def _merge_field(old: Any, patch: Any) -> Any:
"""One-level structural merge used by :func:`derive`.
dict field + dict patch -> merge by key (spec values replace; dict
values patch the existing spec/dict)
dataclass field + dict patch -> replace() with recursively merged fields
anything else -> the patch value wins
"""
if is_dataclass(old) and isinstance(patch, dict):
merged = {k: _merge_field(getattr(old, k), v) for k, v in patch.items()}
return replace(old, **merged)
if isinstance(old, dict) and isinstance(patch, dict):
out = dict(old)
for k, v in patch.items():
out[k] = _merge_field(old[k], v) if k in old else v
return out
return patch
def derive(base: ModelCard, **delta: Any) -> ModelCard:
"""A variant as an explicit diff against a base card.
Additive only: keys merge, nothing is deleted. A variant that must remove a
component or loop is a different architecture — declare it fresh. The
derived card re-validates, so an invalid diff fails at declaration.
"""
valid = {f.name for f in fields(ModelCard)}
unknown = set(delta) - valid
if unknown:
raise CardError(f"derive: unknown card fields {sorted(unknown)}")
merged = {k: _merge_field(getattr(base, k), v) for k, v in delta.items()}
return replace(base, **merged).validate()
-92
View File
@@ -1,92 +0,0 @@
"""CLI: describe / generate / verify — the three agent-facing verbs.
python -m fastvideo2 describe wan2.1-t2v-1.3b
python -m fastvideo2 generate wan2.1-t2v-1.3b --prompt "a cat surfing" --out cat.mp4
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 2 [--bless]
``describe`` prints the card as JSON plus its digest — machine-readable
capability discovery. ``verify`` appends typed results to the evidence ledger
and exits non-zero on any failed gate.
"""
from __future__ import annotations
import argparse
import json
import sys
def _describe(args) -> int:
from fastvideo2.registry import resolve
card, _ = resolve(args.model)
print(card.to_json())
print(f'// digest: {card.digest()}', file=sys.stderr)
return 0
def _generate(args) -> int:
import fastvideo2
kwargs = {k: getattr(args, k) for k in
("seed", "num_steps", "guidance_scale", "height", "width", "num_frames", "shift")
if getattr(args, k) is not None}
model = fastvideo2.load(args.model, root=args.root, device=args.device)
result = model.generate(args.prompt, **kwargs)
result.save(args.out, fps=args.fps)
steps = [t for t in result.trace if "/denoise." in t["label"]]
print(f"video {result.video.shape} -> {args.out}")
print(f"total {result.seconds:.1f}s; denoise steps {len(steps)}, "
f"mean {sum(t['seconds'] for t in steps) / max(len(steps), 1):.2f}s/step")
return 0
def _verify(args) -> int:
from fastvideo2.verify import LEDGER, verify
results = verify(args.model, tier=args.tier, root=args.root, device=args.device,
bless=args.bless, anchor=args.anchor)
for r in results:
mark = {"pass": "PASS ", "blessed": "BLESS", "fail": "FAIL "}[r.status]
print(f" {mark} {r.gate:14s} {r.detail or json.dumps(r.metrics)[:120]}")
print(f"ledger: {LEDGER}")
return 0 if all(r.ok for r in results) else 1
def main(argv: list[str] | None = None) -> int:
p = argparse.ArgumentParser(prog="fastvideo2", description=__doc__)
sub = p.add_subparsers(dest="cmd", required=True)
d = sub.add_parser("describe", help="print a card as JSON + digest")
d.add_argument("model")
d.set_defaults(fn=_describe)
g = sub.add_parser("generate", help="run one request, save an mp4")
g.add_argument("model")
g.add_argument("--prompt", required=True)
g.add_argument("--out", default="out.mp4")
g.add_argument("--root", default=None, help="local checkpoint dir (else HF cache)")
g.add_argument("--device", default=None)
g.add_argument("--seed", type=int, default=0)
g.add_argument("--num-steps", dest="num_steps", type=int, default=None)
g.add_argument("--guidance-scale", dest="guidance_scale", type=float, default=None)
g.add_argument("--height", type=int, default=None)
g.add_argument("--width", type=int, default=None)
g.add_argument("--num-frames", dest="num_frames", type=int, default=None)
g.add_argument("--shift", type=float, default=None)
g.add_argument("--fps", type=int, default=16)
g.set_defaults(fn=_generate)
v = sub.add_parser("verify", help="run tiered gates; append to the evidence ledger")
v.add_argument("model")
v.add_argument("--tier", type=int, default=3, choices=(0, 1, 2, 3))
v.add_argument("--root", default=None)
v.add_argument("--device", default=None)
v.add_argument("--bless", action="store_true",
help="write the T1 fingerprint baseline for this environment")
v.add_argument("--anchor", action="store_true",
help="also certify components against the official Wan2.1 goldens")
v.set_defaults(fn=_verify)
args = p.parse_args(argv)
return args.fn(args)
if __name__ == "__main__":
raise SystemExit(main())
-4
View File
@@ -1,4 +0,0 @@
"""Dreamverse runtime on fastvideo2 — session WS + fMP4 segments. See server.py."""
from fastvideo2.dreamverse.server import build_app, main
__all__ = ["build_app", "main"]
-3
View File
@@ -1,3 +0,0 @@
from fastvideo2.dreamverse.server import main
main()
@@ -1,100 +0,0 @@
"""Dreamverse runtime anchor: boot the server with DUMMY prompt keys, drive
one full session over the protocol, and assert the streaming contract:
segment_start -> live step_complete x3 -> media_init -> binary fMP4 chunks
(first chunk carries an ISO-BMFF `ftyp` box) -> media_segment_complete ->
segment_complete with a latents sha.
Usage (cluster): python -m fastvideo2.dreamverse.gates.dreamverse_anchor
"""
from __future__ import annotations
import asyncio
import json
import os
import subprocess
import sys
import time
import urllib.request
MODEL = "fastwan-qad-fp8-1.3b"
PORT = 8019
PROMPT = "A raccoon in a field of sunflowers, warm light, mid-shot."
def main() -> int:
env = dict(os.environ, CEREBRAS_API_KEY="dummy", GROQ_API_KEY="dummy")
server = subprocess.Popen(
[sys.executable, "-m", "fastvideo2.dreamverse", "--model", MODEL,
"--port", str(PORT)], env=env,
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
try:
for _ in range(360):
try:
with urllib.request.urlopen(
f"http://127.0.0.1:{PORT}/health", timeout=5) as r:
if json.loads(r.read())["model"] == MODEL:
break
except Exception:
time.sleep(2)
else:
raise RuntimeError("server never became healthy")
import websockets
async def session() -> dict:
counts = {"steps": 0, "chunks": 0, "ftyp": False}
async with websockets.connect(
f"ws://127.0.0.1:{PORT}/ws", max_size=None) as ws:
await ws.send(json.dumps({"type": "session_init_v2",
"enhancement": True}))
for expected in ("queue_status", "gpu_assigned", "stream_start"):
got = json.loads(await ws.recv())["type"]
assert got == expected, (got, expected)
await ws.send(json.dumps({"type": "segment_prompt_source",
"prompt": PROMPT, "seed": 7}))
while True:
raw = await asyncio.wait_for(ws.recv(), timeout=900)
if isinstance(raw, bytes):
if counts["chunks"] == 0:
counts["ftyp"] = b"ftyp" in raw[:64]
counts["chunks"] += 1
continue
msg = json.loads(raw)
t = msg["type"]
if t == "step_complete":
counts["steps"] += 1
elif t == "segment_complete":
counts["latents_sha"] = msg["latents_sha"]
counts["frames"] = msg["frames"]
break
elif t == "error":
raise RuntimeError(msg)
await ws.send(json.dumps({"type": "leave"}))
assert json.loads(await ws.recv())["type"] == "stream_complete"
return counts
c = asyncio.run(session())
finally:
server.terminate()
server.wait(timeout=30)
ok = (c["steps"] == 3 and c["chunks"] >= 1 and c["ftyp"]
and c.get("frames", 0) == 81)
print(f"steps={c['steps']} chunks={c['chunks']} ftyp={c['ftyp']} "
f"frames={c.get('frames')} sha={c.get('latents_sha')} "
f"{'OK' if ok else 'FAIL'}")
from fastvideo2.verify import GateResult, append_ledger, env_fingerprint
append_ledger([GateResult(gate="anchor.dreamverse-runtime",
status="pass" if ok else "fail", model_id=MODEL,
card_digest="-",
metrics={"steps": float(c["steps"]),
"chunks": float(c["chunks"]),
"ftyp": 1.0 if c["ftyp"] else 0.0},
tolerances={}, env=env_fingerprint(),
detail=f"latents {c.get('latents_sha')}")])
return 0 if ok else 1
if __name__ == "__main__":
sys.exit(main())
-269
View File
@@ -1,269 +0,0 @@
"""Dreamverse runtime on fastvideo2 — the realtime video session server.
Ported from ``apps/dreamverse`` (fastvideo-main): the WebSocket session
protocol (``session_init_v2`` → per-segment prompts → fMP4 fragments over
the socket), the ffmpeg fragmented-MP4 encoder (verbatim flags from
``entrypoints/streaming/stream.py``: libx264, zerolatency,
``empty_moov+default_base_moof+frag_keyframe+faststart``), and an optional
Cerebras/Groq prompt enhancer (boots with dummy keys — enhancement simply
stays off, the same bring-up shortcut the GB200 deploys used).
Deliberately re-based for v2.1 (the original is LTX2-specific — audio,
refine stage, continuation-state, LoRA stack): segments generate through the
fastvideo2 SDK on the FastWan 3-step DMD student (seconds per segment on
GB200), and per-step progress is LIVE via the engine's ``on_step`` hook
(``step_complete`` per denoise step — the original emits one terminal event
per segment). Message names follow the upstream protocol so their web client
schema maps directly; LTX2-only fields are ignored.
Run: python -m fastvideo2.dreamverse --model fastwan-qad-fp8-1.3b --port 8009
ffmpeg: FASTVIDEO_FFMPEG_BIN, or PATH, or the imageio-ffmpeg bundled binary.
"""
from __future__ import annotations
import asyncio
import hashlib
import json
import os
import shutil
import subprocess
import threading
import urllib.request
import uuid
from typing import Any
def find_ffmpeg() -> str:
p = os.environ.get("FASTVIDEO_FFMPEG_BIN") or shutil.which("ffmpeg")
if p:
return p
try: # the GB200 bring-up shortcut: pip-installed bundled binary
import imageio_ffmpeg
return imageio_ffmpeg.get_ffmpeg_exe()
except ImportError as e:
raise RuntimeError("no ffmpeg (set FASTVIDEO_FFMPEG_BIN, install "
"ffmpeg, or `pip install imageio-ffmpeg`)") from e
def fmp4_encode(frames: Any, *, fps: int, ffmpeg: str) -> list[bytes]:
"""One segment -> fragmented-MP4 byte chunks (upstream's exact flags)."""
t, h, w, _ = frames.shape
args = [ffmpeg, "-hide_banner", "-loglevel", "error",
"-f", "rawvideo", "-pix_fmt", "rgb24", "-s", f"{w}x{h}",
"-r", str(fps), "-i", "-",
"-c:v", "libx264", "-preset", "ultrafast", "-tune", "zerolatency",
"-pix_fmt", "yuv420p",
"-movflags", "empty_moov+default_base_moof+frag_keyframe+faststart",
"-f", "mp4", "-"]
proc = subprocess.Popen(args, stdin=subprocess.PIPE, stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL, bufsize=0)
out: list[bytes] = []
def _read() -> None:
while True:
chunk = proc.stdout.read(65536)
if not chunk:
break
out.append(chunk)
reader = threading.Thread(target=_read, daemon=True)
reader.start()
for i in range(t):
proc.stdin.write(frames[i].tobytes())
proc.stdin.close()
proc.wait(timeout=120)
reader.join(timeout=30)
return out
class PromptEnhancer:
"""Cerebras-or-Groq chat call (upstream's provider pair, gpt-oss-120b).
Dummy/missing keys or any failure -> pass the prompt through unchanged."""
def __init__(self) -> None:
self.cerebras = os.environ.get("CEREBRAS_API_KEY", "")
self.groq = os.environ.get("GROQ_API_KEY", "")
self.enabled = any(k and k != "dummy" for k in (self.cerebras, self.groq))
def enhance(self, prompt: str, history: list[str]) -> str:
if not self.enabled:
return prompt
targets = []
if self.cerebras and self.cerebras != "dummy":
targets.append(("https://api.cerebras.ai/v1/chat/completions",
self.cerebras, "gpt-oss-120b"))
if self.groq and self.groq != "dummy":
targets.append(("https://api.groq.com/openai/v1/chat/completions",
self.groq, "openai/gpt-oss-120b"))
system = ("Rewrite the user's next-video-segment prompt into one vivid, "
"concrete shot description. Prior segments: "
+ " | ".join(history[-3:]))
for url, key, model_name in targets:
try:
req = urllib.request.Request(
url, method="POST",
headers={"Authorization": f"Bearer {key}",
"Content-Type": "application/json"},
data=json.dumps({"model": model_name, "temperature": 1.0,
"messages": [{"role": "system", "content": system},
{"role": "user", "content": prompt}]
}).encode())
with urllib.request.urlopen(req, timeout=20) as r:
return json.loads(r.read())["choices"][0]["message"]["content"].strip()
except Exception:
continue
return prompt
def build_app(model: Any) -> Any:
from fastapi import FastAPI
from starlette.routing import WebSocketRoute
from starlette.websockets import WebSocketDisconnect
from fastvideo2.engine import Request
from fastvideo2.engine import run as engine_run
from fastvideo2.sdk import Result
app = FastAPI(title="dreamverse-fv2", version="0.1")
ffmpeg = find_ffmpeg()
enhancer = PromptEnhancer()
gen_lock = threading.Lock()
@app.get("/health")
def health() -> dict:
return {"status": "ok", "model": model.model_id,
"enhancer": enhancer.enabled, "ffmpeg": ffmpeg}
@app.get("/readyz")
def readyz() -> dict:
return {"ready": True}
async def ws_session(ws) -> None:
await ws.accept()
try:
init = await ws.receive_json()
except WebSocketDisconnect:
return
if init.get("type") != "session_init_v2":
await ws.send_json({"type": "error", "code": "bad_init",
"error": "expected session_init_v2"})
await ws.close()
return
session_id = uuid.uuid4().hex[:12]
enhancement_on = bool(init.get("enhancement", False)) and enhancer.enabled
history: list[str] = []
segment_idx = 0
await ws.send_json({"type": "queue_status", "position": 0})
await ws.send_json({"type": "gpu_assigned", "session_id": session_id})
await ws.send_json({"type": "stream_start", "session_id": session_id,
"model": model.model_id})
loop = asyncio.get_running_loop()
while True:
try:
msg = await ws.receive_json()
except WebSocketDisconnect:
return
mtype = msg.get("type")
if mtype == "leave":
await ws.send_json({"type": "stream_complete",
"segments": segment_idx})
await ws.close()
return
if mtype == "enhancement_updated":
enhancement_on = bool(msg.get("enabled")) and enhancer.enabled
continue
if mtype != "segment_prompt_source":
await ws.send_json({"type": "error", "code": "bad_message",
"error": f"unsupported type {mtype!r}"})
continue
prompt = str(msg.get("prompt", ""))
if not prompt:
await ws.send_json({"type": "error", "code": "bad_prompt",
"error": "prompt required"})
continue
if enhancement_on:
prompt = await loop.run_in_executor(
None, enhancer.enhance, prompt, history)
history.append(prompt)
await ws.send_json({"type": "segment_start", "segment": segment_idx,
"prompt": prompt})
q: asyncio.Queue = asyncio.Queue()
def on_step(label: str, seconds: float, meta: dict) -> None:
loop.call_soon_threadsafe(
q.put_nowait, {"type": "step_complete", "label": label,
"seconds": round(seconds, 4)})
def generate(p: str = prompt, seed: Any = msg.get("seed", 0),
steps: Any = msg.get("num_steps")) -> None:
try:
req = Request(prompt=p, request_id=f"{session_id}-{segment_idx}",
seed=int(seed), num_steps=steps)
with gen_lock:
out = engine_run(model.instance, model.pipeline, req,
on_step=on_step)
result = Result(outputs=out.outputs, trace=out.trace,
request=req.resolve(model.card),
model_id=model.model_id,
card_digest=model.card.digest(),
fps=model.card.sampling_defaults.fps)
chunks = fmp4_encode(result.video, fps=result.fps,
ffmpeg=ffmpeg)
import torch
sha = hashlib.sha256(result.latents.detach().to(
torch.float32).cpu().numpy().tobytes()).hexdigest()[:16]
loop.call_soon_threadsafe(
q.put_nowait, {"__chunks": chunks, "latents_sha": sha,
"frames": int(result.video.shape[0])})
except Exception as e:
loop.call_soon_threadsafe(
q.put_nowait, {"type": "error", "code": "generation",
"error": f"{type(e).__name__}: {e}"})
threading.Thread(target=generate, daemon=True).start()
while True:
ev = await q.get()
if "__chunks" in ev:
await ws.send_json({"type": "media_init",
"segment": segment_idx,
"mime": 'video/mp4; codecs="avc1"'})
for chunk in ev["__chunks"]:
await ws.send_bytes(chunk)
await ws.send_json({"type": "media_segment_complete",
"segment": segment_idx})
await ws.send_json({"type": "segment_complete",
"segment": segment_idx,
"frames": ev["frames"],
"latents_sha": ev["latents_sha"]})
segment_idx += 1
break
await ws.send_json(ev)
if ev.get("type") == "error":
break
app.router.routes.append(WebSocketRoute("/ws", ws_session))
return app
def main(argv: list[str] | None = None) -> None:
import argparse
import uvicorn
import fastvideo2 as fv2
p = argparse.ArgumentParser("fastvideo2.dreamverse")
p.add_argument("--model", default="fastwan-qad-fp8-1.3b")
p.add_argument("--host", default="127.0.0.1")
p.add_argument("--port", type=int, default=8009)
p.add_argument("--device", default=None)
args = p.parse_args(argv)
model = fv2.load(args.model, device=args.device)
uvicorn.run(build_app(model), host=args.host, port=args.port)
if __name__ == "__main__":
main()
-178
View File
@@ -1,178 +0,0 @@
"""The engine — drives a pipeline over one resident instance, one request at a
time, with the identity chain attached.
Identity chain: every unit of work is named ``request/stage`` for one-shot
stages and ``request/stage/loop.step`` for loop steps. The same name goes to
(a) the returned trace (typed timings, machine-readable) and (b) NVTX ranges
when CUDA is present — so Nsight correlates kernels to model-level identity
with no extra instrumentation.
Deliberately absent (this is the one-shot MVP): queueing, admission, batching,
sessions, cancellation. Sessions with forkable state are the next consumer of
the loop contract, not a reason to grow this file now.
"""
from __future__ import annotations
import contextlib
from dataclasses import dataclass, field, replace
from typing import Any
from fastvideo2.card import ModelCard
from fastvideo2.loading import load_component, resolve_weights
from fastvideo2.loop import LoopRunner, build_loop
from fastvideo2.pipeline import ComponentStage, LoopStage, Pipeline, run_component_stage
@dataclass(frozen=True)
class Request:
"""One generation request. ``None`` fields resolve from the card's
sampling defaults via :meth:`resolve`."""
prompt: str
request_id: str = "req0"
negative_prompt: str | None = None
seed: int = 0
num_steps: int | None = None
guidance_scale: float | None = None
height: int | None = None
width: int | None = None
num_frames: int | None = None
shift: float | None = None
capture_trajectory: bool = False
def resolve(self, card: ModelCard) -> "Request":
d = card.sampling_defaults
fill = {
"negative_prompt": d.negative_prompt,
"num_steps": d.num_steps,
"guidance_scale": d.guidance_scale,
"height": d.height,
"width": d.width,
"num_frames": d.num_frames,
"shift": d.shift,
}
patch = {k: v for k, v in fill.items() if getattr(self, k) is None}
return replace(self, **patch)
@dataclass
class Output:
request_id: str
outputs: dict[str, Any]
trace: list[dict] = field(default_factory=list) # [{label, seconds, ...meta}]
@property
def seconds(self) -> float:
return sum(t["seconds"] for t in self.trace)
class Instance:
"""A resident, loaded card: components materialize lazily and are shared
by reference; loops are built from the card's declared specs."""
def __init__(self, card: ModelCard, root: str | None = None, device: str = "cpu"):
self.card = card
self.device = device
self.root = resolve_weights(card, root)
self._components: dict[str, Any] = {}
self._source_roots: dict[str, str] = {}
self._loops: dict[str, Any] = {}
def component(self, component_id: str) -> Any:
if component_id not in self._components:
spec = self.card.components.get(component_id)
if spec is None:
raise KeyError(f"component {component_id!r} not declared on card {self.card.model_id!r}")
root = self._source_root(spec.source) if spec.source else self.root
self._components[component_id] = load_component(spec, root, self.device)
return self._components[component_id]
def _source_root(self, source: str) -> str:
"""Resolve a per-component weights source (e.g. the official-layout
transformer repo) through the same snapshot cache as card weights."""
if source not in self._source_roots:
from huggingface_hub import snapshot_download
self._source_roots[source] = snapshot_download(source)
return self._source_roots[source]
def loop(self, loop_id: str) -> Any:
if loop_id not in self._loops:
spec = self.card.loops.get(loop_id)
if spec is None:
raise KeyError(f"loop {loop_id!r} not declared on card {self.card.model_id!r}")
self._loops[loop_id] = build_loop(spec)
return self._loops[loop_id]
def load(card: ModelCard, root: str | None = None, device: str | None = None) -> Instance:
"""The public entrypoint: card + weights root -> resident instance."""
card.validate()
if device is None:
device = _detect_device()
return Instance(card, root=root, device=device)
def _detect_device() -> str:
try:
import torch
if torch.cuda.is_available():
return "cuda"
except ImportError:
pass
return "cpu"
@contextlib.contextmanager
def _nvtx(name: str):
"""NVTX range when CUDA is live; free otherwise."""
pushed = False
try:
import torch
if torch.cuda.is_available():
torch.cuda.nvtx.range_push(name)
pushed = True
except ImportError:
pass
try:
yield
finally:
if pushed:
import torch
torch.cuda.nvtx.range_pop()
def run(instance: Instance, pipeline: Pipeline, request: Request,
on_step: Any = None) -> Output:
"""Run one request through a pipeline to completion. ``on_step(label,
seconds, meta)``, when given, fires live after every loop step (serving
streams progress through it)."""
import time
req = request.resolve(instance.card)
trace: list[dict] = []
slots: dict[str, Any] = {}
for name in pipeline.inputs: # request-provided slots, by attribute name
slots[name] = getattr(req, name)
for stage in pipeline.stages:
chain = f"{req.request_id}/{stage.stage_id}"
if isinstance(stage, ComponentStage):
with _nvtx(chain):
t0 = time.perf_counter()
run_component_stage(stage, instance, slots, req)
trace.append({"label": chain, "seconds": time.perf_counter() - t0})
elif isinstance(stage, LoopStage):
loop = instance.loop(stage.loop_id)
def observe(label: str, seconds: float, meta: dict, _chain: str = chain) -> None:
trace.append({"label": f"{_chain}/{label}", "seconds": seconds, **meta})
if on_step is not None:
on_step(f"{_chain}/{label}", seconds, meta)
inputs = {k: slots[k] for k in stage.reads}
with _nvtx(chain):
runner = LoopRunner(loop, req, instance, inputs, observe=observe)
slots[stage.writes[0]] = runner.run()
else:
raise TypeError(f"unknown stage kind {type(stage).__name__}")
outputs = {name: slots[slot] for name, slot in pipeline.outputs.items()}
return Output(request_id=req.request_id, outputs=outputs, trace=trace)
-18
View File
@@ -1,18 +0,0 @@
# Evidence ledger
Typed verification records, committed with the code they vouch for.
- `ledger.jsonl` — append-only `GateResult` records: gate, status, card digest,
metrics, tolerances, environment fingerprint, timestamp. Written only by
`python -m fastvideo2 verify`; never edited by hand.
- `<model_id>.fingerprints.json` — the blessed T1 component baseline for one
card digest in one environment.
- `sample_*.mp4` — eyeballable artifacts from full-scale runs (e.g.
`sample_wan21_seed7.mp4`, 50 steps / 81 frames on GB200, byte-identical
between the production pipeline and `reference.py` at the same seed). Re-bless deliberately (`verify --bless`)
when the card or environment legitimately changes; a digest mismatch is a
failure, not a skip.
Ownership rule: baselines and gate tolerances are human-owned. Agents run the
gates and append evidence; they do not re-bless baselines to make a failure
disappear.
@@ -1,82 +0,0 @@
# FastWan variants — bitwise alignment vs fastvideo-main
Authority for FastWan artifacts is **fastvideo main** (they were distilled in
that stack); the alignment target was bit-exactness against main's own serving
path, pinned to main's exposed knobs so the goldens measure the artifact, not
the accelerator stack: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN` (QAD) /
`VIDEO_SPARSE_ATTN` (VSA), FP8 per-tensor dynamic quant (QAD), no torch.compile,
no FSDP, single GB200.
Goldens captured at main commit `c459a1897899ffcec3be7765534d81000b9bb9c1`;
the vendored forward (`wan21/model_fv.py`) was read at `e3f47dc2de2d…` — all
10 numerics-relevant source files verified byte-identical between the two.
## Final anchor results (all rows target 0.0 — bitwise)
| row | fastwan-qad-fp8-1.3b | fastwan-t2v-1.3b (VSA) |
|---|---|---|
| dit bf16 probes t∈{1000,757,522} | 0.0 / 0.0 / 0.0 | — |
| dit fp8 probes t∈{1000,757,522} | 0.0 / 0.0 / 0.0 | — |
| dit vsa probes t∈{1000,757,522} | — | 0.0 / 0.0 / 0.0 |
| text_encoder (e2e + probe prompts) | 0.0 / 0.0 | 0.0 / 0.0 |
| e2e step-1 / step-2 latent chain | 0.0 / 0.0 | 0.0 / 0.0 |
| e2e final latents (81f, 480×832, 3 steps) | 0.0 | 0.0 |
Ledger: `anchor.fastwan-qad-main` and `anchor.fastwan-vsa-main`, both `pass`
(card digests `1c8e6f7d1380552d`, `0bd5c7771e10ce44`).
## Root causes found by the gates (in discovery order)
1. **fp8 quantization is device-sensitive.** main converts weights to fp8 on
the GPU (post-materialization); quantizing the *identical* bf16 weights on
CPU produces different fp8 codes often enough to move a full forward by
~4e-2 rel (fp8's coarse grid amplifies conversion-tie differences).
Fix: `FP8Linear` defers quantization to first forward on the serving
device (`layers/fp8.py`).
2. **0-dim sigma tensors demote the renoise mixing to bf16.** In torch type
promotion 0-dim tensors act as scalars, so `(1-σ)*x0 + σ*ε` with a 0-dim
fp32 σ ran in bf16 (two roundings); main's `[B,1,1,1]` fp32 σ promotes the
arithmetic to fp32 with one final bf16 cast. 3.2e-3 per step, compounding
to 8.2e-2 over 3 steps. Fix: non-0-dim σ in `WanDMDLoop`.
3. **main's DMD sigma table is NOT the one the code appears to prepare.**
`DmdDenoisingStage.__init__` hardcodes a fresh internal
`FlowMatchEulerDiscreteScheduler(shift=8.0)`; the pipeline scheduler that
`TimestepPreparationStage.set_timesteps(n)` configured is never consulted,
and the config `flow_shift` is ignored. Lookups run against the 1000-entry
warped **init** table: σ(1000)=1.0, σ(757)=0.7567567, σ(522)=0.5217391
(confirmed in the capture manifests). `dmd_inference_table` reproduces
this exactly; a canary T0 test guards it.
Also confirmed en route: the CPU-generator RNG stream (initial fp32 draw +
bf16 renoise draws), the fp64 x0 math, and the flash/dense forward at full
81-frame geometry are each independently bitwise (triage decomposition in
session evidence).
## Caveats
- Text parity holds for ASCII prompts; main's ftfy cleaning diverges from
official's on CJK width-folding (see wan21 report — main measured 4.19e-1
vs official on the Chinese negative prompt). FastWan cards reuse the wan21
text stage; DMD uses no negative prompt. Add main's clean fn + a CJK golden
before serving non-ASCII prompts against these cards.
- VAE decode is not bitwise-gated (shared component; wan21 anchors cover it);
golden videos are in the goldens dirs for SSIM-level comparison.
- Committed goldens are trimmed (per-step model outputs and the reproducible
step-0 input dropped); the full set regenerates via
`gates/capture_fastvideo_main.py {qad,vsa}` — one command, pinned config.
## SFWan (self-forcing causal) — added 2026-07-23
`sfwan-t2v-1.3b` (wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers) anchored bitwise on
the FIRST complete run: all **35/35 chunk-rollout forwards** (7 blocks x
4 warped DMD steps + context pass) hash-match main's CausalDMDDenosingStage
exactly, text 0.0, e2e final latents 0.0 (`anchor.sfwan-main` pass).
Causal-specific semantics vendored (each different from BOTH other Wan
forwards): per-frame temb `[B, T_temb, 6, dim]`; ALL-bf16 modulation (no
fp32 promotion anywhere); plain bf16 LayerNorms; fp64 RoPE multipliers at
absolute positions (start_frame offsets); block-causal KV cache
(21-frame global window, `.detach()` on writes — training rollout reuses
this same module); cached text cross-attention; warp table =
SelfForcingFlowMatchScheduler(shift 5, extra_one_step) rows
`[1000, 937.5, 833.33, 625]` self-indexing their own sigmas.
@@ -1,51 +0,0 @@
{
"repo": "FastVideo/FastWan-QAD-FP8-1.3B",
"snapshot": "3de0eec0e2562923d38a87344127a86a35a3c11d",
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
"fastvideo_src": "/mnt/FastVideo",
"torch": "2.12.0+cu130",
"flash_attn": "2.8.3",
"python": "3.12.13",
"gpu": "NVIDIA GB200",
"attention_backend": "FLASH_ATTN",
"quant": "FP8 per-tensor (dynamic act, post-load weight quant from bf16)",
"vsa_sparsity": null,
"seed": 1234,
"e2e_prompt": "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
"probe_prompt": "A cat and a dog baking a cake together in a kitchen.",
"probe_timesteps": [
1000,
757,
522
],
"probe_latent_bcfhw": [
1,
16,
5,
60,
104
],
"e2e_latent_btchw": [
1,
21,
16,
60,
104
],
"dmd_denoising_steps": [
1000,
757,
522
],
"scheduler": {
"class": "FlowMatchEulerDiscreteScheduler",
"shift": 8.0,
"table_len": 1000,
"sigma_lookup": {
"1000": 1.0,
"757": 0.7567567229270935,
"522": 0.52173912525177
}
},
"notes": "no compile, no fsdp, single GPU; DmdDenoisingStage's INTERNAL scheduler (hardcoded shift 8.0) is the sigma authority; committed goldens trimmed: e2e_step0 dropped (input is the seeded draw, reproducible) and per-step outputs dropped (triage-only) \u2014 full set regenerable via capture_fastvideo_main.py"
}

Some files were not shown because too many files have changed in this diff Show More