Compare commits
50
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8b30a222a3 | ||
|
|
624a638cb3 | ||
|
|
1241ee0600 | ||
|
|
cbdd5e2c70 | ||
|
|
8d1c8fa01a | ||
|
|
78f9ce3c37 | ||
|
|
54def005c4 | ||
|
|
8309acae69 | ||
|
|
7926926ad1 | ||
|
|
eb9d439bab | ||
|
|
d6affff3a1 | ||
|
|
1bd77cb466 | ||
|
|
35900c32ba | ||
|
|
2945890df4 | ||
|
|
4bd8457072 | ||
|
|
f13d017fc1 | ||
|
|
5fe5b3b12f | ||
|
|
ec61d3d0bf | ||
|
|
c665dd83c5 | ||
|
|
b4ffc8e22e | ||
|
|
455cd83bed | ||
|
|
4747d51265 | ||
|
|
7ae57e6fbd | ||
|
|
8931f1b5d7 | ||
|
|
42eb451b50 | ||
|
|
f52e88e397 | ||
|
|
5cdbf06a23 | ||
|
|
9e799d5bf4 | ||
|
|
468c4d505e | ||
|
|
99e8c404b6 | ||
|
|
4e6a6a3d17 | ||
|
|
5784e0fcb5 | ||
|
|
576a006d06 | ||
|
|
423f8d7eff | ||
|
|
a91482ed14 | ||
|
|
f0c838a78d | ||
|
|
13c76b710b | ||
|
|
4bd3eff562 | ||
|
|
041e21821a | ||
|
|
d254b7455d | ||
|
|
b0f1b8ad96 | ||
|
|
5d778c94a1 | ||
|
|
65f3b946b9 | ||
|
|
191fcbf46c | ||
|
|
755a4e4470 | ||
|
|
9709b7513b | ||
|
|
32cd603515 | ||
|
|
d4bdd3621a | ||
|
|
e2f8322842 | ||
|
|
6966f9e0bc |
@@ -1,6 +1,8 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
# Buildkite only launches Modal; remote jobs initialize their own submodules.
|
||||
BUILDKITE_GIT_SUBMODULES: false
|
||||
|
||||
notify:
|
||||
- github_commit_status:
|
||||
|
||||
@@ -11,9 +11,7 @@ on:
|
||||
- 'requirements-mkdocs.txt'
|
||||
- 'scripts/check_docs_links.py'
|
||||
- '.github/workflows/infra-docs.yml'
|
||||
# Run the trusted base-branch workflow so fork PRs can be skipped without
|
||||
# waiting for maintainer approval.
|
||||
pull_request_target:
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
@@ -26,19 +24,21 @@ on:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
build:
|
||||
# MkDocs executes repository code; only trusted same-repository PRs run it.
|
||||
if: github.event_name == 'push' || github.event.pull_request.head.repo.full_name == github.repository
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Setup Python
|
||||
uses: actions/setup-python@v5
|
||||
@@ -52,7 +52,6 @@ jobs:
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
if: github.event_name == 'push'
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
- name: Build documentation
|
||||
@@ -62,22 +61,17 @@ jobs:
|
||||
run: python scripts/check_docs_links.py
|
||||
|
||||
- name: Upload artifact
|
||||
if: github.event_name == 'push'
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./site
|
||||
|
||||
deploy:
|
||||
permissions:
|
||||
pages: write
|
||||
id-token: write
|
||||
concurrency: pages
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
if: github.event_name == 'push'
|
||||
if: github.ref == 'refs/heads/main'
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
|
||||
@@ -6,6 +6,7 @@ results/
|
||||
wandb/
|
||||
*.ipynb
|
||||
*.jpg
|
||||
!examples/dataset/lingbotworld2/image.jpg
|
||||
*.safetensors
|
||||
*.mp4
|
||||
*.png
|
||||
@@ -34,9 +35,15 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# Cosmos3 local parity assets (symlinked from main worktree)
|
||||
/official_weights/
|
||||
/converted_weights/
|
||||
/cosmos-framework
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), check out the [Blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py).
|
||||
|
||||
@@ -8,13 +8,15 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
|
||||
* **[Sliding Tile Attention (STA)](sta/index.md)**: STA kernel support is kept in
|
||||
`fastvideo-kernel`; full FastVideo STA pipeline workflow is archived in
|
||||
`sta_do_not_delete`.
|
||||
* **[Attn-QAT Training](../training/attn_qat.md)**: Runtime-JIT Triton forward
|
||||
and backward kernels for role-local quantization-aware training.
|
||||
* **Backend development guide**: See the developer guide at
|
||||
[Attention Backend Development](../contributing/attention_backend.md).
|
||||
|
||||
## General Build Instructions
|
||||
|
||||
These instructions apply to building the `fastvideo-kernel` package from
|
||||
source, which includes both STA and VSA kernels.
|
||||
source, which includes STA, VSA, and Attn-QAT kernels.
|
||||
|
||||
### Prerequisites
|
||||
|
||||
|
||||
@@ -255,10 +255,9 @@ If you add a new CI test category:
|
||||
|
||||
### Documentation
|
||||
|
||||
`.github/workflows/infra-docs.yml` builds documentation for same-repository PRs
|
||||
that touch `docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow
|
||||
itself. Fork PRs skip this executable build instead of waiting for maintainer
|
||||
approval. On pushes to `main`, it also deploys the built site to GitHub Pages.
|
||||
`.github/workflows/infra-docs.yml` builds documentation for PRs that touch
|
||||
`docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow itself. On
|
||||
pushes to `main`, it also deploys the built site to GitHub Pages.
|
||||
|
||||
The docs job:
|
||||
|
||||
|
||||
@@ -333,15 +333,25 @@ surfaces:
|
||||
use_distill:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
scheduler_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
text_encoder_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
tokenizer_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
transformer_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
vae_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
expand_timesteps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
@@ -421,6 +431,8 @@ surfaces:
|
||||
frame_receptive_field: "MagiHuman internal data-proxy receptive-field setting."
|
||||
image_conditioning: "MagiHuman preset variant marker for reference-image conditioning."
|
||||
ref_audio_offset: "MagiHuman internal data-proxy audio alignment offset."
|
||||
scheduler_sigma_min: "Z-Image scheduler parity invariant; not part of the public typed inference API."
|
||||
scheduler_use_reference_discrete_timesteps: "Z-Image scheduler parity invariant; not part of the public typed inference API."
|
||||
sr_local_attn_layers: "MagiHuman SR internal sparse-attention layer selection."
|
||||
text_offset: "MagiHuman internal data-proxy text alignment offset."
|
||||
vae_stride: "MagiHuman internal VAE/data-proxy stride setting."
|
||||
@@ -438,6 +450,7 @@ surfaces:
|
||||
grid_sizes: request.inputs.grid_sizes
|
||||
pose: request.inputs.pose
|
||||
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
|
||||
action_path: request.inputs.action_path
|
||||
refine_from: request.inputs.refine_from
|
||||
stage1_video: request.inputs.stage1_video
|
||||
prompt: request.prompt
|
||||
@@ -447,6 +460,7 @@ surfaces:
|
||||
output_video_name: request.output.output_video_name
|
||||
num_videos_per_prompt: request.sampling.num_videos_per_prompt
|
||||
seed: request.sampling.seed
|
||||
max_sequence_length: request.sampling.max_sequence_length
|
||||
num_frames: request.sampling.num_frames
|
||||
height: request.sampling.height
|
||||
width: request.sampling.width
|
||||
@@ -456,7 +470,10 @@ surfaces:
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
num_inference_steps_sr: request.sampling.num_inference_steps_sr
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
batch_cfg: request.sampling.batch_cfg
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
cfg_normalization: request.sampling.cfg_normalization
|
||||
cfg_truncation: request.sampling.cfg_truncation
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
@@ -508,7 +525,6 @@ surfaces:
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
latents: "Pre-generated diffusion latents supplied by parity/debug harnesses; not a public input."
|
||||
max_sequence_length: "Model-specific text-encoder sequence cap; not part of the public typed inference API."
|
||||
|
||||
sampling_param_extensions: {}
|
||||
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
|
||||
|
||||
!!! tip "Attn-QAT DMD2 workflow"
|
||||
The modular trainer also provides a Wan2.1 MixKit recipe that first
|
||||
fine-tunes with fake-quantized attention, then distills the student to
|
||||
timesteps `[1000, 757, 522]` while teacher and critic remain on Flash
|
||||
Attention. See [Attn-QAT Training](../training/attn_qat.md).
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
We provide two distilled models:
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# Attn-QAT Training
|
||||
|
||||
Attn-QAT simulates low-bit attention during training while keeping the rest of
|
||||
the training method unchanged. In the modular `fastvideo/train` framework it is
|
||||
a per-role model option, not a separate training method: supervised fine-tuning
|
||||
and DMD2 still own their losses and optimizer cadence.
|
||||
|
||||
This guide covers the QAD Wan2.1-T2V-1.3B MixKit workflow:
|
||||
|
||||
1. run a 4,000-step supervised Attn-QAT fine-tune;
|
||||
2. export the stage-1 DCP checkpoint to Diffusers format; and
|
||||
3. distill the student to three denoising steps with DMD2.
|
||||
|
||||
The ready-to-run configs and wrappers are in
|
||||
`examples/train/scenario/qad_wan2_1_mixkit/`.
|
||||
|
||||
## Role-local attention backends
|
||||
|
||||
A DMD2 run owns three independent model roles. Configure the attention backend
|
||||
on each role so fake quantization is applied only to the student:
|
||||
|
||||
```yaml
|
||||
models:
|
||||
student:
|
||||
attention_backend: ATTN_QAT_TRAIN
|
||||
teacher:
|
||||
attention_backend: FLASH_ATTN
|
||||
critic:
|
||||
attention_backend: FLASH_ATTN
|
||||
```
|
||||
|
||||
The override is active only while that role's transformer is constructed, then
|
||||
the previous process-wide backend is restored. This lets student, teacher, and
|
||||
critic use different implementations in one process. Invalid role-level names
|
||||
fail during configuration instead of silently selecting another backend.
|
||||
|
||||
See [Training Infrastructure](train_infra.md) for the complete model-role
|
||||
configuration reference.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Install FastVideo and make the `fastvideo-kernel` Python package importable.
|
||||
`ATTN_QAT_TRAIN` intentionally fails instead of falling back to dense
|
||||
attention when its kernel cannot be loaded.
|
||||
- Prepare the precomputed MixKit VAE latents and text embeddings.
|
||||
- Run the commands below from the repository root. The supplied recipe expects
|
||||
four GPUs by default; set `NUM_GPUS` to override it.
|
||||
|
||||
Download the published preprocessed dataset:
|
||||
|
||||
```bash
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
|
||||
```
|
||||
|
||||
## Stage 1: supervised Attn-QAT fine-tuning
|
||||
|
||||
The stage-1 config uses `ATTN_QAT_TRAIN` on the student, sequence parallelism
|
||||
across four GPUs, FP32 master weights, and 4,000 optimizer steps:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
|
||||
```
|
||||
|
||||
Pass a dataset directory as the first positional argument when it differs from
|
||||
the default:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh \
|
||||
/path/to/combined_parquet_dataset
|
||||
```
|
||||
|
||||
The wrapper calls `examples/train/run.sh`; the YAML file remains the source of
|
||||
truth for optimizer, validation, checkpointing, and distributed settings.
|
||||
|
||||
## Export the stage-1 checkpoint
|
||||
|
||||
Modular training checkpoints use Distributed Checkpoint (DCP) format. Export
|
||||
the student before using it to initialize stage 2:
|
||||
|
||||
```bash
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh \
|
||||
checkpoints/wan_t2v_qat_finetune/checkpoint-4000 \
|
||||
checkpoints/wan_t2v_qat_finetune/diffusers
|
||||
```
|
||||
|
||||
Both arguments are optional; the command above shows their defaults.
|
||||
|
||||
## Stage 2: three-step DMD2 distillation
|
||||
|
||||
Stage 2 loads the exported student weights, keeps Attn-QAT on the student, and
|
||||
uses Flash Attention for the teacher and critic:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh \
|
||||
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset \
|
||||
checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
|
||||
```
|
||||
|
||||
The migrated recipe preserves these behaviors:
|
||||
|
||||
| Behavior | Modular configuration |
|
||||
|---|---|
|
||||
| Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` |
|
||||
| Teacher and critic full-precision attention | Role-local `FLASH_ATTN` |
|
||||
| Generator update every five critic steps | `method.generator_update_interval: 5` |
|
||||
| Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` |
|
||||
| Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` |
|
||||
| Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` |
|
||||
| Stage handoff | DCP checkpoint to Diffusers export to student override weights |
|
||||
|
||||
The timestep ratios apply to randomly sampled teacher and critic score
|
||||
timesteps; `dmd_denoising_steps` separately controls the student rollout. See
|
||||
[DMD Distillation](../distillation/dmd.md) for general DMD concepts.
|
||||
|
||||
## Architecture-specific Triton routing
|
||||
|
||||
The training kernel is runtime-JIT-compiled Triton code and selects its route on
|
||||
every call. It supports different query and key/value sequence lengths for
|
||||
cross-attention; key and value must have the same sequence length.
|
||||
|
||||
| Hardware/configuration | Route |
|
||||
|---|---|
|
||||
| SM100, validated non-causal BF16 QAT configuration with head dimension 128 | Large-tile forward and split 64x64 backward; optimized backward requires a 16-aligned KV length |
|
||||
| SM120, including RTX 5090 | Previous forward tiling with joined quantized/STE P@V operations and a shallower backward pipeline for long sequences |
|
||||
| Unsupported configurations | Previous Triton implementation |
|
||||
|
||||
Warp specialization is disabled automatically on SM100 and SM120 because the
|
||||
Triton 3.7 NVWS compiler pass aborts for this kernel on Blackwell. No user
|
||||
setting is required.
|
||||
|
||||
The available tuning and comparison controls are:
|
||||
|
||||
| Environment variable | Default | Effect |
|
||||
|---|---|---|
|
||||
| `FASTVIDEO_ATTN_QAT_FWD_MODE` | `fast` | Selects `fast`, `balanced`, or `reference` forward tiling on the SM100 optimized route |
|
||||
| `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` | `0` | Set to `1` to recompute reference-order softmax statistics and keep `dV` bitwise-compatible on the SM100 optimized route |
|
||||
| `FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED` | `1` | Set to `0` to force the previous SM100 forward and backward for comparison |
|
||||
| `FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV` | `1` | Set to `0` to compare SM120 against the split P@V path |
|
||||
|
||||
The first invocation JIT-compiles the selected configuration; later calls reuse
|
||||
the Triton cache. To measure the production shape, run
|
||||
`python benchmarks/benchmark_attn_qat_train.py` from `fastvideo-kernel/`.
|
||||
|
||||
For import and backend-selection failures, see [Debugging](../utilities/debugging.md).
|
||||
@@ -62,6 +62,18 @@ bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
|
||||
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
|
||||
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
|
||||
|
||||
## Attention Quantization-Aware Training
|
||||
|
||||
Attn-QAT fine-tunes a model while simulating low-bit attention in the forward
|
||||
and backward passes. The modular trainer can select the backend per model role,
|
||||
so a later DMD2 stage can keep fake quantization on the student while the
|
||||
teacher and critic use Flash Attention.
|
||||
|
||||
The ready-to-run Wan2.1 MixKit workflow includes supervised fine-tuning,
|
||||
checkpoint export, and three-step DMD2 distillation:
|
||||
|
||||
**→ [Follow the Attn-QAT training guide](attn_qat.md)**
|
||||
|
||||
## LoRA Finetuning
|
||||
|
||||
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
|
||||
@@ -166,6 +178,7 @@ Ready-to-run training scripts are available for multiple models:
|
||||
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
|
||||
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
|
||||
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
|
||||
| Wan2.1 T2V 1.3B Attn-QAT | QAT SFT + DMD2 | `examples/train/scenario/qad_wan2_1_mixkit/` |
|
||||
|
||||
Each example includes:
|
||||
|
||||
|
||||
@@ -43,6 +43,9 @@ Ready-to-run examples with preprocessing scripts, training launchers, and valida
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
For the complete two-stage Wan2.1 MixKit quantization-aware workflow, see
|
||||
**[Attn-QAT Training](attn_qat.md)**.
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
@@ -59,9 +62,11 @@ FastVideo supports several training approaches:
|
||||
| **Full finetune** | Adapt entire model to a new domain or style |
|
||||
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
|
||||
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
|
||||
| **Attn-QAT** | Train with fake-quantized attention, optionally followed by DMD2 distillation |
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
|
||||
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
|
||||
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
3. **Train with quantized attention**: Follow the [Attn-QAT two-stage recipe](attn_qat.md)
|
||||
4. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
|
||||
@@ -81,6 +81,7 @@ Common model parameters:
|
||||
| `disable_custom_init_weights` | `false` | Skip custom weight initialization (use for teacher/critic) |
|
||||
| `flow_shift` | `3.0` | Timestep shifting factor |
|
||||
| `enable_gradient_checkpointing_type` | `null` | Gradient checkpointing (`"full"` or `null`) |
|
||||
| `attention_backend` | `null` | Optional role-local backend for Wan models (for example `ATTN_QAT_TRAIN`); overrides the process default only while this role's transformer is built |
|
||||
|
||||
Which roles are needed depends on the training method:
|
||||
|
||||
@@ -298,6 +299,8 @@ method:
|
||||
| `dmd_denoising_steps` | *(required)* | Timestep schedule for student rollout |
|
||||
| `generator_update_interval` | `1` | Update student every N critic steps |
|
||||
| `real_score_guidance_scale` | `1.0` | CFG scale for teacher predictions |
|
||||
| `min_timestep_ratio` | `0.0` | Lower bound for randomly sampled teacher/critic score timesteps |
|
||||
| `max_timestep_ratio` | `1.0` | Upper bound for randomly sampled teacher/critic score timesteps |
|
||||
| `fake_score_learning_rate` | *(required)* | Critic optimizer learning rate |
|
||||
| `fake_score_betas` | *(required)* | Critic optimizer Adam betas |
|
||||
| `fake_score_lr_scheduler` | *(required)* | Critic LR scheduler type |
|
||||
|
||||
@@ -78,7 +78,10 @@ If forcing a backend fails, verify optional dependencies are installed:
|
||||
- `SAGE_ATTN_THREE`: upstream `sageattn3` package
|
||||
- `ATTN_QAT_INFER`: `fastvideo-kernel` checkout/source install that exposes
|
||||
`attn_qat_infer`
|
||||
- `ATTN_QAT_TRAIN`: `fastvideo-kernel` install exposing `fastvideo_kernel`
|
||||
- `ATTN_QAT_TRAIN`: `fastvideo-kernel`; its runtime-JIT Triton implementation
|
||||
selects an optimized route on SM100, joins the quantized and STE P@V paths on
|
||||
SM120, and retains the previous route for unsupported configurations. See
|
||||
[Attn-QAT Training](../training/attn_qat.md) for architecture controls.
|
||||
|
||||
As a fallback, use:
|
||||
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
# LingBot World 2 Example Dataset
|
||||
|
||||
These files were copied unchanged from the LingBot World 2 repository for the
|
||||
FastVideo causal-fast inference example.
|
||||
|
||||
- Repository: `https://github.com/Robbyant/lingbot-world-v2.git`
|
||||
- Source commit: `94f43115de8d4a4f9f282126528c300a0b232c5f`
|
||||
- Source directory: `examples/03`
|
||||
|
||||
## Files
|
||||
|
||||
- `image.jpg`: source image for image-to-video generation. SHA-256:
|
||||
`6ee3dacfef32cfef504dd698adb8a660cf15f686535c52fed4903fef27c0edd0`
|
||||
- `poses.npy`: camera-to-world trajectory matrices. SHA-256:
|
||||
`bd0a23a696e184b0b43e7767eb432bfe644690560fe327fa96961affc941c404`
|
||||
- `intrinsics.npy`: camera intrinsic parameters. SHA-256:
|
||||
`821fca6cf957ae8fbb1181307f02479efb1705e04c9e05734cd02fb43462e082`
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.2 MiB |
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,81 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — image-to-video (I2V) path through
|
||||
# FastVideo's native Cosmos3 pipeline. The input image conditions latent frame 0
|
||||
# (kept clean during denoising); the rest of the clip is generated to follow it.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_i2v"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
image_path = os.environ.get("COSMOS3_IMAGE_PATH", "assets/images/cyclist.jpg")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A mountain biker rides forward along the sunlit forest trail, wheels "
|
||||
"kicking up dust as trees and dappled light sweep past, smooth cinematic "
|
||||
"tracking shot from behind."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,77 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — this example exercises the
|
||||
# text-to-video (T2V) path through FastVideo's native Cosmos3 pipeline.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A golden retriever puppy runs across a sunlit meadow toward the camera, "
|
||||
"ears flopping and wildflowers swaying in the breeze. Shallow depth of "
|
||||
"field, warm afternoon light, smooth cinematic tracking shot."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,78 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — text-to-image (T2I) path through
|
||||
# FastVideo's native Cosmos3 pipeline. T2I is the single-frame case
|
||||
# (num_frames=1); the canonical Cosmos3 T2I resolution is 960x960 (the model's
|
||||
# "720" bucket, UniPC flow_shift=10.0). Point COSMOS3_MODEL_PATH at a local
|
||||
# diffusers checkpoint (e.g. ``official_weights/cosmos3``) to skip the download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2i"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A photograph of a red panda sitting on a mossy log in a misty bamboo "
|
||||
"forest, soft golden morning light filtering through the leaves, shallow "
|
||||
"depth of field, crisp fur detail, serene atmosphere."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# T2I is single-frame; canonical Cosmos3 T2I is 960x960. Overridable
|
||||
# via env for quick smoke runs.
|
||||
num_frames=1,
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "960")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "960")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate image: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,67 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
# t2vs (text -> video + sound). The Cosmos3 denoise stage generates a joint
|
||||
# [vision | sound] latent and AVAE-decodes the sound to a waveform muxed into the
|
||||
# mp4. The joint-sound path is gated on COSMOS3_T2VS (set here for the example).
|
||||
os.environ.setdefault("COSMOS3_T2VS", "1")
|
||||
|
||||
from fastvideo import VideoGenerator # noqa: E402
|
||||
from fastvideo.api import ( # noqa: E402
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2vs"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(text_encoder=True, pin_cpu_memory=True, dit=False, vae=False),
|
||||
),
|
||||
)
|
||||
|
||||
load_start = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start
|
||||
|
||||
prompt = (
|
||||
"Ocean waves crash against a rocky shore at sunset, white foam spraying "
|
||||
"into the air as seagulls wheel overhead. Golden light, cinematic wide "
|
||||
"shot, the rhythmic roar of the surf."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True, return_frames=False),
|
||||
)
|
||||
|
||||
start = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video+sound: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Generate a five-second Dense LingBot-Video clip with the official defaults."""
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
"""Parse the converted checkpoint and output paths for the sample."""
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
type=Path,
|
||||
required=True,
|
||||
help="Path to a converted Dense LingBot-Video checkpoint.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=Path,
|
||||
default=Path("outputs/lingbot-video/dense-t2v"),
|
||||
help="Directory for the generated video.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Load the converted Dense checkpoint and generate the default T2V sample."""
|
||||
args = parse_args()
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(args.model_path),
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
try:
|
||||
generator.generate({
|
||||
"prompt": "A red fox runs through fresh snow at sunrise.",
|
||||
"output": {
|
||||
"output_path": str(args.output_path),
|
||||
"save_video": True,
|
||||
"return_frames": False,
|
||||
},
|
||||
})
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,52 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run LingBot World 2 14B causal-fast I2V generation with FastVideo."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
|
||||
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Load the native FastVideo LingBot World 2 causal-fast pipeline and generate one video."""
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
os.environ["LINGBOTWORLD2_MODEL_PATH"],
|
||||
num_gpus=8,
|
||||
sp_size=8,
|
||||
hsdp_shard_dim=8,
|
||||
use_fsdp_inference=True,
|
||||
dit_layerwise_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=True,
|
||||
override_pipeline_cls_name="LingBotWorld2CausalFastPipeline",
|
||||
)
|
||||
|
||||
try:
|
||||
generator.generate_video(
|
||||
"A serene lakeside scene with a lone tree standing in calm water, surrounded by distant snow-capped mountains under a bright blue sky with drifting white clouds; gentle ripples reflect the tree and sky, creating a tranquil, meditative atmosphere.",
|
||||
image_path=str(DATASET_DIR / "image.jpg"),
|
||||
action_path=str(DATASET_DIR),
|
||||
output_path=str(OUTPUT_PATH),
|
||||
save_video=True,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=65,
|
||||
num_inference_steps=4,
|
||||
guidance_scale=1.0,
|
||||
negative_prompt="",
|
||||
fps=16,
|
||||
seed=42,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,96 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run Z-Image-Turbo text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I want the official Z-Image-Turbo defaults and a deterministic PNG from
|
||||
a local or Hugging Face checkpoint."
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_PROMPT = (
|
||||
"Young Chinese woman in red Hanfu, intricate embroidery. Impeccable makeup, red floral forehead pattern. "
|
||||
"Elaborate high bun, golden phoenix headdress, red flowers, beads. Holds round folding fan with lady, trees, bird. "
|
||||
"Neon lightning-bolt lamp (⚡️), bright yellow glow, above extended left palm. Soft-lit outdoor night background, "
|
||||
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights."
|
||||
)
|
||||
DEFAULT_REVISION = "f332072aa78be7aecdf3ee76d5c247082da564a6"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run Z-Image-Turbo text-to-image generation.")
|
||||
parser.add_argument("--model-path", default="Tongyi-MAI/Z-Image-Turbo")
|
||||
parser.add_argument("--revision", default=DEFAULT_REVISION)
|
||||
parser.add_argument("--output", default="outputs/zimage/zimage_turbo.png")
|
||||
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
|
||||
parser.add_argument("--negative-prompt", default="")
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=8)
|
||||
parser.add_argument("--guidance-scale", type=float, default=0.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=512)
|
||||
parser.add_argument("--cfg-normalization", action=argparse.BooleanOptionalAction, default=False)
|
||||
parser.add_argument("--cfg-truncation", type=float, default=1.0)
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
revision=args.revision,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=1),
|
||||
use_fsdp_inference=False,
|
||||
),
|
||||
# The model registry selects the native zimage_turbo preset.
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
))
|
||||
try:
|
||||
generator.generate(
|
||||
GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
cfg_normalization=args.cfg_normalization,
|
||||
cfg_truncation=args.cfg_truncation,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
))
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -91,3 +91,7 @@ examples/train/
|
||||
```
|
||||
|
||||
See `configs/README.md` and `scenario/README.md` for details.
|
||||
|
||||
The featured QAD Wan2.1 MixKit scenario runs Attn-QAT supervised fine-tuning,
|
||||
checkpoint export, and three-step DMD2. See the
|
||||
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
|
||||
|
||||
@@ -25,6 +25,7 @@ models:
|
||||
disable_custom_init_weights: false # default: false
|
||||
flow_shift: 3.0 # default: 3.0
|
||||
enable_gradient_checkpointing_type: null # default: null (falls back to training.model)
|
||||
attention_backend: null # default: null (global/default); role-local when set
|
||||
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
@@ -52,6 +53,8 @@ method:
|
||||
rollout_mode: simulate # required: "simulate" or "data_latent"
|
||||
generator_update_interval: 5 # default: 1
|
||||
dmd_denoising_steps: [1000, 750, 500, 250] # SDE timestep schedule
|
||||
min_timestep_ratio: 0.0 # score-model timestep lower bound
|
||||
max_timestep_ratio: 1.0 # score-model timestep upper bound
|
||||
|
||||
# Critic optimizer (all required — no fallback)
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# LTX-2.3 T2V overfitting test config.
|
||||
#
|
||||
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
|
||||
# verify the LTX-2 training plugin works end-to-end on an LTX-2.3
|
||||
# checkpoint (gated attention, cross-attention AdaLN, 4096-d
|
||||
# post-connector text embeddings, no in-DiT caption projection).
|
||||
# Uses the distilled checkpoint (validation is 8 sampling steps).
|
||||
#
|
||||
# Preprocess data first (writes data/ltx2_3_overfit_preprocessed):
|
||||
# CUDA_VISIBLE_DEVICES=0 \
|
||||
# LTX2_OVERFIT_MODEL=FastVideo/LTX-2.3-Distilled-Diffusers \
|
||||
# LTX2_OVERFIT_OUTPUT_DIR=data/ltx2_3_overfit_preprocessed \
|
||||
# python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_3_t2v.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.ltx2.LTX2Model
|
||||
init_from: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: data/ltx2_3_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
# LTX2Model requires 0.0: CFG dropout would zero post-connector
|
||||
# embeddings, which is not the model's unconditional input.
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 11 # (81 - 1) / 8 + 1
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/ltx2_3_overfit
|
||||
# A full training-state checkpoint is ~150GB for the 13B trainable
|
||||
# video branch; disable saves for the overfit smoke run.
|
||||
training_state_checkpointing_steps: 0
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: fastvideo_ltx2
|
||||
run_name: ltx2_3_overfit
|
||||
|
||||
model:
|
||||
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
|
||||
# so the default noise-minus-clean target reproduces the official
|
||||
# unweighted masked-MSE (mask is all-ones for plain T2V).
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
|
||||
dataset_file: data/ltx2_3_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 50
|
||||
sampling_steps: [8]
|
||||
guidance_scale: 1.0
|
||||
num_frames: 81
|
||||
|
||||
# Required so the LTX2T2VConfig pipeline config is resolved from
|
||||
# init_from (without a `pipeline:` key the loader falls back to a
|
||||
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,90 @@
|
||||
# LTX-2 T2V overfitting test config.
|
||||
#
|
||||
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
|
||||
# verify the LTX-2 training plugin works end-to-end. Uses the
|
||||
# distilled checkpoint (validation is 8 sampling steps, single pass).
|
||||
#
|
||||
# Preprocess data first (writes data/ltx2_overfit_preprocessed):
|
||||
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_t2v.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.ltx2.LTX2Model
|
||||
init_from: FastVideo/LTX2-Distilled-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: data/ltx2_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
# LTX2Model requires 0.0: CFG dropout would zero post-connector
|
||||
# embeddings, which is not the model's unconditional input.
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 11 # (81 - 1) / 8 + 1
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/ltx2_overfit
|
||||
# A full training-state checkpoint is ~150GB for the 13B trainable
|
||||
# video branch; disable saves for the overfit smoke run.
|
||||
training_state_checkpointing_steps: 0
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: fastvideo_ltx2
|
||||
run_name: ltx2_overfit
|
||||
|
||||
model:
|
||||
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
|
||||
# so the default noise-minus-clean target reproduces the official
|
||||
# unweighted masked-MSE (mask is all-ones for plain T2V).
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
|
||||
dataset_file: data/ltx2_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 50
|
||||
sampling_steps: [8]
|
||||
guidance_scale: 1.0
|
||||
num_frames: 81
|
||||
|
||||
# Required so the LTX2T2VConfig pipeline config is resolved from
|
||||
# init_from (without a `pipeline:` key the loader falls back to a
|
||||
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
|
||||
pipeline: {}
|
||||
@@ -5,9 +5,12 @@ all configs, scripts, and data needed to run a complete workflow.
|
||||
|
||||
```
|
||||
scenario/
|
||||
└── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
|
||||
├── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
|
||||
└── qad_wan2_1_mixkit/ # Attn-QAT SFT → export → DMD2
|
||||
```
|
||||
|
||||
See the `usage.md` inside each scenario for step-by-step instructions.
|
||||
See the `usage.md` inside each scenario for step-by-step instructions. The QAD
|
||||
workflow is also documented in the website-visible
|
||||
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
|
||||
|
||||
For single-step configs, see `examples/train/configs/`.
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
#!/usr/bin/env bash
|
||||
# Export a modular-trainer DCP checkpoint for stage-2 initialization.
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
|
||||
cd "${REPO_ROOT}"
|
||||
|
||||
CHECKPOINT_DIR=${1:-checkpoints/wan_t2v_qat_finetune/checkpoint-4000}
|
||||
OUTPUT_DIR=${2:-checkpoints/wan_t2v_qat_finetune/diffusers}
|
||||
|
||||
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
--role student \
|
||||
--checkpoint "${CHECKPOINT_DIR}" \
|
||||
--output-dir "${OUTPUT_DIR}"
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
#!/usr/bin/env bash
|
||||
# Run the modular Attn-QAT finetune recipe.
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
|
||||
cd "${REPO_ROOT}"
|
||||
|
||||
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
export NUM_GPUS
|
||||
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
|
||||
|
||||
bash "${REPO_ROOT}/examples/train/run.sh" \
|
||||
"${SCRIPT_DIR}/stage1_attn_qat_finetune.yaml" \
|
||||
--training.data.data_path "${DATA_DIR}" \
|
||||
--training.distributed.num_gpus "${NUM_GPUS}" \
|
||||
--training.distributed.sp_size "${NUM_GPUS}" \
|
||||
--training.distributed.hsdp_replicate_dim 1 \
|
||||
--training.distributed.hsdp_shard_dim "${NUM_GPUS}"
|
||||
+27
@@ -0,0 +1,27 @@
|
||||
#!/usr/bin/env bash
|
||||
# Run modular DMD2 with Attn-QAT on the student only.
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
|
||||
cd "${REPO_ROOT}"
|
||||
|
||||
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
|
||||
INIT_WEIGHTS=${2:-checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors}
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
export NUM_GPUS
|
||||
|
||||
if [[ ! -f "${INIT_WEIGHTS}" ]]; then
|
||||
echo "Missing exported stage-1 weights: ${INIT_WEIGHTS}" >&2
|
||||
echo "Run export_stage1.sh before stage 2." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
bash "${REPO_ROOT}/examples/train/run.sh" \
|
||||
"${SCRIPT_DIR}/stage2_attn_qat_dmd.yaml" \
|
||||
--models.student.transformer_override_safetensor "${INIT_WEIGHTS}" \
|
||||
--training.data.data_path "${DATA_DIR}" \
|
||||
--training.distributed.num_gpus "${NUM_GPUS}" \
|
||||
--training.distributed.sp_size 1 \
|
||||
--training.distributed.hsdp_replicate_dim "${NUM_GPUS}" \
|
||||
--training.distributed.hsdp_shard_dim 1
|
||||
@@ -0,0 +1,72 @@
|
||||
# QAD stage 1: Attn-QAT finetune of Wan2.1-T2V-1.3B on MixKit.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
# Role-local: does not change the backend used by other models loaded in
|
||||
# this process.
|
||||
attention_backend: ATTN_QAT_TRAIN
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 4
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.1
|
||||
seed: 1000
|
||||
num_latent_t: 20
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: checkpoints/wan_t2v_qat_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 0
|
||||
|
||||
tracker:
|
||||
project_name: wan_t2v_qat_finetune
|
||||
run_name: wan_t2v_qat_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_pipeline.WanPipeline
|
||||
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
|
||||
every_steps: 50
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 5.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1
|
||||
@@ -0,0 +1,101 @@
|
||||
# QAD stage 2: distill the Attn-QAT student to three sampling steps.
|
||||
#
|
||||
# Export the stage-1 DCP checkpoint first, then point
|
||||
# models.student.transformer_override_safetensor at the exported weight file.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
transformer_override_safetensor: checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
|
||||
trainable: true
|
||||
attention_backend: ATTN_QAT_TRAIN
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
attention_backend: FLASH_ATTN
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
attention_backend: FLASH_ATTN
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
|
||||
rollout_mode: data_latent
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 757, 522]
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
# The modular DMD method uses standard CFG:
|
||||
# uncond + scale * (cond - uncond). This is equivalent to the legacy
|
||||
# recipe's cond + 2.0 * (cond - uncond).
|
||||
real_score_guidance_scale: 3.0
|
||||
|
||||
# The legacy recipe inherited these values from its global optimizer.
|
||||
fake_score_learning_rate: 2.0e-6
|
||||
fake_score_betas: [0.9, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 20
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 77
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 2000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: checkpoints/wan_t2v_distill_dmd_qat
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: wan_t2v_distill_dmd_qat
|
||||
run_name: wan_t2v_distill_dmd_qat
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_dmd_pipeline.WanDMDPipeline
|
||||
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
|
||||
every_steps: 200
|
||||
sampling_steps: [3]
|
||||
sampling_timesteps: [1000, 757, 522]
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,26 @@
|
||||
# QAD Wan2.1 MixKit Attn-QAT
|
||||
|
||||
This scenario runs entirely on the modular `fastvideo/train` stack. The
|
||||
student's attention backend is configured per role, so DMD2 can keep the
|
||||
teacher and critic on Flash Attention while the student uses the fake-quantized
|
||||
Attn-QAT kernel.
|
||||
|
||||
From the repository root:
|
||||
|
||||
```bash
|
||||
# 1. Download the preprocessed MixKit data.
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
|
||||
|
||||
# 2. Stage 1: Attn-QAT supervised finetune.
|
||||
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
|
||||
|
||||
# 3. Export the stage-1 DCP checkpoint to a Diffusers weight file.
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh
|
||||
|
||||
# 4. Stage 2: three-step Attn-QAT DMD2 distillation.
|
||||
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh
|
||||
```
|
||||
|
||||
The two YAML configs are also directly runnable through `examples/train/run.sh`.
|
||||
The wrapper scripts only provide dataset/checkpoint paths and derive distributed
|
||||
dimensions from `NUM_GPUS`.
|
||||
@@ -52,7 +52,7 @@ for the full parameter reference.
|
||||
## Train (QAT finetune)
|
||||
|
||||
With the data in place, run the quantization-aware finetune. The 4-bit attention
|
||||
path is **config-driven** — selected purely by an env var, no monkey-patching:
|
||||
path is **config-driven** and selected by an environment variable:
|
||||
|
||||
```bash
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
|
||||
@@ -60,10 +60,19 @@ bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
|
||||
NUM_GPUS=4 bash .../mixkit/finetune_qat.sh data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/
|
||||
```
|
||||
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` routes attention through the
|
||||
fake-quantized Triton kernel (straight-through estimator), so the DiT learns to
|
||||
absorb FP4 attention error. This kernel is Triton, so it runs on both `sm_100`
|
||||
(B200/GB200) and `sm_120` (RTX 5090).
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` keeps the fake-quantized Triton
|
||||
forward and backward (straight-through estimator). Both kernels ship in
|
||||
`fastvideo-kernel`: SM100 automatically uses the optimized Triton path for the
|
||||
production non-causal, head-dimension-128 configuration, while SM120 GPUs such
|
||||
as RTX 5090 join the quantized and STE P@V operations and use a shallower
|
||||
backward pipeline for long sequences. Set
|
||||
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
|
||||
The script defaults
|
||||
`FASTVIDEO_ATTN_QAT_FWD_EXACT_M=0` for throughput; set it to `1` to reproduce
|
||||
the previous forward softmax statistic and bitwise-compatible `dV`.
|
||||
|
||||
For the website-visible modular SFT-to-DMD2 workflow, see the
|
||||
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
|
||||
|
||||
## Train stage 2 (QAT DMD distillation to 3 steps)
|
||||
|
||||
@@ -71,8 +80,8 @@ Distill the QAT-finetuned generator down to **3 sampling steps**. Only the
|
||||
generator is quantized (Attn-QAT); the teacher (`real_score`) and critic
|
||||
(`fake_score`) stay full precision. This is enforced in the loader
|
||||
(`component_loader.py`, via the `_loading_teacher_critic_model` flag), so the
|
||||
same global `ATTN_QAT_TRAIN` env reaches **only** the generator — no per-model
|
||||
flags or monkey-patching.
|
||||
same global `ATTN_QAT_TRAIN` env reaches **only** the generator, with no
|
||||
per-model flags.
|
||||
|
||||
```bash
|
||||
# generator init = the stage-1 finetune checkpoint
|
||||
|
||||
@@ -2,20 +2,23 @@
|
||||
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
|
||||
# (Attn-QAT) attention.
|
||||
#
|
||||
# The 4-bit attention path is selected purely by env var (config-driven, no
|
||||
# monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN routes attention
|
||||
# through the fake-quantized Triton kernel (straight-through estimator), so the
|
||||
# DiT learns to absorb FP4 attention error instead of fighting it.
|
||||
# The 4-bit attention path is selected by env var:
|
||||
# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN keeps the fake-quantized Triton
|
||||
# forward and backward, so the DiT learns to absorb FP4 attention error instead
|
||||
# of fighting it. SM100 selects the optimized kernels; RTX 5090 keeps the
|
||||
# previous Triton implementation.
|
||||
#
|
||||
# Data: run download_mixkit_data.sh first (preprocessed Parquet).
|
||||
#
|
||||
# Verified end-to-end on Blackwell (GB200/sm_100): the ATTN_QAT_TRAIN backend is
|
||||
# selected (not a fallback), forward+backward run, loss/grad are healthy, and
|
||||
# validation generates videos. The kernel is Triton so it runs on sm_100 and
|
||||
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
|
||||
# validation generates videos.
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../../../.." && pwd)"
|
||||
export PYTHONPATH="${REPO_ROOT}/fastvideo-kernel/python${PYTHONPATH:+:${PYTHONPATH}}"
|
||||
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
|
||||
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
|
||||
export WANDB_MODE=${WANDB_MODE:-online}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
@@ -23,6 +26,8 @@ MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
|
||||
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
MAX_TRAIN_STEPS=${MAX_TRAIN_STEPS:-4000}
|
||||
VALIDATION_SAMPLING_STEPS=${VALIDATION_SAMPLING_STEPS:-50}
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
@@ -30,15 +35,15 @@ torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
--hsdp_replicate_dim 1 --hsdp_shard_dim "${NUM_GPUS}" \
|
||||
--model_path "${MODEL_PATH}" --pretrained_model_name_or_path "${MODEL_PATH}" \
|
||||
--data_path "${DATA_DIR}" --dataloader_num_workers 1 \
|
||||
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--max_train_steps "${MAX_TRAIN_STEPS}" --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
|
||||
--enable_gradient_checkpointing_type full \
|
||||
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
|
||||
--validation_steps 200 --validation_sampling_steps 50 --validation_guidance_scale 3.0 \
|
||||
--learning_rate 5e-5 --mixed_precision bf16 --weight_decay 1e-4 --max_grad_norm 1.0 \
|
||||
--validation_steps 50 --validation_sampling_steps "${VALIDATION_SAMPLING_STEPS}" --validation_guidance_scale 5.0 \
|
||||
--learning_rate 1e-6 --mixed_precision bf16 --weight_decay 0.01 --max_grad_norm 1.0 \
|
||||
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
|
||||
--tracker_project_name wan_t2v_qat_finetune --output_dir checkpoints/wan_t2v_qat_finetune \
|
||||
--inference_mode False --training_cfg_rate 0.1 --not_apply_cfg_solver \
|
||||
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 \
|
||||
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 --flow_shift 1 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -108,6 +108,33 @@ out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
|
||||
|
||||
## Benchmark
|
||||
|
||||
### Attn-QAT training
|
||||
|
||||
The default shape matches one sequence-parallel rank of the 4-GPU
|
||||
Wan2.1-T2V-1.3B MixKit recipe (`B=1, H=3, L=31200, D=128`):
|
||||
|
||||
```bash
|
||||
cd fastvideo-kernel
|
||||
python benchmarks/benchmark_attn_qat_train.py
|
||||
```
|
||||
|
||||
The benchmark reports both conventional attention FLOPs and the extra matrix
|
||||
multiplications executed by the QAT straight-through path. Override
|
||||
`--peak-tflops` when running on a GPU other than RTX 5090.
|
||||
|
||||
The QAT kernel is entirely Triton and routes by architecture at runtime. SM100
|
||||
uses a large-tile forward and split 64x64 backward for the production
|
||||
non-causal, head-dimension-128 configuration with a 16-aligned KV length. SM120
|
||||
(including RTX 5090) keeps the previous tiling but joins the quantized and STE
|
||||
P@V operations and uses a shallower backward pipeline for long sequences. Set
|
||||
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
|
||||
Unsupported configurations retain the previous implementation. Set
|
||||
`FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED=0` to benchmark that previous path on SM100. Forward tuning is available through
|
||||
`FASTVIDEO_ATTN_QAT_FWD_MODE=fast|balanced|reference`; exact reference-order
|
||||
softmax statistics are controlled by `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` and are
|
||||
disabled by default for maximum throughput. Set it to `1` for reference-order
|
||||
statistics and bitwise-compatible `dV`.
|
||||
|
||||
### VSA (block-sparse) TFLOPs
|
||||
|
||||
After building/installing `fastvideo-kernel`, run:
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Benchmark the Attn-QAT training kernel on a single GPU.
|
||||
|
||||
Defaults model one rank of the 4-GPU Wan2.1-T2V-1.3B MixKit recipe:
|
||||
``B=1, H=12/4, L=20*30*52, D=128``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import statistics
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel.triton_kernels.attn_qat_train import attention
|
||||
|
||||
|
||||
RTX_5090_DENSE_BF16_TFLOPS = 209.5
|
||||
|
||||
|
||||
def _qat_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
||||
consumer_blackwell = torch.cuda.get_device_capability()[0] == 12
|
||||
return attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
False,
|
||||
q.shape[-1]**-0.5,
|
||||
True, # use_qat_qkv_backward
|
||||
False, # smooth_k
|
||||
not consumer_blackwell, # warp_specialize
|
||||
True, # IS_QAT
|
||||
False, # two_level_quant_P
|
||||
True, # fake_quant_P
|
||||
True, # use_high_prec_o
|
||||
False, # smooth_q
|
||||
False, # use_global_sf_P
|
||||
False, # use_global_sf_QKV
|
||||
)
|
||||
|
||||
|
||||
def _measure_ms(fn: Callable[[], object], warmup: int, repeat: int) -> tuple[float, float, float]:
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
samples = []
|
||||
for _ in range(repeat):
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
fn()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
samples.append(start.elapsed_time(end))
|
||||
return statistics.median(samples), min(samples), max(samples)
|
||||
|
||||
|
||||
def _format_result(
|
||||
label: str,
|
||||
timing_ms: tuple[float, float, float],
|
||||
algorithmic_flops: int,
|
||||
executed_matmul_flops: int,
|
||||
peak_tflops: float,
|
||||
) -> str:
|
||||
median_ms, min_ms, max_ms = timing_ms
|
||||
algorithmic_tflops = algorithmic_flops / (median_ms * 1e9)
|
||||
executed_tflops = executed_matmul_flops / (median_ms * 1e9)
|
||||
return (
|
||||
f"{label}: {median_ms:.3f} ms (min={min_ms:.3f}, max={max_ms:.3f}), "
|
||||
f"algorithmic={algorithmic_tflops:.2f} TFLOPS/{100 * algorithmic_tflops / peak_tflops:.2f}% MFU, "
|
||||
f"executed_matmul={executed_tflops:.2f} TFLOPS/{100 * executed_tflops / peak_tflops:.2f}% MFU"
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--batch-size", type=int, default=1)
|
||||
parser.add_argument("--heads", type=int, default=3, help="Heads per SP rank; Wan 1.3B has 12 total.")
|
||||
parser.add_argument("--query-length", type=int, default=31_200)
|
||||
parser.add_argument("--kv-length", type=int, help="Defaults to --query-length.")
|
||||
parser.add_argument("--head-dim", type=int, default=128)
|
||||
parser.add_argument("--warmup", type=int, default=3)
|
||||
parser.add_argument("--repeat", type=int, default=10)
|
||||
parser.add_argument(
|
||||
"--peak-tflops",
|
||||
type=float,
|
||||
default=RTX_5090_DENSE_BF16_TFLOPS,
|
||||
help="Dense BF16 Tensor TFLOPS with FP32 accumulation; default is RTX 5090 boost-clock peak.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
kv_length = args.kv_length or args.query_length
|
||||
|
||||
torch.manual_seed(0)
|
||||
q_shape = (args.batch_size, args.heads, args.query_length, args.head_dim)
|
||||
kv_shape = (args.batch_size, args.heads, kv_length, args.head_dim)
|
||||
q = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
|
||||
k = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
|
||||
v = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
|
||||
grad_out = torch.randn_like(q)
|
||||
|
||||
compile_start = time.perf_counter()
|
||||
output = _qat_attention(q, k, v)
|
||||
torch.cuda.synchronize()
|
||||
compile_seconds = time.perf_counter() - compile_start
|
||||
|
||||
forward_ms = _measure_ms(lambda: _qat_attention(q, k, v), args.warmup, args.repeat)
|
||||
backward_ms = _measure_ms(
|
||||
lambda: torch.autograd.grad(output, (q, k, v), grad_out, retain_graph=True),
|
||||
args.warmup,
|
||||
args.repeat,
|
||||
)
|
||||
|
||||
base_flops = args.batch_size * args.heads * args.query_length * kv_length * args.head_dim
|
||||
# Conventional attention FLOPs are 4*base forward and 10*base backward.
|
||||
# QAT additionally computes the STE high-precision P@V path in forward and
|
||||
# the quantized-P dV path in backward, for 6*base and 14*base matmul FLOPs.
|
||||
print(f"device: {torch.cuda.get_device_name()}")
|
||||
print(f"q: {q_shape}; k/v: {kv_shape}; compile+first-forward: {compile_seconds:.3f} s")
|
||||
print(_format_result("forward", forward_ms, 4 * base_flops, 6 * base_flops, args.peak_tflops))
|
||||
print(_format_result("backward", backward_ms, 10 * base_flops, 14 * base_flops, args.peak_tflops))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -20,14 +20,68 @@ def supports_host_descriptor():
|
||||
return is_cuda() and torch.cuda.get_device_capability()[0] >= 9
|
||||
|
||||
|
||||
def is_sm100(device=None):
|
||||
return is_cuda() and torch.cuda.get_device_capability(device) == (10, 0)
|
||||
|
||||
|
||||
def is_blackwell():
|
||||
return is_cuda() and torch.cuda.get_device_capability()[0] == 10
|
||||
|
||||
|
||||
def is_consumer_blackwell():
|
||||
return is_cuda() and torch.cuda.get_device_capability()[0] == 12
|
||||
|
||||
|
||||
def is_hopper():
|
||||
return is_cuda() and torch.cuda.get_device_capability()[0] == 9
|
||||
|
||||
|
||||
def _sm100_optimization_enabled():
|
||||
return os.environ.get("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1") != "0"
|
||||
|
||||
|
||||
def _sm100_exact_m_enabled():
|
||||
return os.environ.get("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "0") != "0"
|
||||
|
||||
|
||||
def _consumer_blackwell_join_qat_pv_enabled():
|
||||
"""Return whether SM120 uses the joined quantized/STE P@V path."""
|
||||
return os.environ.get("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1") != "0"
|
||||
|
||||
|
||||
def _use_sm100_optimized_qat(
|
||||
device,
|
||||
head_dim: int,
|
||||
causal: bool,
|
||||
is_qat: bool,
|
||||
fake_quant_p: bool,
|
||||
two_level_quant_p: bool,
|
||||
use_global_sf_p: bool,
|
||||
) -> bool:
|
||||
"""Return whether this call matches the validated SM100 fast path."""
|
||||
return (
|
||||
_sm100_optimization_enabled()
|
||||
and is_sm100(device)
|
||||
and head_dim == 128
|
||||
and not causal
|
||||
and is_qat
|
||||
and fake_quant_p
|
||||
and not two_level_quant_p
|
||||
and not use_global_sf_p
|
||||
)
|
||||
|
||||
|
||||
def _select_sm100_forward_config(n_ctx_q: int, n_ctx_kv: int, mode: str):
|
||||
n_ctx = max(n_ctx_q, n_ctx_kv)
|
||||
if mode == "reference":
|
||||
return 32, 32, 4, 4 if n_ctx >= 16_384 else 5
|
||||
if n_ctx <= 2_048:
|
||||
return 32, 32, 4, 5
|
||||
if mode == "balanced":
|
||||
return 64, 32, 4, 4
|
||||
return 128, 128, 8, 3
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _mul_alpha(acc, alpha, BM: tl.constexpr, BN: tl.constexpr):
|
||||
acc0, acc1 = acc.reshape([BM, 2, BN // 2]).permute(0, 2, 1).split()
|
||||
@@ -47,7 +101,8 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
|
||||
IS_QAT: tl.constexpr,
|
||||
fake_quant_P: tl.constexpr = True,
|
||||
two_level_quant_P: tl.constexpr = False,
|
||||
use_global_sf_P: tl.constexpr = True):
|
||||
use_global_sf_P: tl.constexpr = True,
|
||||
JOIN_QAT_PV: tl.constexpr = False):
|
||||
# range of values handled by this stage (kv blocks)
|
||||
if STAGE == 1:
|
||||
lo, hi = 0, start_m * BLOCK_M
|
||||
@@ -113,10 +168,19 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
|
||||
v = desc_v.load([offsetv_y, 0])
|
||||
v = tl.where(kv_valid[:, None], v, 0.0)
|
||||
p = p.to(dtype)
|
||||
# note that this non transposed v for FP8 is only supported on Blackwell
|
||||
acc = tl.dot(p, v.to(dtype), acc)
|
||||
if IS_QAT:
|
||||
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
|
||||
# Keep the quantized and STE paths in one tensor-core operation. They
|
||||
# share V, so joining along M avoids issuing two small, independent
|
||||
# dot operations for every KV tile.
|
||||
if IS_QAT and JOIN_QAT_PV:
|
||||
joined_p = tl.join(p, high_prec_p).permute(2, 0, 1).reshape([2 * BLOCK_M, BLOCK_N])
|
||||
joined_acc = tl.join(acc, high_prec_acc).permute(2, 0, 1).reshape([2 * BLOCK_M, HEAD_DIM])
|
||||
joined_acc = tl.dot(joined_p, v.to(dtype), joined_acc)
|
||||
acc, high_prec_acc = joined_acc.reshape([2, BLOCK_M, HEAD_DIM]).permute(1, 2, 0).split()
|
||||
else:
|
||||
# note that this non transposed v for FP8 is only supported on Blackwell
|
||||
acc = tl.dot(p, v.to(dtype), acc)
|
||||
if IS_QAT:
|
||||
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
|
||||
# update m_i and l_i
|
||||
# place this at the end of the loop to reduce register pressure
|
||||
l_i = l_i * alpha + l_ij
|
||||
@@ -204,6 +268,7 @@ def _attn_fwd(sm_scale, M,
|
||||
fake_quant_P: tl.constexpr = True,
|
||||
two_level_quant_P: tl.constexpr = False,
|
||||
use_global_sf_P: tl.constexpr = True,
|
||||
JOIN_QAT_PV: tl.constexpr = False,
|
||||
):
|
||||
dtype = tl.float8e5 if FP8_OUTPUT else tl.bfloat16
|
||||
tl.static_assert(BLOCK_N <= HEAD_DIM)
|
||||
@@ -268,7 +333,7 @@ def _attn_fwd(sm_scale, M,
|
||||
offset_y_kv, dtype, start_m, qk_scale,
|
||||
BLOCK_M, HEAD_DIM, BLOCK_N,
|
||||
4 - STAGE, offs_m, offs_n, N_CTX_KV,
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
|
||||
)
|
||||
# stage 2: on-band
|
||||
if STAGE & 2:
|
||||
@@ -278,7 +343,7 @@ def _attn_fwd(sm_scale, M,
|
||||
offset_y_kv, dtype, start_m, qk_scale,
|
||||
BLOCK_M, HEAD_DIM, BLOCK_N,
|
||||
2, offs_m, offs_n, N_CTX_KV,
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
|
||||
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
|
||||
)
|
||||
# epilogue
|
||||
m_i += tl.math.log2(l_i)
|
||||
@@ -292,6 +357,52 @@ def _attn_fwd(sm_scale, M,
|
||||
desc_high_prec_o.store([off_hz, start_m * BLOCK_M, 0], high_prec_acc[None, :, :])
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_exact_m(
|
||||
desc_q,
|
||||
desc_k,
|
||||
M,
|
||||
sm_scale,
|
||||
N_CTX_Q,
|
||||
N_CTX_KV,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
"""Reproduce the legacy 32x32 forward softmax statistic exactly."""
|
||||
start_m = tl.program_id(0) * BLOCK_M
|
||||
off_hz = tl.program_id(1)
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
q_valid = offs_m < N_CTX_Q
|
||||
|
||||
q_base = off_hz * N_CTX_Q
|
||||
kv_base = off_hz * N_CTX_KV
|
||||
q = desc_q.load([q_base + start_m, 0])
|
||||
q = tl.where(q_valid[:, None], q, 0.0)
|
||||
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.full([BLOCK_M], 1.0, tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504
|
||||
|
||||
for start_n in tl.range(0, N_CTX_KV, BLOCK_N):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N)
|
||||
kv_valid = start_n + offs_n < N_CTX_KV
|
||||
k = desc_k.load([kv_base + start_n, 0])
|
||||
k = tl.where(kv_valid[:, None], k, 0.0)
|
||||
qk = tl.dot(q, tl.trans(k))
|
||||
qk = tl.where(kv_valid[None, :], qk, -1.0e6)
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, axis=1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p.to(tl.bfloat16), axis=1)
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
m_i = m_ij
|
||||
|
||||
m_i += tl.math.log2(l_i)
|
||||
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i, mask=q_valid)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO,
|
||||
Delta,
|
||||
@@ -835,10 +946,33 @@ class _attention(torch.autograd.Function):
|
||||
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
|
||||
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
|
||||
|
||||
# Triton 3.7's NVWS pass aborts for this kernel on Blackwell. Keep the
|
||||
# architecture guard next to the kernel so direct callers and the
|
||||
# FastVideo backend follow the same supported path on sm_100/sm_120.
|
||||
consumer_blackwell = is_consumer_blackwell()
|
||||
blackwell = is_blackwell() or consumer_blackwell
|
||||
warp_specialize = warp_specialize and not blackwell
|
||||
|
||||
# Support different sequence lengths for q and k/v (needed for cross attention)
|
||||
N_CTX_Q = q.shape[2] # Query sequence length
|
||||
N_CTX_KV = k.shape[2] # Key/Value sequence length (may differ from query)
|
||||
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
|
||||
sm100_optimized = (
|
||||
q.dtype == torch.bfloat16
|
||||
and k.dtype == q.dtype
|
||||
and v.dtype == q.dtype
|
||||
and k.device == q.device
|
||||
and v.device == q.device
|
||||
and _use_sm100_optimized_qat(
|
||||
q.device,
|
||||
HEAD_DIM_K,
|
||||
causal,
|
||||
IS_QAT,
|
||||
fake_quant_P,
|
||||
two_level_quant_P,
|
||||
use_global_sf_P,
|
||||
)
|
||||
)
|
||||
|
||||
# smoothing k from SageAttn
|
||||
ctx.k_mean = None
|
||||
@@ -924,7 +1058,20 @@ class _attention(torch.autograd.Function):
|
||||
else:
|
||||
extra_kern_args["maxnreg"] = 80
|
||||
|
||||
BLOCK_M, BLOCK_N = 32, 32
|
||||
qkv_block_m, qkv_block_n = 32, 32
|
||||
fwd_block_m, fwd_block_n = 32, 32
|
||||
fwd_num_warps, fwd_num_stages = 4, 2
|
||||
fwd_mode = "legacy"
|
||||
if sm100_optimized:
|
||||
fwd_mode = os.environ.get("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast").lower()
|
||||
if fwd_mode not in {"fast", "balanced", "reference"}:
|
||||
raise ValueError(
|
||||
f"FASTVIDEO_ATTN_QAT_FWD_MODE={fwd_mode!r} "
|
||||
"(want fast|balanced|reference)"
|
||||
)
|
||||
fwd_block_m, fwd_block_n, fwd_num_warps, fwd_num_stages = _select_sm100_forward_config(
|
||||
N_CTX_Q, N_CTX_KV, fwd_mode
|
||||
)
|
||||
if IS_QAT:
|
||||
fake_q = torch.empty_like(q)
|
||||
fake_k = torch.empty_like(k)
|
||||
@@ -942,8 +1089,8 @@ class _attention(torch.autograd.Function):
|
||||
desc_v = fake_v
|
||||
|
||||
H = q.shape[1]
|
||||
grid_1 = (triton.cdiv(q.shape[2], BLOCK_M), q.shape[0] * q.shape[1], 1)
|
||||
grid_2 = (triton.cdiv(k.shape[2], BLOCK_N), q.shape[0] * q.shape[1], 1)
|
||||
grid_1 = (triton.cdiv(q.shape[2], qkv_block_m), q.shape[0] * q.shape[1], 1)
|
||||
grid_2 = (triton.cdiv(k.shape[2], qkv_block_n), q.shape[0] * q.shape[1], 1)
|
||||
|
||||
fake_quantize_q[grid_1](
|
||||
q, fake_q,
|
||||
@@ -952,7 +1099,7 @@ class _attention(torch.autograd.Function):
|
||||
fake_q.stride(0), fake_q.stride(1),
|
||||
fake_q.stride(2), fake_q.stride(3),
|
||||
H, N_CTX_Q,
|
||||
BLOCK_M=BLOCK_M, HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_M=qkv_block_m, HEAD_DIM=HEAD_DIM_K,
|
||||
use_global_sf=use_global_sf_QKV,
|
||||
)
|
||||
fake_quantize_kv[grid_2](
|
||||
@@ -962,14 +1109,14 @@ class _attention(torch.autograd.Function):
|
||||
fake_k.stride(0), fake_k.stride(1),
|
||||
fake_k.stride(2), fake_k.stride(3),
|
||||
H, N_CTX_KV,
|
||||
BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_N=qkv_block_n, HEAD_DIM=HEAD_DIM_K,
|
||||
use_global_sf=use_global_sf_QKV,
|
||||
)
|
||||
|
||||
# Apply pre-hook to set block shapes on tensor descriptors
|
||||
_host_descriptor_pre_hook({
|
||||
"BLOCK_M": BLOCK_M,
|
||||
"BLOCK_N": BLOCK_N,
|
||||
"BLOCK_M": fwd_block_m,
|
||||
"BLOCK_N": fwd_block_n,
|
||||
"HEAD_DIM": HEAD_DIM_K,
|
||||
"desc_q": desc_q,
|
||||
"desc_k": desc_k,
|
||||
@@ -986,7 +1133,7 @@ class _attention(torch.autograd.Function):
|
||||
N_CTX_Q=N_CTX_Q,
|
||||
N_CTX_KV=N_CTX_KV,
|
||||
HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N,
|
||||
BLOCK_M=fwd_block_m, BLOCK_N=fwd_block_n,
|
||||
FP8_OUTPUT=q.dtype == torch.float8_e5m2,
|
||||
STAGE=stage,
|
||||
warp_specialize=warp_specialize,
|
||||
@@ -995,10 +1142,40 @@ class _attention(torch.autograd.Function):
|
||||
fake_quant_P=fake_quant_P,
|
||||
two_level_quant_P=two_level_quant_P,
|
||||
use_global_sf_P=use_global_sf_P,
|
||||
num_warps=4,
|
||||
num_stages=2,
|
||||
JOIN_QAT_PV=(consumer_blackwell and _consumer_blackwell_join_qat_pv_enabled()),
|
||||
num_warps=fwd_num_warps,
|
||||
num_stages=fwd_num_stages,
|
||||
**extra_kern_args
|
||||
)
|
||||
|
||||
exact_m = _sm100_exact_m_enabled()
|
||||
if (
|
||||
sm100_optimized
|
||||
and fwd_mode != "reference"
|
||||
and exact_m
|
||||
and (fwd_block_m != 32 or fwd_block_n != 32)
|
||||
):
|
||||
# The large forward tile changes a legal reduction order. Restore
|
||||
# the legacy statistic so dV remains bitwise-compatible while the
|
||||
# two output paths retain the faster large-tile PV computation.
|
||||
assert isinstance(desc_q, TensorDescriptor)
|
||||
assert isinstance(desc_k, TensorDescriptor)
|
||||
desc_q.block_shape = [32, HEAD_DIM_K]
|
||||
desc_k.block_shape = [32, HEAD_DIM_K]
|
||||
stats_grid = (triton.cdiv(N_CTX_Q, 32), q.shape[0] * q.shape[1])
|
||||
_attn_fwd_exact_m[stats_grid](
|
||||
desc_q,
|
||||
desc_k,
|
||||
M,
|
||||
sm_scale,
|
||||
N_CTX_Q,
|
||||
N_CTX_KV,
|
||||
HEAD_DIM=HEAD_DIM_K,
|
||||
BLOCK_M=32,
|
||||
BLOCK_N=32,
|
||||
num_warps=8,
|
||||
num_stages=4,
|
||||
)
|
||||
o_for_bwd = high_prec_o if IS_QAT and use_high_prec_o else o
|
||||
|
||||
if IS_QAT:
|
||||
@@ -1018,6 +1195,7 @@ class _attention(torch.autograd.Function):
|
||||
ctx.smooth_q = smooth_q
|
||||
ctx.use_global_sf_P = use_global_sf_P
|
||||
ctx.warp_specialize = warp_specialize
|
||||
ctx.sm100_optimized = sm100_optimized
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
@@ -1032,7 +1210,10 @@ class _attention(torch.autograd.Function):
|
||||
N_CTX_KV = k.shape[2]
|
||||
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
|
||||
PRE_BLOCK = 128
|
||||
NUM_STAGES = 3
|
||||
# Long video sequences are occupancy-bound on consumer Blackwell: a
|
||||
# third software-pipeline stage consumes shared memory without hiding
|
||||
# additional latency. Shorter sequences retain the deeper pipeline.
|
||||
NUM_STAGES = 2 if is_consumer_blackwell() and max(N_CTX_Q, N_CTX_KV) >= 8192 else 3
|
||||
NUM_WARPS = 4
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 32, 32, 32
|
||||
if not ctx.use_qat_qkv_backward:
|
||||
@@ -1057,7 +1238,85 @@ class _attention(torch.autograd.Function):
|
||||
# _, q_m = triton_group_mean(q)
|
||||
q_m = q_m.repeat_interleave(q.shape[2] // q_m.shape[2], dim=2) # B,H,L,D
|
||||
|
||||
if N_CTX_Q == N_CTX_KV:
|
||||
sm100_optimized_backward = (
|
||||
getattr(ctx, "sm100_optimized", False)
|
||||
and ctx.use_qat_qkv_backward
|
||||
and not ctx.smooth_k
|
||||
and not ctx.smooth_q
|
||||
and N_CTX_KV % 16 == 0
|
||||
)
|
||||
if sm100_optimized_backward:
|
||||
# Keeping dQ and dK/dV in separate programs allows 64x64 tiles
|
||||
# without carrying all three fp32 accumulators at once. On SM100
|
||||
# this is substantially faster than the legacy 32x32 combined
|
||||
# self-attention program with the same math and BF16 parity bounds.
|
||||
block_m, block_n = 64, 64
|
||||
grid_dq = ((N_CTX_Q + block_m - 1) // block_m, 1, BATCH * N_HEAD)
|
||||
_attn_bwd_dq_cross[grid_dq](
|
||||
q,
|
||||
arg_k,
|
||||
v,
|
||||
ctx.sm_scale,
|
||||
do,
|
||||
dq,
|
||||
M,
|
||||
delta,
|
||||
q.stride(0),
|
||||
k.stride(0),
|
||||
q.stride(1),
|
||||
k.stride(1),
|
||||
q.stride(2),
|
||||
k.stride(2),
|
||||
q.stride(3),
|
||||
k.stride(3),
|
||||
N_HEAD,
|
||||
N_CTX_Q,
|
||||
N_CTX_KV,
|
||||
ctx.k_mean,
|
||||
BLOCK_M2=block_m,
|
||||
BLOCK_N2=block_n,
|
||||
HEAD_DIM=ctx.HEAD_DIM,
|
||||
SMOOTH_K=False,
|
||||
warp_specialize=False,
|
||||
num_warps=8,
|
||||
num_stages=2,
|
||||
)
|
||||
grid_dkdv = ((N_CTX_KV + block_n - 1) // block_n, 1, BATCH * N_HEAD)
|
||||
_attn_bwd_dkdv_cross[grid_dkdv](
|
||||
q,
|
||||
arg_k,
|
||||
v,
|
||||
ctx.sm_scale,
|
||||
do,
|
||||
dk,
|
||||
dv,
|
||||
M,
|
||||
delta,
|
||||
q_m,
|
||||
q.stride(0),
|
||||
k.stride(0),
|
||||
q.stride(1),
|
||||
k.stride(1),
|
||||
q.stride(2),
|
||||
k.stride(2),
|
||||
q.stride(3),
|
||||
k.stride(3),
|
||||
N_HEAD,
|
||||
N_CTX_Q,
|
||||
N_CTX_KV,
|
||||
BLOCK_M1=block_m,
|
||||
BLOCK_N1=block_n,
|
||||
HEAD_DIM=ctx.HEAD_DIM,
|
||||
IS_QAT=True,
|
||||
two_level_quant_P=False,
|
||||
fake_quant_P=True,
|
||||
SMOOTH_Q=False,
|
||||
use_global_sf_P=False,
|
||||
warp_specialize=False,
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
elif N_CTX_Q == N_CTX_KV:
|
||||
# Use existing kernel for self-attention (same sequence lengths)
|
||||
grid = ((N_CTX_KV + BLOCK_N1 - 1) // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
@@ -1074,10 +1333,10 @@ class _attention(torch.autograd.Function):
|
||||
IS_QAT=ctx.IS_QAT,
|
||||
SMOOTH_K=ctx.smooth_k,
|
||||
two_level_quant_P=ctx.two_level_quant_P,
|
||||
fake_quant_P=ctx.fake_quant_P,
|
||||
SMOOTH_Q=ctx.smooth_q,
|
||||
use_global_sf_P=ctx.use_global_sf_P,
|
||||
warp_specialize=ctx.warp_specialize,
|
||||
fake_quant_P=ctx.fake_quant_P,
|
||||
SMOOTH_Q=ctx.smooth_q,
|
||||
use_global_sf_P=ctx.use_global_sf_P,
|
||||
warp_specialize=ctx.warp_specialize,
|
||||
num_warps=NUM_WARPS,
|
||||
num_stages=NUM_STAGES
|
||||
)
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel.triton_kernels import attn_qat_train as kernel
|
||||
|
||||
|
||||
def _production_route_kwargs():
|
||||
return {
|
||||
"device": torch.device("cuda"),
|
||||
"head_dim": 128,
|
||||
"causal": False,
|
||||
"is_qat": True,
|
||||
"fake_quant_p": True,
|
||||
"two_level_quant_p": False,
|
||||
"use_global_sf_p": False,
|
||||
}
|
||||
|
||||
|
||||
def test_sm100_production_configuration_uses_optimized_route(monkeypatch):
|
||||
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
|
||||
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", raising=False)
|
||||
|
||||
assert kernel._use_sm100_optimized_qat(**_production_route_kwargs())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("override", "value"),
|
||||
[
|
||||
("head_dim", 64),
|
||||
("causal", True),
|
||||
("is_qat", False),
|
||||
("fake_quant_p", False),
|
||||
("two_level_quant_p", True),
|
||||
("use_global_sf_p", True),
|
||||
],
|
||||
)
|
||||
def test_unsupported_configuration_keeps_legacy_route(monkeypatch, override, value):
|
||||
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
|
||||
kwargs = _production_route_kwargs()
|
||||
kwargs[override] = value
|
||||
|
||||
assert not kernel._use_sm100_optimized_qat(**kwargs)
|
||||
|
||||
|
||||
def test_non_sm100_and_debug_switch_keep_legacy_route(monkeypatch):
|
||||
kwargs = _production_route_kwargs()
|
||||
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: False)
|
||||
assert not kernel._use_sm100_optimized_qat(**kwargs)
|
||||
|
||||
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "0")
|
||||
assert not kernel._use_sm100_optimized_qat(**kwargs)
|
||||
|
||||
|
||||
def test_exact_m_is_opt_in(monkeypatch):
|
||||
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", raising=False)
|
||||
assert not kernel._sm100_exact_m_enabled()
|
||||
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
|
||||
assert kernel._sm100_exact_m_enabled()
|
||||
|
||||
|
||||
def test_sm120_joined_pv_is_enabled_by_default_and_can_be_disabled(monkeypatch):
|
||||
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", raising=False)
|
||||
assert kernel._consumer_blackwell_join_qat_pv_enabled()
|
||||
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "0")
|
||||
assert not kernel._consumer_blackwell_join_qat_pv_enabled()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("n_ctx", "mode", "expected"),
|
||||
[
|
||||
(2_048, "fast", (32, 32, 4, 5)),
|
||||
(4_096, "fast", (128, 128, 8, 3)),
|
||||
(4_096, "balanced", (64, 32, 4, 4)),
|
||||
(4_096, "reference", (32, 32, 4, 5)),
|
||||
(31_200, "reference", (32, 32, 4, 4)),
|
||||
],
|
||||
)
|
||||
def test_sm100_forward_config_selection(n_ctx, mode, expected):
|
||||
assert kernel._select_sm100_forward_config(n_ctx, n_ctx, mode) == expected
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0),
|
||||
reason="SM100 parity test",
|
||||
)
|
||||
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
|
||||
def test_sm100_optimized_forward_backward_matches_legacy(monkeypatch, q_length, kv_length):
|
||||
torch.manual_seed(7)
|
||||
q_shape = (1, 1, q_length, 128)
|
||||
kv_shape = (1, 1, kv_length, 128)
|
||||
inputs = [
|
||||
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
|
||||
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
|
||||
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
|
||||
]
|
||||
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
|
||||
flags = (
|
||||
True, # use_qat_qkv_backward
|
||||
False, # smooth_k
|
||||
True, # warp_specialize (disabled internally on Blackwell)
|
||||
True, # IS_QAT
|
||||
False, # two_level_quant_P
|
||||
True, # fake_quant_P
|
||||
True, # use_high_prec_o
|
||||
False, # smooth_q
|
||||
False, # use_global_sf_P
|
||||
False, # use_global_sf_QKV
|
||||
)
|
||||
|
||||
def run(optimized: bool):
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1" if optimized else "0")
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast")
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
|
||||
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
|
||||
output = kernel.attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
False,
|
||||
1.0 / math.sqrt(q_shape[-1]),
|
||||
*flags,
|
||||
)
|
||||
output.backward(grad_out)
|
||||
return output.detach(), q.grad, k.grad, v.grad
|
||||
|
||||
legacy = run(False)
|
||||
optimized = run(True)
|
||||
|
||||
assert (optimized[0].float() - legacy[0].float()).abs().max().item() <= 1e-2
|
||||
assert (optimized[1].float() - legacy[1].float()).abs().max().item() <= 4e-3
|
||||
assert (optimized[2].float() - legacy[2].float()).abs().max().item() <= 4e-3
|
||||
assert torch.equal(optimized[3], legacy[3])
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 12,
|
||||
reason="SM120 parity test",
|
||||
)
|
||||
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
|
||||
def test_sm120_joined_pv_forward_backward_matches_split_path(monkeypatch, q_length, kv_length):
|
||||
torch.manual_seed(11)
|
||||
q_shape = (1, 1, q_length, 128)
|
||||
kv_shape = (1, 1, kv_length, 128)
|
||||
inputs = [
|
||||
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
|
||||
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
|
||||
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
|
||||
]
|
||||
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
|
||||
flags = (
|
||||
True, # use_qat_qkv_backward
|
||||
False, # smooth_k
|
||||
True, # warp_specialize (disabled internally on Blackwell)
|
||||
True, # IS_QAT
|
||||
False, # two_level_quant_P
|
||||
True, # fake_quant_P
|
||||
True, # use_high_prec_o
|
||||
False, # smooth_q
|
||||
False, # use_global_sf_P
|
||||
False, # use_global_sf_QKV
|
||||
)
|
||||
|
||||
def run(joined_pv: bool):
|
||||
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1" if joined_pv else "0")
|
||||
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
|
||||
output = kernel.attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
False,
|
||||
1.0 / math.sqrt(q_shape[-1]),
|
||||
*flags,
|
||||
)
|
||||
output.backward(grad_out)
|
||||
return output.detach(), q.grad, k.grad, v.grad
|
||||
|
||||
split = run(False)
|
||||
joined = run(True)
|
||||
|
||||
assert torch.equal(joined[0], split[0])
|
||||
assert torch.equal(joined[1], split[1])
|
||||
assert torch.equal(joined[2], split[2])
|
||||
assert torch.equal(joined[3], split[3])
|
||||
@@ -304,6 +304,8 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
for key in _LTX2_REFINE_FLAT_KEYS:
|
||||
if key in refine:
|
||||
kwargs[f"ltx2_refine_{key}"] = refine[key]
|
||||
if "enabled" in refine:
|
||||
kwargs["refine_enabled"] = refine["enabled"]
|
||||
kwargs.update(preset_overrides)
|
||||
kwargs.update(deepcopy(normalized.pipeline.experimental))
|
||||
return FastVideoArgs.from_kwargs(**kwargs)
|
||||
|
||||
@@ -51,8 +51,9 @@ class SamplingParam:
|
||||
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
|
||||
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
# Camera control inputs (LingBotWorld and LingBotWorld2)
|
||||
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
action_path: str | None = None # Directory containing poses.npy and intrinsics.npy
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
@@ -89,7 +90,13 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
batch_cfg: bool = False
|
||||
guidance_scale_2: float | None = None
|
||||
# Z-Image CFG controls. ``cfg_normalization=True`` caps the guided
|
||||
# prediction norm at the positive-prediction norm; ``cfg_truncation``
|
||||
# disables CFG above the normalized-noise threshold.
|
||||
cfg_normalization: bool = False
|
||||
cfg_truncation: float | None = 1.0
|
||||
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
|
||||
use_embedded_guidance: bool = False
|
||||
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
|
||||
@@ -323,6 +330,24 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_scale,
|
||||
help="Classifier-free guidance scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cfg-normalization",
|
||||
action=StoreBoolean,
|
||||
default=SamplingParam.cfg_normalization,
|
||||
help="Cap Z-Image CFG prediction norm to the positive-prediction norm",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cfg-truncation",
|
||||
type=float,
|
||||
default=SamplingParam.cfg_truncation,
|
||||
help="Disable Z-Image CFG above this normalized-noise threshold",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch-cfg",
|
||||
action=StoreBoolean,
|
||||
default=SamplingParam.batch_cfg,
|
||||
help="Evaluate conditional and unconditional CFG branches in one batch",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
|
||||
@@ -130,6 +130,7 @@ class InputConfig:
|
||||
keyboard_cond: Any | None = None
|
||||
grid_sizes: Any | None = None
|
||||
c2ws_plucker_emb: Any | None = None
|
||||
action_path: str | None = None
|
||||
refine_from: str | None = None
|
||||
stage1_video: Any | None = None
|
||||
|
||||
@@ -138,6 +139,7 @@ class InputConfig:
|
||||
class SamplingConfig:
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
max_sequence_length: int | None = None
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
@@ -147,7 +149,10 @@ class SamplingConfig:
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
batch_cfg: bool = False
|
||||
guidance_scale_2: float | None = None
|
||||
cfg_normalization: bool = False
|
||||
cfg_truncation: float | None = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
use_embedded_guidance: bool | None = None
|
||||
|
||||
@@ -18,14 +18,14 @@ from fastvideo.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_project_root = Path(__file__).resolve().parent.parent.parent.parent
|
||||
_kernel_root = _project_root / "fastvideo-kernel"
|
||||
_kernel_python_root = _kernel_root / "python"
|
||||
_kernel_python_root = _project_root / "fastvideo-kernel" / "python"
|
||||
_attn_qat_train_attention: Callable[..., torch.Tensor] | None = None
|
||||
_attn_qat_train_import_attempted = False
|
||||
_attn_qat_train_import_error: ImportError | None = None
|
||||
|
||||
|
||||
def _ensure_kernel_paths() -> None:
|
||||
for path in (_project_root, _kernel_root, _kernel_python_root):
|
||||
for path in (_project_root, _kernel_python_root):
|
||||
path_str = str(path)
|
||||
if path_str not in sys.path:
|
||||
sys.path.insert(0, path_str)
|
||||
@@ -34,6 +34,7 @@ def _ensure_kernel_paths() -> None:
|
||||
def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
global _attn_qat_train_attention
|
||||
global _attn_qat_train_import_attempted
|
||||
global _attn_qat_train_import_error
|
||||
|
||||
if _attn_qat_train_import_attempted:
|
||||
return _attn_qat_train_attention
|
||||
@@ -42,8 +43,11 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
_ensure_kernel_paths()
|
||||
|
||||
try:
|
||||
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
|
||||
except ImportError:
|
||||
triton_qat = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train")
|
||||
_attn_qat_train_attention = triton_qat.attention
|
||||
logger.info("ATTN_QAT_TRAIN loaded FastVideo's architecture-optimized Triton kernel")
|
||||
except ImportError as exc:
|
||||
_attn_qat_train_import_error = exc
|
||||
_attn_qat_train_attention = None
|
||||
|
||||
return _attn_qat_train_attention
|
||||
@@ -60,8 +64,9 @@ def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
sm_scale: float | None = None) -> torch.Tensor:
|
||||
attention = _get_attn_qat_train_attention()
|
||||
if attention is None:
|
||||
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
|
||||
"Please ensure the FastVideo kernel package is installed.")
|
||||
detail = f" Original import error: {_attn_qat_train_import_error}" if _attn_qat_train_import_error else ""
|
||||
raise ImportError("ATTN_QAT_TRAIN requires FastVideo's fastvideo-kernel package. Install it or make "
|
||||
f"fastvideo-kernel/python importable.{detail}")
|
||||
|
||||
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
@@ -69,7 +74,11 @@ def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
|
||||
use_qat_qkv_backward = True
|
||||
smooth_k = False
|
||||
warp_specialize = True
|
||||
# Triton 3.7's NVWS pass aborts while compiling this kernel on Blackwell.
|
||||
# The kernel has a supported non-warp-specialized path, so use it on both
|
||||
# datacenter (sm_100) and consumer (sm_120) Blackwell GPUs.
|
||||
capability_major = torch.cuda.get_device_capability()[0]
|
||||
warp_specialize = capability_major not in (10, 12)
|
||||
is_qat = True
|
||||
two_level_quant_p_sage3 = False
|
||||
fake_quant_p_bwd = True
|
||||
@@ -106,7 +115,7 @@ class AttnQatTrainBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 96, 128, 160, 192, 224, 256]
|
||||
return [128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
|
||||
@@ -32,6 +32,27 @@ def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
|
||||
None
|
||||
|
||||
|
||||
def coerce_attn_backend(attn_backend: AttentionBackendEnum | str | None, ) -> AttentionBackendEnum | None:
|
||||
"""Normalize an explicit backend selection.
|
||||
|
||||
Environment-variable parsing remains permissive via
|
||||
:func:`backend_name_to_enum`, but typed/config-driven call sites should
|
||||
fail fast on typos instead of silently falling back to another backend.
|
||||
"""
|
||||
if attn_backend is None or isinstance(attn_backend, AttentionBackendEnum):
|
||||
return attn_backend
|
||||
if not isinstance(attn_backend, str) or not attn_backend.strip():
|
||||
raise ValueError("attention backend must be a non-empty string, "
|
||||
f"an AttentionBackendEnum, or None; got {attn_backend!r}")
|
||||
|
||||
backend_name = attn_backend.strip().upper()
|
||||
backend = backend_name_to_enum(backend_name)
|
||||
if backend is None:
|
||||
raise ValueError(f"Unknown attention backend {attn_backend!r}. "
|
||||
f"Expected one of {sorted(AttentionBackendEnum.__members__)}")
|
||||
return backend
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
@@ -69,6 +90,11 @@ def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None
|
||||
'''
|
||||
global forced_attn_backend
|
||||
forced_attn_backend = attn_backend
|
||||
# Backend selection is cached by tensor shape/dtype, while the global
|
||||
# override is intentionally not part of that cache key. Invalidate cached
|
||||
# resolutions whenever the override changes so independently constructed
|
||||
# role models can bind different attention implementations.
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
|
||||
|
||||
@@ -12,12 +12,16 @@ from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
|
||||
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
|
||||
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
|
||||
"ZImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 VFM Transformer FastVideo dataclass configs.
|
||||
|
||||
Architecture is 1:1 with the published ``nvidia/Cosmos3-Nano`` checkpoint
|
||||
(``transformer/config.json``; class ``Cosmos3OmniTransformer`` / framework
|
||||
``Cosmos3VFMNetwork``). Field values match that config so the FastVideo native
|
||||
DiT builds a parameter tree matching the checkpoint's state-dict surface
|
||||
(814 tensors / 44 patterns, validated 2026-06-06).
|
||||
|
||||
Reference of record: ``cosmos-framework`` (NVIDIA). The checkpoint is a single
|
||||
``layers`` ModuleList of dual-pathway (understanding/text + generation/vision)
|
||||
decoder blocks; per layer: ``self_attn`` with und (``to_{q,k,v}``/``to_out``)
|
||||
and gen (``add_{q,k,v}_proj``/``to_add_out``) projections + QK-norms, plus
|
||||
``mlp`` (und) and ``mlp_moe_gen`` (gen), and four RMSNorms. Top level adds
|
||||
``embed_tokens``/``norm``/``norm_moe_gen``/``lm_head``/``proj_in``/``proj_out``/
|
||||
``time_embedder`` and dormant ``action_*``/``audio_*`` heads. The checkpoint
|
||||
remap lives in ``scripts/checkpoint_conversion/cosmos3_convert.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_cosmos3_transformer_block(name: str, module) -> bool:
|
||||
"""FSDP shard boundary: the dual-pathway decoder blocks ``layers.{i}``."""
|
||||
del module
|
||||
parts = name.split(".")
|
||||
return "layers" in parts and parts[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ArchConfig(DiTArchConfig):
|
||||
"""Architecture config for the Cosmos3 omni DiT (Cosmos3-Nano).
|
||||
|
||||
1:1 with ``transformer/config.json``. The action/sound heads ship in the
|
||||
checkpoint, so they are constructed for strict-load parity even though the
|
||||
PR1 video path (T2V/I2V/T2I) leaves them dormant.
|
||||
"""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_cosmos3_transformer_block])
|
||||
|
||||
# Conversion is owned by scripts/checkpoint_conversion/cosmos3_convert.py;
|
||||
# the native module tree is the source of truth for parameter names.
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
# ---- Backbone (Qwen3-VL-text) ----
|
||||
hidden_size: int = 4096
|
||||
num_hidden_layers: int = 36
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 8 # GQA (4 query groups)
|
||||
head_dim: int = 128
|
||||
intermediate_size: int = 12288
|
||||
hidden_act: str = "silu"
|
||||
vocab_size: int = 151936
|
||||
rms_norm_eps: float = 1e-6
|
||||
attention_bias: bool = False
|
||||
qk_norm_for_diffusion: bool = True
|
||||
qk_norm_for_text: bool = True
|
||||
use_moe: bool = True # dual-pathway weights; sparse routing unused
|
||||
joint_attn_implementation: str = "two_way"
|
||||
freeze_und: bool = False
|
||||
|
||||
# ---- Position embedding (unified 3D MRoPE) ----
|
||||
position_embedding_type: str = "unified_3d_mrope"
|
||||
rope_theta: float = 5_000_000.0
|
||||
max_position_embeddings: int = 262144
|
||||
mrope_section: list[int] = field(default_factory=lambda: [24, 20, 20])
|
||||
mrope_interleaved: bool = True
|
||||
unified_3d_mrope_reset_spatial_ids: bool = True
|
||||
temporal_modality_margin: int = 15000 # unified_3d_mrope_temporal_modality_margin
|
||||
|
||||
# ---- VAE / patch geometry ----
|
||||
latent_patch_size: int = 2
|
||||
latent_channel: int = 48
|
||||
patch_latent_dim: int = 192 # latent_patch_size**2 * latent_channel
|
||||
|
||||
# ---- Diffusion conditioning ----
|
||||
timestep_scale: float = 0.001
|
||||
|
||||
# ---- Temporal / FPS modulation ----
|
||||
base_fps: float = 24.0
|
||||
temporal_compression_factor: int = 4
|
||||
enable_fps_modulation: bool = True
|
||||
video_temporal_causal: bool = False
|
||||
|
||||
# ---- Action generation head (dormant in PR1 video path) ----
|
||||
action_gen: bool = True
|
||||
action_dim: int = 64
|
||||
max_action_dim: int = 64
|
||||
num_embodiment_domains: int = 32
|
||||
|
||||
# ---- Sound generation head (dormant in PR1 video path) ----
|
||||
sound_gen: bool = True
|
||||
sound_dim: int = 64
|
||||
sound_latent_fps: float = 25.0
|
||||
temporal_compression_factor_sound: int = 1
|
||||
|
||||
# ---- BaseDiT bookkeeping ----
|
||||
in_channels: int = 48
|
||||
out_channels: int = 48
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# Video DiT contract: latent channels == VAE z_dim.
|
||||
self.num_channels_latents = self.latent_channel
|
||||
if not self.out_channels:
|
||||
self.out_channels = self.in_channels
|
||||
# Derived: patchify packs latent_patch_size**2 spatial patches * channels.
|
||||
self.patch_latent_dim = self.latent_patch_size**2 * self.latent_channel
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VideoConfig(DiTConfig):
|
||||
"""Pipeline-level Cosmos3 DiT config (T2V / I2V / T2I share this surface)."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=Cosmos3ArchConfig)
|
||||
prefix: str = "Cosmos3"
|
||||
@@ -0,0 +1,70 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Architecture configuration for LingBot-Video Dense and MoE DiTs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_lingbot_video_block(name: str, module: object) -> bool:
|
||||
"""Select top-level transformer blocks for FSDP and compilation."""
|
||||
del module
|
||||
parts = name.split(".")
|
||||
return len(parts) == 2 and parts[0] == "blocks" and parts[1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoArchConfig(DiTArchConfig):
|
||||
"""One-to-one representation of the released transformer config JSON."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_lingbot_video_block])
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
)
|
||||
param_names_mapping: dict = field(default_factory=lambda: {r"^(.*)$": r"\1"})
|
||||
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
hidden_size: int = 2048
|
||||
num_attention_heads: int = 16
|
||||
depth: int = 24
|
||||
intermediate_size: int = 6144
|
||||
text_dim: int = 2560
|
||||
freq_dim: int = 256
|
||||
norm_eps: float = 1e-6
|
||||
rope_theta: float = 256.0
|
||||
axes_dims: tuple[int, int, int] = (32, 48, 48)
|
||||
axes_lens: tuple[int, int, int] = (8192, 1024, 1024)
|
||||
qkv_bias: bool = False
|
||||
out_bias: bool = True
|
||||
patch_embed_bias: bool = True
|
||||
timestep_mlp_bias: bool = True
|
||||
num_experts: int = 0
|
||||
num_experts_per_tok: int = 8
|
||||
moe_intermediate_size: int = 512
|
||||
decoder_sparse_step: int = 1
|
||||
mlp_only_layers: tuple[int, ...] = ()
|
||||
n_shared_experts: int | None = None
|
||||
score_func: str = "sigmoid"
|
||||
norm_topk_prob: bool = True
|
||||
n_group: int | None = None
|
||||
topk_group: int | None = None
|
||||
routed_scaling_factor: float = 1.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Populate FastVideo loader fields from the released architecture."""
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.in_channels
|
||||
self.attention_head_dim = self.hidden_size // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoConfig(DiTConfig):
|
||||
"""FastVideo component configuration for LingBot-Video transformers."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LingBotVideoArchConfig)
|
||||
@@ -0,0 +1,55 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2CausalFastArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
model_type: str = "i2v"
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
text_len: int = 512
|
||||
in_dim: int = 36
|
||||
dim: int = 5120
|
||||
ffn_dim: int = 13824
|
||||
freq_dim: int = 256
|
||||
text_dim: int = 4096
|
||||
out_dim: int = 16
|
||||
num_heads: int = 40
|
||||
num_layers: int = 40
|
||||
qk_norm: bool = True
|
||||
cross_attn_norm: bool = True
|
||||
eps: float = 1e-6
|
||||
|
||||
local_attn_size: int = 18
|
||||
sink_size: int = 6
|
||||
chunk_size: int = 4
|
||||
sample_shift: float = 10.0
|
||||
num_train_timesteps: int = 1000
|
||||
timesteps_index: tuple[int, int, int, int] = (0, 250, 500, 750)
|
||||
max_area: int = 480 * 832
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.hidden_size = self.dim
|
||||
self.num_attention_heads = self.num_heads
|
||||
self.attention_head_dim = self.dim // self.num_heads
|
||||
self.in_channels = self.in_dim
|
||||
self.out_channels = self.out_dim
|
||||
self.num_channels_latents = self.out_dim
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2CausalFastVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=LingBotWorld2CausalFastArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def is_zimage_block(name: str, module) -> bool:
|
||||
parts = name.split(".")
|
||||
return len(parts) >= 2 and parts[-2] in {"noise_refiner", "context_refiner", "layers"} and parts[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class ZImageDiTArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_zimage_block])
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.TORCH_SDPA, )
|
||||
|
||||
all_patch_size: tuple[int, ...] = (2, )
|
||||
all_f_patch_size: tuple[int, ...] = (1, )
|
||||
in_channels: int = 16
|
||||
dim: int = 3840
|
||||
n_layers: int = 30
|
||||
n_refiner_layers: int = 2
|
||||
n_heads: int = 30
|
||||
n_kv_heads: int = 30
|
||||
norm_eps: float = 1e-5
|
||||
qk_norm: bool = True
|
||||
cap_feat_dim: int = 2560
|
||||
rope_theta: float = 256.0
|
||||
t_scale: float = 1000.0
|
||||
axes_dims: tuple[int, ...] = (32, 48, 48)
|
||||
axes_lens: tuple[int, ...] = (1536, 512, 512)
|
||||
|
||||
adaln_embed_dim: int = 256
|
||||
frequency_embedding_size: int = 256
|
||||
timestep_mid_size: int = 1024
|
||||
max_period: int = 10000
|
||||
seq_multi_of: int = 32
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
if len(self.all_patch_size) != len(self.all_f_patch_size):
|
||||
raise ValueError("all_patch_size and all_f_patch_size must have equal length")
|
||||
if self.dim % self.n_heads:
|
||||
raise ValueError("dim must be divisible by n_heads")
|
||||
if self.dim // self.n_heads != sum(self.axes_dims):
|
||||
raise ValueError("attention head dimension must equal sum(axes_dims)")
|
||||
if len(self.axes_dims) != len(self.axes_lens) or any(dim % 2 for dim in self.axes_dims):
|
||||
raise ValueError("RoPE axes require matching lengths and even dimensions")
|
||||
|
||||
self.hidden_size = self.dim
|
||||
self.num_attention_heads = self.n_heads
|
||||
self.num_channels_latents = self.in_channels
|
||||
self.out_channels = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class ZImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=ZImageDiTArchConfig)
|
||||
prefix: str = "ZImage"
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput, EncoderCo
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5ArchConfig, LingBotWorld2UMT5Config
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
|
||||
@@ -9,6 +10,7 @@ from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
@@ -17,5 +19,6 @@ __all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
|
||||
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig"
|
||||
]
|
||||
|
||||
@@ -83,6 +83,7 @@ class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
is_chat_model: bool = False
|
||||
treat_empty_as_dot: bool = False
|
||||
chat_template_enable_thinking: bool = field(default=False, kw_only=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Qwen3-VL text-only encoder configuration used by LingBot-Video."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextArchConfig, Qwen3TextConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoQwen3VLTextArchConfig(Qwen3TextArchConfig):
|
||||
"""Exact Qwen3-VL language-model architecture released with LingBot-Video."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["LingBotVideoQwen3VLTextModel"])
|
||||
vocab_size: int = 151936
|
||||
hidden_size: int = 2560
|
||||
intermediate_size: int = 9728
|
||||
num_hidden_layers: int = 36
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 8
|
||||
max_position_embeddings: int = 262144
|
||||
rms_norm_eps: float = 1e-6
|
||||
rope_theta: float = 5000000.0
|
||||
rope_scaling: dict | None = None
|
||||
mrope_interleaved: bool = True
|
||||
mrope_section: tuple[int, int, int] = (24, 20, 20)
|
||||
bos_token_id: int = 151643
|
||||
eos_token_id: int = 151645
|
||||
pad_token_id: int = 151643
|
||||
text_len: int = 37698
|
||||
output_hidden_states: bool = True
|
||||
require_processor: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Match the official processor call used by LingBotVideoPipeline."""
|
||||
self.tokenizer_kwargs = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"padding": "longest",
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoQwen3VLTextConfig(Qwen3TextConfig):
|
||||
"""FastVideo loader config for the LingBot-Video text-only Qwen3-VL path."""
|
||||
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LingBotVideoQwen3VLTextArchConfig)
|
||||
prefix: str = "language_model"
|
||||
is_chat_model: bool = False
|
||||
@@ -0,0 +1,40 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2UMT5ArchConfig(TextEncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["LingBotWorld2T5EncoderModel"])
|
||||
vocab_size: int = 256384
|
||||
dim: int = 4096
|
||||
dim_attn: int = 4096
|
||||
dim_ffn: int = 10240
|
||||
num_heads: int = 64
|
||||
num_layers: int = 24
|
||||
num_buckets: int = 32
|
||||
text_len: int = 512
|
||||
hidden_size: int = 4096
|
||||
dropout: float = 0.1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"add_special_tokens": True,
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2UMT5Config(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LingBotWorld2UMT5ArchConfig)
|
||||
|
||||
prefix: str = "text_encoder"
|
||||
@@ -1,5 +1,6 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos3vae import Cosmos3VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
@@ -16,6 +17,7 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Cosmos3VAEConfig",
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
"""Cosmos3 (Wan2.2-TI2V-5B) VAE config and checkpoint-key mapping.
|
||||
|
||||
The Cosmos3 checkpoint VAE is literally ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
|
||||
(diffusers ``AutoencoderKLWan``), so this config locks the Wan2.2 geometry:
|
||||
residual down/up blocks, ``patch_size=2``, ``z_dim=48``, ``base_dim=160``,
|
||||
``decoder_base_dim=256``, and ``scale_factor_spatial=16``. The 48-dim
|
||||
``latents_mean``/``latents_std`` are taken verbatim from the Cosmos3
|
||||
checkpoint's ``vae/config.json`` (identical to the canonical Wan2.2-TI2V-5B
|
||||
statistics).
|
||||
|
||||
Mirrors the :class:`Cosmos25VAEArchConfig` pattern. ``param_names_mapping`` /
|
||||
``map_official_key`` translate the *official* Wan2.2 VAE state-dict keys
|
||||
(nested-residual naming, e.g. ``encoder.downsamples.{b}.downsamples.{j}`` and
|
||||
``decoder.upsamples.{b}.upsamples.{j}``) into FastVideo's ``AutoencoderKLWan``
|
||||
key space. The standard diffusers checkpoint already ships native FastVideo
|
||||
keys, so these helpers exist for parity tooling and official ``.pth`` loading.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig, WanVAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VAEArchConfig(WanVAEArchConfig):
|
||||
# Wan2.2-TI2V-5B geometry (differs from the Wan2.1 WanVAEArchConfig
|
||||
# defaults: residual blocks, patch_size=2, z_dim=48, base_dim=160,
|
||||
# decoder_base_dim=256, scale_factor_spatial=16, 12 patch channels).
|
||||
_name_or_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
base_dim: int = 160
|
||||
decoder_base_dim: int | None = 256
|
||||
z_dim: int = 48
|
||||
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple[float, ...] = ()
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
is_residual: bool = True
|
||||
in_channels: int = 12
|
||||
out_channels: int = 12
|
||||
patch_size: int | None = 2
|
||||
scale_factor_temporal: int = 4
|
||||
scale_factor_spatial: int = 16
|
||||
clip_output: bool = False
|
||||
|
||||
# 48-dim statistics copied verbatim from the Cosmos3 checkpoint
|
||||
# (official_weights/cosmos3/vae/config.json).
|
||||
latents_mean: tuple[float, ...] = (
|
||||
-0.2289,
|
||||
-0.0052,
|
||||
-0.1323,
|
||||
-0.2339,
|
||||
-0.2799,
|
||||
0.0174,
|
||||
0.1838,
|
||||
0.1557,
|
||||
-0.1382,
|
||||
0.0542,
|
||||
0.2813,
|
||||
0.0891,
|
||||
0.157,
|
||||
-0.0098,
|
||||
0.0375,
|
||||
-0.1825,
|
||||
-0.2246,
|
||||
-0.1207,
|
||||
-0.0698,
|
||||
0.5109,
|
||||
0.2665,
|
||||
-0.2108,
|
||||
-0.2158,
|
||||
0.2502,
|
||||
-0.2055,
|
||||
-0.0322,
|
||||
0.1109,
|
||||
0.1567,
|
||||
-0.0729,
|
||||
0.0899,
|
||||
-0.2799,
|
||||
-0.123,
|
||||
-0.0313,
|
||||
-0.1649,
|
||||
0.0117,
|
||||
0.0723,
|
||||
-0.2839,
|
||||
-0.2083,
|
||||
-0.052,
|
||||
0.3748,
|
||||
0.0152,
|
||||
0.1957,
|
||||
0.1433,
|
||||
-0.2944,
|
||||
0.3573,
|
||||
-0.0548,
|
||||
-0.1681,
|
||||
-0.0667,
|
||||
)
|
||||
latents_std: tuple[float, ...] = (
|
||||
0.4765,
|
||||
1.0364,
|
||||
0.4514,
|
||||
1.1677,
|
||||
0.5313,
|
||||
0.499,
|
||||
0.4818,
|
||||
0.5013,
|
||||
0.8158,
|
||||
1.0344,
|
||||
0.5894,
|
||||
1.0901,
|
||||
0.6885,
|
||||
0.6165,
|
||||
0.8454,
|
||||
0.4978,
|
||||
0.5759,
|
||||
0.3523,
|
||||
0.7135,
|
||||
0.6804,
|
||||
0.5833,
|
||||
1.4146,
|
||||
0.8986,
|
||||
0.5659,
|
||||
0.7069,
|
||||
0.5338,
|
||||
0.4889,
|
||||
0.4917,
|
||||
0.4069,
|
||||
0.4999,
|
||||
0.6866,
|
||||
0.4093,
|
||||
0.5709,
|
||||
0.6065,
|
||||
0.6415,
|
||||
0.4944,
|
||||
0.5726,
|
||||
1.2042,
|
||||
0.5458,
|
||||
1.6887,
|
||||
0.3971,
|
||||
1.06,
|
||||
0.3943,
|
||||
0.5537,
|
||||
0.5444,
|
||||
0.4089,
|
||||
0.7468,
|
||||
0.7744,
|
||||
)
|
||||
|
||||
# Simple 1:1 renames. The nested-residual block remapping (encoder
|
||||
# downsamples / decoder upsamples / middle / head) is handled by
|
||||
# ``map_official_key()``.
|
||||
param_names_mapping: dict[str, str] = field(
|
||||
default_factory=lambda: {
|
||||
r"^conv1\.(.*)$": r"quant_conv.\1",
|
||||
r"^conv2\.(.*)$": r"post_quant_conv.\1",
|
||||
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
|
||||
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
|
||||
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
|
||||
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
|
||||
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
|
||||
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def map_official_key(key: str) -> str | None:
|
||||
"""Map a single official Wan2.2 VAE key into FastVideo key space.
|
||||
|
||||
Handles the residual (Wan2.2) module layout where each down/up block
|
||||
is a nested ``Sequential`` (``downsamples.{b}.downsamples.{j}`` /
|
||||
``upsamples.{b}.upsamples.{j}``) rather than the flat Wan2.1 indexing.
|
||||
Returns ``None`` for keys with no FastVideo counterpart.
|
||||
"""
|
||||
|
||||
def map_residual_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^residual\.0\.gamma$", sub):
|
||||
return f"{prefix}.norm1.gamma"
|
||||
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv1.{m.group(1)}"
|
||||
if re.match(r"^residual\.3\.gamma$", sub):
|
||||
return f"{prefix}.norm2.gamma"
|
||||
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv2.{m.group(1)}"
|
||||
m = re.match(r"^shortcut\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv_shortcut.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_attn_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^norm\.gamma$", sub):
|
||||
return f"{prefix}.norm.gamma"
|
||||
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.to_qkv.{m.group(1)}"
|
||||
m = re.match(r"^proj\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.proj.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_resample_subkey(prefix: str, sub: str) -> str | None:
|
||||
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.resample.1.{m.group(1)}"
|
||||
m = re.match(r"^time_conv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.time_conv.{m.group(1)}"
|
||||
return None
|
||||
|
||||
m = re.match(r"^conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^conv2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"post_quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_in.{m.group(2)}"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.norm_out.gamma"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_out.{m.group(2)}"
|
||||
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
|
||||
if m:
|
||||
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0", m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
|
||||
|
||||
# Encoder: downsamples.{block}.downsamples.{j}.* (nested residual layout)
|
||||
m = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
|
||||
if sub.startswith("resample.") or sub.startswith("time_conv."):
|
||||
return map_resample_subkey(f"encoder.down_blocks.{block_i}.downsampler", sub)
|
||||
return map_residual_subkey(f"encoder.down_blocks.{block_i}.resnets.{res_i}", sub)
|
||||
|
||||
# Decoder: upsamples.{block}.upsamples.{j}.* (nested residual layout)
|
||||
m = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
|
||||
if sub.startswith("resample.") or sub.startswith("time_conv."):
|
||||
return map_resample_subkey(f"decoder.up_blocks.{block_i}.upsampler", sub)
|
||||
return map_residual_subkey(f"decoder.up_blocks.{block_i}.resnets.{res_i}", sub)
|
||||
|
||||
return None
|
||||
|
||||
# ``__post_init__`` (scaling_factor / shift_factor / compression ratios) is
|
||||
# inherited unchanged from ``WanVAEArchConfig``.
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VAEConfig(WanVAEConfig):
|
||||
"""Cosmos3 VAE config (reuses FastVideo's Wan2.2 ``AutoencoderKLWan``).
|
||||
|
||||
Subclasses :class:`WanVAEConfig` so the model reads the same runtime flags
|
||||
(``use_feature_cache``, ``use_light_vae``, tiling) and only swaps in the
|
||||
Cosmos3 = Wan2.2 ``arch_config``.
|
||||
"""
|
||||
|
||||
arch_config: Cosmos3VAEArchConfig = field(default_factory=Cosmos3VAEArchConfig)
|
||||
|
||||
use_feature_cache: bool = True
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
# ``__post_init__`` (blend_num_frames) is inherited from ``WanVAEConfig``.
|
||||
@@ -7,6 +7,8 @@ from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyua
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
|
||||
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -19,5 +21,6 @@ __all__ = [
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"Kandinsky5I2VConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig", "MatrixGame2I2V480PConfig",
|
||||
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 pipeline configuration.
|
||||
|
||||
Reference of record: the official ``cosmos-framework`` / ``nvidia/Cosmos3-Nano``
|
||||
checkpoint (``model_index.json``). Cosmos3 is structurally different from
|
||||
Cosmos 2.5:
|
||||
|
||||
- Dual-pathway (UND + GEN) DiT lives entirely inside ``Cosmos3VFMTransformer``
|
||||
(``Cosmos3VideoConfig``).
|
||||
- No separate text encoder — the Qwen3-VL-text backbone is inside the DiT, so
|
||||
``text_encoder_configs`` is the empty tuple. The Qwen2 tokenizer is loaded as
|
||||
the ``text_tokenizer`` checkpoint module by the component loader.
|
||||
- VAE is Wan2.2 ``AutoencoderKLWan`` (z_dim=48, scale_factor_spatial=16),
|
||||
configured by ``Cosmos3VAEConfig`` (the checkpoint's exact latents_mean/std).
|
||||
- Scheduler is FastVideo-native ``UniPCMultistepScheduler`` configured for
|
||||
pure flow matching (flow_prediction, use_flow_sigmas), equivalent to the
|
||||
framework's ``FlowUniPCMultistepScheduler``. The checkpoint's diffusers-style
|
||||
scheduler config (karras/sigma_min/max) is coerced to the flow setup in
|
||||
``Cosmos3OmniDiffusersPipeline.initialize_pipeline``.
|
||||
- T2I default ``flow_shift`` is 3.0 (set per-request by ``_set_flow_shift``);
|
||||
T2V/I2V use the engine-init default of 1.0 baked into this config.
|
||||
"""
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.cosmos3 import (Cosmos3ArchConfig, Cosmos3VideoConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.vaes import Cosmos3VAEConfig # Wan2.2 AutoencoderKLWan
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3Config(PipelineConfig):
|
||||
"""Configuration for the Cosmos3 video generation pipeline (T2V/I2V/T2I).
|
||||
|
||||
Wires the framework-parity-verified Cosmos3 components: the native
|
||||
``Cosmos3VideoConfig`` DiT, the Wan2.2 ``Cosmos3VAEConfig`` VAE, the Qwen2
|
||||
tokenizer (loaded as ``text_tokenizer``), and the UniPC scheduler.
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=lambda: Cosmos3VideoConfig(arch_config=Cosmos3ArchConfig()))
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=Cosmos3VAEConfig)
|
||||
|
||||
# No separate text encoder: the Qwen3-VL-text backbone lives inside the DiT
|
||||
# and the pipeline tokenizes in Cosmos3DenoisingStage, so all three
|
||||
# text-encoder lists are empty (the generic text-encode stage is not used).
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=tuple)
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=tuple)
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(default_factory=tuple)
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
|
||||
|
||||
embedded_cfg_scale: float = 0.0
|
||||
# T2V/I2V engine-init flow_shift (framework text2video/image2video default);
|
||||
# T2I overrides to 3.0 per request via Cosmos3DenoisingStage._set_flow_shift.
|
||||
flow_shift: float = 10.0
|
||||
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,73 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Dense LingBot-Video T2V pipeline configuration."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
PROMPT_CROP_START = 140
|
||||
PROMPT_TEMPLATE = ("<|im_start|>system\nGiven a user input that may include a text prompt alone, "
|
||||
"a text prompt with an image reference, or a text prompt with a video reference "
|
||||
'or a video reference alone, generate an "Enhanced prompt" that provides detailed '
|
||||
"visual descriptions suitable for video generation. Evaluate the level of detail "
|
||||
"in the user's input: if it is simple, enrich it by adding specifics about colors, "
|
||||
"shapes, sizes, textures, lighting, motion dynamics, camera movement, temporal "
|
||||
"progression, and spatial relationships to create vivid, concrete, and temporally "
|
||||
"coherent scenes to create vivid and concrete scenes. Please generate only the "
|
||||
"enhanced description for the prompt below and avoid including any additional "
|
||||
"commentary or evaluations:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n"
|
||||
"<|im_start|>assistant\n")
|
||||
|
||||
|
||||
def preprocess_lingbot_video_prompt(prompt: str) -> str:
|
||||
"""Apply the released T2V system/user/assistant prompt template."""
|
||||
return PROMPT_TEMPLATE.format(prompt)
|
||||
|
||||
|
||||
def postprocess_lingbot_video_text(
|
||||
outputs: BaseEncoderOutput,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Select the final hidden state, crop the template, and trim batch-one padding."""
|
||||
if outputs.hidden_states is None:
|
||||
raise ValueError("LingBot-Video requires text-encoder hidden states")
|
||||
prompt_embeds = outputs.hidden_states[-1][:, PROMPT_CROP_START:]
|
||||
prompt_mask = attention_mask[:, PROMPT_CROP_START:]
|
||||
if prompt_embeds.shape[0] == 1:
|
||||
true_length = int(prompt_mask[0].sum().item())
|
||||
prompt_embeds = prompt_embeds[:, :true_length]
|
||||
prompt_mask = prompt_mask[:, :true_length]
|
||||
return prompt_embeds, prompt_mask
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoT2VConfig(PipelineConfig):
|
||||
"""Released Dense T2V component wiring and numerical precision policy."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=LingBotVideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (LingBotVideoQwen3VLTextConfig(), ))
|
||||
preprocess_text_funcs: tuple[Callable, ...] = field(default_factory=lambda: (preprocess_lingbot_video_prompt, ))
|
||||
postprocess_text_funcs: tuple[Callable, ...] = field(default_factory=lambda: (postprocess_lingbot_video_text, ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
vae_decode_precision: str | None = "fp32"
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
flow_shift: float | None = 3.0
|
||||
embedded_cfg_scale: float | None = None
|
||||
scheduler_step_in_fp32: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Load only the VAE decoder for the T2V workload."""
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import html
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import ftfy
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, LingBotWorld2UMT5Config
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.wan import Wan2_2_I2V_A14B_Config
|
||||
|
||||
|
||||
def lingbotworld2_whitespace_preprocess(prompt: str) -> str:
|
||||
text = ftfy.fix_text(prompt)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return " ".join(text.strip().split())
|
||||
|
||||
|
||||
def lingbotworld2_t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
assert outputs.last_hidden_state is not None
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld2CausalFastI2V480PConfig(Wan2_2_I2V_A14B_Config):
|
||||
dit_config: DiTConfig = field(default_factory=LingBotWorld2CausalFastVideoConfig)
|
||||
vae_config: WanVAEConfig = field(default_factory=WanVAEConfig)
|
||||
text_encoder_configs: tuple = field(default_factory=lambda: (LingBotWorld2UMT5Config(), ))
|
||||
preprocess_text_funcs: tuple = field(default_factory=lambda: (lingbotworld2_whitespace_preprocess, ))
|
||||
postprocess_text_funcs: tuple = field(default_factory=lambda: (lingbotworld2_t5_postprocess_text, ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
flow_shift: float | None = 10.0
|
||||
boundary_ratio: float | None = 0.947
|
||||
vae_precision: str = "fp32"
|
||||
vae_decode_precision: str | None = "fp32"
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
dit_precision: str = "bf16"
|
||||
is_causal: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,52 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def _zimage_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.hidden_states is None:
|
||||
raise RuntimeError("Z-Image requires Qwen3 hidden states")
|
||||
return outputs.hidden_states[-2]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ZImagePipelineConfig(PipelineConfig):
|
||||
"""Configuration for the native Z-Image text-to-image pipeline."""
|
||||
|
||||
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
|
||||
transformer_arch: str = "ZImageTransformer2DModel"
|
||||
vae_arch: str = "AutoencoderKL"
|
||||
text_encoder_archs: tuple[str, ...] = ("Qwen3Model", )
|
||||
tokenizer_archs: tuple[str, ...] = ("Qwen2Tokenizer", )
|
||||
|
||||
dit_config: ZImageDiTConfig = field(default_factory=ZImageDiTConfig)
|
||||
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (Qwen3TextConfig(chat_template_enable_thinking=True), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (_zimage_text_postprocess, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
embedded_cfg_scale: float = 0.0
|
||||
flow_shift: float | None = 3.0
|
||||
scheduler_step_in_fp32: bool = True
|
||||
scheduler_sigma_min: float = 0.0
|
||||
scheduler_use_reference_discrete_timesteps: bool = True
|
||||
@@ -112,6 +112,40 @@ def _infer_latent_batch_size(batch: ForwardBatch) -> int:
|
||||
return latent_batch_size
|
||||
|
||||
|
||||
def _resolve_output_size(
|
||||
samples: torch.Tensor,
|
||||
fallback: tuple[int, int, int],
|
||||
*,
|
||||
pixel_output: bool,
|
||||
) -> tuple[int, int, int]:
|
||||
"""Report the final decoded video's `(height, width, frames)`.
|
||||
|
||||
Refiner stages can produce a different resolution from the base request, so
|
||||
pixel outputs use the final `[batch, channels, frames, height, width]` tensor.
|
||||
Latent and audio outputs keep the requested fallback because their tensor
|
||||
dimensions do not describe decoded pixels.
|
||||
"""
|
||||
if pixel_output and samples.ndim == 5:
|
||||
return (int(samples.shape[-2]), int(samples.shape[-1]), int(samples.shape[-3]))
|
||||
return fallback
|
||||
|
||||
|
||||
def _validate_request_stage_overrides(model_path: str, request: GenerationRequest) -> None:
|
||||
"""Validate typed stage overrides against the model's registered preset."""
|
||||
if not request.stage_overrides:
|
||||
return
|
||||
from fastvideo.api.presets import validate_preset_selection
|
||||
from fastvideo.registry import get_preset_selection
|
||||
preset_name, model_family = get_preset_selection(model_path)
|
||||
if preset_name is None or model_family is None:
|
||||
raise ValueError(f"Model {model_path!r} has no preset for stage override validation")
|
||||
validate_preset_selection(
|
||||
preset_name,
|
||||
model_family,
|
||||
stage_overrides=request.stage_overrides,
|
||||
)
|
||||
|
||||
|
||||
class VideoGenerator:
|
||||
"""
|
||||
A unified class for generating videos using diffusion models.
|
||||
@@ -443,6 +477,7 @@ class VideoGenerator:
|
||||
self,
|
||||
request: GenerationRequest,
|
||||
) -> GenerationResult | list[GenerationResult]:
|
||||
_validate_request_stage_overrides(self.fastvideo_args.model_path, request)
|
||||
if isinstance(request.prompt, list):
|
||||
if request.inputs.prompt_path is not None:
|
||||
raise ValueError("request.prompt list cannot be combined with request.inputs.prompt_path")
|
||||
@@ -808,6 +843,15 @@ class VideoGenerator:
|
||||
# 2. Audio-only workload — `samples` is a 1×3×1×8×8 placeholder
|
||||
# no caller will use; skip the grid loop and save a `.wav`.
|
||||
# 3. Pixel video / image — the historical happy path.
|
||||
# `GenerationResult.size` describes the produced media, not only the
|
||||
# base-stage request. Refiner pipelines can change the final pixel
|
||||
# dimensions, so derive this result metadata from the decoded output.
|
||||
output_size = _resolve_output_size(
|
||||
samples,
|
||||
(target_height, target_width, batch.num_frames),
|
||||
pixel_output=not is_latent_output and not audio_only,
|
||||
)
|
||||
|
||||
postprocess_start = time.perf_counter()
|
||||
frames: list[np.ndarray] | None
|
||||
if is_latent_output or audio_only:
|
||||
@@ -916,7 +960,7 @@ class VideoGenerator:
|
||||
"audio": output_batch.extra.get("audio"),
|
||||
"audio_sample_rate": output_batch.extra.get("audio_sample_rate"),
|
||||
"ltx2_audio_latents": output_batch.extra.get("ltx2_audio_latents"),
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"size": output_size,
|
||||
"generation_time": gen_time,
|
||||
"e2e_latency": e2e_time,
|
||||
"logging_info": logging_info,
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 sound tokenizer (AVAE) — decode path.
|
||||
|
||||
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
|
||||
checkpoint is **decoder-only** (``decoder.*``) in ``AutoencoderOobleck`` naming
|
||||
with SnakeBeta activations and ``weight_g``/``weight_v`` weight-norm — exactly
|
||||
FastVideo's native :class:`~fastvideo.models.vaes.oobleck.OobleckDecoder`
|
||||
(verified bit-exact vs the framework in ``test_cosmos3_avae_parity``). Text-to-
|
||||
video+sound (t2vs) only needs DECODE: the DiT generates the sound latent and
|
||||
this module decodes it to a waveform, so only the decoder is ported (the
|
||||
SpectrogramConvNeXt encoder is not exported in the checkpoint).
|
||||
|
||||
Mirrors the framework ``AVAEModel.decode``: run the Oobleck decoder, then clamp
|
||||
to [-1, 1]. The VAE bottleneck's decode is the identity (the DiT already emits
|
||||
the post-bottleneck latent), so there is no bottleneck step here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.oobleck import OobleckDecoder
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundVAEArchConfig:
|
||||
"""Cosmos3 AVAE decoder constants (from ``sound_tokenizer/config.json``)."""
|
||||
|
||||
dec_dim: int = 320 # decoder base channels
|
||||
vocoder_input_dim: int = 64 # latent channels in
|
||||
dec_c_mults: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
|
||||
dec_strides: list[int] = field(default_factory=lambda: [2, 4, 5, 6, 8])
|
||||
audio_channels: int = 2 # stereo
|
||||
sampling_rate: int = 48000
|
||||
|
||||
@property
|
||||
def hop_size(self) -> int:
|
||||
return int(np.prod(self.dec_strides)) # 1920
|
||||
|
||||
|
||||
class Cosmos3SoundVAE(nn.Module):
|
||||
"""Decoder-only Cosmos3 AVAE: latent ``[B, z, T]`` -> waveform ``[B, C, N]``."""
|
||||
|
||||
def __init__(self, arch: Cosmos3SoundVAEArchConfig | None = None) -> None:
|
||||
super().__init__()
|
||||
self.arch = arch or Cosmos3SoundVAEArchConfig()
|
||||
self.decoder = OobleckDecoder(
|
||||
channels=self.arch.dec_dim,
|
||||
input_channels=self.arch.vocoder_input_dim,
|
||||
audio_channels=self.arch.audio_channels,
|
||||
# The framework builds decoder blocks from ``reversed(dec_strides)``
|
||||
# (deepest first), so block strides are e.g. [8,6,5,4,2].
|
||||
upsampling_ratios=list(reversed(self.arch.dec_strides)),
|
||||
channel_multiples=list(self.arch.dec_c_mults),
|
||||
)
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> int:
|
||||
return self.arch.sampling_rate
|
||||
|
||||
@property
|
||||
def audio_channels(self) -> int:
|
||||
return self.arch.audio_channels
|
||||
|
||||
@property
|
||||
def hop_size(self) -> int:
|
||||
return self.arch.hop_size
|
||||
|
||||
def get_latent_num_samples(self, num_audio_samples: int) -> int:
|
||||
"""Latent length for a given audio length (``AVAEInterface``: ``N // hop``)."""
|
||||
return int(num_audio_samples) // self.arch.hop_size
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode normalized latent ``[B, z, T]`` to waveform ``[B, C, N]`` in [-1, 1].
|
||||
|
||||
Matches ``AVAEModel.decode``: Oobleck decoder then clamp to [-1, 1] (the
|
||||
VAE bottleneck decode is identity).
|
||||
"""
|
||||
audio = self.decoder(latent) # [B, C, N]
|
||||
return audio.clamp(-1.0, 1.0)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
model_path: str,
|
||||
*,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
) -> "Cosmos3SoundVAE":
|
||||
"""Build + load the decoder from a ``sound_tokenizer`` directory.
|
||||
|
||||
Reads ``config.json`` (``dec_dim`` / ``vocoder_input_dim`` /
|
||||
``dec_c_mults`` / ``dec_strides`` / ``sampling_rate`` / ``stereo``) and
|
||||
loads the ``decoder.*`` weights (the checkpoint is decoder-only).
|
||||
"""
|
||||
from safetensors.torch import load_file
|
||||
|
||||
cfg_path = os.path.join(model_path, "config.json")
|
||||
with open(cfg_path) as f:
|
||||
cfg = json.load(f)
|
||||
arch = Cosmos3SoundVAEArchConfig(
|
||||
dec_dim=int(cfg["dec_dim"]),
|
||||
vocoder_input_dim=int(cfg["vocoder_input_dim"]),
|
||||
dec_c_mults=list(cfg["dec_c_mults"]),
|
||||
dec_strides=list(cfg["dec_strides"]),
|
||||
audio_channels=2 if cfg.get("stereo", True) else 1,
|
||||
sampling_rate=int(cfg.get("sampling_rate", 48000)),
|
||||
)
|
||||
model = cls(arch)
|
||||
|
||||
weights_path = os.path.join(model_path, "diffusion_pytorch_model.safetensors")
|
||||
state = load_file(weights_path)
|
||||
# Decoder-only checkpoint: strip the ``decoder.`` prefix.
|
||||
dec_state = {k[len("decoder."):]: v for k, v in state.items() if k.startswith("decoder.")}
|
||||
model.decoder.load_state_dict(dec_state, strict=True)
|
||||
logger.info("Loaded Cosmos3 sound AVAE decoder (%d params) from %s",
|
||||
sum(p.numel() for p in model.parameters()), model_path)
|
||||
|
||||
if torch_dtype is not None:
|
||||
model = model.to(dtype=torch_dtype)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
EntryClass = Cosmos3SoundVAE
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,811 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native LingBot-Video Dense and MoE diffusion transformers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig, _is_lingbot_video_block
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_all_to_all_4D,
|
||||
sequence_model_parallel_shard,
|
||||
)
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotVideoTransformerOutput:
|
||||
"""Output container matching the released transformer contract."""
|
||||
|
||||
sample: torch.Tensor
|
||||
|
||||
|
||||
_FP32_MODULE_NAMES = (
|
||||
"time_embedder",
|
||||
"time_modulation",
|
||||
"scale_shift_table",
|
||||
"norm",
|
||||
"norm1",
|
||||
"norm2",
|
||||
"norm_q",
|
||||
"norm_k",
|
||||
"norm_post_attn",
|
||||
"norm_post_ffn",
|
||||
"norm_out",
|
||||
"norm_out_modulation",
|
||||
"router",
|
||||
)
|
||||
|
||||
|
||||
def _keep_in_fp32(name: str) -> bool:
|
||||
"""Return whether a released checkpoint module keeps fp32 parameters."""
|
||||
return any(module_name in name.split(".") for module_name in _FP32_MODULE_NAMES)
|
||||
|
||||
|
||||
def _sequence_parallel_world_size() -> int:
|
||||
"""Use standalone single-rank behavior before distributed initialization."""
|
||||
return get_sp_world_size() if model_parallel_is_initialized() else 1
|
||||
|
||||
|
||||
class LingBotVideoLinear(ReplicatedLinear):
|
||||
"""Replicated FastVideo linear with the tensor-only official call contract."""
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Return the projected tensor while preserving the normal weight surface."""
|
||||
output, _ = super().forward(hidden_states)
|
||||
return output
|
||||
|
||||
|
||||
class LingBotVideoRMSNorm(nn.Module):
|
||||
"""RMSNorm with fp32 accumulation and input-dtype output."""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize the last dimension using the official accumulation order."""
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return (self.weight * hidden_states).to(input_dtype)
|
||||
|
||||
|
||||
def _apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply complex 3D rotary embeddings to `(B, S, H, D)` tensors."""
|
||||
with torch.amp.autocast("cuda", enabled=False):
|
||||
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
||||
output = torch.view_as_real(x_complex * freqs_cis.unsqueeze(2)).flatten(3)
|
||||
return output.type_as(x)
|
||||
|
||||
|
||||
class LingBotVideoRotaryEmbedding(nn.Module):
|
||||
"""Complex64 rotary table indexed by temporal and spatial positions."""
|
||||
|
||||
def __init__(self, axes_dims: tuple[int, ...], axes_lens: tuple[int, ...], theta: float) -> None:
|
||||
super().__init__()
|
||||
self.axes_dims = tuple(axes_dims)
|
||||
self.axes_lens = list(axes_lens)
|
||||
self.theta = theta
|
||||
self.freqs_cis: list[torch.Tensor] | None = None
|
||||
|
||||
@staticmethod
|
||||
def _precompute(dims: tuple[int, ...], lengths: tuple[int, ...], theta: float) -> list[torch.Tensor]:
|
||||
"""Build the per-axis complex frequency tables on CPU."""
|
||||
tables: list[torch.Tensor] = []
|
||||
for dim, length in zip(dims, lengths, strict=True):
|
||||
frequencies = 1.0 / (theta**(torch.arange(0, dim, 2, dtype=torch.float64, device="cpu") / dim))
|
||||
positions = torch.arange(length, device=frequencies.device, dtype=torch.float64)
|
||||
phases = torch.outer(positions, frequencies).float()
|
||||
tables.append(torch.polar(torch.ones_like(phases), phases).to(torch.complex64))
|
||||
return tables
|
||||
|
||||
def forward(
|
||||
self,
|
||||
position_ids: torch.Tensor,
|
||||
maxima: tuple[int, ...] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Gather and concatenate rotary frequencies for `(S, 3)` positions."""
|
||||
device = position_ids.device
|
||||
if maxima is None:
|
||||
maxima = tuple(int(value) for value in position_ids.max(dim=0).values.tolist())
|
||||
rebuild = self.freqs_cis is None or any(maximum >= length
|
||||
for maximum, length in zip(maxima, self.axes_lens, strict=True))
|
||||
if rebuild:
|
||||
for index, maximum in enumerate(maxima):
|
||||
if maximum >= self.axes_lens[index]:
|
||||
self.axes_lens[index] = int(maximum * 1.5) + 1
|
||||
self.freqs_cis = self._precompute(self.axes_dims, tuple(self.axes_lens), self.theta)
|
||||
self.freqs_cis = [table.to(device) for table in self.freqs_cis]
|
||||
elif self.freqs_cis[0].device != device:
|
||||
self.freqs_cis = [table.to(device) for table in self.freqs_cis]
|
||||
return torch.cat(
|
||||
[self.freqs_cis[index][position_ids[:, index]] for index in range(len(self.axes_dims))],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
|
||||
def _make_joint_position_ids(
|
||||
text_len: int,
|
||||
grid_t: int,
|
||||
grid_h: int,
|
||||
grid_w: int,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Create official `[video; text]` 3D positions for one sample."""
|
||||
temporal = torch.arange(grid_t, device=device, dtype=torch.int32) + text_len + 1
|
||||
height = torch.arange(grid_h, device=device, dtype=torch.int32)
|
||||
width = torch.arange(grid_w, device=device, dtype=torch.int32)
|
||||
video_positions = torch.stack(torch.meshgrid(temporal, height, width, indexing="ij"), dim=-1).flatten(0, 2)
|
||||
text_temporal = torch.arange(text_len, device=device, dtype=torch.int32) + 1
|
||||
text_positions = torch.stack(
|
||||
[text_temporal, torch.zeros_like(text_temporal),
|
||||
torch.zeros_like(text_temporal)], dim=-1)
|
||||
return torch.cat([video_positions, text_positions], dim=0)
|
||||
|
||||
|
||||
class LingBotVideoTextEmbedder(nn.Module):
|
||||
"""Project Qwen3-VL hidden states into the DiT hidden dimension."""
|
||||
|
||||
def __init__(self, text_dim: int, hidden_size: int) -> None:
|
||||
super().__init__()
|
||||
self.norm = LingBotVideoRMSNorm(text_dim, eps=1e-6)
|
||||
self.linear_1 = LingBotVideoLinear(text_dim, hidden_size, bias=True)
|
||||
self.linear_2 = LingBotVideoLinear(hidden_size, hidden_size, bias=True)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply RMSNorm followed by the released two-layer SiLU projection."""
|
||||
hidden_states = self.norm(hidden_states)
|
||||
return self.linear_2(F.silu(self.linear_1(hidden_states)))
|
||||
|
||||
|
||||
class LingBotVideoAttention(nn.Module):
|
||||
"""Joint video-text attention shared by the Dense and MoE variants."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
norm_eps: float,
|
||||
qkv_bias: bool,
|
||||
out_bias: bool,
|
||||
) -> None:
|
||||
"""Create released QKV projections, per-head norms, and output projection."""
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = hidden_size // num_heads
|
||||
self.to_q = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
|
||||
self.to_k = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
|
||||
self.to_v = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
|
||||
self.norm_q = LingBotVideoRMSNorm(self.head_dim, norm_eps)
|
||||
self.norm_k = LingBotVideoRMSNorm(self.head_dim, norm_eps)
|
||||
self.to_out = LingBotVideoLinear(hidden_size, hidden_size, bias=out_bias)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rotary_emb: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
original_seq_len: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Project QKV, apply rotary embeddings, and run non-causal SDPA."""
|
||||
batch, sequence, _ = hidden_states.shape
|
||||
query = self.to_q(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
|
||||
key = self.to_k(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
|
||||
value = self.to_v(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
|
||||
query = _apply_rotary_emb(self.norm_q(query), rotary_emb)
|
||||
key = _apply_rotary_emb(self.norm_k(key), rotary_emb)
|
||||
|
||||
if original_seq_len is not None:
|
||||
# Attention needs the full unpadded joint sequence while projections
|
||||
# and residual blocks stay sharded over tokens.
|
||||
qkv = torch.cat([query, key, value], dim=0)
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
|
||||
padded_seq_len = qkv.shape[1]
|
||||
query, key, value = qkv[:, :original_seq_len].chunk(3, dim=0)
|
||||
output = F.scaled_dot_product_attention(
|
||||
query.transpose(1, 2),
|
||||
key.transpose(1, 2),
|
||||
value.transpose(1, 2),
|
||||
attn_mask=attention_mask,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
).transpose(1, 2)
|
||||
if original_seq_len is not None:
|
||||
output = F.pad(output, (0, 0, 0, 0, 0, padded_seq_len - original_seq_len))
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
return self.to_out(output.reshape(batch, sequence, -1).type_as(hidden_states))
|
||||
|
||||
|
||||
class LingBotVideoMLP(nn.Module):
|
||||
"""Dense SwiGLU feed-forward network."""
|
||||
|
||||
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
|
||||
super().__init__()
|
||||
self.gate_proj = LingBotVideoLinear(hidden_size, intermediate_size, bias=False)
|
||||
self.up_proj = LingBotVideoLinear(hidden_size, intermediate_size, bias=False)
|
||||
self.down_proj = LingBotVideoLinear(intermediate_size, hidden_size, bias=False)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply the released SiLU-gated MLP ordering."""
|
||||
return self.down_proj(F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states))
|
||||
|
||||
|
||||
class LingBotVideoRouter(nn.Module):
|
||||
"""Released token-choice router with bias-only expert selection correction."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_experts: int,
|
||||
top_k: int,
|
||||
score_func: str,
|
||||
norm_topk_prob: bool,
|
||||
n_group: int | None,
|
||||
topk_group: int | None,
|
||||
route_scale: float,
|
||||
) -> None:
|
||||
"""Create fp32-routed expert scores with the released persistent bias."""
|
||||
super().__init__()
|
||||
self.num_experts = num_experts
|
||||
self.top_k = top_k
|
||||
self.score_func = score_func
|
||||
self.norm_topk_prob = norm_topk_prob
|
||||
self.n_group = n_group
|
||||
self.topk_group = topk_group
|
||||
self.route_scale = route_scale
|
||||
self.weight = nn.Parameter(torch.empty(num_experts, hidden_size))
|
||||
self.register_buffer("e_score_correction_bias", torch.zeros(num_experts), persistent=True)
|
||||
|
||||
def _group_limited_topk(self, scores_for_choice: torch.Tensor) -> torch.Tensor:
|
||||
"""Restrict token choices to groups with the two strongest expert scores."""
|
||||
sequence_length = scores_for_choice.shape[0]
|
||||
experts_per_group = self.num_experts // self.n_group
|
||||
grouped = scores_for_choice.view(sequence_length, self.n_group, experts_per_group)
|
||||
group_scores = grouped.topk(2, dim=-1)[0].sum(dim=-1)
|
||||
group_indices = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
|
||||
group_mask = torch.zeros_like(group_scores)
|
||||
group_mask.scatter_(1, group_indices, 1)
|
||||
score_mask = (group_mask.unsqueeze(-1).expand(sequence_length, self.n_group,
|
||||
experts_per_group).reshape(sequence_length, -1))
|
||||
masked_scores = scores_for_choice.masked_fill(~score_mask.bool(), float("-inf"))
|
||||
return torch.topk(masked_scores, k=self.top_k, dim=-1, sorted=False)[1]
|
||||
|
||||
def forward(self,
|
||||
tokens: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Score in fp32, select with correction bias, and weight without it."""
|
||||
with torch.amp.autocast(tokens.device.type, enabled=False):
|
||||
logits = F.linear(tokens.float(), self.weight.float())
|
||||
scores = F.softmax(logits, dim=-1) if self.score_func == "softmax" else logits.sigmoid()
|
||||
scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0)
|
||||
if self.n_group is not None and self.n_group > 1:
|
||||
top_indices = self._group_limited_topk(scores_for_choice)
|
||||
else:
|
||||
top_indices = torch.topk(scores_for_choice, k=self.top_k, dim=-1, sorted=False)[1]
|
||||
top_scores = scores.gather(1, top_indices)
|
||||
if self.top_k > 1 and self.norm_topk_prob:
|
||||
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-20)
|
||||
top_scores = top_scores * self.route_scale
|
||||
return top_indices, top_scores.to(tokens.dtype), logits, scores, scores_for_choice
|
||||
|
||||
|
||||
class LingBotVideoGroupedExperts(nn.Module):
|
||||
"""Released grouped-expert parameter layout: w1/w3 `[E,I,H]`, w2 `[E,H,I]`."""
|
||||
|
||||
def __init__(self, num_experts: int, hidden_size: int, intermediate_size: int) -> None:
|
||||
super().__init__()
|
||||
self.num_experts = num_experts
|
||||
self.w1 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size))
|
||||
self.w2 = nn.Parameter(torch.empty(num_experts, hidden_size, intermediate_size))
|
||||
self.w3 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size))
|
||||
|
||||
|
||||
def _round_up_to_multiple(value: int, multiple: int) -> int:
|
||||
"""Round an integer up to the next multiple."""
|
||||
return ((value + multiple - 1) // multiple) * multiple
|
||||
|
||||
|
||||
class LingBotVideoSparseMoeBlock(nn.Module):
|
||||
"""Token-choice sparse feed-forward block matching the released state surface."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
intermediate_size: int,
|
||||
num_experts: int,
|
||||
top_k: int,
|
||||
moe_intermediate_size: int,
|
||||
score_func: str,
|
||||
norm_topk_prob: bool,
|
||||
n_group: int | None,
|
||||
topk_group: int | None,
|
||||
routed_scaling_factor: float,
|
||||
n_shared_experts: int | None,
|
||||
) -> None:
|
||||
"""Create routed and optional shared experts with released parameter names."""
|
||||
super().__init__()
|
||||
del intermediate_size
|
||||
self.hidden_size = hidden_size
|
||||
self.num_experts = num_experts
|
||||
self.router = LingBotVideoRouter(
|
||||
hidden_size,
|
||||
num_experts,
|
||||
top_k,
|
||||
score_func,
|
||||
norm_topk_prob,
|
||||
n_group,
|
||||
topk_group,
|
||||
routed_scaling_factor,
|
||||
)
|
||||
self.experts = LingBotVideoGroupedExperts(num_experts, hidden_size, moe_intermediate_size)
|
||||
self.shared_experts = None
|
||||
if n_shared_experts is not None and n_shared_experts > 0:
|
||||
self.shared_experts = LingBotVideoMLP(hidden_size, moe_intermediate_size * n_shared_experts)
|
||||
|
||||
@staticmethod
|
||||
def _reorder_tokens(
|
||||
tokens: torch.Tensor,
|
||||
top_scores: torch.Tensor,
|
||||
top_indices: torch.Tensor,
|
||||
num_experts: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
|
||||
"""Pack active token choices into stable expert-major order."""
|
||||
num_tokens = tokens.shape[0]
|
||||
top_k = top_indices.shape[1]
|
||||
flat_scores = top_scores.reshape(-1)
|
||||
flat_indices = top_indices.reshape(-1)
|
||||
active_positions = torch.where(flat_scores != 0)[0]
|
||||
active_experts = flat_indices[active_positions]
|
||||
counts = torch.zeros(num_experts, device=tokens.device, dtype=torch.int64)
|
||||
counts.scatter_add_(0, active_experts, torch.ones_like(active_experts, dtype=torch.int64))
|
||||
sort_order = torch.argsort(active_experts, stable=True)
|
||||
sorted_positions = active_positions[sort_order]
|
||||
sorted_scores = flat_scores[sorted_positions]
|
||||
original_token_indices = sorted_positions // top_k
|
||||
permuted_tokens = tokens[original_token_indices]
|
||||
return permuted_tokens, counts, sorted_positions, sorted_scores, num_tokens, top_k
|
||||
|
||||
@staticmethod
|
||||
def _pad_grouped_tokens(tokens: torch.Tensor,
|
||||
counts: torch.Tensor,
|
||||
align: int = 8) -> tuple[torch.Size, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Align each expert segment for `torch._grouped_mm` and retain unpad indices."""
|
||||
num_tokens = tokens.shape[0]
|
||||
num_experts = int(counts.shape[0])
|
||||
max_length = _round_up_to_multiple(num_tokens + num_experts * align, align)
|
||||
# Build the small padding index on CPU after one synchronization instead
|
||||
# of reading three CUDA scalars for every expert.
|
||||
counts_cpu = counts.to(device="cpu", dtype=torch.int64)
|
||||
total_per_expert = torch.clamp_min(counts_cpu, align)
|
||||
aligned_counts_cpu = ((total_per_expert + align - 1) // align * align).to(torch.int32)
|
||||
write_offsets = torch.cumsum(aligned_counts_cpu, dim=0) - aligned_counts_cpu
|
||||
start_indices = torch.cumsum(counts_cpu, dim=0) - counts_cpu
|
||||
permuted_indices_cpu = torch.full((max_length, ), num_tokens, dtype=torch.int64, device="cpu")
|
||||
for expert_index in range(num_experts):
|
||||
length = int(counts_cpu[expert_index])
|
||||
if length == 0:
|
||||
continue
|
||||
write_start = int(write_offsets[expert_index])
|
||||
start = int(start_indices[expert_index])
|
||||
permuted_indices_cpu[write_start:write_start + length] = torch.arange(start,
|
||||
start + length,
|
||||
device="cpu",
|
||||
dtype=torch.int64)
|
||||
permuted_indices = permuted_indices_cpu.to(tokens.device)
|
||||
aligned_counts = aligned_counts_cpu.to(tokens.device)
|
||||
tokens_with_pad = torch.vstack((tokens, tokens.new_zeros((tokens.shape[-1], ))))
|
||||
input_shape = tokens_with_pad.shape
|
||||
return input_shape, tokens_with_pad[permuted_indices], permuted_indices, aligned_counts
|
||||
|
||||
@staticmethod
|
||||
def _unpad_grouped_tokens(output: torch.Tensor, input_shape: torch.Size,
|
||||
permuted_indices: torch.Tensor) -> torch.Tensor:
|
||||
"""Undo per-expert alignment while dropping the shared padding row."""
|
||||
unpermuted = output.new_empty(input_shape)
|
||||
unpermuted[permuted_indices, :] = output
|
||||
return unpermuted[:-1]
|
||||
|
||||
def _run_grouped_experts(self, tokens: torch.Tensor, counts: torch.Tensor) -> torch.Tensor:
|
||||
"""Use the released bf16 grouped matmuls on CUDA and an eager CPU fallback."""
|
||||
if tokens.device.type == "cpu" or not hasattr(torch, "_grouped_mm"):
|
||||
return self._run_experts_for_loop(tokens, counts)
|
||||
input_shape, padded_tokens, permuted_indices, aligned_counts = self._pad_grouped_tokens(tokens, counts)
|
||||
offsets = torch.cumsum(aligned_counts, dim=0, dtype=torch.int32)
|
||||
hidden = F.silu(
|
||||
torch._grouped_mm(
|
||||
padded_tokens.bfloat16(),
|
||||
self.experts.w1.bfloat16().transpose(-2, -1),
|
||||
offs=offsets,
|
||||
))
|
||||
hidden = hidden * torch._grouped_mm(
|
||||
padded_tokens.bfloat16(),
|
||||
self.experts.w3.bfloat16().transpose(-2, -1),
|
||||
offs=offsets,
|
||||
)
|
||||
output = torch._grouped_mm(
|
||||
hidden,
|
||||
self.experts.w2.bfloat16().transpose(-2, -1),
|
||||
offs=offsets,
|
||||
).type_as(padded_tokens)
|
||||
return self._unpad_grouped_tokens(output, input_shape, permuted_indices)
|
||||
|
||||
def _run_experts_for_loop(self, tokens: torch.Tensor, counts: torch.Tensor) -> torch.Tensor:
|
||||
"""Evaluate contiguous expert segments eagerly for CPU correctness tests."""
|
||||
splits = torch.split(tokens, counts.tolist(), dim=0)
|
||||
outputs: list[torch.Tensor] = []
|
||||
for expert_index, expert_tokens in enumerate(splits):
|
||||
if expert_tokens.numel() == 0:
|
||||
continue
|
||||
hidden = F.silu(expert_tokens @ self.experts.w1[expert_index].transpose(-2, -1))
|
||||
hidden = hidden * (expert_tokens @ self.experts.w3[expert_index].transpose(-2, -1))
|
||||
outputs.append(hidden @ self.experts.w2[expert_index].transpose(-2, -1))
|
||||
if not outputs:
|
||||
return tokens.new_zeros(tokens.shape)
|
||||
return torch.cat(outputs, dim=0)
|
||||
|
||||
@staticmethod
|
||||
def _restore_tokens(
|
||||
expert_output: torch.Tensor,
|
||||
sorted_positions: torch.Tensor,
|
||||
sorted_scores: torch.Tensor,
|
||||
num_tokens: int,
|
||||
top_k: int,
|
||||
) -> torch.Tensor:
|
||||
"""Restore token order and combine expert outputs with fp32 weighted sums."""
|
||||
hidden_size = expert_output.shape[-1]
|
||||
unsorted = torch.zeros(
|
||||
(num_tokens * top_k, hidden_size),
|
||||
dtype=expert_output.dtype,
|
||||
device=expert_output.device,
|
||||
)
|
||||
unsorted[sorted_positions] = expert_output
|
||||
unsorted = unsorted.reshape(num_tokens, top_k, hidden_size)
|
||||
scores_unsorted = torch.zeros(
|
||||
num_tokens * top_k,
|
||||
dtype=sorted_scores.dtype,
|
||||
device=sorted_scores.device,
|
||||
)
|
||||
scores_unsorted[sorted_positions] = sorted_scores
|
||||
scores_unsorted = scores_unsorted.reshape(num_tokens, top_k, 1)
|
||||
return (unsorted.float() * scores_unsorted).sum(dim=1).to(expert_output.dtype)
|
||||
|
||||
def _run_selected_experts(
|
||||
self,
|
||||
tokens: torch.Tensor,
|
||||
top_scores: torch.Tensor,
|
||||
top_indices: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Dispatch routed choices, execute experts, and restore token-major order."""
|
||||
permuted_tokens, counts, sorted_positions, sorted_scores, num_tokens, top_k = self._reorder_tokens(
|
||||
tokens, top_scores, top_indices, self.router.num_experts)
|
||||
expert_output = self._run_grouped_experts(permuted_tokens, counts)
|
||||
return self._restore_tokens(expert_output, sorted_positions, sorted_scores, num_tokens, top_k)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, padding_mask: torch.Tensor | None = None) -> torch.Tensor:
|
||||
"""Route token choices, zero padded routes, and add optional shared experts."""
|
||||
batch = hidden_states.shape[0]
|
||||
tokens = hidden_states.view(-1, self.hidden_size)
|
||||
top_indices, top_scores, logits, scores, scores_for_choice = self.router(tokens)
|
||||
del logits, scores, scores_for_choice
|
||||
if padding_mask is not None:
|
||||
mask = padding_mask.unsqueeze(-1).to(top_scores.dtype)
|
||||
top_scores = top_scores * mask
|
||||
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-9)
|
||||
top_scores = top_scores * self.router.route_scale
|
||||
output = self._run_selected_experts(tokens, top_scores, top_indices)
|
||||
output = output.view(batch, -1, self.hidden_size)
|
||||
if self.shared_experts is not None:
|
||||
output = output + self.shared_experts(hidden_states)
|
||||
return output
|
||||
|
||||
|
||||
class LingBotVideoBlock(nn.Module):
|
||||
"""One Dense or sparse LingBot-Video transformer block."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
intermediate_size: int,
|
||||
norm_eps: float,
|
||||
qkv_bias: bool,
|
||||
out_bias: bool,
|
||||
num_experts: int,
|
||||
num_experts_per_tok: int,
|
||||
moe_intermediate_size: int,
|
||||
decoder_sparse_step: int,
|
||||
mlp_only_layers: tuple[int, ...] | list[int],
|
||||
n_shared_experts: int | None,
|
||||
score_func: str,
|
||||
norm_topk_prob: bool,
|
||||
n_group: int | None,
|
||||
topk_group: int | None,
|
||||
routed_scaling_factor: float,
|
||||
layer_idx: int,
|
||||
) -> None:
|
||||
"""Select the released Dense or sparse feed-forward structure for one layer."""
|
||||
super().__init__()
|
||||
self.layer_idx = layer_idx
|
||||
self.scale_shift_table = nn.Parameter(torch.zeros(1, 6 * hidden_size))
|
||||
self.norm1 = LingBotVideoRMSNorm(hidden_size, norm_eps)
|
||||
self.attn = LingBotVideoAttention(hidden_size, num_attention_heads, norm_eps, qkv_bias, out_bias)
|
||||
self.norm_post_attn = LingBotVideoRMSNorm(hidden_size, norm_eps)
|
||||
self.norm2 = LingBotVideoRMSNorm(hidden_size, norm_eps)
|
||||
if layer_idx not in mlp_only_layers and (num_experts > 0 and (layer_idx + 1) % decoder_sparse_step == 0):
|
||||
self.ffn = LingBotVideoSparseMoeBlock(
|
||||
hidden_size,
|
||||
intermediate_size,
|
||||
num_experts,
|
||||
num_experts_per_tok,
|
||||
moe_intermediate_size,
|
||||
score_func,
|
||||
norm_topk_prob,
|
||||
n_group,
|
||||
topk_group,
|
||||
routed_scaling_factor,
|
||||
n_shared_experts,
|
||||
)
|
||||
else:
|
||||
self.ffn = LingBotVideoMLP(hidden_size, intermediate_size)
|
||||
self.norm_post_ffn = LingBotVideoRMSNorm(hidden_size, norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
temb6: torch.Tensor,
|
||||
rotary_emb: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
moe_padding_mask: torch.Tensor | None = None,
|
||||
original_seq_len: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run attention and configured feed-forward residual branches with fp32 AdaLN."""
|
||||
expected_tokens = hidden_states.shape[0] * hidden_states.shape[1]
|
||||
if temb6.ndim != 2 or temb6.shape[0] != expected_tokens:
|
||||
raise ValueError("LingBotVideoBlock expects token-level temb6 with shape "
|
||||
f"(B*S, 6D); got {tuple(temb6.shape)} for {tuple(hidden_states.shape)}.")
|
||||
modulation = temb6.view(*hidden_states.shape[:2], -1) + self.scale_shift_table.unsqueeze(0)
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = modulation.chunk(6, dim=-1)
|
||||
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
|
||||
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
|
||||
bulk_dtype = self.attn.to_q.weight.dtype
|
||||
|
||||
attention_input = (self.norm1(hidden_states) * scale_msa + shift_msa).to(bulk_dtype)
|
||||
attention_output = self.attn(attention_input, rotary_emb, attention_mask, original_seq_len)
|
||||
hidden_states = hidden_states + (gate_msa * self.norm_post_attn(attention_output)).to(hidden_states.dtype)
|
||||
mlp_input = (self.norm2(hidden_states) * scale_mlp + shift_mlp).to(bulk_dtype)
|
||||
if isinstance(self.ffn, LingBotVideoSparseMoeBlock):
|
||||
mlp_output = self.ffn(mlp_input, padding_mask=moe_padding_mask)
|
||||
else:
|
||||
mlp_output = self.ffn(mlp_input)
|
||||
mlp_output = self.norm_post_ffn(mlp_output)
|
||||
return hidden_states + (gate_mlp * mlp_output).to(hidden_states.dtype)
|
||||
|
||||
|
||||
class LingBotVideoTimestepEmbedding(nn.Module):
|
||||
"""Two-layer timestep embedding with released parameter names."""
|
||||
|
||||
def __init__(self, input_dim: int, hidden_size: int, bias: bool) -> None:
|
||||
super().__init__()
|
||||
self.linear_1 = LingBotVideoLinear(input_dim, hidden_size, bias=bias)
|
||||
self.linear_2 = LingBotVideoLinear(hidden_size, hidden_size, bias=bias)
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply the official linear-SiLU-linear timestep projection."""
|
||||
return self.linear_2(F.silu(self.linear_1(sample)))
|
||||
|
||||
|
||||
class LingBotVideoTransformer3DModel(BaseDiT):
|
||||
"""LingBot-Video DiT with a source-compatible Dense or MoE state surface."""
|
||||
|
||||
_fsdp_shard_conditions = [_is_lingbot_video_block]
|
||||
_compile_conditions = [_is_lingbot_video_block]
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
)
|
||||
param_names_mapping = {r"^(.*)$": r"\1"}
|
||||
reverse_param_names_mapping = {r"^(.*)$": r"\1"}
|
||||
|
||||
def _get_parameter_dtype(self, name: str, default_dtype: torch.dtype) -> torch.dtype:
|
||||
"""Select the released mixed-precision dtype while loading each parameter."""
|
||||
return torch.float32 if _keep_in_fp32(name) else default_dtype
|
||||
|
||||
def __init__(self, config: LingBotVideoConfig, hf_config: dict[str, Any]) -> None:
|
||||
"""Construct the Dense or MoE variant from its released transformer config."""
|
||||
config.update_model_arch(hf_config)
|
||||
super().__init__(config, hf_config)
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
if head_dim != sum(config.axes_dims):
|
||||
raise ValueError(f"head_dim {head_dim} != sum(axes_dims) {sum(config.axes_dims)}")
|
||||
sp_world_size = _sequence_parallel_world_size()
|
||||
assert config.num_attention_heads % sp_world_size == 0, (
|
||||
f"The number of attention heads ({config.num_attention_heads}) must be divisible by "
|
||||
f"the sequence parallel size ({sp_world_size})")
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.in_channels
|
||||
self.patch_embedder = LingBotVideoLinear(
|
||||
config.in_channels * math.prod(config.patch_size),
|
||||
config.hidden_size,
|
||||
bias=config.patch_embed_bias,
|
||||
)
|
||||
self.time_proj = Timesteps(config.freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = LingBotVideoTimestepEmbedding(config.freq_dim, config.hidden_size,
|
||||
config.timestep_mlp_bias)
|
||||
self.time_modulation = nn.Sequential(nn.SiLU(), LingBotVideoLinear(config.hidden_size, 6 * config.hidden_size))
|
||||
self.text_embedder = LingBotVideoTextEmbedder(config.text_dim, config.hidden_size)
|
||||
self.rope = LingBotVideoRotaryEmbedding(tuple(config.axes_dims), tuple(config.axes_lens), config.rope_theta)
|
||||
self.blocks = nn.ModuleList([
|
||||
LingBotVideoBlock(
|
||||
config.hidden_size,
|
||||
config.num_attention_heads,
|
||||
config.intermediate_size,
|
||||
config.norm_eps,
|
||||
config.qkv_bias,
|
||||
config.out_bias,
|
||||
config.num_experts,
|
||||
config.num_experts_per_tok,
|
||||
config.moe_intermediate_size,
|
||||
config.decoder_sparse_step,
|
||||
config.mlp_only_layers,
|
||||
config.n_shared_experts,
|
||||
config.score_func,
|
||||
config.norm_topk_prob,
|
||||
config.n_group,
|
||||
config.topk_group,
|
||||
config.routed_scaling_factor,
|
||||
layer_index,
|
||||
) for layer_index in range(config.depth)
|
||||
])
|
||||
self.norm_out = nn.LayerNorm(config.hidden_size, elementwise_affine=False, eps=config.norm_eps)
|
||||
self.norm_out_modulation = nn.Sequential(nn.SiLU(),
|
||||
LingBotVideoLinear(config.hidden_size, 2 * config.hidden_size))
|
||||
self.proj_out = LingBotVideoLinear(config.hidden_size, math.prod(config.patch_size) * config.out_channels)
|
||||
self.__post_init__()
|
||||
|
||||
def to(self, *args: Any, **kwargs: Any) -> LingBotVideoTransformer3DModel:
|
||||
"""Cast bulk weights while retaining the released fp32-sensitive modules."""
|
||||
device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs)
|
||||
if dtype is None or dtype == torch.float32:
|
||||
return super().to(*args, **kwargs)
|
||||
if not torch.is_floating_point(torch.empty((), dtype=dtype)):
|
||||
return super().to(*args, **kwargs)
|
||||
if device is not None:
|
||||
super().to(device=device, non_blocking=non_blocking)
|
||||
for name, parameter in self.named_parameters():
|
||||
if torch.is_floating_point(parameter):
|
||||
target_dtype = torch.float32 if _keep_in_fp32(name) else dtype
|
||||
parameter.data = parameter.data.to(target_dtype, non_blocking=non_blocking)
|
||||
if parameter.grad is not None:
|
||||
parameter.grad.data = parameter.grad.data.to(target_dtype, non_blocking=non_blocking)
|
||||
for name, buffer in self.named_buffers():
|
||||
if torch.is_floating_point(buffer):
|
||||
target_dtype = torch.float32 if _keep_in_fp32(name) else dtype
|
||||
buffer.data = buffer.data.to(target_dtype, non_blocking=non_blocking)
|
||||
return self
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
guidance: torch.Tensor | None = None,
|
||||
encoder_attention_mask: torch.Tensor | None = None,
|
||||
return_dict: bool = True,
|
||||
**kwargs: Any,
|
||||
) -> LingBotVideoTransformerOutput | tuple[torch.Tensor]:
|
||||
"""Denoise video latents with joint video-text attention."""
|
||||
del encoder_hidden_states_image, guidance, kwargs
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
if len(encoder_hidden_states) != 1:
|
||||
raise ValueError("LingBot-Video expects one text-encoder output tensor.")
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
batch, channels, frames, height, width = hidden_states.shape
|
||||
patch_t, patch_h, patch_w = self.config.patch_size
|
||||
grid_t, grid_h, grid_w = frames // patch_t, height // patch_h, width // patch_w
|
||||
video_tokens = grid_t * grid_h * grid_w
|
||||
text_tokens = encoder_hidden_states.shape[1]
|
||||
device = hidden_states.device
|
||||
if encoder_attention_mask is None:
|
||||
encoder_attention_mask = torch.ones((batch, text_tokens), device=device, dtype=torch.bool)
|
||||
text_lengths = encoder_attention_mask.sum(dim=-1).long()
|
||||
|
||||
patches = hidden_states.reshape(batch, channels, grid_t, patch_t, grid_h, patch_h, grid_w, patch_w)
|
||||
patches = patches.permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(batch, video_tokens,
|
||||
patch_t * patch_h * patch_w * channels)
|
||||
video_hidden = self.patch_embedder(patches)
|
||||
text_hidden = self.text_embedder(encoder_hidden_states)
|
||||
joint = torch.cat([video_hidden, text_hidden], dim=1)
|
||||
|
||||
rotary_parts: list[torch.Tensor] = []
|
||||
for index in range(batch):
|
||||
real_text_length = int(text_lengths[index].item())
|
||||
positions = _make_joint_position_ids(real_text_length, grid_t, grid_h, grid_w, device)
|
||||
maxima = (real_text_length + grid_t, grid_h - 1, grid_w - 1)
|
||||
rotary = self.rope(positions, maxima=maxima)
|
||||
if real_text_length < text_tokens:
|
||||
padding = torch.zeros(
|
||||
text_tokens - real_text_length,
|
||||
rotary.shape[-1],
|
||||
device=device,
|
||||
dtype=rotary.dtype,
|
||||
)
|
||||
rotary = torch.cat([rotary, padding], dim=0)
|
||||
rotary_parts.append(rotary)
|
||||
rotary_emb = torch.stack(rotary_parts, dim=0)
|
||||
|
||||
attention_mask = None
|
||||
moe_padding_mask = None
|
||||
if bool((text_lengths < text_tokens).any()):
|
||||
key_mask = torch.cat(
|
||||
[
|
||||
torch.ones(batch, video_tokens, dtype=torch.bool, device=device),
|
||||
encoder_attention_mask.bool(),
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
attention_mask = key_mask[:, None, None, :]
|
||||
moe_padding_mask = key_mask
|
||||
timestep_projection = self.time_proj(timestep.float())
|
||||
timestep_embedding = self.time_embedder(timestep_projection)
|
||||
token_embedding = timestep_embedding.unsqueeze(1).expand(batch, joint.shape[1], -1)
|
||||
original_joint_length: int | None = None
|
||||
if _sequence_parallel_world_size() > 1:
|
||||
# Match the official CP order: project token modulation before
|
||||
# placing every token-aligned tensor on the same padded shard.
|
||||
temb6 = self.time_modulation(token_embedding.reshape(-1,
|
||||
self.hidden_size)).reshape(batch, joint.shape[1], -1)
|
||||
joint, original_joint_length = sequence_model_parallel_shard(joint, dim=1)
|
||||
rotary_emb, _ = sequence_model_parallel_shard(rotary_emb, dim=1)
|
||||
token_embedding, _ = sequence_model_parallel_shard(token_embedding, dim=1)
|
||||
temb6, _ = sequence_model_parallel_shard(temb6, dim=1)
|
||||
if moe_padding_mask is None:
|
||||
moe_padding_mask = torch.ones(batch, original_joint_length, dtype=torch.bool, device=device)
|
||||
moe_padding_mask, _ = sequence_model_parallel_shard(moe_padding_mask, dim=1)
|
||||
moe_padding_mask = moe_padding_mask.reshape(-1)
|
||||
temb6 = temb6.reshape(-1, 6 * self.hidden_size)
|
||||
else:
|
||||
temb6 = self.time_modulation(token_embedding.reshape(-1, self.hidden_size))
|
||||
if moe_padding_mask is not None:
|
||||
moe_padding_mask = moe_padding_mask.reshape(-1)
|
||||
|
||||
for block in self.blocks:
|
||||
joint = block(
|
||||
joint,
|
||||
temb6,
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
moe_padding_mask,
|
||||
original_joint_length,
|
||||
)
|
||||
|
||||
final_modulation = self.norm_out_modulation(token_embedding.reshape(-1, self.hidden_size))
|
||||
shift, scale = final_modulation.reshape(*joint.shape[:2], -1).chunk(2, dim=-1)
|
||||
final_hidden = self.norm_out(joint) * (1.0 + scale) + shift
|
||||
projected = self.proj_out(final_hidden.to(self.proj_out.weight.dtype))
|
||||
if original_joint_length is not None:
|
||||
projected = sequence_model_parallel_all_gather_with_unpad(projected, original_joint_length, dim=1)
|
||||
projected = projected[:, :video_tokens]
|
||||
output_channels = self.config.out_channels
|
||||
output = projected.reshape(batch, grid_t, grid_h, grid_w, patch_t, patch_h, patch_w, output_channels)
|
||||
output = output.permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(batch, output_channels, frames, height, width)
|
||||
if not return_dict:
|
||||
return (output, )
|
||||
return LingBotVideoTransformerOutput(sample=output)
|
||||
|
||||
|
||||
EntryClass = LingBotVideoTransformer3DModel
|
||||
@@ -0,0 +1,5 @@
|
||||
from .causal_fast import LingBotWorld2CausalFastTransformer3DModel
|
||||
|
||||
__all__ = ["LingBotWorld2CausalFastTransformer3DModel"]
|
||||
|
||||
EntryClass = LingBotWorld2CausalFastTransformer3DModel
|
||||
@@ -0,0 +1,204 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import numpy as np
|
||||
import os
|
||||
import torch
|
||||
from scipy.interpolate import interp1d
|
||||
from scipy.spatial.transform import Rotation, Slerp
|
||||
|
||||
|
||||
# --- Official Code (Leave Unchanged) ---
|
||||
|
||||
def interpolate_camera_poses(
|
||||
src_indices: np.ndarray,
|
||||
src_rot_mat: np.ndarray,
|
||||
src_trans_vec: np.ndarray,
|
||||
tgt_indices: np.ndarray,
|
||||
) -> torch.Tensor:
|
||||
# interpolate translation
|
||||
interp_func_trans = interp1d(
|
||||
src_indices,
|
||||
src_trans_vec,
|
||||
axis=0,
|
||||
kind='linear',
|
||||
bounds_error=False,
|
||||
fill_value="extrapolate",
|
||||
)
|
||||
interpolated_trans_vec = interp_func_trans(tgt_indices)
|
||||
|
||||
# interpolate rotation
|
||||
src_quat_vec = Rotation.from_matrix(src_rot_mat)
|
||||
# ensure there is no sudden change in qw
|
||||
quats = src_quat_vec.as_quat().copy() # [N, 4]
|
||||
for i in range(1, len(quats)):
|
||||
if np.dot(quats[i], quats[i-1]) < 0:
|
||||
quats[i] = -quats[i]
|
||||
src_quat_vec = Rotation.from_quat(quats)
|
||||
slerp_func_rot = Slerp(src_indices, src_quat_vec)
|
||||
interpolated_rot_quat = slerp_func_rot(tgt_indices)
|
||||
interpolated_rot_mat = interpolated_rot_quat.as_matrix()
|
||||
|
||||
poses = np.zeros((len(tgt_indices), 4, 4))
|
||||
poses[:, :3, :3] = interpolated_rot_mat
|
||||
poses[:, :3, 3] = interpolated_trans_vec
|
||||
poses[:, 3, 3] = 1.0
|
||||
return torch.from_numpy(poses).float()
|
||||
|
||||
|
||||
def SE3_inverse(T: torch.Tensor) -> torch.Tensor:
|
||||
Rot = T[:, :3, :3] # [B,3,3]
|
||||
trans = T[:, :3, 3:] # [B,3,1]
|
||||
R_inv = Rot.transpose(-1, -2)
|
||||
t_inv = -torch.bmm(R_inv, trans)
|
||||
T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1)
|
||||
T_inv[:, :3, :3] = R_inv
|
||||
T_inv[:, :3, 3:] = t_inv
|
||||
return T_inv
|
||||
|
||||
|
||||
def compute_relative_poses(
|
||||
c2ws_mat: torch.Tensor,
|
||||
framewise: bool = False,
|
||||
normalize_trans: bool = True,
|
||||
) -> torch.Tensor:
|
||||
ref_w2cs = SE3_inverse(c2ws_mat[0:1])
|
||||
relative_poses = torch.matmul(ref_w2cs, c2ws_mat)
|
||||
# ensure identity matrix for 1st frame
|
||||
relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype)
|
||||
if framewise:
|
||||
# compute pose between i and i+1
|
||||
relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:])
|
||||
relative_poses[1:] = relative_poses_framewise
|
||||
if normalize_trans: # note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning."
|
||||
translations = relative_poses[:, :3, 3] # [f, 3]
|
||||
max_norm = torch.norm(translations, dim=-1).max()
|
||||
# only normlaize when moving
|
||||
if max_norm > 0:
|
||||
relative_poses[:, :3, 3] = translations / max_norm
|
||||
return relative_poses
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def create_meshgrid(n_frames: int, height: int, width: int, bias: float = 0.5, device='cuda', dtype=torch.float32) -> torch.Tensor:
|
||||
x_range = torch.arange(width, device=device, dtype=dtype)
|
||||
y_range = torch.arange(height, device=device, dtype=dtype)
|
||||
grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij')
|
||||
grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2]
|
||||
grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2]
|
||||
return grid_xy
|
||||
|
||||
|
||||
def get_plucker_embeddings(
|
||||
c2ws_mat: torch.Tensor,
|
||||
Ks: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
):
|
||||
n_frames = c2ws_mat.shape[0]
|
||||
device = c2ws_mat.device
|
||||
dtype = c2ws_mat.dtype
|
||||
|
||||
grid_y, grid_x = torch.meshgrid(
|
||||
torch.arange(height, device=device, dtype=dtype) + 0.5,
|
||||
torch.arange(width, device=device, dtype=dtype) + 0.5,
|
||||
indexing='ij',
|
||||
)
|
||||
x_flat = grid_x.reshape(-1)
|
||||
y_flat = grid_y.reshape(-1)
|
||||
|
||||
fx, fy, cx, cy = Ks[0, 0], Ks[0, 1], Ks[0, 2], Ks[0, 3]
|
||||
dirs = torch.stack([
|
||||
(x_flat - cx) / fx,
|
||||
(y_flat - cy) / fy,
|
||||
torch.ones_like(x_flat),
|
||||
], dim=-1)
|
||||
dirs = dirs / dirs.norm(dim=-1, keepdim=True)
|
||||
|
||||
rays_d = (c2ws_mat[:, :3, :3] @ dirs.T).transpose(1, 2)
|
||||
rays_o = c2ws_mat[:, :3, 3].unsqueeze(1).expand_as(rays_d)
|
||||
plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1)
|
||||
return plucker_embeddings.view(n_frames, height, width, 6)
|
||||
|
||||
|
||||
def get_Ks_transformed(
|
||||
Ks: torch.Tensor,
|
||||
height_org: int,
|
||||
width_org: int,
|
||||
height_resize: int,
|
||||
width_resize: int,
|
||||
height_final: int,
|
||||
width_final: int,
|
||||
):
|
||||
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
|
||||
|
||||
scale_x = width_resize / width_org
|
||||
scale_y = height_resize / height_org
|
||||
|
||||
fx_resize = fx * scale_x
|
||||
fy_resize = fy * scale_y
|
||||
cx_resize = cx * scale_x
|
||||
cy_resize = cy * scale_y
|
||||
|
||||
crop_offset_x = (width_resize - width_final) / 2
|
||||
crop_offset_y = (height_resize - height_final) / 2
|
||||
|
||||
cx_final = cx_resize - crop_offset_x
|
||||
cy_final = cy_resize - crop_offset_y
|
||||
|
||||
Ks_transformed = torch.zeros_like(Ks)
|
||||
Ks_transformed[:, 0:1] = fx_resize
|
||||
Ks_transformed[:, 1:2] = fy_resize
|
||||
Ks_transformed[:, 2:3] = cx_final
|
||||
Ks_transformed[:, 3:4] = cy_final
|
||||
|
||||
return Ks_transformed
|
||||
|
||||
|
||||
# --- Custom ---
|
||||
|
||||
def prepare_camera_embedding(
|
||||
action_path: str,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
spatial_scale: int = 8,
|
||||
) -> tuple[torch.Tensor, int]:
|
||||
c2ws = np.load(os.path.join(action_path, "poses.npy"))
|
||||
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
|
||||
num_frames = min(num_frames, len_c2ws)
|
||||
c2ws = c2ws[:num_frames]
|
||||
|
||||
Ks = torch.from_numpy(
|
||||
np.load(os.path.join(action_path, "intrinsics.npy"))
|
||||
).float()
|
||||
Ks = get_Ks_transformed(
|
||||
Ks,
|
||||
height_org=480,
|
||||
width_org=832,
|
||||
height_resize=height,
|
||||
width_resize=width,
|
||||
height_final=height,
|
||||
width_final=width,
|
||||
)
|
||||
Ks = Ks[0] # use first frame
|
||||
|
||||
len_c2ws = len(c2ws)
|
||||
num_latent_frames = (len_c2ws - 1) // 4 + 1
|
||||
c2ws_infer = interpolate_camera_poses(
|
||||
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
|
||||
src_rot_mat=c2ws[:, :3, :3],
|
||||
src_trans_vec=c2ws[:, :3, 3],
|
||||
tgt_indices=np.linspace(0, len_c2ws - 1, num_latent_frames),
|
||||
)
|
||||
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
|
||||
Ks = Ks.repeat(num_latent_frames, 1)
|
||||
plucker = get_plucker_embeddings(c2ws_infer, Ks, height, width) # [F, H, W, 6]
|
||||
|
||||
# reshpae
|
||||
latent_height = height // spatial_scale
|
||||
latent_width = width // spatial_scale
|
||||
plucker = plucker.view(num_latent_frames, latent_height, spatial_scale, latent_width, spatial_scale, 6)
|
||||
plucker = plucker.permute(0, 1, 3, 5, 2, 4).contiguous()
|
||||
plucker = plucker.view(num_latent_frames, latent_height, latent_width, 6 * spatial_scale * spatial_scale)
|
||||
c2ws_plucker_emb = plucker.permute(3, 0, 1, 2).contiguous().unsqueeze(0)
|
||||
|
||||
return c2ws_plucker_emb, num_frames
|
||||
@@ -0,0 +1,776 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot World 2 causal-fast DiT implemented inside FastVideo."""
|
||||
|
||||
import math
|
||||
import warnings
|
||||
from typing import Any
|
||||
|
||||
from einops import rearrange
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.dits.lingbotworld2 import (
|
||||
LingBotWorld2CausalFastVideoConfig,
|
||||
)
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather,
|
||||
sequence_model_parallel_all_to_all_4D,
|
||||
)
|
||||
from fastvideo.distributed.parallel_state import get_sp_parallel_rank, get_sp_world_size
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
try:
|
||||
import flash_attn_interface
|
||||
|
||||
FLASH_ATTN_3_AVAILABLE = True
|
||||
except ModuleNotFoundError:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except ModuleNotFoundError:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim: int, position: torch.Tensor) -> torch.Tensor:
|
||||
"""Build Wan/LingBot World 2 sinusoidal timestep embeddings."""
|
||||
assert dim % 2 == 0
|
||||
half = dim // 2
|
||||
position = position.type(torch.float64)
|
||||
sinusoid = torch.outer(
|
||||
position,
|
||||
torch.pow(10000, -torch.arange(half, device=position.device).to(position).div(half)),
|
||||
)
|
||||
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
|
||||
|
||||
@torch.amp.autocast("cuda", enabled=False)
|
||||
def rope_params(max_seq_len: int, dim: int, theta: int = 10000) -> torch.Tensor:
|
||||
"""Return complex RoPE frequencies used by the released LingBot World 2 model."""
|
||||
assert dim % 2 == 0
|
||||
freqs = torch.outer(
|
||||
torch.arange(max_seq_len),
|
||||
1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float64).div(dim)),
|
||||
)
|
||||
return torch.polar(torch.ones_like(freqs), freqs)
|
||||
|
||||
|
||||
def flash_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q_lens: torch.Tensor | None = None,
|
||||
k_lens: torch.Tensor | None = None,
|
||||
dropout_p: float = 0.0,
|
||||
softmax_scale: float | None = None,
|
||||
q_scale: float | None = None,
|
||||
causal: bool = False,
|
||||
window_size: tuple[int, int] = (-1, -1),
|
||||
deterministic: bool = False,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
version: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run LingBot World 2-compatible FlashAttention on packed varlen inputs."""
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
assert dtype in half_dtypes
|
||||
assert q.device.type == "cuda" and q.size(-1) <= 256
|
||||
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x: torch.Tensor) -> torch.Tensor:
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor([lq] * b, dtype=torch.int32, device=q.device)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens, strict=True)]))
|
||||
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor([lk] * b, dtype=torch.int32, device=k.device)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens, strict=True)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens, strict=True)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
warnings.warn("FlashAttention 3 is not available; using FlashAttention 2.")
|
||||
|
||||
cu_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32
|
||||
).to(q.device, non_blocking=True)
|
||||
cu_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32
|
||||
).to(k.device, non_blocking=True)
|
||||
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
x = flash_attn_interface.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=cu_q,
|
||||
cu_seqlens_k=cu_k,
|
||||
seqused_q=None,
|
||||
seqused_k=None,
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
).unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=cu_q,
|
||||
cu_seqlens_k=cu_k,
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic,
|
||||
).unflatten(0, (b, lq))
|
||||
return x.type(out_dtype)
|
||||
|
||||
|
||||
def attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q_lens: torch.Tensor | None = None,
|
||||
k_lens: torch.Tensor | None = None,
|
||||
dropout_p: float = 0.0,
|
||||
softmax_scale: float | None = None,
|
||||
q_scale: float | None = None,
|
||||
causal: bool = False,
|
||||
window_size: tuple[int, int] = (-1, -1),
|
||||
deterministic: bool = False,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
fa_version: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Dispatch LingBot World 2 attention to FlashAttention when available."""
|
||||
if FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE:
|
||||
return flash_attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
q_scale=q_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic,
|
||||
dtype=dtype,
|
||||
version=fa_version,
|
||||
)
|
||||
|
||||
if q_lens is not None or k_lens is not None:
|
||||
warnings.warn("Padding masks are disabled without FlashAttention.")
|
||||
q = q.transpose(1, 2).to(dtype)
|
||||
k = k.transpose(1, 2).to(dtype)
|
||||
v = v.transpose(1, 2).to(dtype)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=None, is_causal=causal, dropout_p=dropout_p
|
||||
)
|
||||
return out.transpose(1, 2).contiguous()
|
||||
|
||||
|
||||
@torch.amp.autocast("cuda", enabled=False)
|
||||
def causal_rope_apply(
|
||||
x: torch.Tensor,
|
||||
grid_sizes: torch.Tensor,
|
||||
freqs: torch.Tensor,
|
||||
start_frame: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Apply LingBot World 2 causal RoPE with the current chunk frame offset."""
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2))
|
||||
freqs_i = torch.cat(
|
||||
[
|
||||
freqs[0][start_frame : start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),
|
||||
],
|
||||
dim=-1,
|
||||
).reshape(seq_len, 1, -1)
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
output.append(x_i)
|
||||
return torch.stack(output).type_as(x)
|
||||
|
||||
|
||||
class WanRMSNorm(nn.Module):
|
||||
"""RMSNorm used by Wan/LingBot World 2 attention projections."""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize the last dimension in fp32 and restore input dtype."""
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class WanLayerNorm(nn.LayerNorm):
|
||||
"""LayerNorm variant that computes in fp32 and returns the input dtype."""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply layer norm in fp32 for LingBot World 2 numerical parity."""
|
||||
return super().forward(x.float()).type_as(x)
|
||||
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
"""LingBot World 2 causal self-attention with rolling KV cache."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
qk_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.q = nn.Linear(dim, dim)
|
||||
self.k = nn.Linear(dim, dim)
|
||||
self.v = nn.Linear(dim, dim)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
grid_sizes: torch.Tensor,
|
||||
freqs: torch.Tensor,
|
||||
kv_cache: dict,
|
||||
current_start: int = 0,
|
||||
max_attention_size: int = 1_000_000,
|
||||
frame_seqlen: int | None = None,
|
||||
seq_lens_int: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Project QKV, update the rolling cache, and attend to its active window."""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
|
||||
if frame_seqlen is None:
|
||||
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
|
||||
current_start_frame = current_start // frame_seqlen
|
||||
if seq_lens_int is None:
|
||||
seq_lens_int = int(seq_lens[0].item() if seq_lens.dim() > 0 else seq_lens.item())
|
||||
|
||||
sp_size = get_sp_world_size()
|
||||
if sp_size > 1:
|
||||
q = sequence_model_parallel_all_to_all_4D(q, scatter_dim=2, gather_dim=1)
|
||||
k = sequence_model_parallel_all_to_all_4D(k, scatter_dim=2, gather_dim=1)
|
||||
v = sequence_model_parallel_all_to_all_4D(v, scatter_dim=2, gather_dim=1)
|
||||
padded_seq_len = s * sp_size
|
||||
roped_query = causal_rope_apply(q, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
|
||||
roped_key = causal_rope_apply(k, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
|
||||
roped_query = roped_query[:, :seq_lens_int]
|
||||
roped_key = roped_key[:, :seq_lens_int]
|
||||
v = v[:, :seq_lens_int]
|
||||
num_new_tokens = seq_lens_int
|
||||
else:
|
||||
padded_seq_len = s
|
||||
roped_query = causal_rope_apply(q, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
|
||||
roped_key = causal_rope_apply(k, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
|
||||
num_new_tokens = roped_query.shape[1]
|
||||
|
||||
current_end = current_start + num_new_tokens
|
||||
sink_tokens = self.sink_size * frame_seqlen
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
|
||||
if self.local_attn_size == -1:
|
||||
local_end_index = current_start + num_new_tokens
|
||||
local_start_index = current_start
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
elif (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size
|
||||
):
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
|
||||
kv_cache["k"][:, sink_tokens : sink_tokens + num_rolled_tokens] = kv_cache["k"][
|
||||
:, sink_tokens + num_evicted_tokens : sink_tokens + num_evicted_tokens + num_rolled_tokens
|
||||
].clone()
|
||||
kv_cache["v"][:, sink_tokens : sink_tokens + num_rolled_tokens] = kv_cache["v"][
|
||||
:, sink_tokens + num_evicted_tokens : sink_tokens + num_evicted_tokens + num_rolled_tokens
|
||||
].clone()
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
else:
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
|
||||
k_cache = kv_cache["k"][:, max(0, local_end_index - max_attention_size) : local_end_index]
|
||||
v_cache = kv_cache["v"][:, max(0, local_end_index - max_attention_size) : local_end_index]
|
||||
x = attention(roped_query, k_cache, v_cache)
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
if sp_size > 1:
|
||||
sp_pad = padded_seq_len - seq_lens_int
|
||||
if sp_pad > 0:
|
||||
x = torch.cat([x, x.new_zeros(b, sp_pad, x.size(2), d)], dim=1)
|
||||
x = sequence_model_parallel_all_to_all_4D(x, scatter_dim=1, gather_dim=2)
|
||||
return self.o(x.flatten(2))
|
||||
|
||||
class WanCrossAttention(CausalWanSelfAttention):
|
||||
"""LingBot World 2 cross-attention with reusable text K/V cache."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
context_lens: torch.Tensor | None,
|
||||
crossattn_cache: dict | None = None,
|
||||
cross_attn_first_call: bool | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Attend hidden states to text context, populating cache on first use."""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
if crossattn_cache is not None:
|
||||
is_first = crossattn_cache["is_init"].item() == 0 if cross_attn_first_call is None else cross_attn_first_call
|
||||
if is_first:
|
||||
crossattn_cache["is_init"].fill_(1)
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
crossattn_cache["k"].copy_(k)
|
||||
crossattn_cache["v"].copy_(v)
|
||||
else:
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
x = flash_attention(q, k, v, k_lens=context_lens)
|
||||
return self.o(x.flatten(2))
|
||||
|
||||
|
||||
class CausalWanAttentionBlock(nn.Module):
|
||||
"""One LingBot World 2 causal transformer block including camera injection."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
qk_norm: bool = True,
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
self.norm1 = WanLayerNorm(dim, eps)
|
||||
self.self_attn = CausalWanSelfAttention(dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
|
||||
self.norm3 = WanLayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WanCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
|
||||
self.norm2 = WanLayerNorm(dim, eps)
|
||||
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(approximate="tanh"), nn.Linear(ffn_dim, dim))
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
self.cam_injector_layer1 = nn.Linear(dim, dim)
|
||||
self.cam_injector_layer2 = nn.Linear(dim, dim)
|
||||
self.cam_scale_layer = nn.Linear(dim, dim)
|
||||
self.cam_shift_layer = nn.Linear(dim, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
e: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
grid_sizes: torch.Tensor,
|
||||
freqs: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
context_lens: torch.Tensor | None,
|
||||
dit_cond_dict: dict[str, Any] | None = None,
|
||||
kv_cache: dict | None = None,
|
||||
crossattn_cache: dict | None = None,
|
||||
current_start: int = 0,
|
||||
max_attention_size: int = 1_000_000,
|
||||
frame_seqlen: int | None = None,
|
||||
cross_attn_first_call: bool | None = None,
|
||||
seq_lens_int: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply self-attention, camera modulation, cross-attention, and FFN."""
|
||||
assert kv_cache is not None
|
||||
assert e.dtype == torch.float32
|
||||
with torch.amp.autocast("cuda", dtype=torch.float32):
|
||||
e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
|
||||
|
||||
y = self.self_attn(
|
||||
self.norm1(x).float() * (1 + e[1].squeeze(2)) + e[0].squeeze(2),
|
||||
seq_lens,
|
||||
grid_sizes,
|
||||
freqs,
|
||||
kv_cache,
|
||||
current_start,
|
||||
max_attention_size,
|
||||
frame_seqlen=frame_seqlen,
|
||||
seq_lens_int=seq_lens_int,
|
||||
)
|
||||
with torch.amp.autocast("cuda", dtype=torch.float32):
|
||||
x = x + y * e[2].squeeze(2)
|
||||
|
||||
if dit_cond_dict is not None and "c2ws_plucker_emb" in dit_cond_dict:
|
||||
c2ws_plucker_emb = dit_cond_dict["c2ws_plucker_emb"]
|
||||
c2ws_hidden_states = self.cam_injector_layer2(
|
||||
F.silu(self.cam_injector_layer1(c2ws_plucker_emb))
|
||||
)
|
||||
c2ws_hidden_states = c2ws_hidden_states + c2ws_plucker_emb
|
||||
x = (1.0 + self.cam_scale_layer(c2ws_hidden_states)) * x + self.cam_shift_layer(c2ws_hidden_states)
|
||||
|
||||
x = x + self.cross_attn(
|
||||
self.norm3(x),
|
||||
context,
|
||||
context_lens,
|
||||
crossattn_cache=crossattn_cache,
|
||||
cross_attn_first_call=cross_attn_first_call,
|
||||
)
|
||||
y = self.ffn(self.norm2(x).float() * (1 + e[4].squeeze(2)) + e[3].squeeze(2))
|
||||
with torch.amp.autocast("cuda", dtype=torch.float32):
|
||||
x = x + y * e[5].squeeze(2)
|
||||
return x
|
||||
|
||||
|
||||
class CausalHead(nn.Module):
|
||||
"""Output projection head for LingBot World 2 causal-fast DiT."""
|
||||
|
||||
def __init__(self, dim: int, out_dim: int, patch_size: tuple[int, int, int], eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
self.norm = WanLayerNorm(dim, eps)
|
||||
self.head = nn.Linear(dim, math.prod(patch_size) * out_dim)
|
||||
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x: torch.Tensor, e: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize, modulate, and project hidden states to latent patches."""
|
||||
assert e.dtype == torch.float32
|
||||
with torch.amp.autocast("cuda", dtype=torch.float32):
|
||||
e = (self.modulation.unsqueeze(0) + e.unsqueeze(2)).chunk(2, dim=2)
|
||||
x = self.head(self.norm(x) * (1 + e[1].squeeze(2)) + e[0].squeeze(2))
|
||||
return x
|
||||
|
||||
|
||||
class LingBotWorld2CausalFastTransformer3DModel(BaseDiT):
|
||||
"""Released LingBot World 2 14B causal-fast model with native FastVideo loading."""
|
||||
|
||||
_fsdp_shard_conditions = [is_blocks]
|
||||
_compile_conditions: list = []
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
param_names_mapping: dict = {}
|
||||
reverse_param_names_mapping: dict = {}
|
||||
lora_param_names_mapping: dict = {}
|
||||
|
||||
def __init__(self, config: LingBotWorld2CausalFastVideoConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.model_type = config.model_type
|
||||
self.patch_size = tuple(config.patch_size)
|
||||
self.text_len = config.text_len
|
||||
self.in_dim = config.in_dim
|
||||
self.dim = config.dim
|
||||
self.hidden_size = config.dim
|
||||
self.ffn_dim = config.ffn_dim
|
||||
self.freq_dim = config.freq_dim
|
||||
self.text_dim = config.text_dim
|
||||
self.out_dim = config.out_dim
|
||||
self.out_channels = config.out_dim
|
||||
self.num_heads = config.num_heads
|
||||
self.num_attention_heads = config.num_heads
|
||||
self.attention_head_dim = config.dim // config.num_heads
|
||||
self.num_layers = config.num_layers
|
||||
self.local_attn_size = config.local_attn_size
|
||||
self.sink_size = config.sink_size
|
||||
self.qk_norm = config.qk_norm
|
||||
self.cross_attn_norm = config.cross_attn_norm
|
||||
self.eps = config.eps
|
||||
self.num_channels_latents = config.out_dim
|
||||
|
||||
control_dim = 6
|
||||
self.patch_embedding = nn.Conv3d(self.in_dim, self.dim, kernel_size=self.patch_size, stride=self.patch_size)
|
||||
self.patch_embedding_wancamctrl = nn.Linear(
|
||||
control_dim * 64 * self.patch_size[0] * self.patch_size[1] * self.patch_size[2],
|
||||
self.dim,
|
||||
)
|
||||
self.c2ws_hidden_states_layer1 = nn.Linear(self.dim, self.dim)
|
||||
self.c2ws_hidden_states_layer2 = nn.Linear(self.dim, self.dim)
|
||||
self.text_embedding = nn.Sequential(
|
||||
nn.Linear(self.text_dim, self.dim),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(self.dim, self.dim),
|
||||
)
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(self.freq_dim, self.dim),
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.dim, self.dim),
|
||||
)
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(self.dim, self.dim * 6))
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
CausalWanAttentionBlock(
|
||||
self.dim,
|
||||
self.ffn_dim,
|
||||
self.num_heads,
|
||||
self.local_attn_size,
|
||||
self.sink_size,
|
||||
self.qk_norm,
|
||||
self.cross_attn_norm,
|
||||
self.eps,
|
||||
)
|
||||
for _ in range(self.num_layers)
|
||||
]
|
||||
)
|
||||
self.head = CausalHead(self.dim, self.out_dim, self.patch_size, self.eps)
|
||||
self.freqs: torch.Tensor | None = None
|
||||
self.init_weights()
|
||||
self.__post_init__()
|
||||
|
||||
def _get_freqs(self, device: torch.device) -> torch.Tensor:
|
||||
"""Materialize the non-persistent RoPE frequency table outside meta init."""
|
||||
if self.freqs is None or self.freqs.is_meta or self.freqs.device != device:
|
||||
d = self.dim // self.num_heads
|
||||
self.freqs = torch.cat(
|
||||
[
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
],
|
||||
dim=1,
|
||||
).to(device)
|
||||
return self.freqs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
timestep: torch.Tensor | None = None,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
guidance=None,
|
||||
*,
|
||||
x: list[torch.Tensor] | None = None,
|
||||
t: torch.Tensor | None = None,
|
||||
context: list[torch.Tensor] | torch.Tensor | None = None,
|
||||
seq_len: int | None = None,
|
||||
y: list[torch.Tensor] | None = None,
|
||||
dit_cond_dict: dict[str, Any] | None = None,
|
||||
kv_cache: list[dict] | None = None,
|
||||
crossattn_cache: list[dict] | None = None,
|
||||
current_start: int = 0,
|
||||
max_attention_size: int = 1_000_000,
|
||||
frame_seqlen: int | None = None,
|
||||
cross_attn_first_call: bool | None = None,
|
||||
**kwargs,
|
||||
) -> list[torch.Tensor]:
|
||||
"""Run one cached causal-fast DiT forward using the released LingBot World 2 ABI."""
|
||||
del encoder_hidden_states_image, guidance, kwargs
|
||||
if x is None:
|
||||
assert isinstance(hidden_states, torch.Tensor)
|
||||
x = [hidden_states[0]]
|
||||
if t is None:
|
||||
assert timestep is not None
|
||||
t = timestep
|
||||
if context is None:
|
||||
context = encoder_hidden_states
|
||||
if isinstance(context, torch.Tensor):
|
||||
context = [u for u in context]
|
||||
assert context is not None
|
||||
assert seq_len is not None
|
||||
assert kv_cache is not None
|
||||
assert crossattn_cache is not None
|
||||
if self.model_type == "i2v":
|
||||
assert y is not None
|
||||
|
||||
device = self.patch_embedding.weight.device
|
||||
freqs = self._get_freqs(device)
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
|
||||
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long, device=u.device) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long, device=device)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat(x)
|
||||
seq_lens_int = int(seq_lens[0].item())
|
||||
sp_size = get_sp_world_size()
|
||||
sp_rank = get_sp_parallel_rank()
|
||||
padded_seq_len = ((seq_lens_int + sp_size - 1) // sp_size) * sp_size
|
||||
sp_pad_len = padded_seq_len - seq_lens_int
|
||||
if sp_pad_len > 0:
|
||||
x = torch.cat([x, x.new_zeros(x.size(0), sp_pad_len, x.size(2))], dim=1)
|
||||
|
||||
if t.dim() == 1:
|
||||
t = t.expand(t.size(0), padded_seq_len)
|
||||
with torch.amp.autocast("cuda", dtype=torch.float32):
|
||||
bt = t.size(0)
|
||||
t = t.flatten()
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t).unflatten(0, (bt, padded_seq_len)).float()
|
||||
)
|
||||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context])
|
||||
)
|
||||
|
||||
if dit_cond_dict is not None and "c2ws_plucker_emb" in dit_cond_dict:
|
||||
c2ws_plucker_emb = dit_cond_dict["c2ws_plucker_emb"]
|
||||
c2ws_plucker_emb = [
|
||||
rearrange(
|
||||
i,
|
||||
"1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)",
|
||||
c1=self.patch_size[0],
|
||||
c2=self.patch_size[1],
|
||||
c3=self.patch_size[2],
|
||||
)
|
||||
for i in c2ws_plucker_emb
|
||||
]
|
||||
c2ws_plucker_emb = torch.cat(c2ws_plucker_emb, dim=1)
|
||||
c2ws_plucker_emb = self.patch_embedding_wancamctrl(c2ws_plucker_emb)
|
||||
c2ws_hidden_states = self.c2ws_hidden_states_layer2(
|
||||
F.silu(self.c2ws_hidden_states_layer1(c2ws_plucker_emb))
|
||||
)
|
||||
c2ws_plucker_emb = c2ws_plucker_emb + c2ws_hidden_states
|
||||
cam_len = c2ws_plucker_emb.size(1)
|
||||
if cam_len < padded_seq_len:
|
||||
c2ws_plucker_emb = torch.cat(
|
||||
[
|
||||
c2ws_plucker_emb,
|
||||
c2ws_plucker_emb.new_zeros(
|
||||
c2ws_plucker_emb.size(0),
|
||||
padded_seq_len - cam_len,
|
||||
c2ws_plucker_emb.size(2),
|
||||
),
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
elif cam_len > padded_seq_len:
|
||||
c2ws_plucker_emb = c2ws_plucker_emb[:, :padded_seq_len, :]
|
||||
if sp_size > 1:
|
||||
c2ws_plucker_emb = torch.chunk(c2ws_plucker_emb, sp_size, dim=1)[sp_rank]
|
||||
dit_cond_dict = dict(dit_cond_dict)
|
||||
dit_cond_dict["c2ws_plucker_emb"] = c2ws_plucker_emb
|
||||
|
||||
if sp_size > 1:
|
||||
x = torch.chunk(x, sp_size, dim=1)[sp_rank]
|
||||
e = torch.chunk(e, sp_size, dim=1)[sp_rank]
|
||||
e0 = torch.chunk(e0, sp_size, dim=1)[sp_rank]
|
||||
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
x = block(
|
||||
x,
|
||||
e=e0,
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
freqs=freqs,
|
||||
context=context,
|
||||
context_lens=context_lens,
|
||||
dit_cond_dict=dit_cond_dict,
|
||||
kv_cache=kv_cache[block_index],
|
||||
crossattn_cache=crossattn_cache[block_index],
|
||||
current_start=current_start,
|
||||
max_attention_size=max_attention_size,
|
||||
frame_seqlen=frame_seqlen,
|
||||
cross_attn_first_call=cross_attn_first_call,
|
||||
seq_lens_int=seq_lens_int,
|
||||
)
|
||||
|
||||
x = self.head(x, e)
|
||||
if sp_size > 1:
|
||||
x = sequence_model_parallel_all_gather(x, dim=1)
|
||||
return [u.float() for u in self.unpatchify(x, grid_sizes)]
|
||||
|
||||
def unpatchify(self, x: torch.Tensor, grid_sizes: torch.Tensor) -> list[torch.Tensor]:
|
||||
"""Reconstruct latent videos from flattened patch tokens."""
|
||||
c = self.out_dim
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist(), strict=True):
|
||||
u = u[: math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = torch.einsum("fhwpqrc->cfphqwr", u)
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
def init_weights(self) -> None:
|
||||
"""Initialize modules for non-meta construction; checkpoint load overwrites them."""
|
||||
if self.patch_embedding.weight.is_meta:
|
||||
return
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
||||
for m in self.text_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=0.02)
|
||||
for m in self.time_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=0.02)
|
||||
nn.init.zeros_(self.head.head.weight)
|
||||
|
||||
|
||||
EntryClass = LingBotWorld2CausalFastTransformer3DModel
|
||||
@@ -0,0 +1,570 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native Z-Image transformer.
|
||||
|
||||
Z-Image attends over padded variable-length image/text streams and requires a
|
||||
key-padding mask. FastVideo's distributed attention wrappers do not expose that
|
||||
mask contract yet, so this implementation uses torch SDPA and is SP=1 only.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
|
||||
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _linear(layer: ReplicatedLinear, x: torch.Tensor) -> torch.Tensor:
|
||||
return layer(x)[0]
|
||||
|
||||
|
||||
def _prepare_attention_mask(attention_mask: torch.Tensor | None, dtype: torch.dtype) -> torch.Tensor | None:
|
||||
if attention_mask is None:
|
||||
return None
|
||||
if attention_mask.ndim == 2:
|
||||
attention_mask = attention_mask[:, None, None, :]
|
||||
if attention_mask.dtype == torch.bool:
|
||||
additive_mask = torch.zeros_like(attention_mask, dtype=dtype)
|
||||
additive_mask.masked_fill_(~attention_mask, float("-inf"))
|
||||
return additive_mask
|
||||
return attention_mask
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
out_size: int,
|
||||
mid_size: int,
|
||||
frequency_embedding_size: int,
|
||||
max_period: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.mlp = nn.ModuleList([
|
||||
ReplicatedLinear(frequency_embedding_size, mid_size, bias=True),
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(mid_size, out_size, bias=True),
|
||||
])
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.max_period = max_period
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t: torch.Tensor, dim: int, max_period: int) -> torch.Tensor:
|
||||
with torch.amp.autocast("cuda", enabled=False):
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) * torch.arange(half, dtype=torch.float32, device=t.device) / half)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t: torch.Tensor) -> torch.Tensor:
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size, self.max_period)
|
||||
weight_dtype = self.mlp[0].weight.dtype
|
||||
if weight_dtype.is_floating_point:
|
||||
t_freq = t_freq.to(weight_dtype)
|
||||
return _linear(self.mlp[2], self.mlp[1](_linear(self.mlp[0], t_freq)))
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-5) -> None:
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
output = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, hidden_dim: int) -> None:
|
||||
super().__init__()
|
||||
self.w1 = ReplicatedLinear(dim, hidden_dim, bias=False)
|
||||
self.w2 = ReplicatedLinear(hidden_dim, dim, bias=False)
|
||||
self.w3 = ReplicatedLinear(dim, hidden_dim, bias=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return _linear(self.w2, F.silu(_linear(self.w1, x)) * _linear(self.w3, x))
|
||||
|
||||
|
||||
def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||
with torch.amp.autocast("cuda", enabled=False):
|
||||
x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2))
|
||||
x_out = torch.view_as_real(x * freqs_cis.unsqueeze(2)).flatten(3)
|
||||
return x_out.type_as(x_in)
|
||||
|
||||
|
||||
class ZImageAttention(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, n_heads: int, n_kv_heads: int, qk_norm: bool = True, eps: float = 1e-5) -> None:
|
||||
super().__init__()
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.head_dim = dim // n_heads
|
||||
|
||||
self.to_q = ReplicatedLinear(dim, n_heads * self.head_dim, bias=False)
|
||||
self.to_k = ReplicatedLinear(dim, n_kv_heads * self.head_dim, bias=False)
|
||||
self.to_v = ReplicatedLinear(dim, n_kv_heads * self.head_dim, bias=False)
|
||||
self.to_out = nn.ModuleList([ReplicatedLinear(n_heads * self.head_dim, dim, bias=False)])
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else None
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
freqs_cis: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
query = _linear(self.to_q, hidden_states).unflatten(-1, (self.n_heads, -1))
|
||||
key = _linear(self.to_k, hidden_states).unflatten(-1, (self.n_kv_heads, -1))
|
||||
value = _linear(self.to_v, hidden_states).unflatten(-1, (self.n_kv_heads, -1))
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
if freqs_cis is not None:
|
||||
query = apply_rotary_emb(query, freqs_cis)
|
||||
key = apply_rotary_emb(key, freqs_cis)
|
||||
|
||||
mask = _prepare_attention_mask(attention_mask, query.dtype)
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query.transpose(1, 2),
|
||||
key.transpose(1, 2),
|
||||
value.transpose(1, 2),
|
||||
attn_mask=mask,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
).transpose(1, 2).contiguous()
|
||||
return _linear(self.to_out[0], hidden_states.flatten(2, 3).to(query.dtype))
|
||||
|
||||
|
||||
class ZImageTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
layer_id: int,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
norm_eps: float,
|
||||
qk_norm: bool,
|
||||
adaln_embed_dim: int,
|
||||
modulation: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.head_dim = dim // n_heads
|
||||
self.layer_id = layer_id
|
||||
self.modulation = modulation
|
||||
|
||||
self.attention = ZImageAttention(dim, n_heads, n_kv_heads, qk_norm, norm_eps)
|
||||
self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8))
|
||||
self.attention_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
self.attention_norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
if modulation:
|
||||
self.adaLN_modulation = nn.ModuleList(
|
||||
[ReplicatedLinear(min(dim, adaln_embed_dim), 4 * dim, bias=True)])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
attn_mask: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
adaln_input: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if self.modulation:
|
||||
assert adaln_input is not None
|
||||
scale_msa, gate_msa, scale_mlp, gate_mlp = _linear(
|
||||
self.adaLN_modulation[0], adaln_input).unsqueeze(1).chunk(4, dim=2)
|
||||
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
|
||||
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
|
||||
|
||||
attn_out = self.attention(
|
||||
self.attention_norm1(x) * scale_msa,
|
||||
attention_mask=attn_mask,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
x = x + gate_msa * self.attention_norm2(attn_out)
|
||||
x = x + gate_mlp * self.ffn_norm2(self.feed_forward(self.ffn_norm1(x) * scale_mlp))
|
||||
else:
|
||||
attn_out = self.attention(
|
||||
self.attention_norm1(x),
|
||||
attention_mask=attn_mask,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
x = x + self.attention_norm2(attn_out)
|
||||
x = x + self.ffn_norm2(self.feed_forward(self.ffn_norm1(x)))
|
||||
return x
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
|
||||
def __init__(self, hidden_size: int, out_channels: int, adaln_embed_dim: int) -> None:
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = ReplicatedLinear(hidden_size, out_channels, bias=True)
|
||||
self.adaLN_modulation = nn.ModuleList([
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(min(hidden_size, adaln_embed_dim), hidden_size, bias=True),
|
||||
])
|
||||
|
||||
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
||||
scale = 1.0 + _linear(self.adaLN_modulation[1], self.adaLN_modulation[0](c))
|
||||
return _linear(self.linear, self.norm_final(x) * scale.unsqueeze(1))
|
||||
|
||||
|
||||
class RopeEmbedder:
|
||||
|
||||
def __init__(self, theta: float, axes_dims: tuple[int, ...], axes_lens: tuple[int, ...]) -> None:
|
||||
if len(axes_dims) != len(axes_lens):
|
||||
raise ValueError("RoPE axes require matching dimensions and lengths")
|
||||
self.theta = theta
|
||||
self.axes_dims = axes_dims
|
||||
self.axes_lens = axes_lens
|
||||
self.freqs_cis: list[torch.Tensor] | None = None
|
||||
|
||||
@staticmethod
|
||||
def precompute_freqs_cis(dim: tuple[int, ...], end: tuple[int, ...], theta: float) -> list[torch.Tensor]:
|
||||
with torch.device("cpu"):
|
||||
freqs_cis = []
|
||||
for axis_dim, axis_end in zip(dim, end):
|
||||
freqs = 1.0 / (theta**(torch.arange(0, axis_dim, 2, dtype=torch.float64) / axis_dim))
|
||||
timestep = torch.arange(axis_end, dtype=torch.float64)
|
||||
angles = torch.outer(timestep, freqs).float()
|
||||
freqs_cis.append(torch.polar(torch.ones_like(angles), angles).to(torch.complex64))
|
||||
return freqs_cis
|
||||
|
||||
def __call__(self, ids: torch.Tensor) -> torch.Tensor:
|
||||
if ids.ndim != 2 or ids.shape[-1] != len(self.axes_dims):
|
||||
raise ValueError("RoPE ids must have shape [sequence, number_of_axes]")
|
||||
if self.freqs_cis is None:
|
||||
self.freqs_cis = [
|
||||
freqs.to(ids.device)
|
||||
for freqs in self.precompute_freqs_cis(self.axes_dims, self.axes_lens, self.theta)
|
||||
]
|
||||
elif self.freqs_cis[0].device != ids.device:
|
||||
self.freqs_cis = [freqs.to(ids.device) for freqs in self.freqs_cis]
|
||||
return torch.cat([self.freqs_cis[i][ids[:, i]] for i in range(len(self.axes_dims))], dim=-1)
|
||||
|
||||
|
||||
class ZImageTransformer2DModel(BaseDiT):
|
||||
_default_config = ZImageDiTConfig()
|
||||
_fsdp_shard_conditions = _default_config.arch_config._fsdp_shard_conditions
|
||||
_compile_conditions = _default_config.arch_config._compile_conditions
|
||||
_supported_attention_backends = (AttentionBackendEnum.TORCH_SDPA, )
|
||||
param_names_mapping = _default_config.arch_config.param_names_mapping
|
||||
reverse_param_names_mapping = _default_config.arch_config.reverse_param_names_mapping
|
||||
|
||||
def __init__(self, config: ZImageDiTConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
arch = config.arch_config
|
||||
|
||||
self.in_channels = arch.in_channels
|
||||
self.out_channels = arch.in_channels
|
||||
self.all_patch_size = tuple(arch.all_patch_size)
|
||||
self.all_f_patch_size = tuple(arch.all_f_patch_size)
|
||||
self.dim = arch.dim
|
||||
self.n_heads = arch.n_heads
|
||||
self.rope_theta = arch.rope_theta
|
||||
self.t_scale = arch.t_scale
|
||||
self.seq_multi_of = arch.seq_multi_of
|
||||
|
||||
self.hidden_size = arch.dim
|
||||
self.num_attention_heads = arch.n_heads
|
||||
self.num_channels_latents = arch.in_channels
|
||||
|
||||
self.all_x_embedder = nn.ModuleDict({
|
||||
f"{patch_size}-{f_patch_size}": ReplicatedLinear(
|
||||
f_patch_size * patch_size * patch_size * arch.in_channels, arch.dim, bias=True)
|
||||
for patch_size, f_patch_size in zip(self.all_patch_size, self.all_f_patch_size)
|
||||
})
|
||||
self.all_final_layer = nn.ModuleDict({
|
||||
f"{patch_size}-{f_patch_size}": FinalLayer(
|
||||
arch.dim,
|
||||
patch_size * patch_size * f_patch_size * self.out_channels,
|
||||
arch.adaln_embed_dim,
|
||||
)
|
||||
for patch_size, f_patch_size in zip(self.all_patch_size, self.all_f_patch_size)
|
||||
})
|
||||
|
||||
block_kwargs = {
|
||||
"dim": arch.dim,
|
||||
"n_heads": arch.n_heads,
|
||||
"n_kv_heads": arch.n_kv_heads,
|
||||
"norm_eps": arch.norm_eps,
|
||||
"qk_norm": arch.qk_norm,
|
||||
"adaln_embed_dim": arch.adaln_embed_dim,
|
||||
}
|
||||
self.noise_refiner = nn.ModuleList([
|
||||
ZImageTransformerBlock(1000 + layer_id, modulation=True, **block_kwargs)
|
||||
for layer_id in range(arch.n_refiner_layers)
|
||||
])
|
||||
self.context_refiner = nn.ModuleList([
|
||||
ZImageTransformerBlock(layer_id, modulation=False, **block_kwargs)
|
||||
for layer_id in range(arch.n_refiner_layers)
|
||||
])
|
||||
self.t_embedder = TimestepEmbedder(
|
||||
min(arch.dim, arch.adaln_embed_dim),
|
||||
mid_size=arch.timestep_mid_size,
|
||||
frequency_embedding_size=arch.frequency_embedding_size,
|
||||
max_period=arch.max_period,
|
||||
)
|
||||
self.cap_embedder = nn.ModuleList([
|
||||
RMSNorm(arch.cap_feat_dim, eps=arch.norm_eps),
|
||||
ReplicatedLinear(arch.cap_feat_dim, arch.dim, bias=True),
|
||||
])
|
||||
self.x_pad_token = nn.Parameter(torch.empty((1, arch.dim)))
|
||||
self.cap_pad_token = nn.Parameter(torch.empty((1, arch.dim)))
|
||||
self.layers = nn.ModuleList([
|
||||
ZImageTransformerBlock(layer_id, modulation=True, **block_kwargs) for layer_id in range(arch.n_layers)
|
||||
])
|
||||
self.axes_dims = tuple(arch.axes_dims)
|
||||
self.axes_lens = tuple(arch.axes_lens)
|
||||
self.rope_embedder = RopeEmbedder(arch.rope_theta, self.axes_dims, self.axes_lens)
|
||||
self.__post_init__()
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
x: list[torch.Tensor],
|
||||
size: list[tuple[int, int, int]],
|
||||
patch_size: int,
|
||||
f_patch_size: int,
|
||||
) -> list[torch.Tensor]:
|
||||
patch_height = patch_width = patch_size
|
||||
patch_frames = f_patch_size
|
||||
if len(x) != len(size):
|
||||
raise ValueError("output batch and original sizes must have equal length")
|
||||
for i, (frames, height, width) in enumerate(size):
|
||||
original_length = (frames // patch_frames) * (height // patch_height) * (width // patch_width)
|
||||
x[i] = (x[i][:original_length].view(
|
||||
frames // patch_frames,
|
||||
height // patch_height,
|
||||
width // patch_width,
|
||||
patch_frames,
|
||||
patch_height,
|
||||
patch_width,
|
||||
self.out_channels,
|
||||
).permute(6, 0, 3, 1, 4, 2, 5).reshape(self.out_channels, frames, height, width))
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def create_coordinate_grid(
|
||||
size: tuple[int, int, int],
|
||||
start: tuple[int, int, int] | None = None,
|
||||
device: torch.device | None = None,
|
||||
) -> torch.Tensor:
|
||||
start = start or (0, ) * len(size)
|
||||
axes = [
|
||||
torch.arange(axis_start, axis_start + span, dtype=torch.int32, device=device)
|
||||
for axis_start, span in zip(start, size)
|
||||
]
|
||||
return torch.stack(torch.meshgrid(axes, indexing="ij"), dim=-1)
|
||||
|
||||
def patchify_and_embed(
|
||||
self,
|
||||
all_image: list[torch.Tensor],
|
||||
all_cap_feats: list[torch.Tensor],
|
||||
patch_size: int,
|
||||
f_patch_size: int,
|
||||
) -> tuple[
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
list[tuple[int, int, int]],
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
]:
|
||||
patch_height = patch_width = patch_size
|
||||
patch_frames = f_patch_size
|
||||
device = all_image[0].device
|
||||
|
||||
image_out = []
|
||||
image_sizes = []
|
||||
image_pos_ids = []
|
||||
image_pad_masks = []
|
||||
cap_pos_ids = []
|
||||
cap_pad_masks = []
|
||||
cap_feats_out = []
|
||||
|
||||
for image, cap_feat in zip(all_image, all_cap_feats):
|
||||
cap_length = len(cap_feat)
|
||||
cap_padding = (-cap_length) % self.seq_multi_of
|
||||
cap_pos_ids.append(
|
||||
self.create_coordinate_grid(
|
||||
(cap_length + cap_padding, 1, 1),
|
||||
start=(1, 0, 0),
|
||||
device=device,
|
||||
).flatten(0, 2))
|
||||
cap_pad_masks.append(
|
||||
torch.cat([
|
||||
torch.zeros(cap_length, dtype=torch.bool, device=device),
|
||||
torch.ones(cap_padding, dtype=torch.bool, device=device),
|
||||
]) if cap_padding else torch.zeros(cap_length, dtype=torch.bool, device=device))
|
||||
cap_feats_out.append(
|
||||
torch.cat([cap_feat, cap_feat[-1:].repeat(cap_padding, 1)]) if cap_padding else cap_feat)
|
||||
|
||||
channels, frames, height, width = image.size()
|
||||
image_sizes.append((frames, height, width))
|
||||
frame_tokens, height_tokens, width_tokens = (
|
||||
frames // patch_frames,
|
||||
height // patch_height,
|
||||
width // patch_width,
|
||||
)
|
||||
image = image.view(
|
||||
channels,
|
||||
frame_tokens,
|
||||
patch_frames,
|
||||
height_tokens,
|
||||
patch_height,
|
||||
width_tokens,
|
||||
patch_width,
|
||||
)
|
||||
image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(
|
||||
frame_tokens * height_tokens * width_tokens,
|
||||
patch_frames * patch_height * patch_width * channels,
|
||||
)
|
||||
|
||||
image_length = len(image)
|
||||
image_padding = (-image_length) % self.seq_multi_of
|
||||
original_pos_ids = self.create_coordinate_grid(
|
||||
(frame_tokens, height_tokens, width_tokens),
|
||||
start=(cap_length + cap_padding + 1, 0, 0),
|
||||
device=device,
|
||||
).flatten(0, 2)
|
||||
if image_padding:
|
||||
padding_pos_ids = self.create_coordinate_grid((1, 1, 1), device=device).flatten(0, 2).repeat(
|
||||
image_padding, 1)
|
||||
image_pos_ids.append(torch.cat([original_pos_ids, padding_pos_ids]))
|
||||
else:
|
||||
image_pos_ids.append(original_pos_ids)
|
||||
image_pad_masks.append(
|
||||
torch.cat([
|
||||
torch.zeros(image_length, dtype=torch.bool, device=device),
|
||||
torch.ones(image_padding, dtype=torch.bool, device=device),
|
||||
]) if image_padding else torch.zeros(image_length, dtype=torch.bool, device=device))
|
||||
image_out.append(
|
||||
torch.cat([image, image[-1:].repeat(image_padding, 1)]) if image_padding else image)
|
||||
|
||||
return (
|
||||
image_out,
|
||||
cap_feats_out,
|
||||
image_sizes,
|
||||
image_pos_ids,
|
||||
cap_pos_ids,
|
||||
image_pad_masks,
|
||||
cap_pad_masks,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _attention_mask(lengths: list[int], device: torch.device) -> torch.Tensor:
|
||||
mask = torch.zeros((len(lengths), max(lengths)), dtype=torch.bool, device=device)
|
||||
for i, length in enumerate(lengths):
|
||||
mask[i, :length] = True
|
||||
return mask
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
guidance=None,
|
||||
patch_size: int = 2,
|
||||
f_patch_size: int = 1,
|
||||
**kwargs,
|
||||
) -> tuple[list[torch.Tensor], dict]:
|
||||
del encoder_hidden_states_image, guidance, kwargs
|
||||
if model_parallel_is_initialized() and get_sp_world_size() != 1:
|
||||
raise NotImplementedError(
|
||||
"Z-Image masked SDPA does not support sequence parallelism; run with sp_size=1")
|
||||
if patch_size not in self.all_patch_size or f_patch_size not in self.all_f_patch_size:
|
||||
raise ValueError(f"unsupported patch sizes: spatial={patch_size}, temporal={f_patch_size}")
|
||||
if isinstance(hidden_states, torch.Tensor):
|
||||
hidden_states = list(hidden_states.unbind(0))
|
||||
if isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = list(encoder_hidden_states.unbind(0))
|
||||
|
||||
device = hidden_states[0].device
|
||||
timestep_embedding = self.t_embedder(timestep * self.t_scale)
|
||||
(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
image_sizes,
|
||||
image_pos_ids,
|
||||
cap_pos_ids,
|
||||
image_inner_pad_masks,
|
||||
cap_inner_pad_masks,
|
||||
) = self.patchify_and_embed(hidden_states, encoder_hidden_states, patch_size, f_patch_size)
|
||||
|
||||
image_lengths = [len(item) for item in hidden_states]
|
||||
if not all(length % self.seq_multi_of == 0 for length in image_lengths):
|
||||
raise ValueError("padded image sequence lengths must be aligned")
|
||||
hidden_states = torch.cat(hidden_states)
|
||||
hidden_states = _linear(self.all_x_embedder[f"{patch_size}-{f_patch_size}"], hidden_states)
|
||||
adaln_input = timestep_embedding.type_as(hidden_states)
|
||||
hidden_states[torch.cat(image_inner_pad_masks)] = self.x_pad_token
|
||||
hidden_states = list(hidden_states.split(image_lengths))
|
||||
image_freqs_cis = list(
|
||||
self.rope_embedder(torch.cat(image_pos_ids)).split([len(item) for item in image_pos_ids]))
|
||||
hidden_states = pad_sequence(hidden_states, batch_first=True, padding_value=0.0)
|
||||
image_freqs_cis = pad_sequence(image_freqs_cis, batch_first=True, padding_value=0.0)
|
||||
image_freqs_cis = image_freqs_cis[:, :hidden_states.shape[1]]
|
||||
image_attn_mask = self._attention_mask(image_lengths, device)
|
||||
for layer in self.noise_refiner:
|
||||
hidden_states = layer(hidden_states, image_attn_mask, image_freqs_cis, adaln_input)
|
||||
|
||||
cap_lengths = [len(item) for item in encoder_hidden_states]
|
||||
if not all(length % self.seq_multi_of == 0 for length in cap_lengths):
|
||||
raise ValueError("padded caption sequence lengths must be aligned")
|
||||
encoder_hidden_states = torch.cat(encoder_hidden_states)
|
||||
encoder_hidden_states = _linear(self.cap_embedder[1], self.cap_embedder[0](encoder_hidden_states))
|
||||
encoder_hidden_states[torch.cat(cap_inner_pad_masks)] = self.cap_pad_token
|
||||
encoder_hidden_states = list(encoder_hidden_states.split(cap_lengths))
|
||||
cap_freqs_cis = list(self.rope_embedder(torch.cat(cap_pos_ids)).split([len(item) for item in cap_pos_ids]))
|
||||
encoder_hidden_states = pad_sequence(encoder_hidden_states, batch_first=True, padding_value=0.0)
|
||||
cap_freqs_cis = pad_sequence(cap_freqs_cis, batch_first=True, padding_value=0.0)
|
||||
cap_freqs_cis = cap_freqs_cis[:, :encoder_hidden_states.shape[1]]
|
||||
cap_attn_mask = self._attention_mask(cap_lengths, device)
|
||||
for layer in self.context_refiner:
|
||||
encoder_hidden_states = layer(encoder_hidden_states, cap_attn_mask, cap_freqs_cis)
|
||||
|
||||
unified = []
|
||||
unified_freqs_cis = []
|
||||
for i, (image_length, cap_length) in enumerate(zip(image_lengths, cap_lengths)):
|
||||
unified.append(
|
||||
torch.cat([hidden_states[i][:image_length], encoder_hidden_states[i][:cap_length]]))
|
||||
unified_freqs_cis.append(
|
||||
torch.cat([image_freqs_cis[i][:image_length], cap_freqs_cis[i][:cap_length]]))
|
||||
unified_lengths = [image_length + cap_length for image_length, cap_length in zip(image_lengths, cap_lengths)]
|
||||
unified = pad_sequence(unified, batch_first=True, padding_value=0.0)
|
||||
unified_freqs_cis = pad_sequence(unified_freqs_cis, batch_first=True, padding_value=0.0)
|
||||
unified_attn_mask = self._attention_mask(unified_lengths, device)
|
||||
for layer in self.layers:
|
||||
unified = layer(unified, unified_attn_mask, unified_freqs_cis, adaln_input)
|
||||
|
||||
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](unified, adaln_input)
|
||||
outputs = self.unpatchify(list(unified.unbind(0)), image_sizes, patch_size, f_patch_size)
|
||||
return outputs, {}
|
||||
|
||||
|
||||
EntryClass = ZImageTransformer2DModel
|
||||
@@ -0,0 +1,221 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Native LingBot-Video Qwen3-VL language model for text-only conditioning."""
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.encoders.qwen3 import (
|
||||
Qwen3Attention,
|
||||
Qwen3DecoderLayer,
|
||||
Qwen3ForCausalLM,
|
||||
Qwen3MLP,
|
||||
)
|
||||
|
||||
|
||||
class LingBotVideoQwen3VLAttention(Qwen3Attention):
|
||||
"""Qwen3-VL attention with the official masked repeat-K/V SDPA path."""
|
||||
|
||||
def _apply_qwen3_vl_rope(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Apply NeoX RoPE with Qwen3-VL's input-dtype multiply-add ordering."""
|
||||
flat_positions = positions.flatten()
|
||||
cos_sin = self.rotary_emb.cos_sin_cache.index_select(0, flat_positions)
|
||||
cos_half, sin_half = cos_sin.chunk(2, dim=-1)
|
||||
cos = torch.cat((cos_half, cos_half), dim=-1).to(query.dtype)
|
||||
sin = torch.cat((sin_half, sin_half), dim=-1).to(query.dtype)
|
||||
if flat_positions.numel() == query.shape[1]:
|
||||
cos = cos.view(1, query.shape[1], 1, self.head_dim)
|
||||
sin = sin.view(1, query.shape[1], 1, self.head_dim)
|
||||
else:
|
||||
cos = cos.view(*query.shape[:2], 1, self.head_dim)
|
||||
sin = sin.view(*query.shape[:2], 1, self.head_dim)
|
||||
|
||||
def rotate(tensor: torch.Tensor) -> torch.Tensor:
|
||||
first, second = tensor.chunk(2, dim=-1)
|
||||
rotated = torch.cat((-second, first), dim=-1)
|
||||
return tensor * cos + rotated * sin
|
||||
|
||||
return rotate(query), rotate(key)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply fused projections, QK norm, RoPE, and causal grouped attention."""
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
query, key, value = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
batch_size, sequence_length = query.shape[:2]
|
||||
query = query.reshape(batch_size, sequence_length, self.num_heads, self.head_dim)
|
||||
key = key.reshape(batch_size, sequence_length, self.num_kv_heads, self.head_dim)
|
||||
value = value.reshape(batch_size, sequence_length, self.num_kv_heads, self.head_dim)
|
||||
query = self.q_norm(query)
|
||||
key = self.k_norm(key)
|
||||
query, key = self._apply_qwen3_vl_rope(positions, query, key)
|
||||
no_padding = attention_mask is None
|
||||
if no_padding:
|
||||
attention_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query.transpose(1, 2),
|
||||
key.transpose(1, 2),
|
||||
value.transpose(1, 2),
|
||||
dropout_p=0.0,
|
||||
is_causal=sequence_length > 1,
|
||||
scale=self.scaling,
|
||||
enable_gqa=self.num_heads != self.num_kv_heads,
|
||||
).transpose(1, 2)
|
||||
else:
|
||||
groups = self.num_heads // self.num_kv_heads
|
||||
key = (key[:, :, :, None, :].expand(-1, -1, -1, groups, -1).reshape(batch_size, sequence_length,
|
||||
self.num_heads, self.head_dim))
|
||||
value = (value[:, :, :, None, :].expand(-1, -1, -1, groups, -1).reshape(batch_size, sequence_length,
|
||||
self.num_heads, self.head_dim))
|
||||
causal_mask = torch.ones(sequence_length, sequence_length, device=query.device, dtype=torch.bool).tril()
|
||||
key_mask = attention_mask.to(device=query.device, dtype=torch.bool)
|
||||
sdpa_mask = causal_mask[None, None, :, :] & key_mask[:, None, None, :]
|
||||
attention_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query.transpose(1, 2),
|
||||
key.transpose(1, 2),
|
||||
value.transpose(1, 2),
|
||||
attn_mask=sdpa_mask,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
).transpose(1, 2)
|
||||
output, _ = self.o_proj(attention_output.reshape(batch_size, sequence_length, -1))
|
||||
return output
|
||||
|
||||
|
||||
class LingBotVideoQwen3VLDecoderLayer(Qwen3DecoderLayer):
|
||||
"""Qwen3-VL decoder layer with explicit official residual rounding order."""
|
||||
|
||||
def __init__(self, config: Any, prefix: str) -> None:
|
||||
"""Build the final Qwen3-VL attention once to avoid orphan parameters."""
|
||||
nn.Module.__init__(self)
|
||||
self.hidden_size = config.hidden_size
|
||||
quant_config = getattr(config, "quant_config", None)
|
||||
self.self_attn = LingBotVideoQwen3VLAttention(
|
||||
config=config,
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=config.num_attention_heads,
|
||||
num_kv_heads=config.num_key_value_heads,
|
||||
rope_theta=config.rope_theta,
|
||||
rope_scaling=config.rope_scaling,
|
||||
max_position_embeddings=config.max_position_embeddings,
|
||||
quant_config=quant_config,
|
||||
bias=config.attention_bias,
|
||||
prefix=f"{prefix}.self_attn",
|
||||
)
|
||||
self.mlp = Qwen3MLP(
|
||||
hidden_size=self.hidden_size,
|
||||
intermediate_size=config.intermediate_size,
|
||||
hidden_act=config.hidden_act,
|
||||
quant_config=quant_config,
|
||||
bias=getattr(config, "mlp_bias", False),
|
||||
prefix=f"{prefix}.mlp",
|
||||
)
|
||||
self.input_layernorm = RMSNorm(self.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = RMSNorm(self.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run attention and MLP with each residual sum rounded before normalization."""
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
hidden_states = self.self_attn(positions, hidden_states, attention_mask)
|
||||
hidden_states = residual + hidden_states
|
||||
residual = hidden_states
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
return residual + hidden_states
|
||||
|
||||
|
||||
class LingBotVideoQwen3VLTextModel(Qwen3ForCausalLM):
|
||||
"""Load the Qwen3-VL language-model subset without its vision tower or LM head."""
|
||||
|
||||
supports_hf_from_pretrained = False
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
"""Construct the exact Qwen3-VL module graph without replacing base layers."""
|
||||
TextEncoder.__init__(self, config)
|
||||
self.quant_config = getattr(config, "quant_config", None)
|
||||
if getattr(config, "lora_config", None) is not None:
|
||||
max_loras = getattr(config.lora_config, "max_loras", 1)
|
||||
lora_vocab_size = getattr(config.lora_config, "lora_extra_vocab_size", 1)
|
||||
lora_vocab = lora_vocab_size * max_loras
|
||||
else:
|
||||
lora_vocab = 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.org_vocab_size = config.vocab_size
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
self.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=self.quant_config,
|
||||
)
|
||||
self.layers = nn.ModuleList(
|
||||
LingBotVideoQwen3VLDecoderLayer(config, prefix=f"{config.prefix}.layers.{index}")
|
||||
for index in range(config.num_hidden_layers))
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None = None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs: Any,
|
||||
) -> BaseEncoderOutput:
|
||||
"""Run explicit Qwen3-VL layers and return the requested hidden-state tuple."""
|
||||
del kwargs
|
||||
output_hidden_states = (output_hidden_states
|
||||
if output_hidden_states is not None else self.config.output_hidden_states)
|
||||
if inputs_embeds is None:
|
||||
if input_ids is None:
|
||||
raise ValueError("input_ids or inputs_embeds is required")
|
||||
hidden_states = self.get_input_embeddings(input_ids)
|
||||
else:
|
||||
hidden_states = inputs_embeds
|
||||
if position_ids is None:
|
||||
position_ids = torch.arange(hidden_states.shape[1], device=hidden_states.device).unsqueeze(0)
|
||||
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
|
||||
attention_mask = None
|
||||
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
|
||||
for layer in self.layers:
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
hidden_states = layer(position_ids, hidden_states, attention_mask)
|
||||
hidden_states = self.norm(hidden_states)
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
hidden_states=all_hidden_states,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
"""Accept either official compound keys or converted native keys."""
|
||||
prefix = "model.language_model."
|
||||
language_weights = ((name[len(prefix):] if name.startswith(prefix) else name, tensor)
|
||||
for name, tensor in weights
|
||||
if name.startswith(prefix) or not name.startswith(("model.", "lm_head.")))
|
||||
return super().load_weights(language_weights)
|
||||
|
||||
|
||||
EntryClass = LingBotVideoQwen3VLTextModel
|
||||
@@ -0,0 +1,269 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot World 2 UMT5 encoder with the released checkpoint's module names."""
|
||||
|
||||
from collections.abc import Iterable
|
||||
import html
|
||||
import math
|
||||
import string
|
||||
|
||||
import ftfy
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5Config
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
|
||||
def basic_clean(text: str) -> str:
|
||||
"""Apply LingBot World 2's ftfy/html cleanup before tokenization."""
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text: str) -> str:
|
||||
"""Collapse all whitespace runs to single spaces."""
|
||||
return " ".join(text.split())
|
||||
|
||||
|
||||
def canonicalize(text: str, keep_punctuation_exact_string: str | None = None) -> str:
|
||||
"""Normalize prompts with LingBot World 2's optional punctuation handling."""
|
||||
text = text.replace("_", " ")
|
||||
if keep_punctuation_exact_string:
|
||||
text = keep_punctuation_exact_string.join(
|
||||
part.translate(str.maketrans("", "", string.punctuation))
|
||||
for part in text.split(keep_punctuation_exact_string)
|
||||
)
|
||||
else:
|
||||
text = text.translate(str.maketrans("", "", string.punctuation))
|
||||
return " ".join(text.lower().split())
|
||||
|
||||
|
||||
def lingbotworld2_whitespace_preprocess(prompt: str) -> str:
|
||||
"""Match the LingBot World 2 source tokenizer's `clean='whitespace'` behavior."""
|
||||
return whitespace_clean(basic_clean(prompt))
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
"""T5 gated GELU approximation used by the source checkpoint."""
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply the tanh GELU approximation."""
|
||||
return 0.5 * x * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
|
||||
|
||||
|
||||
class T5LayerNorm(nn.Module):
|
||||
"""T5 RMS-style layer norm with source-compatible parameter name."""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize in fp32 and apply the learned scale."""
|
||||
x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
if self.weight.dtype in (torch.float16, torch.bfloat16):
|
||||
x = x.type_as(self.weight)
|
||||
return self.weight * x
|
||||
|
||||
|
||||
class T5Attention(nn.Module):
|
||||
"""LingBot World 2 source T5 attention block."""
|
||||
|
||||
def __init__(self, dim: int, dim_attn: int, num_heads: int, dropout: float = 0.1):
|
||||
super().__init__()
|
||||
assert dim_attn % num_heads == 0
|
||||
self.dim = dim
|
||||
self.dim_attn = dim_attn
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim_attn // num_heads
|
||||
self.q = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.k = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.v = nn.Linear(dim, dim_attn, bias=False)
|
||||
self.o = nn.Linear(dim_attn, dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
context: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
pos_bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Project QKV, add relative bias/mask, and return attended states."""
|
||||
context = x if context is None else context
|
||||
b, n, c = x.size(0), self.num_heads, self.head_dim
|
||||
q = self.q(x).view(b, -1, n, c)
|
||||
k = self.k(context).view(b, -1, n, c)
|
||||
v = self.v(context).view(b, -1, n, c)
|
||||
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
|
||||
if pos_bias is not None:
|
||||
attn_bias += pos_bias
|
||||
if mask is not None:
|
||||
assert mask.ndim in (2, 3)
|
||||
mask = mask.view(b, 1, 1, -1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
|
||||
attn = torch.einsum("binc,bjnc->bnij", q, k) + attn_bias
|
||||
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
||||
x = torch.einsum("bnij,bjnc->binc", attn, v)
|
||||
return self.dropout(self.o(x.reshape(b, -1, n * c)))
|
||||
|
||||
|
||||
class T5FeedForward(nn.Module):
|
||||
"""LingBot World 2 source T5 gated feed-forward block."""
|
||||
|
||||
def __init__(self, dim: int, dim_ffn: int, dropout: float = 0.1):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.dim_ffn = dim_ffn
|
||||
self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
|
||||
self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
|
||||
self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply the gated feed-forward projection."""
|
||||
x = self.fc1(x) * self.gate(x)
|
||||
x = self.dropout(x)
|
||||
x = self.fc2(x)
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class T5RelativeEmbedding(nn.Module):
|
||||
"""Per-block relative position embedding used by LingBot World 2 UMT5."""
|
||||
|
||||
def __init__(self, num_buckets: int, num_heads: int, bidirectional: bool, max_dist: int = 128):
|
||||
super().__init__()
|
||||
self.num_buckets = num_buckets
|
||||
self.num_heads = num_heads
|
||||
self.bidirectional = bidirectional
|
||||
self.max_dist = max_dist
|
||||
self.embedding = nn.Embedding(num_buckets, num_heads)
|
||||
|
||||
def forward(self, lq: int, lk: int) -> torch.Tensor:
|
||||
"""Build a relative-position bias tensor for attention logits."""
|
||||
device = self.embedding.weight.device
|
||||
rel_pos = torch.arange(lk, device=device).unsqueeze(0) - torch.arange(lq, device=device).unsqueeze(1)
|
||||
rel_pos = self._relative_position_bucket(rel_pos)
|
||||
rel_pos_embeds = self.embedding(rel_pos)
|
||||
return rel_pos_embeds.permute(2, 0, 1).unsqueeze(0).contiguous()
|
||||
|
||||
def _relative_position_bucket(self, rel_pos: torch.Tensor) -> torch.Tensor:
|
||||
"""Map token offsets to T5 relative-position buckets."""
|
||||
if self.bidirectional:
|
||||
num_buckets = self.num_buckets // 2
|
||||
rel_buckets = (rel_pos > 0).long() * num_buckets
|
||||
rel_pos = torch.abs(rel_pos)
|
||||
else:
|
||||
num_buckets = self.num_buckets
|
||||
rel_buckets = 0
|
||||
rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
|
||||
max_exact = num_buckets // 2
|
||||
rel_pos_large = max_exact + (
|
||||
torch.log(rel_pos.float() / max_exact) / math.log(self.max_dist / max_exact) * (num_buckets - max_exact)
|
||||
).long()
|
||||
rel_pos_large = torch.min(rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
|
||||
rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
|
||||
return rel_buckets
|
||||
|
||||
|
||||
class T5SelfAttention(nn.Module):
|
||||
"""One source-compatible UMT5 encoder block."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_attn: int,
|
||||
dim_ffn: int,
|
||||
num_heads: int,
|
||||
num_buckets: int,
|
||||
shared_pos: bool = False,
|
||||
dropout: float = 0.1,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm1 = T5LayerNorm(dim)
|
||||
self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
||||
self.norm2 = T5LayerNorm(dim)
|
||||
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
|
||||
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(num_buckets, num_heads, bidirectional=True)
|
||||
|
||||
def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None, pos_bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
"""Run self-attention and feed-forward residual updates."""
|
||||
e = pos_bias if self.pos_embedding is None else self.pos_embedding(x.size(1), x.size(1))
|
||||
x = self._fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
|
||||
return self._fp16_clamp(x + self.ffn(self.norm2(x)))
|
||||
|
||||
@staticmethod
|
||||
def _fp16_clamp(x: torch.Tensor) -> torch.Tensor:
|
||||
if x.dtype == torch.float16 and torch.isinf(x).any():
|
||||
clamp = torch.finfo(x.dtype).max - 1000
|
||||
x = torch.clamp(x, min=-clamp, max=clamp)
|
||||
return x
|
||||
|
||||
|
||||
class LingBotWorld2T5EncoderModel(TextEncoder):
|
||||
"""FastVideo-native LingBot World 2 UMT5 encoder for the released `.pth` weights."""
|
||||
|
||||
fall_back_to_pt_during_load = True
|
||||
allow_patterns_overrides = ["*.pt"]
|
||||
|
||||
def __init__(self, config: LingBotWorld2UMT5Config, prefix: str = ""):
|
||||
super().__init__(config)
|
||||
del prefix
|
||||
arch = config.arch_config
|
||||
self.token_embedding = nn.Embedding(arch.vocab_size, arch.dim)
|
||||
self.dropout = nn.Dropout(arch.dropout)
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
T5SelfAttention(
|
||||
arch.dim,
|
||||
arch.dim_attn,
|
||||
arch.dim_ffn,
|
||||
arch.num_heads,
|
||||
arch.num_buckets,
|
||||
shared_pos=False,
|
||||
dropout=arch.dropout,
|
||||
)
|
||||
for _ in range(arch.num_layers)
|
||||
]
|
||||
)
|
||||
self.norm = T5LayerNorm(arch.dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
"""Encode token ids and return source-compatible hidden states."""
|
||||
del position_ids, inputs_embeds, output_hidden_states, kwargs
|
||||
assert input_ids is not None
|
||||
x = self.dropout(self.token_embedding(input_ids))
|
||||
for block in self.blocks:
|
||||
x = block(x, attention_mask)
|
||||
x = self.dropout(self.norm(x))
|
||||
return BaseEncoderOutput(last_hidden_state=x, attention_mask=attention_mask)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
"""Load source `.pth` weights whose names already match this module."""
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if name not in params_dict:
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
|
||||
|
||||
EntryClass = LingBotWorld2T5EncoderModel
|
||||
@@ -327,17 +327,15 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> nn.Module:
|
||||
from transformers import AutoModelForCausalLM
|
||||
from transformers import AutoModel
|
||||
|
||||
if device.type == "cpu" and torch.cuda.is_available():
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
device = get_local_torch_device()
|
||||
|
||||
return AutoModelForCausalLM.from_pretrained(
|
||||
# FastVideo uses Qwen3 only as a text encoder. Loading the body avoids
|
||||
# materializing an unused LM head and full-vocabulary logits, including
|
||||
# for checkpoints whose metadata names Qwen3ForCausalLM.
|
||||
return AutoModel.from_pretrained(
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval().to(device)
|
||||
|
||||
@@ -368,9 +366,18 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
residual = None
|
||||
|
||||
if position_ids is None:
|
||||
# Expand to [batch_size, seq_len]: the rotary layer flattens
|
||||
# positions to ``num_tokens`` and reshapes q/k to
|
||||
# ``(num_tokens, -1, head_dim)``. A bare [1, seq_len] only matches
|
||||
# ``num_tokens`` when batch_size == 1; for batched inputs it folds
|
||||
# the batch dim into the head dim and misaligns RoPE. Expanding to
|
||||
# ``batch_size * seq_len`` tokens keeps the layout correct.
|
||||
position_ids = torch.arange(
|
||||
0, hidden_states.shape[1], device=hidden_states.device
|
||||
).unsqueeze(0)
|
||||
0,
|
||||
hidden_states.shape[1],
|
||||
device=hidden_states.device,
|
||||
dtype=torch.long,
|
||||
).unsqueeze(0).expand(hidden_states.shape[0], -1)
|
||||
|
||||
all_hidden_states: tuple[Any, ...] | None = (
|
||||
() if output_hidden_states else None
|
||||
@@ -405,6 +412,20 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
) -> set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
stacked_params_mapping = self.config.arch_config.stacked_params_mapping
|
||||
# A fused destination is initialized either by one already-fused tensor
|
||||
# or after every split source projection has been loaded. Include
|
||||
# auxiliary quantization parameters (for example scale_weight) rather
|
||||
# than limiting completeness checks to weight/bias tensors.
|
||||
expected_stacked_shards = {
|
||||
(name, shard_id)
|
||||
for name in params_dict
|
||||
for param_name, _, shard_id in stacked_params_mapping
|
||||
if param_name in name
|
||||
}
|
||||
fused_param_names = {name for name, _ in expected_stacked_shards}
|
||||
loaded_stacked_shards: set[tuple[str, str | int]] = set()
|
||||
loaded_fused_params: set[str] = set()
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("model."):
|
||||
@@ -423,37 +444,66 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
continue
|
||||
name = kv_scale_name
|
||||
|
||||
for (
|
||||
param_name,
|
||||
weight_name,
|
||||
shard_id,
|
||||
) in self.config.arch_config.stacked_params_mapping:
|
||||
matched_stacked_param = False
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
matched_stacked_param = True
|
||||
target_name = name.replace(weight_name, param_name)
|
||||
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if target_name.endswith(".bias") and target_name not in params_dict:
|
||||
break
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
if target_name not in params_dict:
|
||||
break
|
||||
|
||||
param = params_dict[name]
|
||||
param = params_dict[target_name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
loaded_params.add(target_name)
|
||||
loaded_stacked_shards.add((target_name, shard_id))
|
||||
break
|
||||
|
||||
if matched_stacked_param:
|
||||
continue
|
||||
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
if name in fused_param_names and name.endswith(".scale_weight"):
|
||||
# Merged scale loaders interpret a missing shard id as shard 0.
|
||||
# An exact fused key is already a complete vector, so copy it
|
||||
# atomically and retain the default loader's shape validation.
|
||||
default_weight_loader(param, loaded_weight)
|
||||
else:
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
|
||||
loaded_params.add(name)
|
||||
if name in fused_param_names:
|
||||
loaded_fused_params.add(name)
|
||||
|
||||
required_split_shards = {
|
||||
(name, shard_id)
|
||||
for name, shard_id in expected_stacked_shards
|
||||
if name not in loaded_fused_params
|
||||
}
|
||||
missing_stacked_shards = required_split_shards - loaded_stacked_shards
|
||||
if missing_stacked_shards:
|
||||
formatted_missing = ", ".join(
|
||||
f"{name}[{shard_id}]"
|
||||
for name, shard_id in sorted(
|
||||
missing_stacked_shards,
|
||||
key=lambda item: (item[0], str(item[1])),
|
||||
)
|
||||
)
|
||||
raise ValueError(
|
||||
"Missing required stacked checkpoint shards: "
|
||||
f"{formatted_missing}"
|
||||
)
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
@@ -92,6 +92,8 @@ class ComponentLoader(ABC):
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_3": (TokenizerLoader, "transformers"),
|
||||
# Cosmos3's model_index names its Qwen2 tokenizer "text_tokenizer".
|
||||
"text_tokenizer": (TokenizerLoader, "transformers"),
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"feature_extractor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
@@ -391,18 +393,22 @@ class TextEncoderLoader(ComponentLoader):
|
||||
model_config.quant_config = quant_cls()
|
||||
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
if getattr(model_cls, "supports_hf_from_pretrained", False):
|
||||
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
|
||||
model_path,
|
||||
model_config, # type: ignore[arg-type]
|
||||
dtype=PRECISION_TO_TYPE[dtype],
|
||||
device=target_device,
|
||||
)
|
||||
return model.eval()
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
if getattr(model_cls, "supports_hf_from_pretrained", False):
|
||||
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
|
||||
model_path,
|
||||
model_config, # type: ignore[arg-type]
|
||||
dtype=PRECISION_TO_TYPE[dtype],
|
||||
device=target_device,
|
||||
)
|
||||
# HF passthrough encoders return before FastVideo's FSDP
|
||||
# wrapping path, so the text stage needs their placement to
|
||||
# put token tensors on the same device.
|
||||
model._fastvideo_input_device = target_device
|
||||
return model.eval()
|
||||
|
||||
with target_device:
|
||||
model = model_cls(model_config) # type: ignore
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
@@ -820,6 +826,24 @@ class VAELoader(ComponentLoader):
|
||||
vae.load_state_dict(sd, strict=False)
|
||||
return vae.eval()
|
||||
|
||||
if class_name == "LingBotWorld2WanVAE":
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
config.pop("_class_name", None)
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
weight_path = os.path.join(model_path, "Wan2.1_VAE.pth")
|
||||
if not os.path.exists(weight_path):
|
||||
raise FileNotFoundError(
|
||||
f"Missing LingBot World 2 VAE weights: {weight_path}"
|
||||
)
|
||||
vae = vae_cls(
|
||||
vae_config,
|
||||
checkpoint_path=weight_path,
|
||||
dtype=dtype,
|
||||
).to(target_device)
|
||||
return vae.eval()
|
||||
|
||||
# LTX-2 uses CausalVideoAutoencoder with nested "vae" config
|
||||
if class_name == "CausalVideoAutoencoder" and "vae" in config:
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
@@ -1137,7 +1161,19 @@ class SchedulerLoader(ComponentLoader):
|
||||
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
# Diffusers checkpoints can carry newer scheduler config keys than the
|
||||
# vendored scheduler accepts (e.g. shift_terminal / sigma_min / sigma_max
|
||||
# from a newer diffusers release). Filter to the class's __init__ params,
|
||||
# mirroring diffusers' ``from_config``, so loading is robust to schema
|
||||
# drift instead of crashing on an unexpected kwarg.
|
||||
import inspect
|
||||
valid_params = set(inspect.signature(scheduler_cls.__init__).parameters)
|
||||
filtered_config = {k: v for k, v in config.items() if k in valid_params}
|
||||
dropped = sorted(set(config) - set(filtered_config))
|
||||
if dropped:
|
||||
logger.warning("Scheduler %s: dropping unsupported config keys %s", class_name, dropped)
|
||||
|
||||
scheduler = scheduler_cls(**filtered_config)
|
||||
if fastvideo_args.pipeline_config.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
|
||||
return scheduler
|
||||
|
||||
@@ -15,13 +15,11 @@ import torch
|
||||
from torch import nn
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
from torch.distributed._tensor import distribute_tensor
|
||||
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
|
||||
MixedPrecisionPolicy, fully_shard)
|
||||
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard)
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.utils import (get_param_names_mapping,
|
||||
hf_to_custom_state_dict)
|
||||
from fastvideo.models.loader.utils import (get_param_names_mapping, hf_to_custom_state_dict)
|
||||
from fastvideo.models.loader.weight_utils import safetensors_weights_iterator
|
||||
from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
|
||||
|
||||
@@ -43,13 +41,16 @@ def _maybe_quantize_model(model: nn.Module) -> None:
|
||||
"""
|
||||
# Defer imports: these modules pull in heavy symbols at module-load time.
|
||||
from fastvideo.layers.quantization.nvfp4_config import (
|
||||
NVFP4QuantizeMethod, convert_model_to_nvfp4,
|
||||
NVFP4QuantizeMethod,
|
||||
convert_model_to_nvfp4,
|
||||
)
|
||||
from fastvideo.layers.quantization.nvfp4_qat_config import (
|
||||
NVFP4QATQuantizeMethod, convert_model_to_fp4,
|
||||
NVFP4QATQuantizeMethod,
|
||||
convert_model_to_fp4,
|
||||
)
|
||||
from fastvideo.layers.quantization.fp8_config import (
|
||||
FP8QuantizeMethod, convert_model_to_fp8,
|
||||
FP8QuantizeMethod,
|
||||
convert_model_to_fp8,
|
||||
)
|
||||
|
||||
for mod in model.modules():
|
||||
@@ -121,10 +122,7 @@ def maybe_load_fsdp_model(
|
||||
"""
|
||||
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
|
||||
# manually casting the inputs to the model
|
||||
mp_policy = MixedPrecisionPolicy(param_dtype,
|
||||
reduce_dtype,
|
||||
output_dtype,
|
||||
cast_forward_inputs=False)
|
||||
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False)
|
||||
|
||||
set_mixed_precision_policy(
|
||||
param_dtype=param_dtype,
|
||||
@@ -137,6 +135,13 @@ def maybe_load_fsdp_model(
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
dtype_selector = getattr(model, "_get_parameter_dtype", None)
|
||||
has_mixed_parameter_dtypes = callable(dtype_selector) and any(
|
||||
dtype_selector(name, param_dtype) != param_dtype for name, _ in model.named_parameters())
|
||||
if training_mode and has_mixed_parameter_dtypes:
|
||||
raise NotImplementedError("FSDP training with model-selected mixed parameter dtypes requires "
|
||||
"separate gradient synchronization for replicated parameters.")
|
||||
|
||||
# Check if we should use FSDP
|
||||
use_fsdp = training_mode or fsdp_inference
|
||||
|
||||
@@ -152,7 +157,7 @@ def maybe_load_fsdp_model(
|
||||
if not training_mode and not fsdp_inference:
|
||||
hsdp_replicate_dim = world_size
|
||||
hsdp_shard_dim = 1
|
||||
|
||||
|
||||
if current_platform.is_npu():
|
||||
with torch.device("cpu"):
|
||||
device_mesh = init_device_mesh(
|
||||
@@ -163,11 +168,11 @@ def maybe_load_fsdp_model(
|
||||
)
|
||||
else:
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
# (Replicate(), Shard(dim=0))
|
||||
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
"cuda",
|
||||
# (Replicate(), Shard(dim=0))
|
||||
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
@@ -188,12 +193,10 @@ def maybe_load_fsdp_model(
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
)
|
||||
if hasattr(model, "materialize_non_persistent_buffers"):
|
||||
model.materialize_non_persistent_buffers(
|
||||
device=device, dtype=default_dtype)
|
||||
model.materialize_non_persistent_buffers(device=device, dtype=default_dtype)
|
||||
for n, p in chain(model.named_parameters(), model.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
@@ -209,8 +212,7 @@ def maybe_load_fsdp_model(
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
if compile_in_loader:
|
||||
compile_kwargs = torch_compile_kwargs or {}
|
||||
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
|
||||
compile_kwargs)
|
||||
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
|
||||
model = torch.compile(model, **compile_kwargs)
|
||||
logger.info("torch.compile enabled for %s", type(model).__name__)
|
||||
return model
|
||||
@@ -254,58 +256,81 @@ def shard_model(
|
||||
"""
|
||||
# Check if we should use size-based filtering
|
||||
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
|
||||
|
||||
|
||||
if not fsdp_shard_conditions:
|
||||
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
|
||||
return
|
||||
|
||||
default_param_dtype = getattr(mp_policy, "param_dtype", None)
|
||||
dtype_selector = getattr(model, "_get_parameter_dtype", None)
|
||||
ignored_params: set[nn.Parameter] = set()
|
||||
if callable(dtype_selector) and default_param_dtype is not None:
|
||||
ignored_params = {
|
||||
parameter
|
||||
for name, parameter in model.named_parameters()
|
||||
if dtype_selector(name, default_param_dtype) != default_param_dtype
|
||||
}
|
||||
named_modules = list(model.named_modules())
|
||||
ignored_params_by_module = {
|
||||
id(module): ignored_params.intersection(set(module.parameters()))
|
||||
for _, module in named_modules
|
||||
}
|
||||
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": mesh,
|
||||
"mp_policy": mp_policy,
|
||||
}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
|
||||
pin_memory=pin_cpu_memory)
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=pin_cpu_memory)
|
||||
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
|
||||
|
||||
if use_size_filtering:
|
||||
# Size-based filtering mode
|
||||
min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000"))
|
||||
logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6)
|
||||
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
|
||||
for n, m in reversed(named_modules):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
# Count all parameters
|
||||
param_count = sum(p.numel() for p in m.parameters(recurse=True))
|
||||
|
||||
|
||||
# Skip small modules
|
||||
if param_count < min_params:
|
||||
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)",
|
||||
n, param_count / 1e6, min_params / 1e6)
|
||||
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)", n, param_count / 1e6,
|
||||
min_params / 1e6)
|
||||
continue
|
||||
|
||||
|
||||
# Shard this module
|
||||
logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6)
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
module_kwargs = fsdp_kwargs
|
||||
local_ignored_params = ignored_params_by_module[id(m)]
|
||||
if local_ignored_params:
|
||||
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
|
||||
fully_shard(m, **module_kwargs)
|
||||
num_layers_sharded += 1
|
||||
else:
|
||||
# Shard all modules matching conditions
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
# Shard all modules matching conditions
|
||||
for n, m in reversed(named_modules):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
module_kwargs = fsdp_kwargs
|
||||
local_ignored_params = ignored_params_by_module[id(m)]
|
||||
if local_ignored_params:
|
||||
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
|
||||
fully_shard(m, **module_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.")
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
root_kwargs = fsdp_kwargs
|
||||
if ignored_params:
|
||||
root_kwargs = {**fsdp_kwargs, "ignored_params": ignored_params}
|
||||
fully_shard(model, **root_kwargs)
|
||||
|
||||
|
||||
# TODO(PY): device mesh for cfg parallel
|
||||
@@ -341,17 +366,17 @@ def load_model_from_full_model_state_dict(
|
||||
"""
|
||||
meta_sd = model.state_dict()
|
||||
named_parameters = dict(model.named_parameters())
|
||||
named_buffers = dict(model.named_buffers())
|
||||
sharded_sd = {}
|
||||
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
||||
full_sd_iterator, param_names_mapping) # type: ignore
|
||||
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator,
|
||||
param_names_mapping) # type: ignore
|
||||
for target_param_name, full_tensor in custom_param_sd.items():
|
||||
meta_sharded_param = meta_sd.get(target_param_name)
|
||||
if meta_sharded_param is None:
|
||||
# Some checkpoints include extra entries that are not part of the
|
||||
# instantiated model's state_dict (e.g. `_extra_state` keys from
|
||||
# some FSDP checkpoint formats). These can be safely skipped.
|
||||
if (target_param_name.endswith("._extra_state")
|
||||
or target_param_name.endswith("_extra_state")):
|
||||
if (target_param_name.endswith("._extra_state") or target_param_name.endswith("_extra_state")):
|
||||
logger.warning(
|
||||
"Skipping non-parameter checkpoint key: %s",
|
||||
target_param_name,
|
||||
@@ -370,8 +395,12 @@ def load_model_from_full_model_state_dict(
|
||||
raise ValueError(
|
||||
f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect."
|
||||
)
|
||||
target_dtype = param_dtype
|
||||
dtype_selector = getattr(model, "_get_parameter_dtype", None)
|
||||
if callable(dtype_selector):
|
||||
target_dtype = dtype_selector(target_param_name, param_dtype)
|
||||
if not hasattr(meta_sharded_param, "device_mesh"):
|
||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
|
||||
target_param = named_parameters.get(target_param_name)
|
||||
weight_loader = getattr(target_param, "weight_loader", None)
|
||||
# Gated on a shape mismatch: only fused/stacked params with a custom
|
||||
@@ -380,9 +409,7 @@ def load_model_from_full_model_state_dict(
|
||||
# fall through to the original `sharded_tensor = full_tensor` below.
|
||||
if target_param is not None and callable(weight_loader) and tuple(target_param.shape) != tuple(
|
||||
full_tensor.shape):
|
||||
loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape),
|
||||
device=device,
|
||||
dtype=param_dtype),
|
||||
loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape), device=device, dtype=target_dtype),
|
||||
requires_grad=False)
|
||||
for attr_name, attr_value in vars(target_param).items():
|
||||
setattr(loaded_param, attr_name, attr_value)
|
||||
@@ -392,7 +419,7 @@ def load_model_from_full_model_state_dict(
|
||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors.
|
||||
sharded_tensor = full_tensor
|
||||
else:
|
||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
|
||||
sharded_tensor = distribute_tensor(
|
||||
full_tensor,
|
||||
meta_sharded_param.device_mesh,
|
||||
@@ -400,36 +427,35 @@ def load_model_from_full_model_state_dict(
|
||||
)
|
||||
if cpu_offload:
|
||||
sharded_tensor = sharded_tensor.cpu()
|
||||
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
|
||||
if target_param_name in named_buffers:
|
||||
sharded_sd[target_param_name] = sharded_tensor
|
||||
else:
|
||||
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
|
||||
|
||||
model.reverse_param_names_mapping = reverse_param_names_mapping
|
||||
unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys())
|
||||
if unused_keys:
|
||||
logger.warning("Found unloaded parameters in meta state dict: %s",
|
||||
unused_keys)
|
||||
logger.warning("Found unloaded parameters in meta state dict: %s", unused_keys)
|
||||
|
||||
# List of allowed parameter name patterns
|
||||
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
|
||||
for new_param_name in unused_keys:
|
||||
if not any(pattern in new_param_name
|
||||
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
|
||||
logger.error("Unsupported new parameter: %s. Allowed patterns: %s",
|
||||
new_param_name, ALLOWED_NEW_PARAM_PATTERNS)
|
||||
raise ValueError(
|
||||
f"New parameter '{new_param_name}' is not supported. "
|
||||
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed."
|
||||
)
|
||||
if not any(pattern in new_param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS):
|
||||
logger.error("Unsupported new parameter: %s. Allowed patterns: %s", new_param_name,
|
||||
ALLOWED_NEW_PARAM_PATTERNS)
|
||||
raise ValueError(f"New parameter '{new_param_name}' is not supported. "
|
||||
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed.")
|
||||
meta_sharded_param = meta_sd.get(new_param_name)
|
||||
target_dtype = param_dtype
|
||||
dtype_selector = getattr(model, "_get_parameter_dtype", None)
|
||||
if callable(dtype_selector):
|
||||
target_dtype = dtype_selector(new_param_name, param_dtype)
|
||||
if not hasattr(meta_sharded_param, "device_mesh"):
|
||||
# Initialize with zeros
|
||||
sharded_tensor = torch.zeros_like(meta_sharded_param,
|
||||
device=device,
|
||||
dtype=param_dtype)
|
||||
sharded_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype)
|
||||
else:
|
||||
# Initialize with zeros and distribute
|
||||
full_tensor = torch.zeros_like(meta_sharded_param,
|
||||
device=device,
|
||||
dtype=param_dtype)
|
||||
full_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype)
|
||||
sharded_tensor = distribute_tensor(
|
||||
full_tensor,
|
||||
meta_sharded_param.device_mesh,
|
||||
|
||||
@@ -18,8 +18,10 @@ def set_default_torch_dtype(dtype: torch.dtype):
|
||||
"""Sets the default torch dtype to the given dtype."""
|
||||
old_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(dtype)
|
||||
yield
|
||||
torch.set_default_dtype(old_dtype)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.set_default_dtype(old_dtype)
|
||||
|
||||
|
||||
def get_param_names_mapping(
|
||||
|
||||
@@ -37,11 +37,19 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
|
||||
# Cosmos3-Nano's checkpoint model_index names the DiT "Cosmos3OmniTransformer";
|
||||
# map that HF class name to FastVideo's native Cosmos3VFMTransformer.
|
||||
"Cosmos3OmniTransformer": ("dits", "cosmos3", "Cosmos3VFMTransformer"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
|
||||
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
|
||||
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
|
||||
"LingBotWorld2CausalFastTransformer3DModel": (
|
||||
"dits",
|
||||
"lingbotworld2",
|
||||
"LingBotWorld2CausalFastTransformer3DModel",
|
||||
),
|
||||
"Gen3CTransformer3DModel": ("dits", "gen3c", "Gen3CTransformer3DModel"),
|
||||
"Kandinsky5Transformer3DModel": ("dits", "kandinsky5", "Kandinsky5Transformer3DModel"),
|
||||
"Flux2Transformer2DModel": ("dits", "flux_2", "Flux2Transformer2DModel"),
|
||||
@@ -53,6 +61,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
|
||||
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"LingBotWorld2CausalFastTransformer3DModel": (
|
||||
"dits",
|
||||
"lingbotworld2",
|
||||
"LingBotWorld2CausalFastTransformer3DModel",
|
||||
),
|
||||
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
|
||||
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
|
||||
# Legacy aliases for older HF model_index.json files
|
||||
@@ -64,6 +77,7 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
# Text-to-image DiT models (2D image generation)
|
||||
_TEXT_TO_IMAGE_DIT_MODELS = {
|
||||
"GlmImageTransformer2DModel": ("dits", "glm_image", "GlmImageTransformer2DModel"),
|
||||
"ZImageTransformer2DModel": ("dits", "zimage", "ZImageTransformer2DModel"),
|
||||
}
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
@@ -72,12 +86,16 @@ _TEXT_ENCODER_MODELS = {
|
||||
("encoders", "clip", "CLIPTextModelWithProjection"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
|
||||
"LingBotWorld2T5EncoderModel": ("encoders", "lingbotworld2_t5", "LingBotWorld2T5EncoderModel"),
|
||||
"T5EncoderModel": ("encoders", "t5_hf", "T5EncoderModel"),
|
||||
"BertModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
|
||||
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"Qwen2_5_VLForConditionalGeneration":
|
||||
("encoders", "reason1", "Reason1TextEncoder"),
|
||||
# Z-Image-Turbo's text_encoder/config.json declares architecture
|
||||
# "Qwen3Model"; route it to the shared Qwen3 encoder (added for Flux2 Klein).
|
||||
"Qwen3Model": ("encoders", "qwen3", "Qwen3ForCausalLM"),
|
||||
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
|
||||
"Qwen3ForCausalLM": ("encoders", "qwen3", "Qwen3ForCausalLM"),
|
||||
"Mistral3ForConditionalGeneration":
|
||||
@@ -98,6 +116,7 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLHYWorld": ("vaes", "hyworldvae", "AutoencoderKLHYWorld"),
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"LingBotWorld2WanVAE": ("vaes", "lingbotworld2_wanvae", "LingBotWorld2WanVAE"),
|
||||
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
|
||||
"AutoencoderKLGen3CTokenizer":
|
||||
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
|
||||
|
||||
@@ -96,6 +96,14 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
The minimum sigma value for the noise schedule.
|
||||
sigma_data (`float`, *optional*):
|
||||
The sigma data value for scaling.
|
||||
use_reference_discrete_timesteps (`bool`, defaults to False):
|
||||
Some reference schedulers (e.g. Z-Image) construct the timestep
|
||||
schedule by linspacing `num_inference_steps + 1` points from
|
||||
`t_max` to `t_min` and dropping the terminal point. Default
|
||||
(`False`) preserves the original `np.linspace(t_max, t_min,
|
||||
num_inference_steps)` (float64) behaviour used by every existing
|
||||
model. Enable this flag only when matching a reference scheduler
|
||||
that expects the +1 + drop-terminal construction.
|
||||
"""
|
||||
|
||||
_compatibles: list[Any] = []
|
||||
@@ -122,6 +130,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
sigma_max: float | None = None,
|
||||
sigma_min: float | None = None,
|
||||
sigma_data: float | None = None,
|
||||
use_reference_discrete_timesteps: bool = False,
|
||||
):
|
||||
if sum([
|
||||
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
|
||||
@@ -155,7 +164,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
|
||||
self.sigmas = sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_min = sigma_min if sigma_min is not None else self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
@@ -350,7 +359,19 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
if timesteps_array is None:
|
||||
t_max = self._sigma_to_t(self.sigma_max)
|
||||
t_min = self._sigma_to_t(self.sigma_min)
|
||||
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
|
||||
if self.config.use_reference_discrete_timesteps:
|
||||
# Some reference schedulers (for example Z-Image) build a
|
||||
# float64 num_steps+1 linspace and drop the terminal point.
|
||||
timesteps_array = np.linspace(
|
||||
t_max,
|
||||
t_min,
|
||||
num_inference_steps + 1,
|
||||
)[:-1]
|
||||
else:
|
||||
# Preserve the original numpy default (float64) here —
|
||||
# casting to float32 silently shifts rounded timestep
|
||||
# values for every existing model that uses this branch.
|
||||
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
|
||||
sigmas_array = timesteps_array / self.config.num_train_timesteps
|
||||
else:
|
||||
sigmas_array = np.array(sigmas).astype(np.float32)
|
||||
|
||||
@@ -0,0 +1,722 @@
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
__all__ = [
|
||||
'Wan2_1_VAE',
|
||||
'LingBotWorld2WanVAE',
|
||||
]
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
|
||||
class CausalConv3d(nn.Conv3d):
|
||||
"""
|
||||
Causal 3d convolusion.
|
||||
"""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._padding = (self.padding[2], self.padding[2], self.padding[1],
|
||||
self.padding[1], 2 * self.padding[0], 0)
|
||||
self.padding = (0, 0, 0)
|
||||
|
||||
def forward(self, x, cache_x=None):
|
||||
padding = list(self._padding)
|
||||
if cache_x is not None and self._padding[4] > 0:
|
||||
cache_x = cache_x.to(x.device)
|
||||
x = torch.cat([cache_x, x], dim=2)
|
||||
padding[4] -= cache_x.shape[2]
|
||||
x = F.pad(x, padding)
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class RMS_norm(nn.Module):
|
||||
|
||||
def __init__(self, dim, channel_first=True, images=True, bias=False):
|
||||
super().__init__()
|
||||
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
||||
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
||||
|
||||
self.channel_first = channel_first
|
||||
self.scale = dim**0.5
|
||||
self.gamma = nn.Parameter(torch.ones(shape))
|
||||
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
|
||||
|
||||
def forward(self, x):
|
||||
return F.normalize(
|
||||
x, dim=(1 if self.channel_first else
|
||||
-1)) * self.scale * self.gamma + self.bias
|
||||
|
||||
|
||||
class Upsample(nn.Upsample):
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Fix bfloat16 support for nearest neighbor interpolation.
|
||||
"""
|
||||
return super().forward(x.float()).type_as(x)
|
||||
|
||||
|
||||
class Resample(nn.Module):
|
||||
|
||||
def __init__(self, dim, mode):
|
||||
assert mode in ('none', 'upsample2d', 'upsample3d', 'downsample2d',
|
||||
'downsample3d')
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mode = mode
|
||||
|
||||
# layers
|
||||
if mode == 'upsample2d':
|
||||
self.resample = nn.Sequential(
|
||||
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
|
||||
nn.Conv2d(dim, dim // 2, 3, padding=1))
|
||||
elif mode == 'upsample3d':
|
||||
self.resample = nn.Sequential(
|
||||
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
|
||||
nn.Conv2d(dim, dim // 2, 3, padding=1))
|
||||
self.time_conv = CausalConv3d(
|
||||
dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
|
||||
|
||||
elif mode == 'downsample2d':
|
||||
self.resample = nn.Sequential(
|
||||
nn.ZeroPad2d((0, 1, 0, 1)),
|
||||
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
||||
elif mode == 'downsample3d':
|
||||
self.resample = nn.Sequential(
|
||||
nn.ZeroPad2d((0, 1, 0, 1)),
|
||||
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
||||
self.time_conv = CausalConv3d(
|
||||
dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
|
||||
|
||||
else:
|
||||
self.resample = nn.Identity()
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
b, c, t, h, w = x.size()
|
||||
if self.mode == 'upsample3d':
|
||||
if feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
if feat_cache[idx] is None:
|
||||
feat_cache[idx] = 'Rep'
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[
|
||||
idx] is not None and feat_cache[idx] != 'Rep':
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
if cache_x.shape[2] < 2 and feat_cache[
|
||||
idx] is not None and feat_cache[idx] == 'Rep':
|
||||
cache_x = torch.cat([
|
||||
torch.zeros_like(cache_x).to(cache_x.device),
|
||||
cache_x
|
||||
],
|
||||
dim=2)
|
||||
if feat_cache[idx] == 'Rep':
|
||||
x = self.time_conv(x)
|
||||
else:
|
||||
x = self.time_conv(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
|
||||
3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
t = x.shape[2]
|
||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
||||
x = self.resample(x)
|
||||
x = rearrange(x, '(b t) c h w -> b c t h w', t=t)
|
||||
|
||||
if self.mode == 'downsample3d':
|
||||
if feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
if feat_cache[idx] is None:
|
||||
feat_cache[idx] = x.clone()
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
|
||||
cache_x = x[:, :, -1:, :, :].clone()
|
||||
# if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx]!='Rep':
|
||||
# # cache last frame of last two chunk
|
||||
# cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
||||
|
||||
x = self.time_conv(
|
||||
torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
return x
|
||||
|
||||
def init_weight(self, conv):
|
||||
conv_weight = conv.weight
|
||||
nn.init.zeros_(conv_weight)
|
||||
c1, c2, t, h, w = conv_weight.size()
|
||||
one_matrix = torch.eye(c1, c2)
|
||||
init_matrix = one_matrix
|
||||
nn.init.zeros_(conv_weight)
|
||||
#conv_weight.data[:,:,-1,1,1] = init_matrix * 0.5
|
||||
conv_weight.data[:, :, 1, 0, 0] = init_matrix #* 0.5
|
||||
conv.weight.data.copy_(conv_weight)
|
||||
nn.init.zeros_(conv.bias.data)
|
||||
|
||||
def init_weight2(self, conv):
|
||||
conv_weight = conv.weight.data
|
||||
nn.init.zeros_(conv_weight)
|
||||
c1, c2, t, h, w = conv_weight.size()
|
||||
init_matrix = torch.eye(c1 // 2, c2)
|
||||
#init_matrix = repeat(init_matrix, 'o ... -> (o 2) ...').permute(1,0,2).contiguous().reshape(c1,c2)
|
||||
conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
|
||||
conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
|
||||
conv.weight.data.copy_(conv_weight)
|
||||
nn.init.zeros_(conv.bias.data)
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
|
||||
def __init__(self, in_dim, out_dim, dropout=0.0):
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
|
||||
# layers
|
||||
self.residual = nn.Sequential(
|
||||
RMS_norm(in_dim, images=False), nn.SiLU(),
|
||||
CausalConv3d(in_dim, out_dim, 3, padding=1),
|
||||
RMS_norm(out_dim, images=False), nn.SiLU(), nn.Dropout(dropout),
|
||||
CausalConv3d(out_dim, out_dim, 3, padding=1))
|
||||
self.shortcut = CausalConv3d(in_dim, out_dim, 1) \
|
||||
if in_dim != out_dim else nn.Identity()
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
h = self.shortcut(x)
|
||||
for layer in self.residual:
|
||||
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = layer(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x + h
|
||||
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
"""
|
||||
Causal self-attention with a single head.
|
||||
"""
|
||||
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
|
||||
# layers
|
||||
self.norm = RMS_norm(dim)
|
||||
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
||||
self.proj = nn.Conv2d(dim, dim, 1)
|
||||
|
||||
# zero out the last layer params
|
||||
nn.init.zeros_(self.proj.weight)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
b, c, t, h, w = x.size()
|
||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
||||
x = self.norm(x)
|
||||
# compute query, key, value
|
||||
q, k, v = self.to_qkv(x).reshape(b * t, 1, c * 3,
|
||||
-1).permute(0, 1, 3,
|
||||
2).contiguous().chunk(
|
||||
3, dim=-1)
|
||||
|
||||
# apply attention
|
||||
x = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
)
|
||||
x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
|
||||
|
||||
# output
|
||||
x = self.proj(x)
|
||||
x = rearrange(x, '(b t) c h w-> b c t h w', t=t)
|
||||
return x + identity
|
||||
|
||||
|
||||
class Encoder3d(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=128,
|
||||
z_dim=4,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_downsample=[True, True, False],
|
||||
dropout=0.0):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_downsample = temperal_downsample
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [1] + dim_mult]
|
||||
scale = 1.0
|
||||
|
||||
# init block
|
||||
self.conv1 = CausalConv3d(3, dims[0], 3, padding=1)
|
||||
|
||||
# downsample blocks
|
||||
downsamples = []
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
# residual (+attention) blocks
|
||||
for _ in range(num_res_blocks):
|
||||
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
||||
if scale in attn_scales:
|
||||
downsamples.append(AttentionBlock(out_dim))
|
||||
in_dim = out_dim
|
||||
|
||||
# downsample block
|
||||
if i != len(dim_mult) - 1:
|
||||
mode = 'downsample3d' if temperal_downsample[
|
||||
i] else 'downsample2d'
|
||||
downsamples.append(Resample(out_dim, mode=mode))
|
||||
scale /= 2.0
|
||||
self.downsamples = nn.Sequential(*downsamples)
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(
|
||||
ResidualBlock(out_dim, out_dim, dropout), AttentionBlock(out_dim),
|
||||
ResidualBlock(out_dim, out_dim, dropout))
|
||||
|
||||
# output blocks
|
||||
self.head = nn.Sequential(
|
||||
RMS_norm(out_dim, images=False), nn.SiLU(),
|
||||
CausalConv3d(out_dim, z_dim, 3, padding=1))
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
if feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv1(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = self.conv1(x)
|
||||
|
||||
## downsamples
|
||||
for layer in self.downsamples:
|
||||
if feat_cache is not None:
|
||||
x = layer(x, feat_cache, feat_idx)
|
||||
else:
|
||||
x = layer(x)
|
||||
|
||||
## middle
|
||||
for layer in self.middle:
|
||||
if isinstance(layer, ResidualBlock) and feat_cache is not None:
|
||||
x = layer(x, feat_cache, feat_idx)
|
||||
else:
|
||||
x = layer(x)
|
||||
|
||||
## head
|
||||
for layer in self.head:
|
||||
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = layer(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
|
||||
class Decoder3d(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=128,
|
||||
z_dim=4,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_upsample=[False, True, True],
|
||||
dropout=0.0):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_upsample = temperal_upsample
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
||||
scale = 1.0 / 2**(len(dim_mult) - 2)
|
||||
|
||||
# init block
|
||||
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(
|
||||
ResidualBlock(dims[0], dims[0], dropout), AttentionBlock(dims[0]),
|
||||
ResidualBlock(dims[0], dims[0], dropout))
|
||||
|
||||
# upsample blocks
|
||||
upsamples = []
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
# residual (+attention) blocks
|
||||
if i == 1 or i == 2 or i == 3:
|
||||
in_dim = in_dim // 2
|
||||
for _ in range(num_res_blocks + 1):
|
||||
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
||||
if scale in attn_scales:
|
||||
upsamples.append(AttentionBlock(out_dim))
|
||||
in_dim = out_dim
|
||||
|
||||
# upsample block
|
||||
if i != len(dim_mult) - 1:
|
||||
mode = 'upsample3d' if temperal_upsample[i] else 'upsample2d'
|
||||
upsamples.append(Resample(out_dim, mode=mode))
|
||||
scale *= 2.0
|
||||
self.upsamples = nn.Sequential(*upsamples)
|
||||
|
||||
# output blocks
|
||||
self.head = nn.Sequential(
|
||||
RMS_norm(out_dim, images=False), nn.SiLU(),
|
||||
CausalConv3d(out_dim, 3, 3, padding=1))
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
## conv1
|
||||
if feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv1(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = self.conv1(x)
|
||||
|
||||
## middle
|
||||
for layer in self.middle:
|
||||
if isinstance(layer, ResidualBlock) and feat_cache is not None:
|
||||
x = layer(x, feat_cache, feat_idx)
|
||||
else:
|
||||
x = layer(x)
|
||||
|
||||
## upsamples
|
||||
for layer in self.upsamples:
|
||||
if feat_cache is not None:
|
||||
x = layer(x, feat_cache, feat_idx)
|
||||
else:
|
||||
x = layer(x)
|
||||
|
||||
## head
|
||||
for layer in self.head:
|
||||
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = layer(x, feat_cache[idx])
|
||||
feat_cache[idx] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
|
||||
def count_conv3d(model):
|
||||
count = 0
|
||||
for m in model.modules():
|
||||
if isinstance(m, CausalConv3d):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
class WanVAE_(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=128,
|
||||
z_dim=4,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_downsample=[True, True, False],
|
||||
dropout=0.0):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_downsample = temperal_downsample
|
||||
self.temperal_upsample = temperal_downsample[::-1]
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_downsample, dropout)
|
||||
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
|
||||
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
|
||||
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_upsample, dropout)
|
||||
|
||||
def forward(self, x):
|
||||
mu, log_var = self.encode(x)
|
||||
z = self.reparameterize(mu, log_var)
|
||||
x_recon = self.decode(z)
|
||||
return x_recon, mu, log_var
|
||||
|
||||
def encode(self, x, scale):
|
||||
self.clear_cache()
|
||||
## cache
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
for i in range(iter_):
|
||||
self._enc_conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.encoder(
|
||||
x[:, :, :1, :, :],
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
else:
|
||||
out_ = self.encoder(
|
||||
x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
mu, log_var = self.conv1(out).chunk(2, dim=1)
|
||||
if isinstance(scale[0], torch.Tensor):
|
||||
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
else:
|
||||
mu = (mu - scale[0]) * scale[1]
|
||||
self.clear_cache()
|
||||
return mu
|
||||
|
||||
def decode(self, z, scale):
|
||||
self.clear_cache()
|
||||
# z: [b,c,t,h,w]
|
||||
if isinstance(scale[0], torch.Tensor):
|
||||
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
else:
|
||||
z = z / scale[1] + scale[0]
|
||||
iter_ = z.shape[2]
|
||||
x = self.conv2(z)
|
||||
for i in range(iter_):
|
||||
self._conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.decoder(
|
||||
x[:, :, i:i + 1, :, :],
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
else:
|
||||
out_ = self.decoder(
|
||||
x[:, :, i:i + 1, :, :],
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
self.clear_cache()
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
std = torch.exp(0.5 * log_var)
|
||||
eps = torch.randn_like(std)
|
||||
return eps * std + mu
|
||||
|
||||
def sample(self, imgs, deterministic=False):
|
||||
mu, log_var = self.encode(imgs)
|
||||
if deterministic:
|
||||
return mu
|
||||
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
|
||||
return mu + std * torch.randn_like(std)
|
||||
|
||||
def clear_cache(self):
|
||||
self._conv_num = count_conv3d(self.decoder)
|
||||
self._conv_idx = [0]
|
||||
self._feat_map = [None] * self._conv_num
|
||||
#cache encode
|
||||
self._enc_conv_num = count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = [0]
|
||||
self._enc_feat_map = [None] * self._enc_conv_num
|
||||
|
||||
|
||||
def _video_vae(pretrained_path=None, z_dim=None, device='cpu', **kwargs):
|
||||
"""
|
||||
Autoencoder3d adapted from Stable Diffusion 1.x, 2.x and XL.
|
||||
"""
|
||||
# params
|
||||
cfg = dict(
|
||||
dim=96,
|
||||
z_dim=z_dim,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_downsample=[False, True, True],
|
||||
dropout=0.0)
|
||||
cfg.update(**kwargs)
|
||||
|
||||
# init model
|
||||
with torch.device('meta'):
|
||||
model = WanVAE_(**cfg)
|
||||
|
||||
# load checkpoint
|
||||
logging.info(f'loading {pretrained_path}')
|
||||
model.load_state_dict(
|
||||
torch.load(pretrained_path, map_location=device), assign=True)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
class Wan2_1_VAE:
|
||||
|
||||
def __init__(self,
|
||||
z_dim=16,
|
||||
vae_pth='cache/vae_step_411000.pth',
|
||||
dtype=torch.float,
|
||||
device="cuda"):
|
||||
self.dtype = dtype
|
||||
self.device = device
|
||||
|
||||
mean = [
|
||||
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
|
||||
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
|
||||
]
|
||||
std = [
|
||||
2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743,
|
||||
3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160
|
||||
]
|
||||
self.mean = torch.tensor(mean, dtype=dtype, device=device)
|
||||
self.std = torch.tensor(std, dtype=dtype, device=device)
|
||||
self.scale = [self.mean, 1.0 / self.std]
|
||||
|
||||
# init model
|
||||
self.model = _video_vae(
|
||||
pretrained_path=vae_pth,
|
||||
z_dim=z_dim,
|
||||
).eval().requires_grad_(False).to(device)
|
||||
|
||||
def encode(self, videos):
|
||||
"""
|
||||
videos: A list of videos each with shape [C, T, H, W].
|
||||
"""
|
||||
with amp.autocast(dtype=self.dtype):
|
||||
return [
|
||||
self.model.encode(u.unsqueeze(0), self.scale).float().squeeze(0)
|
||||
for u in videos
|
||||
]
|
||||
|
||||
def decode(self, zs):
|
||||
with amp.autocast(dtype=self.dtype):
|
||||
return [
|
||||
self.model.decode(u.unsqueeze(0),
|
||||
self.scale).float().clamp_(-1, 1).squeeze(0)
|
||||
for u in zs
|
||||
]
|
||||
|
||||
|
||||
class LingBotWorld2WanVAE(nn.Module):
|
||||
"""FastVideo-facing wrapper around the exact LingBot World 2 Wan2.1 VAE computation."""
|
||||
|
||||
handles_latent_denorm = True
|
||||
|
||||
def __init__(self, config, checkpoint_path=None, dtype=torch.float):
|
||||
"""Load the official LingBot World 2 VAE weights and expose FastVideo VAE APIs."""
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.dtype = dtype
|
||||
z_dim = int(getattr(config, "z_dim", 16))
|
||||
mean = torch.tensor(getattr(config, "latents_mean"), dtype=dtype)
|
||||
std = torch.tensor(getattr(config, "latents_std"), dtype=dtype)
|
||||
self.register_buffer("shift_factor", mean.view(1, z_dim, 1, 1, 1), persistent=False)
|
||||
self.register_buffer("scaling_factor", (1.0 / std).view(1, z_dim, 1, 1, 1), persistent=False)
|
||||
self.scale = [mean, 1.0 / std]
|
||||
|
||||
if checkpoint_path is None:
|
||||
self.model = WanVAE_(dim=96, z_dim=z_dim, dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2, attn_scales=[],
|
||||
temperal_downsample=[False, True, True],
|
||||
dropout=0.0)
|
||||
else:
|
||||
self.model = _video_vae(pretrained_path=checkpoint_path, z_dim=z_dim)
|
||||
self.model.eval().requires_grad_(False)
|
||||
|
||||
def _scale_for(self, device: torch.device) -> list[torch.Tensor]:
|
||||
"""Return source VAE scale tensors on the active device."""
|
||||
return [u.to(device) for u in self.scale]
|
||||
|
||||
def encode(self, videos: torch.Tensor):
|
||||
"""Encode `[B,C,T,H,W]` videos and return a FastVideo-style mean tensor."""
|
||||
if videos.ndim != 5:
|
||||
raise ValueError(f"LingBotWorld2WanVAE.encode expects 5D input, got {videos.shape}")
|
||||
scale = self._scale_for(videos.device)
|
||||
latents = []
|
||||
with amp.autocast(dtype=self.dtype):
|
||||
for video in videos:
|
||||
latent = self.model.encode(video.unsqueeze(0), scale).float().squeeze(0)
|
||||
latents.append(latent)
|
||||
normalized = torch.stack(latents, dim=0)
|
||||
return SimpleNamespace(mean=normalized)
|
||||
|
||||
def decode(self, latents: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode normalized LingBot World 2 latents to clamped `[-1,1]` video tensors."""
|
||||
if latents.ndim != 5:
|
||||
raise ValueError(f"LingBotWorld2WanVAE.decode expects 5D input, got {latents.shape}")
|
||||
scale = self._scale_for(latents.device)
|
||||
videos = []
|
||||
with amp.autocast(dtype=self.dtype):
|
||||
for latent in latents:
|
||||
video = self.model.decode(latent.unsqueeze(0), scale).float().clamp_(-1, 1).squeeze(0)
|
||||
videos.append(video)
|
||||
return torch.stack(videos, dim=0)
|
||||
|
||||
|
||||
EntryClass = LingBotWorld2WanVAE
|
||||
@@ -94,6 +94,9 @@ class OobleckDecoderBlock(nn.Module):
|
||||
input_dim, output_dim,
|
||||
kernel_size=2 * stride, stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
# Clean L*stride upsample for both parities; a no-op (0) for even
|
||||
# strides (Stable Audio), needed for odd strides (Cosmos3: 5).
|
||||
output_padding=stride % 2,
|
||||
))
|
||||
self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1)
|
||||
self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3)
|
||||
|
||||
@@ -35,7 +35,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs,
|
||||
"""
|
||||
# Get pipeline type
|
||||
model_path = fastvideo_args.model_path
|
||||
model_path = maybe_download_model(model_path)
|
||||
model_path = maybe_download_model(model_path, revision=fastvideo_args.revision)
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
|
||||
|
||||
@@ -0,0 +1,708 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native Cosmos3 video pipeline (T2V / I2V / T2I).
|
||||
|
||||
This replaces the earlier vllm-omni-derived skeleton with a native, stage-based
|
||||
:class:`ComposedPipelineBase` pipeline that wires the framework-parity-verified
|
||||
Cosmos3 components:
|
||||
|
||||
* tokenizer: Qwen2 ``Qwen2TokenizerFast`` + chat template (the only allowed
|
||||
third-party model-adjacent dependency; tokenizers are explicitly permitted),
|
||||
* VAE: FastVideo-native ``AutoencoderKLWan`` (Wan2.2) via ``Cosmos3VAEConfig``;
|
||||
encode normalizes ``(mu - mean) * inv_std`` and decode denormalizes + clamps,
|
||||
* sequence-packing: :func:`pack_cosmos3_video_sequence` (native, parity-tested),
|
||||
* DiT: ``Cosmos3VFMTransformer`` (native, bit-identical to the framework),
|
||||
* scheduler: FastVideo-native ``UniPCMultistepScheduler`` configured for pure
|
||||
flow matching (``flow_prediction`` + ``use_flow_sigmas``), numerically
|
||||
equivalent to the framework's ``FlowUniPCMultistepScheduler`` (parity-tested
|
||||
in ``test_cosmos3_scheduler_parity``).
|
||||
|
||||
The denoise/CFG glue is a faithful port of the framework's
|
||||
``Cosmos3OmniDiffusersPipeline`` math (mirrored in the framework-equivalent
|
||||
``diffusers_cosmos3.pipeline``): per UniPC timestep, run a SEQUENTIAL conditional
|
||||
then unconditional pass (each repacks the sequence with the prompt / negative
|
||||
prompt token ids, forwards the DiT, and zeros the prediction on conditioning
|
||||
frames), then combine ``v = uncond + guidance * (cond - uncond)`` and take one
|
||||
``scheduler.step(model_output=v, timestep, sample=latent)``. ``timestep_scale``
|
||||
is applied to the per-token timesteps *inside* the DiT (its ``forward`` already
|
||||
multiplies ``vision_timesteps * timestep_scale`` before the time embedder), so
|
||||
the loop passes raw scheduler timesteps to the packer.
|
||||
|
||||
The pure denoise math lives in :class:`Cosmos3DenoiseEngine` and the free
|
||||
function :func:`cosmos3_get_cfg_velocity` so it can be unit-/parity-tested
|
||||
directly against the framework oracle without constructing the full pipeline.
|
||||
|
||||
No diffusers/transformers *model* classes are imported at runtime here; only the
|
||||
Qwen2 tokenizer (loaded by the component loader) and the UniPC scheduler are
|
||||
third-party, both explicitly allowed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_unipc_multistep import (
|
||||
UniPCMultistepScheduler, )
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3ActionItem,
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3SoundItem,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# System prompts, verbatim from the framework (diffusers_cosmos3.pipeline).
|
||||
_SYSTEM_PROMPT_IMAGE = "You are a helpful assistant who will generate images from a give prompt."
|
||||
_SYSTEM_PROMPT_VIDEO = "You are a helpful assistant who will generate videos from a give prompt."
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Special-token resolution (Qwen2 chat tokenizer)
|
||||
# ===========================================================================
|
||||
def cosmos3_special_tokens(tokenizer: Any) -> dict[str, int]:
|
||||
"""Resolve the Cosmos3 generation special tokens from a Qwen2 tokenizer.
|
||||
|
||||
Mirrors the framework's ``llm_special_tokens``:
|
||||
``start_of_generation=<|vision_start|>``, ``end_of_generation=<|vision_end|>``,
|
||||
``eos_token_id=tokenizer.eos_token_id``.
|
||||
"""
|
||||
return {
|
||||
"start_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_start|>")),
|
||||
"end_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_end|>")),
|
||||
"eos_token_id": int(tokenizer.eos_token_id),
|
||||
}
|
||||
|
||||
|
||||
def cosmos3_tokenize_caption(
|
||||
tokenizer: Any,
|
||||
caption: str,
|
||||
*,
|
||||
is_video: bool = False,
|
||||
use_system_prompt: bool = False,
|
||||
) -> list[int]:
|
||||
"""Tokenize a caption with the Qwen2 chat template (framework-faithful).
|
||||
|
||||
Optionally prepends an image/video system prompt; always adds the
|
||||
generation prompt and disables ``add_vision_id`` (matching the framework's
|
||||
``tokenize_caption``).
|
||||
"""
|
||||
conversations: list[dict[str, str]] = []
|
||||
if use_system_prompt:
|
||||
conversations.append({
|
||||
"role": "system",
|
||||
"content": _SYSTEM_PROMPT_VIDEO if is_video else _SYSTEM_PROMPT_IMAGE,
|
||||
})
|
||||
conversations.append({"role": "user", "content": caption})
|
||||
token_ids = tokenizer.apply_chat_template(
|
||||
conversations,
|
||||
tokenize=True,
|
||||
add_generation_prompt=True,
|
||||
add_vision_id=False,
|
||||
)
|
||||
return list(token_ids)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Reasoning (VLM text generation) — und (causal) pathway + lm_head
|
||||
# ===========================================================================
|
||||
def cosmos3_generate_reasoner_text(
|
||||
transformer: Any,
|
||||
input_ids: list[int],
|
||||
max_new_tokens: int,
|
||||
*,
|
||||
eos_token_id: int | list[int] | None = None,
|
||||
) -> list[int]:
|
||||
"""Greedy text reasoning via the und (causal) backbone + ``lm_head``.
|
||||
|
||||
Mirrors the framework ``generate_reasoner_text`` (text-only prefill, greedy):
|
||||
only the und-pathway weights (no ``_moe_gen``) + ``embed_tokens`` / ``norm`` /
|
||||
``lm_head`` participate; the generation pathway and the VFM multimodal
|
||||
embedders are bypassed (no vision/sound/action tokens). Token-for-token
|
||||
identical to the framework reasoner (``test_cosmos3_reasoning_parity``).
|
||||
|
||||
Re-prefills each step (no KV cache) — correctness-first; a KV-cache fast path
|
||||
is a later optimization. Returns the newly generated token ids.
|
||||
"""
|
||||
device = next(transformer.parameters()).device
|
||||
ids = [int(x) for x in input_ids]
|
||||
eos: set[int] = set()
|
||||
if eos_token_id is not None:
|
||||
eos = {int(eos_token_id)} if isinstance(eos_token_id, int) else {int(x) for x in eos_token_id}
|
||||
|
||||
new_tokens: list[int] = []
|
||||
for _ in range(int(max_new_tokens)):
|
||||
n = len(ids)
|
||||
pos = torch.arange(n).unsqueeze(0).expand(3, -1).contiguous().to(device)
|
||||
out = transformer(
|
||||
text_ids=torch.tensor(ids, device=device, dtype=torch.long),
|
||||
text_indexes=torch.arange(n, device=device),
|
||||
position_ids=pos,
|
||||
sequence_length=n,
|
||||
split_lens=[n],
|
||||
attn_modes=["causal"],
|
||||
vision_tokens=[],
|
||||
vision_token_shapes=[],
|
||||
vision_sequence_indexes=torch.empty(0, dtype=torch.long, device=device),
|
||||
vision_timesteps=torch.empty(0, device=device),
|
||||
vision_mse_loss_indexes=torch.empty(0, dtype=torch.long, device=device),
|
||||
vision_noisy_frame_indexes=[],
|
||||
)
|
||||
logits = transformer.lm_head(out["last_hidden_state"][n - 1]) # [vocab]
|
||||
nxt = int(logits.argmax().item())
|
||||
ids.append(nxt)
|
||||
new_tokens.append(nxt)
|
||||
if nxt in eos:
|
||||
break
|
||||
return new_tokens
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# VAE encode/decode bridge (normalize / denormalize, matching the framework)
|
||||
# ===========================================================================
|
||||
@dataclass
|
||||
class _VaeNorm:
|
||||
"""Cached ``mean`` / ``inv_std`` for VAE (de)normalization."""
|
||||
|
||||
mean: torch.Tensor # [z_dim]
|
||||
inv_std: torch.Tensor # [z_dim]
|
||||
|
||||
@classmethod
|
||||
def from_vae(cls, vae: Any, dtype: torch.dtype) -> _VaeNorm:
|
||||
mean = torch.tensor(list(vae.config.latents_mean), dtype=dtype)
|
||||
std = torch.tensor(list(vae.config.latents_std), dtype=dtype)
|
||||
return cls(mean=mean, inv_std=1.0 / std)
|
||||
|
||||
|
||||
def cosmos3_vae_encode(vae: Any, video: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
|
||||
"""Encode ``[B, 3, T, H, W]`` pixels in [-1, 1] to NORMALIZED latents.
|
||||
|
||||
Matches the framework ``DiffusersWan22VAE.encode``: take the posterior mode
|
||||
and apply ``(mu - mean) * inv_std``. FastVideo's ``AutoencoderKLWan.encode``
|
||||
returns a ``DiagonalGaussianDistribution``; we read ``.mode()``.
|
||||
"""
|
||||
in_dtype = video.dtype
|
||||
device = video.device
|
||||
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
raw_mu = vae.encode(video).mode()
|
||||
return ((raw_mu - mean) * inv_std).to(in_dtype)
|
||||
|
||||
|
||||
def cosmos3_vae_decode(vae: Any, latents: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
|
||||
"""Decode NORMALIZED latents ``[B, z, T, H, W]`` to pixels ``[B, 3, T, H, W]``.
|
||||
|
||||
Inverts the normalization (``z / inv_std + mean``) then calls
|
||||
``vae.decode`` (which already clamps to [-1, 1]).
|
||||
"""
|
||||
in_dtype = latents.dtype
|
||||
device = latents.device
|
||||
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
z_raw = latents / inv_std + mean
|
||||
out = vae.decode(z_raw)
|
||||
if isinstance(out, tuple):
|
||||
out = out[0]
|
||||
if hasattr(out, "sample"):
|
||||
out = out.sample
|
||||
return out.to(in_dtype)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Per-vision-item packing geometry
|
||||
# ===========================================================================
|
||||
@dataclass
|
||||
class Cosmos3VisionSpec:
|
||||
"""Geometry + conditioning for one vision item in a denoise run.
|
||||
|
||||
Args:
|
||||
condition_frame_indexes: Latent-frame indices kept clean.
|
||||
shape: ``(C, T, H, W)`` of the latent for this item.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int, int, int]
|
||||
condition_frame_indexes: list[int]
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Pure denoise/CFG math (parity oracle target)
|
||||
# ===========================================================================
|
||||
def _split_flat_latent(flat: torch.Tensor, specs: list[Any]) -> list[torch.Tensor]:
|
||||
"""Split a flat vector into per-item tensors via each spec's ``numel``/``shape``.
|
||||
|
||||
Shared by vision (``[C, T, H, W]``), sound (``[C, T]``), and action
|
||||
(``[T, D]``) specs — every spec exposes ``numel`` and ``shape``.
|
||||
"""
|
||||
out: list[torch.Tensor] = []
|
||||
offset = 0
|
||||
for spec in specs:
|
||||
out.append(flat[offset:offset + spec.numel].reshape(spec.shape))
|
||||
offset += spec.numel
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundSpec:
|
||||
"""Geometry + conditioning for one sound item in a denoise run.
|
||||
|
||||
Args:
|
||||
shape: ``(C, T)`` of the sound latent (channels, temporal frames).
|
||||
condition_frame_indexes: Latent-frame indices kept clean (``[]`` for t2vs).
|
||||
fps: Sound latent FPS (``sound_latent_fps``); used iff fps modulation is on.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int]
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ActionSpec:
|
||||
"""Geometry + conditioning for one action item in a denoise run.
|
||||
|
||||
Args:
|
||||
shape: ``(T, action_dim)`` of the action latent.
|
||||
condition_frame_indexes: Frame indices kept clean (conditioning actions).
|
||||
domain_id: Embodiment domain id for the domain-aware action projection.
|
||||
fps: Action FPS; used iff fps modulation is on.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int]
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
domain_id: int = 0
|
||||
fps: float | None = None
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
def cosmos3_get_cfg_velocity(
|
||||
*,
|
||||
transformer: Any,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
special_tokens: dict[str, int],
|
||||
latent_patch_size: int,
|
||||
temporal_modality_margin: int,
|
||||
reset_spatial_ids: bool,
|
||||
enable_fps_modulation: bool,
|
||||
base_fps: float,
|
||||
temporal_compression_factor: int,
|
||||
include_end_of_generation_token: bool = False,
|
||||
fps_per_item: list[float] | None = None,
|
||||
normalize_cfg: bool = False,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Sequential-CFG velocity for one denoise step (framework math).
|
||||
|
||||
Replicates the framework ``get_cfg_velocity``:
|
||||
|
||||
1. split ``flat_latent`` into per-vision-item ``[C, T, H, W]`` latents,
|
||||
2. run a conditional pass (prompt tokens) and an unconditional pass
|
||||
(negative-prompt tokens); each repacks via
|
||||
:func:`pack_cosmos3_video_sequence`, forwards the DiT to obtain
|
||||
``preds_vision`` (a list of ``[1, C, T, H, W]`` unpatchified noisy-frame
|
||||
predictions), and zeros the prediction on conditioning frames
|
||||
(``pred * (1 - condition_mask)``),
|
||||
3. combine ``v = uncond + guidance * (cond - uncond)`` (optionally
|
||||
norm-rescaled), returned flattened to match ``flat_latent``.
|
||||
|
||||
``timestep`` is a scalar tensor (raw scheduler timestep); ``timestep_scale``
|
||||
is applied inside the DiT, so it is passed through unscaled here.
|
||||
"""
|
||||
assert timestep.numel() == 1, "timestep must be a scalar"
|
||||
timestep_value = float(timestep.reshape(()).item())
|
||||
|
||||
# Combined flat layout: [all vision | all action | all sound], matching the
|
||||
# framework per-sample concat order ([vision_i | action_i | sound_i]); single
|
||||
# sample here.
|
||||
vision_total = sum(spec.numel for spec in specs)
|
||||
action_total = sum(spec.numel for spec in action_specs) if action_specs else 0
|
||||
noise_x_vision = _split_flat_latent(flat_latent[:vision_total], specs)
|
||||
noise_x_action = (_split_flat_latent(flat_latent[vision_total:vision_total +
|
||||
action_total], action_specs) if action_specs else None)
|
||||
noise_x_sound = (_split_flat_latent(flat_latent[vision_total +
|
||||
action_total:], sound_specs) if sound_specs else None)
|
||||
device = next(transformer.parameters()).device
|
||||
|
||||
def _run(token_ids: list[int]) -> torch.Tensor:
|
||||
sound_items: list[Cosmos3SoundItem] = []
|
||||
if sound_specs is not None and noise_x_sound is not None:
|
||||
sound_items = [
|
||||
Cosmos3SoundItem(
|
||||
latent=noise_x_sound[i],
|
||||
condition_frame_indexes=list(ss.condition_frame_indexes),
|
||||
fps=(sound_fps_per_item[i] if sound_fps_per_item is not None else None),
|
||||
) for i, ss in enumerate(sound_specs)
|
||||
]
|
||||
action_items: list[Cosmos3ActionItem] = []
|
||||
if action_specs is not None and noise_x_action is not None:
|
||||
action_items = [
|
||||
Cosmos3ActionItem(
|
||||
latent=noise_x_action[i],
|
||||
condition_frame_indexes=list(asp.condition_frame_indexes),
|
||||
domain_id=asp.domain_id,
|
||||
fps=(action_fps_per_item[i] if action_fps_per_item is not None else None),
|
||||
) for i, asp in enumerate(action_specs)
|
||||
]
|
||||
samples = [
|
||||
Cosmos3SampleInputs(
|
||||
text_ids=list(token_ids),
|
||||
vision=Cosmos3VisionItem(
|
||||
latent=latent,
|
||||
condition_frame_indexes=list(spec.condition_frame_indexes),
|
||||
fps=(fps_per_item[i] if fps_per_item is not None else None),
|
||||
),
|
||||
sound=(sound_items[i] if i < len(sound_items) else None),
|
||||
action=(action_items[i] if i < len(action_items) else None),
|
||||
timestep=timestep_value,
|
||||
) for i, (latent, spec) in enumerate(zip(noise_x_vision, specs, strict=False))
|
||||
]
|
||||
packed = pack_cosmos3_video_sequence(
|
||||
samples,
|
||||
special_tokens,
|
||||
latent_patch_size=latent_patch_size,
|
||||
include_end_of_generation_token=include_end_of_generation_token,
|
||||
temporal_modality_margin=temporal_modality_margin,
|
||||
reset_spatial_ids=reset_spatial_ids,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=temporal_compression_factor,
|
||||
)
|
||||
out = transformer(**packed.to_dit_kwargs(device=device))
|
||||
|
||||
# Vision velocity: zero on conditioning frames, per item, flattened.
|
||||
vision_vel = torch.zeros(vision_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds = out.get("preds_vision")
|
||||
if preds is not None:
|
||||
items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds, packed.vision_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 5 else pred # [C, T, H, W]
|
||||
keep = (1.0 - cond_mask).to(dtype=pred.dtype, device=pred.device) # [T,1,1]
|
||||
items.append(pred * keep if keep.sum() > 0 else torch.zeros_like(pred))
|
||||
vision_vel = torch.cat([v.reshape(-1) for v in items]).to(flat_latent.dtype)
|
||||
|
||||
parts = [vision_vel]
|
||||
|
||||
if action_specs:
|
||||
# Action velocity: preds_action are per-item [T, D], already zero on
|
||||
# clean frames; zero on cond frames defensively.
|
||||
action_vel = torch.zeros(action_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds_a = out.get("preds_action")
|
||||
if preds_a is not None:
|
||||
a_items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds_a, packed.action_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [T, D]
|
||||
keep = (1.0 - cond_mask).reshape(-1, 1).to(dtype=pred.dtype, device=pred.device) # [T, 1]
|
||||
a_items.append(pred * keep)
|
||||
action_vel = torch.cat([v.reshape(-1) for v in a_items]).to(flat_latent.dtype)
|
||||
parts.append(action_vel)
|
||||
|
||||
if sound_specs:
|
||||
# Sound velocity: preds_sound are per-item [C, T], already zero on clean
|
||||
# frames (unpack fills only noisy frames); zero on cond frames defensively.
|
||||
sound_total = sum(spec.numel for spec in sound_specs)
|
||||
sound_vel = torch.zeros(sound_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds_s = out.get("preds_sound")
|
||||
if preds_s is not None:
|
||||
s_items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds_s, packed.sound_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [C, T]
|
||||
keep = (1.0 - cond_mask).reshape(1, -1).to(dtype=pred.dtype, device=pred.device) # [1, T]
|
||||
s_items.append(pred * keep)
|
||||
sound_vel = torch.cat([v.reshape(-1) for v in s_items]).to(flat_latent.dtype)
|
||||
parts.append(sound_vel)
|
||||
|
||||
return vision_vel if len(parts) == 1 else torch.cat(parts)
|
||||
|
||||
cond_v = _run(cond_token_ids)
|
||||
uncond_v = _run(uncond_token_ids)
|
||||
v_pred = uncond_v + guidance * (cond_v - uncond_v)
|
||||
if normalize_cfg:
|
||||
scale = (torch.norm(cond_v) / (torch.norm(v_pred) + 1e-8)).clamp(min=0.0, max=1.0)
|
||||
v_pred = v_pred * scale
|
||||
return v_pred
|
||||
|
||||
|
||||
class Cosmos3DenoiseEngine:
|
||||
"""Stateless denoise driver tying CFG velocity to UniPC stepping.
|
||||
|
||||
Holds the transformer + scheduler + packing constants and runs the full
|
||||
UniPC denoise loop. Kept separate from the pipeline so it can be exercised
|
||||
in isolation (smoke + parity tests) with stub or real components.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transformer: Any,
|
||||
scheduler: Any,
|
||||
special_tokens: dict[str, int],
|
||||
latent_patch_size: int,
|
||||
temporal_modality_margin: int,
|
||||
reset_spatial_ids: bool,
|
||||
enable_fps_modulation: bool,
|
||||
base_fps: float,
|
||||
temporal_compression_factor: int,
|
||||
include_end_of_generation_token: bool = False,
|
||||
) -> None:
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.special_tokens = special_tokens
|
||||
self.latent_patch_size = latent_patch_size
|
||||
self.temporal_modality_margin = temporal_modality_margin
|
||||
self.reset_spatial_ids = reset_spatial_ids
|
||||
self.enable_fps_modulation = enable_fps_modulation
|
||||
self.base_fps = base_fps
|
||||
self.temporal_compression_factor = temporal_compression_factor
|
||||
self.include_end_of_generation_token = include_end_of_generation_token
|
||||
|
||||
def velocity(
|
||||
self,
|
||||
*,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
fps_per_item: list[float] | None = None,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
return cosmos3_get_cfg_velocity(
|
||||
transformer=self.transformer,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
specs=specs,
|
||||
cond_token_ids=cond_token_ids,
|
||||
uncond_token_ids=uncond_token_ids,
|
||||
special_tokens=self.special_tokens,
|
||||
latent_patch_size=self.latent_patch_size,
|
||||
temporal_modality_margin=self.temporal_modality_margin,
|
||||
reset_spatial_ids=self.reset_spatial_ids,
|
||||
enable_fps_modulation=self.enable_fps_modulation,
|
||||
base_fps=self.base_fps,
|
||||
temporal_compression_factor=self.temporal_compression_factor,
|
||||
include_end_of_generation_token=self.include_end_of_generation_token,
|
||||
fps_per_item=fps_per_item,
|
||||
sound_specs=sound_specs,
|
||||
sound_fps_per_item=sound_fps_per_item,
|
||||
action_specs=action_specs,
|
||||
action_fps_per_item=action_fps_per_item,
|
||||
)
|
||||
|
||||
def denoise(
|
||||
self,
|
||||
*,
|
||||
flat_latent: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
fps_per_item: list[float] | None = None,
|
||||
progress_bar: Any | None = None,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run the full UniPC denoise loop, returning the final flat latent.
|
||||
|
||||
For each timestep: compute the sequential-CFG velocity, then
|
||||
``scheduler.step(model_output=v, timestep, sample=latent.unsqueeze(0))``
|
||||
(the framework steps with a leading batch axis), squeezing back to flat.
|
||||
For t2vs the flat latent is ``[vision | sound]`` and the velocity covers
|
||||
both; the scheduler steps the combined vector jointly.
|
||||
"""
|
||||
latent = flat_latent
|
||||
iterator = progress_bar(timesteps) if progress_bar is not None else timesteps
|
||||
for t in iterator:
|
||||
v_pred = self.velocity(
|
||||
flat_latent=latent,
|
||||
timestep=t.reshape(1),
|
||||
guidance=guidance,
|
||||
specs=specs,
|
||||
cond_token_ids=cond_token_ids,
|
||||
uncond_token_ids=uncond_token_ids,
|
||||
fps_per_item=fps_per_item,
|
||||
sound_specs=sound_specs,
|
||||
sound_fps_per_item=sound_fps_per_item,
|
||||
action_specs=action_specs,
|
||||
action_fps_per_item=action_fps_per_item,
|
||||
)
|
||||
stepped = self.scheduler.step(
|
||||
model_output=v_pred,
|
||||
timestep=t,
|
||||
sample=latent.unsqueeze(0),
|
||||
return_dict=False,
|
||||
)[0]
|
||||
latent = stepped.squeeze(0)
|
||||
return latent
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Pipeline (ComposedPipelineBase)
|
||||
# ===========================================================================
|
||||
class Cosmos3OmniDiffusersPipeline(ComposedPipelineBase):
|
||||
"""Cosmos3 video generation pipeline (T2V / I2V / T2I).
|
||||
|
||||
Stage-based ``ComposedPipelineBase`` pipeline. The required modules
|
||||
(``transformer`` / ``vae`` / ``scheduler`` / ``text_tokenizer``) are loaded
|
||||
from the ``nvidia/Cosmos3-Nano`` checkpoint by the component loader. The
|
||||
class name matches the checkpoint ``model_index.json`` ``_class_name`` so
|
||||
the registry resolves it directly.
|
||||
|
||||
The denoise/CFG/VAE math is delegated to module-level helpers
|
||||
(:func:`cosmos3_get_cfg_velocity`, :class:`Cosmos3DenoiseEngine`,
|
||||
:func:`cosmos3_vae_encode` / :func:`cosmos3_vae_decode`) which are
|
||||
framework-parity tested in ``tests/local_tests/cosmos3``.
|
||||
"""
|
||||
|
||||
is_video_pipeline = True
|
||||
# ``vision_encoder`` / ``sound_tokenizer`` ship in the checkpoint but the
|
||||
# video path does not need them; they are intentionally omitted here.
|
||||
_required_config_modules = ["text_tokenizer", "vae", "transformer", "scheduler"]
|
||||
|
||||
# Engine-init flow_shift (T2V/I2V); T2I overrides to 3.0 per request.
|
||||
_engine_init_flow_shift: float = 1.0
|
||||
# Class-attribute defaults so ``__new__``-based unit tests can read these
|
||||
# before ``initialize_pipeline`` runs.
|
||||
scheduler: Any = None
|
||||
_base_scheduler_config: Any = None
|
||||
_current_flow_shift: float | None = None
|
||||
|
||||
@staticmethod
|
||||
def _flow_scheduler_config(config: Any) -> dict[str, Any]:
|
||||
"""Coerce a loaded UniPC config to the framework's flow-matching setup.
|
||||
|
||||
The checkpoint ``scheduler_config.json`` carries diffusers-style fields
|
||||
(``use_karras_sigmas=True``, ``sigma_min``/``sigma_max``, beta schedule)
|
||||
that do not describe the framework sampler. The framework uses
|
||||
``FlowUniPCMultistepScheduler`` (pure flow matching: ``shift`` +
|
||||
``num_train_timesteps`` only). FastVideo's vendored UniPC checks
|
||||
``use_karras_sigmas`` *before* ``use_flow_sigmas``, so leaving karras on
|
||||
builds diffusion-style sigmas and the denoise diverges to NaN. Force the
|
||||
flow config here (parity-verified in ``test_cosmos3_scheduler_parity``).
|
||||
"""
|
||||
cfg = dict(config)
|
||||
cfg.update(
|
||||
use_karras_sigmas=False,
|
||||
use_exponential_sigmas=False,
|
||||
use_beta_sigmas=False,
|
||||
use_flow_sigmas=True,
|
||||
prediction_type="flow_prediction",
|
||||
predict_x0=True,
|
||||
final_sigmas_type="zero",
|
||||
)
|
||||
return cfg
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Bind the loaded scheduler + snapshot its config so per-request
|
||||
flow_shift rebuilds are cheap and the engine-init shift is applied."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
engine_shift = getattr(pipeline_config, "flow_shift", None)
|
||||
if engine_shift is not None:
|
||||
self._engine_init_flow_shift = float(engine_shift)
|
||||
scheduler = self.get_module("scheduler")
|
||||
if scheduler is not None:
|
||||
# Rebuild from a flow-coerced config so the runtime scheduler matches
|
||||
# the framework sampler (the loaded checkpoint config is diffusers-style).
|
||||
flow_config = self._flow_scheduler_config(scheduler.config)
|
||||
self.scheduler = UniPCMultistepScheduler.from_config(flow_config)
|
||||
if isinstance(self.modules, dict):
|
||||
self.modules["scheduler"] = self.scheduler
|
||||
self._base_scheduler_config = self.scheduler.config
|
||||
self._current_flow_shift = float(getattr(self.scheduler.config, "flow_shift", 1.0))
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Wire the Cosmos3 stages.
|
||||
|
||||
The whole text->latent->denoise->decode flow is custom (sequential CFG
|
||||
with per-pass repacking), so a single :class:`Cosmos3DenoisingStage`
|
||||
owns it. ``InputValidationStage`` runs first for the standard checks.
|
||||
"""
|
||||
from fastvideo.pipelines.stages import InputValidationStage
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=Cosmos3DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
tokenizer=self.get_module("text_tokenizer"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
# -- Scheduler control --------------------------------------------------
|
||||
|
||||
def _set_flow_shift(self, target_shift: float) -> None:
|
||||
"""Set UniPC ``flow_shift`` to ``target_shift``.
|
||||
|
||||
Lazily builds a default UniPC scheduler when called before
|
||||
``initialize_pipeline`` (e.g. the ``__new__``-based scheduler-parity
|
||||
tests); otherwise rebuilds from the snapshotted base config only when
|
||||
the target differs from the current shift.
|
||||
"""
|
||||
target = float(target_shift)
|
||||
base_config = self._base_scheduler_config
|
||||
if base_config is None:
|
||||
self.scheduler = UniPCMultistepScheduler(
|
||||
num_train_timesteps=1000,
|
||||
solver_order=2,
|
||||
prediction_type="flow_prediction",
|
||||
use_flow_sigmas=True,
|
||||
flow_shift=target,
|
||||
)
|
||||
self._base_scheduler_config = self.scheduler.config
|
||||
self._current_flow_shift = target
|
||||
return
|
||||
current = self._current_flow_shift
|
||||
if current is not None and target == float(current):
|
||||
return
|
||||
self.scheduler = UniPCMultistepScheduler.from_config(base_config, flow_shift=target)
|
||||
if isinstance(self.modules, dict):
|
||||
self.modules["scheduler"] = self.scheduler
|
||||
self._current_flow_shift = target
|
||||
|
||||
# -- Tokenization -------------------------------------------------------
|
||||
|
||||
def tokenize_caption(self, caption: str, *, is_video: bool = False, use_system_prompt: bool = False) -> list[int]:
|
||||
return cosmos3_tokenize_caption(self.get_module("text_tokenizer"),
|
||||
caption,
|
||||
is_video=is_video,
|
||||
use_system_prompt=use_system_prompt)
|
||||
|
||||
|
||||
# Entry point for the pipeline registry. The class name matches the checkpoint
|
||||
# ``model_index.json`` ``_class_name`` so ``resolve_pipeline_cls`` finds it.
|
||||
EntryClass = Cosmos3OmniDiffusersPipeline
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 (Cosmos3-Nano) inference presets.
|
||||
|
||||
Defaults track the official ``cosmos-framework`` ``sample_args`` for the video
|
||||
paths (``text2video`` / ``image2video``: guidance=6.0, num_steps=35, shift=10.0,
|
||||
fps=24, num_frames=189) and ``text2image`` (guidance=4.0, num_steps=50,
|
||||
shift=3.0). The default resolution is 16:9 at a VAE-aligned 704x1280 (spatial
|
||||
compression 16 -> 44x80 latent grid).
|
||||
"""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Cosmos3 sequential-CFG UniPC denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
# Framework video negative prompt (Cosmos quality prompt).
|
||||
COSMOS3_VIDEO_NEGATIVE_PROMPT = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, "
|
||||
"fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
|
||||
COSMOS3_NANO = InferencePreset(
|
||||
name="cosmos3_nano",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano text-to-video",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 189,
|
||||
"fps": 24,
|
||||
"guidance_scale": 6.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
COSMOS3_NANO_I2V = InferencePreset(
|
||||
name="cosmos3_nano_i2v",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano image-to-video",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 189,
|
||||
"fps": 24,
|
||||
"guidance_scale": 6.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
COSMOS3_NANO_T2I = InferencePreset(
|
||||
name="cosmos3_nano_t2i",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano text-to-image",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 24,
|
||||
"guidance_scale": 4.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": "",
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (COSMOS3_NANO, COSMOS3_NANO_I2V, COSMOS3_NANO_T2I)
|
||||
@@ -0,0 +1,549 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native Cosmos3 sequence packing (video subset).
|
||||
|
||||
Numerical-parity port of the official ``cosmos_framework`` data packer
|
||||
(``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``) restricted
|
||||
to the VIDEO generation path that the FastVideo Cosmos3 DiT consumes (T2V / I2V
|
||||
/ T2I). It builds, per sample, two splits:
|
||||
|
||||
* a ``causal`` text split (prompt token ids, plus the trailing ``eos`` and
|
||||
``start_of_generation`` markers the framework appends when a generation
|
||||
modality follows), and
|
||||
* a ``full`` vision split (VAE latent patch tokens).
|
||||
|
||||
The 3D-MRoPE position ids ``[3, seq]`` are produced exactly like the framework:
|
||||
text tokens broadcast a single monotone id across the (t, h, w) axes, the
|
||||
temporal offset is bumped by ``temporal_modality_margin`` at the text->vision
|
||||
boundary, and vision tokens lay out a (T, H, W) grid with spatial ids reset per
|
||||
segment. Condition frames (I2V cond frame 0, T2I single conditioned frame, ...)
|
||||
are kept in the packed sequence and rope grid but excluded from the MSE-loss /
|
||||
timestep bookkeeping, mirroring the framework.
|
||||
|
||||
The output ``Cosmos3PackedSequence`` maps 1:1 onto the
|
||||
``Cosmos3VFMTransformer.forward`` kwargs via :meth:`to_dit_kwargs`. This module
|
||||
is pure torch/python; it imports no diffusers/transformers model classes.
|
||||
|
||||
Reference of record: ``cosmos_framework`` (NVIDIA), the parity oracle used by
|
||||
``tests/local_tests/cosmos3/test_cosmos3_packing_parity.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.dits.cosmos3 import (
|
||||
compute_mrope_position_ids_text,
|
||||
compute_mrope_position_ids_vision,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Cosmos3VisionItem",
|
||||
"Cosmos3SampleInputs",
|
||||
"Cosmos3PackedSequence",
|
||||
"pack_cosmos3_video_sequence",
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inputs
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass
|
||||
class Cosmos3VisionItem:
|
||||
"""One vision latent for a sample.
|
||||
|
||||
Args:
|
||||
latent: VAE latent ``[C, T, H, W]`` (a leading batch axis of size 1 is
|
||||
accepted and squeezed).
|
||||
condition_frame_indexes: Latent-frame indices that are *conditioned*
|
||||
(clean) rather than noisy. ``[]`` for T2V, ``[0]`` for I2V, and the
|
||||
single conditioned frame for T2I.
|
||||
fps: Frames-per-second for this clip; only used when
|
||||
``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundItem:
|
||||
"""One sound latent for a sample (t2vs).
|
||||
|
||||
Args:
|
||||
latent: AVAE sound latent ``[C, T]`` (channels, temporal frames).
|
||||
condition_frame_indexes: Latent-frame indices that are *conditioned*
|
||||
(clean). ``[]`` for t2vs (all frames generated).
|
||||
fps: Sound latent FPS (``sound_latent_fps``, e.g. 25); only used when
|
||||
``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ActionItem:
|
||||
"""One action latent for a sample (action-conditioned world model).
|
||||
|
||||
Args:
|
||||
latent: Action latent ``[T, action_dim]`` (per-frame action vectors).
|
||||
condition_frame_indexes: Frame indices kept clean (conditioning actions).
|
||||
domain_id: Embodiment domain id (scalar / ``[1]``) for the
|
||||
domain-aware action projection.
|
||||
fps: Action FPS; only used when ``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
domain_id: int = 0
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SampleInputs:
|
||||
"""Per-sample packing inputs (text prompt + vision item, +sound, +action)."""
|
||||
|
||||
text_ids: list[int]
|
||||
vision: Cosmos3VisionItem
|
||||
timestep: float
|
||||
sound: Cosmos3SoundItem | None = None
|
||||
action: Cosmos3ActionItem | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Output
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass
|
||||
class Cosmos3PackedSequence:
|
||||
"""Packed-sequence inputs consumed by ``Cosmos3VFMTransformer.forward``.
|
||||
|
||||
Field names mirror the framework ``PackedSequence`` (+ its ``vision``
|
||||
``ModalityData``) so the parity test can compare field-by-field.
|
||||
"""
|
||||
|
||||
# Sequence structure.
|
||||
sample_lens: list[int]
|
||||
split_lens: list[int]
|
||||
attn_modes: list[str]
|
||||
sequence_length: int
|
||||
is_image_batch: bool
|
||||
|
||||
# Text modality.
|
||||
text_ids: torch.Tensor
|
||||
text_indexes: torch.Tensor
|
||||
position_ids: torch.Tensor # [3, sequence_length]
|
||||
|
||||
# Vision modality.
|
||||
vision_tokens: list[torch.Tensor]
|
||||
vision_token_shapes: list[tuple[int, int, int]]
|
||||
vision_sequence_indexes: torch.Tensor
|
||||
vision_timesteps: torch.Tensor
|
||||
vision_mse_loss_indexes: torch.Tensor
|
||||
vision_noisy_frame_indexes: list[torch.Tensor]
|
||||
vision_condition_mask: list[torch.Tensor]
|
||||
fps_vision: torch.Tensor | None = None
|
||||
|
||||
# Sound modality (t2vs); empty/None when no sound.
|
||||
sound_tokens: list[torch.Tensor] = field(default_factory=list)
|
||||
sound_token_shapes: list[tuple[int, int, int]] = field(default_factory=list)
|
||||
sound_sequence_indexes: torch.Tensor | None = None
|
||||
sound_timesteps: torch.Tensor | None = None
|
||||
sound_mse_loss_indexes: torch.Tensor | None = None
|
||||
sound_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
|
||||
sound_condition_mask: list[torch.Tensor] = field(default_factory=list)
|
||||
fps_sound: torch.Tensor | None = None
|
||||
|
||||
# Action modality (action-conditioned world model); empty/None when no action.
|
||||
action_tokens: list[torch.Tensor] = field(default_factory=list)
|
||||
action_token_shapes: list[tuple[int, ...]] = field(default_factory=list)
|
||||
action_sequence_indexes: torch.Tensor | None = None
|
||||
action_timesteps: torch.Tensor | None = None
|
||||
action_mse_loss_indexes: torch.Tensor | None = None
|
||||
action_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
|
||||
action_condition_mask: list[torch.Tensor] = field(default_factory=list)
|
||||
action_domain_id: list[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
def to_dit_kwargs(self, device: torch.device | str | None = None) -> dict[str, Any]:
|
||||
"""Return the kwargs dict for ``Cosmos3VFMTransformer.forward``.
|
||||
|
||||
Packing is device-agnostic (ids/indexes/position-ids are built on CPU).
|
||||
When ``device`` is given, every tensor input is moved to it so the DiT
|
||||
forward runs on a single device (e.g. the model's GPU at inference).
|
||||
"""
|
||||
|
||||
def _mv(x: Any) -> Any:
|
||||
return x.to(device) if (device is not None and torch.is_tensor(x)) else x
|
||||
|
||||
return dict(
|
||||
text_ids=_mv(self.text_ids),
|
||||
text_indexes=_mv(self.text_indexes),
|
||||
position_ids=_mv(self.position_ids),
|
||||
sequence_length=int(self.sequence_length),
|
||||
split_lens=list(self.split_lens),
|
||||
attn_modes=list(self.attn_modes),
|
||||
vision_tokens=[_mv(t) for t in self.vision_tokens],
|
||||
vision_token_shapes=list(self.vision_token_shapes),
|
||||
vision_sequence_indexes=_mv(self.vision_sequence_indexes),
|
||||
vision_timesteps=_mv(self.vision_timesteps),
|
||||
vision_mse_loss_indexes=_mv(self.vision_mse_loss_indexes),
|
||||
vision_noisy_frame_indexes=[_mv(t) for t in self.vision_noisy_frame_indexes],
|
||||
fps_vision=self.fps_vision,
|
||||
sound_tokens=[_mv(t) for t in self.sound_tokens],
|
||||
sound_token_shapes=list(self.sound_token_shapes),
|
||||
sound_sequence_indexes=_mv(self.sound_sequence_indexes),
|
||||
sound_timesteps=_mv(self.sound_timesteps),
|
||||
sound_mse_loss_indexes=_mv(self.sound_mse_loss_indexes),
|
||||
sound_noisy_frame_indexes=[_mv(t) for t in self.sound_noisy_frame_indexes],
|
||||
fps_sound=_mv(self.fps_sound),
|
||||
action_tokens=[_mv(t) for t in self.action_tokens],
|
||||
action_token_shapes=list(self.action_token_shapes),
|
||||
action_sequence_indexes=_mv(self.action_sequence_indexes),
|
||||
action_timesteps=_mv(self.action_timesteps),
|
||||
action_mse_loss_indexes=_mv(self.action_mse_loss_indexes),
|
||||
action_noisy_frame_indexes=[_mv(t) for t in self.action_noisy_frame_indexes],
|
||||
action_domain_id=[_mv(t) for t in self.action_domain_id],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Packing
|
||||
# ---------------------------------------------------------------------------
|
||||
def pack_cosmos3_video_sequence(
|
||||
samples: list[Cosmos3SampleInputs],
|
||||
special_tokens: dict[str, int],
|
||||
*,
|
||||
latent_patch_size: int = 2,
|
||||
include_end_of_generation_token: bool = False,
|
||||
temporal_modality_margin: int = 15_000,
|
||||
reset_spatial_ids: bool = True,
|
||||
enable_fps_modulation: bool = False,
|
||||
base_fps: float = 24.0,
|
||||
temporal_compression_factor: int = 4,
|
||||
initial_mrope_temporal_offset: int | float = 0,
|
||||
) -> Cosmos3PackedSequence:
|
||||
"""Pack prompts + vision latents into the Cosmos3 DiT packed-sequence inputs.
|
||||
|
||||
Video subset of ``cosmos_framework`` ``pack_input_sequence`` under
|
||||
``unified_3d_mrope``: each sample is ``[causal text, full vision]``.
|
||||
|
||||
Args:
|
||||
samples: Per-sample text prompt token ids + vision item + timestep.
|
||||
special_tokens: Must contain ``eos_token_id`` and
|
||||
``start_of_generation`` (and ``end_of_generation`` if
|
||||
``include_end_of_generation_token``). ``bos_token_id`` is honored if
|
||||
present (prepended) to match the framework.
|
||||
latent_patch_size: Latent patch size used by the DiT.
|
||||
include_end_of_generation_token: Append the framework's end-of-generation
|
||||
marker after the vision split.
|
||||
temporal_modality_margin: Temporal-offset bump applied at the
|
||||
text->vision boundary (``unified_3d_mrope_temporal_modality_margin``).
|
||||
reset_spatial_ids: Reset vision spatial ids to 0 per segment.
|
||||
enable_fps_modulation: Use float, fps-scaled temporal positions.
|
||||
base_fps: Base FPS used when ``enable_fps_modulation``.
|
||||
temporal_compression_factor: VAE temporal compression factor.
|
||||
initial_mrope_temporal_offset: Per-sample starting temporal offset.
|
||||
|
||||
Returns:
|
||||
A :class:`Cosmos3PackedSequence`.
|
||||
"""
|
||||
assert "eos_token_id" in special_tokens, "special_tokens must contain eos_token_id"
|
||||
assert "start_of_generation" in special_tokens, "special_tokens must contain start_of_generation"
|
||||
if latent_patch_size < 1:
|
||||
raise ValueError(f"latent_patch_size must be >= 1, got {latent_patch_size}")
|
||||
|
||||
# Build-time accumulators (concatenated across samples).
|
||||
sample_lens: list[int] = []
|
||||
split_lens: list[int] = []
|
||||
attn_modes: list[str] = []
|
||||
|
||||
text_ids: list[int] = []
|
||||
text_indexes: list[int] = []
|
||||
position_id_blocks: list[torch.Tensor] = [] # each [3, n]
|
||||
|
||||
vision_tokens: list[torch.Tensor] = []
|
||||
vision_token_shapes: list[tuple[int, int, int]] = []
|
||||
vision_sequence_indexes: list[int] = []
|
||||
vision_timesteps: list[float] = []
|
||||
vision_mse_loss_indexes: list[int] = []
|
||||
vision_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
vision_condition_mask: list[torch.Tensor] = []
|
||||
fps_values: list[float] = []
|
||||
|
||||
sound_tokens: list[torch.Tensor] = []
|
||||
sound_token_shapes: list[tuple[int, int, int]] = []
|
||||
sound_sequence_indexes: list[int] = []
|
||||
sound_timesteps: list[float] = []
|
||||
sound_mse_loss_indexes: list[int] = []
|
||||
sound_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
sound_condition_mask: list[torch.Tensor] = []
|
||||
sound_fps_values: list[float] = []
|
||||
|
||||
action_tokens: list[torch.Tensor] = []
|
||||
action_token_shapes: list[tuple[int, ...]] = []
|
||||
action_sequence_indexes: list[int] = []
|
||||
action_timesteps: list[float] = []
|
||||
action_mse_loss_indexes: list[int] = []
|
||||
action_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
action_condition_mask: list[torch.Tensor] = []
|
||||
action_domain_id: list[torch.Tensor] = []
|
||||
|
||||
curr = 0 # running position in the packed sequence
|
||||
is_image_batch = True
|
||||
|
||||
for sample in samples:
|
||||
temporal_offset: int | float = initial_mrope_temporal_offset
|
||||
sample_len = 0
|
||||
|
||||
# ---- 1. Text split (causal) ----
|
||||
if "bos_token_id" in special_tokens:
|
||||
shifted_text_ids = [special_tokens["bos_token_id"], *sample.text_ids]
|
||||
else:
|
||||
shifted_text_ids = list(sample.text_ids)
|
||||
# The video path always has a following generation modality, so the
|
||||
# framework appends eos + start_of_generation.
|
||||
shifted_text_ids = [*shifted_text_ids, special_tokens["eos_token_id"], special_tokens["start_of_generation"]]
|
||||
text_split_len = len(shifted_text_ids)
|
||||
|
||||
text_ids.extend(shifted_text_ids)
|
||||
text_indexes.extend(range(curr, curr + text_split_len))
|
||||
|
||||
text_mrope, temporal_offset = compute_mrope_position_ids_text(
|
||||
num_tokens=text_split_len,
|
||||
temporal_offset=int(temporal_offset),
|
||||
)
|
||||
position_id_blocks.append(text_mrope)
|
||||
|
||||
attn_modes.append("causal")
|
||||
split_lens.append(text_split_len)
|
||||
curr += text_split_len
|
||||
sample_len += text_split_len
|
||||
|
||||
# End of text modality: bump temporal offset before vision.
|
||||
temporal_offset += temporal_modality_margin
|
||||
# Sound shares the vision temporal start (parallel temporal positions).
|
||||
vision_start_temporal_offset = temporal_offset
|
||||
|
||||
# ---- 2. Vision split (full) ----
|
||||
latent = sample.vision.latent
|
||||
latent = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
|
||||
_c, latent_t, latent_h, latent_w = latent.shape
|
||||
patch_h = math.ceil(latent_h / latent_patch_size)
|
||||
patch_w = math.ceil(latent_w / latent_patch_size)
|
||||
num_vision_tokens = latent_t * patch_h * patch_w
|
||||
|
||||
vision_tokens.append(sample.vision.latent)
|
||||
vision_token_shapes.append((latent_t, patch_h, patch_w))
|
||||
vision_sequence_indexes.extend(range(curr, curr + num_vision_tokens))
|
||||
|
||||
condition_set = {idx for idx in sample.vision.condition_frame_indexes if 0 <= idx < latent_t}
|
||||
cond_mask = torch.zeros((latent_t, 1, 1), device=latent.device, dtype=latent.dtype)
|
||||
for frame_idx in condition_set:
|
||||
cond_mask[frame_idx, 0, 0] = 1.0
|
||||
vision_condition_mask.append(cond_mask)
|
||||
|
||||
noisy_frames = torch.tensor(
|
||||
[idx for idx in range(latent_t) if idx not in condition_set],
|
||||
device=latent.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
vision_noisy_frame_indexes.append(noisy_frames)
|
||||
|
||||
# MSE-loss indices + per-token timesteps cover only the noisy frames.
|
||||
frame_token_stride = patch_h * patch_w
|
||||
for frame_idx in range(latent_t):
|
||||
if frame_idx in condition_set:
|
||||
continue
|
||||
frame_start = curr + frame_idx * frame_token_stride
|
||||
vision_mse_loss_indexes.extend(range(frame_start, frame_start + frame_token_stride))
|
||||
vision_timesteps.extend([float(sample.timestep)] * frame_token_stride)
|
||||
|
||||
vision_fps = sample.vision.fps if enable_fps_modulation else None
|
||||
if vision_fps is not None:
|
||||
fps_values.append(float(vision_fps))
|
||||
vision_mrope, temporal_offset = compute_mrope_position_ids_vision(
|
||||
grid_t=latent_t,
|
||||
grid_h=patch_h,
|
||||
grid_w=patch_w,
|
||||
temporal_offset=temporal_offset,
|
||||
fps=vision_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=temporal_compression_factor,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
)
|
||||
position_id_blocks.append(vision_mrope)
|
||||
|
||||
curr += num_vision_tokens
|
||||
sample_len += num_vision_tokens
|
||||
|
||||
# ---- 2a2. Action split: shares the vision "full" split ----
|
||||
# Mirrors framework ``_pack_action_tokens``: action latent [T, D] -> T
|
||||
# tokens (token shape (T,)), domain-aware, 3D-MRoPE at the vision temporal
|
||||
# offset with ``start_frame_offset=1`` (parallel to vision; tcf=1; does
|
||||
# not advance the offset).
|
||||
action_split_len = 0
|
||||
if sample.action is not None:
|
||||
action_latent = sample.action.latent # [T, D]
|
||||
action_t = int(action_latent.shape[0])
|
||||
action_split_len = action_t
|
||||
|
||||
action_tokens.append(action_latent)
|
||||
action_token_shapes.append((action_t, ))
|
||||
action_sequence_indexes.extend(range(curr, curr + action_t))
|
||||
action_domain_id.append(torch.tensor([int(sample.action.domain_id)], dtype=torch.long))
|
||||
|
||||
a_cond_set = {idx for idx in sample.action.condition_frame_indexes if 0 <= idx < action_t}
|
||||
a_cond_mask = torch.zeros((action_t, 1), device=action_latent.device, dtype=action_latent.dtype)
|
||||
for fi in a_cond_set:
|
||||
a_cond_mask[fi, 0] = 1.0
|
||||
action_condition_mask.append(a_cond_mask)
|
||||
|
||||
a_noisy = torch.tensor([idx for idx in range(action_t) if idx not in a_cond_set],
|
||||
device=action_latent.device,
|
||||
dtype=torch.long)
|
||||
action_noisy_frame_indexes.append(a_noisy)
|
||||
|
||||
for fi in range(action_t):
|
||||
if fi in a_cond_set:
|
||||
continue
|
||||
action_mse_loss_indexes.append(curr + fi)
|
||||
action_timesteps.append(float(sample.timestep))
|
||||
|
||||
action_fps = sample.action.fps if enable_fps_modulation else None
|
||||
action_mrope, _ = compute_mrope_position_ids_vision(
|
||||
grid_t=action_t,
|
||||
grid_h=1,
|
||||
grid_w=1,
|
||||
temporal_offset=vision_start_temporal_offset,
|
||||
fps=action_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=1, # action is at frame rate
|
||||
base_temporal_compression_factor=temporal_compression_factor,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
start_frame_offset=1,
|
||||
)
|
||||
position_id_blocks.append(action_mrope)
|
||||
curr += action_t
|
||||
sample_len += action_t
|
||||
|
||||
# ---- 2b. Sound split (t2vs): shares the vision "full" split ----
|
||||
# Mirrors framework ``_pack_sound_tokens``: sound latent [C, T] -> T
|
||||
# tokens (token shape (T,1,1)), packed right after vision, with 3D-MRoPE
|
||||
# temporal positions starting at the vision temporal offset (parallel to
|
||||
# vision, start_frame_offset=0, tcf=1) and NOT advancing it.
|
||||
sound_split_len = 0
|
||||
if sample.sound is not None:
|
||||
sound_latent = sample.sound.latent
|
||||
sound_latent = sound_latent.squeeze(0) if sound_latent.dim() == 3 else sound_latent # [C, T]
|
||||
_sc, sound_t = sound_latent.shape
|
||||
sound_split_len = sound_t
|
||||
|
||||
sound_tokens.append(sound_latent)
|
||||
sound_token_shapes.append((sound_t, 1, 1))
|
||||
sound_sequence_indexes.extend(range(curr, curr + sound_t))
|
||||
|
||||
s_cond_set = {idx for idx in sample.sound.condition_frame_indexes if 0 <= idx < sound_t}
|
||||
s_cond_mask = torch.zeros((sound_t, 1), device=sound_latent.device, dtype=sound_latent.dtype)
|
||||
for fi in s_cond_set:
|
||||
s_cond_mask[fi, 0] = 1.0
|
||||
sound_condition_mask.append(s_cond_mask)
|
||||
|
||||
s_noisy = torch.tensor([idx for idx in range(sound_t) if idx not in s_cond_set],
|
||||
device=sound_latent.device,
|
||||
dtype=torch.long)
|
||||
sound_noisy_frame_indexes.append(s_noisy)
|
||||
|
||||
for fi in range(sound_t):
|
||||
if fi in s_cond_set:
|
||||
continue
|
||||
sound_mse_loss_indexes.append(curr + fi) # 1 token per sound frame
|
||||
sound_timesteps.append(float(sample.timestep))
|
||||
|
||||
sound_fps = sample.sound.fps if enable_fps_modulation else None
|
||||
if sound_fps is not None:
|
||||
sound_fps_values.append(float(sound_fps))
|
||||
sound_mrope, _ = compute_mrope_position_ids_vision(
|
||||
grid_t=sound_t,
|
||||
grid_h=1,
|
||||
grid_w=1,
|
||||
temporal_offset=vision_start_temporal_offset,
|
||||
fps=sound_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=1, # sound latent already at sound_latent_fps
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
start_frame_offset=0,
|
||||
)
|
||||
position_id_blocks.append(sound_mrope)
|
||||
curr += sound_t
|
||||
sample_len += sound_t
|
||||
|
||||
# ---- 3. Optional end-of-generation marker ----
|
||||
eov_len = 0
|
||||
if include_end_of_generation_token:
|
||||
assert "end_of_generation" in special_tokens, ("special_tokens must contain end_of_generation when "
|
||||
"include_end_of_generation_token=True")
|
||||
text_ids.append(special_tokens["end_of_generation"])
|
||||
text_indexes.append(curr)
|
||||
eov_dtype = torch.float32 if enable_fps_modulation else torch.long
|
||||
eov_ids = torch.full((3, 1), temporal_offset, dtype=eov_dtype)
|
||||
position_id_blocks.append(eov_ids)
|
||||
temporal_offset += 1
|
||||
curr += 1
|
||||
eov_len = 1
|
||||
sample_len += 1
|
||||
|
||||
# Vision + action + sound + any trailing eov marker share one "full" split.
|
||||
attn_modes.append("full")
|
||||
split_lens.append(num_vision_tokens + action_split_len + sound_split_len + eov_len)
|
||||
sample_lens.append(sample_len)
|
||||
|
||||
if latent_t != 1:
|
||||
is_image_batch = False
|
||||
|
||||
sequence_length = sum(sample_lens)
|
||||
|
||||
# position_ids: float iff any block is float (fps modulation path).
|
||||
any_float = any(b.dtype.is_floating_point for b in position_id_blocks)
|
||||
if any_float:
|
||||
position_id_blocks = [b.to(torch.float32) for b in position_id_blocks]
|
||||
position_ids = torch.cat(position_id_blocks, dim=1) # [3, sequence_length]
|
||||
|
||||
timesteps_dtype = torch.float32
|
||||
return Cosmos3PackedSequence(
|
||||
sample_lens=sample_lens,
|
||||
split_lens=split_lens,
|
||||
attn_modes=attn_modes,
|
||||
sequence_length=sequence_length,
|
||||
is_image_batch=is_image_batch,
|
||||
text_ids=torch.tensor(text_ids, dtype=torch.long),
|
||||
text_indexes=torch.tensor(text_indexes, dtype=torch.long),
|
||||
position_ids=position_ids,
|
||||
vision_tokens=vision_tokens,
|
||||
vision_token_shapes=vision_token_shapes,
|
||||
vision_sequence_indexes=torch.tensor(vision_sequence_indexes, dtype=torch.long),
|
||||
vision_timesteps=torch.tensor(vision_timesteps, dtype=timesteps_dtype),
|
||||
vision_mse_loss_indexes=torch.tensor(vision_mse_loss_indexes, dtype=torch.long),
|
||||
vision_noisy_frame_indexes=vision_noisy_frame_indexes,
|
||||
vision_condition_mask=vision_condition_mask,
|
||||
fps_vision=(torch.tensor(fps_values, dtype=torch.float32) if fps_values else None),
|
||||
sound_tokens=sound_tokens,
|
||||
sound_token_shapes=sound_token_shapes,
|
||||
sound_sequence_indexes=(torch.tensor(sound_sequence_indexes, dtype=torch.long) if sound_tokens else None),
|
||||
sound_timesteps=(torch.tensor(sound_timesteps, dtype=timesteps_dtype) if sound_tokens else None),
|
||||
sound_mse_loss_indexes=(torch.tensor(sound_mse_loss_indexes, dtype=torch.long) if sound_tokens else None),
|
||||
sound_noisy_frame_indexes=sound_noisy_frame_indexes,
|
||||
sound_condition_mask=sound_condition_mask,
|
||||
fps_sound=(torch.tensor(sound_fps_values, dtype=torch.float32) if sound_fps_values else None),
|
||||
action_tokens=action_tokens,
|
||||
action_token_shapes=action_token_shapes,
|
||||
action_sequence_indexes=(torch.tensor(action_sequence_indexes, dtype=torch.long) if action_tokens else None),
|
||||
action_timesteps=(torch.tensor(action_timesteps, dtype=timesteps_dtype) if action_tokens else None),
|
||||
action_mse_loss_indexes=(torch.tensor(action_mse_loss_indexes, dtype=torch.long) if action_tokens else None),
|
||||
action_noisy_frame_indexes=action_noisy_frame_indexes,
|
||||
action_condition_mask=action_condition_mask,
|
||||
action_domain_id=action_domain_id,
|
||||
)
|
||||
@@ -19,7 +19,7 @@ class GlmImageDecodingStage(DecodingStage):
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast = (vae_dtype != torch.float32 and not fastvideo_args.disable_autocast)
|
||||
|
||||
latents = self._denormalize_latents(latents)
|
||||
latents = self._denormalize_latents(latents, fastvideo_args)
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Dense LingBot-Video inference pipeline."""
|
||||
|
||||
from fastvideo.pipelines.basic.lingbot_video.lingbot_video_pipeline import LingBotVideoPipeline
|
||||
|
||||
__all__ = ["LingBotVideoPipeline"]
|
||||
@@ -0,0 +1,106 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage-composed LingBot-Video Dense and MoE/refiner T2V pipeline."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.basic.lingbot_video.stages import (
|
||||
LingBotVideoDenoisingStage,
|
||||
LingBotVideoInputValidationStage,
|
||||
LingBotVideoLatentPreparationStage,
|
||||
LingBotVideoRefinerPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
|
||||
|
||||
class LingBotVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""T2V pipeline with optional released MoE pixel-space refinement."""
|
||||
|
||||
is_video_pipeline = True
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
|
||||
|
||||
def load_modules(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
loaded_modules: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Load the optional refiner DiT and the VAE encoder only when declared."""
|
||||
model_index = self._load_config(self.model_path)
|
||||
required = list(type(self)._required_config_modules)
|
||||
load_refiner = "transformer_2" in model_index and getattr(fastvideo_args, "refine_enabled", None) is not False
|
||||
if load_refiner:
|
||||
required.append("transformer_2")
|
||||
fastvideo_args.pipeline_config.vae_config.load_encoder = True
|
||||
self._required_config_modules = required
|
||||
return super().load_modules(fastvideo_args, loaded_modules)
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Apply the released runtime flow shift to the loaded scheduler."""
|
||||
shift = fastvideo_args.pipeline_config.flow_shift
|
||||
if shift is None:
|
||||
raise ValueError("LingBot-Video requires a flow shift")
|
||||
self.get_module("scheduler").set_shift(float(shift))
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Create base generation and the optional decoded-video refiner stages."""
|
||||
refiner = self.get_module("transformer_2")
|
||||
self.add_stage(
|
||||
"input_validation_stage",
|
||||
LingBotVideoInputValidationStage(refiner_enabled=refiner is not None),
|
||||
)
|
||||
self.add_stage(
|
||||
"prompt_encoding_stage",
|
||||
TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
self.add_stage("conditioning_stage", ConditioningStage())
|
||||
self.add_stage(
|
||||
"timestep_preparation_stage",
|
||||
TimestepPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
self.add_stage(
|
||||
"latent_preparation_stage",
|
||||
LingBotVideoLatentPreparationStage(transformer=self.get_module("transformer")),
|
||||
)
|
||||
self.add_stage(
|
||||
"denoising_stage",
|
||||
LingBotVideoDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"decoding_stage",
|
||||
DecodingStage(vae=self.get_module("vae"), pipeline=self),
|
||||
)
|
||||
if refiner is not None:
|
||||
self.add_stage(
|
||||
"refiner_preparation_stage",
|
||||
LingBotVideoRefinerPreparationStage(
|
||||
vae=self.get_module("vae"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"refiner_denoising_stage",
|
||||
LingBotVideoDenoisingStage(
|
||||
transformer=refiner,
|
||||
scheduler=self.get_module("scheduler"),
|
||||
refiner=True,
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"refiner_decoding_stage",
|
||||
DecodingStage(vae=self.get_module("vae"), pipeline=self),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = LingBotVideoPipeline
|
||||
@@ -0,0 +1,92 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Official LingBot-Video T2V inference presets."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
DEFAULT_NEGATIVE_PROMPT = ('{"universal_negative": {"visual_quality": ["low quality", "worst quality", "blurry", '
|
||||
'"pixelated", "jpeg artifacts", "low resolution", "unstable color", "color flicker", '
|
||||
'"underexposed", "overexposed", "invisible subject", "subject hidden in darkness"], '
|
||||
'"artistic_style": ["painting", "illustration", "drawing", "cartoon", "3d render", '
|
||||
'"cgi", "sketch", "digital art"], "composition_and_content": ["text", "watermark", '
|
||||
'"signature", "logo", "subtitles", "pillarboxed", "side bars", "portrait image in '
|
||||
'landscape frame"], "temporal_and_motion_stability": ["flickering", "jittery", '
|
||||
'"motion blur", "temporal inconsistency", "warping", "morphing", "incoherent motion", '
|
||||
'"unnatural movement", "static object with sudden jump", "frame-to-frame inconsistency"], '
|
||||
'"material_and_structure": ["plastic-like glass", "unrealistic texture", "deformed '
|
||||
'bottle", "liquid freezing improperly", "distorted reflections"]}}')
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="LingBot-Video batched-CFG denoising",
|
||||
allowed_overrides=frozenset({"num_inference_steps", "guidance_scale"}),
|
||||
)
|
||||
|
||||
_REFINE_STAGE = PresetStageSpec(
|
||||
name="refine",
|
||||
kind="refinement",
|
||||
description="LingBot-Video pixel-space resize, VAE re-encode, and refiner denoising",
|
||||
allowed_overrides=frozenset({
|
||||
"height_sr",
|
||||
"width_sr",
|
||||
"num_inference_steps_sr",
|
||||
"guidance_scale_2",
|
||||
"t_thresh",
|
||||
}),
|
||||
)
|
||||
|
||||
LINGBOT_VIDEO_DENSE_T2V = InferencePreset(
|
||||
name="lingbot_video_dense_t2v",
|
||||
version=1,
|
||||
model_family="lingbot_video",
|
||||
description="LingBot-Video Dense 1.3B text-to-video at 480p",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 40,
|
||||
"guidance_scale": 3.0,
|
||||
"batch_cfg": True,
|
||||
"seed": 42,
|
||||
"negative_prompt": DEFAULT_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
LINGBOT_VIDEO_MOE_REFINER_T2V = InferencePreset(
|
||||
name="lingbot_video_moe_refiner_t2v",
|
||||
version=1,
|
||||
model_family="lingbot_video",
|
||||
description="LingBot-Video MoE 30B-A3B T2V with 1080p refiner",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, _REFINE_STAGE),
|
||||
defaults={
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"height_sr": 1088,
|
||||
"width_sr": 1920,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 40,
|
||||
"num_inference_steps_sr": 8,
|
||||
"guidance_scale": 3.0,
|
||||
"guidance_scale_2": 3.0,
|
||||
"batch_cfg": True,
|
||||
"t_thresh": 0.85,
|
||||
"seed": 42,
|
||||
"negative_prompt": DEFAULT_NEGATIVE_PROMPT,
|
||||
},
|
||||
stage_defaults={
|
||||
"refine": {
|
||||
"height_sr": 1088,
|
||||
"width_sr": 1920,
|
||||
"num_inference_steps_sr": 8,
|
||||
"guidance_scale_2": 3.0,
|
||||
"t_thresh": 0.85,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (LINGBOT_VIDEO_DENSE_T2V, LINGBOT_VIDEO_MOE_REFINER_T2V)
|
||||
@@ -0,0 +1,345 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot-Video stages whose contracts differ from shared Wan behavior."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
|
||||
LINGBOT_VIDEO_REFINER_TAIL_STEPS = 2
|
||||
|
||||
|
||||
def _compute_refiner_sigmas(
|
||||
sigma_max: float,
|
||||
sigma_min: float,
|
||||
num_inference_steps: int,
|
||||
shift: float,
|
||||
t_thresh: float,
|
||||
) -> np.ndarray:
|
||||
"""Build the released truncated schedule plus its two-step low-noise tail."""
|
||||
if not 0.0 < t_thresh <= 1.0:
|
||||
raise ValueError(f"LingBot-Video refiner t_thresh must be in (0, 1], got {t_thresh}")
|
||||
if num_inference_steps < 1:
|
||||
raise ValueError("LingBot-Video refiner requires at least one inference step")
|
||||
base = np.linspace(sigma_max, sigma_min, num_inference_steps + 1).copy()[:-1]
|
||||
shifted = shift * base / (1.0 + (shift - 1.0) * base)
|
||||
sigmas = shifted[shifted <= t_thresh + 1e-6]
|
||||
if sigmas.size == 0 or abs(float(sigmas[0]) - t_thresh) > 1e-6:
|
||||
sigmas = np.concatenate(([t_thresh], sigmas))
|
||||
tail = np.linspace(
|
||||
float(sigmas[-1]),
|
||||
min(sigma_min, float(sigmas[-1])),
|
||||
LINGBOT_VIDEO_REFINER_TAIL_STEPS + 2,
|
||||
)[1:-1]
|
||||
sigmas = np.concatenate((sigmas, tail))
|
||||
if sigmas.size > 1 and not np.all(np.diff(sigmas) < 0.0):
|
||||
raise ValueError(f"LingBot-Video refiner sigmas must descend strictly, got {sigmas.tolist()}")
|
||||
return sigmas.astype(np.float32)
|
||||
|
||||
|
||||
class LingBotVideoInputValidationStage(InputValidationStage):
|
||||
"""Validate released shape constraints and construct the official CUDA RNG."""
|
||||
|
||||
def __init__(self, refiner_enabled: bool = False) -> None:
|
||||
self.refiner_enabled = refiner_enabled
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Use one device-local generator, matching the official production runner."""
|
||||
del fastvideo_args
|
||||
if batch.seed is None:
|
||||
raise ValueError("LingBot-Video requires a seed")
|
||||
if batch.num_videos_per_prompt != 1:
|
||||
raise ValueError("LingBot-Video currently supports one video per prompt")
|
||||
batch.seeds = [batch.seed]
|
||||
batch.generator = torch.Generator(device=get_local_torch_device()).manual_seed(batch.seed)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Run shared validation, then enforce LingBot temporal and spatial geometry."""
|
||||
batch = super().forward(batch, fastvideo_args)
|
||||
if not isinstance(batch.num_frames, int):
|
||||
raise TypeError("LingBot-Video num_frames must be an integer")
|
||||
if batch.num_frames != 1 and (batch.num_frames - 1) % 4 != 0:
|
||||
raise ValueError(f"num_frames must be 1 or 4n+1, got {batch.num_frames}")
|
||||
if not isinstance(batch.height, int) or not isinstance(batch.width, int):
|
||||
raise TypeError("LingBot-Video height and width must be integers")
|
||||
if batch.height % 16 != 0 or batch.width % 16 != 0:
|
||||
raise ValueError(f"height and width must be divisible by 16, got {batch.height}x{batch.width}")
|
||||
if isinstance(batch.prompt, list) and len(batch.prompt) != 1:
|
||||
raise ValueError("LingBot-Video currently supports prompt batch size one")
|
||||
if self.refiner_enabled and fastvideo_args.output_type == "latent":
|
||||
raise ValueError("LingBot-Video refinement requires decoded pixel output")
|
||||
return batch
|
||||
|
||||
|
||||
class LingBotVideoLatentPreparationStage(PipelineStage):
|
||||
"""Prepare fp32 latents in the released 4x temporal and 8x spatial geometry."""
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Generate or validate one normalized fp32 latent video."""
|
||||
del fastvideo_args
|
||||
if not all(isinstance(value, int) for value in (batch.num_frames, batch.height, batch.width)):
|
||||
raise TypeError("latent geometry must contain integer frames, height, and width")
|
||||
shape = (
|
||||
1,
|
||||
self.transformer.num_channels_latents,
|
||||
(batch.num_frames - 1) // 4 + 1,
|
||||
batch.height // 8,
|
||||
batch.width // 8,
|
||||
)
|
||||
device = get_local_torch_device()
|
||||
if batch.latents is None:
|
||||
batch.latents = torch.randn(
|
||||
shape,
|
||||
generator=batch.generator,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
else:
|
||||
if tuple(batch.latents.shape) != shape:
|
||||
raise ValueError(f"supplied latent shape {tuple(batch.latents.shape)} does not match {shape}")
|
||||
batch.latents = batch.latents.to(device=device, dtype=torch.float32)
|
||||
batch.raw_latent_shape = shape
|
||||
return batch
|
||||
|
||||
|
||||
class LingBotVideoRefinerPreparationStage(PipelineStage):
|
||||
"""Resize and encode the base video, then initialize the released refiner state."""
|
||||
|
||||
performance_component_metric = "vae_encode_time_s"
|
||||
|
||||
def __init__(self, vae, scheduler) -> None:
|
||||
self.vae = vae
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _resize_video(video: torch.Tensor, height: int, width: int) -> torch.Tensor:
|
||||
"""Bicubic-resize every decoded frame using the released tensor layout."""
|
||||
batch, channels, frames, source_height, source_width = video.shape
|
||||
flat = video.permute(0, 2, 1, 3, 4).reshape(batch * frames, channels, source_height, source_width)
|
||||
resized = F.interpolate(flat, size=(height, width), mode="bicubic", align_corners=False).clamp(0.0, 1.0)
|
||||
return resized.reshape(batch, frames, channels, height, width).permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
def _encode_video(
|
||||
self,
|
||||
video: torch.Tensor,
|
||||
generator: torch.Generator,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Encode `[0,1]` pixels and convert Wan VAE latents to normalized DiT space."""
|
||||
video = video.to(device=device, dtype=torch.float32).mul(2.0).sub(1.0)
|
||||
with torch.autocast(device_type=device.type, dtype=torch.bfloat16, enabled=device.type == "cuda"):
|
||||
encoded = self.vae.encode(video)
|
||||
if hasattr(encoded, "latent_dist"):
|
||||
latents = encoded.latent_dist.sample(generator)
|
||||
elif hasattr(encoded, "sample") and callable(encoded.sample):
|
||||
latents = encoded.sample(generator)
|
||||
elif isinstance(encoded, tuple | list):
|
||||
latents = encoded[0]
|
||||
else:
|
||||
latents = encoded
|
||||
mean = torch.tensor(self.vae.config.latents_mean, device=device, dtype=torch.float32).view(1, -1, 1, 1, 1)
|
||||
std = torch.tensor(self.vae.config.latents_std, device=device, dtype=torch.float32).view(1, -1, 1, 1, 1)
|
||||
return ((latents.float() - mean) / std).to(latents)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Prepare high-resolution refiner latents and its exact truncated sigma schedule."""
|
||||
if batch.output is None or batch.output.ndim != 5:
|
||||
raise ValueError("LingBot-Video refinement requires a decoded base video")
|
||||
if not isinstance(batch.height_sr, int) or not isinstance(batch.width_sr, int):
|
||||
raise TypeError("LingBot-Video refinement requires integer height_sr and width_sr")
|
||||
if batch.height_sr % 16 != 0 or batch.width_sr % 16 != 0:
|
||||
raise ValueError("LingBot-Video refiner height_sr and width_sr must be divisible by 16")
|
||||
if batch.seed is None:
|
||||
raise ValueError("LingBot-Video refinement requires a seed")
|
||||
device = get_local_torch_device()
|
||||
if isinstance(self.vae, torch.nn.Module):
|
||||
self.vae.to(device)
|
||||
generator = torch.Generator(device=device).manual_seed(batch.seed)
|
||||
resized = self._resize_video(batch.output, batch.height_sr, batch.width_sr)
|
||||
encoded = self._encode_video(resized, generator, device)
|
||||
noise = torch.randn(encoded.shape, generator=generator, device=device, dtype=encoded.dtype)
|
||||
batch.latents = ((1.0 - batch.t_thresh) * encoded + batch.t_thresh * noise).float()
|
||||
batch.generator = generator
|
||||
batch.height = batch.height_sr
|
||||
batch.width = batch.width_sr
|
||||
batch.raw_latent_shape = tuple(batch.latents.shape)
|
||||
batch.extra["lingbot_video_base_shape"] = tuple(batch.output.shape)
|
||||
batch.output = None
|
||||
|
||||
shift = fastvideo_args.pipeline_config.flow_shift
|
||||
if shift is None:
|
||||
raise ValueError("LingBot-Video refinement requires a flow shift")
|
||||
sigmas = _compute_refiner_sigmas(
|
||||
float(self.scheduler.sigma_max),
|
||||
float(self.scheduler.sigma_min),
|
||||
batch.num_inference_steps_sr,
|
||||
float(shift),
|
||||
float(batch.t_thresh),
|
||||
)
|
||||
self.scheduler.set_timesteps(len(sigmas), device=device, sigmas=sigmas, shift=1.0)
|
||||
batch.timesteps = self.scheduler.timesteps
|
||||
if getattr(fastvideo_args, "vae_cpu_offload", False):
|
||||
self.vae.to("cpu")
|
||||
return batch
|
||||
|
||||
|
||||
class LingBotVideoDenoisingStage(PipelineStage):
|
||||
"""Run the released batched-CFG bf16 DiT loop with fp32 scheduler state."""
|
||||
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler, refiner: bool = False) -> None:
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.refiner = refiner
|
||||
|
||||
@staticmethod
|
||||
def _pad_condition(
|
||||
embeds: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
length: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Right-pad one condition stream to a shared batched-CFG length."""
|
||||
pad_length = length - embeds.shape[1]
|
||||
if pad_length < 0 or embeds.shape[:2] != mask.shape:
|
||||
raise ValueError("invalid LingBot-Video prompt embedding/mask shapes")
|
||||
if pad_length == 0:
|
||||
return embeds, mask
|
||||
embed_padding = embeds.new_zeros(embeds.shape[0], pad_length, embeds.shape[2])
|
||||
mask_padding = mask.new_zeros(mask.shape[0], pad_length)
|
||||
return (
|
||||
torch.cat((embeds, embed_padding), dim=1),
|
||||
torch.cat((mask, mask_padding), dim=1),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _transformer_timestep(timestep: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
|
||||
"""Reproduce the official divide-cast-multiply timestep rounding."""
|
||||
sigma = timestep.float() / 1000.0
|
||||
if dtype in (torch.bfloat16, torch.float16):
|
||||
sigma = sigma.to(dtype)
|
||||
return (sigma * 1000.0).float()
|
||||
|
||||
def _prepare_conditions(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Pack conditional then unconditional text streams for one batched CFG call."""
|
||||
prompt = batch.prompt_embeds[0].to(device=device, dtype=dtype)
|
||||
if batch.prompt_attention_mask is None or not batch.prompt_attention_mask:
|
||||
raise ValueError("LingBot-Video requires a prompt attention mask")
|
||||
prompt_mask = batch.prompt_attention_mask[0].to(device=device)
|
||||
if not self._uses_cfg(batch) or not batch.batch_cfg:
|
||||
return prompt, prompt_mask
|
||||
negative, negative_mask = self._negative_condition(batch, prompt, prompt_mask, dtype, device)
|
||||
target_length = max(prompt.shape[1], negative.shape[1])
|
||||
prompt, prompt_mask = self._pad_condition(prompt, prompt_mask, target_length)
|
||||
negative, negative_mask = self._pad_condition(negative, negative_mask, target_length)
|
||||
return (
|
||||
torch.cat((prompt, negative), dim=0),
|
||||
torch.cat((prompt_mask, negative_mask), dim=0),
|
||||
)
|
||||
|
||||
def _negative_condition(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
prompt: torch.Tensor,
|
||||
prompt_mask: torch.Tensor,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Use the refiner's zero-cloned null condition or the encoded negative prompt."""
|
||||
if self.refiner:
|
||||
return torch.zeros_like(prompt), prompt_mask.clone()
|
||||
if batch.negative_prompt_embeds is None or not batch.negative_prompt_embeds:
|
||||
raise ValueError("LingBot-Video CFG requires negative prompt embeddings")
|
||||
if batch.negative_attention_mask is None or not batch.negative_attention_mask:
|
||||
raise ValueError("LingBot-Video CFG requires a negative prompt mask")
|
||||
return (
|
||||
batch.negative_prompt_embeds[0].to(device=device, dtype=dtype),
|
||||
batch.negative_attention_mask[0].to(device=device),
|
||||
)
|
||||
|
||||
def _uses_cfg(self, batch: ForwardBatch) -> bool:
|
||||
"""Enable guidance independently for the base or refiner scale."""
|
||||
scale = batch.guidance_scale_2 if self.refiner else batch.guidance_scale
|
||||
return scale is not None and scale > 1.0
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Denoise latents while keeping scheduler samples and predictions in fp32."""
|
||||
if batch.latents is None or batch.timesteps is None:
|
||||
raise ValueError("LingBot-Video denoising requires latents and timesteps")
|
||||
device = get_local_torch_device()
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
condition, condition_mask = self._prepare_conditions(batch, transformer_dtype, device)
|
||||
latents = batch.latents.to(device=device, dtype=torch.float32)
|
||||
do_cfg = self._uses_cfg(batch)
|
||||
negative = negative_mask = None
|
||||
if do_cfg and not batch.batch_cfg:
|
||||
negative, negative_mask = self._negative_condition(batch, condition, condition_mask, transformer_dtype,
|
||||
device)
|
||||
trajectory: list[torch.Tensor] = []
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
for timestep in batch.timesteps:
|
||||
timestep_batch = self._transformer_timestep(timestep, transformer_dtype).expand(1).to(device)
|
||||
latent_input = latents
|
||||
if do_cfg and batch.batch_cfg:
|
||||
latent_input = torch.cat((latents, latents), dim=0)
|
||||
timestep_batch = torch.cat((timestep_batch, timestep_batch), dim=0)
|
||||
autocast_enabled = device.type == "cuda" and transformer_dtype != torch.float32
|
||||
with torch.autocast(device_type=device.type, dtype=transformer_dtype, enabled=autocast_enabled):
|
||||
prediction = self.transformer(
|
||||
latent_input,
|
||||
timestep_batch,
|
||||
condition,
|
||||
encoder_attention_mask=condition_mask,
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if do_cfg:
|
||||
if batch.batch_cfg:
|
||||
conditional, unconditional = prediction.chunk(2, dim=0)
|
||||
else:
|
||||
with torch.autocast(
|
||||
device_type=device.type,
|
||||
dtype=transformer_dtype,
|
||||
enabled=autocast_enabled,
|
||||
):
|
||||
unconditional = self.transformer(
|
||||
latents,
|
||||
timestep_batch,
|
||||
negative,
|
||||
encoder_attention_mask=negative_mask,
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
conditional = prediction
|
||||
guidance_scale = batch.guidance_scale_2 if self.refiner else batch.guidance_scale
|
||||
if guidance_scale is None:
|
||||
raise ValueError("LingBot-Video CFG requires a guidance scale")
|
||||
prediction = unconditional + guidance_scale * (conditional - unconditional)
|
||||
latents = self.scheduler.step(
|
||||
prediction,
|
||||
timestep,
|
||||
latents,
|
||||
return_dict=False,
|
||||
generator=batch.generator,
|
||||
)[0].float()
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory.append(latents.detach().cpu())
|
||||
trajectory_timesteps.append(timestep.detach().cpu())
|
||||
batch.latents = latents
|
||||
if trajectory:
|
||||
batch.trajectory_latents = torch.stack(trajectory, dim=1)
|
||||
batch.trajectory_timesteps = trajectory_timesteps
|
||||
return batch
|
||||
@@ -0,0 +1,5 @@
|
||||
from .causal_fast_pipeline import LingBotWorld2CausalFastPipeline
|
||||
|
||||
__all__ = ["LingBotWorld2CausalFastPipeline"]
|
||||
|
||||
EntryClass = LingBotWorld2CausalFastPipeline
|
||||
@@ -0,0 +1,365 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot World 2 causal-fast image-to-video pipeline."""
|
||||
|
||||
import math
|
||||
import os
|
||||
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.lingbotworld2.cam_utils import (
|
||||
compute_relative_poses,
|
||||
get_Ks_transformed,
|
||||
get_plucker_embeddings,
|
||||
interpolate_camera_poses,
|
||||
)
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler, )
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LingBotWorld2TextEncodingStage(TextEncodingStage):
|
||||
"""Keep LingBot World 2 T5 attention masks so DiT context matches the source pipeline."""
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Initialize mask storage before running the shared text-encoding stage."""
|
||||
if batch.prompt_attention_mask is None:
|
||||
batch.prompt_attention_mask = []
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class LingBotWorld2CausalFastGenerationStage(PipelineStage):
|
||||
"""Prepare LingBot World 2 conditions and run the released causal-fast sampling loop."""
|
||||
|
||||
def __init__(self, transformer, scheduler, vae) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self._cross_attn_initialized = False
|
||||
|
||||
def _convert_flow_pred_to_x0(
|
||||
self,
|
||||
flow_pred: torch.Tensor,
|
||||
xt: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Convert LingBot World 2 flow prediction to x0 using the scheduler sigma."""
|
||||
original_dtype = flow_pred.dtype
|
||||
flow_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(flow_pred.device),
|
||||
[flow_pred, xt, self.scheduler.sigmas, self.scheduler.timesteps],
|
||||
)
|
||||
timestep_id = torch.argmin((timesteps - timestep).abs())
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
return (xt - sigma_t * flow_pred).to(original_dtype)
|
||||
|
||||
def _initialize_self_kv_cache(
|
||||
self,
|
||||
batch_size: int,
|
||||
kv_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> list[dict]:
|
||||
"""Allocate per-block self-attention KV cache tensors."""
|
||||
head_dim = self.transformer.dim // self.transformer.num_heads
|
||||
num_heads = self.transformer.num_heads // get_sp_world_size()
|
||||
shape = [batch_size, kv_size, num_heads, head_dim]
|
||||
return [{
|
||||
"k": torch.zeros(shape, dtype=dtype, device=device),
|
||||
"v": torch.zeros(shape, dtype=dtype, device=device),
|
||||
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index": torch.tensor([0], dtype=torch.long, device=device),
|
||||
} for _ in range(self.transformer.num_layers)]
|
||||
|
||||
def _initialize_crossattn_cache(
|
||||
self,
|
||||
batch_size: int,
|
||||
max_sequence_length: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> list[dict]:
|
||||
"""Allocate per-block text cross-attention KV cache tensors."""
|
||||
head_dim = self.transformer.dim // self.transformer.num_heads
|
||||
shape = [batch_size, max_sequence_length, self.transformer.num_heads, head_dim]
|
||||
return [{
|
||||
"k": torch.zeros(shape, dtype=dtype, device=device),
|
||||
"v": torch.zeros(shape, dtype=dtype, device=device),
|
||||
"is_init": torch.tensor([0], dtype=torch.bool, device=device),
|
||||
} for _ in range(self.transformer.num_layers)]
|
||||
|
||||
@staticmethod
|
||||
def _prompt_context(batch: ForwardBatch, device: torch.device) -> list[torch.Tensor]:
|
||||
"""Slice padded text encoder states back to LingBot World 2's unpadded context list."""
|
||||
assert batch.prompt_embeds
|
||||
context_tensor = batch.prompt_embeds[0].to(device)
|
||||
if batch.prompt_attention_mask:
|
||||
mask = batch.prompt_attention_mask[0].to(device)
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
return [u[:v] for u, v in zip(context_tensor, seq_lens, strict=True)]
|
||||
return [u for u in context_tensor]
|
||||
|
||||
def _prepare_image_tensor(self, batch: ForwardBatch, device: torch.device) -> torch.Tensor:
|
||||
"""Return the source-style normalized image tensor `[C,H,W]`."""
|
||||
image = batch.pil_image
|
||||
if image is None:
|
||||
raise ValueError("LingBot World 2 causal-fast requires `image_path` or `pil_image`.")
|
||||
if isinstance(image, torch.Tensor):
|
||||
if image.ndim == 5:
|
||||
return image[0, :, 0].to(device)
|
||||
if image.ndim == 4:
|
||||
return image[0].to(device)
|
||||
return image.to(device)
|
||||
return TF.to_tensor(image).sub_(0.5).div_(0.5).to(device)
|
||||
|
||||
def _prepare_camera(
|
||||
self,
|
||||
action_path: str,
|
||||
c2ws: np.ndarray,
|
||||
h: int,
|
||||
w: int,
|
||||
lat_f: int,
|
||||
lat_h: int,
|
||||
lat_w: int,
|
||||
chunk_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Build the LingBot World 2 camera Plucker tensor for latent chunks."""
|
||||
Ks = torch.from_numpy(np.load(os.path.join(action_path, "intrinsics.npy"))).float()
|
||||
Ks = get_Ks_transformed(
|
||||
Ks,
|
||||
height_org=480,
|
||||
width_org=832,
|
||||
height_resize=h,
|
||||
width_resize=w,
|
||||
height_final=h,
|
||||
width_final=w,
|
||||
)
|
||||
Ks = Ks[0]
|
||||
len_c2ws = len(c2ws)
|
||||
len_c2ws_ = int((len_c2ws - 1) // 4) + 1
|
||||
len_c2ws_ = int(len_c2ws_ - (len_c2ws_ % chunk_size))
|
||||
c2ws_infer = interpolate_camera_poses(
|
||||
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
|
||||
src_rot_mat=c2ws[:, :3, :3],
|
||||
src_trans_vec=c2ws[:, :3, 3],
|
||||
tgt_indices=np.linspace(0, len_c2ws - 1, len_c2ws_),
|
||||
)
|
||||
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
|
||||
Ks = Ks.repeat(len(c2ws_infer), 1)
|
||||
c2ws_plucker_emb = get_plucker_embeddings(c2ws_infer.to(device), Ks.to(device), h, w)
|
||||
c2ws_plucker_emb = rearrange(
|
||||
c2ws_plucker_emb,
|
||||
"f (h c1) (w c2) c -> (f h w) (c c1 c2)",
|
||||
c1=int(h // lat_h),
|
||||
c2=int(w // lat_w),
|
||||
)
|
||||
c2ws_plucker_emb = c2ws_plucker_emb[None, ...]
|
||||
return rearrange(
|
||||
c2ws_plucker_emb,
|
||||
"b (f h w) c -> b c f h w",
|
||||
f=lat_f,
|
||||
h=lat_h,
|
||||
w=lat_w,
|
||||
).to(device=device, dtype=dtype)
|
||||
|
||||
def _encode_condition_video(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
h: int,
|
||||
w: int,
|
||||
frames: int,
|
||||
mask: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> torch.Tensor:
|
||||
"""Encode the first-frame conditioning video and prepend mask channels."""
|
||||
device = get_local_torch_device()
|
||||
self.vae = self.vae.to(device)
|
||||
video_condition = torch.concat(
|
||||
[
|
||||
torch.nn.functional.interpolate(
|
||||
img[None].cpu(),
|
||||
size=(h, w),
|
||||
mode="bicubic",
|
||||
).transpose(0, 1),
|
||||
torch.zeros(3, frames - 1, h, w),
|
||||
],
|
||||
dim=1,
|
||||
).to(device)
|
||||
vae_dtype = torch.float32
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=False):
|
||||
encoder_output = self.vae.encode(video_condition.unsqueeze(0).to(torch.float32))
|
||||
latent_condition = encoder_output.mean
|
||||
if not bool(getattr(self.vae, "handles_latent_denorm", False)):
|
||||
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
|
||||
latent_condition -= self.vae.shift_factor.to(latent_condition.device, latent_condition.dtype)
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(latent_condition.device,
|
||||
latent_condition.dtype)
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
return torch.concat([mask, latent_condition[0]], dim=0)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Execute LingBot World 2 causal-fast generation and store final latents on the batch."""
|
||||
device = get_local_torch_device()
|
||||
cfg = fastvideo_args.pipeline_config.dit_config.arch_config
|
||||
chunk_size = int(cfg.chunk_size)
|
||||
max_sequence_length = int(batch.max_sequence_length or cfg.text_len)
|
||||
action_path = batch.action_path
|
||||
if action_path is None:
|
||||
raise ValueError("LingBot World 2 causal-fast requires `action_path`.")
|
||||
|
||||
c2ws = np.load(os.path.join(action_path, "poses.npy"))
|
||||
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
|
||||
frame_num = ((int(batch.num_frames) - 1) // 4) * 4 + 1
|
||||
frame_num = min(frame_num, len_c2ws)
|
||||
c2ws = c2ws[:frame_num]
|
||||
|
||||
img = self._prepare_image_tensor(batch, device)
|
||||
h0, w0 = img.shape[1:]
|
||||
aspect_ratio = h0 / w0
|
||||
lat_h = round(np.sqrt(cfg.max_area * aspect_ratio) // 8 // cfg.patch_size[1] * cfg.patch_size[1])
|
||||
lat_w = round(np.sqrt(cfg.max_area / aspect_ratio) // 8 // cfg.patch_size[2] * cfg.patch_size[2])
|
||||
h = lat_h * 8
|
||||
w = lat_w * 8
|
||||
lat_f = (frame_num - 1) // 4 + 1
|
||||
lat_f = int(lat_f - (lat_f % chunk_size))
|
||||
frames = (lat_f - 1) * 4 + 1
|
||||
batch.height = h
|
||||
batch.width = w
|
||||
batch.num_frames = frames
|
||||
|
||||
seed = int(batch.seed if batch.seed is not None else 42)
|
||||
seed_g = torch.Generator(device=device)
|
||||
seed_g.manual_seed(seed)
|
||||
noise = torch.randn(16, lat_f, lat_h, lat_w, dtype=torch.float32, generator=seed_g, device=device)
|
||||
|
||||
mask = torch.ones(1, frames, lat_h, lat_w, device=device)
|
||||
mask[:, 1:] = 0
|
||||
mask = torch.concat([torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1), mask[:, 1:]], dim=1)
|
||||
mask = mask.view(1, mask.shape[1] // 4, 4, lat_h, lat_w).transpose(1, 2)[0]
|
||||
|
||||
self.scheduler.set_timesteps(cfg.num_train_timesteps, shift=cfg.sample_shift)
|
||||
timesteps = self.scheduler.timesteps[list(cfg.timesteps_index)].to(device)
|
||||
context = self._prompt_context(batch, device)
|
||||
c2ws_plucker_emb = self._prepare_camera(
|
||||
action_path,
|
||||
c2ws,
|
||||
h,
|
||||
w,
|
||||
lat_f,
|
||||
lat_h,
|
||||
lat_w,
|
||||
chunk_size,
|
||||
torch.bfloat16,
|
||||
device,
|
||||
)
|
||||
y = self._encode_condition_video(img, h, w, frames, mask, fastvideo_args).to(device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
transformer_dtype = torch.bfloat16
|
||||
frame_seqlen = int(noise.shape[-2] * noise.shape[-1] // 4)
|
||||
kv_size = frame_seqlen * cfg.local_attn_size if cfg.local_attn_size > -1 else frame_seqlen * lat_f
|
||||
self_kv_cache = self._initialize_self_kv_cache(1, kv_size, transformer_dtype, device)
|
||||
cross_kv_cache = self._initialize_crossattn_cache(1, max_sequence_length, transformer_dtype, device)
|
||||
|
||||
self.transformer = self.transformer.to(device)
|
||||
self._cross_attn_initialized = False
|
||||
pred_latent_chunks = []
|
||||
latents_chunk = noise.split(chunk_size, dim=1)
|
||||
condition_chunk = y.split(chunk_size, dim=1)
|
||||
c2ws_plucker_emb_chunk = c2ws_plucker_emb.split(chunk_size, dim=2)
|
||||
max_seq_len = int(math.ceil(chunk_size * lat_h * lat_w // 4))
|
||||
|
||||
with torch.amp.autocast("cuda", dtype=transformer_dtype):
|
||||
for chunk_id, current_latent in enumerate(latents_chunk):
|
||||
current_condition = condition_chunk[chunk_id]
|
||||
current_c2ws_plucker_emb = c2ws_plucker_emb_chunk[chunk_id]
|
||||
dit_cond_dict = {"c2ws_plucker_emb": current_c2ws_plucker_emb.chunk(1, dim=0)}
|
||||
kwargs = {
|
||||
"context": [context[0]],
|
||||
"seq_len": max_seq_len,
|
||||
"y": [current_condition],
|
||||
"dit_cond_dict": dit_cond_dict,
|
||||
"kv_cache": self_kv_cache,
|
||||
"crossattn_cache": cross_kv_cache,
|
||||
"current_start": chunk_id * chunk_size * frame_seqlen,
|
||||
"max_attention_size": kv_size,
|
||||
"frame_seqlen": frame_seqlen,
|
||||
}
|
||||
x0 = current_latent
|
||||
for timestep_idx, timestep_value in enumerate(timesteps):
|
||||
timestep = torch.stack([timestep_value]).to(device)
|
||||
noise_pred = self.transformer(
|
||||
x=[current_latent.to(device)],
|
||||
t=timestep,
|
||||
cross_attn_first_call=not self._cross_attn_initialized,
|
||||
**kwargs,
|
||||
)[0]
|
||||
self._cross_attn_initialized = True
|
||||
x0 = self._convert_flow_pred_to_x0(noise_pred, current_latent, timestep_value)
|
||||
if timestep_idx < len(timesteps) - 1:
|
||||
next_timestep = timesteps[timestep_idx + 1].reshape(1)
|
||||
current_latent = self.scheduler.add_noise(
|
||||
x0,
|
||||
torch.randn(x0.shape, generator=seed_g, device=x0.device, dtype=x0.dtype),
|
||||
next_timestep,
|
||||
)
|
||||
pred_latent_chunks.append(x0)
|
||||
context_timestep = torch.stack([timesteps[-1] * 0.0]).to(device)
|
||||
self.transformer(x=[x0], t=context_timestep, cross_attn_first_call=False, **kwargs)
|
||||
|
||||
batch.latents = torch.cat(pred_latent_chunks, dim=1).unsqueeze(0)
|
||||
return batch
|
||||
|
||||
|
||||
class LingBotWorld2CausalFastPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""FastVideo pipeline for LingBot World 2 14B causal-fast I2V generation."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
if "scheduler" not in self.modules:
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(shift=1.0, use_dynamic_shifting=False)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up the LingBot World 2 causal-fast pipeline stages."""
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=LingBotWorld2TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
|
||||
self.add_stage(
|
||||
stage_name="lingbotworld2_causal_fast_generation_stage",
|
||||
stage=LingBotWorld2CausalFastGenerationStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
),
|
||||
)
|
||||
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = LingBotWorld2CausalFastPipeline
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBotWorld2 causal-fast pipeline preset."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Causal-fast denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
LINGBOTWORLD2_CAUSAL_FAST_I2V = InferencePreset(
|
||||
name="lingbotworld2_causal_fast_i2v",
|
||||
version=1,
|
||||
model_family="lingbotworld2",
|
||||
description="LingBot World 2 14B causal-fast I2V",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"fps": 16,
|
||||
"seed": 42,
|
||||
"num_frames": 65,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"negative_prompt": "",
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (LINGBOTWORLD2_CAUSAL_FAST_I2V, )
|
||||
@@ -0,0 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Z-Image pipeline."""
|
||||
|
||||
from fastvideo.pipelines.basic.zimage.zimage_pipeline import ZImagePipeline
|
||||
|
||||
__all__ = ["ZImagePipeline"]
|
||||
@@ -0,0 +1,40 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Z-Image inference presets."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Z-Image denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
"cfg_normalization",
|
||||
"cfg_truncation",
|
||||
}),
|
||||
)
|
||||
|
||||
ZIMAGE_TURBO = InferencePreset(
|
||||
name="zimage_turbo",
|
||||
version=1,
|
||||
model_family="zimage",
|
||||
description="Z-Image-Turbo text-to-image generation",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"seed": 42,
|
||||
"guidance_scale": 0.0,
|
||||
"num_inference_steps": 8,
|
||||
"negative_prompt": "",
|
||||
"max_sequence_length": 512,
|
||||
"cfg_normalization": False,
|
||||
"cfg_truncation": 1.0,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (ZIMAGE_TURBO, )
|
||||
@@ -0,0 +1,337 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline stages for the native Z-Image text-to-image path."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
class ZImageInputValidationStage(InputValidationStage):
|
||||
"""Validate the image geometry and reproduce the official device RNG."""
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> None:
|
||||
del fastvideo_args
|
||||
assert batch.seed is not None
|
||||
batch.seeds = [batch.seed]
|
||||
device = get_local_torch_device()
|
||||
batch.generator = torch.Generator(device=device).manual_seed(batch.seed)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.do_classifier_free_guidance and batch.negative_prompt is None and not batch.negative_prompt_embeds:
|
||||
batch.negative_prompt = ""
|
||||
batch = super().forward(batch, fastvideo_args)
|
||||
if batch.num_frames != 1:
|
||||
raise ValueError(f"Z-Image is text-to-image and requires num_frames=1, got {batch.num_frames}")
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("Z-Image requires height and width")
|
||||
if batch.height % 16 or batch.width % 16:
|
||||
raise ValueError("Z-Image height and width must be divisible by 16; "
|
||||
f"got {batch.height}x{batch.width}")
|
||||
return batch
|
||||
|
||||
|
||||
class ZImageConditioningStage(PipelineStage):
|
||||
"""Trim padded Qwen states and materialize variable-length CFG streams."""
|
||||
|
||||
@staticmethod
|
||||
def _trim_embeddings(
|
||||
embeds: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None,
|
||||
) -> list[torch.Tensor]:
|
||||
if attention_mask is None:
|
||||
return list(embeds.unbind(0))
|
||||
return [
|
||||
sample[mask.to(device=sample.device, dtype=torch.bool)]
|
||||
for sample, mask in zip(embeds, attention_mask, strict=True)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _repeat(items: list[torch.Tensor], count: int) -> list[torch.Tensor]:
|
||||
return [item for item in items for _ in range(count)]
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
del fastvideo_args
|
||||
if len(batch.prompt_embeds) != 1:
|
||||
raise ValueError(f"Z-Image expects one text encoder, got {len(batch.prompt_embeds)}")
|
||||
|
||||
prompt_mask = batch.prompt_attention_mask[0] if batch.prompt_attention_mask else None
|
||||
prompt_embeds = self._trim_embeddings(batch.prompt_embeds[0], prompt_mask)
|
||||
batch.extra["zimage_prompt_embeds"] = self._repeat(prompt_embeds, batch.num_videos_per_prompt)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
if not batch.negative_prompt_embeds:
|
||||
raise ValueError("Z-Image CFG requires negative prompt embeddings")
|
||||
negative_mask = batch.negative_attention_mask[0] if batch.negative_attention_mask else None
|
||||
negative_embeds = self._trim_embeddings(batch.negative_prompt_embeds[0], negative_mask)
|
||||
batch.extra["zimage_negative_prompt_embeds"] = self._repeat(
|
||||
negative_embeds,
|
||||
batch.num_videos_per_prompt,
|
||||
)
|
||||
else:
|
||||
batch.extra["zimage_negative_prompt_embeds"] = []
|
||||
return batch
|
||||
|
||||
|
||||
class ZImageLatentPreparationStage(PipelineStage):
|
||||
"""Create the official fp32 image latents on the transformer device."""
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
self.transformer = transformer
|
||||
|
||||
@staticmethod
|
||||
def _randn(
|
||||
shape: tuple[int, ...],
|
||||
generators: torch.Generator | list[torch.Generator] | None,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(generators, list):
|
||||
if len(generators) != shape[0]:
|
||||
raise ValueError(f"generator list length {len(generators)} does not match batch size {shape[0]}")
|
||||
sample_shape = (1, *shape[1:])
|
||||
return torch.cat([
|
||||
torch.randn(sample_shape, generator=generator, device=device, dtype=torch.float32)
|
||||
for generator in generators
|
||||
])
|
||||
return torch.randn(shape, generator=generators, device=device, dtype=torch.float32)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("Z-Image requires height and width before latent preparation")
|
||||
|
||||
prompt_embeds = batch.extra.get("zimage_prompt_embeds")
|
||||
if not isinstance(prompt_embeds, list) or not prompt_embeds:
|
||||
raise ValueError("Z-Image conditioning must run before latent preparation")
|
||||
|
||||
channels = int(getattr(self.transformer, "in_channels", 16))
|
||||
spatial_ratio = int(fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio)
|
||||
shape = (
|
||||
len(prompt_embeds),
|
||||
channels,
|
||||
1,
|
||||
batch.height // spatial_ratio,
|
||||
batch.width // spatial_ratio,
|
||||
)
|
||||
device = get_local_torch_device()
|
||||
if batch.latents is None:
|
||||
latents = self._randn(shape, batch.generator, device)
|
||||
else:
|
||||
latents = batch.latents
|
||||
if latents.ndim == 4:
|
||||
latents = latents.unsqueeze(2)
|
||||
if tuple(latents.shape) != shape:
|
||||
raise ValueError(f"Expected Z-Image latents with shape {shape}, got {tuple(latents.shape)}")
|
||||
latents = latents.to(device=device, dtype=torch.float32)
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = shape
|
||||
return batch
|
||||
|
||||
|
||||
class ZImageTimestepPreparationStage(PipelineStage):
|
||||
"""Apply the native scheduler's zero endpoint and discrete schedule."""
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _calculate_shift(
|
||||
image_seq_len: int,
|
||||
base_seq_len: int,
|
||||
max_seq_len: int,
|
||||
base_shift: float,
|
||||
max_shift: float,
|
||||
) -> float:
|
||||
slope = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
return image_seq_len * slope + base_shift - slope * base_seq_len
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("Z-Image latents must be prepared before timesteps")
|
||||
if batch.timesteps is not None and batch.sigmas is not None:
|
||||
raise ValueError("Only one of timesteps or sigmas may be supplied")
|
||||
|
||||
scheduler = self.scheduler
|
||||
sigma_min = float(fastvideo_args.pipeline_config.scheduler_sigma_min)
|
||||
use_reference_timesteps = bool(fastvideo_args.pipeline_config.scheduler_use_reference_discrete_timesteps)
|
||||
scheduler.sigma_min = sigma_min
|
||||
scheduler.register_to_config(
|
||||
sigma_min=sigma_min,
|
||||
use_reference_discrete_timesteps=use_reference_timesteps,
|
||||
)
|
||||
config = scheduler.config
|
||||
image_seq_len = (batch.latents.shape[-2] // 2) * (batch.latents.shape[-1] // 2)
|
||||
mu = self._calculate_shift(
|
||||
image_seq_len,
|
||||
int(config.get("base_image_seq_len", 256)),
|
||||
int(config.get("max_image_seq_len", 4096)),
|
||||
float(config.get("base_shift", 0.5)),
|
||||
float(config.get("max_shift", 1.15)),
|
||||
)
|
||||
device = get_local_torch_device()
|
||||
|
||||
if batch.timesteps is not None:
|
||||
if "timesteps" not in inspect.signature(scheduler.set_timesteps).parameters:
|
||||
raise ValueError(f"{type(scheduler).__name__} does not accept custom timesteps")
|
||||
scheduler.set_timesteps(timesteps=batch.timesteps, device=device, mu=mu)
|
||||
elif batch.sigmas is not None:
|
||||
if "sigmas" not in inspect.signature(scheduler.set_timesteps).parameters:
|
||||
raise ValueError(f"{type(scheduler).__name__} does not accept custom sigmas")
|
||||
scheduler.set_timesteps(sigmas=batch.sigmas, device=device, mu=mu)
|
||||
else:
|
||||
scheduler.set_timesteps(batch.num_inference_steps, device=device, mu=mu)
|
||||
|
||||
batch.timesteps = scheduler.timesteps
|
||||
batch.num_inference_steps = len(batch.timesteps)
|
||||
return batch
|
||||
|
||||
|
||||
class ZImageDenoisingStage(PipelineStage):
|
||||
"""Run the native Z-Image flow-matching loop."""
|
||||
|
||||
performance_component_metric = "transformer_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None or batch.timesteps is None:
|
||||
raise ValueError("Z-Image denoising requires latents and timesteps")
|
||||
|
||||
latents = batch.latents.float()
|
||||
positive = batch.extra.get("zimage_prompt_embeds")
|
||||
negative = batch.extra.get("zimage_negative_prompt_embeds", [])
|
||||
if not isinstance(positive, list) or not positive:
|
||||
raise ValueError("Z-Image denoising requires prompt embeddings")
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
batch_size = latents.shape[0]
|
||||
|
||||
for index, timestep_value in enumerate(batch.timesteps):
|
||||
if timestep_value.item() == 0 and index == len(batch.timesteps) - 1:
|
||||
continue
|
||||
|
||||
timestep = timestep_value.expand(batch_size).to(device=device, dtype=torch.float32)
|
||||
timestep = (1000.0 - timestep) / 1000.0
|
||||
current_guidance_scale = float(batch.guidance_scale)
|
||||
if (batch.do_classifier_free_guidance and batch.cfg_truncation is not None and batch.cfg_truncation <= 1.0
|
||||
and timestep[0].item() > batch.cfg_truncation):
|
||||
current_guidance_scale = 0.0
|
||||
apply_cfg = batch.do_classifier_free_guidance and current_guidance_scale > 0.0
|
||||
|
||||
if apply_cfg:
|
||||
if not isinstance(negative, list) or len(negative) != batch_size:
|
||||
raise ValueError("Z-Image CFG requires one negative embedding per image")
|
||||
model_latents = latents.to(target_dtype).repeat(2, 1, 1, 1, 1)
|
||||
model_embeddings = positive + negative
|
||||
model_timestep = timestep.repeat(2)
|
||||
else:
|
||||
model_latents = latents.to(target_dtype)
|
||||
model_embeddings = positive
|
||||
model_timestep = timestep
|
||||
|
||||
with (
|
||||
torch.autocast(
|
||||
device_type=device.type,
|
||||
enabled=False,
|
||||
),
|
||||
trace_step(index),
|
||||
set_forward_context(
|
||||
current_timestep=index,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
),
|
||||
):
|
||||
model_outputs = self.transformer(
|
||||
hidden_states=model_latents,
|
||||
encoder_hidden_states=model_embeddings,
|
||||
timestep=model_timestep,
|
||||
)[0]
|
||||
|
||||
if apply_cfg:
|
||||
positive_outputs = model_outputs[:batch_size]
|
||||
negative_outputs = model_outputs[batch_size:]
|
||||
guided_outputs = []
|
||||
for positive_output, negative_output in zip(
|
||||
positive_outputs,
|
||||
negative_outputs,
|
||||
strict=True,
|
||||
):
|
||||
positive_fp32 = positive_output.float()
|
||||
prediction = positive_fp32 + current_guidance_scale * (positive_fp32 - negative_output.float())
|
||||
if batch.cfg_normalization:
|
||||
positive_norm = torch.linalg.vector_norm(positive_fp32)
|
||||
prediction_norm = torch.linalg.vector_norm(prediction)
|
||||
if prediction_norm > positive_norm:
|
||||
prediction = prediction * (positive_norm / prediction_norm)
|
||||
guided_outputs.append(prediction)
|
||||
noise_pred = torch.stack(guided_outputs)
|
||||
else:
|
||||
noise_pred = torch.stack([output.float() for output in model_outputs])
|
||||
|
||||
noise_pred = -noise_pred
|
||||
latents = self.scheduler.step(
|
||||
noise_pred,
|
||||
timestep_value,
|
||||
latents,
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
|
||||
class ZImageDecodingStage(PipelineStage):
|
||||
"""Apply the official latent transform and decode one image frame."""
|
||||
|
||||
performance_component_metric = "vae_decode_time_s"
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae = vae
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("Z-Image decoding requires latents")
|
||||
if fastvideo_args.output_type == "latent":
|
||||
# FastVideo standardizes image and video latents as [B,C,T,H,W].
|
||||
# Tongyi's image-only API returns the equivalent tensor with T squeezed.
|
||||
batch.output = batch.latents
|
||||
return batch
|
||||
|
||||
latents = batch.latents
|
||||
if latents.ndim != 5 or latents.shape[2] != 1:
|
||||
raise ValueError(f"Expected Z-Image latents [B,C,1,H,W], got {tuple(latents.shape)}")
|
||||
|
||||
device = get_local_torch_device()
|
||||
self.vae = self.vae.to(device)
|
||||
vae_dtype = getattr(self.vae, "dtype", None)
|
||||
if vae_dtype is None:
|
||||
vae_dtype = next(self.vae.parameters()).dtype
|
||||
config = self.vae.config
|
||||
scaling_factor = float(config.scaling_factor)
|
||||
shift_factor = float(config.shift_factor or 0.0)
|
||||
latents_2d = latents.squeeze(2).to(device=device, dtype=vae_dtype)
|
||||
latents_2d = latents_2d / scaling_factor + shift_factor
|
||||
decoded = self.vae.decode(latents_2d, return_dict=False)[0]
|
||||
decoded = (decoded / 2 + 0.5).clamp(0, 1)
|
||||
batch.output = decoded.unsqueeze(2).float()
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
return batch
|
||||
@@ -0,0 +1,63 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.configs.pipelines.zimage import ZImagePipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
|
||||
from .stages import (
|
||||
ZImageConditioningStage,
|
||||
ZImageDecodingStage,
|
||||
ZImageDenoisingStage,
|
||||
ZImageInputValidationStage,
|
||||
ZImageLatentPreparationStage,
|
||||
ZImageTimestepPreparationStage,
|
||||
)
|
||||
|
||||
|
||||
class ZImagePipeline(ComposedPipelineBase):
|
||||
"""Native Z-Image text-to-image pipeline."""
|
||||
|
||||
pipeline_config_cls: type[ZImagePipelineConfig] = ZImagePipelineConfig
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"transformer",
|
||||
"vae",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
scheduler = self.get_module("scheduler")
|
||||
transformer = self.get_module("transformer")
|
||||
|
||||
self.add_stage("input_validation_stage", ZImageInputValidationStage())
|
||||
self.add_stage(
|
||||
"text_encoding_stage",
|
||||
TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
self.add_stage("zimage_conditioning_stage", ZImageConditioningStage())
|
||||
self.add_stage(
|
||||
"latent_preparation_stage",
|
||||
ZImageLatentPreparationStage(transformer=transformer),
|
||||
)
|
||||
self.add_stage(
|
||||
"timestep_preparation_stage",
|
||||
ZImageTimestepPreparationStage(scheduler=scheduler),
|
||||
)
|
||||
self.add_stage(
|
||||
"denoising_stage",
|
||||
ZImageDenoisingStage(transformer=transformer, scheduler=scheduler),
|
||||
)
|
||||
self.add_stage(
|
||||
"decoding_stage",
|
||||
ZImageDecodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = ZImagePipeline
|
||||
@@ -303,7 +303,8 @@ class ComposedPipelineBase(ABC):
|
||||
self.modules[module_name] = module
|
||||
|
||||
def _load_config(self, model_path: str) -> dict[str, Any]:
|
||||
model_path = maybe_download_model(self.model_path)
|
||||
revision = getattr(self.fastvideo_args, "revision", None)
|
||||
model_path = maybe_download_model(self.model_path, revision=revision)
|
||||
self.model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
|
||||
@@ -147,8 +147,9 @@ class ForwardBatch:
|
||||
camera_trajectory: str | None = None # Camera trajectory file/identifier
|
||||
action_list: list[str] | None = None # List of actions (e.g., ['forward', 'left'])
|
||||
action_speed_list: list[float] | None = None # Speed for each action
|
||||
# Camera control inputs (LingBotWorld)
|
||||
# Camera control inputs (LingBotWorld and LingBotWorld2)
|
||||
c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
action_path: str | None = None # Directory containing poses.npy and intrinsics.npy
|
||||
|
||||
# Camera control inputs (GEN3C)
|
||||
trajectory_type: str | None = None
|
||||
@@ -177,7 +178,10 @@ class ForwardBatch:
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
batch_cfg: bool = False
|
||||
guidance_scale_2: float | None = None
|
||||
cfg_normalization: bool = False
|
||||
cfg_truncation: float | None = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
eta: float = 0.0
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
@@ -0,0 +1,258 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Preprocess LTX-2 overfit data into parquet format.
|
||||
|
||||
Encodes videos with the LTX-2 causal video VAE and captions with the
|
||||
Gemma text encoder (feature extractor + embedding connector) into the
|
||||
t2v parquet schema expected by the training framework.
|
||||
|
||||
The stored text embeddings are POST-connector: the connector replaces
|
||||
pad positions with learnable registers and returns an all-valid mask,
|
||||
so the parquet collate's ones/zeros mask stays semantically correct
|
||||
and training needs no text encoder at all. Captions are encoded via
|
||||
the encoder's forward() (the exact inference path), which handles both
|
||||
LTX-2.0 (shared 3840-d features) and LTX-2.3 (separate 4096-d video /
|
||||
2048-d audio feature extractors).
|
||||
|
||||
Videos are resampled to TRAIN_FPS and the preprocessed clip is also
|
||||
saved as an mp4 next to the parquet so overfit tests can use it as
|
||||
the SSIM reference.
|
||||
|
||||
Usage:
|
||||
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from typing import Any
|
||||
|
||||
import av
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.utils import maybe_download_model, verify_model_config_and_directory
|
||||
|
||||
# --- Config ---
|
||||
NUM_FRAMES = 81 # 8k+1 for temporal compression ratio 8
|
||||
MAX_HEIGHT = 480 # divisible by 32 (spatial compression)
|
||||
MAX_WIDTH = 832
|
||||
TRAIN_FPS = 24.0 # matches the LTX-2 preset fps used at validation
|
||||
|
||||
DATA_DIR = os.environ.get("LTX2_OVERFIT_DATA_DIR", "data/cats")
|
||||
CAPTION_JSON = os.environ.get("LTX2_OVERFIT_CAPTION_JSON", "videos2caption_1_sample.json")
|
||||
VIDEO_SUBDIR = os.environ.get("LTX2_OVERFIT_VIDEO_SUBDIR", "video")
|
||||
OUTPUT_DIR = os.environ.get("LTX2_OVERFIT_OUTPUT_DIR", "data/ltx2_overfit_preprocessed")
|
||||
MODEL_REPO = os.environ.get("LTX2_OVERFIT_MODEL", "FastVideo/LTX2-Distilled-Diffusers")
|
||||
# The train dataloader samples with drop_last=True across data-parallel
|
||||
# groups, so the dataset must hold at least num_sp_groups * batch_size
|
||||
# rows or every rank gets zero batches. Replicate the overfit sample so
|
||||
# a 4-GPU FSDP run still sees one batch per rank.
|
||||
NUM_COPIES = int(os.environ.get("LTX2_OVERFIT_NUM_COPIES", "4"))
|
||||
|
||||
|
||||
def _init_single_process_distributed() -> None:
|
||||
"""FastVideo component loaders expect an initialized distributed env."""
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29511")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
from fastvideo.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel, )
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
|
||||
|
||||
def load_video(path: str, num_frames: int, target_fps: float, height: int,
|
||||
width: int) -> tuple[torch.Tensor, np.ndarray]:
|
||||
"""Load a video as [1, C, T, H, W] in [-1, 1], resampled to target_fps.
|
||||
|
||||
Also returns the uint8 RGB frames [T, H, W, C] for reference-video export.
|
||||
"""
|
||||
with av.open(path) as container:
|
||||
if not container.streams.video:
|
||||
raise RuntimeError(f"No video stream found in {path}")
|
||||
stream = container.streams.video[0]
|
||||
native_fps = float(stream.average_rate or target_fps)
|
||||
decoded = [torch.from_numpy(frame.to_ndarray(format="rgb24")) for frame in container.decode(video=0)]
|
||||
if not decoded:
|
||||
raise RuntimeError(f"Could not read any frames from {path}")
|
||||
raw = torch.stack(decoded).permute(0, 3, 1, 2) # [T, C, H, W] uint8
|
||||
step = native_fps / target_fps
|
||||
wanted = [min(int(round(i * step)), raw.shape[0] - 1) for i in range(num_frames)]
|
||||
|
||||
frames = raw[wanted].float() # [T, C, H, W] in [0, 255]
|
||||
src_h, src_w = frames.shape[2], frames.shape[3]
|
||||
scale = max(height / src_h, width / src_w)
|
||||
new_h, new_w = int(round(src_h * scale)), int(round(src_w * scale))
|
||||
frames = torch.nn.functional.interpolate(frames, size=(new_h, new_w), mode="bilinear", antialias=True)
|
||||
top = (new_h - height) // 2
|
||||
left = (new_w - width) // 2
|
||||
frames = frames[:, :, top:top + height, left:left + width]
|
||||
|
||||
frames_np = (frames.permute(0, 2, 3, 1).clamp(0, 255).round().to(torch.uint8).numpy())
|
||||
video = frames / 127.5 - 1.0 # [0,255] -> [-1,1]
|
||||
video = video.permute(1, 0, 2, 3).unsqueeze(0) # [1,C,T,H,W]
|
||||
return video, frames_np
|
||||
|
||||
|
||||
def main() -> None:
|
||||
_init_single_process_distributed()
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.loader.component_loader import (
|
||||
PipelineComponentLoader, )
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
model_path = maybe_download_model(MODEL_REPO)
|
||||
model_index = verify_model_config_and_directory(model_path)
|
||||
|
||||
os.makedirs(OUTPUT_DIR, exist_ok=True)
|
||||
# The map-style dataset caches parquet file metadata; a stale cache
|
||||
# next to a regenerated parquet can crash or serve old rows.
|
||||
shutil.rmtree(os.path.join(OUTPUT_DIR, "map_style_cache"), ignore_errors=True)
|
||||
|
||||
with open(os.path.join(DATA_DIR, CAPTION_JSON)) as f:
|
||||
caption_data = json.load(f)
|
||||
|
||||
pipeline_config = LTX2T2VConfig()
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=model_path,
|
||||
pipeline_config=pipeline_config,
|
||||
num_gpus=1,
|
||||
tp_size=1,
|
||||
sp_size=1,
|
||||
hsdp_shard_dim=1,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
def load_component(name: str) -> Any:
|
||||
transformers_or_diffusers, _ = model_index[name]
|
||||
return PipelineComponentLoader.load_module(
|
||||
module_name=name,
|
||||
component_model_path=os.path.join(model_path, name),
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
print("Loading LTX-2 VAE...")
|
||||
vae = load_component("vae")
|
||||
vae_dtype = next(vae.parameters()).dtype
|
||||
print(f"VAE loaded ({sum(p.numel() for p in vae.parameters())/1e6:.0f}M, {vae_dtype})")
|
||||
|
||||
print("Loading Gemma text encoder + tokenizer...")
|
||||
text_encoder = load_component("text_encoder")
|
||||
tokenizer = load_component("tokenizer")
|
||||
tokenizer.padding_side = "left"
|
||||
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
encoder_config = pipeline_config.text_encoder_configs[0]
|
||||
tokenizer_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if "max_length" not in tokenizer_kwargs:
|
||||
tokenizer_kwargs["max_length"] = encoder_config.arch_config.text_len
|
||||
preprocess_text = pipeline_config.preprocess_text_funcs[0]
|
||||
|
||||
# --- Process each video ---
|
||||
records = []
|
||||
for idx, item in enumerate(caption_data):
|
||||
video_name = item["path"]
|
||||
record_id = f"{idx:04d}_{video_name}"
|
||||
caption = item["cap"][0] if isinstance(item["cap"], list) else item["cap"]
|
||||
video_path = os.path.join(DATA_DIR, VIDEO_SUBDIR, video_name)
|
||||
|
||||
print(f"\nProcessing: {video_name}")
|
||||
print(f" Caption: {caption[:80]}...")
|
||||
|
||||
video, frames_np = load_video(video_path, NUM_FRAMES, TRAIN_FPS, MAX_HEIGHT, MAX_WIDTH)
|
||||
video = video.to(device=device, dtype=vae_dtype)
|
||||
print(f" Video shape: {video.shape}")
|
||||
|
||||
with torch.no_grad():
|
||||
# LTX-2 encode() returns a deterministic distribution whose
|
||||
# mean is already per-channel normalized; store as-is.
|
||||
latent = vae.encode(video).mean.squeeze(0).float().cpu()
|
||||
print(f" Latent shape: {latent.shape}")
|
||||
|
||||
with torch.no_grad():
|
||||
# Encode through forward() — the inference text path. The
|
||||
# two-step preprocess_text_embeddings + run_connectors route
|
||||
# breaks on LTX-2.3, whose separate audio feature extractor
|
||||
# is narrower than the video one.
|
||||
text_inputs = tokenizer([preprocess_text(caption)], **tokenizer_kwargs)
|
||||
encoder_out = text_encoder(
|
||||
input_ids=text_inputs["input_ids"].to(device),
|
||||
attention_mask=text_inputs["attention_mask"].to(device),
|
||||
)
|
||||
text_embedding = encoder_out.last_hidden_state.squeeze(0).float().cpu()
|
||||
print(f" Text embedding shape: {text_embedding.shape}")
|
||||
|
||||
record = {
|
||||
"id": record_id,
|
||||
"vae_latent_bytes": latent.numpy().tobytes(),
|
||||
"vae_latent_shape": list(latent.shape),
|
||||
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
|
||||
"text_embedding_bytes": (text_embedding.numpy().tobytes()),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
|
||||
"file_name": video_name,
|
||||
"caption": caption,
|
||||
"media_type": "video",
|
||||
"width": MAX_WIDTH,
|
||||
"height": MAX_HEIGHT,
|
||||
"num_frames": NUM_FRAMES,
|
||||
"duration_sec": NUM_FRAMES / TRAIN_FPS,
|
||||
"fps": TRAIN_FPS,
|
||||
}
|
||||
records.append(record)
|
||||
|
||||
# Save the preprocessed clip so overfit tests can compare
|
||||
# validation output against the memorization target.
|
||||
import imageio
|
||||
ref_path = os.path.join(OUTPUT_DIR, f"training_sample_{idx}.mp4")
|
||||
with imageio.get_writer(ref_path, fps=TRAIN_FPS) as writer:
|
||||
for frame in frames_np:
|
||||
writer.append_data(frame)
|
||||
print(f" Wrote reference clip to {ref_path}")
|
||||
|
||||
# Clean up
|
||||
del text_encoder, tokenizer, vae
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Write parquet (replicated NUM_COPIES times; see comment at top)
|
||||
replicated = []
|
||||
for copy_idx in range(max(1, NUM_COPIES)):
|
||||
for r in records:
|
||||
row = dict(r)
|
||||
row["id"] = f"{r['id']}_copy{copy_idx}"
|
||||
replicated.append(row)
|
||||
table = pa.table(
|
||||
{k: [r[k] for r in replicated]
|
||||
for k in replicated[0]},
|
||||
schema=pyarrow_schema_t2v,
|
||||
)
|
||||
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
|
||||
pq.write_table(table, output_path)
|
||||
print(f"\nWrote {len(replicated)} records "
|
||||
f"({len(records)} unique x {max(1, NUM_COPIES)} copies) to {output_path}")
|
||||
|
||||
# Write validation prompts for the validation callback
|
||||
val_prompts = {
|
||||
"data": [{
|
||||
"caption": (item["cap"][0] if isinstance(item["cap"], list) else item["cap"]),
|
||||
} for item in caption_data],
|
||||
}
|
||||
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
|
||||
with open(val_path, "w") as f:
|
||||
json.dump(val_prompts, f, indent=2)
|
||||
print(f"Wrote validation prompts to {val_path}")
|
||||
|
||||
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user