Compare commits
50
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c59f568d56 | ||
|
|
fdea9e9898 | ||
|
|
9b19098434 | ||
|
|
24e2e73454 | ||
|
|
e5bd10678f | ||
|
|
5c749a6911 | ||
|
|
5623c1ba1b | ||
|
|
c72fc2de2c | ||
|
|
41e2d3ee8b | ||
|
|
1ec268cf80 | ||
|
|
637d0bf943 | ||
|
|
417e653604 | ||
|
|
9de2ab7505 | ||
|
|
7b4f757fc2 | ||
|
|
35fe3db250 | ||
|
|
2114d19ced | ||
|
|
895cef22a3 | ||
|
|
284c50fc96 | ||
|
|
94d9dda5cb | ||
|
|
f661bbf104 | ||
|
|
a2a204c25d | ||
|
|
4533447e58 | ||
|
|
c94d625844 | ||
|
|
7427b9d2d3 | ||
|
|
9dab81f180 | ||
|
|
840b43b4c3 | ||
|
|
ace34f866a | ||
|
|
b5f90a6b77 | ||
|
|
ed8f118acd | ||
|
|
b379ea96b4 | ||
|
|
6771f4ec2e | ||
|
|
b48eafe515 | ||
|
|
1d199050af | ||
|
|
8b65f7ce6c | ||
|
|
2d2657821e | ||
|
|
ded49ab1c8 | ||
|
|
10f5372f42 | ||
|
|
ed4125e7d8 | ||
|
|
26e78f8fb1 | ||
|
|
12f94fd6e7 | ||
|
|
facd035b97 | ||
|
|
c3c6c0cb2d | ||
|
|
1ea2517e22 | ||
|
|
0c63528c59 | ||
|
|
b063f8ca41 | ||
|
|
e7fff0173a | ||
|
|
d82abc271e | ||
|
|
970409962f | ||
|
|
055586703d | ||
|
|
5d89f86675 |
@@ -436,7 +436,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
|
||||
+8
-2
@@ -37,6 +37,11 @@ logs/
|
||||
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/**
|
||||
@@ -72,8 +77,7 @@ docs/distillation/examples/
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
@@ -127,6 +131,8 @@ apps/dreamverse/web/.env.production.local
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.png
|
||||
|
||||
# Editor logs and local Python version pins (accidentally committed)
|
||||
*.nvimlog
|
||||
|
||||
@@ -458,6 +458,8 @@ surfaces:
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# 🌊 AnyFlow Any-Step Video Distillation
|
||||
|
||||
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
|
||||
|
||||
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
|
||||
|
||||
```
|
||||
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
|
||||
```
|
||||
|
||||
for any ``t > r``.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
|
||||
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
|
||||
|
||||
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
|
||||
|
||||
## ⚙️ Inference
|
||||
|
||||
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
|
||||
|
||||
## 🧠 Algorithm
|
||||
|
||||
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
|
||||
|
||||
### Stage 1 — Pretrain (flow-map central-difference)
|
||||
|
||||
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
|
||||
|
||||
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
|
||||
|
||||
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
|
||||
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
|
||||
- the remainder is free.
|
||||
|
||||
The student forward at ``(t, r)`` is trained against the central-difference target
|
||||
|
||||
```
|
||||
target = (eps - x_0) - (t - r) · dF/dt
|
||||
```
|
||||
|
||||
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
|
||||
|
||||
### Stage 2 — On-policy DMD
|
||||
|
||||
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
|
||||
|
||||
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
|
||||
|
||||
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Stage 1 — pretrain
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
|
||||
```
|
||||
|
||||
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
|
||||
|
||||
- Global batch size: 32 (8 GPUs × 4 per-GPU)
|
||||
- Learning rate: 5e-5
|
||||
- Flow shift: 5.0
|
||||
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
|
||||
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
|
||||
- ``weight_type``: ``beta08``
|
||||
- ``fuse_guidance_scale``: 3.0
|
||||
- Training steps: 6000
|
||||
|
||||
### Stage 2 — on-policy
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
|
||||
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
|
||||
```
|
||||
|
||||
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
|
||||
|
||||
**Key configuration**:
|
||||
|
||||
- Global batch size: 8 (8 GPUs × 1 per-GPU)
|
||||
- Learning rate: 2e-6
|
||||
- Flow shift: 5.0
|
||||
- ``student_sample_steps``: 4
|
||||
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
|
||||
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
|
||||
- ``real_score_guidance_scale``: 3.0
|
||||
- ``generator_update_interval``: 5 (DMD2 alternation)
|
||||
- Training steps: 4000
|
||||
|
||||
## 🔌 Loading published AnyFlow checkpoints
|
||||
|
||||
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
|
||||
|
||||
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
|
||||
|
||||
## 🧭 Note on ``fuse_guidance_scale``
|
||||
|
||||
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
|
||||
|
||||
```
|
||||
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
|
||||
```
|
||||
|
||||
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
|
||||
|
||||
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
|
||||
@@ -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,140 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _safe_filename(text: str, max_len: int = 100) -> str:
|
||||
"""Make a stable, filesystem-friendly filename base."""
|
||||
s = text[:max_len].strip()
|
||||
s = s.replace(os.sep, "_")
|
||||
if os.altsep:
|
||||
s = s.replace(os.altsep, "_")
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
|
||||
s = s.strip(" .")
|
||||
return s or "prompt"
|
||||
|
||||
|
||||
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
|
||||
if not os.path.isdir(out_dir):
|
||||
return
|
||||
|
||||
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
|
||||
for fn in os.listdir(out_dir):
|
||||
if pattern.match(fn):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.remove(os.path.join(out_dir, fn))
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
help="Local Diffusers checkpoint dir or HF repo id.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
default="outputs/flux_dev/samples",
|
||||
help="Directory for saved PNG outputs.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--prompt",
|
||||
action="append",
|
||||
default=None,
|
||||
help="Prompt. Repeat for multiple images.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
|
||||
)
|
||||
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
|
||||
p.add_argument("--height", type=int, default=1024, help="Output height.")
|
||||
p.add_argument("--width", type=int, default=1024, help="Output width.")
|
||||
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
|
||||
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
|
||||
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
|
||||
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": args.num_gpus,
|
||||
"workload_type": "t2i",
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
"use_fsdp_inference": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
**init_kwargs,
|
||||
)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
generation_kwargs = {
|
||||
"output_path": output_path,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": args.steps,
|
||||
"guidance_scale": args.guidance,
|
||||
"use_embedded_guidance": True,
|
||||
"true_cfg_scale": 1.0,
|
||||
"seed": seed,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
print(f"[flux] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,107 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
|
||||
text-to-image generation command, saved as a PNG."
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/landscape.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default=("A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style."),
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
|
||||
want a minimal edit command (text + image -> edited image), saved as a PNG."
|
||||
|
||||
GLM-Image is a single unified pipeline: passing a condition image switches it
|
||||
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
|
||||
write pass), so the generator config is identical to `basic_glm_image.py` — the
|
||||
`inputs.pil_image` on the request is what selects the edit mode.
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
default="assets/images/couple.jpg",
|
||||
help="Condition image to edit.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/edited.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="Change the background to a snowy mountain landscape at golden hour.",
|
||||
help="Edit instruction.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
condition = Image.open(args.image).convert("RGB")
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
# The pipeline is registered as t2i; passing inputs.pil_image below switches
|
||||
# it to the edit path.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
inputs=InputConfig(pil_image=condition),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,105 @@
|
||||
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
|
||||
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
|
||||
# rollout from pure noise. Teacher provides the real score, critic
|
||||
# learns the fake score; both inherited from DMD2Method.
|
||||
#
|
||||
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
|
||||
# or with the NVIDIA-released checkpoint
|
||||
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
|
||||
# the paper weights (the delta_embedder rename is handled by the
|
||||
# param_names_mapping in WanVideoArchConfig).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: <PATH_TO_PRETRAIN_CKPT>
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 3.0
|
||||
dmd_denoising_steps: [999, 937, 833, 624]
|
||||
warp_denoising_step: false
|
||||
|
||||
# AnyFlow rollout knobs.
|
||||
student_sample_steps: 4
|
||||
use_mean_velocity: true
|
||||
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
|
||||
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
|
||||
|
||||
# Critic optimizer (DMD2 inherited).
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
attn_kind: vsa
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.0, 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: outputs/wan2.1_anyflow_onpolicy
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_onpolicy
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -0,0 +1,83 @@
|
||||
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
|
||||
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
|
||||
# can be sampled at arbitrary NFE in the on-policy stage.
|
||||
#
|
||||
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
|
||||
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
|
||||
diffusion_ratio: 0.5
|
||||
consistency_ratio: 0.25
|
||||
epsilon: 5 # finite-difference step in absolute train-timestep units
|
||||
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
|
||||
fuse_guidance_scale: 3.0
|
||||
# shift is taken from pipeline.flow_shift below.
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 4
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
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: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_pretrain
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_pretrain
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
# Enable AnyFlow dual-timestep conditioning. The student loads from
|
||||
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
|
||||
# get initialized identically to time_embedder via deep-copy in
|
||||
# WanTimeTextImageEmbedding.__init__.
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxSamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 1024
|
||||
width: int = 1024
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 3.5
|
||||
use_embedded_guidance: bool = True
|
||||
true_cfg_scale: float = 1.0
|
||||
@@ -90,6 +90,10 @@ class SamplingParam:
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
# 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).
|
||||
true_cfg_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -325,6 +329,18 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-embedded-guidance",
|
||||
action="store_true",
|
||||
default=SamplingParam.use_embedded_guidance,
|
||||
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--true-cfg-scale",
|
||||
type=float,
|
||||
default=SamplingParam.true_cfg_scale,
|
||||
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
|
||||
@@ -150,6 +150,7 @@ class SamplingConfig:
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
use_embedded_guidance: bool | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention.utils.flash_attn_default import (
|
||||
fa_version,
|
||||
flash_attn_func_compilable,
|
||||
)
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -17,119 +19,6 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
|
||||
# kernel package's FASTVIDEO_VSA_CUTEDSL: its CuTeDSL kernels JIT-compile per
|
||||
# shape family and can fail at runtime on some arch/shape combinations, so it
|
||||
# is never auto-selected just because it is installed. Below sm90 a capability
|
||||
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
|
||||
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
|
||||
# JIT, observed on sm_89).
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer have a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
|
||||
# breaks the graph at the call site (observed: wanvideo.py self-attn,
|
||||
# once per layer every step), which fragments the compiled region and
|
||||
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
|
||||
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
# opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
# (correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
|
||||
if fa_version in ("2", "3"):
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
|
||||
# call shape used by FlashAttentionImpl.forward's default branch
|
||||
# (see `flash_attn_func_compilable(...)` call site below). The
|
||||
# masked/no-pad and varlen / cross-attn paths use different
|
||||
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
|
||||
# are intentionally out of scope for this PR — wrapping them is a
|
||||
# natural follow-up. The wrapper's signature is the contract: any
|
||||
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
|
||||
# return_attn_probs, ...) raises TypeError at the call site, so
|
||||
# silent loss of kwargs is not a failure mode.
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
|
||||
# same dtype/device as q (head dim taken from v).
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
|
||||
# registered torch.library custom op (with an FA4 backward on sm90+;
|
||||
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
|
||||
# is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -309,9 +198,17 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
# Route through the *_compilable wrappers so dynamo sees one
|
||||
# traceable node for each masked entry point (the unpad/pad
|
||||
# bookkeeping runs eager inside the custom op). On FA2 these
|
||||
# wrappers go through ops with full register_autograd, so
|
||||
# training also backprops through the op (no graph break on
|
||||
# the training path); on FA3/FA4 they carve out to the
|
||||
# autograd.Function for grad-enabled calls — see
|
||||
# fastvideo/attention/utils/flash_attn_no_pad.py.
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
flash_attn_no_pad_compilable as flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
|
||||
|
||||
The FA4/cute path (`fa_version == "4"`) is already a registered
|
||||
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
|
||||
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
|
||||
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
|
||||
self-attn, once per layer every step), which fragments the compiled region
|
||||
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
|
||||
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
(correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
boundary changes, so numerics are unchanged (SSIM-gated).
|
||||
|
||||
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
|
||||
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
|
||||
backprops *through* the op (no graph break on the training path either).
|
||||
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
|
||||
because FA3's private backward signature wants validation on a real Hopper
|
||||
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
|
||||
can mirror FA2.
|
||||
|
||||
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
|
||||
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
|
||||
traceable FA default call without pulling in backend dispatch logic. The
|
||||
backend (`attention/backends/flash_attn.py`) just imports
|
||||
`flash_attn_func_compilable` and `fa_version` from here.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
|
||||
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
|
||||
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
|
||||
#
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
|
||||
# kernels JIT-compile per shape family and can fail at runtime on some
|
||||
# arch/shape combinations, so it is never auto-selected just because it is
|
||||
# installed.
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer has a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if fa_version == "2":
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
|
||||
# shape used by FlashAttentionImpl.forward's default branch (see
|
||||
# `flash_attn_func_compilable(...)` call site in
|
||||
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
|
||||
# cross-attn paths use different entry points
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
|
||||
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
|
||||
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
|
||||
# deterministic, return_attn_probs, ...) raises TypeError at the call
|
||||
# site, so silent loss of kwargs is not a failure mode.
|
||||
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
|
||||
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
|
||||
# dropout mask (always None here since dropout_p is fixed at 0).
|
||||
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
|
||||
return out, softmax_lse
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del softmax_scale, causal
|
||||
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
|
||||
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
|
||||
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
|
||||
out = q.new_empty(b, sq, hq, v.shape[-1])
|
||||
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_default_setup_context(ctx, inputs, output):
|
||||
q, k, v, softmax_scale, causal = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(q, k, v, out, lse)
|
||||
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
|
||||
# should differentiate through it. Mark it non-differentiable so
|
||||
# autograd errors loudly if a caller wires it into a loss, rather
|
||||
# than silently producing zero/None grads through the `del grad_lse`
|
||||
# in our backward.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
|
||||
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
|
||||
# float in its C++ schema and rejects None at the binding boundary. Resolve
|
||||
# the default here so the value saved on ctx (and passed to backward) is
|
||||
# always a real float — matches what FA2's own autograd.Function does.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = q.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
|
||||
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
|
||||
# We only differentiate `out`; softmax_lse is saved-for-backward, not
|
||||
# a real differentiable output. (Mirrors the FP4 cute template.)
|
||||
del grad_lse
|
||||
q, k, v, out, lse = ctx.saved_tensors
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
|
||||
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
|
||||
# rng_state) are pinned to the same defaults the forward wrapper
|
||||
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
|
||||
# all of them explicitly. `rng_state=None` is correct for our
|
||||
# `dropout_p=0` configuration.
|
||||
_fa2_backward(
|
||||
grad_out,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
out,
|
||||
lse,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=False,
|
||||
rng_state=None,
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
_flash_attn_default_backward,
|
||||
setup_context=_flash_attn_default_setup_context,
|
||||
)
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Backward is registered: autograd flows through the op (training
|
||||
# path is also traceable; no carve-out needed). Public API matches
|
||||
# `flash_attn_func` — returns just `out`; we drop the saved-for-
|
||||
# backward `lse` here so callers see the original single-tensor
|
||||
# contract.
|
||||
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
return out
|
||||
elif fa_version == "3":
|
||||
# FA3 path: same forward+fake custom op as the original PR #1373, with
|
||||
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
|
||||
# above) wants a Hopper box for grad-check validation, which we don't
|
||||
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
|
||||
# inference traceable + training correct (via the original
|
||||
# autograd.Function path + a pre-PR-style graph break on training).
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
@@ -24,7 +24,7 @@ from flash_attn.bert_padding import pad_input, unpad_input
|
||||
from fastvideo import envs
|
||||
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
@@ -39,20 +39,28 @@ def _resolve_flash_attn_varlen_func() -> Any:
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
return flash_attn_varlen_func_cute, "4"
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
return flash_attn_varlen_func_interface, "3"
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
return flash_attn_varlen_func_flash, "2"
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
|
||||
|
||||
# FA2-only: the private varlen backward we register against the custom ops
|
||||
# below. FA3 / FA4 have different private signatures and validation paths
|
||||
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
|
||||
# pattern from PR #1373 until their setup PRs land.
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
from flash_attn.flash_attn_interface import (
|
||||
_flash_attn_varlen_backward as _fa2_varlen_backward, )
|
||||
|
||||
|
||||
def flash_attn_no_pad(
|
||||
@@ -191,3 +199,473 @@ def flash_attn_varlen_qk_no_pad(
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# torch.compile traceability + register_autograd parity for the masked /
|
||||
# varlen attention paths.
|
||||
#
|
||||
# Wraps the two entry points `FlashAttentionImpl.forward` calls
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
|
||||
# `torch.library.custom_op`s so dynamo sees one traceable node — the
|
||||
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
|
||||
# eager inside the op, and the op's outputs are the statically-shaped
|
||||
# padded tensors. This mirrors the FA2 default-path wrapper in
|
||||
# `fastvideo/attention/backends/flash_attn.py`.
|
||||
#
|
||||
# Autograd: on FA2 we register a real backward (`register_autograd`)
|
||||
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
|
||||
# — re-unpadding the saved padded tensors using the saved mask. The
|
||||
# `softmax_lse` from the varlen forward is naturally unpadded
|
||||
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
|
||||
# the way out (statically shaped) and re-unpad in backward. So
|
||||
# training backprops *through* the op (no graph break on the training
|
||||
# path either).
|
||||
#
|
||||
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
|
||||
# custom op has forward + fake only, and `*_compilable` falls back to
|
||||
# the original autograd.Function for grad-enabled calls. Those legs
|
||||
# are gated on Hopper-class / Blackwell-class boxes for backward
|
||||
# validation and ship as separate follow-ups.
|
||||
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, s, _three, h, d = qkv.shape
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz, h, d] -> [b, s, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
|
||||
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
|
||||
# shaped so register_fake matches.
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
out = qkv.new_empty(b, s, h, d)
|
||||
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
|
||||
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's varlen backward requires a concrete float for softmax_scale.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
# lse is saved-for-backward, not differentiated.
|
||||
del grad_lse
|
||||
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
|
||||
b, s, _three, h, d = qkv.shape
|
||||
|
||||
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
|
||||
# reuse those for out / dout / lse below via direct indexing instead
|
||||
# of redundant `unpad_input` calls (each of which would re-run
|
||||
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
|
||||
|
||||
# Direct-index variants reuse `indices` (computed above).
|
||||
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_s,
|
||||
max_seqlen_k=max_s,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# Re-pad each grad and stack into dqkv.
|
||||
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
|
||||
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
|
||||
return dqkv, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
_flash_attn_no_pad_backward,
|
||||
setup_context=_flash_attn_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, sq, h, d = query.shape
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: [h, nnz_q] -> [b, h, sq]
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
out = query.new_empty(b, sq, h, value.shape[-1])
|
||||
lse = query.new_empty(b, h, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
|
||||
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic) = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
if softmax_scale is None:
|
||||
softmax_scale = query.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
del grad_lse
|
||||
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
|
||||
b, sq, h, d = query.shape
|
||||
sk = key.shape[1]
|
||||
|
||||
# One `unpad_input` call per distinct mask; reuse the returned
|
||||
# indices via direct indexing for everything else that shares
|
||||
# the same mask (v with k_mask; out/dout/lse with q_mask; the
|
||||
# final repad of dk/dv also reuses k_indices). Avoids ~4
|
||||
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
|
||||
# syncs.
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
|
||||
|
||||
# out / dout / lse follow q's shape, so index with q_indices.
|
||||
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# k_indices is already available from the unpad_input above —
|
||||
# no need to recompute it for the dk/dv repad.
|
||||
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dq_padded = _repad(dq_unpad, q_indices, b, sq)
|
||||
dk_padded = _repad(dk_unpad, k_indices, b, sk)
|
||||
dv_padded = _repad(dv_unpad, k_indices, b, sk)
|
||||
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
|
||||
# softmax_scale, deterministic.
|
||||
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
_flash_attn_varlen_qk_no_pad_backward,
|
||||
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- public dispatchers (FA2: autograd flows through the op) -----
|
||||
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
|
||||
full register_autograd on FA2 — both inference and training go through
|
||||
the op, no graph break on either)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
return out
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
|
||||
op, full register_autograd on FA2)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
return out
|
||||
|
||||
else:
|
||||
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
|
||||
# Same pattern as the parked varlen-extension and the FA3 default leg in
|
||||
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
|
||||
# versions is a follow-up gated on Hopper / Blackwell box validation.
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_no_pad( # type: ignore[no-untyped-call]
|
||||
qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
return qkv.new_empty(b, s, h, d)
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# output ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
return query.new_empty(b, sq, h, value.shape[-1])
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and qkv.requires_grad:
|
||||
return flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
|
||||
return flash_attn_varlen_qk_no_pad(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -15,6 +17,7 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
|
||||
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -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,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DArchConfig(DiTArchConfig):
|
||||
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = None
|
||||
num_layers: int = 19
|
||||
num_single_layers: int = 38
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
guidance_embeds: bool = True
|
||||
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
|
||||
prefix: str = "flux"
|
||||
@@ -0,0 +1,61 @@
|
||||
# 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 "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTArchConfig(DiTArchConfig):
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
hidden_size: int = 4096
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_layers: int = 30
|
||||
|
||||
text_embed_dim: int = 1472
|
||||
time_embed_dim: int = 512
|
||||
condition_dim: int = 256
|
||||
|
||||
prior_vq_quantizer_codebook_size: int = 16384
|
||||
|
||||
patch_size: int = 2
|
||||
|
||||
max_height: int = 2048
|
||||
max_width: int = 2048
|
||||
|
||||
qk_norm: str = "layer_norm"
|
||||
eps: float = 1e-5
|
||||
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
|
||||
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
|
||||
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
|
||||
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
|
||||
prefix: str = "GlmImage"
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -19,6 +20,11 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
|
||||
# same internal layout as time_embedder. The regex is harmless on plain
|
||||
# Wan checkpoints (no delta_embedder keys to match).
|
||||
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
@@ -86,6 +92,14 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
|
||||
rope_cache_policy: str = "absolute"
|
||||
|
||||
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
|
||||
# the legacy single-timestep forward (no delta_embedder allocated, no
|
||||
# extra computation on the embedder forward path).
|
||||
r_embedder: bool = False
|
||||
r_embedder_fusion: Literal["additive", "gated"] = "additive"
|
||||
r_embedder_gate_value: float = 0.25
|
||||
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
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
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -15,10 +17,12 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Cosmos3VAEConfig",
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
"GlmImageVAEConfig",
|
||||
]
|
||||
|
||||
@@ -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``.
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import (AutoencoderKLArchConfig, AutoencoderKLVAEConfig)
|
||||
|
||||
_GLM_IMAGE_LATENTS_MEAN: tuple[float, ...] = (
|
||||
-0.2080078125,
|
||||
1.875,
|
||||
-0.470703125,
|
||||
-1.265625,
|
||||
-1.421875,
|
||||
0.77734375,
|
||||
-0.3671875,
|
||||
-0.9453125,
|
||||
0.318359375,
|
||||
0.7734375,
|
||||
-0.1884765625,
|
||||
-0.022216796875,
|
||||
-0.220703125,
|
||||
-1.59375,
|
||||
-0.81640625,
|
||||
-0.255859375,
|
||||
)
|
||||
_GLM_IMAGE_LATENTS_STD: tuple[float, ...] = (
|
||||
3.0625,
|
||||
2.203125,
|
||||
2.265625,
|
||||
4.84375,
|
||||
2.5,
|
||||
3.9375,
|
||||
2.203125,
|
||||
3.03125,
|
||||
2.1875,
|
||||
2.046875,
|
||||
2.71875,
|
||||
2.390625,
|
||||
2.390625,
|
||||
2.453125,
|
||||
2.25,
|
||||
2.15625,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEArchConfig(AutoencoderKLArchConfig):
|
||||
act_fn: str = "silu"
|
||||
block_out_channels: tuple[int, ...] = (128, 512, 1024, 1024)
|
||||
down_block_types: tuple[str, ...] = (
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
)
|
||||
force_upcast: bool = True
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
latents_mean: tuple[float, ...] = _GLM_IMAGE_LATENTS_MEAN
|
||||
latents_std: tuple[float, ...] = _GLM_IMAGE_LATENTS_STD
|
||||
layers_per_block: int = 3
|
||||
mid_block_add_attention: bool = False
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 1024
|
||||
scaling_factor: float = 0.18215
|
||||
shift_factor: float | None = None
|
||||
use_quant_conv: bool = False
|
||||
use_post_quant_conv: bool = False
|
||||
|
||||
temporal_compression_ratio: int = 1
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEConfig(AutoencoderKLVAEConfig):
|
||||
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
|
||||
|
||||
use_tiling: bool = True
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
tile_sample_min_height: int = 512
|
||||
tile_sample_min_width: int = 512
|
||||
tile_sample_stride_height: int = 384
|
||||
tile_sample_stride_width: int = 384
|
||||
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
@@ -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,74 @@
|
||||
# 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.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
T5LargeConfig,
|
||||
)
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def _flux_clip_pooled_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""CLIP branch for FLUX: Diffusers uses pooled prompt embeddings only."""
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError(
|
||||
"FLUX CLIP conditioning requires pooler_output. Ensure the CLIP text encoder returns pooled features.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
def _flux_t5_sequence_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.last_hidden_state is None:
|
||||
raise RuntimeError("FLUX T5 conditioning requires last_hidden_state.")
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxPipelineConfig(PipelineConfig):
|
||||
"""Pipeline layout for Diffusers FLUX.1-dev (CLIP + T5 + packed DiT + FlowMatch)."""
|
||||
|
||||
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
|
||||
transformer_arch: str = "FluxTransformer2DModel"
|
||||
vae_arch: str = "AutoencoderKL"
|
||||
text_encoder_archs: tuple[str, ...] = ("CLIPTextModel", "T5EncoderModel")
|
||||
tokenizer_archs: tuple[str, ...] = ("CLIPTokenizer", "T5TokenizerFast")
|
||||
|
||||
dit_config: FluxDiTConfig = field(default_factory=FluxDiTConfig)
|
||||
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
|
||||
|
||||
embedded_cfg_scale: float = 3.5
|
||||
flow_shift: float | None = None
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (CLIPTextConfig(), T5LargeConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str],
|
||||
...] = field(default_factory=lambda: (preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
|
||||
default_factory=lambda: (_flux_clip_pooled_postprocess, _flux_t5_sequence_postprocess))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", "bf16"))
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
te_cfgs = list(self.text_encoder_configs)
|
||||
if len(te_cfgs) >= 1:
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("max_length", 77)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
if len(te_cfgs) >= 2:
|
||||
cap = 512
|
||||
te_cfgs[1].tokenizer_kwargs["max_length"] = min(int(te_cfgs[1].tokenizer_kwargs.get("max_length", cap)),
|
||||
cap)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
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.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def glm_image_t5_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
assert torch.isnan(hidden_state).sum() == 0, "T5 hidden states contain NaN"
|
||||
|
||||
max_len = 512
|
||||
prompt_embeds = [u[:min(v, max_len)] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_len - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0)
|
||||
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageConfig(PipelineConfig):
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=GlmImageDiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (glm_image_t5_postprocess, ))
|
||||
|
||||
flow_shift: float | None = 1.0
|
||||
embedded_cfg_scale: float = 7.5
|
||||
@@ -295,6 +295,7 @@ def get_1d_rotary_pos_embed(
|
||||
interpolation_factor: float = 1.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
use_real: bool = True,
|
||||
freqs_dtype: torch.dtype | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
@@ -319,12 +320,16 @@ def get_1d_rotary_pos_embed(
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos).float()
|
||||
|
||||
# freqs_dtype is an alias for dtype (Diffusers-compatible calling convention).
|
||||
if freqs_dtype is not None:
|
||||
dtype = freqs_dtype
|
||||
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2, device=pos.device)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
freqs_cos = freqs.cos() # [S, D/2]
|
||||
freqs_sin = freqs.sin() # [S, D/2]
|
||||
@@ -445,6 +450,21 @@ def get_nd_rotary_pos_embed(
|
||||
return cos, sin
|
||||
|
||||
|
||||
_ROTARY_POS_EMBED_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
# Bound the table cache so long-running servers / causal models (which vary
|
||||
# start_frame per frame) cannot grow it without limit; entries are large float64
|
||||
# tensors. Least-recently-used eviction keeps the active resolution(s) hot while
|
||||
# capping memory.
|
||||
_ROTARY_POS_EMBED_CACHE_MAXSIZE = 16
|
||||
|
||||
|
||||
def _hashable(value: Any) -> Any:
|
||||
"""Return a hashable view of a scalar or sequence for use in a cache key."""
|
||||
if isinstance(value, list | tuple):
|
||||
return tuple(value)
|
||||
return value
|
||||
|
||||
|
||||
def get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
@@ -495,6 +515,31 @@ def get_rotary_pos_embed(
|
||||
sp_rank = 0
|
||||
sp_world_size = 1
|
||||
|
||||
# Memoize on every output-affecting argument; the table is constant across
|
||||
# denoising steps, so this avoids recomputing the float64 cos/sin tables.
|
||||
cache_key = (
|
||||
_hashable(rope_sizes),
|
||||
tuple(rope_dim_list),
|
||||
rope_theta,
|
||||
_hashable(theta_rescale_factor),
|
||||
_hashable(interpolation_factor),
|
||||
shard_dim,
|
||||
sp_rank,
|
||||
sp_world_size,
|
||||
dtype,
|
||||
start_frame,
|
||||
use_real,
|
||||
)
|
||||
cached = _ROTARY_POS_EMBED_CACHE.get(cache_key)
|
||||
if cached is not None:
|
||||
# Move to most-recently-used position so the active table is not evicted
|
||||
# when several resolutions / buckets share the process (LRU recency).
|
||||
# Pop with a default: a concurrent eviction between the get() above and
|
||||
# here would otherwise raise KeyError on the hit path.
|
||||
if _ROTARY_POS_EMBED_CACHE.pop(cache_key, None) is not None:
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = cached
|
||||
return cached
|
||||
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -508,6 +553,14 @@ def get_rotary_pos_embed(
|
||||
start_frame=start_frame,
|
||||
use_real=use_real,
|
||||
)
|
||||
# The returned tensors are shared cache entries: callers must never mutate
|
||||
# them in place. Note .to(device) is an identity alias when the tensor is
|
||||
# already on the target device (e.g. CPU runs), so it does NOT guarantee a
|
||||
# copy — treat the tables as read-only and copy before any in-place op.
|
||||
# Reached only on a miss, so evict the least-recently-used entry at capacity.
|
||||
if len(_ROTARY_POS_EMBED_CACHE) >= _ROTARY_POS_EMBED_CACHE_MAXSIZE:
|
||||
_ROTARY_POS_EMBED_CACHE.pop(next(iter(_ROTARY_POS_EMBED_CACHE)))
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = (freqs_cos, freqs_sin)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
|
||||
@@ -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,578 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.layers.rotary_embedding import apply_rotary_emb, get_1d_rotary_pos_embed
|
||||
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.configs.models import DiTConfig
|
||||
from fastvideo.forward_context import get_forward_context, set_forward_context
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.sd3 import (
|
||||
CombinedTimestepTextProjEmbeddings,
|
||||
SD3AdaLayerNormContinuous,
|
||||
SD3AdaLayerNormZero,
|
||||
SD3FeedForward,
|
||||
SD3TextProjection,
|
||||
SD3TimestepEmbedding,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DModelOutput:
|
||||
sample: torch.Tensor
|
||||
|
||||
|
||||
class FluxPosEmbed(nn.Module):
|
||||
"""1D RoPE axes concatenated per Diffusers `FluxPosEmbed`."""
|
||||
|
||||
def __init__(self, theta: int, axes_dim: list[int]) -> None:
|
||||
super().__init__()
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
n_axes = ids.shape[-1]
|
||||
cos_out: list[torch.Tensor] = []
|
||||
sin_out: list[torch.Tensor] = []
|
||||
pos = ids.float()
|
||||
is_mps = ids.device.type == "mps"
|
||||
is_npu = ids.device.type == "npu"
|
||||
freqs_dtype = torch.float32 if (is_mps or is_npu) else torch.float64
|
||||
for i in range(n_axes):
|
||||
cos, sin = get_1d_rotary_pos_embed(
|
||||
self.axes_dim[i],
|
||||
pos[:, i],
|
||||
theta=self.theta,
|
||||
use_real=True,
|
||||
freqs_dtype=freqs_dtype,
|
||||
)
|
||||
cos_out.append(cos)
|
||||
sin_out.append(sin)
|
||||
freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device)
|
||||
freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
class FluxCombinedTimestepGuidanceTextProjEmbeddings(nn.Module):
|
||||
def __init__(self, embedding_dim: int, pooled_projection_dim: int) -> None:
|
||||
super().__init__()
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=256,
|
||||
flip_sin_to_cos=True,
|
||||
downscale_freq_shift=0,
|
||||
)
|
||||
self.timestep_embedder = SD3TimestepEmbedding(
|
||||
in_channels=256,
|
||||
time_embed_dim=embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
self.guidance_embedder = SD3TimestepEmbedding(
|
||||
in_channels=256,
|
||||
time_embed_dim=embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
self.text_embedder = SD3TextProjection(
|
||||
pooled_projection_dim,
|
||||
embedding_dim,
|
||||
act_fn="silu",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
guidance: torch.Tensor,
|
||||
pooled_projection: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype))
|
||||
guidance_proj = self.time_proj(guidance)
|
||||
guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype))
|
||||
time_guidance_emb = timesteps_emb + guidance_emb
|
||||
pooled_projections = self.text_embedder(pooled_projection)
|
||||
return time_guidance_emb + pooled_projections
|
||||
|
||||
|
||||
class FluxAdaLayerNormZeroSingle(nn.Module):
|
||||
def __init__(self, embedding_dim: int, bias: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = ReplicatedLinear(embedding_dim, 3 * embedding_dim, bias=bias)
|
||||
self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
emb: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
emb, _ = self.linear(self.silu(emb))
|
||||
shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1)
|
||||
x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]
|
||||
return x, gate_msa
|
||||
|
||||
|
||||
class FluxJointAttention(nn.Module):
|
||||
"""Joint attention: text tokens precede image tokens (Diffusers order)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.heads = num_attention_heads
|
||||
self.head_dim = attention_head_dim
|
||||
self.inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_added_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_added_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
|
||||
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_q_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_k_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.add_v_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
|
||||
self.to_out = nn.ModuleList(
|
||||
[
|
||||
ReplicatedLinear(self.inner_dim, dim, bias=True),
|
||||
nn.Dropout(0.0),
|
||||
]
|
||||
)
|
||||
self.to_add_out = ReplicatedLinear(self.inner_dim, dim, bias=True)
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=attention_head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
batch_size = hidden_states.shape[0]
|
||||
text_seq_len = encoder_hidden_states.shape[1]
|
||||
img_seq_len = hidden_states.shape[1]
|
||||
|
||||
q, _ = self.to_q(hidden_states)
|
||||
k, _ = self.to_k(hidden_states)
|
||||
v, _ = self.to_v(hidden_states)
|
||||
q = q.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
k = k.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
v = v.view(batch_size, img_seq_len, self.heads, self.head_dim)
|
||||
q = self.norm_q(q)
|
||||
k = self.norm_k(k)
|
||||
|
||||
enc_q, _ = self.add_q_proj(encoder_hidden_states)
|
||||
enc_k, _ = self.add_k_proj(encoder_hidden_states)
|
||||
enc_v, _ = self.add_v_proj(encoder_hidden_states)
|
||||
enc_q = enc_q.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_k = enc_k.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_v = enc_v.view(batch_size, text_seq_len, self.heads, self.head_dim)
|
||||
enc_q = self.norm_added_q(enc_q)
|
||||
enc_k = self.norm_added_k(enc_k)
|
||||
|
||||
q = torch.cat([enc_q, q], dim=1)
|
||||
k = torch.cat([enc_k, k], dim=1)
|
||||
v = torch.cat([enc_v, v], dim=1)
|
||||
|
||||
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
|
||||
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
|
||||
|
||||
joint_out, _ = self.attn(q, k, v)
|
||||
joint_out = joint_out.reshape(batch_size, text_seq_len + img_seq_len, self.inner_dim)
|
||||
|
||||
enc_out = joint_out[:, :text_seq_len]
|
||||
img_out = joint_out[:, text_seq_len:]
|
||||
|
||||
img_out, _ = self.to_out[0](img_out)
|
||||
img_out = self.to_out[1](img_out)
|
||||
enc_out, _ = self.to_add_out(enc_out)
|
||||
return img_out, enc_out
|
||||
|
||||
|
||||
class FluxSingleStreamAttention(nn.Module):
|
||||
"""Self-attention on concatenated text+image sequence (single blocks)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.heads = num_attention_heads
|
||||
self.head_dim = attention_head_dim
|
||||
self.inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
|
||||
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=attention_head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
batch_size, seq_len, _ = hidden_states.shape
|
||||
q, _ = self.to_q(hidden_states)
|
||||
k, _ = self.to_k(hidden_states)
|
||||
v, _ = self.to_v(hidden_states)
|
||||
q = q.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
k = k.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
v = v.view(batch_size, seq_len, self.heads, self.head_dim)
|
||||
q = self.norm_q(q)
|
||||
k = self.norm_k(k)
|
||||
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
|
||||
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
|
||||
out, _ = self.attn(q, k, v)
|
||||
return out.reshape(batch_size, seq_len, self.inner_dim)
|
||||
|
||||
|
||||
class FluxTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = SD3AdaLayerNormZero(dim)
|
||||
self.norm1_context = SD3AdaLayerNormZero(dim)
|
||||
self.attn = FluxJointAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
||||
self.ff = SD3FeedForward(
|
||||
dim=dim,
|
||||
dim_out=dim,
|
||||
activation_fn="gelu-approximate",
|
||||
)
|
||||
self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
||||
self.ff_context = SD3FeedForward(
|
||||
dim=dim,
|
||||
dim_out=dim,
|
||||
activation_fn="gelu-approximate",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del joint_attention_kwargs
|
||||
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb)
|
||||
(norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp) = self.norm1_context(
|
||||
encoder_hidden_states, emb=temb
|
||||
)
|
||||
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
|
||||
attn_output = gate_msa.unsqueeze(1) * attn_output
|
||||
hidden_states = hidden_states + attn_output
|
||||
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
|
||||
|
||||
context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output
|
||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
||||
|
||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||
norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states + (c_gate_mlp.unsqueeze(1) * context_ff_output)
|
||||
if encoder_hidden_states.dtype == torch.float16:
|
||||
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
||||
|
||||
return encoder_hidden_states, hidden_states
|
||||
|
||||
|
||||
class FluxSingleTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.norm = FluxAdaLayerNormZeroSingle(dim)
|
||||
self.proj_mlp = ReplicatedLinear(dim, mlp_hidden_dim, bias=True)
|
||||
self.act_mlp = nn.GELU(approximate="tanh")
|
||||
self.proj_out = ReplicatedLinear(dim + mlp_hidden_dim, dim, bias=True)
|
||||
self.attn = FluxSingleStreamAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del joint_attention_kwargs
|
||||
text_seq_len = encoder_hidden_states.shape[1]
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
residual = hidden_states
|
||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||
mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)[0])
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||
gate = gate.unsqueeze(1)
|
||||
hidden_states = gate * self.proj_out(hidden_states)[0]
|
||||
hidden_states = residual + hidden_states
|
||||
if hidden_states.dtype == torch.float16:
|
||||
hidden_states = hidden_states.clip(-65504, 65504)
|
||||
encoder_hidden_states = hidden_states[:, :text_seq_len]
|
||||
hidden_states = hidden_states[:, text_seq_len:]
|
||||
return encoder_hidden_states, hidden_states
|
||||
|
||||
|
||||
class FluxTransformer2DModel(BaseDiT):
|
||||
"""FastVideo FLUX transformer; load Diffusers FLUX safetensors 1:1."""
|
||||
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: (n.startswith("transformer_blocks.") or n.startswith("single_transformer_blocks."))
|
||||
and n.split(".")[-1].isdigit(),
|
||||
]
|
||||
_compile_conditions = _fsdp_shard_conditions
|
||||
# HF weight names already match this module layout (cf. SGLang regex maps).
|
||||
param_names_mapping: dict[str, Any] = {}
|
||||
reverse_param_names_mapping: dict[str, Any] = {}
|
||||
lora_param_names_mapping: dict[str, Any] = {}
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
|
||||
del kwargs
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.fastvideo_config = config
|
||||
self.hf_config = hf_config
|
||||
arch = config.arch_config
|
||||
|
||||
out_ch = arch.out_channels
|
||||
self.out_channels = out_ch if out_ch is not None else arch.in_channels
|
||||
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
|
||||
self.hidden_size = self.inner_dim
|
||||
self.num_attention_heads = arch.num_attention_heads
|
||||
self.num_channels_latents = arch.in_channels
|
||||
|
||||
axes_list = list(arch.axes_dims_rope)
|
||||
self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_list)
|
||||
if arch.guidance_embeds:
|
||||
self.time_text_embed = FluxCombinedTimestepGuidanceTextProjEmbeddings(
|
||||
embedding_dim=self.inner_dim,
|
||||
pooled_projection_dim=arch.pooled_projection_dim,
|
||||
)
|
||||
else:
|
||||
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
|
||||
embedding_dim=self.inner_dim,
|
||||
pooled_projection_dim=arch.pooled_projection_dim,
|
||||
)
|
||||
self.context_embedder = ReplicatedLinear(arch.joint_attention_dim, self.inner_dim)
|
||||
self.x_embedder = ReplicatedLinear(arch.in_channels, self.inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
FluxTransformerBlock(
|
||||
dim=self.inner_dim,
|
||||
num_attention_heads=arch.num_attention_heads,
|
||||
attention_head_dim=arch.attention_head_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
)
|
||||
for _ in range(arch.num_layers)
|
||||
]
|
||||
)
|
||||
self.single_transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
FluxSingleTransformerBlock(
|
||||
dim=self.inner_dim,
|
||||
num_attention_heads=arch.num_attention_heads,
|
||||
attention_head_dim=arch.attention_head_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
)
|
||||
for _ in range(arch.num_single_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm_out = SD3AdaLayerNormContinuous(
|
||||
self.inner_dim,
|
||||
self.inner_dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
bias=True,
|
||||
norm_type="layer_norm",
|
||||
)
|
||||
self.proj_out = ReplicatedLinear(
|
||||
self.inner_dim,
|
||||
arch.patch_size * arch.patch_size * self.out_channels,
|
||||
bias=True,
|
||||
)
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | None = None,
|
||||
pooled_projections: torch.Tensor | None = None,
|
||||
timestep: torch.LongTensor | torch.Tensor | None = None,
|
||||
img_ids: torch.Tensor | None = None,
|
||||
txt_ids: torch.Tensor | None = None,
|
||||
guidance: torch.Tensor | None = None,
|
||||
joint_attention_kwargs: dict[str, Any] | None = None,
|
||||
return_dict: bool = True,
|
||||
controlnet_block_samples: Any | None = None,
|
||||
controlnet_single_block_samples: Any | None = None,
|
||||
controlnet_blocks_repeat: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> FluxTransformer2DModelOutput | tuple[torch.Tensor, ...]:
|
||||
del kwargs
|
||||
if encoder_hidden_states is None:
|
||||
raise ValueError("encoder_hidden_states must be provided")
|
||||
if pooled_projections is None:
|
||||
raise ValueError("pooled_projections must be provided")
|
||||
if timestep is None:
|
||||
raise ValueError("timestep must be provided")
|
||||
if img_ids is None or txt_ids is None:
|
||||
raise ValueError("img_ids and txt_ids must be provided")
|
||||
|
||||
arch = self.fastvideo_config.arch_config
|
||||
if arch.guidance_embeds and guidance is None:
|
||||
raise ValueError("guidance must be provided when guidance_embeds=True")
|
||||
|
||||
if timestep.dim() == 0:
|
||||
timestep = timestep[None]
|
||||
if timestep.dim() > 1:
|
||||
timestep = timestep.reshape(-1)
|
||||
if timestep.shape[0] == 1 and hidden_states.shape[0] > 1:
|
||||
timestep = timestep.expand(hidden_states.shape[0])
|
||||
|
||||
try:
|
||||
get_forward_context()
|
||||
forward_context = nullcontext()
|
||||
except AssertionError:
|
||||
if timestep.numel() == 0:
|
||||
ts0 = 0
|
||||
elif torch.is_floating_point(timestep):
|
||||
ts0 = int(round(timestep[0].item() * 1000))
|
||||
else:
|
||||
ts0 = int(timestep[0].item())
|
||||
forward_context = set_forward_context(current_timestep=ts0, attn_metadata=None)
|
||||
|
||||
with forward_context:
|
||||
hidden_states, _ = self.x_embedder(hidden_states)
|
||||
|
||||
ts = timestep.to(hidden_states.dtype) * 1000
|
||||
g = None if guidance is None else guidance.to(hidden_states.dtype) * 1000
|
||||
|
||||
if arch.guidance_embeds:
|
||||
assert g is not None
|
||||
temb = self.time_text_embed(ts, g, pooled_projections)
|
||||
else:
|
||||
temb = self.time_text_embed(timestep=ts, pooled_projection=pooled_projections)
|
||||
|
||||
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
||||
|
||||
if txt_ids.ndim == 3:
|
||||
txt_ids = txt_ids[0]
|
||||
if img_ids.ndim == 3:
|
||||
img_ids = img_ids[0]
|
||||
|
||||
ids = torch.cat((txt_ids, img_ids), dim=0)
|
||||
image_rotary_emb = self.pos_embed(ids)
|
||||
|
||||
jkwargs = joint_attention_kwargs or {}
|
||||
|
||||
for idx, block in enumerate(self.transformer_blocks):
|
||||
encoder_hidden_states, hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
joint_attention_kwargs=jkwargs,
|
||||
)
|
||||
if controlnet_block_samples:
|
||||
interval = len(self.transformer_blocks) / len(controlnet_block_samples)
|
||||
interval = int(math.ceil(interval))
|
||||
if controlnet_blocks_repeat:
|
||||
hidden_states = hidden_states + controlnet_block_samples[idx % len(controlnet_block_samples)]
|
||||
else:
|
||||
hidden_states = hidden_states + controlnet_block_samples[idx // interval]
|
||||
|
||||
for idx, block in enumerate(self.single_transformer_blocks):
|
||||
encoder_hidden_states, hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
joint_attention_kwargs=jkwargs,
|
||||
)
|
||||
if controlnet_single_block_samples:
|
||||
interval = len(self.single_transformer_blocks) / len(controlnet_single_block_samples)
|
||||
interval = int(math.ceil(interval))
|
||||
hidden_states = hidden_states + controlnet_single_block_samples[idx // interval]
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
output, _ = self.proj_out(hidden_states)
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
return FluxTransformer2DModelOutput(sample=output)
|
||||
|
||||
|
||||
EntryClass = FluxTransformer2DModel
|
||||
@@ -0,0 +1,776 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from fastvideo.layers.mlp import MLP
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.layers.layernorm import ScaleResidualLayerNormScaleShift
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class GlmImageLayerKVCache:
|
||||
|
||||
def __init__(self):
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.mode: Optional[str] = None
|
||||
|
||||
def store(self, k: torch.Tensor, v: torch.Tensor):
|
||||
# Append along seq (dim=1).
|
||||
if self.k_cache is None:
|
||||
self.k_cache = k
|
||||
self.v_cache = v
|
||||
else:
|
||||
self.k_cache = torch.cat([self.k_cache, k], dim=1)
|
||||
self.v_cache = torch.cat([self.v_cache, v], dim=1)
|
||||
|
||||
def get(self):
|
||||
return self.k_cache, self.v_cache
|
||||
|
||||
def clear(self):
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.mode = None
|
||||
|
||||
|
||||
class GlmImageKVCache:
|
||||
|
||||
def __init__(self, num_layers: int):
|
||||
self.num_layers = num_layers
|
||||
self.caches = [GlmImageLayerKVCache() for _ in range(num_layers)]
|
||||
|
||||
def __getitem__(self, layer_idx: int) -> GlmImageLayerKVCache:
|
||||
return self.caches[layer_idx]
|
||||
|
||||
def set_mode(self, mode: Optional[str]):
|
||||
if mode is not None and mode not in ["write", "read", "skip"]:
|
||||
raise ValueError(
|
||||
f"Invalid mode: {mode}, must be one of 'write', 'read', 'skip'"
|
||||
)
|
||||
for cache in self.caches:
|
||||
cache.mode = mode
|
||||
|
||||
def clear(self):
|
||||
for cache in self.caches:
|
||||
cache.clear()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Timestep and Text Projection
|
||||
# =============================================================================
|
||||
class GlmImageTimestepEmbedding(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
time_embed_dim: int,
|
||||
act_fn: str = "silu",
|
||||
out_dim: int = None,
|
||||
):
|
||||
super().__init__()
|
||||
if out_dim is None:
|
||||
out_dim = time_embed_dim
|
||||
self.linear_1 = ReplicatedLinear(in_channels, time_embed_dim, bias=True)
|
||||
if act_fn == "silu":
|
||||
self.act = nn.SiLU()
|
||||
elif act_fn == "gelu":
|
||||
self.act = nn.GELU(approximate="tanh")
|
||||
else:
|
||||
self.act = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(time_embed_dim, out_dim, bias=True)
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
sample, _ = self.linear_1(sample)
|
||||
sample = self.act(sample)
|
||||
sample, _ = self.linear_2(sample)
|
||||
return sample
|
||||
|
||||
|
||||
class GlmImageTextProjection(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
hidden_size: int,
|
||||
out_features: int = None,
|
||||
act_fn: str = "silu",
|
||||
):
|
||||
super().__init__()
|
||||
if out_features is None:
|
||||
out_features = hidden_size
|
||||
self.linear_1 = ReplicatedLinear(in_features, hidden_size, bias=True)
|
||||
if act_fn == "silu":
|
||||
self.act_1 = nn.SiLU()
|
||||
elif act_fn == "gelu_tanh":
|
||||
self.act_1 = nn.GELU(approximate="tanh")
|
||||
else:
|
||||
self.act_1 = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(hidden_size, out_features, bias=True)
|
||||
|
||||
def forward(self, caption: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.linear_1(caption)
|
||||
hidden_states = self.act_1(hidden_states)
|
||||
hidden_states, _ = self.linear_2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class GlmImageCombinedTimestepSizeEmbeddings(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
condition_dim: int,
|
||||
pooled_projection_dim: int,
|
||||
timesteps_dim: int = 256,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=timesteps_dim, flip_sin_to_cos=True, downscale_freq_shift=0
|
||||
)
|
||||
self.condition_proj = Timesteps(
|
||||
num_channels=condition_dim, flip_sin_to_cos=True, downscale_freq_shift=0
|
||||
)
|
||||
self.timestep_embedder = GlmImageTimestepEmbedding(
|
||||
in_channels=timesteps_dim, time_embed_dim=embedding_dim
|
||||
)
|
||||
self.condition_embedder = GlmImageTextProjection(
|
||||
pooled_projection_dim, embedding_dim, act_fn="silu"
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
target_size: torch.Tensor,
|
||||
crop_coords: torch.Tensor,
|
||||
hidden_dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
|
||||
crop_coords_proj = self.condition_proj(crop_coords.flatten()).view(
|
||||
crop_coords.size(0), -1
|
||||
)
|
||||
target_size_proj = self.condition_proj(target_size.flatten()).view(
|
||||
target_size.size(0), -1
|
||||
)
|
||||
|
||||
condition_proj = torch.cat([crop_coords_proj, target_size_proj], dim=1)
|
||||
|
||||
timesteps_emb = self.timestep_embedder(
|
||||
timesteps_proj.to(dtype=hidden_dtype)
|
||||
)
|
||||
condition_emb = self.condition_embedder(
|
||||
condition_proj.to(dtype=hidden_dtype)
|
||||
)
|
||||
|
||||
conditioning = timesteps_emb + condition_emb
|
||||
return conditioning
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Image Projector
|
||||
# =============================================================================
|
||||
class GlmImageImageProjector(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 16,
|
||||
hidden_size: int = 2560,
|
||||
patch_size: int = 2,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.proj = nn.Linear(in_channels * patch_size**2, hidden_size)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
post_patch_height = height // self.patch_size
|
||||
post_patch_width = width // self.patch_size
|
||||
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size,
|
||||
channel,
|
||||
post_patch_height,
|
||||
self.patch_size,
|
||||
post_patch_width,
|
||||
self.patch_size,
|
||||
)
|
||||
hidden_states = (
|
||||
hidden_states.permute(0, 2, 4, 1, 3, 5).flatten(3, 5).flatten(1, 2)
|
||||
)
|
||||
hidden_states = self.proj(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# AdaLayerNorm
|
||||
# =============================================================================
|
||||
class GlmImageAdaLayerNormZero(nn.Module):
|
||||
|
||||
def __init__(self, embedding_dim: int, dim: int) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
|
||||
self.norm_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
|
||||
self.linear = ReplicatedLinear(embedding_dim, 12 * dim, bias=True)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, ...]:
|
||||
dtype = hidden_states.dtype
|
||||
norm_hidden_states = self.norm(hidden_states).to(dtype=dtype)
|
||||
norm_encoder_hidden_states = self.norm_context(encoder_hidden_states).to(
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
emb, _ = self.linear(temb)
|
||||
(
|
||||
shift_msa,
|
||||
c_shift_msa,
|
||||
scale_msa,
|
||||
c_scale_msa,
|
||||
gate_msa,
|
||||
c_gate_msa,
|
||||
shift_mlp,
|
||||
c_shift_mlp,
|
||||
scale_mlp,
|
||||
c_scale_mlp,
|
||||
gate_mlp,
|
||||
c_gate_mlp,
|
||||
) = emb.chunk(12, dim=1)
|
||||
|
||||
hidden_states = norm_hidden_states * (
|
||||
1 + scale_msa.unsqueeze(1)
|
||||
) + shift_msa.unsqueeze(1)
|
||||
encoder_hidden_states = norm_encoder_hidden_states * (
|
||||
1 + c_scale_msa.unsqueeze(1)
|
||||
) + c_shift_msa.unsqueeze(1)
|
||||
|
||||
return (
|
||||
hidden_states,
|
||||
gate_msa,
|
||||
shift_mlp,
|
||||
scale_mlp,
|
||||
gate_mlp,
|
||||
encoder_hidden_states,
|
||||
c_gate_msa,
|
||||
c_shift_mlp,
|
||||
c_scale_mlp,
|
||||
c_gate_mlp,
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Attention
|
||||
# =============================================================================
|
||||
class GlmImageAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
heads: int,
|
||||
dim_head: int,
|
||||
out_dim: int,
|
||||
bias: bool = True,
|
||||
qk_norm: str = "layer_norm",
|
||||
elementwise_affine: bool = False,
|
||||
eps: float = 1e-5,
|
||||
supported_attention_backends: set[AttentionBackendEnum] | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
self.num_kv_heads = self.heads
|
||||
self.dim_head = dim_head
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.inner_kv_dim = self.inner_dim
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
|
||||
self.to_q = ReplicatedLinear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_k = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
|
||||
self.to_v = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
|
||||
|
||||
self.to_out = nn.ModuleList(
|
||||
[ReplicatedLinear(self.inner_dim, self.out_dim, bias=True)]
|
||||
)
|
||||
|
||||
if qk_norm is None:
|
||||
self.norm_q = None
|
||||
self.norm_k = None
|
||||
elif qk_norm == "layer_norm":
|
||||
self.norm_q = nn.LayerNorm(
|
||||
dim_head, eps=eps, elementwise_affine=elementwise_affine
|
||||
)
|
||||
self.norm_k = nn.LayerNorm(
|
||||
dim_head, eps=eps, elementwise_affine=elementwise_affine
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown qk_norm: {qk_norm}")
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=self.heads,
|
||||
head_size=dim_head,
|
||||
num_kv_heads=self.heads,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
kv_cache: Optional[GlmImageLayerKVCache] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
dtype = encoder_hidden_states.dtype
|
||||
|
||||
batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape
|
||||
batch_size, image_seq_length, embed_dim = hidden_states.shape
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
|
||||
# 1. QKV projections
|
||||
query, _ = self.to_q(hidden_states)
|
||||
key, _ = self.to_k(hidden_states)
|
||||
value, _ = self.to_v(hidden_states)
|
||||
|
||||
query = query.unflatten(2, (self.heads, -1))
|
||||
key = key.unflatten(2, (self.heads, -1))
|
||||
value = value.unflatten(2, (self.heads, -1))
|
||||
|
||||
# 2. QK normalization
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query).to(dtype=dtype)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key).to(dtype=dtype)
|
||||
|
||||
# 3. Rotational positional embeddings applied to latent stream
|
||||
if image_rotary_emb is not None:
|
||||
cos, sin = image_rotary_emb
|
||||
|
||||
query[:, text_seq_length:, :, :] = _apply_rotary_emb(
|
||||
query[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
|
||||
)
|
||||
key[:, text_seq_length:, :, :] = _apply_rotary_emb(
|
||||
key[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
|
||||
)
|
||||
|
||||
# 4. KV Cache handling
|
||||
if kv_cache is not None:
|
||||
if kv_cache.mode == "write":
|
||||
kv_cache.store(key, value)
|
||||
elif kv_cache.mode == "read":
|
||||
# Prepend cached condition k/v along seq (dim=1).
|
||||
k_cache, v_cache = kv_cache.get()
|
||||
key = torch.cat([k_cache, key], dim=1) if k_cache is not None else key
|
||||
value = (
|
||||
torch.cat([v_cache, value], dim=1) if v_cache is not None else value
|
||||
)
|
||||
elif kv_cache.mode == "skip":
|
||||
pass
|
||||
|
||||
hidden_states = self.attn(query, key, value)
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# 6. Output projection
|
||||
hidden_states, _ = self.to_out[0](hidden_states)
|
||||
|
||||
encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Transformer Block
|
||||
# =============================================================================
|
||||
class GlmImageTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int = 2560,
|
||||
num_attention_heads: int = 64,
|
||||
attention_head_dim: int = 40,
|
||||
time_embed_dim: int = 512,
|
||||
supported_attention_backends: set[AttentionBackendEnum] | None = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# 1. Attention
|
||||
self.norm1 = GlmImageAdaLayerNormZero(time_embed_dim, dim)
|
||||
|
||||
self.attn1 = GlmImageAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
out_dim=dim,
|
||||
bias=True,
|
||||
qk_norm="layer_norm",
|
||||
elementwise_affine=False,
|
||||
eps=1e-5,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn1",
|
||||
)
|
||||
|
||||
# 2. Feedforward with fused ScaleResidualLayerNorm
|
||||
self.norm2 = ScaleResidualLayerNormScaleShift(
|
||||
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
|
||||
)
|
||||
self.norm2_context = ScaleResidualLayerNormScaleShift(
|
||||
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
|
||||
)
|
||||
self.ff = MLP(input_dim=dim, mlp_hidden_dim=dim * 4, output_dim=dim, act_type="gelu_pytorch_tanh")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[
|
||||
Union[
|
||||
Tuple[torch.Tensor, torch.Tensor],
|
||||
List[Tuple[torch.Tensor, torch.Tensor]],
|
||||
]
|
||||
] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
kv_cache: Optional[GlmImageLayerKVCache] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# 1. Timestep conditioning
|
||||
(
|
||||
norm_hidden_states,
|
||||
gate_msa,
|
||||
shift_mlp,
|
||||
scale_mlp,
|
||||
gate_mlp,
|
||||
norm_encoder_hidden_states,
|
||||
c_gate_msa,
|
||||
c_shift_mlp,
|
||||
c_scale_mlp,
|
||||
c_gate_mlp,
|
||||
) = self.norm1(hidden_states, encoder_hidden_states, temb)
|
||||
|
||||
# 2. Attention
|
||||
if attention_kwargs is None:
|
||||
attention_kwargs = {}
|
||||
|
||||
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
kv_cache=kv_cache,
|
||||
**attention_kwargs,
|
||||
)
|
||||
|
||||
# 3. Feedforward (fused residual + norm + scale/shift)
|
||||
norm_hidden_states, hidden_states = self.norm2(
|
||||
hidden_states,
|
||||
attn_hidden_states,
|
||||
gate_msa.unsqueeze(1),
|
||||
shift_mlp.unsqueeze(1),
|
||||
scale_mlp.unsqueeze(1),
|
||||
)
|
||||
norm_encoder_hidden_states, encoder_hidden_states = self.norm2_context(
|
||||
encoder_hidden_states,
|
||||
attn_encoder_hidden_states,
|
||||
c_gate_msa.unsqueeze(1),
|
||||
c_shift_mlp.unsqueeze(1),
|
||||
c_scale_mlp.unsqueeze(1),
|
||||
)
|
||||
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
ff_output_context = self.ff(norm_encoder_hidden_states)
|
||||
hidden_states = hidden_states + ff_output * gate_mlp.unsqueeze(1)
|
||||
encoder_hidden_states = (
|
||||
encoder_hidden_states + ff_output_context * c_gate_mlp.unsqueeze(1)
|
||||
)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Rotary Positional Embedding
|
||||
# =============================================================================
|
||||
class GlmImageRotaryPosEmbed(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, patch_size: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.patch_size = patch_size
|
||||
self.theta = theta
|
||||
self._cache_key: tuple | None = None
|
||||
self._cache_value: tuple[torch.Tensor, torch.Tensor] | None = None
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
_, _, raw_h, raw_w = hidden_states.shape
|
||||
height = raw_h // self.patch_size
|
||||
width = raw_w // self.patch_size
|
||||
device = hidden_states.device
|
||||
|
||||
cache_key = (height, width, device.type,
|
||||
device.index if device.index is not None else -1)
|
||||
if self._cache_key == cache_key and self._cache_value is not None:
|
||||
return self._cache_value
|
||||
|
||||
dim_h, dim_w = self.dim // 2, self.dim // 2
|
||||
h_inv_freq = 1.0 / (
|
||||
self.theta
|
||||
** (
|
||||
torch.arange(0, dim_h, 2, dtype=torch.float32, device=device)[
|
||||
: (dim_h // 2)
|
||||
]
|
||||
/ dim_h
|
||||
)
|
||||
)
|
||||
w_inv_freq = 1.0 / (
|
||||
self.theta
|
||||
** (
|
||||
torch.arange(0, dim_w, 2, dtype=torch.float32, device=device)[
|
||||
: (dim_w // 2)
|
||||
]
|
||||
/ dim_w
|
||||
)
|
||||
)
|
||||
h_seq = torch.arange(height, device=device)
|
||||
w_seq = torch.arange(width, device=device)
|
||||
freqs_h = torch.outer(h_seq, h_inv_freq).unsqueeze(1).expand(height, width, -1)
|
||||
freqs_w = torch.outer(w_seq, w_inv_freq).unsqueeze(0).expand(height, width, -1)
|
||||
|
||||
freqs = torch.cat([freqs_h, freqs_w], dim=-1).reshape(height * width, -1)
|
||||
result = (freqs.cos(), freqs.sin())
|
||||
self._cache_key = cache_key
|
||||
self._cache_value = result
|
||||
return result
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Final AdaLayerNorm
|
||||
# =============================================================================
|
||||
class GlmImageAdaLayerNormContinuous(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
conditioning_embedding_dim: int,
|
||||
elementwise_affine: bool = True,
|
||||
eps: float = 1e-5,
|
||||
bias: bool = True,
|
||||
norm_type: str = "layer_norm",
|
||||
):
|
||||
super().__init__()
|
||||
self.linear = nn.Linear(
|
||||
conditioning_embedding_dim, embedding_dim * 2, bias=bias
|
||||
)
|
||||
if norm_type == "layer_norm":
|
||||
self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias)
|
||||
elif norm_type == "rms_norm":
|
||||
self.norm = nn.RMSNorm(embedding_dim, eps, elementwise_affine)
|
||||
else:
|
||||
raise ValueError(f"unknown norm_type {norm_type}")
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, conditioning_embedding: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
emb = self.linear(conditioning_embedding.to(x.dtype))
|
||||
scale, shift = torch.chunk(emb, 2, dim=1)
|
||||
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||
return x
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Main Model
|
||||
# =============================================================================
|
||||
class GlmImageTransformer2DModel(BaseDiT):
|
||||
|
||||
_fsdp_shard_conditions = GlmImageDiTConfig().arch_config._fsdp_shard_conditions
|
||||
_compile_conditions = GlmImageDiTConfig().arch_config._compile_conditions
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
param_names_mapping = GlmImageDiTConfig().arch_config.param_names_mapping
|
||||
reverse_param_names_mapping = {}
|
||||
lora_param_names_mapping = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: GlmImageDiTConfig,
|
||||
hf_config: dict[str, Any],
|
||||
):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
arch_config = config.arch_config
|
||||
|
||||
self.in_channels = arch_config.in_channels
|
||||
self.out_channels = arch_config.out_channels
|
||||
self.patch_size = arch_config.patch_size
|
||||
self.num_layers = arch_config.num_layers
|
||||
self.attention_head_dim = arch_config.attention_head_dim
|
||||
self.num_attention_heads = arch_config.num_attention_heads
|
||||
self.text_embed_dim = arch_config.text_embed_dim
|
||||
self.time_embed_dim = arch_config.time_embed_dim
|
||||
|
||||
# GlmImage uses 2 additional SDXL-like conditions - target_size, crop_coords
|
||||
# Each of these are sincos embeddings of shape 2 * condition_dim
|
||||
pooled_projection_dim = 2 * 2 * arch_config.condition_dim
|
||||
inner_dim = arch_config.num_attention_heads * arch_config.attention_head_dim
|
||||
|
||||
self.hidden_size = inner_dim
|
||||
self.num_channels_latents = arch_config.out_channels
|
||||
|
||||
# 1. RoPE
|
||||
self.rotary_emb = GlmImageRotaryPosEmbed(
|
||||
arch_config.attention_head_dim, arch_config.patch_size, theta=10000.0
|
||||
)
|
||||
|
||||
# 2. Patch & Text-timestep embedding
|
||||
self.image_projector = GlmImageImageProjector(
|
||||
arch_config.in_channels, inner_dim, arch_config.patch_size
|
||||
)
|
||||
self.glyph_projector = MLP(
|
||||
input_dim=arch_config.text_embed_dim,
|
||||
mlp_hidden_dim=inner_dim,
|
||||
output_dim=inner_dim,
|
||||
act_type="gelu",
|
||||
)
|
||||
self.prior_token_embedding = nn.Embedding(
|
||||
arch_config.prior_vq_quantizer_codebook_size, inner_dim
|
||||
)
|
||||
self.prior_projector = MLP(
|
||||
input_dim=inner_dim,
|
||||
mlp_hidden_dim=inner_dim,
|
||||
output_dim=inner_dim,
|
||||
act_type="silu",
|
||||
)
|
||||
|
||||
self.time_condition_embed = GlmImageCombinedTimestepSizeEmbeddings(
|
||||
embedding_dim=arch_config.time_embed_dim,
|
||||
condition_dim=arch_config.condition_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
timesteps_dim=arch_config.time_embed_dim,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
GlmImageTransformerBlock(
|
||||
inner_dim,
|
||||
arch_config.num_attention_heads,
|
||||
arch_config.attention_head_dim,
|
||||
arch_config.time_embed_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"transformer_blocks.{i}",
|
||||
)
|
||||
for i in range(arch_config.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output projection
|
||||
self.norm_out = GlmImageAdaLayerNormContinuous(
|
||||
inner_dim, arch_config.time_embed_dim, elementwise_affine=False
|
||||
)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim,
|
||||
arch_config.patch_size * arch_config.patch_size * arch_config.out_channels,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
prior_token_id: torch.Tensor,
|
||||
prior_token_drop: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
target_size: torch.Tensor,
|
||||
crop_coords: torch.Tensor,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
kv_caches: Optional[GlmImageKVCache] = None,
|
||||
kv_caches_mode: Optional[str] = None,
|
||||
freqs_cis: Optional[
|
||||
Union[
|
||||
Tuple[torch.Tensor, torch.Tensor],
|
||||
List[Tuple[torch.Tensor, torch.Tensor]],
|
||||
]
|
||||
] = None,
|
||||
guidance: torch.Tensor = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if kv_caches is not None:
|
||||
kv_caches.set_mode(kv_caches_mode)
|
||||
|
||||
batch_size, num_channels, height, width = hidden_states.shape
|
||||
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
# 1. RoPE
|
||||
image_rotary_emb = freqs_cis
|
||||
if image_rotary_emb is None:
|
||||
image_rotary_emb = self.rotary_emb(hidden_states)
|
||||
|
||||
# 2. Patch & Timestep embeddings
|
||||
p = self.patch_size
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
hidden_states = self.image_projector(hidden_states)
|
||||
encoder_hidden_states = self.glyph_projector(encoder_hidden_states)
|
||||
prior_embedding = self.prior_token_embedding(prior_token_id)
|
||||
# Zero dropped priors by multiply: boolean indexing + .any() syncs each step.
|
||||
keep = (~prior_token_drop).to(device=prior_embedding.device, dtype=prior_embedding.dtype)
|
||||
while keep.dim() < prior_embedding.dim():
|
||||
keep = keep.unsqueeze(-1)
|
||||
prior_embedding = prior_embedding * keep
|
||||
prior_hidden_states = self.prior_projector(prior_embedding)
|
||||
hidden_states = hidden_states + prior_hidden_states
|
||||
|
||||
temb = self.time_condition_embed(
|
||||
timestep, target_size, crop_coords, hidden_states.dtype
|
||||
)
|
||||
temb = F.silu(temb)
|
||||
|
||||
# 3. Transformer blocks
|
||||
for idx, block in enumerate(self.transformer_blocks):
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
attention_kwargs,
|
||||
kv_cache=kv_caches[idx] if kv_caches is not None else None,
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
# 5. Unpatchify
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_height, post_patch_width, -1, p, p
|
||||
)
|
||||
output = hidden_states.permute(0, 3, 1, 4, 2, 5).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output.float()
|
||||
|
||||
|
||||
EntryClass = GlmImageTransformer2DModel
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import copy
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
@@ -59,6 +60,11 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
time_freq_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: int | None = None,
|
||||
*,
|
||||
r_embedder: bool = False,
|
||||
r_embedder_fusion: str = "additive",
|
||||
r_embedder_gate_value: float = 0.25,
|
||||
r_embedder_deltatime_type: str = "r",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -77,14 +83,57 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
|
||||
|
||||
# AnyFlow dual-timestep support. When r_embedder is False the forward
|
||||
# path bypasses delta_embedder entirely and the output is byte-identical
|
||||
# to the legacy single-timestep implementation.
|
||||
self._r_embedder_enabled = bool(r_embedder)
|
||||
self._r_embedder_fusion = r_embedder_fusion
|
||||
self._r_embedder_deltatime_type = r_embedder_deltatime_type
|
||||
if self._r_embedder_enabled:
|
||||
if r_embedder_fusion not in ("additive", "gated"):
|
||||
raise ValueError(
|
||||
"r_embedder_fusion must be one of {additive, gated}, "
|
||||
f"got {r_embedder_fusion!r}")
|
||||
if r_embedder_deltatime_type not in ("r", "t-r"):
|
||||
raise ValueError(
|
||||
"r_embedder_deltatime_type must be one of {r, t-r}, "
|
||||
f"got {r_embedder_deltatime_type!r}")
|
||||
# Deep-copy preserves identical initialization with time_embedder,
|
||||
# matching AnyFlow reference setup_flowmap_model() behavior.
|
||||
self.delta_embedder = copy.deepcopy(self.time_embedder)
|
||||
# Non-persistent buffer — gate is a hyperparameter, not learned.
|
||||
self.register_buffer(
|
||||
"_r_embedder_gate",
|
||||
torch.tensor(float(r_embedder_gate_value)),
|
||||
persistent=False,
|
||||
)
|
||||
else:
|
||||
self.delta_embedder = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: torch.Tensor | None = None,
|
||||
timestep_seq_len: int | None = None,
|
||||
r_timestep: torch.Tensor | None = None,
|
||||
):
|
||||
temb = self.time_embedder(timestep, timestep_seq_len)
|
||||
|
||||
if self._r_embedder_enabled and r_timestep is not None:
|
||||
assert self.delta_embedder is not None
|
||||
if self._r_embedder_deltatime_type == "r":
|
||||
delta_input = r_timestep
|
||||
else:
|
||||
delta_input = timestep - r_timestep
|
||||
delta_emb = self.delta_embedder(delta_input, timestep_seq_len)
|
||||
gate = self._r_embedder_gate
|
||||
if self._r_embedder_fusion == "gated":
|
||||
temb = (1.0 - gate) * temb + gate * delta_emb
|
||||
else:
|
||||
# Additive (additional channel, no convex blend).
|
||||
temb = temb + gate * delta_emb
|
||||
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
if self.text_embedder is not None:
|
||||
@@ -595,6 +644,10 @@ class WanTransformer3DModel(BaseDiT):
|
||||
time_freq_dim=config.freq_dim,
|
||||
text_embed_dim=config.text_dim,
|
||||
image_embed_dim=config.image_dim,
|
||||
r_embedder=config.r_embedder,
|
||||
r_embedder_fusion=config.r_embedder_fusion,
|
||||
r_embedder_gate_value=config.r_embedder_gate_value,
|
||||
r_embedder_deltatime_type=config.r_embedder_deltatime_type,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
@@ -636,6 +689,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
guidance=None,
|
||||
r_timestep: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
orig_dtype = hidden_states.dtype
|
||||
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
@@ -684,8 +738,17 @@ class WanTransformer3DModel(BaseDiT):
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
# AnyFlow dual-timestep — match timestep's flattening so embedder
|
||||
# sees aligned shapes.
|
||||
if r_timestep is not None and r_timestep.dim() == 2:
|
||||
r_timestep = r_timestep.flatten()
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states_image,
|
||||
timestep_seq_len=ts_seq_len,
|
||||
r_timestep=r_timestep)
|
||||
if ts_seq_len is not None:
|
||||
# batch_size, seq_len, 6, inner_dim
|
||||
timestep_proj = timestep_proj.unflatten(2, (6, -1))
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Sole HF-import boundary for GLM-Image's AR encoder (lazy-wrapper exception E001)."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GlmImageARLoader(nn.Module):
|
||||
|
||||
def __init__(self, model_path: str, processor_path: str | None = None,
|
||||
*, torch_dtype: torch.dtype = torch.bfloat16,
|
||||
trust_remote_code: bool = True) -> None:
|
||||
super().__init__()
|
||||
from transformers import (AutoProcessor,
|
||||
GlmImageForConditionalGeneration)
|
||||
logger.info("Loading GLM-Image AR encoder from %s", model_path)
|
||||
self._model = GlmImageForConditionalGeneration.from_pretrained(
|
||||
model_path,
|
||||
torch_dtype=torch_dtype,
|
||||
trust_remote_code=trust_remote_code,
|
||||
)
|
||||
if processor_path is not None:
|
||||
logger.info("Loading GLM-Image processor from %s", processor_path)
|
||||
self.processor = AutoProcessor.from_pretrained(
|
||||
processor_path, trust_remote_code=trust_remote_code)
|
||||
else:
|
||||
self.processor = None
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(self, *args: Any, **kwargs: Any) -> torch.Tensor:
|
||||
return self._model.generate(*args, **kwargs)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_image_features(self, pixel_values: torch.Tensor,
|
||||
image_grid_thw: torch.Tensor) -> Any:
|
||||
return self._model.get_image_features(pixel_values, image_grid_thw)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_image_tokens(self, image_embeds: torch.Tensor,
|
||||
image_grid_thw: torch.Tensor) -> torch.Tensor:
|
||||
return self._model.get_image_tokens(image_embeds, image_grid_thw)
|
||||
|
||||
@property
|
||||
def config(self): # type: ignore[no-untyped-def]
|
||||
return self._model.config
|
||||
|
||||
@property
|
||||
def generation_config(self): # type: ignore[no-untyped-def]
|
||||
return self._model.generation_config
|
||||
|
||||
def to(self, *args, **kwargs): # type: ignore[override]
|
||||
self._model = self._model.to(*args, **kwargs)
|
||||
return super().to(*args, **kwargs)
|
||||
|
||||
def eval(self): # type: ignore[override]
|
||||
self._model = self._model.eval()
|
||||
return super().eval()
|
||||
@@ -42,7 +42,7 @@ def load_independent(files: list[str], device: str):
|
||||
"""Before-PR behavior: every rank reads every tensor from disk to GPU."""
|
||||
for st_file in files:
|
||||
with safe_open(st_file, framework="pt", device=device) as f:
|
||||
for name in f:
|
||||
for name in f.keys(): # noqa: SIM118
|
||||
param = f.get_tensor(name)
|
||||
yield name, param
|
||||
|
||||
@@ -54,7 +54,7 @@ def load_broadcast(files: list[str], device: str, node_group,
|
||||
handles = []
|
||||
for st_file in files:
|
||||
with safe_open(st_file, framework="pt", device=device) as f:
|
||||
for name in f:
|
||||
for name in f.keys(): # noqa: SIM118
|
||||
if local_rank == 0:
|
||||
param = f.get_tensor(name)
|
||||
else:
|
||||
|
||||
@@ -92,9 +92,13 @@ 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"),
|
||||
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
|
||||
"processor": (ProcessorLoader, "transformers"),
|
||||
"upsampler": (UpsamplerLoader, "diffusers"),
|
||||
"upsampler_2": (UpsamplerLoader, "diffusers"),
|
||||
# Stable Audio's `StableAudioMultiConditioner` bundles T5 +
|
||||
@@ -523,6 +527,39 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
)
|
||||
|
||||
|
||||
class VisionLanguageEncoderLoader(ComponentLoader):
|
||||
"""Loader for vision-language autoregressive encoders."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.models.encoders.glm_image_ar_loader import (
|
||||
GlmImageARLoader)
|
||||
|
||||
logger.info("Loading vision-language encoder from %s", model_path)
|
||||
target_device = get_local_torch_device()
|
||||
loader = GlmImageARLoader(
|
||||
model_path,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
).to(target_device).eval()
|
||||
return loader
|
||||
|
||||
|
||||
class ProcessorLoader(ComponentLoader):
|
||||
"""Loader for HF processors that pair with vision-language encoders."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
from transformers import AutoProcessor
|
||||
|
||||
logger.info("Loading processor from %s", model_path)
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
)
|
||||
logger.info("Loaded processor: %s", processor.__class__.__name__)
|
||||
return processor
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
@@ -1102,7 +1139,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
|
||||
|
||||
@@ -37,6 +37,9 @@ _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"),
|
||||
@@ -61,6 +64,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
"MatrixGame3WanModel": ("dits", "matrixgame3", "MatrixGame3WanModel"),
|
||||
}
|
||||
|
||||
# Text-to-image DiT models (2D image generation)
|
||||
_TEXT_TO_IMAGE_DIT_MODELS = {
|
||||
"GlmImageTransformer2DModel": ("dits", "glm_image", "GlmImageTransformer2DModel"),
|
||||
}
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"CLIPTextModelWithProjection":
|
||||
@@ -134,6 +142,7 @@ _UPSAMPLERS = {
|
||||
_LEGACY_FAST_VIDEO_MODELS = {
|
||||
**_TEXT_TO_VIDEO_DIT_MODELS,
|
||||
**_IMAGE_TO_VIDEO_DIT_MODELS,
|
||||
**_TEXT_TO_IMAGE_DIT_MODELS,
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Flow-map any-step Euler scheduler for AnyFlow.
|
||||
|
||||
The model predicts the *average* velocity ``u_θ(x_t, t, r)`` from time
|
||||
``t`` back to time ``r``, so one Euler step is
|
||||
|
||||
x_r = x_t - ((t - r) / num_train_timesteps) * u_θ(x_t, t, r)
|
||||
|
||||
regardless of how far apart ``t`` and ``r`` are. The scheduler also
|
||||
provides the AnyFlow training-time helpers ``apply_shift`` (flow-matching
|
||||
shift transform) and ``get_train_weight`` (per-timestep loss weight,
|
||||
including ``beta08``).
|
||||
|
||||
Standalone — does not depend on diffusers' ConfigMixin/SchedulerMixin.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.schedulers.base import BaseScheduler
|
||||
|
||||
|
||||
WeightType = Literal["uniform", "gaussian", "beta08"]
|
||||
|
||||
|
||||
class FlowMapEulerDiscreteScheduler(BaseScheduler):
|
||||
"""Minimal flow-map scheduler.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
num_train_timesteps:
|
||||
Discretization granularity for training. ``t`` is expressed in
|
||||
absolute units in ``[0, num_train_timesteps]``.
|
||||
shift:
|
||||
Flow-matching shift parameter (Wan video default: ``5.0``). Set
|
||||
to ``1.0`` for an identity shift.
|
||||
"""
|
||||
|
||||
order: int = 1
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
) -> None:
|
||||
self.num_train_timesteps = int(num_train_timesteps)
|
||||
self.shift = float(shift)
|
||||
self.timesteps: torch.Tensor = torch.empty(0)
|
||||
self.sigmas: torch.Tensor = torch.empty(0)
|
||||
super().__init__()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# BaseScheduler abstract surface
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.shift = float(shift)
|
||||
|
||||
def scale_model_input(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
timestep: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
# Flow-matching has no per-step input scaling; pass through.
|
||||
del timestep
|
||||
return sample
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public helpers used by the AnyFlow pretrain method.
|
||||
|
||||
def apply_shift(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
*,
|
||||
shift: float | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply the flow-matching shift: ``t' = s * t / (1 + (s - 1) * t)``.
|
||||
|
||||
Operates in the normalized ``[0, 1]`` domain — callers should pass
|
||||
``t / num_train_timesteps`` (or sample ``t`` directly from
|
||||
``[0, 1]``).
|
||||
"""
|
||||
s = self.shift if shift is None else float(shift)
|
||||
if s == 1.0:
|
||||
return t
|
||||
return s * t / (1.0 + (s - 1.0) * t)
|
||||
|
||||
def get_train_weight(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
*,
|
||||
weight_type: WeightType = "beta08",
|
||||
) -> torch.Tensor:
|
||||
"""Per-timestep training weight, renormalized so the total weight
|
||||
mass equals ``num_train_timesteps`` (matching AnyFlow reference's
|
||||
``scheduling_flowmap_euler_discrete.py``).
|
||||
|
||||
``beta08``: ``w(t) = t * sqrt(1 - t)`` (in normalized t-space).
|
||||
"""
|
||||
# Auto-detect domain: if t was given in absolute units, normalize.
|
||||
t_f = t.float()
|
||||
max_val = t_f.max() if t_f.numel() > 0 else torch.tensor(0.0)
|
||||
if max_val > 1.0 + 1e-6:
|
||||
t_norm = t_f / self.num_train_timesteps
|
||||
else:
|
||||
t_norm = t_f
|
||||
t_norm = t_norm.clamp(min=0.0, max=1.0)
|
||||
|
||||
if weight_type == "uniform":
|
||||
w = torch.ones_like(t_norm)
|
||||
elif weight_type == "gaussian":
|
||||
w = torch.exp(-0.5 * ((t_norm - 0.5) / 0.2) ** 2)
|
||||
elif weight_type == "beta08":
|
||||
w = t_norm.pow(1.0) * (1.0 - t_norm).clamp_min(0.0).pow(0.5)
|
||||
else:
|
||||
raise ValueError(f"Unknown weight_type: {weight_type!r}")
|
||||
|
||||
denom = w.sum().clamp_min(1e-8)
|
||||
return w * (float(self.num_train_timesteps) / denom)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
*,
|
||||
num_inference_steps: int,
|
||||
device: torch.device | str = "cpu",
|
||||
custom_timesteps: list[float] | torch.Tensor | None = None,
|
||||
) -> None:
|
||||
"""Build a descending timestep schedule ending at 0.
|
||||
|
||||
With ``num_inference_steps=N`` the schedule has ``N + 1`` entries
|
||||
``[T_max, ..., 0]`` so a rollout consumes ``N`` Euler steps.
|
||||
|
||||
``custom_timesteps`` overrides the linspace+shift schedule with
|
||||
a pinned list (in absolute train-timestep units), useful for the
|
||||
AnyFlow paper's hand-tuned ``[999, 937, 833, 624, 0]`` schedule.
|
||||
"""
|
||||
if num_inference_steps <= 0:
|
||||
raise ValueError(
|
||||
"num_inference_steps must be positive, "
|
||||
f"got {num_inference_steps}")
|
||||
device = torch.device(device)
|
||||
|
||||
if custom_timesteps is not None:
|
||||
ts = torch.as_tensor(
|
||||
custom_timesteps, dtype=torch.float32, device=device)
|
||||
if ts.ndim != 1:
|
||||
raise ValueError(
|
||||
"custom_timesteps must be 1-D, got shape "
|
||||
f"{tuple(ts.shape)}")
|
||||
if not torch.all(ts[:-1] >= ts[1:]):
|
||||
raise ValueError(
|
||||
"custom_timesteps must be descending (largest first)")
|
||||
else:
|
||||
ts_norm = torch.linspace(
|
||||
1.0, 0.0, num_inference_steps + 1, device=device)
|
||||
ts_norm = self.apply_shift(ts_norm)
|
||||
ts = ts_norm * self.num_train_timesteps
|
||||
|
||||
self.timesteps = ts
|
||||
self.sigmas = ts / self.num_train_timesteps
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
*,
|
||||
sample: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
r_timestep: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""One Euler step from ``t`` to ``r``.
|
||||
|
||||
``model_output`` is the average-velocity prediction
|
||||
``u_θ(x_t, t, r)``. Both ``timestep`` and ``r_timestep`` are in
|
||||
absolute train-timestep units (``[0, num_train_timesteps]``).
|
||||
"""
|
||||
t = timestep.to(sample.device, dtype=sample.dtype)
|
||||
r = r_timestep.to(sample.device, dtype=sample.dtype)
|
||||
dt_norm = (t - r) / float(self.num_train_timesteps)
|
||||
# Broadcast dt over channel/spatial dims.
|
||||
view: list[int] = [-1] + [1] * (sample.ndim - 1)
|
||||
return sample - dt_norm.view(*view) * model_output
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Linear flow-matching interpolation: ``x_t = (1 - σ) * x_0 + σ * ε``,
|
||||
where ``σ = t / num_train_timesteps``.
|
||||
"""
|
||||
sigma = (timestep.to(original_samples.device,
|
||||
dtype=original_samples.dtype)
|
||||
/ float(self.num_train_timesteps))
|
||||
view: list[int] = [-1] + [1] * (original_samples.ndim - 1)
|
||||
sigma = sigma.view(*view)
|
||||
return (1.0 - sigma) * original_samples + sigma * noise
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages.flux_stages import (
|
||||
FluxConditioningStage,
|
||||
FluxDecodingStage,
|
||||
FluxDenoisingStage,
|
||||
FluxInputValidationStage,
|
||||
FluxLatentPreparationStage,
|
||||
FluxTimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
|
||||
|
||||
class FluxPipeline(ComposedPipelineBase):
|
||||
"""FLUX.1-dev T2I (Diffusers module layout, packed latents, embedded guidance)."""
|
||||
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"transformer",
|
||||
"vae",
|
||||
"text_encoder",
|
||||
"text_encoder_2",
|
||||
"tokenizer",
|
||||
"tokenizer_2",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage", stage=FluxInputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2"),
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2"),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(stage_name="flux_conditioning_stage", stage=FluxConditioningStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=FluxTimestepPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=FluxLatentPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=FluxDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=FluxDecodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = FluxPipeline
|
||||
@@ -0,0 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GLM-Image pipeline package."""
|
||||
|
||||
from fastvideo.pipelines.basic.glm_image.glm_image_pipeline import (
|
||||
GlmImagePipeline, )
|
||||
|
||||
__all__ = ["GlmImagePipeline"]
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.basic.glm_image.stages import (
|
||||
GlmImageBeforeDenoisingStage,
|
||||
GlmImageConditionEncodingStage,
|
||||
GlmImageDecodingStage,
|
||||
GlmImageDenoisingStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages import InputValidationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
pipeline_name = "GlmImagePipeline"
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"vision_language_encoder",
|
||||
"processor",
|
||||
]
|
||||
|
||||
_optional_config_modules: list[str] = []
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=1.0)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="glm_image_before_denoising_stage",
|
||||
stage=GlmImageBeforeDenoisingStage(
|
||||
vae=self.get_module("vae"),
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
processor=self.get_module("processor"),
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vision_language_encoder=self.get_module("vision_language_encoder"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="glm_image_condition_encoding_stage",
|
||||
stage=GlmImageConditionEncodingStage(
|
||||
vae=self.get_module("vae"),
|
||||
transformer=self.get_module("transformer"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=GlmImageDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=GlmImageDecodingStage(
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = GlmImagePipeline
|
||||
@@ -0,0 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GLM-Image pipeline stages."""
|
||||
|
||||
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import (GlmImageBeforeDenoisingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.condition_encoding import (GlmImageConditionEncodingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.decoding import (GlmImageDecodingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.denoising import (GlmImageDenoisingStage)
|
||||
|
||||
__all__ = [
|
||||
"GlmImageBeforeDenoisingStage",
|
||||
"GlmImageConditionEncodingStage",
|
||||
"GlmImageDecodingStage",
|
||||
"GlmImageDenoisingStage",
|
||||
]
|
||||
@@ -0,0 +1,269 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from math import sqrt
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len: int,
|
||||
base_seq_len: int = 256,
|
||||
base_shift: float = 0.25,
|
||||
max_shift: float = 0.75,
|
||||
) -> float:
|
||||
return (image_seq_len / base_seq_len)**0.5 * max_shift + base_shift
|
||||
|
||||
|
||||
def get_glyph_texts(prompt: str | list[str]) -> list[str] | list[list[str]]:
|
||||
if isinstance(prompt, str):
|
||||
prompts: list[str] = [prompt]
|
||||
is_batch = False
|
||||
else:
|
||||
prompts = prompt
|
||||
is_batch = True
|
||||
out: list[list[str]] = []
|
||||
for p in prompts:
|
||||
out.append(
|
||||
re.findall(r"'([^']*)'", p) + re.findall(r"“([^“”]*)”", p) + re.findall(r'"([^"]*)"', p) +
|
||||
re.findall(r"「([^「」]*)」", p))
|
||||
return out if is_batch else out[0]
|
||||
|
||||
|
||||
def compute_glyph_embeds(
|
||||
prompts: list[str],
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
max_sequence_length: int = 2048,
|
||||
) -> torch.Tensor:
|
||||
all_glyph_texts = get_glyph_texts(prompts)
|
||||
all_glyph_embeds = []
|
||||
for glyph_texts in all_glyph_texts:
|
||||
if len(glyph_texts) == 0:
|
||||
glyph_texts = [""]
|
||||
input_ids = tokenizer(
|
||||
glyph_texts,
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
).input_ids
|
||||
input_ids = [[tokenizer.pad_token_id] * ((len(input_ids) + 1) % 2) + ids for ids in input_ids]
|
||||
max_length = max(len(ids) for ids in input_ids)
|
||||
attention_mask = torch.tensor(
|
||||
[[1] * len(ids) + [0] * (max_length - len(ids)) for ids in input_ids],
|
||||
device=device,
|
||||
)
|
||||
input_ids_t = torch.tensor(
|
||||
[ids + [tokenizer.pad_token_id] * (max_length - len(ids)) for ids in input_ids],
|
||||
device=device,
|
||||
)
|
||||
outputs = text_encoder(input_ids_t, attention_mask=attention_mask)
|
||||
glyph_embeds = outputs.last_hidden_state[attention_mask.bool()].unsqueeze(0)
|
||||
all_glyph_embeds.append(glyph_embeds)
|
||||
|
||||
max_seq_len = max(emb.size(1) for emb in all_glyph_embeds)
|
||||
padded = []
|
||||
for emb in all_glyph_embeds:
|
||||
if emb.size(1) < max_seq_len:
|
||||
pad = torch.zeros(emb.size(0), max_seq_len - emb.size(1), emb.size(2), device=device, dtype=emb.dtype)
|
||||
emb = torch.cat([pad, emb], dim=1)
|
||||
padded.append(emb)
|
||||
return torch.cat(padded, dim=0).to(device=device, dtype=dtype)
|
||||
|
||||
|
||||
def _grid_dims(height: int, width: int) -> tuple[int, int, int, int]:
|
||||
th, tw = height // 32, width // 32
|
||||
ratio = th / tw
|
||||
pth = int(sqrt(ratio) * 16)
|
||||
ptw = int(sqrt(1 / ratio) * 16)
|
||||
return th, tw, pth, ptw
|
||||
|
||||
|
||||
def _upsample_d32_to_d16(tokens: torch.Tensor, th: int, tw: int) -> torch.Tensor:
|
||||
tokens = tokens.view(1, 1, th, tw).float()
|
||||
tokens = torch.nn.functional.interpolate(tokens, scale_factor=2, mode="nearest").long()
|
||||
return tokens.view(1, -1)
|
||||
|
||||
|
||||
class GlmImageBeforeDenoisingStage(PipelineStage):
|
||||
|
||||
def __init__(self,
|
||||
vae,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
processor,
|
||||
transformer,
|
||||
scheduler,
|
||||
vision_language_encoder=None) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
self.processor = processor
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
if isinstance(vision_language_encoder, tuple):
|
||||
self.vision_language_encoder, self.vl_processor = (vision_language_encoder[0], vision_language_encoder[1]
|
||||
or processor)
|
||||
else:
|
||||
self.vision_language_encoder = vision_language_encoder
|
||||
self.vl_processor = processor
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
th, tw, pth, ptw = _grid_dims(batch.height, batch.width)
|
||||
|
||||
if batch.seed is not None:
|
||||
torch.manual_seed(batch.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(batch.seed)
|
||||
|
||||
# 1-3. AR token generation. I2I prepends the condition image and uses a
|
||||
# single-scale target grid; T2I is multi-scale.
|
||||
is_t2i = batch.pil_image is None
|
||||
if self.vision_language_encoder is not None:
|
||||
content = [{"type": "text", "text": batch.prompt}]
|
||||
if not is_t2i:
|
||||
content.insert(0, {"type": "image", "image": batch.pil_image})
|
||||
messages = [{"role": "user", "content": content}]
|
||||
inputs = self.vl_processor.apply_chat_template(messages,
|
||||
tokenize=True,
|
||||
target_h=batch.height,
|
||||
target_w=batch.width,
|
||||
return_dict=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
if is_t2i:
|
||||
up_h, up_w = th, tw
|
||||
large_start, large_count = pth * ptw, th * tw
|
||||
max_new = large_count + (pth * ptw) + 1
|
||||
else:
|
||||
# Condition grid(s) first, target grid last.
|
||||
_, t_h, t_w = inputs["image_grid_thw"][-1].tolist()
|
||||
up_h, up_w = int(t_h), int(t_w)
|
||||
large_start, large_count = 0, up_h * up_w
|
||||
max_new = large_count + 1
|
||||
|
||||
outputs = self.vision_language_encoder.generate(**inputs, max_new_tokens=max_new, do_sample=True)
|
||||
gen_tokens = outputs[0][inputs.input_ids.shape[-1]:]
|
||||
if gen_tokens.shape[0] >= large_start + large_count:
|
||||
large_tokens = gen_tokens[large_start:large_start + large_count]
|
||||
else:
|
||||
available = gen_tokens[large_start:]
|
||||
large_tokens = torch.zeros(large_count, dtype=gen_tokens.dtype, device=gen_tokens.device)
|
||||
if available.shape[0] > 0:
|
||||
large_tokens[:min(available.shape[0], large_count)] = available[:large_count]
|
||||
logger.warning("AR generated %d tokens, expected %d. Padding with zeros.", gen_tokens.shape[0],
|
||||
large_start + large_count)
|
||||
batch.prior_token_id = _upsample_d32_to_d16(large_tokens, up_h, up_w)
|
||||
batch.prior_token_drop = torch.zeros(batch.prior_token_id.shape, dtype=torch.bool, device=device)
|
||||
|
||||
if not is_t2i:
|
||||
self._compute_source_prior_tokens(batch, inputs)
|
||||
else:
|
||||
num_prior_tokens = 4 * th * tw
|
||||
logger.warning("No vision_language_encoder provided; using random dropped priors.")
|
||||
batch.prior_token_id = torch.randint(0, 16384, (1, num_prior_tokens), device=device)
|
||||
batch.prior_token_drop = torch.ones(batch.prior_token_id.shape, dtype=torch.bool, device=device)
|
||||
|
||||
# 4. Glyph T5 encoding.
|
||||
prompts = [batch.prompt] if isinstance(batch.prompt, str) else list(batch.prompt)
|
||||
prompt_embeds = compute_glyph_embeds(prompts, self.tokenizer, self.text_encoder, device, dtype)
|
||||
|
||||
# 5. CFG-side negative encoding.
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompts = [batch.negative_prompt or ""] * len(prompts)
|
||||
neg_embeds = compute_glyph_embeds(neg_prompts, self.tokenizer, self.text_encoder, device, dtype)
|
||||
L_pos, L_neg = prompt_embeds.shape[1], neg_embeds.shape[1]
|
||||
max_L = max(L_pos, L_neg)
|
||||
if L_pos < max_L:
|
||||
pad = torch.zeros(prompt_embeds.shape[0],
|
||||
max_L - L_pos,
|
||||
prompt_embeds.shape[2],
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
prompt_embeds = torch.cat([pad, prompt_embeds], dim=1)
|
||||
if L_neg < max_L:
|
||||
pad = torch.zeros(neg_embeds.shape[0], max_L - L_neg, neg_embeds.shape[2], device=device, dtype=dtype)
|
||||
neg_embeds = torch.cat([pad, neg_embeds], dim=1)
|
||||
# Row 0 conditional (positive), row 1 unconditional (negative).
|
||||
prompt_embeds = torch.cat([prompt_embeds, neg_embeds], dim=0)
|
||||
att_pos = torch.ones((1, max_L), device=device)
|
||||
att_neg = torch.ones((1, max_L), device=device)
|
||||
if L_pos < max_L:
|
||||
att_pos[:, :max_L - L_pos] = 0
|
||||
if L_neg < max_L:
|
||||
att_neg[:, :max_L - L_neg] = 0
|
||||
attention_mask = torch.cat([att_pos, att_neg], dim=0)
|
||||
else:
|
||||
attention_mask = torch.ones((1, prompt_embeds.shape[1]), device=device)
|
||||
|
||||
batch.prompt_embeds = [prompt_embeds]
|
||||
batch.attention_mask = attention_mask
|
||||
|
||||
# 6. Latents + dynamic flow shift.
|
||||
if batch.seed is not None:
|
||||
torch.manual_seed(batch.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(batch.seed)
|
||||
batch.latents = torch.randn((1, 16, 1, batch.height // 8, batch.width // 8), device=device, dtype=dtype)
|
||||
|
||||
# Integer-cast linspace timesteps with resolution-dependent shift applied to
|
||||
# sigmas only; the DiT is conditioned on the unshifted integer timesteps.
|
||||
ntt = self.scheduler.config.num_train_timesteps
|
||||
patch_size = self.transformer.patch_size
|
||||
image_seq_len = ((batch.height // 8) * (batch.width // 8)) // (patch_size**2)
|
||||
sched_timesteps = np.linspace(ntt, 1.0, batch.num_inference_steps + 1)[:-1].astype(np.int64).astype(np.float32)
|
||||
sched_sigmas = sched_timesteps / ntt
|
||||
self.scheduler.set_shift(calculate_shift(image_seq_len))
|
||||
self.scheduler.set_timesteps(batch.num_inference_steps,
|
||||
device=device,
|
||||
sigmas=sched_sigmas.tolist(),
|
||||
timesteps=sched_timesteps.tolist())
|
||||
batch.timesteps = self.scheduler.timesteps
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def _compute_source_prior_tokens(self, batch: ForwardBatch, inputs) -> None:
|
||||
image_grid_thw = inputs["image_grid_thw"]
|
||||
num_condition_images = image_grid_thw.shape[0] - 1
|
||||
source_grids = image_grid_thw[:num_condition_images]
|
||||
image_features = self.vision_language_encoder.get_image_features(inputs["pixel_values"], source_grids)
|
||||
image_feature_parts = getattr(image_features, "pooler_output", image_features)
|
||||
embed = torch.cat(image_feature_parts, dim=0)
|
||||
src_ids_d32 = self.vision_language_encoder.get_image_tokens(embed, source_grids)
|
||||
split_sizes = source_grids.prod(dim=-1).tolist()
|
||||
upsampled = [
|
||||
_upsample_d32_to_d16(ids, int(grid[1]), int(grid[2])).squeeze(0)
|
||||
for ids, grid in zip(torch.split(src_ids_d32, split_sizes), source_grids, strict=False)
|
||||
]
|
||||
src_grids_up = source_grids.clone()
|
||||
src_grids_up[:, 1] *= 2
|
||||
src_grids_up[:, 2] *= 2
|
||||
batch.extra["glm_prior_token_image_ids"] = torch.cat(upsampled, dim=0)
|
||||
batch.extra["glm_source_image_grid_thw"] = src_grids_up
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_not_empty)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("prior_token_id", batch.prior_token_id, V.is_tensor)
|
||||
return result
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
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.image_processor import ImageProcessor
|
||||
from fastvideo.models.dits.glm_image import GlmImageKVCache
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
_CONDITION_MULTIPLE_OF = 16 # vae_scale_factor (8) * DiT patch_size (2)
|
||||
|
||||
|
||||
class GlmImageConditionEncodingStage(PipelineStage):
|
||||
|
||||
def __init__(self, vae, transformer) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
self.transformer = transformer
|
||||
self.image_processor = ImageProcessor(vae_scale_factor=_CONDITION_MULTIPLE_OF)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.pil_image is None:
|
||||
return batch
|
||||
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
self.vae.to(device)
|
||||
|
||||
prior_ids = batch.extra["glm_prior_token_image_ids"].to(device)
|
||||
if prior_ids.dim() == 1:
|
||||
prior_ids = prior_ids.unsqueeze(0)
|
||||
|
||||
# Latent patch count must match the source prior tokens; mismatch is fatal.
|
||||
src_grid = batch.extra["glm_source_image_grid_thw"][0]
|
||||
cond_h = int(src_grid[1]) * _CONDITION_MULTIPLE_OF
|
||||
cond_w = int(src_grid[2]) * _CONDITION_MULTIPLE_OF
|
||||
cond_img = self.image_processor.preprocess(batch.pil_image, cond_h, cond_w).to(device=device,
|
||||
dtype=torch.float32)
|
||||
latent = self.vae.encode(cond_img).latent_dist.mode()
|
||||
# NOTE: at runtime self.vae.config is a diffusers FrozenDict with flat
|
||||
# latents_mean/latents_std fields. Access the flat fields directly.
|
||||
cfg = self.vae.config
|
||||
mean = torch.tensor(cfg.latents_mean, device=device, dtype=torch.float32).view(1, -1, 1, 1)
|
||||
std = torch.tensor(cfg.latents_std, device=device, dtype=torch.float32).view(1, -1, 1, 1)
|
||||
latent = ((latent - mean) / std).to(dtype)
|
||||
|
||||
kv_caches = GlmImageKVCache(num_layers=self.transformer.num_layers)
|
||||
empty_text = batch.prompt_embeds[0][:1, :0, :].to(device=device, dtype=dtype)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None, forward_batch=batch):
|
||||
self.transformer(
|
||||
hidden_states=latent,
|
||||
encoder_hidden_states=empty_text,
|
||||
prior_token_id=prior_ids,
|
||||
prior_token_drop=torch.zeros((prior_ids.shape[0], ), dtype=torch.bool, device=device),
|
||||
timestep=torch.zeros((1, ), device=device),
|
||||
target_size=torch.tensor([tuple(cond_img.shape[-2:])], device=device, dtype=torch.long),
|
||||
crop_coords=torch.zeros((1, 2), device=device, dtype=torch.long),
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="write",
|
||||
)
|
||||
batch.extra["glm_kv_caches"] = kv_caches
|
||||
return batch
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
class GlmImageDecodingStage(DecodingStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor, fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
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)
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
if not vae_autocast:
|
||||
latents = latents.to(vae_dtype)
|
||||
decoded = self.vae.decode(latents)
|
||||
|
||||
image = decoded.sample if hasattr(decoded, "sample") else decoded
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image.unsqueeze(2)
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CFG convention (both denoise paths): row 0 conditional (positive), row 1 unconditional."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadata
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class GlmImageDenoisingStage(DenoisingStage):
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.min_dims(1)])
|
||||
latents = getattr(batch, "latent", getattr(batch, "latents", None))
|
||||
result.add_check("latents", latents, [V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps, V.positive_int)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_cfg = guidance_scale > 1.0
|
||||
|
||||
latents = getattr(batch, "latent", getattr(batch, "latents", None))
|
||||
if latents is None:
|
||||
raise ValueError("No latents found in batch.")
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
text_attention_mask = getattr(batch, "attention_mask", None)
|
||||
timesteps = batch.timesteps
|
||||
|
||||
patch_size = self.transformer.patch_size
|
||||
_, _, h, w = latents.shape
|
||||
image_seq_length = (h // patch_size) * (w // patch_size)
|
||||
text_seq_length = prompt_embeds.shape[1] if prompt_embeds.dim() >= 2 else 0
|
||||
|
||||
first_block = self.transformer.transformer_blocks[0]
|
||||
backend = getattr(first_block.attn1.attn, "backend", None)
|
||||
sdpa = backend == AttentionBackendEnum.TORCH_SDPA and text_attention_mask is not None
|
||||
|
||||
kv_caches = batch.extra.get("glm_kv_caches")
|
||||
if kv_caches is None:
|
||||
self._denoise_t2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, device, dtype)
|
||||
else:
|
||||
self._denoise_i2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype)
|
||||
return batch
|
||||
|
||||
def _denoise_t2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, device, dtype) -> None:
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
bs = 2 if do_cfg else 1
|
||||
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long).repeat(bs, 1)
|
||||
crop_coords = torch.zeros((bs, 2), device=device, dtype=torch.long)
|
||||
|
||||
prior_token_id = batch.prior_token_id
|
||||
if do_cfg and prior_token_id.shape[0] == 1:
|
||||
prior_token_id = prior_token_id.repeat(2, 1)
|
||||
if do_cfg:
|
||||
prior_token_drop = torch.tensor([False, True], device=device)
|
||||
else:
|
||||
prior_token_drop = getattr(batch, "prior_token_drop", torch.tensor([False], device=device))
|
||||
|
||||
attention_mask_kv = None
|
||||
if sdpa:
|
||||
if (text_attention_mask.shape[0] == 1 and bs > 1):
|
||||
text_attention_mask = text_attention_mask.repeat(bs, 1)
|
||||
mix_attn_mask = torch.ones((bs, text_seq_length + image_seq_length), device=device, dtype=torch.float32)
|
||||
mix_attn_mask[:, :text_seq_length] = (text_attention_mask.float().to(device))
|
||||
attention_mask_kv = (mix_attn_mask > 0).unsqueeze(1).unsqueeze(2)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = torch.cat([latents] * 2) if do_cfg else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t).to(dtype)
|
||||
t_expand = t.expand(latent_model_input.shape[0]) - 1
|
||||
|
||||
attn_metadata = (SDPAMetadata(current_timestep=i, attn_mask=attention_mask_kv)
|
||||
if attention_mask_kv is not None else None)
|
||||
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
prior_token_id,
|
||||
prior_token_drop,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
)
|
||||
|
||||
if do_cfg:
|
||||
noise_pred_cond, noise_pred_uncond = noise_pred.float().chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
guidance_rescale = getattr(batch, "guidance_rescale", 0.0)
|
||||
if guidance_rescale > 0.0:
|
||||
dims = list(range(1, noise_pred_cond.ndim))
|
||||
std_text = noise_pred_cond.std(dim=dims, keepdim=True)
|
||||
std_cfg = noise_pred.std(dim=dims, keepdim=True)
|
||||
rescaled = noise_pred * (std_text / std_cfg)
|
||||
noise_pred = (guidance_rescale * rescaled + (1 - guidance_rescale) * noise_pred)
|
||||
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
progress_bar.update()
|
||||
|
||||
batch.latents = latents.unsqueeze(2)
|
||||
|
||||
def _denoise_i2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype) -> None:
|
||||
"""Two separate transformer calls (cond reads the cache, uncond skips it):
|
||||
the cache mode is one global flag with batch-1 k/v, so a 2-row CFG call
|
||||
cannot express both."""
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long)
|
||||
crop_coords = torch.zeros((1, 2), device=device, dtype=torch.long)
|
||||
prior_token_id = batch.prior_token_id[:1]
|
||||
drop_keep = torch.zeros((1, ), dtype=torch.bool, device=device)
|
||||
drop_all = torch.ones((1, ), dtype=torch.bool, device=device)
|
||||
cache_len = kv_caches[0].k_cache.shape[1] if kv_caches[0].k_cache is not None else 0
|
||||
|
||||
def _mask(row: int, with_cache: bool):
|
||||
if not sdpa:
|
||||
return None
|
||||
prefix = cache_len if with_cache else 0
|
||||
m = torch.ones((1, prefix + text_seq_length + image_seq_length), device=device, dtype=torch.float32)
|
||||
m[:, prefix:prefix + text_seq_length] = text_attention_mask[row:row + 1].float().to(device)
|
||||
return (m > 0).unsqueeze(1).unsqueeze(2)
|
||||
|
||||
cond_mask = _mask(0, with_cache=True)
|
||||
uncond_mask = _mask(1, with_cache=False) if do_cfg else None
|
||||
pos_embeds = prompt_embeds[:1]
|
||||
neg_embeds = prompt_embeds[1:2] if do_cfg else None
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = self.scheduler.scale_model_input(latents, t).to(dtype)
|
||||
t_expand = t.expand(1) - 1
|
||||
|
||||
cond_meta = (SDPAMetadata(current_timestep=i, attn_mask=cond_mask) if cond_mask is not None else None)
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=cond_meta,
|
||||
forward_batch=batch):
|
||||
noise_pred = self.transformer(latent_model_input,
|
||||
pos_embeds,
|
||||
prior_token_id,
|
||||
drop_keep,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="read")
|
||||
|
||||
if do_cfg:
|
||||
uncond_meta = (SDPAMetadata(current_timestep=i, attn_mask=uncond_mask)
|
||||
if uncond_mask is not None else None)
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=uncond_meta,
|
||||
forward_batch=batch):
|
||||
noise_pred_uncond = self.transformer(latent_model_input,
|
||||
neg_embeds,
|
||||
prior_token_id,
|
||||
drop_all,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="skip")
|
||||
noise_pred = noise_pred_uncond.float() + guidance_scale * (noise_pred.float() -
|
||||
noise_pred_uncond.float())
|
||||
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
progress_bar.update()
|
||||
|
||||
kv_caches.clear()
|
||||
batch.latents = latents.unsqueeze(2)
|
||||
@@ -109,6 +109,10 @@ class ForwardBatch:
|
||||
max_sequence_length: int | None = None
|
||||
prompt_template: dict[str, Any] | None = None
|
||||
do_classifier_free_guidance: bool = False
|
||||
# When True, ``guidance_scale`` is passed into models that use embedded guidance (e.g. FLUX)
|
||||
# and must not imply classic dual-forward CFG. Use ``true_cfg_scale > 1`` for true CFG.
|
||||
use_embedded_guidance: bool = False
|
||||
true_cfg_scale: float = 1.0
|
||||
|
||||
# Batch info
|
||||
batch_size: int | None = None
|
||||
@@ -252,9 +256,12 @@ class ForwardBatch:
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
|
||||
# Enable CFG for standard guidance_scale and LTX-2 text CFG scales.
|
||||
# LTX-2 text CFG scales; FLUX uses ``use_embedded_guidance`` so ``guidance_scale > 1`` alone
|
||||
# does not enable classifier-free guidance.
|
||||
ltx2_text_cfg_enabled = (self.ltx2_cfg_scale_video != 1.0 or self.ltx2_cfg_scale_audio != 1.0)
|
||||
if self.guidance_scale > 1.0 or ltx2_text_cfg_enabled:
|
||||
if self.use_embedded_guidance:
|
||||
self.do_classifier_free_guidance = (self.true_cfg_scale > 1.0) or ltx2_text_cfg_enabled
|
||||
elif self.guidance_scale > 1.0 or ltx2_text_cfg_enabled:
|
||||
self.do_classifier_free_guidance = True
|
||||
if self.negative_prompt_embeds is None:
|
||||
self.negative_prompt_embeds = []
|
||||
|
||||
@@ -0,0 +1,331 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 video denoising stage.
|
||||
|
||||
The Cosmos3 video path is monolithic by design: each CFG pass repacks the whole
|
||||
text+vision sequence (the conditional pass carries prompt tokens, the
|
||||
unconditional pass carries negative-prompt tokens), so the standard
|
||||
encode/condition/denoise/decode stage split does not apply. This single stage
|
||||
owns the full flow, delegating the framework-parity-tested math to
|
||||
``fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline``:
|
||||
|
||||
1. resolve mode (T2I / I2V / T2V) + per-mode defaults, set ``flow_shift``;
|
||||
2. tokenize the prompt + negative prompt with the Qwen2 chat template;
|
||||
3. VAE-encode the conditioning frame(s) for I2V / T2I (kept clean), build the
|
||||
initial noise (clean condition frames + pure noise elsewhere);
|
||||
4. run the UniPC denoise loop with sequential CFG
|
||||
(``Cosmos3DenoiseEngine.denoise``);
|
||||
5. VAE-decode + ``(1 + x) / 2`` clamp to [0, 1].
|
||||
|
||||
This mirrors the framework ``Cosmos3OmniDiffusersPipeline.__call__``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import weakref
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3DenoiseEngine,
|
||||
Cosmos3SoundSpec,
|
||||
Cosmos3VisionSpec,
|
||||
_VaeNorm,
|
||||
cosmos3_special_tokens,
|
||||
cosmos3_tokenize_caption,
|
||||
cosmos3_vae_decode,
|
||||
cosmos3_vae_encode,
|
||||
)
|
||||
from fastvideo.pipelines.basic.cosmos3.presets import (
|
||||
COSMOS3_VIDEO_NEGATIVE_PROMPT, )
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Cosmos3DenoisingStage(PipelineStage):
|
||||
"""Full Cosmos3 video denoise: tokenize + encode + denoise + decode."""
|
||||
|
||||
def __init__(self, *, transformer, scheduler, vae, tokenizer, pipeline=None) -> None:
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.tokenizer = tokenizer
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline is not None else None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Geometry helpers
|
||||
# ------------------------------------------------------------------
|
||||
@staticmethod
|
||||
def _latent_frames(num_frames: int, temporal_factor: int) -> int:
|
||||
return (int(num_frames) - 1) // int(temporal_factor) + 1
|
||||
|
||||
@staticmethod
|
||||
def _flow_shift_for_resolution(height: int, width: int) -> float:
|
||||
"""UniPC ``flow_shift`` for a given pixel resolution.
|
||||
|
||||
Mirrors the framework's ``_RESOLUTION_SHIFT_DEFAULTS`` (8B VLM backbone,
|
||||
which Cosmos3-Nano uses): the shift is keyed by the named resolution
|
||||
bucket the (H, W) belongs to, regardless of task (T2V/I2V/T2I):
|
||||
|
||||
"256" -> 3.0, "480" -> 5.0, "704"/"720"/"768" -> 10.0
|
||||
|
||||
We invert the framework's ``{IMAGE,VIDEO}_RES_SIZE_INFO`` tables by the
|
||||
longest side: <=320 is the 256 bucket, 640-832 the 480 bucket, and
|
||||
960-1360 the 704/720/768 buckets.
|
||||
"""
|
||||
long_side = max(int(height), int(width))
|
||||
if long_side <= 480:
|
||||
return 3.0
|
||||
if long_side <= 896:
|
||||
return 5.0
|
||||
return 10.0
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
arch = pipeline_config.dit_config.arch_config
|
||||
device = self.transformer.embed_tokens.weight.device
|
||||
dtype = self.transformer.embed_tokens.weight.dtype
|
||||
|
||||
num_frames = int(batch.num_frames) if batch.num_frames is not None else 1
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
fps = float(batch.fps) if batch.fps is not None else float(arch.base_fps)
|
||||
guidance = float(batch.guidance_scale)
|
||||
|
||||
is_t2i = num_frames == 1 and batch.preprocessed_image is None and batch.pil_image is None
|
||||
is_i2v = (batch.preprocessed_image is not None or batch.pil_image is not None) and not is_t2i
|
||||
|
||||
# Resolution-based flow_shift, set on the owning pipeline (rebuilds the
|
||||
# scheduler). The framework picks the UniPC shift purely from the named
|
||||
# resolution bucket (``_RESOLUTION_SHIFT_DEFAULTS``), NOT from the task,
|
||||
# so T2V/I2V/T2I at the same resolution share a shift.
|
||||
pipe = self.pipeline() if self.pipeline is not None else None
|
||||
flow_shift = self._flow_shift_for_resolution(height, width)
|
||||
if pipe is not None and hasattr(pipe, "_set_flow_shift"):
|
||||
pipe._set_flow_shift(flow_shift)
|
||||
scheduler = pipe.scheduler
|
||||
else:
|
||||
scheduler = self.scheduler
|
||||
|
||||
# ---- Tokenize prompt + negative prompt ----
|
||||
prompt = batch.prompt if isinstance(batch.prompt, str) else (batch.prompt[0] if batch.prompt else "")
|
||||
negative_prompt = batch.negative_prompt
|
||||
if negative_prompt is None:
|
||||
negative_prompt = "" if is_t2i else COSMOS3_VIDEO_NEGATIVE_PROMPT
|
||||
if isinstance(negative_prompt, list):
|
||||
negative_prompt = negative_prompt[0] if negative_prompt else ""
|
||||
|
||||
special_tokens = cosmos3_special_tokens(self.tokenizer)
|
||||
is_video = not is_t2i
|
||||
cond_ids = cosmos3_tokenize_caption(self.tokenizer, prompt, is_video=is_video, use_system_prompt=False)
|
||||
uncond_ids = cosmos3_tokenize_caption(self.tokenizer,
|
||||
negative_prompt,
|
||||
is_video=is_video,
|
||||
use_system_prompt=False)
|
||||
|
||||
# ---- VAE normalization constants + geometry ----
|
||||
norm = _VaeNorm.from_vae(self.vae, dtype)
|
||||
temporal_factor = int(arch.temporal_compression_factor)
|
||||
spatial_factor = int(self.vae.config.scale_factor_spatial)
|
||||
latent_t = self._latent_frames(num_frames, temporal_factor)
|
||||
latent_h = height // spatial_factor
|
||||
latent_w = width // spatial_factor
|
||||
latent_channel = int(arch.latent_channel)
|
||||
latent_shape = (latent_channel, latent_t, latent_h, latent_w)
|
||||
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list):
|
||||
generator = generator[0] if generator else None
|
||||
|
||||
# ---- Conditioning latent (I2V / T2I) + condition mask ----
|
||||
condition_frame_indexes: list[int] = []
|
||||
clean_latent: torch.Tensor | None = None
|
||||
if is_i2v or (is_t2i and (batch.preprocessed_image is not None or batch.pil_image is not None)):
|
||||
image = batch.preprocessed_image if batch.preprocessed_image is not None else batch.pil_image
|
||||
cond_pixels = self._image_to_video_tensor(image, num_frames, height, width, device, dtype)
|
||||
clean_latent = cosmos3_vae_encode(self.vae, cond_pixels, norm).squeeze(0).float() # [C, T, H, W]
|
||||
condition_frame_indexes = [0]
|
||||
|
||||
# ---- Initial noise (clean condition frames + pure noise elsewhere) ----
|
||||
pure_noise = randn_tensor(latent_shape, generator=generator, device=device, dtype=dtype).float()
|
||||
if clean_latent is not None:
|
||||
cond_mask = torch.zeros((latent_t, 1, 1), device=device, dtype=pure_noise.dtype)
|
||||
for idx in condition_frame_indexes:
|
||||
if 0 <= idx < latent_t:
|
||||
cond_mask[idx, 0, 0] = 1.0
|
||||
clean = clean_latent.to(device=device, dtype=pure_noise.dtype)
|
||||
init_latent = cond_mask * clean + (1.0 - cond_mask) * pure_noise
|
||||
else:
|
||||
init_latent = pure_noise
|
||||
|
||||
spec = Cosmos3VisionSpec(
|
||||
shape=latent_shape,
|
||||
condition_frame_indexes=condition_frame_indexes,
|
||||
)
|
||||
|
||||
# ---- Scheduler timesteps ----
|
||||
scheduler.set_timesteps(int(batch.num_inference_steps), device=device)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
engine = Cosmos3DenoiseEngine(
|
||||
transformer=self.transformer,
|
||||
scheduler=scheduler,
|
||||
special_tokens=special_tokens,
|
||||
latent_patch_size=int(arch.latent_patch_size),
|
||||
temporal_modality_margin=int(arch.temporal_modality_margin),
|
||||
reset_spatial_ids=bool(arch.unified_3d_mrope_reset_spatial_ids),
|
||||
enable_fps_modulation=bool(arch.enable_fps_modulation),
|
||||
base_fps=float(arch.base_fps),
|
||||
temporal_compression_factor=temporal_factor,
|
||||
include_end_of_generation_token=False,
|
||||
)
|
||||
|
||||
flat_latent = init_latent.reshape(-1)
|
||||
fps_per_item = [fps] if bool(arch.enable_fps_modulation) else None
|
||||
|
||||
# ---- t2vs: jointly generate sound (combined [vision | sound] latent) ----
|
||||
# Mirrors the framework: a placeholder audio sized to the video duration
|
||||
# sets the sound latent length; sound shares the denoise/CFG with vision.
|
||||
with_audio = is_video and os.environ.get("COSMOS3_T2VS", "") not in ("", "0")
|
||||
sound_specs = None
|
||||
sound_fps_per_item = None
|
||||
sound_vae = None
|
||||
sound_shape: tuple[int, int] | None = None
|
||||
if with_audio:
|
||||
sound_vae = self._get_sound_vae(pipe, device, dtype)
|
||||
sound_dim = int(arch.sound_dim)
|
||||
sound_latent_fps = float(arch.sound_latent_fps)
|
||||
# Framework ``create_placeholder_audio`` + ``get_latent_num_samples``.
|
||||
num_audio_samples = int(num_frames / fps * sound_vae.sample_rate)
|
||||
sound_latent_t = max(1, sound_vae.get_latent_num_samples(num_audio_samples))
|
||||
sound_shape = (sound_dim, sound_latent_t)
|
||||
sound_noise = randn_tensor((sound_dim, sound_latent_t), generator=generator, device=device,
|
||||
dtype=dtype).float()
|
||||
flat_latent = torch.cat([flat_latent, sound_noise.reshape(-1)])
|
||||
sound_specs = [Cosmos3SoundSpec(shape=sound_shape, condition_frame_indexes=[], fps=sound_latent_fps)]
|
||||
sound_fps_per_item = [sound_latent_fps] if bool(arch.enable_fps_modulation) else None
|
||||
|
||||
final_flat = engine.denoise(
|
||||
flat_latent=flat_latent,
|
||||
timesteps=timesteps,
|
||||
guidance=guidance,
|
||||
specs=[spec],
|
||||
cond_token_ids=cond_ids,
|
||||
uncond_token_ids=uncond_ids,
|
||||
fps_per_item=fps_per_item,
|
||||
progress_bar=lambda it: tqdm(it, desc="Cosmos3 denoising"),
|
||||
sound_specs=sound_specs,
|
||||
sound_fps_per_item=sound_fps_per_item,
|
||||
)
|
||||
|
||||
# ---- Decode vision: [C, T, H, W] -> pixels [B, 3, T, H, W] in [0, 1] ----
|
||||
vision_flat = final_flat[:spec.numel]
|
||||
result_latent = vision_flat.reshape(latent_shape).unsqueeze(0).to(device=device, dtype=dtype)
|
||||
decoded = cosmos3_vae_decode(self.vae, result_latent, norm) # [B, 3, T, H, W] in [-1, 1]
|
||||
video = ((1.0 + decoded) / 2.0).clamp(0.0, 1.0)
|
||||
|
||||
batch.latents = result_latent
|
||||
batch.output = video
|
||||
|
||||
# ---- Decode sound: AVAE latent [C, T] -> waveform [C, N] in [-1, 1] ----
|
||||
if with_audio and sound_vae is not None and sound_shape is not None:
|
||||
sound_latent = final_flat[spec.numel:].reshape(sound_shape).unsqueeze(0).to(device=device, dtype=dtype)
|
||||
waveform = sound_vae.decode(sound_latent) # [1, C_audio, N]
|
||||
batch.extra["audio"] = waveform[0].detach().float().cpu() # [C_audio, N]
|
||||
batch.extra["audio_sample_rate"] = int(sound_vae.sample_rate)
|
||||
return batch
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Image preprocessing
|
||||
# ------------------------------------------------------------------
|
||||
@staticmethod
|
||||
def _resize_and_center_crop(img: torch.Tensor, target_h: int, target_w: int) -> torch.Tensor:
|
||||
"""Aspect-ratio-preserving resize + center crop, matching the framework
|
||||
(``cosmos_framework.inference.vision._resize_and_center_crop``)."""
|
||||
import math
|
||||
|
||||
import torchvision.transforms.functional as TF
|
||||
orig_h, orig_w = img.shape[-2], img.shape[-1]
|
||||
scaling_ratio = max(target_w / orig_w, target_h / orig_h)
|
||||
resize_h = int(math.ceil(scaling_ratio * orig_h))
|
||||
resize_w = int(math.ceil(scaling_ratio * orig_w))
|
||||
img = TF.resize(img, [resize_h, resize_w])
|
||||
return TF.center_crop(img, [target_h, target_w])
|
||||
|
||||
@staticmethod
|
||||
def _get_sound_vae(pipe: Any, device: torch.device, dtype: torch.dtype) -> Any:
|
||||
"""Lazily load + cache the Cosmos3 sound AVAE decoder from the checkpoint.
|
||||
|
||||
The video path does not load ``sound_tokenizer``; t2vs needs only its
|
||||
decoder, so we load it on first use from ``<model_path>/sound_tokenizer``.
|
||||
"""
|
||||
cached = getattr(pipe, "_sound_vae", None) if pipe is not None else None
|
||||
if cached is not None:
|
||||
return cached
|
||||
from fastvideo.models.audio.cosmos3_avae import Cosmos3SoundVAE
|
||||
model_path = pipe.model_path
|
||||
sound_dir = os.path.join(model_path, "sound_tokenizer")
|
||||
sound_vae = Cosmos3SoundVAE.from_pretrained(sound_dir, torch_dtype=dtype).to(device)
|
||||
if pipe is not None:
|
||||
pipe._sound_vae = sound_vae
|
||||
return sound_vae
|
||||
|
||||
@classmethod
|
||||
def _image_to_video_tensor(
|
||||
cls,
|
||||
image: Any,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Build the I2V conditioning pixel video ``[1, 3, T, H, W]`` in [-1, 1].
|
||||
|
||||
Faithful to the framework (``cosmos_framework.inference.vision``):
|
||||
``load_conditioning_image`` (aspect-preserving resize + center crop +
|
||||
uint8 quantization, then ``/127.5 - 1``) followed by
|
||||
``build_conditioned_video_batch``, which fills frame 0 with the image and
|
||||
**repeats the last conditioning frame** for the rest of the clip (a static
|
||||
video), NOT zeros. The whole clip is VAE-encoded by the caller; only the
|
||||
latent condition frame(s) are kept clean by the condition mask, but the
|
||||
VAE is temporal, so the repeated (not zeroed) frames change the condition
|
||||
latent — zero-filling here produces a wrong conditioning latent.
|
||||
"""
|
||||
import numpy as np
|
||||
|
||||
if hasattr(image, "convert"): # PIL.Image: framework-exact preprocessing.
|
||||
arr = np.array(image.convert("RGB"))
|
||||
img = torch.from_numpy(arr).permute(2, 0, 1).float() # [3, H, W] in [0, 255]
|
||||
# Resize + center crop + uint8 quantization, then -> [-1, 1]
|
||||
# (load_conditioning_image / load_conditioning_image_pixels).
|
||||
img = cls._resize_and_center_crop(img.unsqueeze(0), height, width).squeeze(0)
|
||||
img = img.round().clamp(0, 255) / 127.5 - 1.0 # [3, H, W] in [-1, 1]
|
||||
elif isinstance(image, torch.Tensor): # already-preprocessed conditioning frame.
|
||||
img = image.float()
|
||||
if img.dim() == 5: # [B,3,T,H,W]
|
||||
img = img[0]
|
||||
if img.dim() == 4: # [3,T,H,W] or [B,3,H,W] -> first frame
|
||||
img = img[:, 0]
|
||||
if img.max() > 1.5: # [0, 255] -> [-1, 1]; otherwise assume already [-1, 1].
|
||||
img = img / 127.5 - 1.0
|
||||
if img.shape[-2:] != (height, width):
|
||||
img = cls._resize_and_center_crop(img.unsqueeze(0), height, width).squeeze(0)
|
||||
else:
|
||||
raise TypeError(f"Unsupported conditioning image type: {type(image)}")
|
||||
|
||||
# Static-repeat video (build_conditioned_video_batch: frame 0 = image,
|
||||
# remaining frames repeat the last conditioning frame). The whole clip is
|
||||
# VAE-encoded by the caller; only the latent condition frame(s) are kept
|
||||
# clean by the condition mask, but the VAE is temporal, so the repeated
|
||||
# (not zeroed) frames change the condition latent — zero-filling here
|
||||
# produces a wrong conditioning latent.
|
||||
img = img.to(device=device, dtype=dtype)
|
||||
video = img.unsqueeze(0).unsqueeze(2).expand(1, 3, num_frames, height, width)
|
||||
return video.contiguous()
|
||||
@@ -0,0 +1,424 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
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.pipelines.stages.timestep_preparation import TimestepPreparationStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _pack_latents(
|
||||
latents: torch.Tensor,
|
||||
batch_size: int,
|
||||
num_channels_latents: int,
|
||||
height: int,
|
||||
width: int,
|
||||
) -> torch.Tensor:
|
||||
"""Diffusers ``_pack_latents`` for FLUX (2×2 spatial pack in latent space)."""
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
return latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
|
||||
|
||||
def _unpack_latents(
|
||||
latents: torch.Tensor,
|
||||
batch_size: int,
|
||||
num_channels_latents: int,
|
||||
height: int,
|
||||
width: int,
|
||||
) -> torch.Tensor:
|
||||
"""Inverse of ``_pack_latents``."""
|
||||
latents = latents.reshape(batch_size, height // 2, width // 2, num_channels_latents, 2, 2)
|
||||
latents = latents.permute(0, 3, 1, 4, 2, 5)
|
||||
return latents.reshape(batch_size, num_channels_latents, height, width)
|
||||
|
||||
|
||||
def _prepare_latent_image_ids(
|
||||
patch_height: int,
|
||||
patch_width: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype = torch.long,
|
||||
) -> torch.Tensor:
|
||||
"""Match Diffusers ``FluxPipeline._prepare_latent_image_ids`` (no batch dim)."""
|
||||
latent_image_ids = torch.zeros(patch_height, patch_width, 3, device=device, dtype=torch.float32)
|
||||
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(patch_height, device=device)[:, None]
|
||||
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(patch_width, device=device)[None, :]
|
||||
h, w, c = latent_image_ids.shape
|
||||
return latent_image_ids.reshape(h * w, c).to(dtype=dtype)
|
||||
|
||||
|
||||
class FluxInputValidationStage(InputValidationStage):
|
||||
"""Require height/width divisible by 16 (VAE scale × 2 for FLUX packing)."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if (batch.height is not None and batch.width is not None and (batch.height % 16 != 0 or batch.width % 16 != 0)):
|
||||
raise ValueError("FLUX expects height and width divisible by 16 "
|
||||
f"(VAE latent grid × 2× packing); got {batch.height}×{batch.width}.")
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class FluxConditioningStage(PipelineStage):
|
||||
"""Build CLIP pooled + T5 sequence + ``text_ids`` (and optional negative for true CFG)."""
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if len(batch.prompt_embeds) < 2:
|
||||
raise ValueError("FluxConditioningStage expects 2 prompt_embeds (CLIP pooled, T5 sequence), "
|
||||
f"got {len(batch.prompt_embeds)}")
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
|
||||
pooled = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
enc = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
seq_len = enc.shape[1]
|
||||
text_ids = torch.zeros(seq_len, 3, device=device, dtype=torch.long)
|
||||
|
||||
batch.extra["flux_pooled_projections"] = pooled
|
||||
batch.extra["flux_encoder_hidden_states"] = enc
|
||||
batch.extra["flux_text_ids"] = text_ids
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
if not batch.negative_prompt_embeds or len(batch.negative_prompt_embeds) < 2:
|
||||
raise ValueError("True CFG requires two negative_prompt_embeds (CLIP, T5).")
|
||||
neg_pooled = batch.negative_prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
neg_enc = batch.negative_prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
batch.extra["flux_negative_pooled_projections"] = neg_pooled
|
||||
batch.extra["flux_negative_encoder_hidden_states"] = neg_enc
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class FluxTimestepPreparationStage(TimestepPreparationStage):
|
||||
"""Flow Match with resolution-dependent ``mu`` from packed image sequence length."""
|
||||
|
||||
@staticmethod
|
||||
def _calculate_mu(
|
||||
image_seq_len: int,
|
||||
base_seq_len: int = 256,
|
||||
max_seq_len: int = 4096,
|
||||
base_shift: float = 0.5,
|
||||
max_shift: float = 1.15,
|
||||
) -> float:
|
||||
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
b = base_shift - m * base_seq_len
|
||||
return float(image_seq_len) * m + b
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
sig = inspect.signature(self.scheduler.set_timesteps)
|
||||
if "mu" not in sig.parameters:
|
||||
logger.warning(
|
||||
"FLUX timestep prep: scheduler %s.set_timesteps does not accept 'mu'; falling back to the base "
|
||||
"timestep schedule. FLUX expects a FlowMatchEulerDiscreteScheduler with resolution-dependent "
|
||||
"dynamic shifting — output quality may degrade.",
|
||||
type(self.scheduler).__name__)
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
cfg = getattr(self.scheduler, "config", None)
|
||||
use_dynamic = bool(getattr(cfg, "use_dynamic_shifting", False)) if cfg is not None else False
|
||||
if not use_dynamic:
|
||||
logger.warning(
|
||||
"FLUX timestep prep: scheduler has use_dynamic_shifting=False; falling back to the base timestep "
|
||||
"schedule and skipping the resolution-dependent 'mu' shift. FLUX requires dynamic shifting for "
|
||||
"correct timesteps — output quality may degrade.")
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("height/width must be set before FluxTimestepPreparationStage")
|
||||
|
||||
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
|
||||
spatial_ratio = int(getattr(vae_arch, "spatial_compression_ratio", 8))
|
||||
h_lat = batch.height // spatial_ratio
|
||||
w_lat = batch.width // spatial_ratio
|
||||
if h_lat % 2 != 0 or w_lat % 2 != 0:
|
||||
raise ValueError(
|
||||
f"Latent spatial dims must be even for FLUX packing; got {h_lat}×{w_lat} from {batch.height}×{batch.width}."
|
||||
)
|
||||
image_seq_len = (h_lat // 2) * (w_lat // 2)
|
||||
|
||||
base_seq_len = int(getattr(cfg, "base_image_seq_len", 256))
|
||||
max_seq_len = int(getattr(cfg, "max_image_seq_len", 4096))
|
||||
base_shift = float(getattr(cfg, "base_shift", 0.5))
|
||||
max_shift = float(getattr(cfg, "max_shift", 1.15))
|
||||
|
||||
device = get_local_torch_device()
|
||||
mu = self._calculate_mu(
|
||||
image_seq_len=image_seq_len,
|
||||
base_seq_len=base_seq_len,
|
||||
max_seq_len=max_seq_len,
|
||||
base_shift=base_shift,
|
||||
max_shift=max_shift,
|
||||
)
|
||||
self.scheduler.set_timesteps(batch.num_inference_steps, device=device, mu=mu)
|
||||
batch.timesteps = self.scheduler.timesteps
|
||||
return batch
|
||||
|
||||
|
||||
class FluxLatentPreparationStage(PipelineStage):
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("height/width required for FluxLatentPreparationStage")
|
||||
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
if not batch.prompt_embeds:
|
||||
raise ValueError("prompt or prompt_embeds must be provided")
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
if isinstance(batch.generator, list) and len(batch.generator) != batch_size:
|
||||
raise ValueError(f"generator list length {len(batch.generator)} does not match batch_size {batch_size}")
|
||||
|
||||
arch = fastvideo_args.pipeline_config.dit_config.arch_config
|
||||
in_channels = int(getattr(arch, "in_channels", 64))
|
||||
num_channels_latents = in_channels // 4
|
||||
|
||||
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
|
||||
spatial_ratio = int(getattr(vae_arch, "spatial_compression_ratio", 8))
|
||||
|
||||
h_lat = batch.height // spatial_ratio
|
||||
w_lat = batch.width // spatial_ratio
|
||||
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
device = get_local_torch_device()
|
||||
|
||||
shape = (batch_size, num_channels_latents, h_lat, w_lat)
|
||||
latents = batch.latents
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=batch.generator, device=device, dtype=dtype)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
latents = latents.to(device=device, dtype=dtype)
|
||||
if latents.shape != shape:
|
||||
raise ValueError(f"Expected latents shape {shape}, got {tuple(latents.shape)}")
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
packed = _pack_latents(latents, batch_size, num_channels_latents, h_lat, w_lat)
|
||||
|
||||
patch_h, patch_w = h_lat // 2, w_lat // 2
|
||||
img_ids = _prepare_latent_image_ids(patch_h, patch_w, device, dtype=torch.long)
|
||||
|
||||
batch.latents = packed
|
||||
batch.raw_latent_shape = shape
|
||||
batch.extra["flux_h_lat"] = h_lat
|
||||
batch.extra["flux_w_lat"] = w_lat
|
||||
batch.extra["flux_num_channels_latents"] = num_channels_latents
|
||||
batch.extra["flux_latent_image_ids"] = img_ids
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class FluxDenoisingStage(PipelineStage):
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _step_kwargs(scheduler_step, batch: ForwardBatch) -> dict[str, Any]:
|
||||
kwargs: dict[str, Any] = {}
|
||||
sig = inspect.signature(scheduler_step)
|
||||
if "generator" in sig.parameters:
|
||||
gen = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
|
||||
kwargs["generator"] = gen
|
||||
return kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("timesteps must be set before FluxDenoisingStage")
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be set before FluxDenoisingStage")
|
||||
|
||||
packed = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
|
||||
pooled = batch.extra["flux_pooled_projections"]
|
||||
enc = batch.extra["flux_encoder_hidden_states"]
|
||||
txt_ids = batch.extra["flux_text_ids"]
|
||||
img_ids = batch.extra["flux_latent_image_ids"]
|
||||
|
||||
neg_pooled = batch.extra.get("flux_negative_pooled_projections")
|
||||
neg_enc = batch.extra.get("flux_negative_encoder_hidden_states")
|
||||
|
||||
true_cfg_scale = float(batch.true_cfg_scale)
|
||||
use_true_cfg = batch.do_classifier_free_guidance and true_cfg_scale > 1.0
|
||||
|
||||
# Prefer the loaded transformer's arch (HF ``guidance_embeds``), not static pipeline defaults.
|
||||
tr_arch = self.transformer.fastvideo_config.arch_config
|
||||
guidance_embeds = bool(getattr(tr_arch, "guidance_embeds", False))
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
bs = packed.shape[0]
|
||||
if guidance_embeds:
|
||||
guidance = torch.full((bs, ), float(batch.guidance_scale), device=device, dtype=torch.float32)
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
step_extras = self._step_kwargs(self.scheduler.step, batch)
|
||||
|
||||
for t in timesteps:
|
||||
t_scalar = t
|
||||
if not isinstance(t_scalar, torch.Tensor):
|
||||
t_scalar = torch.tensor([t_scalar], device=device, dtype=torch.float32)
|
||||
t_scalar = t_scalar.to(device=device, dtype=torch.float32)
|
||||
|
||||
timestep_model = t_scalar.expand(bs).float() / 1000.0
|
||||
timestep_model = timestep_model.to(dtype=target_dtype)
|
||||
|
||||
ts_ctx = int(t_scalar.reshape(-1)[0].item())
|
||||
with (
|
||||
torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled and device.type == "cuda",
|
||||
),
|
||||
set_forward_context(
|
||||
current_timestep=ts_ctx,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
),
|
||||
):
|
||||
if use_true_cfg:
|
||||
assert neg_enc is not None and neg_pooled is not None
|
||||
n_neg = self.transformer(
|
||||
hidden_states=packed,
|
||||
encoder_hidden_states=neg_enc,
|
||||
pooled_projections=neg_pooled,
|
||||
timestep=timestep_model,
|
||||
guidance=guidance,
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
n_pos = self.transformer(
|
||||
hidden_states=packed,
|
||||
encoder_hidden_states=enc,
|
||||
pooled_projections=pooled,
|
||||
timestep=timestep_model,
|
||||
guidance=guidance,
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = n_neg + true_cfg_scale * (n_pos - n_neg)
|
||||
else:
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=packed,
|
||||
encoder_hidden_states=enc,
|
||||
pooled_projections=pooled,
|
||||
timestep=timestep_model,
|
||||
guidance=guidance,
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
packed = self.scheduler.step(
|
||||
noise_pred,
|
||||
t_scalar,
|
||||
packed,
|
||||
return_dict=False,
|
||||
**step_extras,
|
||||
)[0]
|
||||
|
||||
batch.latents = packed
|
||||
return batch
|
||||
|
||||
|
||||
class FluxDecodingStage(PipelineStage):
|
||||
"""Unpack latents, apply VAE scaling/shift, decode to pixels (5D output ``B×3×1×H×W``)."""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
@staticmethod
|
||||
def _denormalize_latents(latents: torch.Tensor, vae: Any) -> torch.Tensor:
|
||||
cfg = getattr(vae, "config", None)
|
||||
sf = getattr(cfg, "scaling_factor", None) if cfg is not None else None
|
||||
sh = getattr(cfg, "shift_factor", None) if cfg is not None else None
|
||||
if sf is None and hasattr(vae, "scaling_factor"):
|
||||
sf = vae.scaling_factor
|
||||
if sh is None and hasattr(vae, "shift_factor"):
|
||||
sh = vae.shift_factor
|
||||
if sf is not None:
|
||||
latents = latents / (sf.to(latents.device, latents.dtype) if isinstance(sf, torch.Tensor) else sf)
|
||||
if sh is not None:
|
||||
latents = latents + (sh.to(latents.device, latents.dtype) if isinstance(sh, torch.Tensor) else sh)
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
packed = batch.latents
|
||||
if packed is None:
|
||||
raise ValueError("latents must be set before FluxDecodingStage")
|
||||
|
||||
h_lat = int(batch.extra["flux_h_lat"])
|
||||
w_lat = int(batch.extra["flux_w_lat"])
|
||||
num_ch = int(batch.extra["flux_num_channels_latents"])
|
||||
raw_shape = batch.raw_latent_shape
|
||||
if raw_shape is None:
|
||||
raise ValueError("raw_latent_shape missing; FluxLatentPreparationStage must run first.")
|
||||
batch_size = int(raw_shape[0])
|
||||
|
||||
infer_device = get_local_torch_device()
|
||||
packed = packed.to(infer_device)
|
||||
|
||||
latents_4d = _unpack_latents(packed, batch_size, num_ch, h_lat, w_lat)
|
||||
latents_4d = self._denormalize_latents(latents_4d, self.vae)
|
||||
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
latents_4d = latents_4d.to(device=vae_device)
|
||||
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
use_cuda_autocast = autocast_enabled and vae_device.type == "cuda"
|
||||
|
||||
with torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=use_cuda_autocast,
|
||||
):
|
||||
if not autocast_enabled:
|
||||
latents_4d = latents_4d.to(dtype=vae_dtype)
|
||||
dec = self.vae.decode(latents_4d)
|
||||
image = dec.sample if hasattr(dec, "sample") else dec[0]
|
||||
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
batch.output = image.unsqueeze(2).detach().float().cpu()
|
||||
return batch
|
||||
@@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
|
||||
Cosmos25Config,
|
||||
Cosmos25_14BConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
@@ -58,11 +59,14 @@ from fastvideo.configs.pipelines.wan import (
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.glm_image import GlmImageConfig
|
||||
from fastvideo.configs.pipelines.flux import FluxPipelineConfig
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.configs.pipelines.stable_audio import (StableAudioOpenSmallConfig, StableAudioT2AConfig)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.api.matrixgame2 import MatrixGame2SamplingParam
|
||||
from fastvideo.api.matrixgame3 import MatrixGame3SamplingParam
|
||||
from fastvideo.api.flux import FluxSamplingParam
|
||||
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -769,6 +773,22 @@ def _register_configs() -> None:
|
||||
default_preset="gen3c_cosmos_7b",
|
||||
)
|
||||
|
||||
# Cosmos 3 (must register before Cosmos 2.5 and generic Cosmos detectors
|
||||
# so the cosmos3 path-detection takes precedence)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Cosmos3Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"nvidia/Cosmos3-Nano",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "cosmos3" in path.lower() or "cosmos-3" in path.lower(),
|
||||
],
|
||||
model_family="cosmos3",
|
||||
default_preset="cosmos3_nano",
|
||||
)
|
||||
|
||||
# Cosmos 2.5 (2B)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -1074,6 +1094,33 @@ def _register_configs() -> None:
|
||||
default_preset="sd35_medium",
|
||||
)
|
||||
|
||||
# GLM-Image
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=GlmImageConfig,
|
||||
hf_model_paths=[
|
||||
"zai-org/GLM-Image",
|
||||
],
|
||||
model_detectors=[lambda path: "glmimage" in path.lower() or "glm-image" in path.lower()],
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
model_family="glm_image",
|
||||
)
|
||||
|
||||
# FLUX.1-dev (Diffusers)
|
||||
register_configs(
|
||||
sampling_param_cls=FluxSamplingParam,
|
||||
pipeline_config_cls=FluxPipelineConfig,
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
hf_model_paths=[
|
||||
"black-forest-labs/FLUX.1-dev",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "fluxpipeline" in path,
|
||||
lambda path: "flux.1-dev" in path or "flux_1_dev" in path,
|
||||
lambda path: "/flux/" in path or path.endswith("/flux"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# --- Part 3: Main Resolver ---
|
||||
|
||||
@@ -1164,6 +1211,8 @@ def _register_presets() -> None:
|
||||
from fastvideo.api.presets import register_preset
|
||||
from fastvideo.pipelines.basic.cosmos.presets import (
|
||||
ALL_PRESETS as COSMOS_PRESETS, )
|
||||
from fastvideo.pipelines.basic.cosmos3.presets import (
|
||||
ALL_PRESETS as COSMOS3_PRESETS, )
|
||||
from fastvideo.pipelines.basic.dreamx_world.presets import (
|
||||
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gamecraft.presets import (
|
||||
@@ -1201,6 +1250,7 @@ def _register_presets() -> None:
|
||||
|
||||
all_preset_groups = (
|
||||
COSMOS_PRESETS,
|
||||
COSMOS3_PRESETS,
|
||||
DREAMX_WORLD_PRESETS,
|
||||
FLUX2_PRESETS,
|
||||
GAMECRAFT_PRESETS,
|
||||
|
||||
@@ -183,6 +183,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
"guidance_scale_2": None,
|
||||
"guidance_rescale": 0.0,
|
||||
"true_cfg_scale": None,
|
||||
"use_embedded_guidance": None,
|
||||
"boundary_ratio": None,
|
||||
"sigmas": None,
|
||||
},
|
||||
|
||||
@@ -31,23 +31,29 @@ def fa_default_impls():
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for the FA2/FA3 default custom-op tests")
|
||||
|
||||
# The registration + dispatcher live in `attention/utils/`, alongside the
|
||||
# FP4 cute template and the masked/varlen wrappers. The backend
|
||||
# (`attention/backends/flash_attn.py`) just imports `flash_attn_func_compilable`
|
||||
# from there, so test references the utils module directly.
|
||||
try:
|
||||
from fastvideo.attention.backends import flash_attn as fa_backend
|
||||
from fastvideo.attention.utils import flash_attn_default as fa_module
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"flash_attn backend not importable: {exc}")
|
||||
pytest.skip(f"flash_attn_default not importable: {exc}")
|
||||
|
||||
if fa_backend.fa_version not in ("2", "3"):
|
||||
if fa_module.fa_version not in ("2", "3"):
|
||||
pytest.skip(
|
||||
f"FA2/FA3 default custom op only exists for fa_version in (2, 3); "
|
||||
f"got {fa_backend.fa_version!r}"
|
||||
f"got {fa_module.fa_version!r}"
|
||||
)
|
||||
|
||||
# compilable dispatcher, the original FA wrapper it falls back to, and
|
||||
# the raw custom op for opcheck.
|
||||
# compilable dispatcher, the original FA wrapper it falls back to, the
|
||||
# raw custom op for opcheck, and the fa_version (FA2 has full register_
|
||||
# autograd; FA3 keeps the carve-out so some tests gate on this).
|
||||
return (
|
||||
fa_backend.flash_attn_func_compilable,
|
||||
fa_backend._fa_default,
|
||||
fa_module.flash_attn_func_compilable,
|
||||
fa_module._fa_default,
|
||||
torch.ops.fastvideo._flash_attn_default_forward,
|
||||
fa_module.fa_version,
|
||||
)
|
||||
|
||||
|
||||
@@ -71,7 +77,7 @@ def test_default_compilable_inference_matches_original(fa_default_impls, dtype,
|
||||
"""No-grad path routes through the custom op and is numerically identical."""
|
||||
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("bfloat16 is not supported on this GPU")
|
||||
compilable, original, _ = fa_default_impls
|
||||
compilable, original, _, _ = fa_default_impls
|
||||
|
||||
torch.manual_seed(0)
|
||||
q, k, v = _qkv(dtype, requires_grad=False)
|
||||
@@ -91,7 +97,7 @@ def test_default_compilable_training_backward_flows(fa_default_impls, dtype, cau
|
||||
"""
|
||||
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("bfloat16 is not supported on this GPU")
|
||||
compilable, original, _ = fa_default_impls
|
||||
compilable, original, _, _ = fa_default_impls
|
||||
|
||||
torch.manual_seed(0)
|
||||
q_ref, k_ref, v_ref = _qkv(dtype, requires_grad=True)
|
||||
@@ -115,7 +121,65 @@ def test_default_compilable_training_backward_flows(fa_default_impls, dtype, cau
|
||||
@pytest.mark.parametrize("causal", [False, True])
|
||||
def test_default_forward_opcheck(fa_default_impls, causal):
|
||||
"""Schema / fake-kernel consistency for the custom op (forward only)."""
|
||||
_, _, op = fa_default_impls
|
||||
_, _, op, _ = fa_default_impls
|
||||
torch.manual_seed(0)
|
||||
q, k, v = _qkv(torch.float16, requires_grad=False)
|
||||
torch.library.opcheck(op, (q, k, v, None, causal))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# FA2-only: backward is registered on the custom op itself. Exercises the #
|
||||
# `register_autograd` wiring directly (not via the dispatcher's carve-out). #
|
||||
# Skipped on FA3 until Kuan-Hao's Modal FA3 setup PR lands and we mirror the #
|
||||
# pattern there. #
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("causal", [False, True])
|
||||
def test_default_op_backward_through_registered_autograd(fa_default_impls, dtype, causal):
|
||||
"""FA2: gradients flow through ``torch.ops.fastvideo._flash_attn_default_forward``
|
||||
itself (no dispatcher carve-out involved) and match the original
|
||||
``flash_attn_func``'s gradients."""
|
||||
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("bfloat16 is not supported on this GPU")
|
||||
_, original, op, fa_version = fa_default_impls
|
||||
if fa_version != "2":
|
||||
pytest.skip(f"register_autograd is only wired for FA2 right now; got {fa_version!r}")
|
||||
|
||||
torch.manual_seed(0)
|
||||
q_ref, k_ref, v_ref = _qkv(dtype, requires_grad=True)
|
||||
q_test, k_test, v_test = _clone(q_ref, k_ref, v_ref)
|
||||
|
||||
# Reference grads via the original autograd.Function.
|
||||
out_ref = original(q_ref, k_ref, v_ref, softmax_scale=None, causal=causal)
|
||||
# Custom-op grads via the registered backward — unpack (out, lse), discard lse.
|
||||
out_test, _ = op(q_test, k_test, v_test, None, causal)
|
||||
|
||||
torch.testing.assert_close(out_test, out_ref,
|
||||
atol=0 if dtype == torch.float16 else 1e-3,
|
||||
rtol=0 if dtype == torch.float16 else 1e-3)
|
||||
|
||||
dout = torch.randn_like(out_ref)
|
||||
dq_ref, dk_ref, dv_ref = torch.autograd.grad(
|
||||
(out_ref * dout).sum(), (q_ref, k_ref, v_ref))
|
||||
dq_test, dk_test, dv_test = torch.autograd.grad(
|
||||
(out_test * dout).sum(), (q_test, k_test, v_test))
|
||||
|
||||
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
|
||||
torch.testing.assert_close(dq_test, dq_ref, atol=atol, rtol=rtol)
|
||||
torch.testing.assert_close(dk_test, dk_ref, atol=atol, rtol=rtol)
|
||||
torch.testing.assert_close(dv_test, dv_ref, atol=atol, rtol=rtol)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("causal", [False, True])
|
||||
def test_default_op_opcheck_with_grad_inputs(fa_default_impls, causal):
|
||||
"""FA2: full ``opcheck`` including ``test_autograd_registration`` —
|
||||
catches a missing/inconsistent backward at unit-test time, which was
|
||||
exactly the gap that #1373's first revision shipped."""
|
||||
_, _, op, fa_version = fa_default_impls
|
||||
if fa_version != "2":
|
||||
pytest.skip(f"autograd registration only wired for FA2; got {fa_version!r}")
|
||||
torch.manual_seed(0)
|
||||
q, k, v = _qkv(torch.float16, requires_grad=True)
|
||||
torch.library.opcheck(op, (q, k, v, None, causal))
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Regression guard for the masked/varlen custom ops.
|
||||
|
||||
`flash_attn_no_pad_compilable` / `flash_attn_varlen_qk_no_pad_compilable` wrap
|
||||
the whole masked-attention functions in `torch.library.custom_op`s so dynamo
|
||||
sees one traceable node (the internal unpad/pad bookkeeping runs eager inside).
|
||||
|
||||
On FA2 these ops register a real backward — `softmax_lse` is padded back to a
|
||||
statically-shaped `[batch, nheads, seqlen]` form on the way out and re-unpadded
|
||||
in backward, which calls FA2's `_flash_attn_varlen_backward`. So training also
|
||||
backprops through the op (no graph break on the training path either).
|
||||
|
||||
On FA3/FA4 these ops are forward+fake only and the `*_compilable` dispatchers
|
||||
carve out to the autograd.Function for grad-enabled calls (PR #1373 pattern).
|
||||
Tests gating on FA2 are skipped on FA3/FA4.
|
||||
|
||||
These tests pin:
|
||||
- inference (no grad): output through the custom op is bit-identical to the
|
||||
original function;
|
||||
- training (requires_grad): gradients through the registered op match the
|
||||
original autograd.Function;
|
||||
- schema/fake-kernel consistency via torch.library.opcheck, both with and
|
||||
without grad-requiring inputs (the latter exercises
|
||||
`test_autograd_registration` — would have caught a missing/broken backward
|
||||
at unit-test time).
|
||||
|
||||
GPU assumptions: requires CUDA and the FA2 `flash_attn` varlen package.
|
||||
Skips on CPU and when flash_attn is unavailable.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def no_pad_impls():
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for the masked/varlen custom-op tests")
|
||||
try:
|
||||
from fastvideo.attention.utils import flash_attn_no_pad as mod
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"flash_attn_no_pad not importable (flash_attn missing?): {exc}")
|
||||
return mod
|
||||
|
||||
|
||||
def _dtype_skip(dtype):
|
||||
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("bfloat16 is not supported on this GPU")
|
||||
|
||||
|
||||
def _fa2_only(mod):
|
||||
if mod._FA_VARLEN_VERSION != "2":
|
||||
pytest.skip(
|
||||
f"register_autograd is only wired for FA2; got "
|
||||
f"_FA_VARLEN_VERSION={mod._FA_VARLEN_VERSION!r}"
|
||||
)
|
||||
|
||||
|
||||
def _padding_mask(batch, seqlen, valid_lens, device):
|
||||
mask = torch.zeros(batch, seqlen, dtype=torch.bool, device=device)
|
||||
for i, n in enumerate(valid_lens):
|
||||
mask[i, :n] = True
|
||||
return mask
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# flash_attn_no_pad (masked self-attention) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_no_pad_inference_matches_original(no_pad_impls, dtype):
|
||||
"""No-grad path through the custom op is bit-identical to the original."""
|
||||
_dtype_skip(dtype)
|
||||
mod = no_pad_impls
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
b, s, h, d = 2, 64, 4, 64
|
||||
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=dtype)
|
||||
mask = _padding_mask(b, s, [64, 48], device)
|
||||
|
||||
with torch.inference_mode():
|
||||
out_ref = mod.flash_attn_no_pad(qkv, mask, causal=False, dropout_p=0.0)
|
||||
out_test = mod.flash_attn_no_pad_compilable(qkv, mask, causal=False, dropout_p=0.0)
|
||||
torch.testing.assert_close(out_test, out_ref, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_no_pad_training_backward_through_registered_autograd(no_pad_impls, dtype):
|
||||
"""FA2: grads flow through the registered op and match the original."""
|
||||
_dtype_skip(dtype)
|
||||
mod = no_pad_impls
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
b, s, h, d = 2, 64, 4, 64
|
||||
mask = _padding_mask(b, s, [64, 48], device)
|
||||
|
||||
qkv_ref = torch.randn(b, s, 3, h, d, device=device, dtype=dtype, requires_grad=True)
|
||||
qkv_test = qkv_ref.detach().clone().requires_grad_(True)
|
||||
|
||||
out_ref = mod.flash_attn_no_pad(qkv_ref, mask, causal=False, dropout_p=0.0)
|
||||
# Go through the compilable wrapper (which on FA2 unconditionally routes
|
||||
# to the registered op — no carve-out).
|
||||
out_test = mod.flash_attn_no_pad_compilable(qkv_test, mask, causal=False, dropout_p=0.0)
|
||||
|
||||
dout = torch.randn_like(out_ref)
|
||||
(dqkv_ref,) = torch.autograd.grad((out_ref * dout).sum(), (qkv_ref,))
|
||||
(dqkv_test,) = torch.autograd.grad((out_test * dout).sum(), (qkv_test,))
|
||||
|
||||
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
|
||||
torch.testing.assert_close(dqkv_test, dqkv_ref, atol=atol, rtol=rtol)
|
||||
|
||||
|
||||
def test_no_pad_forward_opcheck(no_pad_impls):
|
||||
"""Schema/fake-kernel consistency for the custom op (forward only)."""
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
b, s, h, d = 2, 64, 4, 64
|
||||
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=torch.float16)
|
||||
mask = _padding_mask(b, s, [64, 48], device)
|
||||
torch.library.opcheck(
|
||||
torch.ops.fastvideo._flash_attn_no_pad_forward,
|
||||
(qkv, mask, False, 0.0, None, False),
|
||||
)
|
||||
|
||||
|
||||
def test_no_pad_opcheck_with_grad_inputs(no_pad_impls):
|
||||
"""FA2: full opcheck including ``test_autograd_registration`` — catches a
|
||||
missing/inconsistent backward at unit-test time."""
|
||||
mod = no_pad_impls
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
b, s, h, d = 2, 64, 4, 64
|
||||
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=torch.float16, requires_grad=True)
|
||||
mask = _padding_mask(b, s, [64, 48], device)
|
||||
torch.library.opcheck(
|
||||
torch.ops.fastvideo._flash_attn_no_pad_forward,
|
||||
(qkv, mask, False, 0.0, None, False),
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# flash_attn_varlen_qk_no_pad (cross-attn / unequal q-k seqlen) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_varlen_qk_inference_matches_original(no_pad_impls, dtype):
|
||||
_dtype_skip(dtype)
|
||||
mod = no_pad_impls
|
||||
# The varlen-qk custom op's real forward goes through FA's varlen func; off
|
||||
# FA2 that path is the pre-existing FA3 `dropout_p` carve-out (out of scope
|
||||
# here), so skip cleanly like the autograd tests until that's fixed.
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(1)
|
||||
device = torch.device("cuda")
|
||||
b, sq, sk, h, d = 2, 48, 64, 4, 64
|
||||
q = torch.randn(b, sq, h, d, device=device, dtype=dtype)
|
||||
k = torch.randn(b, sk, h, d, device=device, dtype=dtype)
|
||||
v = torch.randn(b, sk, h, d, device=device, dtype=dtype)
|
||||
qmask = _padding_mask(b, sq, [48, 40], device)
|
||||
kmask = _padding_mask(b, sk, [64, 56], device)
|
||||
|
||||
with torch.inference_mode():
|
||||
out_ref = mod.flash_attn_varlen_qk_no_pad(q, k, v, qmask, kmask, causal=False, dropout_p=0.0)
|
||||
out_test = mod.flash_attn_varlen_qk_no_pad_compilable(q, k, v, qmask, kmask, causal=False, dropout_p=0.0)
|
||||
torch.testing.assert_close(out_test, out_ref, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_varlen_qk_training_backward_through_registered_autograd(no_pad_impls, dtype):
|
||||
"""FA2: grads through q, k, v all flow via the registered op and match."""
|
||||
_dtype_skip(dtype)
|
||||
mod = no_pad_impls
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(1)
|
||||
device = torch.device("cuda")
|
||||
b, sq, sk, h, d = 2, 48, 64, 4, 64
|
||||
qmask = _padding_mask(b, sq, [48, 40], device)
|
||||
kmask = _padding_mask(b, sk, [64, 56], device)
|
||||
|
||||
q_ref = torch.randn(b, sq, h, d, device=device, dtype=dtype, requires_grad=True)
|
||||
k_ref = torch.randn(b, sk, h, d, device=device, dtype=dtype, requires_grad=True)
|
||||
v_ref = torch.randn(b, sk, h, d, device=device, dtype=dtype, requires_grad=True)
|
||||
q_test = q_ref.detach().clone().requires_grad_(True)
|
||||
k_test = k_ref.detach().clone().requires_grad_(True)
|
||||
v_test = v_ref.detach().clone().requires_grad_(True)
|
||||
|
||||
out_ref = mod.flash_attn_varlen_qk_no_pad(q_ref, k_ref, v_ref, qmask, kmask, causal=False, dropout_p=0.0)
|
||||
out_test = mod.flash_attn_varlen_qk_no_pad_compilable(q_test, k_test, v_test, qmask, kmask,
|
||||
causal=False, dropout_p=0.0)
|
||||
|
||||
dout = torch.randn_like(out_ref)
|
||||
dq_ref, dk_ref, dv_ref = torch.autograd.grad((out_ref * dout).sum(), (q_ref, k_ref, v_ref))
|
||||
dq_test, dk_test, dv_test = torch.autograd.grad((out_test * dout).sum(), (q_test, k_test, v_test))
|
||||
|
||||
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
|
||||
torch.testing.assert_close(dq_test, dq_ref, atol=atol, rtol=rtol)
|
||||
torch.testing.assert_close(dk_test, dk_ref, atol=atol, rtol=rtol)
|
||||
torch.testing.assert_close(dv_test, dv_ref, atol=atol, rtol=rtol)
|
||||
|
||||
|
||||
def test_varlen_qk_forward_opcheck(no_pad_impls):
|
||||
mod = no_pad_impls
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(1)
|
||||
device = torch.device("cuda")
|
||||
b, sq, sk, h, d = 2, 48, 64, 4, 64
|
||||
q = torch.randn(b, sq, h, d, device=device, dtype=torch.float16)
|
||||
k = torch.randn(b, sk, h, d, device=device, dtype=torch.float16)
|
||||
v = torch.randn(b, sk, h, d, device=device, dtype=torch.float16)
|
||||
qmask = _padding_mask(b, sq, [48, 40], device)
|
||||
kmask = _padding_mask(b, sk, [64, 56], device)
|
||||
torch.library.opcheck(
|
||||
torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward,
|
||||
(q, k, v, qmask, kmask, False, 0.0, None, False),
|
||||
)
|
||||
|
||||
|
||||
def test_varlen_qk_opcheck_with_grad_inputs(no_pad_impls):
|
||||
"""FA2: full opcheck with requires_grad inputs (autograd-registration check)."""
|
||||
mod = no_pad_impls
|
||||
_fa2_only(mod)
|
||||
torch.manual_seed(1)
|
||||
device = torch.device("cuda")
|
||||
b, sq, sk, h, d = 2, 48, 64, 4, 64
|
||||
q = torch.randn(b, sq, h, d, device=device, dtype=torch.float16, requires_grad=True)
|
||||
k = torch.randn(b, sk, h, d, device=device, dtype=torch.float16, requires_grad=True)
|
||||
v = torch.randn(b, sk, h, d, device=device, dtype=torch.float16, requires_grad=True)
|
||||
qmask = _padding_mask(b, sq, [48, 40], device)
|
||||
kmask = _padding_mask(b, sk, [64, 56], device)
|
||||
torch.library.opcheck(
|
||||
torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward,
|
||||
(q, k, v, qmask, kmask, False, 0.0, None, False),
|
||||
)
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Guard the Modal FA4 defaults that keep CI lanes on their intended backend.
|
||||
|
||||
Pure text/AST analysis: no fastvideo imports, no torch, no Modal client.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
MODAL_ROOT = REPO_ROOT / "fastvideo" / "tests" / "modal"
|
||||
PR_TEST = MODAL_ROOT / "pr_test.py"
|
||||
LAUNCH_L40S_JOB = MODAL_ROOT / "launch_l40s_job.py"
|
||||
SSIM_TEST = MODAL_ROOT / "ssim_test.py"
|
||||
|
||||
|
||||
def _function_strings(path: Path, function_name: str) -> str:
|
||||
source = path.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source)
|
||||
for node in tree.body:
|
||||
if isinstance(node, ast.FunctionDef) and node.name == function_name:
|
||||
return "\n".join(
|
||||
child.value
|
||||
for child in ast.walk(node)
|
||||
if isinstance(child, ast.Constant)
|
||||
and isinstance(child.value, str)
|
||||
)
|
||||
raise AssertionError(f"{function_name} not found in {path}")
|
||||
|
||||
|
||||
def test_generic_l40s_launcher_defaults_fa4_off():
|
||||
source = LAUNCH_L40S_JOB.read_text(encoding="utf-8")
|
||||
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0")' in source
|
||||
|
||||
|
||||
def test_ssim_launcher_keeps_fa4_enabled_by_default():
|
||||
source = SSIM_TEST.read_text(encoding="utf-8")
|
||||
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1")' in source
|
||||
|
||||
|
||||
def test_pr_model_load_and_training_lanes_disable_fa4():
|
||||
lanes = {
|
||||
"run_transformer_tests": "pytest ./fastvideo/tests/transformers -vs",
|
||||
"run_training_tests": "pytest ./fastvideo/tests/training/Vanilla -srP",
|
||||
"run_training_lora_tests": "pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP",
|
||||
"run_training_tests_VSA": "pytest ./fastvideo/tests/training/VSA -srP",
|
||||
"run_distill_dmd_tests": "pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs",
|
||||
"run_self_forcing_tests": "pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs",
|
||||
"run_train_framework_tests": "pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs",
|
||||
"seed_grad_norm_references": "pytest ./fastvideo/tests/train/methods -vs -rs",
|
||||
}
|
||||
|
||||
for function_name, pytest_command in lanes.items():
|
||||
function_strings = _function_strings(PR_TEST, function_name)
|
||||
assert "FASTVIDEO_FA4=0" in function_strings
|
||||
assert pytest_command in function_strings
|
||||
@@ -97,9 +97,10 @@ image = (
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
|
||||
# references. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4). Generic ad hoc jobs should follow the
|
||||
# product default unless a caller opts in through the local env or
|
||||
# --env-vars.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0"),
|
||||
})
|
||||
)
|
||||
|
||||
|
||||
@@ -73,8 +73,10 @@ ci_env_secret = modal.Secret.from_dict({
|
||||
**({
|
||||
"UV_TORCH_BACKEND": uv_torch_backend_override
|
||||
} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
|
||||
# SSIM/perf baselines. Caller override wins.
|
||||
# FA4 is opt-in (FASTVIDEO_FA4). Keep the default enabled for
|
||||
# inference/perf parity; model-load and training lanes that do not exercise
|
||||
# FA4 explicitly set FASTVIDEO_FA4=0 in their command strings below.
|
||||
# Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
})
|
||||
|
||||
@@ -184,7 +186,8 @@ def run_vae_tests():
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_transformer_tests():
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs"
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/transformers -vs"
|
||||
)
|
||||
|
||||
|
||||
@@ -197,7 +200,8 @@ def run_transformer_tests():
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_tests():
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
)
|
||||
|
||||
|
||||
@@ -210,7 +214,8 @@ def run_training_tests():
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_lora_tests():
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
|
||||
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
|
||||
)
|
||||
|
||||
|
||||
@@ -220,7 +225,7 @@ def run_training_lora_tests():
|
||||
secrets=[wandb_secret, ci_env_secret])
|
||||
def run_training_tests_VSA():
|
||||
run_test(
|
||||
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
"wandb login $WANDB_API_KEY && FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/VSA -srP"
|
||||
)
|
||||
|
||||
|
||||
@@ -254,7 +259,7 @@ def run_inference_lora_tests():
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[ci_env_secret])
|
||||
def run_distill_dmd_tests():
|
||||
run_test(
|
||||
"pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
|
||||
@app.function(gpu="L40S:2",
|
||||
@@ -263,7 +268,8 @@ def run_distill_dmd_tests():
|
||||
secrets=[wandb_secret, ci_env_secret])
|
||||
def run_self_forcing_tests():
|
||||
run_test(
|
||||
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
|
||||
"wandb login $WANDB_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
|
||||
)
|
||||
|
||||
|
||||
@@ -323,7 +329,8 @@ def run_dreamverse_app_tests():
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_train_framework_tests():
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
|
||||
)
|
||||
|
||||
|
||||
@@ -349,7 +356,8 @@ def seed_grad_norm_references():
|
||||
the local command and the ``_DEVICE_MAPPINGS`` table.
|
||||
"""
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
|
||||
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
|
||||
"FASTVIDEO_FA4=0 FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Unit tests for the get_rotary_pos_embed memoization cache."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.layers.rotary_embedding import (
|
||||
_ROTARY_POS_EMBED_CACHE,
|
||||
_ROTARY_POS_EMBED_CACHE_MAXSIZE,
|
||||
get_rotary_pos_embed,
|
||||
)
|
||||
|
||||
|
||||
def _rope_dim_list(hidden_size: int, heads_num: int) -> list[int]:
|
||||
"""Return the default 3-axis rope_dim_list used by the video DiTs."""
|
||||
d = hidden_size // heads_num
|
||||
return [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
|
||||
|
||||
def _call(
|
||||
rope_sizes=(21, 30, 52),
|
||||
hidden_size=1536,
|
||||
heads_num=12,
|
||||
rope_dim_list="default",
|
||||
rope_theta=10000.0,
|
||||
dtype=torch.float64,
|
||||
start_frame=0,
|
||||
use_real=True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Thin wrapper around get_rotary_pos_embed with DiT-like defaults."""
|
||||
if rope_dim_list == "default":
|
||||
rope_dim_list = _rope_dim_list(hidden_size, heads_num)
|
||||
return get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
rope_dim_list,
|
||||
rope_theta,
|
||||
dtype=dtype,
|
||||
start_frame=start_frame,
|
||||
use_real=use_real,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_cache():
|
||||
"""Isolate every test by clearing the module-level cache around it."""
|
||||
_ROTARY_POS_EMBED_CACHE.clear()
|
||||
yield
|
||||
_ROTARY_POS_EMBED_CACHE.clear()
|
||||
|
||||
|
||||
def test_repeated_call_hits_cache():
|
||||
"""A second identical call returns the exact same tensor objects."""
|
||||
cos1, sin1 = _call()
|
||||
cos2, sin2 = _call()
|
||||
assert cos1 is cos2 and sin1 is sin2
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == 1
|
||||
|
||||
|
||||
def test_many_identical_calls_keep_single_entry():
|
||||
"""Many identical calls never grow the cache beyond one entry."""
|
||||
for _ in range(10):
|
||||
_call()
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rope_sizes,hidden_size,heads_num,dtype,use_real",
|
||||
[
|
||||
((21, 30, 52), 1536, 12, torch.float64, True),
|
||||
((21, 45, 80), 5120, 40, torch.float64, True),
|
||||
((1, 16, 16), 1536, 12, torch.float32, True),
|
||||
((4, 8, 8), 1536, 12, torch.float64, False),
|
||||
],
|
||||
)
|
||||
def test_cached_matches_fresh_recompute(rope_sizes, hidden_size, heads_num,
|
||||
dtype, use_real):
|
||||
"""Cached tensors are bitwise-equal to a fresh uncached recompute."""
|
||||
cos_cached, sin_cached = _call(rope_sizes=rope_sizes,
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
dtype=dtype,
|
||||
use_real=use_real)
|
||||
_ROTARY_POS_EMBED_CACHE.clear()
|
||||
cos_fresh, sin_fresh = _call(rope_sizes=rope_sizes,
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
dtype=dtype,
|
||||
use_real=use_real)
|
||||
assert torch.equal(cos_cached, cos_fresh)
|
||||
assert torch.equal(sin_cached, sin_fresh)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs_a,kwargs_b",
|
||||
[
|
||||
({"rope_sizes": (21, 30, 52)}, {"rope_sizes": (21, 45, 80)}),
|
||||
({"dtype": torch.float64}, {"dtype": torch.float32}),
|
||||
({"use_real": True}, {"use_real": False}),
|
||||
({"start_frame": 0}, {"start_frame": 3}),
|
||||
({"rope_theta": 10000.0}, {"rope_theta": 5000.0}),
|
||||
({"shard_dim": 0}, {"shard_dim": 1}),
|
||||
],
|
||||
)
|
||||
def test_distinct_args_create_distinct_entries(kwargs_a, kwargs_b):
|
||||
"""Any output-affecting argument difference yields a separate cache entry."""
|
||||
_call(**kwargs_a)
|
||||
_call(**kwargs_b)
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == 2
|
||||
|
||||
|
||||
def test_none_rope_dim_list_shares_key_with_equivalent_list():
|
||||
"""None rope_dim_list and its derived explicit list map to one entry."""
|
||||
# head_dim must be divisible by 3 for the None branch to stay valid.
|
||||
hidden_size, heads_num = 1536, 16 # head_dim == 96 -> [32, 32, 32]
|
||||
_call(rope_dim_list=None, hidden_size=hidden_size, heads_num=heads_num)
|
||||
before = len(_ROTARY_POS_EMBED_CACHE)
|
||||
_call(rope_dim_list=[32, 32, 32], hidden_size=hidden_size, heads_num=heads_num)
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == before == 1
|
||||
|
||||
|
||||
def test_use_real_controls_last_dim():
|
||||
"""use_real=True spans full head_dim; use_real=False spans half."""
|
||||
cos_full, _ = _call(use_real=True)
|
||||
cos_half, _ = _call(use_real=False)
|
||||
assert cos_full.shape[-1] == 128
|
||||
assert cos_half.shape[-1] == 64
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rope_sizes", [(1, 1, 1), (1, 30, 52), (21, 1, 1)])
|
||||
def test_degenerate_grid_shapes(rope_sizes):
|
||||
"""Degenerate single-element axes still produce a correctly sized table."""
|
||||
cos, sin = _call(rope_sizes=rope_sizes)
|
||||
expected = rope_sizes[0] * rope_sizes[1] * rope_sizes[2]
|
||||
assert cos.shape[0] == expected
|
||||
assert sin.shape[0] == expected
|
||||
|
||||
|
||||
def test_scalar_and_list_factors_are_hashable_and_distinct():
|
||||
"""List-valued rescale factors are hashable and keyed apart from scalars."""
|
||||
_call(theta_rescale_factor=1.0)
|
||||
_call(theta_rescale_factor=[1.0, 1.0, 1.0])
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == 2
|
||||
|
||||
|
||||
def test_caller_device_copy_does_not_corrupt_cache():
|
||||
"""The .to()/.float() copy callers perform must not mutate cached tensors."""
|
||||
cos, _ = _call()
|
||||
snapshot = cos.clone()
|
||||
_ = cos.to("cpu").float()
|
||||
cos_again, _ = _call()
|
||||
assert torch.equal(cos_again, snapshot)
|
||||
|
||||
|
||||
def test_start_frame_offsets_values():
|
||||
"""A non-zero start_frame shifts the temporal positions, changing output."""
|
||||
cos0, _ = _call(start_frame=0)
|
||||
cos3, _ = _call(start_frame=3)
|
||||
assert not torch.equal(cos0, cos3)
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == 2
|
||||
|
||||
|
||||
def test_cache_is_bounded_and_evicts_oldest():
|
||||
"""The cache caps at the max size and evicts the oldest entry first."""
|
||||
# Tiny grids keep this lightweight; each start_frame is a distinct key.
|
||||
overshoot = _ROTARY_POS_EMBED_CACHE_MAXSIZE + 4
|
||||
for frame in range(overshoot):
|
||||
_call(rope_sizes=(2, 2, 2), start_frame=frame)
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) <= _ROTARY_POS_EMBED_CACHE_MAXSIZE
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == _ROTARY_POS_EMBED_CACHE_MAXSIZE
|
||||
# The earliest-inserted frames must have been evicted; the latest survive.
|
||||
surviving = {key[-2] for key in _ROTARY_POS_EMBED_CACHE} # start_frame slot
|
||||
assert overshoot - 1 in surviving
|
||||
assert 0 not in surviving
|
||||
|
||||
|
||||
def test_cache_hit_refreshes_recency():
|
||||
"""Re-accessing an entry protects it from eviction over an untouched one."""
|
||||
_call(rope_sizes=(2, 2, 2), start_frame=0) # entry we will keep hot
|
||||
for frame in range(1, _ROTARY_POS_EMBED_CACHE_MAXSIZE):
|
||||
_call(rope_sizes=(2, 2, 2), start_frame=frame)
|
||||
assert len(_ROTARY_POS_EMBED_CACHE) == _ROTARY_POS_EMBED_CACHE_MAXSIZE
|
||||
_call(rope_sizes=(2, 2, 2), start_frame=0) # hit -> frame 0 becomes most recent
|
||||
_call(rope_sizes=(2, 2, 2), start_frame=99) # miss -> evicts now-oldest (frame 1)
|
||||
surviving = {key[-2] for key in _ROTARY_POS_EMBED_CACHE}
|
||||
assert 0 in surviving
|
||||
assert 1 not in surviving
|
||||
@@ -62,12 +62,29 @@ def resolve_inference_device_reference_folder(logger: Logger) -> str:
|
||||
return device_reference_folder
|
||||
|
||||
|
||||
def _find_reference_video(reference_folder: str, prompt: str) -> str:
|
||||
def _find_reference_media(
|
||||
reference_folder: str,
|
||||
prompt: str,
|
||||
*,
|
||||
media_extension: str,
|
||||
) -> str:
|
||||
"""Pick a reference file whose basename contains the prompt prefix."""
|
||||
prompt_prefix = prompt[:100].strip()
|
||||
allowed = (media_extension.lower(), ".mp4", ".png", ".jpg", ".jpeg")
|
||||
matches: list[str] = []
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith(".mp4") and prompt_prefix in filename:
|
||||
return os.path.join(reference_folder, filename)
|
||||
raise FileNotFoundError("Reference video missing")
|
||||
low = filename.lower()
|
||||
if not any(low.endswith(ext) for ext in allowed):
|
||||
continue
|
||||
if prompt_prefix in filename:
|
||||
matches.append(filename)
|
||||
if not matches:
|
||||
raise FileNotFoundError("Reference media missing")
|
||||
preferred = media_extension.lower().lstrip(".")
|
||||
for name in matches:
|
||||
if name.lower().endswith(f".{preferred}"):
|
||||
return os.path.join(reference_folder, name)
|
||||
return os.path.join(reference_folder, matches[0])
|
||||
|
||||
|
||||
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
|
||||
@@ -80,21 +97,23 @@ def _assert_similarity(
|
||||
*,
|
||||
logger: Logger,
|
||||
output_dir: str,
|
||||
output_video_name: str,
|
||||
output_media_name: str,
|
||||
reference_folder: str,
|
||||
prompt: str,
|
||||
num_inference_steps: int,
|
||||
min_acceptable_ssim: float,
|
||||
model_id: str,
|
||||
attention_backend_name: str,
|
||||
media_extension: str,
|
||||
) -> None:
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
generated_media_path = os.path.join(output_dir, output_media_name)
|
||||
artifact_kind = "image" if media_extension.lower() in (".png", ".jpg", ".jpeg") else "video"
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing: %s", reference_folder)
|
||||
xfail_missing_reference_in_bootstrap_mode(
|
||||
generated_artifact_path=generated_video_path,
|
||||
generated_artifact_path=generated_media_path,
|
||||
reference_folder=reference_folder,
|
||||
artifact_kind="video",
|
||||
artifact_kind=artifact_kind,
|
||||
)
|
||||
error_msg = (
|
||||
f"Reference video folder does not exist: {reference_folder}\n"
|
||||
@@ -104,28 +123,32 @@ def _assert_similarity(
|
||||
raise FileNotFoundError(error_msg)
|
||||
|
||||
try:
|
||||
reference_video_path = _find_reference_video(reference_folder, prompt)
|
||||
reference_media_path = _find_reference_media(
|
||||
reference_folder,
|
||||
prompt,
|
||||
media_extension=media_extension,
|
||||
)
|
||||
except FileNotFoundError as error:
|
||||
logger.error(
|
||||
"Reference video not found for prompt: %s with backend: %s",
|
||||
"Reference media not found for prompt: %s with backend: %s",
|
||||
prompt,
|
||||
attention_backend_name,
|
||||
)
|
||||
xfail_missing_reference_in_bootstrap_mode(
|
||||
generated_artifact_path=generated_video_path,
|
||||
generated_artifact_path=generated_media_path,
|
||||
reference_folder=reference_folder,
|
||||
artifact_kind="video",
|
||||
artifact_kind=artifact_kind,
|
||||
)
|
||||
raise error
|
||||
|
||||
logger.info(
|
||||
"Computing SSIM between %s and %s",
|
||||
reference_video_path,
|
||||
generated_video_path,
|
||||
reference_media_path,
|
||||
generated_media_path,
|
||||
)
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_video_path,
|
||||
generated_video_path,
|
||||
reference_media_path,
|
||||
generated_media_path,
|
||||
use_ms_ssim=True,
|
||||
)
|
||||
|
||||
@@ -136,8 +159,8 @@ def _assert_similarity(
|
||||
success = write_ssim_results(
|
||||
output_dir,
|
||||
ssim_values,
|
||||
reference_video_path,
|
||||
generated_video_path,
|
||||
reference_media_path,
|
||||
generated_media_path,
|
||||
num_inference_steps,
|
||||
prompt,
|
||||
)
|
||||
@@ -222,6 +245,7 @@ def run_text_to_video_similarity_test(
|
||||
min_acceptable_ssim: float,
|
||||
init_kwargs_override: dict[str, object] | None = None,
|
||||
generation_kwargs_override: dict[str, object] | None = None,
|
||||
media_extension: str = ".mp4",
|
||||
) -> None:
|
||||
with attention_backend(attention_backend_name):
|
||||
output_dir = build_generated_output_dir(
|
||||
@@ -230,9 +254,9 @@ def run_text_to_video_similarity_test(
|
||||
model_id,
|
||||
attention_backend_name,
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
output_media_name = f"{prompt[:100].strip()}{media_extension}"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
_remove_stale_generated_video(output_dir, output_media_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
@@ -274,13 +298,14 @@ def run_text_to_video_similarity_test(
|
||||
_assert_similarity(
|
||||
logger=logger,
|
||||
output_dir=output_dir,
|
||||
output_video_name=output_video_name,
|
||||
output_media_name=output_media_name,
|
||||
reference_folder=reference_folder,
|
||||
prompt=prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
min_acceptable_ssim=min_acceptable_ssim,
|
||||
model_id=model_id,
|
||||
attention_backend_name=attention_backend_name,
|
||||
media_extension=media_extension,
|
||||
)
|
||||
|
||||
|
||||
@@ -306,9 +331,9 @@ def run_image_to_video_similarity_test(
|
||||
model_id,
|
||||
attention_backend_name,
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
output_media_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
_remove_stale_generated_video(output_dir, output_media_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
@@ -351,11 +376,12 @@ def run_image_to_video_similarity_test(
|
||||
_assert_similarity(
|
||||
logger=logger,
|
||||
output_dir=output_dir,
|
||||
output_video_name=output_video_name,
|
||||
output_media_name=output_media_name,
|
||||
reference_folder=reference_folder,
|
||||
prompt=prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
min_acceptable_ssim=min_acceptable_ssim,
|
||||
model_id=model_id,
|
||||
attention_backend_name=attention_backend_name,
|
||||
media_extension=".mp4",
|
||||
)
|
||||
|
||||
BIN
Binary file not shown.
|
After Width: | Height: | Size: 90 KiB |
@@ -0,0 +1,134 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.api.flux import FluxSamplingParam
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
# MS-SSIM gate (see module docstring).
|
||||
FLUX_T2I_MIN_SSIM = 0.98
|
||||
|
||||
FLUX_MODEL_PATH = os.getenv(
|
||||
"FLUX_T2I_MODEL_DIR",
|
||||
"black-forest-labs/FLUX.1-dev",
|
||||
)
|
||||
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
("RTX 4090", "RTX4090"),
|
||||
("4090", "RTX4090"),
|
||||
),
|
||||
device_name=get_cuda_device_name(),
|
||||
fallback_device_prefix="L40S",
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
# Folder token must match Hub path with slashes → double underscore (SD3.5
|
||||
# pattern in ``test_sd35_similarity.py``).
|
||||
MODEL_ID = "black-forest-labs__FLUX.1-dev"
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
]
|
||||
|
||||
FLUX_DEFAULT_PARAMS: dict[str, object] = {
|
||||
"num_gpus": 1,
|
||||
"model_path": FLUX_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": 256,
|
||||
"width": 256,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": 8,
|
||||
"guidance_scale": 3.5,
|
||||
"seed": 0,
|
||||
}
|
||||
|
||||
_flux_full_defaults = FluxSamplingParam()
|
||||
FLUX_FULL_QUALITY_PARAMS: dict[str, object] = {
|
||||
"num_gpus": 1,
|
||||
"model_path": FLUX_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": _flux_full_defaults.height,
|
||||
"width": _flux_full_defaults.width,
|
||||
"num_frames": 1,
|
||||
"fps": _flux_full_defaults.fps,
|
||||
"num_inference_steps": _flux_full_defaults.num_inference_steps,
|
||||
"guidance_scale": _flux_full_defaults.guidance_scale,
|
||||
"seed": _flux_full_defaults.seed,
|
||||
}
|
||||
|
||||
FLUX_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: FLUX_DEFAULT_PARAMS,
|
||||
}
|
||||
FLUX_FULL_QUALITY_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: FLUX_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="FLUX T2I SSIM test requires CUDA",
|
||||
)
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(FLUX_MODEL_TO_PARAMS.keys()))
|
||||
def test_flux_t2i_similarity(
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
is_hf_repo = "/" in FLUX_MODEL_PATH and not FLUX_MODEL_PATH.startswith("/")
|
||||
if not is_hf_repo and not os.path.isdir(FLUX_MODEL_PATH):
|
||||
pytest.skip(
|
||||
f"FLUX weights not found at {FLUX_MODEL_PATH} "
|
||||
f"(set FLUX_T2I_MODEL_DIR to override)"
|
||||
)
|
||||
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=FLUX_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FLUX_FULL_QUALITY_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=FLUX_T2I_MIN_SSIM,
|
||||
media_extension=".png",
|
||||
init_kwargs_override={
|
||||
"workload_type": "t2i",
|
||||
"use_fsdp_inference": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"save_video": True,
|
||||
"use_embedded_guidance": True,
|
||||
"true_cfg_scale": 1.0,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,141 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""SSIM-based regression test for GLM-Image generation."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / "glm_image"))
|
||||
GLM_IMAGE_MODEL_PATH = os.getenv("GLM_IMAGE_MODEL_DIR", str(LOCAL_WEIGHTS_DIR))
|
||||
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
("B200", "B200"),
|
||||
),
|
||||
device_name=get_cuda_device_name(),
|
||||
fallback_device_prefix="L40S",
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
MODEL_ID = "zai-org__GLM-Image"
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style.",
|
||||
]
|
||||
|
||||
GLM_IMAGE_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": GLM_IMAGE_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": 256,
|
||||
"width": 256,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.5,
|
||||
"seed": 0,
|
||||
"neg_prompt": "",
|
||||
}
|
||||
|
||||
GLM_IMAGE_FULL_QUALITY_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": GLM_IMAGE_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 1.5,
|
||||
"seed": 0,
|
||||
"neg_prompt": "",
|
||||
}
|
||||
|
||||
GLM_IMAGE_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: GLM_IMAGE_PARAMS,
|
||||
}
|
||||
GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: GLM_IMAGE_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
required = ["transformer", "vae", "text_encoder",
|
||||
"vision_language_encoder", "processor", "tokenizer",
|
||||
"scheduler"]
|
||||
return all((LOCAL_WEIGHTS_DIR / r).exists() for r in required)
|
||||
|
||||
|
||||
def _upstream_glm_image_available() -> bool:
|
||||
try:
|
||||
import transformers
|
||||
except ImportError:
|
||||
return False
|
||||
return hasattr(transformers, "GlmImageForConditionalGeneration")
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="GLM-Image SSIM test requires CUDA",
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image full weights not found at {LOCAL_WEIGHTS_DIR}.",
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
not _upstream_glm_image_available(),
|
||||
reason="GLM-Image needs transformers>=5.0.0rc0 (ships the AR encoder).",
|
||||
)
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(GLM_IMAGE_MODEL_TO_PARAMS.keys()))
|
||||
def test_glm_image_similarity(
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=GLM_IMAGE_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"trust_remote_code": True,
|
||||
"use_fsdp_inference": False,
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"save_video": True,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,231 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""AnyFlow on-policy method tests (CPU-only, no Wan instantiation).
|
||||
|
||||
The full AnyFlowMethod requires a real student/teacher/critic trio plus
|
||||
DMD2's optimizer wiring — too heavyweight for a unit test. These tests
|
||||
exercise the rollout-shape helpers and source-level invariants via
|
||||
``object.__new__`` bypassing of ``__init__``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.distribution_matching.anyflow import AnyFlowMethod
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers: build a "naked" AnyFlowMethod that skips __init__.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _naked_method(
|
||||
*,
|
||||
student_sample_steps: int = 4,
|
||||
use_mean_velocity: bool = True,
|
||||
t_list_override: list[float] | None = None,
|
||||
denoising_step_list: list[float] | None = None,
|
||||
) -> AnyFlowMethod:
|
||||
method = AnyFlowMethod.__new__(AnyFlowMethod)
|
||||
method._student_sample_steps = int(student_sample_steps) # type: ignore[attr-defined]
|
||||
method._use_mean_velocity = bool(use_mean_velocity) # type: ignore[attr-defined]
|
||||
method._t_list_override = ( # type: ignore[attr-defined]
|
||||
list(t_list_override) if t_list_override else None)
|
||||
method._dmd_score_r = 0.0 # type: ignore[attr-defined]
|
||||
method._real_score_guidance = 1.0 # type: ignore[attr-defined]
|
||||
method.cuda_generator = None # type: ignore[attr-defined]
|
||||
method._cfg_uncond = None # type: ignore[attr-defined]
|
||||
method._denoising_step_list_cache = None # type: ignore[attr-defined]
|
||||
|
||||
# Stub _get_denoising_step_list — DMD2 reads method_config but for these
|
||||
# focused tests we want a deterministic schedule.
|
||||
raw = denoising_step_list or [999, 750, 500, 250]
|
||||
cached = torch.tensor(raw, dtype=torch.long)
|
||||
|
||||
def _stub(self, device: torch.device) -> torch.Tensor:
|
||||
return cached.to(device=device)
|
||||
|
||||
bound = _stub.__get__(method, AnyFlowMethod)
|
||||
method._get_denoising_step_list = bound # type: ignore[assignment]
|
||||
return method
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schedule construction.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_rollout_schedule_uses_t_list_override_verbatim() -> None:
|
||||
method = _naked_method(
|
||||
t_list_override=[999.0, 937.0, 833.0, 624.0, 0.0])
|
||||
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
|
||||
torch.testing.assert_close(
|
||||
schedule,
|
||||
torch.tensor([999.0, 937.0, 833.0, 624.0, 0.0],
|
||||
dtype=torch.float32),
|
||||
)
|
||||
|
||||
|
||||
def test_get_rollout_schedule_falls_back_to_denoising_step_list() -> None:
|
||||
method = _naked_method(
|
||||
denoising_step_list=[999, 750, 500, 250])
|
||||
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
|
||||
# Tail must be a 0 boundary so the final Euler step lands at t=0.
|
||||
assert float(schedule[-1].item()) == 0.0
|
||||
# Original step list preserved at the front.
|
||||
torch.testing.assert_close(
|
||||
schedule[:4],
|
||||
torch.tensor([999.0, 750.0, 500.0, 250.0], dtype=torch.float32),
|
||||
)
|
||||
|
||||
|
||||
def test_get_rollout_schedule_does_not_double_append_zero() -> None:
|
||||
method = _naked_method(
|
||||
denoising_step_list=[999, 750, 500, 0])
|
||||
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
|
||||
assert schedule.numel() == 4 # no extra boundary inserted.
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# t_list_override validation.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_anyflow_method_rejects_ascending_t_list_override() -> None:
|
||||
"""__init__ validates that t_list_override is descending. We test the
|
||||
validation logic by stitching together a minimal cfg + role_models
|
||||
path; if construction fails for an unrelated reason we still catch
|
||||
the descending check via the explicit error message."""
|
||||
src = inspect.getsource(AnyFlowMethod.__init__)
|
||||
assert 't_list_override must be descending' in src
|
||||
assert 'descending' in src
|
||||
|
||||
|
||||
def test_anyflow_method_rejects_non_positive_student_sample_steps() -> None:
|
||||
src = inspect.getsource(AnyFlowMethod.__init__)
|
||||
assert 'student_sample_steps must be positive' in src
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Rollout dynamics — stubbed student.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _SpyStudent:
|
||||
"""Stand-in student that records every (t, r) pair seen during a
|
||||
rollout and predicts a constant velocity field."""
|
||||
|
||||
def __init__(self, num_train_timesteps: int = 1000) -> None:
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.seen: list[tuple[float, float]] = []
|
||||
# A single trainable parameter so callers can verify gradient flow.
|
||||
self.param = torch.nn.Parameter(torch.zeros(1))
|
||||
|
||||
def predict_velocity_with_r(
|
||||
self,
|
||||
noisy: torch.Tensor,
|
||||
t: torch.Tensor,
|
||||
r: torch.Tensor,
|
||||
batch: Any,
|
||||
*,
|
||||
conditional: bool = True,
|
||||
cfg_uncond: Any = None,
|
||||
attn_kind: str = "vsa",
|
||||
) -> torch.Tensor:
|
||||
del batch, conditional, cfg_uncond, attn_kind
|
||||
self.seen.append((float(t.flatten()[0].item()),
|
||||
float(r.flatten()[0].item())))
|
||||
# Constant velocity field of magnitude param so we can backprop.
|
||||
return self.param * torch.ones_like(noisy)
|
||||
|
||||
|
||||
def _make_batch(shape: tuple[int, ...]) -> SimpleNamespace:
|
||||
batch = SimpleNamespace()
|
||||
batch.latents = torch.randn(*shape)
|
||||
batch.dmd_latent_vis_dict = {}
|
||||
return batch
|
||||
|
||||
|
||||
def test_rollout_uses_mean_velocity_r_equals_t_next() -> None:
|
||||
"""With use_mean_velocity=True, r at step i must equal t at step i+1."""
|
||||
method = _naked_method(
|
||||
student_sample_steps=4,
|
||||
use_mean_velocity=True,
|
||||
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
|
||||
)
|
||||
student = _SpyStudent()
|
||||
method.student = student # type: ignore[assignment]
|
||||
|
||||
batch = _make_batch((1, 2, 4, 4, 4))
|
||||
_ = method._student_rollout(batch, with_grad=False)
|
||||
|
||||
# 4 forwards = 4 (t, r) pairs.
|
||||
assert len(student.seen) == 4
|
||||
for i in range(3):
|
||||
# r at step i must equal t at step i+1.
|
||||
assert student.seen[i][1] == student.seen[i + 1][0]
|
||||
|
||||
|
||||
def test_rollout_use_mean_velocity_false_uses_r_equal_t() -> None:
|
||||
method = _naked_method(
|
||||
student_sample_steps=2,
|
||||
use_mean_velocity=False,
|
||||
t_list_override=[999.0, 500.0, 0.0],
|
||||
)
|
||||
student = _SpyStudent()
|
||||
method.student = student # type: ignore[assignment]
|
||||
batch = _make_batch((1, 2, 4, 4, 4))
|
||||
_ = method._student_rollout(batch, with_grad=False)
|
||||
for t_seen, r_seen in student.seen:
|
||||
assert t_seen == r_seen
|
||||
|
||||
|
||||
def test_rollout_with_grad_true_produces_differentiable_output() -> None:
|
||||
method = _naked_method(
|
||||
student_sample_steps=4,
|
||||
use_mean_velocity=True,
|
||||
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
|
||||
)
|
||||
student = _SpyStudent()
|
||||
method.student = student # type: ignore[assignment]
|
||||
batch = _make_batch((1, 2, 4, 4, 4))
|
||||
out = method._student_rollout(batch, with_grad=True)
|
||||
assert out.requires_grad, (
|
||||
"Rollout output must keep a gradient so the DMD loss can backprop "
|
||||
"through the chosen step.")
|
||||
out.sum().backward()
|
||||
assert student.param.grad is not None
|
||||
assert student.param.grad.abs().sum() > 0
|
||||
|
||||
|
||||
def test_rollout_with_grad_false_blocks_gradient_completely() -> None:
|
||||
method = _naked_method(
|
||||
student_sample_steps=4,
|
||||
use_mean_velocity=True,
|
||||
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
|
||||
)
|
||||
student = _SpyStudent()
|
||||
method.student = student # type: ignore[assignment]
|
||||
batch = _make_batch((1, 2, 4, 4, 4))
|
||||
out = method._student_rollout(batch, with_grad=False)
|
||||
assert not out.requires_grad
|
||||
|
||||
|
||||
def test_broadcast_grad_step_index_in_range() -> None:
|
||||
method = _naked_method(student_sample_steps=4)
|
||||
for _ in range(20):
|
||||
idx = method._broadcast_grad_step_index(
|
||||
num_steps=4, device=torch.device("cpu"))
|
||||
assert 0 <= idx < 4
|
||||
|
||||
|
||||
def test_broadcast_grad_step_index_rejects_non_positive_num_steps() -> None:
|
||||
method = _naked_method()
|
||||
with pytest.raises(ValueError, match="num_steps must be positive"):
|
||||
method._broadcast_grad_step_index(
|
||||
num_steps=0, device=torch.device("cpu"))
|
||||
@@ -0,0 +1,604 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""AnyFlow pretrain method tests.
|
||||
|
||||
CPU-only unit tests covering:
|
||||
- Config flag defaults (bit-identity preserved on legacy paths).
|
||||
- ``WanTimeTextImageEmbedding`` dual-timestep forward (additive default = bit-identical
|
||||
to legacy; gated mode reproduces AnyFlow's ``(1 - g) * temb + g * delta_emb`` fusion).
|
||||
- ``WanTransformer3DModel.forward`` accepts ``r_timestep``.
|
||||
- ``FlowMapEulerDiscreteScheduler`` numerics: ``apply_shift``, ``get_train_weight``,
|
||||
``step``.
|
||||
- ``(t, r)`` per-batch sampling distribution.
|
||||
- Central-difference target math.
|
||||
- AnyFlow HF checkpoint key remap (``remap_anyflow_keys``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 1: r_embedder config flags default to bit-identity preservation.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_wan_arch_defaults_preserve_bit_identity() -> None:
|
||||
cfg = WanVideoConfig()
|
||||
arch = cfg.arch_config
|
||||
assert arch.r_embedder is False
|
||||
assert arch.r_embedder_fusion == "additive"
|
||||
assert arch.r_embedder_gate_value == 0.25
|
||||
assert arch.r_embedder_deltatime_type == "r"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 2: WanTimeTextImageEmbedding dual-timestep forward.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _init_uninitialized_weights(module: torch.nn.Module, seed: int = 0) -> None:
|
||||
"""FastVideo's ``ReplicatedLinear`` allocates weights with
|
||||
``torch.empty`` and relies on a downstream ``load_weights`` pass to
|
||||
populate them. Unit tests bypass that pass, so weights start as
|
||||
uninitialized garbage (typically NaN/Inf). Manually init every
|
||||
Linear / RMSNorm / LayerNorm parameter so the forward produces
|
||||
deterministic finite outputs.
|
||||
"""
|
||||
torch.manual_seed(seed)
|
||||
with torch.no_grad():
|
||||
for p in module.parameters():
|
||||
if p.ndim >= 2:
|
||||
# Xavier-uniform scaled by inverse fan-in for stable forwards.
|
||||
torch.nn.init.xavier_uniform_(p)
|
||||
else:
|
||||
p.zero_()
|
||||
|
||||
|
||||
def _make_embedder(
|
||||
*,
|
||||
r_embedder: bool,
|
||||
fusion: str = "additive",
|
||||
gate: float = 0.25,
|
||||
deltatime_type: str = "r",
|
||||
init_seed: int = 0,
|
||||
):
|
||||
from fastvideo.models.dits.wanvideo import WanTimeTextImageEmbedding
|
||||
|
||||
emb = WanTimeTextImageEmbedding(
|
||||
dim=32,
|
||||
time_freq_dim=64,
|
||||
text_embed_dim=16,
|
||||
image_embed_dim=None,
|
||||
r_embedder=r_embedder,
|
||||
r_embedder_fusion=fusion,
|
||||
r_embedder_gate_value=gate,
|
||||
r_embedder_deltatime_type=deltatime_type,
|
||||
)
|
||||
_init_uninitialized_weights(emb, seed=init_seed)
|
||||
emb.eval()
|
||||
return emb
|
||||
|
||||
|
||||
def test_embedder_default_path_no_delta_module() -> None:
|
||||
"""When r_embedder=False, delta_embedder must not be allocated."""
|
||||
emb = _make_embedder(r_embedder=False)
|
||||
assert emb.delta_embedder is None
|
||||
|
||||
|
||||
def test_embedder_default_path_is_bit_identical_to_legacy() -> None:
|
||||
"""With r_embedder=False, forward output must match the legacy single-t path
|
||||
(no r_timestep kwarg, no extra computation)."""
|
||||
torch.manual_seed(0)
|
||||
emb = _make_embedder(r_embedder=False)
|
||||
t = torch.randint(0, 1000, (2,), dtype=torch.long)
|
||||
txt = torch.randn(2, 4, 16)
|
||||
temb_a, proj_a, _, _ = emb(t, txt)
|
||||
# Calling without r_timestep again must be deterministic-equal.
|
||||
temb_b, proj_b, _, _ = emb(t, txt)
|
||||
torch.testing.assert_close(temb_a, temb_b)
|
||||
torch.testing.assert_close(proj_a, proj_b)
|
||||
assert temb_a.shape == (2, 32)
|
||||
assert proj_a.shape == (2, 32 * 6)
|
||||
|
||||
|
||||
def test_embedder_enabled_without_r_timestep_is_bit_identical_to_legacy() -> None:
|
||||
"""Even with r_embedder=True, if r_timestep is None at call time the
|
||||
forward must skip the delta path entirely so existing call sites that
|
||||
don't pass r_timestep stay byte-equal to the legacy result."""
|
||||
torch.manual_seed(0)
|
||||
emb_legacy = _make_embedder(r_embedder=False)
|
||||
torch.manual_seed(0)
|
||||
emb_dual = _make_embedder(r_embedder=True, fusion="additive")
|
||||
t = torch.randint(0, 1000, (2,), dtype=torch.long)
|
||||
txt = torch.randn(2, 4, 16)
|
||||
temb_legacy, proj_legacy, _, _ = emb_legacy(t, txt)
|
||||
temb_dual, proj_dual, _, _ = emb_dual(t, txt) # No r_timestep.
|
||||
torch.testing.assert_close(temb_legacy, temb_dual)
|
||||
torch.testing.assert_close(proj_legacy, proj_dual)
|
||||
|
||||
|
||||
def test_embedder_gated_fusion_formula() -> None:
|
||||
"""Gated mode: rt_emb = (1 - g) * temb_t + g * delta_emb (with delta_input=r)."""
|
||||
torch.manual_seed(0)
|
||||
emb = _make_embedder(r_embedder=True, fusion="gated", gate=0.25)
|
||||
t = torch.tensor([500, 500], dtype=torch.long)
|
||||
r = torch.tensor([100, 100], dtype=torch.long)
|
||||
txt = torch.randn(2, 4, 16)
|
||||
|
||||
temb_t = emb.time_embedder(t)
|
||||
delta_emb = emb.delta_embedder(r)
|
||||
expected = 0.75 * temb_t + 0.25 * delta_emb
|
||||
|
||||
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
|
||||
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_embedder_additive_fusion_formula() -> None:
|
||||
"""Additive mode: rt_emb = temb_t + g * delta_emb."""
|
||||
torch.manual_seed(0)
|
||||
emb = _make_embedder(r_embedder=True, fusion="additive", gate=0.3)
|
||||
t = torch.tensor([700, 700], dtype=torch.long)
|
||||
r = torch.tensor([200, 200], dtype=torch.long)
|
||||
txt = torch.randn(2, 4, 16)
|
||||
|
||||
temb_t = emb.time_embedder(t)
|
||||
delta_emb = emb.delta_embedder(r)
|
||||
expected = temb_t + 0.3 * delta_emb
|
||||
|
||||
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
|
||||
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_embedder_deltatime_type_t_minus_r() -> None:
|
||||
"""When deltatime_type='t-r', delta_embedder consumes (t - r)."""
|
||||
torch.manual_seed(0)
|
||||
emb = _make_embedder(
|
||||
r_embedder=True, fusion="gated", gate=0.5, deltatime_type="t-r")
|
||||
t = torch.tensor([800, 800], dtype=torch.long)
|
||||
r = torch.tensor([300, 300], dtype=torch.long)
|
||||
txt = torch.randn(2, 4, 16)
|
||||
|
||||
temb_t = emb.time_embedder(t)
|
||||
delta_emb = emb.delta_embedder(t - r)
|
||||
expected = 0.5 * temb_t + 0.5 * delta_emb
|
||||
|
||||
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
|
||||
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_embedder_invalid_fusion_raises() -> None:
|
||||
with pytest.raises(ValueError, match="r_embedder_fusion"):
|
||||
_make_embedder(r_embedder=True, fusion="bogus")
|
||||
|
||||
|
||||
def test_embedder_invalid_deltatime_type_raises() -> None:
|
||||
with pytest.raises(ValueError, match="r_embedder_deltatime_type"):
|
||||
_make_embedder(r_embedder=True, fusion="gated", deltatime_type="2t-r")
|
||||
|
||||
|
||||
def test_embedder_gate_not_in_state_dict() -> None:
|
||||
"""Gate is a non-persistent buffer; it must not appear in state_dict so
|
||||
checkpoints stay portable across different gate hyperparameters."""
|
||||
emb = _make_embedder(r_embedder=True, fusion="gated", gate=0.25)
|
||||
keys = list(emb.state_dict().keys())
|
||||
assert not any("_r_embedder_gate" in k for k in keys)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 3: WanTransformer3DModel threads r_timestep through.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_wan_transformer_forward_signature_has_r_timestep() -> None:
|
||||
"""The forward signature must declare r_timestep explicitly (not
|
||||
swallowed by **kwargs) so callers and type checkers can see it."""
|
||||
import inspect
|
||||
|
||||
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
|
||||
|
||||
sig = inspect.signature(WanTransformer3DModel.forward)
|
||||
assert "r_timestep" in sig.parameters
|
||||
param = sig.parameters["r_timestep"]
|
||||
assert param.default is None
|
||||
|
||||
|
||||
def test_wan_transformer_init_propagates_r_embedder_config() -> None:
|
||||
"""When the arch config sets r_embedder=True the WanTransformer3DModel
|
||||
constructor must instantiate the embedder with the delta path active.
|
||||
|
||||
We avoid full WanTransformer3DModel instantiation (which requires
|
||||
distributed init) by reading the source's __init__ to confirm it
|
||||
forwards the four arch flags to WanTimeTextImageEmbedding.
|
||||
"""
|
||||
import inspect
|
||||
|
||||
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
|
||||
|
||||
src = inspect.getsource(WanTransformer3DModel.__init__)
|
||||
# All four arch config fields must be passed to WanTimeTextImageEmbedding.
|
||||
assert "r_embedder=config.r_embedder" in src
|
||||
assert "r_embedder_fusion=config.r_embedder_fusion" in src
|
||||
assert "r_embedder_gate_value=config.r_embedder_gate_value" in src
|
||||
assert "r_embedder_deltatime_type=config.r_embedder_deltatime_type" in src
|
||||
|
||||
|
||||
def test_wan_transformer_forward_threads_r_timestep_to_embedder() -> None:
|
||||
"""The forward must pass r_timestep into the embedder call (verified via
|
||||
source inspection to avoid heavyweight distributed bring-up)."""
|
||||
import inspect
|
||||
|
||||
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
|
||||
|
||||
src = inspect.getsource(WanTransformer3DModel.forward)
|
||||
assert "r_timestep=r_timestep" in src, (
|
||||
"WanTransformer3DModel.forward must forward r_timestep into "
|
||||
"self.condition_embedder")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 4: FlowMapEulerDiscreteScheduler numerics.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _scheduler(*, shift: float = 1.0, n_train: int = 1000):
|
||||
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
|
||||
FlowMapEulerDiscreteScheduler, )
|
||||
return FlowMapEulerDiscreteScheduler(
|
||||
num_train_timesteps=n_train, shift=shift)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_set_timesteps_descending() -> None:
|
||||
sched = _scheduler(shift=5.0)
|
||||
sched.set_timesteps(num_inference_steps=4, device=torch.device("cpu"))
|
||||
ts = sched.timesteps
|
||||
# N inference steps → N + 1 boundary entries.
|
||||
assert ts.numel() == 5
|
||||
assert torch.all(ts[:-1] >= ts[1:]) # descending
|
||||
assert ts[-1].item() == 0.0
|
||||
assert ts[0].item() == pytest.approx(1000.0, abs=1e-3)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_set_timesteps_custom_overrides_schedule() -> None:
|
||||
sched = _scheduler(shift=5.0)
|
||||
custom = [999.0, 937.0, 833.0, 624.0, 0.0]
|
||||
sched.set_timesteps(
|
||||
num_inference_steps=4,
|
||||
device=torch.device("cpu"),
|
||||
custom_timesteps=custom,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
sched.timesteps, torch.tensor(custom, dtype=torch.float32))
|
||||
|
||||
|
||||
def test_flow_map_scheduler_custom_timesteps_must_be_descending() -> None:
|
||||
sched = _scheduler(shift=5.0)
|
||||
with pytest.raises(ValueError, match="descending"):
|
||||
sched.set_timesteps(
|
||||
num_inference_steps=4,
|
||||
device=torch.device("cpu"),
|
||||
custom_timesteps=[100.0, 500.0, 900.0],
|
||||
)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_apply_shift_endpoints_invariant() -> None:
|
||||
"""apply_shift fixes the endpoints {0, 1} and produces non-trivial
|
||||
motion in the interior for shift != 1."""
|
||||
sched = _scheduler(shift=5.0)
|
||||
t = torch.tensor([0.0, 0.5, 1.0])
|
||||
shifted = sched.apply_shift(t)
|
||||
torch.testing.assert_close(
|
||||
shifted, torch.tensor([0.0, 5.0 / 6.0, 1.0]), rtol=1e-6, atol=1e-6)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_apply_shift_shift_one_is_identity() -> None:
|
||||
sched = _scheduler(shift=1.0)
|
||||
t = torch.linspace(0.0, 1.0, 100)
|
||||
torch.testing.assert_close(sched.apply_shift(t), t)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_step_one_euler_iteration_matches_formula() -> None:
|
||||
"""One step: x_r = x_t - ((t - r) / N) * model_output."""
|
||||
sched = _scheduler(shift=1.0)
|
||||
sched.set_timesteps(num_inference_steps=4, device=torch.device("cpu"))
|
||||
|
||||
torch.manual_seed(0)
|
||||
x_t = torch.randn(2, 4, 1, 8, 8)
|
||||
v = torch.randn_like(x_t)
|
||||
t = torch.tensor([750.0, 500.0])
|
||||
r = torch.tensor([500.0, 250.0])
|
||||
|
||||
out = sched.step(v, sample=x_t, timestep=t, r_timestep=r)
|
||||
expected = x_t - ((t - r) / 1000.0).view(-1, 1, 1, 1, 1) * v
|
||||
torch.testing.assert_close(out, expected, rtol=1e-6, atol=1e-6)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_get_train_weight_beta08_shape_and_renorm() -> None:
|
||||
"""beta08: t * sqrt(1-t), renormalized so sum equals num_train_timesteps.
|
||||
The interior of the schedule must dominate the endpoints (monotone up
|
||||
then monotone down)."""
|
||||
sched = _scheduler()
|
||||
t = torch.linspace(0.001, 0.999, 1000)
|
||||
w = sched.get_train_weight(t, weight_type="beta08")
|
||||
assert torch.allclose(w.sum(), torch.tensor(1000.0), rtol=1e-3)
|
||||
assert torch.all(w >= 0.0)
|
||||
# Endpoints smaller than the middle bump.
|
||||
mid = len(w) // 2
|
||||
assert w[0] < w[mid]
|
||||
assert w[-1] < w[mid]
|
||||
|
||||
|
||||
def test_flow_map_scheduler_get_train_weight_uniform_is_constant_norm() -> None:
|
||||
sched = _scheduler()
|
||||
t = torch.linspace(0.0, 1.0, 1000)
|
||||
w = sched.get_train_weight(t, weight_type="uniform")
|
||||
assert torch.allclose(w.sum(), torch.tensor(1000.0), rtol=1e-3)
|
||||
# All entries equal to 1.0 after renormalization.
|
||||
torch.testing.assert_close(w, torch.ones_like(w), rtol=1e-6, atol=1e-6)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_get_train_weight_accepts_absolute_units() -> None:
|
||||
"""When t is provided in [0, num_train_timesteps] the helper auto-
|
||||
normalizes; result must match the [0, 1] call."""
|
||||
sched = _scheduler()
|
||||
t_norm = torch.linspace(0.001, 0.999, 1000)
|
||||
t_abs = t_norm * 1000
|
||||
w_norm = sched.get_train_weight(t_norm, weight_type="beta08")
|
||||
w_abs = sched.get_train_weight(t_abs, weight_type="beta08")
|
||||
torch.testing.assert_close(w_norm, w_abs, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_flow_map_scheduler_add_noise_matches_flow_matching_formula() -> None:
|
||||
"""Linear flow-matching: x_t = (1 - sigma) * x_0 + sigma * eps, where
|
||||
sigma = t / num_train_timesteps."""
|
||||
sched = _scheduler()
|
||||
torch.manual_seed(0)
|
||||
x0 = torch.randn(2, 4, 1, 4, 4)
|
||||
eps = torch.randn_like(x0)
|
||||
t = torch.tensor([250.0, 750.0])
|
||||
out = sched.add_noise(x0, eps, t)
|
||||
sigma = (t / 1000.0).view(-1, 1, 1, 1, 1)
|
||||
expected = (1.0 - sigma) * x0 + sigma * eps
|
||||
torch.testing.assert_close(out, expected, rtol=1e-6, atol=1e-6)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 6: (t, r) per-batch sampling.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sample_pair_timesteps_partitions_batch_correctly() -> None:
|
||||
"""For batch=8 with diffusion=0.5/consistency=0.25: 4 r=t, 2 r=0,
|
||||
2 free entries."""
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_sample_pair_timesteps, )
|
||||
|
||||
torch.manual_seed(42)
|
||||
t, r, is_diffusion, is_consistency = _sample_pair_timesteps(
|
||||
batch_size=8,
|
||||
diffusion_ratio=0.5,
|
||||
consistency_ratio=0.25,
|
||||
device=torch.device("cpu"),
|
||||
generator=None,
|
||||
)
|
||||
assert t.shape == (8,)
|
||||
assert r.shape == (8,)
|
||||
assert int(is_diffusion.sum()) == 4
|
||||
assert int(is_consistency.sum()) == 2
|
||||
# The masks must be disjoint.
|
||||
assert not torch.any(is_diffusion & is_consistency)
|
||||
|
||||
diff_idx = torch.nonzero(is_diffusion).flatten()
|
||||
torch.testing.assert_close(r[diff_idx], t[diff_idx])
|
||||
|
||||
cons_idx = torch.nonzero(is_consistency).flatten()
|
||||
torch.testing.assert_close(r[cons_idx], torch.zeros(2))
|
||||
|
||||
free_idx = torch.nonzero(~(is_diffusion | is_consistency)).flatten()
|
||||
assert torch.all(r[free_idx] <= t[free_idx])
|
||||
assert torch.all(r[free_idx] >= 0.0)
|
||||
assert torch.all(t[free_idx] <= 1.0)
|
||||
|
||||
|
||||
def test_sample_pair_timesteps_t_max_r_min_ordering() -> None:
|
||||
"""For the free fraction (ratios=0), t and r come from max/min of two
|
||||
uniform draws so r <= t holds."""
|
||||
torch.manual_seed(0)
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_sample_pair_timesteps, )
|
||||
|
||||
for _ in range(50):
|
||||
t, r, is_diff, is_cons = _sample_pair_timesteps(
|
||||
batch_size=4,
|
||||
diffusion_ratio=0.0,
|
||||
consistency_ratio=0.0,
|
||||
device=torch.device("cpu"),
|
||||
generator=None,
|
||||
)
|
||||
assert torch.all(t >= r)
|
||||
assert torch.all(r >= 0.0)
|
||||
assert torch.all(t <= 1.0)
|
||||
assert int(is_diff.sum()) == 0
|
||||
assert int(is_cons.sum()) == 0
|
||||
|
||||
|
||||
def test_sample_pair_timesteps_rejects_ratios_summing_above_one() -> None:
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_sample_pair_timesteps, )
|
||||
|
||||
with pytest.raises(ValueError, match="must be <= 1"):
|
||||
_sample_pair_timesteps(
|
||||
batch_size=8,
|
||||
diffusion_ratio=0.7,
|
||||
consistency_ratio=0.4,
|
||||
device=torch.device("cpu"),
|
||||
generator=None,
|
||||
)
|
||||
|
||||
|
||||
def test_sample_pair_timesteps_rejects_negative_ratios() -> None:
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_sample_pair_timesteps, )
|
||||
|
||||
with pytest.raises(ValueError, match="non-negative"):
|
||||
_sample_pair_timesteps(
|
||||
batch_size=8,
|
||||
diffusion_ratio=-0.1,
|
||||
consistency_ratio=0.25,
|
||||
device=torch.device("cpu"),
|
||||
generator=None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 7: central-difference target.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _StubStudent:
|
||||
"""Stand-in student for unit-testing the central-difference helper.
|
||||
|
||||
The "velocity prediction" is a closed-form function of x and t
|
||||
(no actual neural network) so we can compare against the analytical
|
||||
derivative.
|
||||
"""
|
||||
|
||||
def __init__(self, alpha: float = 0.3) -> None:
|
||||
self.alpha = float(alpha)
|
||||
|
||||
def predict_velocity_with_r(
|
||||
self,
|
||||
noisy: torch.Tensor,
|
||||
t: torch.Tensor,
|
||||
r: torch.Tensor,
|
||||
batch,
|
||||
*,
|
||||
conditional: bool = True,
|
||||
attn_kind: str = "dense",
|
||||
cfg_uncond=None,
|
||||
) -> torch.Tensor:
|
||||
del batch, conditional, attn_kind, cfg_uncond, r
|
||||
# f(x, t) = x + alpha * (t / 1000) → dF/dt = alpha / 1000.
|
||||
view = [-1] + [1] * (noisy.ndim - 1)
|
||||
return noisy + self.alpha * (t.view(*view).float() / 1000.0)
|
||||
|
||||
|
||||
def test_central_difference_dF_dt_linear_function() -> None:
|
||||
"""For f(x, t) = x + alpha * (t / N), the central-difference estimate
|
||||
of dF/dt must be alpha / N (in absolute t-units that's alpha / N
|
||||
velocity per t-unit, and our helper divides the difference by
|
||||
2*delta -> exactly alpha / N)."""
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_central_difference_dF_dt, )
|
||||
|
||||
student = _StubStudent(alpha=0.3)
|
||||
x = torch.randn(2, 1, 4, 4, 4)
|
||||
latents = torch.zeros_like(x)
|
||||
noise = torch.zeros_like(x) # v_pred = noise - latents = 0
|
||||
t = torch.tensor([500.0, 250.0])
|
||||
r = torch.tensor([100.0, 0.0])
|
||||
|
||||
dF = _central_difference_dF_dt(
|
||||
student=student,
|
||||
batch=None,
|
||||
noisy=x,
|
||||
latents=latents,
|
||||
noise=noise,
|
||||
t=t,
|
||||
r=r,
|
||||
delta=5.0,
|
||||
num_train_timesteps=1000.0,
|
||||
)
|
||||
|
||||
expected = torch.full_like(x, 0.3 / 1000.0)
|
||||
torch.testing.assert_close(dF, expected, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_central_difference_dF_dt_guidance_scaling() -> None:
|
||||
"""With guidance_scale != 1, the result must be divided by guidance."""
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_central_difference_dF_dt, )
|
||||
|
||||
student = _StubStudent(alpha=0.6)
|
||||
x = torch.zeros(1, 1, 4, 4, 4)
|
||||
latents = torch.zeros_like(x)
|
||||
noise = torch.zeros_like(x)
|
||||
t = torch.tensor([500.0])
|
||||
r = torch.tensor([100.0])
|
||||
|
||||
dF_g1 = _central_difference_dF_dt(
|
||||
student=student, batch=None, noisy=x, latents=latents, noise=noise,
|
||||
t=t, r=r, delta=5.0, num_train_timesteps=1000.0, guidance_scale=1.0)
|
||||
dF_g3 = _central_difference_dF_dt(
|
||||
student=student, batch=None, noisy=x, latents=latents, noise=noise,
|
||||
t=t, r=r, delta=5.0, num_train_timesteps=1000.0, guidance_scale=3.0)
|
||||
|
||||
torch.testing.assert_close(dF_g1 / 3.0, dF_g3, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_central_difference_dF_dt_rejects_zero_delta() -> None:
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
_central_difference_dF_dt, )
|
||||
|
||||
with pytest.raises(ValueError, match="delta must be positive"):
|
||||
_central_difference_dF_dt(
|
||||
student=_StubStudent(),
|
||||
batch=None,
|
||||
noisy=torch.zeros(1, 1, 4, 4, 4),
|
||||
latents=torch.zeros(1, 1, 4, 4, 4),
|
||||
noise=torch.zeros(1, 1, 4, 4, 4),
|
||||
t=torch.tensor([500.0]),
|
||||
r=torch.tensor([100.0]),
|
||||
delta=0.0,
|
||||
num_train_timesteps=1000.0,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 9: param_names_mapping handles AnyFlow checkpoints + is a no-op on plain
|
||||
# Wan checkpoints. We don't ship a separate remap_anyflow_keys helper because
|
||||
# the existing param_names_mapping regex mechanism does the job (delta_embedder
|
||||
# rename is a no-op when the source state dict doesn't contain those keys).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_param_names_mapping_includes_delta_embedder_rename() -> None:
|
||||
"""The Wan arch config's param_names_mapping must rename HF AnyFlow
|
||||
delta_embedder weights into FastVideo's internal mlp.fc_in/fc_out layout."""
|
||||
arch = WanVideoConfig().arch_config
|
||||
mapping_keys = list(arch.param_names_mapping.keys())
|
||||
assert any("delta_embedder" in k for k in mapping_keys), (
|
||||
"WanVideoArchConfig.param_names_mapping must include delta_embedder "
|
||||
"rename so HF AnyFlow checkpoints load without a separate adapter")
|
||||
|
||||
|
||||
def test_param_names_mapping_default_doesnt_break_plain_wan_keys() -> None:
|
||||
"""The new delta_embedder regex must not match any key in a plain
|
||||
pretrained Wan2.1 checkpoint (those don't have delta_embedder)."""
|
||||
import re
|
||||
|
||||
plain_wan_keys = [
|
||||
"patch_embedding.weight",
|
||||
"condition_embedder.time_embedder.linear_1.weight",
|
||||
"condition_embedder.time_embedder.linear_2.weight",
|
||||
"condition_embedder.time_proj.weight",
|
||||
"condition_embedder.text_embedder.linear_1.weight",
|
||||
"blocks.0.attn1.to_q.weight",
|
||||
"blocks.0.ffn.net.0.proj.weight",
|
||||
]
|
||||
arch = WanVideoConfig().arch_config
|
||||
delta_regexes = [
|
||||
k for k in arch.param_names_mapping if "delta_embedder" in k
|
||||
]
|
||||
assert delta_regexes, "expected at least one delta_embedder regex"
|
||||
|
||||
for plain in plain_wan_keys:
|
||||
for rx in delta_regexes:
|
||||
assert re.match(rx, plain) is None, (
|
||||
f"delta_embedder regex {rx!r} unexpectedly matched plain "
|
||||
f"Wan key {plain!r}")
|
||||
@@ -0,0 +1,130 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU smoke test for AnyFlow pretrain + on-policy.
|
||||
|
||||
Mirrors ``test_distill_dmd.py`` — runs the new YAML-driven training
|
||||
entrypoint via ``torchrun`` for two iterations to verify end-to-end
|
||||
wiring (model load, optimizer build, forward, backward, step, save).
|
||||
|
||||
CPU-only environments are skipped; the test is intended to fire on the
|
||||
Buildkite ``/test distillation`` lane and on local boxes with at least
|
||||
2 H100/H200 GPUs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[4]
|
||||
PRETRAIN_YAML = (
|
||||
REPO_ROOT
|
||||
/ "examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml")
|
||||
ONPOLICY_YAML = (
|
||||
REPO_ROOT
|
||||
/ "examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml")
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
|
||||
|
||||
def _have_enough_gpus() -> bool:
|
||||
"""Return True iff at least 2 CUDA devices are visible. The smoke test
|
||||
needs HSDP/FSDP with a non-trivial world size; single-GPU bring-up
|
||||
races against the new framework's distributed barriers."""
|
||||
try:
|
||||
import torch
|
||||
except Exception:
|
||||
return False
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
return torch.cuda.device_count() >= 2
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _have_enough_gpus(),
|
||||
reason="AnyFlow smoke test requires >= 2 CUDA devices")
|
||||
|
||||
|
||||
def _run_torchrun(config_path: Path, *, output_dir: Path) -> None:
|
||||
if not config_path.exists():
|
||||
pytest.fail(f"YAML config missing: {config_path}")
|
||||
|
||||
env = os.environ.copy()
|
||||
env.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
env.setdefault("MASTER_PORT", "29551")
|
||||
env.setdefault("WANDB_MODE", "offline")
|
||||
env.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE,
|
||||
"--master_port", env["MASTER_PORT"],
|
||||
"-m", "fastvideo.train.entrypoint.train",
|
||||
"--config", str(config_path),
|
||||
"--training.loop.max_train_steps", "2",
|
||||
"--training.checkpoint.output_dir", str(output_dir),
|
||||
"--training.distributed.num_gpus", NUM_GPUS_PER_NODE,
|
||||
"--training.distributed.hsdp_shard_dim", NUM_GPUS_PER_NODE,
|
||||
"--training.data.train_batch_size", "1",
|
||||
]
|
||||
process = subprocess.run(cmd, capture_output=True, text=True, env=env)
|
||||
if process.stdout:
|
||||
print("STDOUT:", process.stdout)
|
||||
if process.stderr:
|
||||
print("STDERR:", process.stderr)
|
||||
if process.returncode != 0:
|
||||
raise subprocess.CalledProcessError(
|
||||
process.returncode, cmd, process.stdout, process.stderr)
|
||||
|
||||
|
||||
def test_anyflow_pretrain_smoke(tmp_path: Path) -> None:
|
||||
"""Two-iteration pretrain — exercises (t, r) sampling, central-difference
|
||||
target, scale balance, optimizer step, and checkpoint save path."""
|
||||
_run_torchrun(PRETRAIN_YAML, output_dir=tmp_path / "pretrain")
|
||||
|
||||
|
||||
def test_anyflow_onpolicy_smoke(tmp_path: Path) -> None:
|
||||
"""Two-iteration on-policy DMD — exercises the multi-step Euler-flow
|
||||
rollout, grad-step broadcast, DMD2 alternating updates."""
|
||||
# Override init_from to the public Wan2.1-T2V-1.3B-Diffusers checkpoint
|
||||
# for the smoke run; param_names_mapping handles the (non-existent)
|
||||
# delta_embedder rename as a no-op.
|
||||
env = os.environ.copy()
|
||||
env.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
env.setdefault("MASTER_PORT", "29552")
|
||||
env.setdefault("WANDB_MODE", "offline")
|
||||
env.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE,
|
||||
"--master_port", env["MASTER_PORT"],
|
||||
"-m", "fastvideo.train.entrypoint.train",
|
||||
"--config", str(ONPOLICY_YAML),
|
||||
"--training.loop.max_train_steps", "2",
|
||||
"--training.checkpoint.output_dir", str(tmp_path / "onpolicy"),
|
||||
"--training.distributed.num_gpus", NUM_GPUS_PER_NODE,
|
||||
"--training.distributed.hsdp_shard_dim", NUM_GPUS_PER_NODE,
|
||||
"--training.data.train_batch_size", "1",
|
||||
"--models.student.init_from", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--method.student_sample_steps", "2",
|
||||
]
|
||||
process = subprocess.run(cmd, capture_output=True, text=True, env=env)
|
||||
if process.stdout:
|
||||
print("STDOUT:", process.stdout)
|
||||
if process.stderr:
|
||||
print("STDERR:", process.stderr)
|
||||
if process.returncode != 0:
|
||||
raise subprocess.CalledProcessError(
|
||||
process.returncode, cmd, process.stdout, process.stderr)
|
||||
@@ -0,0 +1,185 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import glob
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers import FluxTransformer2DModel as HFFluxTransformer2DModel
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29517")
|
||||
|
||||
_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))
|
||||
_DEFAULT_FLUX_TRANSFORMER = os.path.join(
|
||||
_REPO_ROOT,
|
||||
"official_weights",
|
||||
"FLUX.1-dev",
|
||||
"transformer",
|
||||
)
|
||||
|
||||
|
||||
def _flux_transformer_path() -> str:
|
||||
return os.environ.get("FLUX_TRANSFORMER_PATH", _DEFAULT_FLUX_TRANSFORMER)
|
||||
|
||||
|
||||
def _prepare_latent_image_ids(
|
||||
height: int,
|
||||
width: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype = torch.long,
|
||||
) -> torch.Tensor:
|
||||
"""Match Diffusers ``FluxPipeline._prepare_latent_image_ids`` (batch omitted)."""
|
||||
latent_image_ids = torch.zeros(height, width, 3)
|
||||
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
|
||||
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
|
||||
h, w, c = latent_image_ids.shape
|
||||
latent_image_ids = latent_image_ids.reshape(h * w, c)
|
||||
return latent_image_ids.to(device=device, dtype=dtype)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def torch_sdpa_attention_backend(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
|
||||
requires_cuda = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="FLUX DiT parity test requires CUDA",
|
||||
)
|
||||
|
||||
requires_weights = pytest.mark.skipif(
|
||||
not glob.glob(os.path.join(_flux_transformer_path(), "*.safetensors")),
|
||||
reason=(
|
||||
f"No safetensors under {_flux_transformer_path()} — download FLUX.1-dev "
|
||||
"transformer or set FLUX_TRANSFORMER_PATH"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_weights
|
||||
@pytest.mark.usefixtures("distributed_setup", "torch_sdpa_attention_backend")
|
||||
def test_flux_transformer_parity_vs_diffusers() -> None:
|
||||
"""Single forward: FastVideo DiT vs Diffusers ``FluxTransformer2DModel``."""
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
transformer_path = _flux_transformer_path()
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=transformer_path,
|
||||
dit_cpu_offload=False,
|
||||
dit_layerwise_offload=False,
|
||||
pipeline_config=PipelineConfig(dit_config=FluxDiTConfig(), dit_precision="bf16"),
|
||||
)
|
||||
args.device = device
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(0)
|
||||
torch.manual_seed(0)
|
||||
|
||||
batch_size = 1
|
||||
latent_h, latent_w = 4, 4
|
||||
img_seq = latent_h * latent_w
|
||||
text_len = 32
|
||||
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
img_seq,
|
||||
64,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
generator=generator,
|
||||
)
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
text_len,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
generator=generator,
|
||||
)
|
||||
pooled_projections = torch.randn(
|
||||
batch_size,
|
||||
768,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
# Diffusers pipeline passes scheduler timesteps / 1000 (float, same dtype as latents).
|
||||
timestep = torch.tensor([512.0], device=device, dtype=precision) / 1000.0
|
||||
guidance = torch.full((batch_size,), 3.5, device=device, dtype=torch.float32)
|
||||
|
||||
txt_ids = torch.zeros(text_len, 3, device=device, dtype=torch.long)
|
||||
img_ids = _prepare_latent_image_ids(latent_h, latent_w, device, dtype=torch.long)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="dummy")
|
||||
|
||||
# One ~12B model at a time avoids peak VRAM from holding both checkpoints.
|
||||
loader = TransformerLoader()
|
||||
fv_model = loader.load(transformer_path, args).to(device=device, dtype=precision)
|
||||
fv_model.eval()
|
||||
with (
|
||||
torch.no_grad(),
|
||||
torch.amp.autocast("cuda", dtype=precision),
|
||||
set_forward_context(
|
||||
current_timestep=512,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
),
|
||||
):
|
||||
fv_out = fv_model(
|
||||
hidden_states=hidden_states.clone(),
|
||||
encoder_hidden_states=encoder_hidden_states.clone(),
|
||||
pooled_projections=pooled_projections.clone(),
|
||||
timestep=timestep.clone(),
|
||||
guidance=guidance.clone(),
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
fv_out_cpu = fv_out.detach().float().cpu()
|
||||
del fv_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
hf_model = (
|
||||
HFFluxTransformer2DModel.from_pretrained(
|
||||
transformer_path,
|
||||
torch_dtype=precision,
|
||||
)
|
||||
.to(device)
|
||||
.eval()
|
||||
)
|
||||
with torch.no_grad(), torch.amp.autocast("cuda", dtype=precision):
|
||||
hf_out = hf_model(
|
||||
hidden_states=hidden_states.clone(),
|
||||
encoder_hidden_states=encoder_hidden_states.clone(),
|
||||
pooled_projections=pooled_projections.clone(),
|
||||
timestep=timestep.clone(),
|
||||
guidance=guidance.clone(),
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
assert hf_out.shape == fv_out_cpu.shape
|
||||
hf_cpu = hf_out.float().cpu()
|
||||
abs_diff = (hf_cpu - fv_out_cpu).abs()
|
||||
print(f"[FLUX DiT parity] max_diff={abs_diff.max():.4f} mean_diff={abs_diff.mean():.4f} "
|
||||
f"median_diff={abs_diff.median():.4f} p99_diff="
|
||||
f"{abs_diff.flatten().kthvalue(int(0.99 * abs_diff.numel())).values:.4f}")
|
||||
# bfloat16 accumulation over 57 transformer layers produces tail errors up to ~0.5
|
||||
# on isolated elements (median=0, mean~0.04 on L40S). atol=0.5 catches real bugs
|
||||
# (wrong weights / missing layers) which produce mean_diff >> 0.1.
|
||||
assert_close(hf_cpu, fv_out_cpu, atol=0.5, rtol=0.0)
|
||||
@@ -77,13 +77,32 @@ def _read_video_frames(path: str) -> torch.Tensor:
|
||||
return torch.stack(frames)
|
||||
|
||||
|
||||
def _read_image_as_single_frame_video(path: str) -> torch.Tensor:
|
||||
"""Read one image as a single-frame ``(1, C, H, W)`` uint8 tensor."""
|
||||
from torchvision.io import read_image
|
||||
|
||||
img = read_image(path)
|
||||
return img.unsqueeze(0)
|
||||
|
||||
|
||||
def _read_visual_frames(path: str) -> torch.Tensor:
|
||||
"""Read a video or a single image as ``(T, C, H, W)`` uint8."""
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
if ext in {".png", ".jpg", ".jpeg", ".webp"}:
|
||||
return _read_image_as_single_frame_video(path)
|
||||
return _read_video_frames(path)
|
||||
|
||||
|
||||
def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
|
||||
"""
|
||||
Compute SSIM between two videos.
|
||||
Compute SSIM between two videos or single-frame image files.
|
||||
|
||||
Image paths (``.png``, ``.jpg``, ``.jpeg``, ``.webp``) are treated as
|
||||
one-frame clips so T2I SSIM can share the same MS-SSIM path as video.
|
||||
|
||||
Args:
|
||||
video1_path: Path to the first video.
|
||||
video2_path: Path to the second video.
|
||||
video1_path: Path to the first video or image.
|
||||
video2_path: Path to the second video or image.
|
||||
use_ms_ssim: Whether to use Multi-Scale Structural Similarity(MS-SSIM) instead of SSIM.
|
||||
"""
|
||||
from pytorch_msssim import ms_ssim, ssim
|
||||
@@ -94,8 +113,8 @@ def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
|
||||
if not os.path.exists(video2_path):
|
||||
raise FileNotFoundError(f"Video2 not found: {video2_path}")
|
||||
|
||||
frames1 = _read_video_frames(video1_path)
|
||||
frames2 = _read_video_frames(video2_path)
|
||||
frames1 = _read_visual_frames(video1_path)
|
||||
frames2 = _read_visual_frames(video2_path)
|
||||
|
||||
# Ensure same number of frames
|
||||
min_frames = min(frames1.shape[0], frames2.shape[0])
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.train.methods.distribution_matching.anyflow import AnyFlowMethod
|
||||
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
|
||||
AnyFlowPretrainMethod, )
|
||||
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
|
||||
from fastvideo.train.methods.distribution_matching.self_forcing import (
|
||||
SelfForcingMethod, )
|
||||
@@ -7,6 +10,8 @@ from fastvideo.train.methods.distribution_matching.streaming_long_tuning import
|
||||
StreamingLongTuningMethod, )
|
||||
|
||||
__all__ = [
|
||||
"AnyFlowMethod",
|
||||
"AnyFlowPretrainMethod",
|
||||
"DMD2Method",
|
||||
"SelfForcingMethod",
|
||||
"StreamingLongTuningMethod",
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""AnyFlow on-policy distillation method.
|
||||
|
||||
Stage 2 of the AnyFlow two-stage recipe. Continues from a pretrained
|
||||
flow-map student and refines it via distribution-matching distillation
|
||||
(DMD2) where the student is rolled out for ``student_sample_steps``
|
||||
Euler-flow steps from pure noise. One randomly-chosen step in the
|
||||
rollout is gradient-enabled (and broadcast across ranks so every worker
|
||||
agrees on which step to gradient-enable); the rest run under
|
||||
``torch.no_grad``.
|
||||
|
||||
Inherits ``DMD2Method`` for the alternating student / critic update
|
||||
machinery and the existing DMD VSD-with-fake-score loss. Overrides
|
||||
``_student_rollout`` to drive the multi-step Euler-flow rollout with
|
||||
``r = t_next`` (mean-velocity sampling — matches the AnyFlow paper's
|
||||
``WanAnyFlowPipeline.training_rollout`` with ``use_mean_velocity=True``).
|
||||
|
||||
Reference: ``pipeline_wan_anyflow.py::training_rollout`` in
|
||||
NVlabs/AnyFlow at commit ``549236a``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Literal
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
|
||||
from fastvideo.train.utils.config import (
|
||||
get_optional_float,
|
||||
get_optional_int,
|
||||
)
|
||||
|
||||
|
||||
class AnyFlowMethod(DMD2Method):
|
||||
"""AnyFlow on-policy distillation (multi-step rollout)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
cfg: Any,
|
||||
role_models: dict[str, Any],
|
||||
) -> None:
|
||||
super().__init__(cfg=cfg, role_models=role_models)
|
||||
mcfg = self.method_config
|
||||
|
||||
student_sample_steps = get_optional_int(mcfg, "student_sample_steps", where="method.student_sample_steps")
|
||||
if student_sample_steps is None:
|
||||
student_sample_steps = 4
|
||||
if int(student_sample_steps) <= 0:
|
||||
raise ValueError("method.student_sample_steps must be positive, "
|
||||
f"got {student_sample_steps}")
|
||||
self._student_sample_steps = int(student_sample_steps)
|
||||
|
||||
use_mean_velocity_raw = mcfg.get("use_mean_velocity", True)
|
||||
if not isinstance(use_mean_velocity_raw, bool):
|
||||
raise ValueError("method.use_mean_velocity must be a bool, "
|
||||
f"got {type(use_mean_velocity_raw).__name__}")
|
||||
self._use_mean_velocity = bool(use_mean_velocity_raw)
|
||||
|
||||
# Optional pinned rollout schedule (descending, absolute t-units).
|
||||
# Falls back to dmd_denoising_steps when absent.
|
||||
raw_t_list = mcfg.get("t_list_override", None)
|
||||
if raw_t_list is None:
|
||||
self._t_list_override: list[float] | None = None
|
||||
else:
|
||||
if not isinstance(raw_t_list, list) or not raw_t_list:
|
||||
raise ValueError("method.t_list_override must be a non-empty list of "
|
||||
f"floats when set, got {raw_t_list!r}")
|
||||
t_list = [float(x) for x in raw_t_list]
|
||||
for i in range(len(t_list) - 1):
|
||||
if t_list[i] < t_list[i + 1]:
|
||||
raise ValueError("method.t_list_override must be descending, "
|
||||
f"got {t_list!r}")
|
||||
self._t_list_override = t_list
|
||||
|
||||
# Scoring conditioning: AnyFlow scores against r=0 for the DMD branch.
|
||||
score_r_raw = mcfg.get("dmd_score_r_value", 0.0)
|
||||
try:
|
||||
self._dmd_score_r = float(score_r_raw)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("method.dmd_score_r_value must be numeric, "
|
||||
f"got {score_r_raw!r}") from exc
|
||||
|
||||
# Optional teacher guidance scale for the DMD loss (carry over from
|
||||
# DMD2Method's behavior; default 1.0).
|
||||
guidance = get_optional_float(mcfg, "real_score_guidance_scale", where="method.real_score_guidance_scale")
|
||||
self._real_score_guidance = float(guidance) if guidance is not None else 1.0
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Rollout schedule
|
||||
|
||||
def _get_rollout_schedule(self, *, device: torch.device) -> torch.Tensor:
|
||||
"""Build the descending timestep schedule used by the on-policy
|
||||
rollout. Length is ``num_steps + 1`` so ``num_steps`` Euler steps
|
||||
consume the full range.
|
||||
|
||||
Order of precedence:
|
||||
1. ``method.t_list_override`` — used verbatim (absolute units).
|
||||
2. ``method.dmd_denoising_steps`` (inherited from DMD2) appended
|
||||
with a final 0 boundary if the last entry isn't already 0.
|
||||
"""
|
||||
if self._t_list_override is not None:
|
||||
return torch.tensor(self._t_list_override, device=device, dtype=torch.float32)
|
||||
|
||||
steps = self._get_denoising_step_list(device).to(dtype=torch.float32)
|
||||
if float(steps[-1].item()) != 0.0:
|
||||
zero = torch.zeros(1, device=device, dtype=torch.float32)
|
||||
steps = torch.cat([steps, zero], dim=0)
|
||||
return steps
|
||||
|
||||
def _broadcast_grad_step_index(
|
||||
self,
|
||||
num_steps: int,
|
||||
*,
|
||||
device: torch.device,
|
||||
) -> int:
|
||||
"""Pick the rollout step that gets gradient enabled. In distributed
|
||||
runs the choice is broadcast from rank 0 so every worker agrees."""
|
||||
if num_steps <= 0:
|
||||
raise ValueError("num_steps must be positive")
|
||||
if dist.is_initialized() and dist.get_rank() != 0:
|
||||
idx_tensor = torch.empty((1, ), dtype=torch.long, device=device)
|
||||
else:
|
||||
idx_tensor = torch.randint(0,
|
||||
num_steps, (1, ),
|
||||
device=device,
|
||||
dtype=torch.long,
|
||||
generator=self.cuda_generator)
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(idx_tensor, src=0)
|
||||
return int(idx_tensor.item())
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Rollout
|
||||
|
||||
def _student_rollout(
|
||||
self,
|
||||
batch: Any,
|
||||
*,
|
||||
with_grad: bool,
|
||||
) -> torch.Tensor:
|
||||
"""Multi-step Euler-flow rollout from pure noise.
|
||||
|
||||
Returns the predicted clean latent ``x_0`` after the chosen
|
||||
gradient step (or the final ``x`` after the last step if
|
||||
``with_grad`` is False — used by the critic path).
|
||||
"""
|
||||
latents = batch.latents
|
||||
if latents is None or latents.ndim != 5:
|
||||
raise RuntimeError("AnyFlow on-policy rollout requires TrainingBatch.latents "
|
||||
"of shape [B, T, C, H, W] for shape templating")
|
||||
device = latents.device
|
||||
dtype = latents.dtype
|
||||
|
||||
schedule = self._get_rollout_schedule(device=device)
|
||||
num_entries = int(schedule.numel())
|
||||
num_steps = num_entries - 1
|
||||
if num_steps <= 0:
|
||||
raise RuntimeError("rollout schedule must have at least two entries "
|
||||
f"(got {num_entries})")
|
||||
if num_steps > self._student_sample_steps:
|
||||
# Trim to the configured cap, keeping the last (=0) boundary.
|
||||
schedule = torch.cat([schedule[:self._student_sample_steps], schedule[-1:]], dim=0)
|
||||
num_steps = self._student_sample_steps
|
||||
|
||||
grad_step = self._broadcast_grad_step_index(num_steps, device=device) if with_grad else -1
|
||||
|
||||
attn_kind: Literal["dense", "vsa"] = "vsa"
|
||||
n_train = float(self.student.num_train_timesteps)
|
||||
|
||||
x = torch.randn(latents.shape, device=device, dtype=dtype, generator=self.cuda_generator)
|
||||
last_pred_x0: torch.Tensor | None = None
|
||||
batch_size = int(latents.shape[0])
|
||||
|
||||
for i in range(num_steps):
|
||||
t_cur = schedule[i].expand(batch_size)
|
||||
t_next = schedule[i + 1].expand(batch_size)
|
||||
r = t_next if self._use_mean_velocity else t_cur
|
||||
|
||||
enable_grad = bool(with_grad) and (i == grad_step)
|
||||
with torch.set_grad_enabled(enable_grad):
|
||||
v = self.student.predict_velocity_with_r(
|
||||
x,
|
||||
t_cur,
|
||||
r,
|
||||
batch,
|
||||
conditional=True,
|
||||
cfg_uncond=self._cfg_uncond,
|
||||
attn_kind=attn_kind,
|
||||
)
|
||||
# Euler step in absolute units: x ← x - ((t_cur - t_next) / N) * v.
|
||||
view = [-1] + [1] * (x.ndim - 1)
|
||||
dt = ((t_cur - t_next) / n_train).view(*view)
|
||||
x = x - dt * v
|
||||
|
||||
if enable_grad:
|
||||
# We treat the rollout output (post-step) as a predicted
|
||||
# clean latent — AnyFlow's last Euler step lands at t=0.
|
||||
last_pred_x0 = x
|
||||
|
||||
if last_pred_x0 is None:
|
||||
# No gradient step taken (with_grad=False path).
|
||||
last_pred_x0 = x
|
||||
|
||||
if hasattr(batch, "dmd_latent_vis_dict"):
|
||||
batch.dmd_latent_vis_dict["generator_timestep"] = (schedule[-1].detach().clone())
|
||||
return last_pred_x0
|
||||
@@ -0,0 +1,404 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""AnyFlow pretrain (flow-map central-difference) training method.
|
||||
|
||||
Stage 1 of the AnyFlow two-stage recipe. Trains a single student network
|
||||
``u_θ(x_t, t, r)`` to predict the average velocity from time ``t`` back
|
||||
to time ``r`` via the central-difference target
|
||||
|
||||
target = (eps - x_0) - ((t - r) / N) * dF/dt
|
||||
|
||||
where ``N = num_train_timesteps`` and ``dF/dt`` is estimated from the
|
||||
student's own forward at ``(t ± δ, r)`` (with one-sided fallback near the
|
||||
schedule endpoints).
|
||||
|
||||
Per-batch ``(t, r)`` sampling follows the AnyFlow paper:
|
||||
|
||||
- ``diffusion_ratio`` fraction: ``r = t`` (recovers plain flow matching).
|
||||
- ``consistency_ratio`` fraction: ``r = 0`` (consistency to clean data).
|
||||
- Remaining fraction: ``(t, r) = (max, min)`` of two independent uniform
|
||||
draws (full reconstruction range).
|
||||
|
||||
Reference: ``trainer_wan_anyflow_pretrain.py`` in NVlabs/AnyFlow at
|
||||
commit ``549236a``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections.abc import Sequence
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.base import LogScalar, TrainingMethod
|
||||
from fastvideo.train.models.base import ModelBase
|
||||
from fastvideo.train.utils.config import (
|
||||
get_optional_float,
|
||||
get_optional_int,
|
||||
)
|
||||
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
|
||||
|
||||
|
||||
def _sample_pair_timesteps(
|
||||
*,
|
||||
batch_size: int,
|
||||
diffusion_ratio: float,
|
||||
consistency_ratio: float,
|
||||
device: torch.device,
|
||||
generator: torch.Generator | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Sample ``(t, r)`` per the AnyFlow paper.
|
||||
|
||||
Two uniform draws ``u1, u2 ∈ [0, 1]`` per sample, then
|
||||
``t = max(u1, u2)`` and ``r = min(u1, u2)``. After the base sample,
|
||||
the first ``diffusion_ratio * B`` entries get ``r = t`` (diffusion
|
||||
branch, plain flow matching), and the next ``consistency_ratio * B``
|
||||
entries get ``r = 0`` (consistency branch).
|
||||
|
||||
Returns
|
||||
-------
|
||||
t, r, is_diffusion, is_consistency
|
||||
Each tensor has shape ``(batch_size,)``. ``t`` and ``r`` are in
|
||||
``[0, 1]`` (i.e. *not* yet shifted and *not* yet in absolute
|
||||
train-timestep units). ``is_diffusion`` and ``is_consistency``
|
||||
are bool masks that partition a subset of the batch — entries
|
||||
outside both masks are the "free" reconstruction fraction.
|
||||
"""
|
||||
if batch_size <= 0:
|
||||
raise ValueError(f"batch_size must be positive, got {batch_size}")
|
||||
if diffusion_ratio < 0.0 or consistency_ratio < 0.0:
|
||||
raise ValueError("diffusion_ratio and consistency_ratio must be non-negative")
|
||||
if diffusion_ratio + consistency_ratio > 1.0:
|
||||
raise ValueError("diffusion_ratio + consistency_ratio must be <= 1, "
|
||||
f"got {diffusion_ratio} + {consistency_ratio}")
|
||||
|
||||
u1 = torch.rand(batch_size, device=device, generator=generator)
|
||||
u2 = torch.rand(batch_size, device=device, generator=generator)
|
||||
t = torch.maximum(u1, u2)
|
||||
r = torch.minimum(u1, u2)
|
||||
|
||||
n_diff = int(diffusion_ratio * batch_size)
|
||||
n_cons = int(consistency_ratio * batch_size)
|
||||
is_diffusion = torch.zeros(batch_size, dtype=torch.bool, device=device)
|
||||
is_consistency = torch.zeros(batch_size, dtype=torch.bool, device=device)
|
||||
is_diffusion[:n_diff] = True
|
||||
is_consistency[n_diff:n_diff + n_cons] = True
|
||||
|
||||
# Override per the AnyFlow paper:
|
||||
# - diffusion entries: r = t (plain flow matching)
|
||||
# - consistency entries: r = 0 (consistency to clean data)
|
||||
r = torch.where(is_diffusion, t, r)
|
||||
r = torch.where(is_consistency, torch.zeros_like(r), r)
|
||||
return t, r, is_diffusion, is_consistency
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _central_difference_dF_dt(
|
||||
*,
|
||||
student: Any,
|
||||
batch: Any,
|
||||
noisy: torch.Tensor,
|
||||
latents: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
t: torch.Tensor,
|
||||
r: torch.Tensor,
|
||||
delta: float,
|
||||
num_train_timesteps: float,
|
||||
attn_kind: str = "dense",
|
||||
guidance_scale: float = 1.0,
|
||||
) -> torch.Tensor:
|
||||
"""Estimate ``dF/dt`` for the AnyFlow central-difference target.
|
||||
|
||||
Computes a symmetric finite difference of the velocity prediction in
|
||||
*absolute train-timestep units*:
|
||||
|
||||
dF/dt ≈ [u_θ(x_{t+δ}, t+δ, r) - u_θ(x_{t-δ}, t-δ, r)] / (2 * δ * guidance)
|
||||
|
||||
The sample is also moved along the flow trajectory by the same
|
||||
finite step (``v_pred * (δ / N)``) to match AnyFlow's reference
|
||||
formulation in ``trainer_wan_anyflow_pretrain.py::compute_central_difference``.
|
||||
Wrapped in ``torch.no_grad`` so the two extra forwards never enter
|
||||
the backward graph.
|
||||
"""
|
||||
if delta <= 0.0:
|
||||
raise ValueError(f"delta must be positive, got {delta}")
|
||||
if guidance_scale <= 0.0:
|
||||
raise ValueError(f"guidance_scale must be positive, got {guidance_scale}")
|
||||
|
||||
v_pred = noise - latents # ground-truth flow velocity
|
||||
delta_x = delta / float(num_train_timesteps)
|
||||
|
||||
t_plus = t + delta
|
||||
noisy_plus = noisy + v_pred * delta_x
|
||||
f_plus = student.predict_velocity_with_r(noisy_plus, t_plus, r, batch, conditional=True, attn_kind=attn_kind)
|
||||
|
||||
t_minus = t - delta
|
||||
noisy_minus = noisy - v_pred * delta_x
|
||||
f_minus = student.predict_velocity_with_r(noisy_minus, t_minus, r, batch, conditional=True, attn_kind=attn_kind)
|
||||
|
||||
return (f_plus - f_minus) / (2.0 * delta * guidance_scale)
|
||||
|
||||
|
||||
class AnyFlowPretrainMethod(TrainingMethod):
|
||||
"""AnyFlow flow-map pretrain method.
|
||||
|
||||
Single-student training; no teacher or critic. The student must
|
||||
implement ``predict_velocity_with_r(noisy, t, r, batch, ...)`` —
|
||||
typically a ``WanModel`` with ``r_embedder=True`` in its arch config.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
cfg: Any,
|
||||
role_models: dict[str, ModelBase],
|
||||
) -> None:
|
||||
super().__init__(cfg=cfg, role_models=role_models)
|
||||
if "student" not in role_models:
|
||||
raise ValueError("AnyFlowPretrainMethod requires role 'student'")
|
||||
if not self.student._trainable:
|
||||
raise ValueError("AnyFlowPretrainMethod requires student to be trainable")
|
||||
|
||||
mcfg = self.method_config
|
||||
self._diffusion_ratio = float(
|
||||
get_optional_float(mcfg, "diffusion_ratio", where="method.diffusion_ratio") or 0.5)
|
||||
self._consistency_ratio = float(
|
||||
get_optional_float(mcfg, "consistency_ratio", where="method.consistency_ratio") or 0.25)
|
||||
if self._diffusion_ratio + self._consistency_ratio > 1.0:
|
||||
raise ValueError("method.diffusion_ratio + method.consistency_ratio must "
|
||||
f"be <= 1, got {self._diffusion_ratio} + "
|
||||
f"{self._consistency_ratio}")
|
||||
|
||||
# δ: finite-difference step in absolute train-timestep units.
|
||||
epsilon = get_optional_int(mcfg, "epsilon", where="method.epsilon")
|
||||
self._fd_epsilon = float(epsilon) if epsilon is not None else 5.0
|
||||
|
||||
# Loss weighting scheme (uniform / gaussian / beta08).
|
||||
raw_weight_type = mcfg.get("weight_type", "beta08")
|
||||
if not isinstance(raw_weight_type, str):
|
||||
raise ValueError("method.weight_type must be a string, got "
|
||||
f"{type(raw_weight_type).__name__}")
|
||||
weight_type = raw_weight_type.strip().lower()
|
||||
if weight_type not in {"uniform", "gaussian", "beta08"}:
|
||||
raise ValueError("method.weight_type must be one of "
|
||||
"{uniform, gaussian, beta08}, "
|
||||
f"got {raw_weight_type!r}")
|
||||
self._weight_type = weight_type
|
||||
|
||||
# Guidance fused into the training target (default 1.0 = unused).
|
||||
fg = get_optional_float(mcfg, "fuse_guidance_scale", where="method.fuse_guidance_scale")
|
||||
self._fuse_guidance_scale = float(fg) if fg is not None else 1.0
|
||||
if self._fuse_guidance_scale <= 0.0:
|
||||
raise ValueError("method.fuse_guidance_scale must be positive, "
|
||||
f"got {self._fuse_guidance_scale}")
|
||||
|
||||
# Flow-map scheduler — uses pipeline_config.flow_shift if present
|
||||
# and falls back to method.shift (and finally 1.0).
|
||||
shift = float(getattr(self.training_config.pipeline_config, "flow_shift", 0.0) or 0.0)
|
||||
if shift <= 0.0:
|
||||
shift_override = get_optional_float(mcfg, "shift", where="method.shift")
|
||||
shift = float(shift_override) if shift_override is not None else 1.0
|
||||
self._shift = shift
|
||||
|
||||
# Lazy-imported to avoid circular imports on package load.
|
||||
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
|
||||
FlowMapEulerDiscreteScheduler, )
|
||||
self._flow_map_scheduler = FlowMapEulerDiscreteScheduler(
|
||||
num_train_timesteps=int(self.student.num_train_timesteps),
|
||||
shift=self._shift,
|
||||
)
|
||||
|
||||
self.student.init_preprocessors(self.training_config)
|
||||
self._init_optimizer_and_scheduler()
|
||||
|
||||
@property
|
||||
def _optimizer_dict(self) -> dict[str, torch.optim.Optimizer]:
|
||||
return {"student": self._student_optimizer}
|
||||
|
||||
@property
|
||||
def _lr_scheduler_dict(self) -> dict[str, Any]:
|
||||
return {"student": self._student_lr_scheduler}
|
||||
|
||||
def get_optimizers(
|
||||
self,
|
||||
iteration: int,
|
||||
) -> Sequence[torch.optim.Optimizer]:
|
||||
del iteration
|
||||
return [self._student_optimizer]
|
||||
|
||||
def get_lr_schedulers(self, iteration: int) -> Sequence[Any]:
|
||||
del iteration
|
||||
return [self._student_lr_scheduler]
|
||||
|
||||
def single_train_step(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
iteration: int,
|
||||
) -> tuple[
|
||||
dict[str, torch.Tensor],
|
||||
dict[str, Any],
|
||||
dict[str, LogScalar],
|
||||
]:
|
||||
del iteration # AnyFlow pretrain has no iteration-dependent dispatch.
|
||||
|
||||
training_batch = self.student.prepare_batch(
|
||||
batch,
|
||||
generator=self.cuda_generator,
|
||||
latents_source="data",
|
||||
)
|
||||
latents = training_batch.latents # [B, T, C, H, W] (post-permute in prepare_batch).
|
||||
if latents is None or latents.ndim != 5:
|
||||
raise RuntimeError("AnyFlow pretrain expects TrainingBatch.latents of shape "
|
||||
"[B, T, C, H, W] after prepare_batch; got "
|
||||
f"{None if latents is None else tuple(latents.shape)}")
|
||||
device = latents.device
|
||||
dtype = latents.dtype
|
||||
batch_size = int(latents.shape[0])
|
||||
|
||||
# AnyFlow (t, r) sampling — overrides the timestep drawn by
|
||||
# WanModel._sample_timesteps inside prepare_batch.
|
||||
t_norm, r_norm, is_diffusion, is_consistency = _sample_pair_timesteps(
|
||||
batch_size=batch_size,
|
||||
diffusion_ratio=self._diffusion_ratio,
|
||||
consistency_ratio=self._consistency_ratio,
|
||||
device=device,
|
||||
generator=self.cuda_generator,
|
||||
)
|
||||
|
||||
sched = self._flow_map_scheduler
|
||||
n_train = float(self.student.num_train_timesteps)
|
||||
t = (sched.apply_shift(t_norm) * n_train).to(device=device, dtype=dtype)
|
||||
r = (sched.apply_shift(r_norm) * n_train).to(device=device, dtype=dtype)
|
||||
|
||||
# Fresh noise drawn from the method's RNG; ignore the noise that
|
||||
# prepare_batch attached (it pairs with the discarded timestep).
|
||||
noise = torch.randn(
|
||||
latents.shape,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
generator=self.cuda_generator,
|
||||
)
|
||||
noisy = sched.add_noise(latents, noise, t)
|
||||
|
||||
# Keep training_batch coherent with the new (t, noisy): downstream
|
||||
# forward_context uses these fields.
|
||||
training_batch.timesteps = t
|
||||
training_batch.noise = noise
|
||||
training_batch.noisy_model_input = noisy.permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Student velocity prediction at (t, r).
|
||||
noise_pred = self.student.predict_velocity_with_r(
|
||||
noisy,
|
||||
t,
|
||||
r,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind="dense",
|
||||
)
|
||||
|
||||
# Optional guidance distillation — fuse CFG into the training target so
|
||||
# the resulting checkpoint can be sampled at guidance_scale=1.0.
|
||||
if self._fuse_guidance_scale != 1.0:
|
||||
with torch.no_grad():
|
||||
noise_pred_uncond = self.student.predict_velocity_with_r(
|
||||
noisy,
|
||||
t,
|
||||
r,
|
||||
training_batch,
|
||||
conditional=False,
|
||||
attn_kind="dense",
|
||||
)
|
||||
g = float(self._fuse_guidance_scale)
|
||||
noise_pred = (noise_pred - (1.0 - g) * noise_pred_uncond) / g
|
||||
|
||||
dF_dt = _central_difference_dF_dt(
|
||||
student=self.student,
|
||||
batch=training_batch,
|
||||
noisy=noisy,
|
||||
latents=latents,
|
||||
noise=noise,
|
||||
t=t,
|
||||
r=r,
|
||||
delta=self._fd_epsilon,
|
||||
num_train_timesteps=n_train,
|
||||
attn_kind="dense",
|
||||
guidance_scale=self._fuse_guidance_scale,
|
||||
)
|
||||
|
||||
# AnyFlow target: target = (eps - x_0) - (t - r) * dF/dt
|
||||
# dF/dt is in (velocity per absolute t-unit); (t - r) is in absolute units;
|
||||
# the product cancels back to velocity units, matching noise_pred.
|
||||
view = [batch_size] + [1] * (latents.ndim - 1)
|
||||
target = (noise - latents) - (t - r).view(*view) * dF_dt
|
||||
|
||||
# Per-sample squared error, then per-timestep weight, then scale-balance
|
||||
# so the non-diffusion branches stay on the same magnitude as the
|
||||
# diffusion branch (matches AnyFlow's stop-grad rescaling).
|
||||
per_sample = torch.mean(
|
||||
((noise_pred.float() - target.float())**2).reshape(batch_size, -1),
|
||||
dim=-1,
|
||||
)
|
||||
weight = sched.get_train_weight(t, weight_type=self._weight_type)
|
||||
per_sample = per_sample * weight
|
||||
|
||||
with torch.no_grad():
|
||||
diff_mask = is_diffusion
|
||||
diff_mean = per_sample[diff_mask].mean() if diff_mask.any() else per_sample.mean()
|
||||
non_diff_mask = ~diff_mask
|
||||
if non_diff_mask.any():
|
||||
scale = diff_mean / (per_sample[non_diff_mask] + 1e-5)
|
||||
else:
|
||||
scale = torch.tensor(1.0, device=device, dtype=per_sample.dtype)
|
||||
non_diff_idx = torch.nonzero(non_diff_mask, as_tuple=False).flatten()
|
||||
if non_diff_idx.numel() > 0:
|
||||
per_sample = per_sample.clone()
|
||||
per_sample[non_diff_idx] = per_sample[non_diff_idx] * scale
|
||||
|
||||
total_loss = per_sample.mean()
|
||||
|
||||
loss_map = {"total_loss": total_loss}
|
||||
metrics: dict[str, LogScalar] = {
|
||||
"diffusion_fraction": float(is_diffusion.float().mean()),
|
||||
"consistency_fraction": float(is_consistency.float().mean()),
|
||||
"scale_weight_mean": float(scale.mean()) if isinstance(scale, torch.Tensor) else float(scale),
|
||||
}
|
||||
outputs = {
|
||||
"student_ctx": (
|
||||
training_batch.timesteps,
|
||||
training_batch.attn_metadata,
|
||||
),
|
||||
}
|
||||
return loss_map, outputs, metrics
|
||||
|
||||
def backward(
|
||||
self,
|
||||
loss_map: dict[str, torch.Tensor],
|
||||
outputs: dict[str, Any],
|
||||
*,
|
||||
grad_accum_rounds: int = 1,
|
||||
) -> None:
|
||||
"""Route the loss backward through the student's forward_context
|
||||
so attn metadata stays attached during gradient computation."""
|
||||
student_ctx = outputs.get("student_ctx")
|
||||
if student_ctx is None:
|
||||
super().backward(loss_map, outputs, grad_accum_rounds=grad_accum_rounds)
|
||||
return
|
||||
self.student.backward(
|
||||
loss_map["total_loss"],
|
||||
student_ctx,
|
||||
grad_accum_rounds=grad_accum_rounds,
|
||||
)
|
||||
|
||||
def _init_optimizer_and_scheduler(self) -> None:
|
||||
tc = self.training_config
|
||||
params = [p for p in self.student.transformer.parameters() if p.requires_grad]
|
||||
(
|
||||
self._student_optimizer,
|
||||
self._student_lr_scheduler,
|
||||
) = build_optimizer_and_scheduler(
|
||||
params=params,
|
||||
optimizer_config=tc.optimizer,
|
||||
loop_config=tc.loop,
|
||||
learning_rate=float(tc.optimizer.learning_rate),
|
||||
betas=tc.optimizer.betas,
|
||||
scheduler_name=str(tc.optimizer.lr_scheduler),
|
||||
)
|
||||
@@ -358,6 +358,53 @@ class WanModel(ModelBase):
|
||||
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
return pred_noise
|
||||
|
||||
def predict_velocity_with_r(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
r_timestep: torch.Tensor,
|
||||
batch: TrainingBatch,
|
||||
*,
|
||||
conditional: bool,
|
||||
cfg_uncond: dict[str, Any] | None = None,
|
||||
attn_kind: Literal["dense", "vsa"] = "dense",
|
||||
) -> torch.Tensor:
|
||||
"""AnyFlow forward: predict average velocity from ``t`` back to ``r``.
|
||||
|
||||
Same plumbing as :meth:`predict_noise` but injects ``r_timestep``
|
||||
into the transformer kwargs. The transformer must have been
|
||||
constructed with an arch config that sets ``r_embedder=True`` for
|
||||
the dual-timestep branch to be active — otherwise ``r_timestep``
|
||||
is silently ignored by the embedder and the forward reduces to
|
||||
the single-timestep path.
|
||||
"""
|
||||
device_type = self.device.type
|
||||
dtype = noisy_latents.dtype
|
||||
if conditional:
|
||||
text_dict = batch.conditional_dict
|
||||
if text_dict is None:
|
||||
raise RuntimeError("Missing conditional_dict in "
|
||||
"TrainingBatch")
|
||||
else:
|
||||
text_dict = self._get_uncond_text_dict(batch, cfg_uncond=cfg_uncond)
|
||||
|
||||
if attn_kind == "dense":
|
||||
attn_metadata = batch.attn_metadata
|
||||
elif attn_kind == "vsa":
|
||||
attn_metadata = batch.attn_metadata_vsa
|
||||
else:
|
||||
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
|
||||
|
||||
with torch.autocast(device_type, dtype=dtype), set_forward_context(
|
||||
current_timestep=batch.timesteps,
|
||||
attn_metadata=attn_metadata,
|
||||
):
|
||||
input_kwargs = (self._build_distill_input_kwargs(noisy_latents, timestep, text_dict))
|
||||
input_kwargs["r_timestep"] = r_timestep
|
||||
transformer = self._get_transformer(timestep)
|
||||
pred_velocity = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
return pred_velocity
|
||||
|
||||
def backward(
|
||||
self,
|
||||
loss: torch.Tensor,
|
||||
|
||||
@@ -182,6 +182,7 @@ nav:
|
||||
- Distillation:
|
||||
- Data Preprocessing: distillation/data_preprocess.md
|
||||
- DMD: distillation/dmd.md
|
||||
- AnyFlow: distillation/anyflow.md
|
||||
- Attention:
|
||||
- Overview: attention/index.md
|
||||
- Video Sparse Attention: attention/vsa/index.md
|
||||
|
||||
+4
-2
@@ -21,7 +21,9 @@ dependencies = [
|
||||
"requests>=2.32.2",
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.57.3",
|
||||
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
|
||||
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
|
||||
"transformers>=5.0.0",
|
||||
# <0.23: tokenizers 0.23 renamed RobertaProcessing's binding args, so
|
||||
# transformers' CLIP-style tokenizer loading dies with
|
||||
# "RobertaProcessing.__new__() got an unexpected keyword argument 'cls'".
|
||||
@@ -227,7 +229,7 @@ skip = "./data,./wandb,apps/fastvideo_studio/package-lock.json,apps/performance_
|
||||
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
|
||||
# Early Drop). codespell lowercases ignore-words entries, so the single
|
||||
# lowercase form silences all case variants.
|
||||
ignore-words-list = "tread,passt"
|
||||
ignore-words-list = "tread,passt,dout"
|
||||
|
||||
[tool.ruff]
|
||||
# Allow lines to be as long as 120.
|
||||
|
||||
@@ -21,7 +21,9 @@ dependencies = [
|
||||
"requests>=2.32.2",
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.57.3",
|
||||
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
|
||||
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
|
||||
"transformers>=5.0.0",
|
||||
"tokenizers>=0.20.1",
|
||||
"sentencepiece>=0.2.0",
|
||||
"timm>=1.0.11",
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 checkpoint strict-load verifier (no weight conversion required).
|
||||
|
||||
The published ``nvidia/Cosmos3-Nano`` checkpoint is diffusers-format and its
|
||||
transformer weight keys map 1:1 (identity) onto FastVideo's native
|
||||
``Cosmos3VFMTransformer`` parameters -- ``needs_conversion=no``. There is no
|
||||
remap to apply; the checkpoint loads directly.
|
||||
|
||||
This utility verifies strict-load completeness (every checkpoint key has a
|
||||
matching DiT parameter of the right shape, and every DiT parameter is provided
|
||||
by the checkpoint) without allocating the full ~30 GB model, by reading
|
||||
safetensors headers and instantiating the DiT on the ``meta`` device.
|
||||
|
||||
Usage:
|
||||
python scripts/checkpoint_conversion/cosmos3_convert.py \
|
||||
--transformer official_weights/cosmos3/transformer
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import os
|
||||
import re
|
||||
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
|
||||
from fastvideo.configs.models.dits.cosmos3 import Cosmos3VideoConfig
|
||||
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
|
||||
|
||||
|
||||
def checkpoint_key_shapes(transformer_dir: str) -> dict[str, tuple[int, ...]]:
|
||||
"""Read ``{key: shape}`` from a sharded safetensors transformer dir."""
|
||||
shards = sorted(glob.glob(os.path.join(transformer_dir, "*.safetensors")))
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"no .safetensors found in {transformer_dir}")
|
||||
shapes: dict[str, tuple[int, ...]] = {}
|
||||
for shard in shards:
|
||||
with safe_open(shard, framework="pt") as handle:
|
||||
for key in handle.keys():
|
||||
shapes[key] = tuple(handle.get_slice(key).get_shape())
|
||||
return shapes
|
||||
|
||||
|
||||
def verify_strict_load(transformer_dir: str) -> None:
|
||||
"""Raise SystemExit if the checkpoint does not strict-load into the DiT."""
|
||||
ckpt = checkpoint_key_shapes(transformer_dir)
|
||||
cfg = Cosmos3VideoConfig()
|
||||
with torch.device("meta"):
|
||||
dit = Cosmos3VFMTransformer(cfg, hf_config={})
|
||||
params = {name: tuple(p.shape) for name, p in dit.named_parameters()}
|
||||
buffers = {name for name, _ in dit.named_buffers()}
|
||||
name_map: dict[str, str] = cfg.arch_config.param_names_mapping
|
||||
|
||||
def remap(key: str) -> str:
|
||||
for pattern, replacement in name_map.items():
|
||||
if re.match(pattern, key):
|
||||
return re.sub(pattern, replacement, key)
|
||||
return key
|
||||
|
||||
mapped = {remap(key): shape for key, shape in ckpt.items()}
|
||||
unexpected = sorted(set(mapped) - set(params) - buffers)
|
||||
missing = sorted(set(params) - set(mapped))
|
||||
mismatched = [(k, mapped[k], params[k]) for k in (set(mapped) & set(params)) if mapped[k] != params[k]]
|
||||
|
||||
if unexpected or missing or mismatched:
|
||||
raise SystemExit("strict-load FAILED: "
|
||||
f"unexpected={unexpected[:10]} missing={missing[:10]} "
|
||||
f"shape_mismatch={mismatched[:10]}")
|
||||
print(f"strict-load OK: {len(ckpt)} checkpoint keys map 1:1 onto "
|
||||
f"{len(params)} DiT params (identity; needs_conversion=no)")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--transformer",
|
||||
default=os.path.join("official_weights", "cosmos3", "transformer"),
|
||||
help="path to the checkpoint transformer/ directory",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
verify_strict_load(args.transformer)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,282 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-side AnyFlow 14B T2V demo at NFE=4 and NFE=50.
|
||||
|
||||
Loads ``nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers`` into FastVideo's
|
||||
``WanTransformer3DModel`` (with the ``param_names_mapping`` regex
|
||||
handling the ``delta_embedder`` rename), runs the
|
||||
``FlowMapEulerDiscreteScheduler`` for the requested NFE schedule, and
|
||||
saves the decoded video as MP4. Matches the prompt / shift / guidance
|
||||
recipe used by the parallel FastGen demo so the videos are directly
|
||||
comparable.
|
||||
|
||||
Memory tactics (single H200, 141 GB HBM):
|
||||
- Encode prompts with UMT5, free the encoder.
|
||||
- Build the FastVideo Wan-14B transformer, load AnyFlow safetensor
|
||||
shards via param_names_mapping translation.
|
||||
- Sample at both NFEs without re-loading the transformer.
|
||||
- Free the transformer, then load the Wan VAE with tiling for decode.
|
||||
|
||||
Configure local checkout paths via env vars (defaults assume a sibling
|
||||
layout next to this repo):
|
||||
|
||||
ANYFLOW_LOCAL — path to ``nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers``
|
||||
(default ``./anyflow-14b``)
|
||||
ANYFLOW_DEMO_OUT — output directory for the rendered MP4s
|
||||
(default ``./demo_videos``)
|
||||
|
||||
Run via::
|
||||
|
||||
PYTHONPATH=$PWD python scripts/demo_anyflow_14b.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
ANYFLOW_LOCAL = Path(os.environ.get("ANYFLOW_LOCAL", "./anyflow-14b")).expanduser()
|
||||
OUT_DIR = Path(os.environ.get("ANYFLOW_DEMO_OUT", "./demo_videos")).expanduser()
|
||||
OUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
SEED = 0
|
||||
DEVICE = torch.device("cuda")
|
||||
DTYPE = torch.bfloat16
|
||||
NUM_FRAMES = 81
|
||||
HEIGHT, WIDTH = 480, 832
|
||||
PROMPT = (
|
||||
"CG game concept digital art, a majestic elephant with a vibrant tusk and sleek fur "
|
||||
"running swiftly towards a herd of its kind. The elephant has a calm yet determined "
|
||||
"expression, with its ears flapping slightly as it moves at high speed. The herd consists "
|
||||
"of several other elephants of various ages and sizes, all moving in unison. The landscape "
|
||||
"is vast savanna with rolling hills, tall grasses, and scattered acacia trees. The sun "
|
||||
"sets behind the horizon, casting a warm golden glow over the scene. Low-angle view, focus "
|
||||
"on the elephant as it accelerates towards the herd."
|
||||
)
|
||||
NEG_PROMPT = "blurry, low quality, distorted"
|
||||
# The published nvidia/AnyFlow-* checkpoints are on-policy distilled with
|
||||
# fuse_guidance_scale=3.0 baked into the weights — inference uses 1.0
|
||||
# (single conditional forward, no CFG; matches AnyFlow's official demo.py).
|
||||
GUIDANCE = 1.0
|
||||
|
||||
|
||||
def banner(msg: str) -> None:
|
||||
print("\n" + "=" * 80)
|
||||
print(msg)
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
def free_gpu() -> None:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
used = torch.cuda.memory_allocated() / 1e9
|
||||
print(f" [mem] allocated {used:.1f} GB after free")
|
||||
|
||||
|
||||
def init_single_rank() -> None:
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29571")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
from fastvideo.distributed import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="nccl")
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
sequence_model_parallel_size=1,
|
||||
data_parallel_size=1,
|
||||
)
|
||||
|
||||
|
||||
def translate_keys(raw: dict, *, mapping: dict[str, str]) -> dict:
|
||||
out: dict = {}
|
||||
for k, v in raw.items():
|
||||
new_k = k
|
||||
for pat, repl in mapping.items():
|
||||
new_k = re.sub(pat, repl, new_k)
|
||||
out[new_k] = v
|
||||
return out
|
||||
|
||||
|
||||
def encode_prompts():
|
||||
banner("(1) Encode prompts via UMT5")
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(str(ANYFLOW_LOCAL), subfolder="tokenizer", use_fast=False)
|
||||
enc = UMT5EncoderModel.from_pretrained(
|
||||
str(ANYFLOW_LOCAL), subfolder="text_encoder", torch_dtype=DTYPE,
|
||||
).to(DEVICE).eval()
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_one(prompts):
|
||||
out = tok(
|
||||
prompts, padding="max_length", max_length=512, truncation=True,
|
||||
return_attention_mask=True, return_tensors="pt")
|
||||
ids = out.input_ids.to(DEVICE)
|
||||
mask = out.attention_mask.to(DEVICE)
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
embeds = enc(ids, mask).last_hidden_state
|
||||
padded = []
|
||||
for i, l in enumerate(seq_lens):
|
||||
e = embeds[i, :l]
|
||||
pad = torch.zeros(512 - e.size(0), e.size(1), device=DEVICE, dtype=embeds.dtype)
|
||||
padded.append(torch.cat([e, pad], dim=0))
|
||||
return torch.stack(padded, dim=0)
|
||||
|
||||
text_e = encode_one([PROMPT])
|
||||
neg_e = encode_one([NEG_PROMPT])
|
||||
print(f" prompts encoded: text={tuple(text_e.shape)} neg={tuple(neg_e.shape)}")
|
||||
del enc, tok
|
||||
free_gpu()
|
||||
return text_e, neg_e
|
||||
|
||||
|
||||
def load_transformer():
|
||||
banner("(2) Build FastVideo Wan-14B + load AnyFlow weights")
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
|
||||
|
||||
cfg = WanVideoConfig()
|
||||
arch = cfg.arch_config
|
||||
# Wan2.1-T2V-14B arch (per AnyFlow checkpoint config.json).
|
||||
arch.num_attention_heads = 40
|
||||
arch.attention_head_dim = 128
|
||||
arch.num_layers = 40
|
||||
arch.ffn_dim = 13824
|
||||
arch.r_embedder = True
|
||||
arch.r_embedder_fusion = "gated"
|
||||
arch.r_embedder_gate_value = 0.25
|
||||
arch.r_embedder_deltatime_type = "r"
|
||||
arch.__post_init__()
|
||||
|
||||
t0 = time.time()
|
||||
model = WanTransformer3DModel(config=cfg, hf_config={}).to(DEVICE, dtype=DTYPE).eval()
|
||||
print(f" transformer built in {time.time() - t0:.1f}s; "
|
||||
f"params: {sum(p.numel() for p in model.parameters())/1e9:.2f}B")
|
||||
|
||||
ckpt_dir = ANYFLOW_LOCAL / "transformer"
|
||||
sd: dict = {}
|
||||
for shard in sorted(ckpt_dir.glob("diffusion_pytorch_model-*.safetensors")):
|
||||
sd.update(load_file(str(shard), device="cpu"))
|
||||
print(f" AnyFlow state dict: {len(sd)} tensors loaded")
|
||||
sd = translate_keys(sd, mapping=arch.param_names_mapping)
|
||||
info = model.load_state_dict(sd, strict=False)
|
||||
print(f" load: missing={len(info.missing_keys)} unexpected={len(info.unexpected_keys)}")
|
||||
del sd
|
||||
free_gpu()
|
||||
return model
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(model, text_e, neg_e, nfe: int) -> torch.Tensor:
|
||||
banner(f"(3) Sample 14B NFE={nfe}")
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
|
||||
FlowMapEulerDiscreteScheduler, )
|
||||
|
||||
scheduler = FlowMapEulerDiscreteScheduler(num_train_timesteps=1000, shift=5.0)
|
||||
scheduler.set_timesteps(num_inference_steps=nfe, device=DEVICE)
|
||||
timesteps = scheduler.timesteps.to(DEVICE, dtype=DTYPE)
|
||||
|
||||
B, C = 1, 16
|
||||
F = (NUM_FRAMES - 1) // 4 + 1 # temporal VAE compression = 4 (81 → 21)
|
||||
H_l, W_l = HEIGHT // 8, WIDTH // 8
|
||||
g = torch.Generator(device=DEVICE).manual_seed(SEED)
|
||||
x = torch.randn(B, C, F, H_l, W_l, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
|
||||
t0 = time.time()
|
||||
for i, (t_cur, t_next) in enumerate(zip(timesteps[:-1], timesteps[1:])):
|
||||
t_in = t_cur.expand(B).to(DTYPE)
|
||||
r_in = t_next.expand(B).to(DTYPE)
|
||||
with set_forward_context(current_timestep=t_in, attn_metadata=None):
|
||||
flow_cond = model(
|
||||
hidden_states=x, encoder_hidden_states=text_e,
|
||||
timestep=t_in, r_timestep=r_in)
|
||||
if GUIDANCE != 1.0:
|
||||
flow_uncond = model(
|
||||
hidden_states=x, encoder_hidden_states=neg_e,
|
||||
timestep=t_in, r_timestep=r_in)
|
||||
flow = flow_uncond + GUIDANCE * (flow_cond - flow_uncond)
|
||||
else:
|
||||
flow = flow_cond
|
||||
x = scheduler.step(
|
||||
flow, sample=x,
|
||||
timestep=t_cur.repeat(B), r_timestep=t_next.repeat(B))
|
||||
print(f" NFE={nfe} sample time: {time.time() - t0:.1f}s "
|
||||
f"({(time.time() - t0) / nfe:.1f}s/step)")
|
||||
xf = x.float()
|
||||
print(f" latents mean={xf.mean().item():+.3f} std={xf.std().item():.3f} "
|
||||
f"range=[{xf.min().item():+.2f}, {xf.max().item():+.2f}] "
|
||||
f"finite={torch.isfinite(xf).all().item()}")
|
||||
return x.detach()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_to_mp4(latents: torch.Tensor, out_path: Path) -> tuple[Path, tuple]:
|
||||
from diffusers.models.autoencoders.autoencoder_kl_wan import AutoencoderKLWan
|
||||
import imageio.v3 as iio
|
||||
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
str(ANYFLOW_LOCAL), subfolder="vae", torch_dtype=DTYPE,
|
||||
).to(DEVICE).eval()
|
||||
try:
|
||||
vae.enable_tiling()
|
||||
print(f" VAE tiling enabled")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
mean = torch.tensor(vae.config.latents_mean, device=DEVICE, dtype=DTYPE).view(1, -1, 1, 1, 1)
|
||||
std = torch.tensor(vae.config.latents_std, device=DEVICE, dtype=DTYPE).view(1, -1, 1, 1, 1)
|
||||
latents_unscaled = latents * std + mean
|
||||
t0 = time.time()
|
||||
frames = vae.decode(latents_unscaled, return_dict=False)[0]
|
||||
print(f" VAE decode time: {time.time() - t0:.1f}s")
|
||||
frames = (frames.clamp(-1, 1) + 1) / 2
|
||||
frames = frames[0].permute(1, 2, 3, 0).float().cpu().numpy()
|
||||
frames = (frames * 255).astype("uint8")
|
||||
iio.imwrite(str(out_path), frames, fps=16, codec="libx264", quality=8)
|
||||
del vae
|
||||
free_gpu()
|
||||
return out_path, frames.shape
|
||||
|
||||
|
||||
def main() -> None:
|
||||
torch.manual_seed(SEED)
|
||||
torch.cuda.manual_seed_all(SEED)
|
||||
init_single_rank()
|
||||
|
||||
text_e, neg_e = encode_prompts()
|
||||
model = load_transformer()
|
||||
|
||||
latents_list = []
|
||||
for nfe in [4, 50]:
|
||||
lat = sample(model, text_e, neg_e, nfe=nfe)
|
||||
latents_list.append((nfe, lat))
|
||||
|
||||
del model, text_e, neg_e
|
||||
free_gpu()
|
||||
|
||||
for nfe, lat in latents_list:
|
||||
banner(f"(4) Decode 14B NFE={nfe}")
|
||||
out_path = OUT_DIR / f"fastvideo_anyflow_14b_nfe{nfe}_seed{SEED}.mp4"
|
||||
path, shape = decode_to_mp4(lat, out_path)
|
||||
print(f" decoded {shape}")
|
||||
print(f" saved: {path} ({path.stat().st_size / 1e6:.2f} MB)")
|
||||
|
||||
banner("DONE")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,404 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo↔AnyFlow numerical parity verification.
|
||||
|
||||
Two checks, both on a single H200:
|
||||
|
||||
(A) Forward parity — load nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers into
|
||||
FastVideo's WanTransformer3DModel (with r_embedder enabled +
|
||||
param_names_mapping handling the delta_embedder rename), forward
|
||||
it on identical inputs against AnyFlow's reference loader, and
|
||||
compare. Expectation: rel-mean diff < 10% in bf16 (bf16 kernel noise).
|
||||
|
||||
(B) Any-step end-to-end sampling — run the new
|
||||
FlowMapEulerDiscreteScheduler through 4 Euler-flow steps on the
|
||||
same loaded weights, confirm the final latent is finite and
|
||||
well-scaled.
|
||||
|
||||
Configure local checkout paths via env vars (defaults assume a sibling
|
||||
layout next to this repo):
|
||||
|
||||
ANYFLOW_LOCAL — path to ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers``
|
||||
(default ``./anyflow-1.3b``)
|
||||
ANYFLOW_REF — path to the NVlabs/AnyFlow reference repo, used to
|
||||
import its loader (default ``./anyflow-ref``)
|
||||
|
||||
Run via::
|
||||
|
||||
PYTHONPATH=$PWD python scripts/verify_anyflow_fastvideo_parity.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
ANYFLOW_LOCAL = Path(os.environ.get("ANYFLOW_LOCAL", "./anyflow-1.3b")).expanduser()
|
||||
ANYFLOW_REF = Path(os.environ.get("ANYFLOW_REF", "./anyflow-ref")).expanduser()
|
||||
sys.path.insert(0, str(ANYFLOW_REF))
|
||||
|
||||
SEED = 1234
|
||||
DEVICE = torch.device("cuda")
|
||||
DTYPE = torch.bfloat16
|
||||
|
||||
|
||||
def banner(msg: str) -> None:
|
||||
print("\n" + "=" * 80)
|
||||
print(msg)
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Distributed bootstrap (single-rank). Required by WanTransformer3DModel
|
||||
# which calls get_sp_world_size() and uses ReplicatedLinear.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def init_single_rank() -> None:
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29551")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
from fastvideo.distributed import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
init_distributed_environment(world_size=1, rank=0, local_rank=0,
|
||||
backend="nccl")
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
sequence_model_parallel_size=1,
|
||||
data_parallel_size=1,
|
||||
)
|
||||
print(" single-rank distributed environment + TP/SP/DP groups initialized")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Translate AnyFlow HF safetensor keys onto FastVideo's WanTransformer3DModel
|
||||
# internal layout, applying the regex from WanVideoArchConfig.param_names_mapping.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def translate_keys(
|
||||
raw: dict[str, torch.Tensor],
|
||||
*,
|
||||
mapping: dict[str, str],
|
||||
) -> dict[str, torch.Tensor]:
|
||||
out: dict[str, torch.Tensor] = {}
|
||||
for k, v in raw.items():
|
||||
new_k = k
|
||||
for pat, repl in mapping.items():
|
||||
new_k = re.sub(pat, repl, new_k)
|
||||
out[new_k] = v
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Build FastVideo WanTransformer3DModel with AnyFlow weights loaded.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_fastvideo_transformer():
|
||||
banner("(1) Build FastVideo WanTransformer3DModel + load AnyFlow weights")
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
|
||||
|
||||
cfg = WanVideoConfig()
|
||||
arch = cfg.arch_config
|
||||
# AnyFlow Wan2.1-T2V-1.3B arch dims.
|
||||
arch.num_attention_heads = 12
|
||||
arch.attention_head_dim = 128
|
||||
arch.num_layers = 30
|
||||
arch.ffn_dim = 8960
|
||||
# AnyFlow dual-timestep.
|
||||
arch.r_embedder = True
|
||||
arch.r_embedder_fusion = "gated"
|
||||
arch.r_embedder_gate_value = 0.25
|
||||
arch.r_embedder_deltatime_type = "r"
|
||||
arch.__post_init__()
|
||||
|
||||
# WanTransformer3DModel takes (config, hf_config). hf_config is
|
||||
# diffusers-style — we provide a minimal dict; only fields the model
|
||||
# actually reads matter.
|
||||
hf_config: dict = {}
|
||||
t0 = time.time()
|
||||
model = WanTransformer3DModel(config=cfg, hf_config=hf_config)
|
||||
model = model.to(DEVICE, dtype=DTYPE).eval()
|
||||
print(f" Wan transformer built in {time.time() - t0:.1f}s; "
|
||||
f"params: {sum(p.numel() for p in model.parameters())/1e9:.2f}B")
|
||||
|
||||
# Load AnyFlow checkpoint.
|
||||
af_path = ANYFLOW_LOCAL / "transformer" / "diffusion_pytorch_model.safetensors"
|
||||
af_raw = load_file(str(af_path), device="cpu")
|
||||
print(f" AnyFlow checkpoint: {len(af_raw)} tensors")
|
||||
|
||||
translated = translate_keys(af_raw, mapping=arch.param_names_mapping)
|
||||
info = model.load_state_dict(translated, strict=False)
|
||||
miss, unex = info.missing_keys, info.unexpected_keys
|
||||
print(f" missing_keys : {len(miss)} (first 5: {miss[:5]})")
|
||||
print(f" unexpected_keys : {len(unex)} (first 5: {unex[:5]})")
|
||||
return model, len(miss), len(unex)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Build AnyFlow reference net.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_anyflow_reference():
|
||||
banner("(2) Build AnyFlow reference loader")
|
||||
from far.models import build_model
|
||||
|
||||
af_net = build_model("FAR_Wan_Transformer3DModel").from_pretrained(
|
||||
str(ANYFLOW_LOCAL),
|
||||
subfolder="transformer",
|
||||
chunk_partition=None,
|
||||
full_chunk_limit=0,
|
||||
compressed_patch_size=[1, 4, 4],
|
||||
).to(DEVICE, dtype=DTYPE).eval()
|
||||
print(f" AnyFlow {type(af_net).__name__} ready")
|
||||
return af_net
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Forward parity test.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def forward_compare(fv_model, af_net):
|
||||
banner("(3) Forward output comparison")
|
||||
B, C, F, H, W = 1, 16, 21, 60, 104
|
||||
SEQ, DIM = 32, 4096
|
||||
g = torch.Generator(device=DEVICE).manual_seed(SEED)
|
||||
x = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
|
||||
# AnyFlow expects per-frame (t, r) [B, F]; FastVideo T2V expects a
|
||||
# single scalar per sample [B] (timestep.dim()==2 is reserved for
|
||||
# Wan2.2 ti2v's per-token timestep schedule). For shared-t T2V the two
|
||||
# are semantically equivalent.
|
||||
t_per_frame = torch.full((B, F), 500.0, device=DEVICE, dtype=DTYPE)
|
||||
r_per_frame = torch.full((B, F), 200.0, device=DEVICE, dtype=DTYPE)
|
||||
t_per_sample = torch.full((B,), 500.0, device=DEVICE, dtype=DTYPE)
|
||||
r_per_sample = torch.full((B,), 200.0, device=DEVICE, dtype=DTYPE)
|
||||
|
||||
with torch.no_grad():
|
||||
# AnyFlow native: takes [B, F, C, H, W] + [B, F] (t, r).
|
||||
x_af = x.permute(0, 2, 1, 3, 4).contiguous()
|
||||
af_out = af_net(
|
||||
x_af,
|
||||
timestep=t_per_frame,
|
||||
r_timestep=r_per_frame,
|
||||
encoder_hidden_states=enc,
|
||||
return_dict=False,
|
||||
is_causal=False,
|
||||
)[0]
|
||||
if af_out.shape[1] != C and af_out.shape[2] == C:
|
||||
af_out = af_out.permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
# FastVideo: [B, C, F, H, W] + [B] (t, r). Must wrap in
|
||||
# set_forward_context so the attention layer can pick up the
|
||||
# current timestep / attn_metadata.
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
with set_forward_context(
|
||||
current_timestep=t_per_sample,
|
||||
attn_metadata=None,
|
||||
):
|
||||
fv_out = fv_model(
|
||||
hidden_states=x,
|
||||
encoder_hidden_states=enc,
|
||||
timestep=t_per_sample,
|
||||
r_timestep=r_per_sample,
|
||||
)
|
||||
|
||||
print(f" AnyFlow out: shape={tuple(af_out.shape)} dtype={af_out.dtype}")
|
||||
print(f" FastVideo : shape={tuple(fv_out.shape)} dtype={fv_out.dtype}")
|
||||
if af_out.shape != fv_out.shape:
|
||||
print(" ❌ shape mismatch")
|
||||
return False
|
||||
diff = (af_out.float() - fv_out.float()).abs()
|
||||
ref = af_out.float().abs().mean().item() + 1e-12
|
||||
print(f" max abs diff : {diff.max().item():.3e}")
|
||||
print(f" mean abs diff: {diff.mean().item():.3e}")
|
||||
print(f" rel mean diff: {diff.mean().item() / ref:.3e}")
|
||||
ok = diff.mean().item() / ref < 0.10
|
||||
print(f" >>> {'PASS' if ok else 'FAIL'} (target rel diff < 10%, bf16 noise)")
|
||||
return ok
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Any-step sampling smoke via FlowMapEulerDiscreteScheduler.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def sample_anystep(fv_model):
|
||||
banner("(4) Any-step 4-step Euler-flow sampling")
|
||||
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
|
||||
FlowMapEulerDiscreteScheduler, )
|
||||
|
||||
scheduler = FlowMapEulerDiscreteScheduler(num_train_timesteps=1000, shift=5.0)
|
||||
scheduler.set_timesteps(num_inference_steps=4, device=DEVICE)
|
||||
timesteps = scheduler.timesteps.to(dtype=DTYPE)
|
||||
|
||||
B, C, F, H, W = 1, 16, 21, 60, 104
|
||||
SEQ, DIM = 32, 4096
|
||||
g = torch.Generator(device=DEVICE).manual_seed(SEED)
|
||||
x = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
t0 = time.time()
|
||||
with torch.no_grad():
|
||||
for t_cur, t_next in zip(timesteps[:-1], timesteps[1:]):
|
||||
t_in = t_cur.expand(B).to(DTYPE)
|
||||
r_in = t_next.expand(B).to(DTYPE)
|
||||
with set_forward_context(current_timestep=t_in, attn_metadata=None):
|
||||
v = fv_model(
|
||||
hidden_states=x,
|
||||
encoder_hidden_states=enc,
|
||||
timestep=t_in,
|
||||
r_timestep=r_in,
|
||||
)
|
||||
x = scheduler.step(
|
||||
v, sample=x,
|
||||
timestep=t_cur.repeat(B),
|
||||
r_timestep=t_next.repeat(B),
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
xf = x.float()
|
||||
print(f" elapsed: {elapsed:.1f}s")
|
||||
print(f" final latent: mean={xf.mean().item():+.3f} std={xf.std().item():.3f} "
|
||||
f"range=[{xf.min().item():+.2f}, {xf.max().item():+.2f}] "
|
||||
f"finite={torch.isfinite(xf).all().item()}")
|
||||
ok = torch.isfinite(xf).all().item() and 0.01 < xf.std().item() < 30
|
||||
print(f" >>> {'PASS' if ok else 'FAIL'}")
|
||||
return ok
|
||||
|
||||
|
||||
NUM_TRAIN_TIMESTEPS = 1000
|
||||
EPSILON = 5.0 # AnyFlow paper default
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def training_step_compare(fv_model, af_net) -> bool:
|
||||
"""Inline replica of AnyFlow's train_bidirection central-difference loss
|
||||
on both code paths with identical synthetic (real, noise, t, r) inputs.
|
||||
|
||||
Compares scalar loss + intermediate flow_pred / target tensors.
|
||||
"""
|
||||
banner("(5) Training-step loss comparison (central-difference target)")
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
|
||||
B, C, F, H, W = 1, 16, 21, 60, 104
|
||||
SEQ, DIM = 32, 4096
|
||||
g = torch.Generator(device=DEVICE).manual_seed(SEED)
|
||||
real = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
|
||||
noise = torch.randn_like(real)
|
||||
t_abs_pf = torch.full((B, F), 500.0, device=DEVICE, dtype=DTYPE)
|
||||
r_abs_pf = torch.full((B, F), 200.0, device=DEVICE, dtype=DTYPE)
|
||||
t_abs_ps = torch.full((B,), 500.0, device=DEVICE, dtype=DTYPE)
|
||||
r_abs_ps = torch.full((B,), 200.0, device=DEVICE, dtype=DTYPE)
|
||||
|
||||
# AnyFlow inline replica: uses [B, F, C, H, W] layout.
|
||||
real_btchw = real.permute(0, 2, 1, 3, 4).contiguous()
|
||||
noise_btchw = noise.permute(0, 2, 1, 3, 4).contiguous()
|
||||
t_norm_pf = (t_abs_pf / NUM_TRAIN_TIMESTEPS).view(B, F, 1, 1, 1).to(DTYPE)
|
||||
noisy_btchw = t_norm_pf * noise_btchw + (1 - t_norm_pf) * real_btchw
|
||||
|
||||
def u_func_af(x_in, t_in, r_in):
|
||||
return af_net(
|
||||
x_in, timestep=t_in, r_timestep=r_in,
|
||||
encoder_hidden_states=enc, return_dict=False, is_causal=False)[0]
|
||||
|
||||
v_pred = noise_btchw - real_btchw
|
||||
eps = EPSILON
|
||||
F_plus = u_func_af(noisy_btchw + v_pred * (eps / NUM_TRAIN_TIMESTEPS),
|
||||
t_abs_pf + eps, r_abs_pf)
|
||||
F_minus = u_func_af(noisy_btchw - v_pred * (eps / NUM_TRAIN_TIMESTEPS),
|
||||
t_abs_pf - eps, r_abs_pf)
|
||||
dF_dt_af = (F_plus - F_minus) / (2 * eps)
|
||||
target_af = ((noise_btchw - real_btchw)
|
||||
- (t_abs_pf - r_abs_pf).view(B, F, 1, 1, 1) * dF_dt_af)
|
||||
flow_af = u_func_af(noisy_btchw, t_abs_pf, r_abs_pf)
|
||||
loss_af = (flow_af.float() - target_af.float()).pow(2).reshape(B, -1).mean(-1)
|
||||
|
||||
# FastVideo inline replica: uses [B, C, F, H, W] layout + [B] t/r.
|
||||
real_bcfhw = real
|
||||
noise_bcfhw = noise
|
||||
t_norm_ps = (t_abs_ps / NUM_TRAIN_TIMESTEPS).view(B, 1, 1, 1, 1).to(DTYPE)
|
||||
noisy_bcfhw = t_norm_ps * noise_bcfhw + (1 - t_norm_ps) * real_bcfhw
|
||||
|
||||
def u_func_fv(x_in, t_in, r_in):
|
||||
with set_forward_context(current_timestep=t_in, attn_metadata=None):
|
||||
return fv_model(
|
||||
hidden_states=x_in,
|
||||
encoder_hidden_states=enc,
|
||||
timestep=t_in,
|
||||
r_timestep=r_in,
|
||||
)
|
||||
|
||||
v_pred_fv = noise_bcfhw - real_bcfhw
|
||||
F_plus_fv = u_func_fv(
|
||||
noisy_bcfhw + v_pred_fv * (eps / NUM_TRAIN_TIMESTEPS),
|
||||
t_abs_ps + eps, r_abs_ps)
|
||||
F_minus_fv = u_func_fv(
|
||||
noisy_bcfhw - v_pred_fv * (eps / NUM_TRAIN_TIMESTEPS),
|
||||
t_abs_ps - eps, r_abs_ps)
|
||||
dF_dt_fv = (F_plus_fv - F_minus_fv) / (2 * eps)
|
||||
target_fv = ((noise_bcfhw - real_bcfhw)
|
||||
- (t_abs_ps - r_abs_ps).view(B, 1, 1, 1, 1) * dF_dt_fv)
|
||||
flow_fv = u_func_fv(noisy_bcfhw, t_abs_ps, r_abs_ps)
|
||||
loss_fv = (flow_fv.float() - target_fv.float()).pow(2).reshape(B, -1).mean(-1)
|
||||
|
||||
# Compare. Align AnyFlow's [B, F, C, H, W] → [B, C, F, H, W].
|
||||
flow_af_aligned = flow_af.permute(0, 2, 1, 3, 4)
|
||||
target_af_aligned = target_af.permute(0, 2, 1, 3, 4)
|
||||
flow_diff = (flow_fv.float() - flow_af_aligned.float()).abs()
|
||||
target_diff = (target_fv.float() - target_af_aligned.float()).abs()
|
||||
|
||||
af_loss_v = loss_af.mean().item()
|
||||
fv_loss_v = loss_fv.mean().item()
|
||||
abs_diff = abs(af_loss_v - fv_loss_v)
|
||||
rel_diff = abs_diff / abs(af_loss_v + 1e-12)
|
||||
print(f" AnyFlow loss : {af_loss_v:.6f}")
|
||||
print(f" FastVideo loss: {fv_loss_v:.6f}")
|
||||
print(f" abs diff : {abs_diff:.3e}")
|
||||
print(f" rel diff : {rel_diff:.3e}")
|
||||
print(f" flow_pred : max abs {flow_diff.max().item():.3e} "
|
||||
f"mean {flow_diff.mean().item():.3e}")
|
||||
print(f" target : max abs {target_diff.max().item():.3e} "
|
||||
f"mean {target_diff.mean().item():.3e}")
|
||||
ok = rel_diff < 0.20
|
||||
print(f" >>> {'PASS' if ok else 'FAIL'} (target rel loss diff < 20%)")
|
||||
return ok
|
||||
|
||||
|
||||
def main() -> None:
|
||||
torch.manual_seed(SEED)
|
||||
torch.cuda.manual_seed_all(SEED)
|
||||
|
||||
init_single_rank()
|
||||
fv_model, n_miss, n_unex = build_fastvideo_transformer()
|
||||
af_net = build_anyflow_reference()
|
||||
forward_ok = forward_compare(fv_model, af_net)
|
||||
sample_ok = sample_anystep(fv_model)
|
||||
train_ok = training_step_compare(fv_model, af_net)
|
||||
banner(
|
||||
f"SUMMARY: missing_keys={n_miss} unexpected_keys={n_unex} "
|
||||
f"forward_parity={forward_ok} sample_smoke={sample_ok} "
|
||||
f"training_parity={train_ok}"
|
||||
)
|
||||
sys.exit(0 if (forward_ok and sample_ok and train_ok) else 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,77 @@
|
||||
# Cosmos3 Audio (PR2) — Port Plan
|
||||
|
||||
Branch: `feat/cosmos3-audio` (stacked on `feat/cosmos3-i2v`, which has T2V/I2V/T2I).
|
||||
Goal: text-to-video+sound (**t2vs**) — generate synchronized audio alongside video.
|
||||
|
||||
## How the framework does audio (studied 2026-06-07)
|
||||
|
||||
- **Sound tokenizer = AVAE** (`cosmos_framework/model/vfm/tokenizers/audio/avae.py`
|
||||
+ `avae_utils/`, ~2268 lines): a 48 kHz **stereo** neural audio codec.
|
||||
- checkpoint: `official_weights/cosmos3/sound_tokenizer/` (`model_type:
|
||||
autoencoder_v2`, ~1.9 GB). enc=`spec_convnext` (enc_dim 192, latent_dim 128,
|
||||
n_fft 64), dec=`oobleck` (dec_dim 320, strides [2,4,5,6,8]), VAE bottleneck,
|
||||
`snakebeta` activations, hop_size 1920.
|
||||
- interface: `encode(audio[1,C,N]) -> latent`, `decode(latent) -> audio`,
|
||||
`get_latent_num_samples(N)`, `sample_rate=48000`, `audio_channels=2`,
|
||||
`sound_latent_fps=25`.
|
||||
- **DiT sound pathway** (`cosmos3_vfm_network.py`, 136 sound/audio refs): the MoT
|
||||
has `sound2llm` / `llm2sound` / `sound_modality_embed` + `pack_sound_latents`
|
||||
and joint vision+sound denoising (`preds_sound`, sound `condition_mask`, sound
|
||||
noise init `cond_mask*x0 + (1-cond_mask)*noise`, velocity `pred*(1-cond_mask)`).
|
||||
- FastVideo's native DiT ALREADY constructs the dormant heads
|
||||
(`audio_proj_in`/`audio_proj_out`/`audio_modality_embed`, gated on
|
||||
`arch.sound_gen`) for strict-load — the forward just doesn't use them yet.
|
||||
- **Inference flow** (`cosmos_framework/inference/sound.py`): t2vs builds a
|
||||
zero **placeholder audio** sized to the video duration (sets sound latent
|
||||
length), `inject_sound_into_batch` upgrades the SequencePlan to has_sound,
|
||||
the omni model denoises vision+sound jointly, then AVAE-decodes the sound
|
||||
latent and `mux_audio_into_video` (PyAV, AAC) muxes it into the mp4
|
||||
(`save_sound` writes a WAV).
|
||||
|
||||
## Components (each: native port + framework parity test, per methodology)
|
||||
|
||||
1. **AVAE codec** — `fastvideo/models/.../cosmos3_avae.py` + config. Port
|
||||
encoder/decoder/bottleneck/snake. Parity: tiny AVAE, framework weights copied
|
||||
in, bit-exact `decode` (and `encode`) on CPU/fp32. **(largest piece)**
|
||||
2. **DiT sound pathway** — activate the dormant heads in `forward`; port
|
||||
`pack_sound_latents` + sound token scatter/proj/modality-embed/velocity.
|
||||
Parity: extend the DiT harness with sound tokens.
|
||||
3. **Sound sequence packing** — extend `sequence_packing.py` with the sound
|
||||
modality (positions, attn mode, condition mask). Parity vs framework
|
||||
`pack_input_sequence` with sound.
|
||||
4. **Pipeline (t2vs)** — placeholder audio -> joint denoise -> split ->
|
||||
AVAE-decode sound -> mux into mp4 / save wav. Extend `Cosmos3DenoisingStage`
|
||||
+ a sound-decode/mux stage.
|
||||
5. **FastVideo AV infra** — audio in `OutputConfig` / a mux stage (check what
|
||||
exists; `cosmos_framework.inference.sound.mux_audio_into_video` is the ref).
|
||||
|
||||
## Open decisions
|
||||
- **D1 (AVAE approach)** — full native port (methodology-consistent; ~2.3k lines)
|
||||
vs a documented lazy-wrapper around the framework AVAE (faster; but pulls heavy
|
||||
deps and bends the "native + no-framework-at-runtime" rule). Default per
|
||||
methodology: native port.
|
||||
- **D2 (scope)** — t2vs (T+video+sound) first; defer audio-conditioned / v2vs.
|
||||
- **D3** — confirm FastVideo can mux/emit audio (output format).
|
||||
|
||||
## Status
|
||||
- [x] Branch forked, framework audio path studied, plan written.
|
||||
- [x] D1: native port (user-chosen). D2: t2vs first.
|
||||
- [x] **AVAE sound decoder (component 1) — DONE** (commit `5f81fb3d5`). Key
|
||||
finding: the checkpoint is decoder-only in AutoencoderOobleck naming with
|
||||
SnakeBeta + weight_g/v == FastVideo's native `OobleckVAE` decoder. Reused it
|
||||
(+ `output_padding=stride%2` for the odd stride 5); `Cosmos3SoundVAE`
|
||||
decoder-only wrapper; bit-exact parity vs the framework OobleckDecoder
|
||||
(`test_cosmos3_avae_parity`); real 1.9 GB checkpoint strict-loads, decodes
|
||||
[1,64,25] -> [1,2,48000] (1 s @ 48 kHz stereo).
|
||||
- [x] **DiT sound pathway (component 2) — DONE** (commit `005d6684a`). Activated
|
||||
the dormant audio heads in the forward (`_encode_sound`/`_decode_sound` mirror);
|
||||
`preds_vision` + `preds_sound` bit-exact (max=mean=0.0).
|
||||
- [x] **Sound sequence packing (component 3) — DONE** (commit `005d6684a`).
|
||||
`Cosmos3SoundItem` + sound fields; sound shares the vision "full" split with
|
||||
parallel MRoPE. Field-by-field + position_ids exact vs framework.
|
||||
- [x] **t2vs pipeline + AV mux (components 4-5) — DONE** (commit `3d8355129`).
|
||||
Joint [vision|sound] denoise, AVAE-decode, stereo 48 kHz AAC mux. t2vs CFG
|
||||
velocity parity max=mean=0.0; real-weights run produces coherent video + real
|
||||
audio (mean -10.2 dB). Example `basic_cosmos3_t2vs_new_api.py`.
|
||||
|
||||
**PR2 (audio/t2vs) COMPLETE** — every component bit-exact vs the framework.
|
||||
@@ -0,0 +1,152 @@
|
||||
# Cosmos3 Port Status
|
||||
|
||||
## Summary
|
||||
|
||||
- model_family: `cosmos3`
|
||||
- workload_types: `T2V, I2V, T2I` supported by `WorkloadType` today; full-omni target also needs audio (AV), VLM reasoning, and action-conditioning, which require framework extensions (Q002, Q003).
|
||||
- official_ref: `https://github.com/NVIDIA/cosmos-framework` — diffusers backend `diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline`; HF `nvidia/Cosmos3-Nano`.
|
||||
- official_ref_dir: `cosmos-framework` (symlink -> `/home/william5lin/FastVideo/cosmos-framework`, commit `003d66d4`)
|
||||
- hf_weights_path: `nvidia/Cosmos3-Nano`
|
||||
- local_weights_dir: `official_weights/cosmos3` (symlink -> `/home/william5lin/FastVideo/official_weights/cosmos3`, 33 GiB / 67 files)
|
||||
- source_layout: `diffusers`
|
||||
- local_tests_readme: `tests/local_tests/cosmos3/README.md`
|
||||
|
||||
## Current Phase
|
||||
|
||||
- phase: `FULL OMNI SUPPORTED — every modality framework-parity verified bit-exact (suite 150 passed, 0 skipped). PR1 video core (T2V/I2V/T2I) + PR2 audio (t2vs) real-weights verified on B200; PR3 action (domain-aware) + PR4 reasoning (text + vision_encoder + deepstack reasoner) bit-exact. Branch chain: feat/cosmos3-tier-a-port (T2V) -> feat/cosmos3-i2v (I2V+T2I+flow_shift) -> feat/cosmos3-audio (t2vs) -> feat/cosmos3-action -> feat/cosmos3-reasoning. Optional follow-ups: real-weights action2world (needs robot-action data) + image-conditioned-reasoning prefill wiring (vision_encoder + get_rope_index, both proven).`
|
||||
- status: `in_progress`
|
||||
- owner: `orchestrator`
|
||||
- last_updated: `2026-06-07`
|
||||
- env: `fv-cosmos3` (conda clone of fv-main; `fastvideo` editable repointed to this worktree). Run tests from the worktree cwd with this env's python.
|
||||
- branch: rebased onto `origin/main` @ `1c627a3f9` (was 33 behind, merge-base 2026-05-22); now 6 commits ahead; `fastvideo` imports clean; Tier-A `13 passed, 2 skipped`.
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|
||||
|---|---|---|---|---|---|---|---|---|---|
|
||||
| transformer | dit | port | `diffusers_cosmos3/transformer.py:Cosmos3OmniTransformer` (model_type `qwen3_vl_text`, MoT + MRoPE) | `model_index.json: transformer`; `cosmos_framework/model/vfm/mot/cosmos3_vfm_network.py`, `omni_mot_model.py` | `fastvideo/models/dits/cosmos3.py` (branch: `Cosmos3VFMTransformer`+`Cosmos3LanguageModel` — reconcile to `Cosmos3OmniTransformer`) | skeleton | not_started | scaffold_skip | I001 |
|
||||
| vae | vae | reuse | diffusers `AutoencoderKLWan` | `model_index.json: vae` | reuse Wan VAE (`fastvideo/models/vaes/`, cf. `cosmos25wanvae.py`) | not_started | passthrough? | not_started | Q001 |
|
||||
| scheduler | generic | reuse (flow-coerced) | framework `FlowUniPCMultistepScheduler` (`cosmos_framework/.../fm_solvers_unipc.py`; checkpoint ships diffusers-style config) | `model_index.json: scheduler`; `cosmos_framework/.../samplers/unipc.py:UniPCSampler` | FastVideo-native `UniPCMultistepScheduler` (flow config), coerced in `initialize_pipeline` | done | n/a | framework-parity DONE (`test_cosmos3_scheduler_parity`: timesteps bit-exact, sigmas ~1e-8, trajectory <~1e-6) | I003 (resolved) |
|
||||
| text_tokenizer | tokenizer | reuse | transformers `Qwen2TokenizerFast` | `model_index.json: text_tokenizer` | reuse (tokenizer = allowed third-party) | not_started | passthrough | scaffold_skip (`test_cosmos3_tokenizer_chat_template`) | - |
|
||||
| vision_encoder | encoder | port | transformers `Qwen3VLVisionModel` | `model_index.json: vision_encoder` | new encoder bucket OR documented lazy-wrapper | not_started | not_started | not_started | Q002 |
|
||||
| sound_tokenizer | generic/vae | port (decode) | framework AVAE `LatentAutoEncoderV2` (`avae_utils`); checkpoint is decoder-only AutoencoderOobleck-named w/ SnakeBeta | `model_index.json: sound_tokenizer` | reuse FastVideo native `OobleckVAE` decoder + `Cosmos3SoundVAE` wrapper (`models/audio/cosmos3_avae.py`) | done (decode) | n/a | DECODE bit-exact vs framework (`test_cosmos3_avae_parity`); real ckpt strict-loads | PR2 (branch feat/cosmos3-audio) |
|
||||
|
||||
## Conversion State
|
||||
|
||||
- conversion_script: `scripts/checkpoint_conversion/cosmos3_convert.py` (branch has it, 246 lines, built vs vllm-omni — repoint/verify vs diffusers checkpoint)
|
||||
- converted_weights_dir: `converted_weights/cosmos3` (n/a while needs_conversion=no)
|
||||
- source_layout: `diffusers`
|
||||
- needs_conversion: `no` (HF already diffusers-format; verify FastVideo loaders consume directly)
|
||||
- strict_load_status: `not_run`
|
||||
- passthrough_components: `vae (AutoencoderKLWan), scheduler (UniPC), text_tokenizer (Qwen2)` likely passthrough
|
||||
- retry_history: `none`
|
||||
|
||||
## Parity Commands
|
||||
|
||||
| Scope | Command | Last Result | Notes |
|
||||
|---|---|---|---|
|
||||
| Tier-A scaffold | `cd <worktree> && <fv-cosmos3 python> -m pytest tests/local_tests/cosmos3/ -q` | `13 passed, 2 skipped` (2026-06-06, post-rebase) | 2 skips: Cosmos3 tokenizer/_tokenize_prompt not yet wired on pipeline |
|
||||
| component | `pytest tests/local_tests/<bucket>/test_cosmos3_<component>_parity.py -v -s` | `not_run` | after env activation + native prototypes |
|
||||
| pipeline | `pytest tests/local_tests/pipelines/test_cosmos3_pipeline_parity.py -v -s` | `not_run` | |
|
||||
|
||||
## Open Questions
|
||||
|
||||
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|
||||
|---|---|---|---|---|---|
|
||||
| Q001 | Does Cosmos3 VAE (`AutoencoderKLWan`) match FastVideo's existing Wan VAE config/instantiation exactly (z_dim, scale factors, latents_mean/std)? | orchestrator | 3 (reuse gate) | open | |
|
||||
| Q002 | `vision_encoder` (`Qwen3VLVisionModel`): native port vs documented lazy-wrapper exception? Needed for I2V/reasoning. | user/orchestrator | 3 | open | |
|
||||
| Q003 | `sound_tokenizer` (`Cosmos3AVAEAudioTokenizer`) + audio output requires `WorkloadType` AV + audio regression metric. | user | 0/10 | open | full-omni scope chosen 2026-06-06; infra extensions pending |
|
||||
|
||||
## Issues And Blockers
|
||||
|
||||
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| I001 | port | transformer | high | Branch DiT (`Cosmos3VFMTransformer`+`Cosmos3LanguageModel`) built vs vllm-omni #3454; official checkpoint loads `Cosmos3OmniTransformer` (diffusers shim). Class/structure reconciliation required. | `model_index.json`; `diffusers_cosmos3/transformer.py`; branch commit `52bb65f49` | orchestrator | resolved | DiT rewritten to checkpoint layout (single `layers` dual-pathway, BaseDiT-conformant); bit-identical framework parity (3d_rope + unified_3d_mrope), commits 59a4a571c/7c4633295 |
|
||||
| I002 | all | tests | medium | Tier-A conftest+tests mirror vllm-omni line-by-line (stubs, `vllm_omni...guardrails`). Must be repointed to `diffusers_cosmos3` / official structures. | `tests/local_tests/cosmos3/conftest.py` | orchestrator | open | |
|
||||
| I003 | inference | scheduler | high | First real-weights T2V was all-black: checkpoint `scheduler_config.json` sets `use_karras_sigmas=true`; vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas -> `scheduler.step` -> NaN latents. DiT/CFG velocity was clean. The scheduler had never been parity-tested vs the framework (`test_cosmos3_denoise_cfg_parity` used diffusers UniPC on both sides). | `result_latent` NaN at denoise step 0 (v_pred clean); ffprobe 3 KB black mp4 | orchestrator | resolved | Coerce loaded config to flow setup in `initialize_pipeline`; switch pipeline+tests to native UniPC (no diffusers at runtime); add `test_cosmos3_scheduler_parity` vs framework `FlowUniPCMultistepScheduler`; repoint denoise_cfg oracle to the framework scheduler. Commit 255311cf2 |
|
||||
|
||||
## Escape Hatches
|
||||
|
||||
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|
|
||||
| E001 | prep | dependency/env | Shared `fv-main` env has `fastvideo` editable-installed from the MAIN worktree; the cosmos3 worktree's `fastvideo` is not importable (PEP660 finder overrides PYTHONPATH), so Tier-A tests skip. How to activate the worktree's `fastvideo` for verification without disrupting ~24 other worktrees sharing the env? | Dedicated conda env for the cosmos3 worktree | resolved | Created fv-cosmos3 (clone of fv-main); repointed fastvideo editable to worktree; run from worktree cwd. Branch also rebased onto origin/main to fix stale import. |
|
||||
|
||||
## Decisions
|
||||
|
||||
| Date | Decision | Rationale | Impact |
|
||||
|---|---|---|---|
|
||||
| 2026-06-06 | Reference source of truth = official diffusers (`Cosmos3OmniDiffusersPipeline` + `cosmos-framework`/`diffusers-cosmos3`), not vllm-omni #3454 | Official weights now public & diffusers-format; the artifact users actually load | Repoint DiT/pipeline/conversion/tests off vllm-omni (I001, I002) |
|
||||
| 2026-06-06 | Resume in worktree `/home/william5lin/FastVideo_cosmos3_port`; weights+reference symlinked (no copy) | Preserve 2,492 lines of Tier-A work; avoid 33 GB duplication | Verification needs worktree `fastvideo` active (E001) |
|
||||
| 2026-06-06 | Scope = full omni (video + audio + reasoning + action) | User choice (revised from branch's original video-only scope) | Adds `vision_encoder`, `sound_tokenizer` ports + `WorkloadType` AV + audio metric |
|
||||
| 2026-06-06 | Downloaded full 34.9 GB (33 GiB) `nvidia/Cosmos3-Nano` | Unblocks May-22 `PENDING` weight status (HF was 401, now public) | Real parity now possible |
|
||||
| 2026-06-06 | Rebased branch onto origin/main (33 commits); resolved registry.py conflict by reconstructing from main + cosmos3 import/entry | Branch was stale; fastvideo failed to import (main removed MatrixGameI2V480PConfig) | Branch imports clean; Tier-A 13 passed/2 skipped |
|
||||
| 2026-06-06 | Reference = cosmos_framework ONLY (full omni); diffusers shim dropped even for video | User directive (Phase 1 found diffusers __call__ is video-only; sound/action/reasoning live only in the framework) | Larger port; ref DiT = `Cosmos3VFMNetwork`/`Cosmos3VFMNetworkConfig` (not diffusers `Cosmos3OmniTransformer`); core model imports in fv-cosmos3 with light deps; TE only in optional dot_product_attention |
|
||||
|
||||
## Handoff Notes
|
||||
|
||||
- Prep (weights/reference/env editable installs) done in MAIN worktree; symlinked into this worktree. Env installs (`diffusers-cosmos3`, `cosmos-framework`) are in shared `fv-main`.
|
||||
- Next: resolve E001 (env), then Phase 1 reference study of `diffusers_cosmos3` pipeline/transformer, then Phase 3 reuse gate (VAE/scheduler/tokenizer) + component dispatch (transformer, vision_encoder, sound_tokenizer).
|
||||
- diffusers 0.36.0 imports the shim OK; checkpoint saved with 0.37.1 — watch `from_pretrained` needs (bump within FastVideo's `diffusers>=0.33.1` pin if required).
|
||||
|
||||
### PR1 (video core) progress — 2026-06-06
|
||||
- Arch config 1:1 with checkpoint, committed `9567efdf0`.
|
||||
- Framework parity-reference harness committed `dd97efda3`: `tests/local_tests/cosmos3/test_cosmos3_reference_forward.py` builds a tiny `Cosmos3VFMNetwork` on CPU/float32 (SDPA monkeypatch; flash2/3/natten are CUDA-only) and forwards `packed_seq -> {last_hidden_state, preds_vision}`. 23 tests pass in fv-cosmos3. This is the ground-truth side for DiT parity. Run: `cd <worktree> && <fv-cosmos3 py> -m pytest tests/local_tests/cosmos3/test_cosmos3_reference_forward.py -q`.
|
||||
- THREE naming conventions to bridge:
|
||||
1. framework-native (`Cosmos3VFMNetwork`): `language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)`, `{q,k}_norm(+_moe_gen)`, `mlp(+_moe_gen)`, `vae2llm`/`llm2vae`, `time_embedder.mlp.{0,2}`.
|
||||
2. diffusers checkpoint (on disk, what we load): `layers.{i}.self_attn.{to_q,to_k,to_v,to_out}` + `{add_q,add_k,add_v}_proj`/`to_add_out`, `{norm_q,norm_k,norm_added_q,norm_added_k}`, `mlp`/`mlp_moe_gen`, `proj_in`/`proj_out`, `time_embedder.linear_{1,2}`.
|
||||
3. FastVideo DiT (our choice). Conversion maps (2)->(3); the DiT parity test copies (1)->(3).
|
||||
- BaseDiT signature is `__init__(self, config: DiTConfig, hf_config: dict)`; the branch `Cosmos3VFMTransformer` uses `fastvideo_args`/SimpleNamespace and does NOT conform — rewrite to conform + match the checkpoint key surface (single `layers` dual-pathway, not split language_model/gen_layers).
|
||||
- Native layers (per cosmos2_5): `ReplicatedLinear`/`MLP`/`RMSNorm` (fastvideo.layers.*), `LocalAttention`/`DistributedAttention` (fastvideo.attention), `apply_rotary_emb` (use_real_unbind_dim=-2 for Cosmos). EntryClass at module bottom; class attrs bound from config; 3D-MRoPE has no reusable util — adapt Cosmos25RotaryPosEmbed.
|
||||
- NEXT: write native `fastvideo/models/dits/cosmos3.py` + fastvideo-vs-framework forward parity test (copy framework weights into the FastVideo DiT, compare outputs), then conversion script (diffusers checkpoint -> FastVideo) + strict-load, then video pipeline/packing.
|
||||
|
||||
### PR1 (video core) acceptance — real-weights E2E — 2026-06-07
|
||||
- First real-weights T2V (`examples/inference/basic/basic_cosmos3_new_api.py`, `COSMOS3_MODEL_PATH=official_weights/cosmos3`) ran mechanically but produced an all-black 3 KB mp4. Instrumenting the denoise loop showed `v_pred` clean at step 0 but `scheduler.step` -> NaN. Root cause I003: checkpoint `scheduler_config.json` is diffusers-style (`use_karras_sigmas=true`), and the vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas instead of flow sigmas -> NaN. The framework actually samples with `FlowUniPCMultistepScheduler` (pure flow: `shift` + `num_train_timesteps`).
|
||||
- Fix (commit `255311cf2`): coerce the loaded scheduler to the flow setup in `Cosmos3OmniDiffusersPipeline.initialize_pipeline`; use FastVideo's native UniPC (not diffusers) in pipeline + tests. Added `test_cosmos3_scheduler_parity.py` (native UniPC flow-config vs framework `FlowUniPCMultistepScheduler`: timesteps bit-exact, sigmas ~1e-8, full trajectory <~1e-6 over shift in {10,3}, steps in {4,10,35}). Repointed `test_cosmos3_denoise_cfg_parity` oracle to the framework scheduler (it previously compared diffusers-vs-diffusers, so the scheduler was never checked against the framework).
|
||||
- Also wired the remaining integration glue (registry alias `Cosmos3OmniTransformer`->`Cosmos3VFMTransformer`; `text_tokenizer`->TokenizerLoader; scheduler config param-filtering; DiT `materialize_non_persistent_buffers` + compute-dtype casts; packing device-move in `to_dit_kwargs`; empty text-preprocess).
|
||||
- Verified: 1280x704, 29 frames, 35 steps on a single B200 -> coherent golden-retriever-in-meadow video matching the prompt (no NaNs; per-frame pixel std ~58; visible temporal motion). Full cosmos3 suite: 95 passed, 0 skipped.
|
||||
- NEXT: PR2 audio (`sound_tokenizer` AVAE) / PR3 action / PR4 reasoning. Optional: I2V/T2I real-weights spot-checks; force-push branch (needs explicit OK).
|
||||
|
||||
### PR1 (video core) — I2V real-weights — 2026-06-07 (branch feat/cosmos3-i2v)
|
||||
- Forked `feat/cosmos3-i2v` off `feat/cosmos3-tier-a-port` (stacked, includes the T2V + scheduler fix).
|
||||
- Studied the framework I2V path: `cosmos_framework.inference.vision.load_conditioning_image` (aspect-preserving resize + center crop + uint8 quantize -> `/127.5-1`) + `build_conditioned_video_batch` (frame 0 = image, remaining frames REPEAT the last conditioning frame -> static video), then VAE-encode; `condition_frame_indexes=[0]` (latent). Condition frames kept clean during sampling exactly as FastVideo already does: init noise `cond_mask*x0 + (1-cond_mask)*noise` (`omni_mot_model._prepare_inference_data`) + velocity zeroed `pred*(1-cond_mask)` each step (`_get_velocity`), no re-injection.
|
||||
- Bug found + fixed (commit `bd8d604fb`): FastVideo's `_image_to_video_tensor` ZERO-filled the non-condition frames; the temporal Wan VAE (4x) makes latent frame 0 depend on several pixel frames, so zero-fill -> wrong conditioning latent. Rewrote it to repeat-fill + framework resize/crop/quantize.
|
||||
- Parity: `test_cosmos3_i2v_conditioning_parity.py` vs framework `load_conditioning_image` + repeat-fill — bit-exact (max abs diff 0.0) across aspect/size/frame cases. Existing `test_cosmos3_denoise_cfg_parity` already covers the I2V cond-mask + velocity math (i2v case).
|
||||
- Example: `examples/inference/basic/basic_cosmos3_i2v_new_api.py` (`InputConfig(image_path=...)`, default `assets/images/cyclist.jpg`).
|
||||
- Verified on B200 (1280x704, 29f, 35 steps, real weights): output frame 0 reproduces the conditioning cyclist image; later frames show coherent forward motion down the trail following the prompt. Full suite 98 passed, 0 skipped.
|
||||
- NEXT: optional T2I real-weights spot-check; then PR2 audio / PR3 action / PR4 reasoning.
|
||||
|
||||
### PR1 (video core) — T2I real-weights + resolution-based flow_shift — 2026-06-07 (branch feat/cosmos3-i2v)
|
||||
- Studied framework T2I: tokenization uses `vlm_config.use_system_prompt` which is `false` in the checkpoint (config.json:199) — matches FastVideo's hardcoded `use_system_prompt=False` for all modes (no divergence). Canonical T2I is 960x960 (inputs/omni/t2i.json), single-frame (num_frames=1).
|
||||
- Bug found + fixed (commit `604dc2637`): the stage chose `flow_shift` by task (`3.0 if is_t2i else 10.0`), but the framework picks it purely by the named resolution bucket (`OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256->3.0, 480->5.0, 720/768->10.0; model default resolution "720"). Task-based only matched T2V@720 / T2I@256 by luck; canonical T2I@960x960 is the "720" bucket -> 10.0, so `is_t2i->3.0` was wrong. Replaced with `_flow_shift_for_resolution(h,w)` (longest-side bucketing), applied to all tasks.
|
||||
- Parity: `test_cosmos3_flow_shift_parity.py` checks the mapping vs framework `{VIDEO,IMAGE}_RES_SIZE_INFO` x `_RESOLUTION_SHIFT_DEFAULTS` (8B rows, 20 cases). Also hardened `_image_to_video_tensor` tensor branch to respect the [-1,1] convention (PIL path stays framework-exact).
|
||||
- Example: `examples/inference/basic/basic_cosmos3_t2i_new_api.py` (num_frames=1, 960x960).
|
||||
- Verified on B200 (real weights, 35 steps): coherent red-panda image matching the prompt, flow_shift=10.0. Full suite 118 passed, 0 skipped.
|
||||
- Video core (T2V/I2V/T2I) is now complete and real-weights verified. NEXT: PR2 audio (sound_tokenizer AVAE + audio output) on a new stacked branch.
|
||||
|
||||
## Full-omni parity summary (every component, max / mean abs diff vs framework)
|
||||
|
||||
All run on CPU / float32 (tiny models, framework weights copied in; framework =
|
||||
oracle). `tests/local_tests/cosmos3/`, suite: 150 passed, 0 skipped.
|
||||
|
||||
| Component / pipeline | Test | max | mean |
|
||||
|---|---|---|---|
|
||||
| Scheduler (UniPC flow) | test_cosmos3_scheduler_parity | timesteps 0; sigmas ~1e-8; traj <~1e-6 | ~1e-7 |
|
||||
| DiT (video, unified_3d_mrope) | test_cosmos3_dit_parity_mrope | 0.0 | 0.0 |
|
||||
| Sequence packing (video) | test_cosmos3_packing_parity | 0.0 (exact) | 0.0 |
|
||||
| VAE (Wan2.2) | test_cosmos3_vae_parity | 0.0 | 0.0 |
|
||||
| Denoise / CFG velocity | test_cosmos3_denoise_cfg_parity | <1e-6 | <1e-7 |
|
||||
| flow_shift (resolution) | test_cosmos3_flow_shift_parity | exact | exact |
|
||||
| I2V conditioning (static-repeat) | test_cosmos3_i2v_conditioning_parity | 0.0 | 0.0 |
|
||||
| AVAE sound decoder | test_cosmos3_avae_parity | 0.0 | 0.0 |
|
||||
| DiT sound pathway + packing | test_cosmos3_sound_parity | 0.0 | 0.0 |
|
||||
| t2vs CFG velocity | test_cosmos3_sound_parity | 0.0 | 0.0 |
|
||||
| DiT action pathway + packing | test_cosmos3_action_parity | 0.0 | 0.0 |
|
||||
| action CFG velocity | test_cosmos3_action_parity | 0.0 | 0.0 |
|
||||
| Reasoner prefill logits (text) | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
|
||||
| Reasoner greedy generation | test_cosmos3_reasoning_parity | token-exact | - |
|
||||
| Deepstack reasoner forward | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
|
||||
| vision_encoder (Qwen3-VL ViT) | test_cosmos3_vision_encoder_parity | 0.0 | 0.0 |
|
||||
|
||||
Real-weights pipelines verified on B200 (`examples/inference/basic/basic_cosmos3*_new_api.py`):
|
||||
T2V (1280x704), I2V (cyclist), T2I (960x960 red panda), t2vs (ocean + stereo
|
||||
48kHz audio), text reasoning (greedy == framework). All coherent / prompt-matching.
|
||||
@@ -0,0 +1,71 @@
|
||||
# Cosmos3 local parity workspace
|
||||
|
||||
## Overview
|
||||
|
||||
This workspace tracks the FastVideo Cosmos3 port. Live port state, component matrix,
|
||||
decisions, and blockers live in `PORT_STATUS.md`.
|
||||
|
||||
- **Reference (2026-06-06): official NVIDIA `cosmos-framework` diffusers backend** —
|
||||
`Cosmos3OmniDiffusersPipeline` from the `diffusers-cosmos3` shim — loading the
|
||||
now-public `nvidia/Cosmos3-Nano` checkpoint.
|
||||
- **Scope: full omni** — T2V / I2V / T2I, audio (sound generation), VLM reasoning,
|
||||
and action-conditioning.
|
||||
- The original Tier-A scaffold was written against vllm-omni PR #3454 before official
|
||||
weights were public; it is being repointed to the diffusers reference (see I001/I002
|
||||
in `PORT_STATUS.md`).
|
||||
|
||||
## Reference code
|
||||
|
||||
Primary (official):
|
||||
|
||||
- Local: `cosmos-framework/` (symlink -> `/home/william5lin/FastVideo/cosmos-framework`,
|
||||
commit `003d66d4`); GitHub <https://github.com/NVIDIA/cosmos-framework>
|
||||
- diffusers shim `cosmos-framework/packages/diffusers-cosmos3/diffusers_cosmos3/`:
|
||||
- `pipeline.py` — `Cosmos3OmniDiffusersPipeline`
|
||||
- `transformer.py` — `Cosmos3OmniTransformer`
|
||||
- `sequence_packing.py`
|
||||
- framework model code: `cosmos_framework/model/vfm/mot/cosmos3_vfm_network.py`,
|
||||
`cosmos_framework/model/vfm/omni_mot_model.py`
|
||||
- Installed editable in shared `fv-main`: `diffusers-cosmos3`, `cosmos-framework`
|
||||
(both `--no-deps`).
|
||||
|
||||
Original Tier-A reference (superseded, kept for diffing during repoint):
|
||||
|
||||
- vllm-omni PR #3454 <https://github.com/vllm-project/vllm-omni/pull/3454>, pinned
|
||||
`8536f5b1`, checkout `/home/william5lin/cosmos3-reference`.
|
||||
- The current `conftest.py` + tests still mirror this suite line-by-line.
|
||||
|
||||
## Weight status
|
||||
|
||||
DOWNLOADED (2026-06-06). `nvidia/Cosmos3-Nano` is now public and diffusers-format
|
||||
(the 2026-05-22 `401` is resolved).
|
||||
|
||||
- Local: `official_weights/cosmos3/` (symlink -> main worktree; 33 GiB, 67 files,
|
||||
`model_index.json` present)
|
||||
- Source: `nvidia/Cosmos3-Nano`, default revision; `source_layout=diffusers`,
|
||||
`needs_conversion=no`
|
||||
- `model_index` class: `Cosmos3OmniDiffusersPipeline` (diffusers 0.37.1)
|
||||
- Token: not required (public repo)
|
||||
|
||||
Components (from `model_index.json`): `transformer` (`Cosmos3OmniTransformer`),
|
||||
`vae` (`AutoencoderKLWan`), `scheduler` (`UniPCMultistepScheduler`),
|
||||
`text_tokenizer` (`Qwen2TokenizerFast`), `vision_encoder` (`Qwen3VLVisionModel`),
|
||||
`sound_tokenizer` (`Cosmos3AVAEAudioTokenizer`).
|
||||
|
||||
## Running the Tier-A scaffold
|
||||
|
||||
```bash
|
||||
PYTHONPATH=/home/william5lin/FastVideo_cosmos3_port \
|
||||
python -m pytest tests/local_tests/cosmos3/ -q
|
||||
```
|
||||
|
||||
NOTE: as of 2026-06-06 these report `15 skipped` because the shared `fv-main` env's
|
||||
editable `fastvideo` resolves to the MAIN worktree (a PEP660 finder overrides
|
||||
`PYTHONPATH`), so the worktree's cosmos3 modules are not importable. Tracked as E001
|
||||
in `PORT_STATUS.md`.
|
||||
|
||||
## SSIM placeholder
|
||||
|
||||
No SSIM references seeded yet. Add SSIM coverage only after a FastVideo inference path
|
||||
can load the Cosmos3 weights and generate stable T2V/I2V/T2I outputs. Audio quality
|
||||
uses a separate metric (not SSIM); see `PORT_STATUS.md` Q003.
|
||||
@@ -0,0 +1,214 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared fixtures for the Cosmos3 native-pipeline local tests.
|
||||
|
||||
These fixtures build the FastVideo-native Cosmos3 pipeline
|
||||
(``fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline.Cosmos3OmniDiffusersPipeline``)
|
||||
via ``__new__`` and wire it with tiny stub components so the runtime call graph
|
||||
(sequential CFG, condition-frame masking, mode dispatch) can be exercised on CPU
|
||||
without real weights or ``cosmos_framework``.
|
||||
|
||||
The stub transformer implements the native DiT's packed-input contract
|
||||
(``{"preds_vision": [[1, C, T, H, W], ...]}``) and records, per call, the first
|
||||
``text_ids`` token so tests can assert the cond/uncond pass order.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.models.schedulers.scheduling_unipc_multistep import (
|
||||
UniPCMultistepScheduler,
|
||||
)
|
||||
from torch import nn
|
||||
|
||||
_LATENT_CHANNEL = 16
|
||||
_LATENT_PATCH_SIZE = 2
|
||||
_SPATIAL_FACTOR = 8
|
||||
_TEMPORAL_FACTOR = 4
|
||||
|
||||
|
||||
def pytest_configure(config: pytest.Config) -> None:
|
||||
"""Register the ``local`` marker used by sibling test files."""
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"local: marker for local-only parity/scaffold tests (skipped in CI)",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stub transformer: records cond/uncond call order; bounded preds_vision.
|
||||
# ---------------------------------------------------------------------------
|
||||
class StubCosmos3Transformer(nn.Module):
|
||||
"""Records each forward's first ``text_ids`` token + returns preds_vision.
|
||||
|
||||
``preds_vision`` is keyed by the first text token (so the conditional and
|
||||
unconditional passes return different velocities) and is zero on
|
||||
conditioning frames, matching the real DiT's unpatchify output.
|
||||
"""
|
||||
|
||||
def __init__(self, latent_channel: int = _LATENT_CHANNEL) -> None:
|
||||
super().__init__()
|
||||
self.latent_channel = latent_channel
|
||||
self.embed_tokens = nn.Embedding(64, 8)
|
||||
self.calls: list[dict[str, Any]] = []
|
||||
|
||||
def forward(self, **kwargs: Any) -> dict[str, Any]:
|
||||
token_ids = kwargs["text_ids"]
|
||||
token = int(token_ids.reshape(-1)[0].item()) if token_ids.numel() else 0
|
||||
self.calls.append({"token": token, "kwargs": dict(kwargs)})
|
||||
scale = 0.01 * (1.0 + (token % 7))
|
||||
preds: list[torch.Tensor] = []
|
||||
for latent, _shape, nfi in zip(kwargs["vision_tokens"], kwargs["vision_token_shapes"],
|
||||
kwargs["vision_noisy_frame_indexes"]):
|
||||
lat = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
|
||||
out = torch.zeros_like(lat)
|
||||
if nfi.numel() > 0:
|
||||
out[:, nfi] = scale * torch.tanh(lat[:, nfi])
|
||||
preds.append(out.unsqueeze(0))
|
||||
return {"preds_vision": preds}
|
||||
|
||||
|
||||
class _StubLatentDist:
|
||||
|
||||
def __init__(self, latents: torch.Tensor) -> None:
|
||||
self._latents = latents
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self._latents
|
||||
|
||||
|
||||
class StubCosmos3VAE:
|
||||
"""Deterministic VAE shaped by the Wan scale factors."""
|
||||
|
||||
def __init__(self, z_dim: int = _LATENT_CHANNEL) -> None:
|
||||
self.config = SimpleNamespace(
|
||||
z_dim=z_dim,
|
||||
scale_factor_temporal=_TEMPORAL_FACTOR,
|
||||
scale_factor_spatial=_SPATIAL_FACTOR,
|
||||
latents_mean=[0.0] * z_dim,
|
||||
latents_std=[1.0] * z_dim,
|
||||
)
|
||||
|
||||
def encode(self, video: torch.Tensor):
|
||||
b, _c, t, h, w = video.shape
|
||||
lt = (t - 1) // self.config.scale_factor_temporal + 1
|
||||
lh = h // self.config.scale_factor_spatial
|
||||
lw = w // self.config.scale_factor_spatial
|
||||
return _StubLatentDist(torch.ones(b, self.config.z_dim, lt, lh, lw, dtype=video.dtype, device=video.device))
|
||||
|
||||
def decode(self, z: torch.Tensor):
|
||||
b, _c, lt, lh, lw = z.shape
|
||||
t = (lt - 1) * self.config.scale_factor_temporal + 1
|
||||
h = lh * self.config.scale_factor_spatial
|
||||
w = lw * self.config.scale_factor_spatial
|
||||
sig = torch.nan_to_num(torch.tanh(z[:, :1, :1, :1, :1])).reshape(b, 1, 1, 1, 1)
|
||||
return torch.clamp(torch.zeros(b, 3, t, h, w, dtype=z.dtype, device=z.device) + sig, -1.0, 1.0)
|
||||
|
||||
|
||||
class StubQwen2Tokenizer:
|
||||
"""Qwen2-shaped chat tokenizer stub (special tokens + chat template)."""
|
||||
|
||||
eos_token_id = 62
|
||||
_SPECIAL = {"<|vision_start|>": 60, "<|vision_end|>": 61}
|
||||
|
||||
def convert_tokens_to_ids(self, token: str) -> int:
|
||||
return self._SPECIAL[token]
|
||||
|
||||
def apply_chat_template(self, conversations, *, tokenize=True, add_generation_prompt=True, add_vision_id=False):
|
||||
user = next((c["content"] for c in conversations if c["role"] == "user"), "")
|
||||
n = max(1, min(8, len(user) % 8 + 1))
|
||||
return [10 + (i % 40) for i in range(n)]
|
||||
|
||||
|
||||
def make_scheduler(flow_shift: float = 10.0) -> UniPCMultistepScheduler:
|
||||
return UniPCMultistepScheduler(
|
||||
num_train_timesteps=1000,
|
||||
solver_order=2,
|
||||
prediction_type="flow_prediction",
|
||||
use_flow_sigmas=True,
|
||||
flow_shift=flow_shift,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pipeline factory — builds the native pipeline via __new__ + stub modules.
|
||||
# ---------------------------------------------------------------------------
|
||||
@pytest.fixture
|
||||
def make_cosmos3_pipeline():
|
||||
"""Return a factory building the native Cosmos3 pipeline wired with stubs."""
|
||||
|
||||
def _make(**overrides: Any):
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # noqa: F401
|
||||
Cosmos3OmniDiffusersPipeline, )
|
||||
|
||||
pipe = Cosmos3OmniDiffusersPipeline.__new__(Cosmos3OmniDiffusersPipeline)
|
||||
scheduler = make_scheduler()
|
||||
pipe.modules = {
|
||||
"transformer": StubCosmos3Transformer(),
|
||||
"vae": StubCosmos3VAE(),
|
||||
"scheduler": scheduler,
|
||||
"text_tokenizer": StubQwen2Tokenizer(),
|
||||
}
|
||||
pipe.scheduler = scheduler
|
||||
pipe._base_scheduler_config = scheduler.config
|
||||
pipe._current_flow_shift = float(scheduler.config.flow_shift)
|
||||
pipe._engine_init_flow_shift = 10.0
|
||||
for key, value in overrides.items():
|
||||
setattr(pipe, key, value)
|
||||
return pipe
|
||||
|
||||
return _make
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def make_cosmos3_stage():
|
||||
"""Return a factory building a ``Cosmos3DenoisingStage`` bound to a pipeline."""
|
||||
|
||||
def _make(pipeline):
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
|
||||
|
||||
return Cosmos3DenoisingStage(
|
||||
transformer=pipeline.modules["transformer"],
|
||||
scheduler=pipeline.modules["scheduler"],
|
||||
vae=pipeline.modules["vae"],
|
||||
tokenizer=pipeline.modules["text_tokenizer"],
|
||||
pipeline=pipeline,
|
||||
)
|
||||
|
||||
return _make
|
||||
|
||||
|
||||
def make_forward_batch(*, num_frames: int, height: int, width: int, image: Any = None, **overrides: Any):
|
||||
"""Build a tiny ``ForwardBatch`` for the Cosmos3 stage."""
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
values: dict[str, Any] = dict(
|
||||
data_type="video",
|
||||
prompt="a calm ocean at sunrise",
|
||||
negative_prompt="",
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=2,
|
||||
guidance_scale=6.0,
|
||||
generator=torch.Generator("cpu").manual_seed(0),
|
||||
preprocessed_image=image,
|
||||
)
|
||||
values.update(overrides)
|
||||
return ForwardBatch(**values)
|
||||
|
||||
|
||||
def make_fastvideo_args():
|
||||
"""Build minimal ``fastvideo_args`` (only ``pipeline_config`` is read)."""
|
||||
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
|
||||
|
||||
cfg = Cosmos3Config()
|
||||
arch = cfg.dit_config.arch_config
|
||||
arch.latent_channel = _LATENT_CHANNEL
|
||||
arch.latent_patch_size = _LATENT_PATCH_SIZE
|
||||
arch.temporal_compression_factor = _TEMPORAL_FACTOR
|
||||
arch.enable_fps_modulation = False
|
||||
return SimpleNamespace(pipeline_config=cfg)
|
||||
@@ -0,0 +1,166 @@
|
||||
# Cosmos3 → FastVideo port — feedback (pitfalls, issues, difficulties)
|
||||
|
||||
Retrospective on porting the full **NVIDIA Cosmos3-Nano** omni world model
|
||||
(video / audio / action generation + text & image reasoning) into FastVideo.
|
||||
Methodology: framework-only reference, native FastVideo port, a bit-exact
|
||||
framework-parity test per component, then real-weights verification. Every
|
||||
modality landed bit-exact (see `PORT_STATUS.md` "Full-omni parity summary").
|
||||
|
||||
This doc records what bit, so the next omni/world-model port (and the `/add-model`
|
||||
skill) can avoid the same traps.
|
||||
|
||||
---
|
||||
|
||||
## 1. The checkpoint's config does NOT describe the runtime — verify against the framework
|
||||
|
||||
The single biggest time sink. The HF checkpoint is "diffusers format", which led
|
||||
to two silent traps:
|
||||
|
||||
- **Scheduler (caused an all-black video).** `scheduler/scheduler_config.json`
|
||||
is a diffusers `UniPCMultistepScheduler` config carrying
|
||||
`use_karras_sigmas=true`, `sigma_min/sigma_max`, a beta schedule, etc. But the
|
||||
framework actually samples with a *flow-matching* `FlowUniPCMultistepScheduler`
|
||||
(shift + num_train_timesteps only). FastVideo's vendored UniPC checks
|
||||
`use_karras_sigmas` **before** `use_flow_sigmas`, so it built diffusion (beta)
|
||||
sigmas → `scheduler.step` → **NaN latents → 3 KB black mp4**. The DiT/CFG
|
||||
velocity was perfectly clean; only the scheduler diverged.
|
||||
- Fix: coerce the loaded config to the flow setup in `initialize_pipeline`.
|
||||
- **Lesson:** treat the checkpoint's generic-format config as *lossy*. Find how
|
||||
the framework actually instantiates the component and match THAT, not the JSON.
|
||||
|
||||
- **`flow_shift` is resolution-based, not task-based.** Natural assumption:
|
||||
"T2I uses a small shift, video a large one." Reality: the framework keys the
|
||||
UniPC shift purely off the named resolution bucket
|
||||
(`_RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256→3, 480→5, 720/768→10). The
|
||||
task-based heuristic only *coincidentally* matched (T2V@720, T2I@256); canonical
|
||||
T2I is 960×960 (the "720" bucket → 10), so `is_t2i→3.0` was wrong.
|
||||
|
||||
## 2. A "parity test" that compares two copies of the wrong thing proves nothing
|
||||
|
||||
The original denoise/CFG test imported **diffusers** `UniPCMultistepScheduler` and
|
||||
used it on BOTH the "oracle" and "FastVideo" sides. So the scheduler was never
|
||||
actually compared against the framework — which is exactly why the black-video
|
||||
scheduler bug sailed through a green test suite.
|
||||
- **Lesson:** the oracle side of every parity test MUST be the official framework
|
||||
object, never a second instance of the unit under test. After writing a parity
|
||||
test, ask: "if the framework were wrong here, would this test fail?"
|
||||
|
||||
## 3. Temporal-VAE conditioning: static-repeat vs zero-fill (silent corruption)
|
||||
|
||||
I2V/T2I condition on the input image. The framework
|
||||
(`build_conditioned_video_batch`) fills frame 0 with the image and **repeats the
|
||||
last conditioning frame across the whole clip** (a static video) before
|
||||
VAE-encoding. The first native cut **zero-filled** the non-condition frames.
|
||||
Because the Wan VAE is temporal (4× compression), latent-frame-0 (the kept-clean
|
||||
condition frame) depends on several *pixel* frames — so zero-filling produced a
|
||||
*wrong* conditioning latent. This is the kind of bug that doesn't crash and can
|
||||
even look plausible at a glance.
|
||||
- **Lesson:** when a conditioning latent feeds a temporal autoencoder, trace the
|
||||
temporal receptive field; "only frame 0 matters" is false under temporal conv.
|
||||
|
||||
## 4. Checkpoint param names ≠ framework module structure (three naming conventions)
|
||||
|
||||
For the DiT there were **three** namings to bridge: framework-native
|
||||
(`Cosmos3VFMNetwork`: `language_model.model.layers.*`, `vae2llm`, `q_proj_moe_gen`,
|
||||
…), the diffusers checkpoint on disk (`layers.*.to_q`, `add_q_proj`, `proj_in`,
|
||||
…), and the FastVideo DiT. The weight map crosses (framework)→(FastVideo) for
|
||||
parity and (checkpoint)→(FastVideo) for loading.
|
||||
|
||||
The **sound tokenizer** was the sharpest example: the checkpoint is **decoder-only**
|
||||
in diffusers `AutoencoderOobleck` naming (`decoder.conv1`, `block.N.conv_t1`,
|
||||
`res_unitM`, `snake1`) — but with **`SnakeBeta`** (learned alpha *and* beta,
|
||||
logscale), NOT diffusers' alpha-only `Snake1d`. So neither "use diffusers
|
||||
AutoencoderOobleck" nor "port the framework `LatentAutoEncoderV2` Sequential
|
||||
module" matched the on-disk keys.
|
||||
- **Lesson:** dump `safetensors` keys + shapes for every sub-checkpoint *first*.
|
||||
The naming reveals which existing native module (if any) already matches.
|
||||
|
||||
## 5. A "matching" native module can still differ on an untested config path
|
||||
|
||||
FastVideo already had a native `OobleckVAE` (Stable Audio) whose decoder matched
|
||||
the Cosmos3 sound decoder bit-for-bit — except `OobleckDecoderBlock.conv_t1`
|
||||
omitted `output_padding = stride % 2`. That omission is a **no-op for Stable
|
||||
Audio's even strides** [2,4,4,8,8], so it had never mattered; Cosmos3 has an
|
||||
**odd** stride (5), where the framework's `output_padding=1` makes the decode one
|
||||
sample longer per odd-stride block (parity diverged 60 vs 59 samples).
|
||||
- **Lesson:** reusing a native module is great, but re-run parity on the *new*
|
||||
model's config — shared code can hide config-specific divergences.
|
||||
|
||||
## 6. Loader / registry plumbing the checkpoint format forces
|
||||
|
||||
- **DiT class alias.** `model_index.json` names the DiT `Cosmos3OmniTransformer`
|
||||
(the diffusers shim class); the registry normalized unknown classes to a
|
||||
generic `TransformersModel`. Needed an explicit registry alias
|
||||
`Cosmos3OmniTransformer → Cosmos3VFMTransformer`.
|
||||
- **Tokenizer module name.** `model_index.json` calls the Qwen2 tokenizer
|
||||
`text_tokenizer` (not `tokenizer`); the component loader had no mapping for that
|
||||
key and tried to load it as a model.
|
||||
- **Scheduler config schema drift.** The vendored UniPC predates
|
||||
`shift_terminal` / `sigma_min` / `sigma_max`; constructing it with the raw
|
||||
checkpoint config crashes on the unexpected kwargs. Filter to the class's
|
||||
`__init__` params (mirroring diffusers `from_config`).
|
||||
- **Meta-device load + non-persistent buffers.** `rotary_emb.inv_freq` is derived
|
||||
from `rope_theta` and is non-persistent (absent from the checkpoint), so after
|
||||
the meta-device FSDP load it stays on the `meta` device → needs a
|
||||
`materialize_non_persistent_buffers` hook to recompute it on the real device.
|
||||
- **dtype boundaries.** Noise/VAE latents arrive fp32; the model runs bf16.
|
||||
Needed explicit casts at `proj_in` and the timestep embedder (no-ops in the
|
||||
fp32 parity tests, required at inference).
|
||||
- **device in packing.** The packer builds ids/positions on CPU; `to_dit_kwargs`
|
||||
must move every tensor to the model device before the forward.
|
||||
|
||||
## 7. The omni model is a Mixture-of-Transformers — modality bookkeeping is the work
|
||||
|
||||
The backbone is a dual-pathway MoT: **und** (causal text) + **gen** (full-attention
|
||||
vision/sound/action). Once the video path worked, each extra modality was the same
|
||||
*shape* of work (a proj-in + modality embed + timestep-scatter encode, a proj-out
|
||||
decode, packing, a CFG-velocity slice) but with per-modality quirks:
|
||||
- sound/action **share the vision "full" split** (preserving the causal+full
|
||||
2-split invariant); the combined flat latent is `[vision | action | sound]` in
|
||||
that order (must match the framework's per-sample concat).
|
||||
- sound MRoPE uses `start_frame_offset=0` (parallel to vision); action uses
|
||||
`start_frame_offset=1`; both at the vision temporal offset, tcf=1, and do NOT
|
||||
advance the offset.
|
||||
- action is **domain-aware** (`DomainAwareLinear`: per-embodiment weight/bias via
|
||||
`nn.Embedding`, indexed by a per-token domain id).
|
||||
- the unpack already zeros clean frames, so the per-step velocity masking is
|
||||
defensive (but kept, to mirror the framework exactly).
|
||||
- **Lesson:** build the first modality (vision) with clean seams for "a modality"
|
||||
and the rest fall out; spend the care on the packing layout + MRoPE offsets,
|
||||
which are the only per-modality novelties.
|
||||
|
||||
## 8. Reasoning reused more than expected; the encoder is just transformers
|
||||
|
||||
- **Text reasoning** needed *no new model code*: it's the und (causal) pathway +
|
||||
`embed_tokens`/`norm`/`lm_head`, all already in the DiT. A text-only forward +
|
||||
`lm_head` is token-for-token identical to the framework reasoner.
|
||||
- **vision_encoder** is a stock `transformers.Qwen3VLVisionModel` (the framework
|
||||
ships its own *copy* of the same class); reusing transformers' (like the Qwen2
|
||||
tokenizer) is bit-exact vs the framework — re-porting a 27-layer ViT would have
|
||||
been wasted effort.
|
||||
- **deepstack** (image-conditioned reasoning) is the one new native piece: inject
|
||||
the 3 vision-encoder deepstack features into the first 3 text layers at the
|
||||
image-token positions.
|
||||
- **Lesson:** before porting a big sub-model, check whether it's literally a
|
||||
stock library class — and whether an existing in-repo module already implements
|
||||
it (audio decoder + vision encoder were both "already there").
|
||||
|
||||
## 9. Running the framework as a CPU parity oracle
|
||||
|
||||
The framework's attention path is flash/natten (CUDA-only). Parity tests run on
|
||||
CPU/float32 via an SDPA monkey-patch (`test_cosmos3_reference_forward._apply_sdpa_patches`).
|
||||
A couple of framework helpers also can't import headless (`cosmos_framework.inference.args`
|
||||
pulls `multistorageclient`), so a constant or two is mirrored in the test with a
|
||||
cited source rather than imported.
|
||||
- **Lesson:** budget for "make the oracle runnable on CPU" — monkeypatch attention,
|
||||
build tiny configs, and accept a small amount of mirrored constants when a
|
||||
framework module won't import in isolation.
|
||||
|
||||
## 10. What made it tractable
|
||||
|
||||
- Tiny CPU/fp32 models + copy-framework-weights-in + bit-exact compare, per
|
||||
component, is a fast and decisive loop (max=mean=0.0 or it's wrong).
|
||||
- A persistent `PORT_STATUS.md` (resumable state, issues, decisions) survived
|
||||
several context resets.
|
||||
- Stacked PR branches (one per modality) kept each parity-verified increment
|
||||
reviewable and the chain bisectable.
|
||||
@@ -0,0 +1,271 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 action pathway vs the framework.
|
||||
|
||||
Covers the action (multi-embodiment world-model) modality at the DiT level:
|
||||
|
||||
* **action packing** — native ``pack_cosmos3_video_sequence`` with a
|
||||
``Cosmos3ActionItem`` vs framework ``pack_input_sequence`` with
|
||||
``has_action``: action tokens share the vision "full" split, with ``(T,)``
|
||||
shapes, a ``(T,1)`` condition mask, and 3D-MRoPE temporal positions at the
|
||||
vision offset with ``start_frame_offset=1`` (parallel to vision); and
|
||||
* **DiT action forward** — the dormant domain-aware ``action_proj_in`` /
|
||||
``action_proj_out`` (``DomainAwareLinear``) + ``action_modality_embed`` heads,
|
||||
now activated, with a per-token embodiment ``domain_id``.
|
||||
|
||||
Framework model + pack is the parity ORACLE (CPU/float32 via SDPA monkey-patch).
|
||||
We assert the native packer matches the framework field-by-field, then that
|
||||
``preds_vision`` AND ``preds_action`` match the framework forward.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_action_parity.py -q -s
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
_framework_to_fastvideo_state_dict,
|
||||
)
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_ACTION_DIM,
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
_apply_sdpa_patches()
|
||||
|
||||
_SPECIAL_TOKENS = {"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62}
|
||||
|
||||
|
||||
def _copy_weights_with_action(vfm, dit) -> None:
|
||||
"""Copy backbone + vision weights AND the domain-aware action heads."""
|
||||
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
|
||||
src = dict(vfm.named_parameters())
|
||||
mapped["action_proj_in.fc.weight"] = src["action2llm.fc.weight"].detach().clone()
|
||||
mapped["action_proj_in.bias.weight"] = src["action2llm.bias.weight"].detach().clone()
|
||||
mapped["action_proj_out.fc.weight"] = src["llm2action.fc.weight"].detach().clone()
|
||||
mapped["action_proj_out.bias.weight"] = src["llm2action.bias.weight"].detach().clone()
|
||||
mapped["action_modality_embed"] = src["action_modality_embed"].detach().clone()
|
||||
dst = dict(dit.named_parameters())
|
||||
with torch.no_grad():
|
||||
for name, tensor in mapped.items():
|
||||
assert name in dst, f"DiT missing param {name!r}"
|
||||
assert dst[name].shape == tensor.shape, f"shape mismatch {name}"
|
||||
dst[name].copy_(tensor.to(dst[name].dtype))
|
||||
|
||||
|
||||
def _framework_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep,
|
||||
is_image_batch):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
gen = GenerationDataClean(
|
||||
batch_size=1,
|
||||
is_image_batch=is_image_batch,
|
||||
x0_tokens_vision=[vision],
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1],
|
||||
x0_tokens_action=[action],
|
||||
fps_action=None,
|
||||
action_domain_id=[torch.tensor([domain_id], dtype=torch.long)],
|
||||
)
|
||||
plans = [SequencePlan(
|
||||
has_text=True, has_vision=True, has_action=True,
|
||||
condition_frame_indexes_vision=list(cond_vision),
|
||||
condition_frame_indexes_action=list(cond_action),
|
||||
)]
|
||||
ps = pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(text_ids)],
|
||||
gen_data_clean=gen,
|
||||
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
# The framework sets action.domain_id on the packed sequence from
|
||||
# gen_data_clean inside the model (_get_velocity); mirror that for the oracle.
|
||||
if ps.action is not None:
|
||||
ps.action.domain_id = [torch.tensor([domain_id], dtype=torch.long)]
|
||||
return ps
|
||||
|
||||
|
||||
def _fastvideo_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep):
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3ActionItem,
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
|
||||
samples = [Cosmos3SampleInputs(
|
||||
text_ids=list(text_ids),
|
||||
vision=Cosmos3VisionItem(latent=vision, condition_frame_indexes=list(cond_vision)),
|
||||
action=Cosmos3ActionItem(latent=action, condition_frame_indexes=list(cond_action), domain_id=domain_id),
|
||||
timestep=float(timestep),
|
||||
)]
|
||||
return pack_cosmos3_video_sequence(
|
||||
samples, _SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE, include_end_of_generation_token=False,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _fv_inputs_with_action(ps) -> dict:
|
||||
kw = _fastvideo_inputs_from_packed_seq(ps)
|
||||
a = ps.action
|
||||
kw.update(
|
||||
action_tokens=list(a.tokens),
|
||||
action_token_shapes=[tuple(x) for x in a.token_shapes],
|
||||
action_sequence_indexes=a.sequence_indexes,
|
||||
action_timesteps=a.timesteps,
|
||||
action_mse_loss_indexes=a.mse_loss_indexes,
|
||||
action_noisy_frame_indexes=list(a.noisy_frame_indexes),
|
||||
action_domain_id=list(a.domain_id),
|
||||
)
|
||||
return kw
|
||||
|
||||
|
||||
def _diffs(a, b):
|
||||
d = (a - b).abs()
|
||||
return d.max().item(), d.mean().item()
|
||||
|
||||
|
||||
# (grid_t, lh, lw, action_t, n_text, cond_vision, cond_action, domain_id)
|
||||
_CASES = [
|
||||
pytest.param(2, 4, 4, 6, 4, [], [], 0, id="a2v_2x2x2_act6_dom0"),
|
||||
pytest.param(3, 8, 4, 9, 5, [], [], 7, id="a2v_3x4x2_act9_dom7"),
|
||||
pytest.param(2, 4, 4, 5, 5, [0], [0], 3, id="ai2v_cond_act5_dom3"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3ActionParity:
|
||||
|
||||
def _build(self, num_layers=2, seed_model=42):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers, action_gen=True)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights_with_action(vfm, dit)
|
||||
return vfm, dit
|
||||
|
||||
def _make_inputs(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom, seed=7):
|
||||
torch.manual_seed(seed)
|
||||
return dict(
|
||||
text_ids=torch.randint(0, 60, (n_text,)).tolist(),
|
||||
vision=torch.randn(1, _LATENT_CHANNEL, grid_t, lh, lw),
|
||||
action=torch.randn(act_t, _ACTION_DIM), # [T, D]
|
||||
cond_vision=cond_v, cond_action=cond_a, domain_id=dom, timestep=500.0,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_packing_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
|
||||
fw = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
|
||||
fv = _fastvideo_pack_action(**ins)
|
||||
assert fv.split_lens == list(fw.split_lens), f"split_lens fv={fv.split_lens} fw={list(fw.split_lens)}"
|
||||
assert fv.attn_modes == list(fw.attn_modes)
|
||||
assert int(fv.sequence_length) == int(fw.sequence_length)
|
||||
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
|
||||
a = fw.action
|
||||
torch.testing.assert_close(fv.action_sequence_indexes, a.sequence_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert fv.action_token_shapes == [tuple(x) for x in a.token_shapes]
|
||||
torch.testing.assert_close(fv.action_timesteps.to(torch.float32), a.timesteps.to(torch.float32))
|
||||
torch.testing.assert_close(fv.action_mse_loss_indexes, a.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
|
||||
for x, y in zip(fv.action_noisy_frame_indexes, a.noisy_frame_indexes):
|
||||
torch.testing.assert_close(x.to(torch.long), y.to(torch.long), rtol=0, atol=0)
|
||||
print(f"\n[action_packing {grid_t}x{lh}x{lw} act={act_t} dom={dom}] position_ids + action fields exact")
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_dit_forward_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
vfm, dit = self._build()
|
||||
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
|
||||
fw_pack = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
|
||||
fv_pack = _fastvideo_pack_action(**ins)
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=fw_pack)
|
||||
fv_out = dit(**fv_pack.to_dit_kwargs())
|
||||
fv_on_fw = dit(**_fv_inputs_with_action(fw_pack))
|
||||
pv_mx, pv_mn = _diffs(fv_out["preds_vision"][0], fw_out["preds_vision"][0])
|
||||
pa_mx, pa_mn = _diffs(fv_out["preds_action"][0], fw_out["preds_action"][0])
|
||||
paf_mx, paf_mn = _diffs(fv_on_fw["preds_action"][0], fw_out["preds_action"][0])
|
||||
print(f"\n[action_dit {grid_t}x{lh}x{lw} act={act_t} dom={dom}] "
|
||||
f"preds_vision max={pv_mx:.3e} mean={pv_mn:.3e} | "
|
||||
f"preds_action max={pa_mx:.3e} mean={pa_mn:.3e} | "
|
||||
f"preds_action(fwpack) max={paf_mx:.3e} mean={paf_mn:.3e}")
|
||||
assert fv_out["preds_action"][0].shape == fw_out["preds_action"][0].shape
|
||||
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
|
||||
torch.testing.assert_close(fv_out["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
|
||||
torch.testing.assert_close(fv_on_fw["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_cfg_velocity_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
"""Combined [vision|action] sequential-CFG velocity (action pipeline glue)
|
||||
matches a framework-DiT oracle."""
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3ActionSpec,
|
||||
Cosmos3VisionSpec,
|
||||
cosmos3_get_cfg_velocity,
|
||||
)
|
||||
|
||||
vfm, dit = self._build()
|
||||
vlat_shape = (_LATENT_CHANNEL, grid_t, lh, lw)
|
||||
action_shape = (act_t, _ACTION_DIM)
|
||||
torch.manual_seed(3)
|
||||
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
|
||||
vis_numel = int(torch.tensor(vlat_shape).prod())
|
||||
act_numel = int(torch.tensor(action_shape).prod())
|
||||
flat = torch.randn(vis_numel + act_numel)
|
||||
guidance, ts = 6.0, 500.0
|
||||
|
||||
def _fw_velocity(ids):
|
||||
vision = flat[:vis_numel].reshape(vlat_shape).unsqueeze(0)
|
||||
action = flat[vis_numel:].reshape(action_shape)
|
||||
ps = _framework_pack_action(text_ids=ids, vision=vision, action=action, cond_vision=cond_v,
|
||||
cond_action=cond_a, domain_id=dom, timestep=ts,
|
||||
is_image_batch=(grid_t == 1))
|
||||
with torch.no_grad():
|
||||
out = vfm(packed_seq=ps)
|
||||
pv = out["preds_vision"][0].squeeze(0) # [C,T,H,W] (zero on clean)
|
||||
pa = out["preds_action"][0] # [T,D] (zero on clean)
|
||||
return torch.cat([pv.reshape(-1), pa.reshape(-1)])
|
||||
|
||||
fw_cond, fw_uncond = _fw_velocity(cond_ids), _fw_velocity(uncond_ids)
|
||||
fw_v = fw_uncond + guidance * (fw_cond - fw_uncond)
|
||||
fv_v = cosmos3_get_cfg_velocity(
|
||||
transformer=dit, flat_latent=flat, timestep=torch.tensor([ts]), guidance=guidance,
|
||||
specs=[Cosmos3VisionSpec(shape=vlat_shape, condition_frame_indexes=list(cond_v))],
|
||||
action_specs=[Cosmos3ActionSpec(shape=action_shape, condition_frame_indexes=list(cond_a), domain_id=dom)],
|
||||
cond_token_ids=cond_ids, uncond_token_ids=uncond_ids,
|
||||
special_tokens=_SPECIAL_TOKENS, latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
|
||||
)
|
||||
assert fv_v.shape == fw_v.shape, f"shape fv={fv_v.shape} fw={fw_v.shape}"
|
||||
mx, mn = _diffs(fv_v, fw_v)
|
||||
print(f"\n[action_cfg_velocity {grid_t}x{lh}x{lw} act={act_t} dom={dom}] max={mx:.3e} mean={mn:.3e}")
|
||||
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,131 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 sound decoder vs the framework AVAE.
|
||||
|
||||
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
|
||||
checkpoint is **decoder-only** (``decoder.*``; the SpectrogramConvNeXt encoder is
|
||||
not exported) in diffusers ``AutoencoderOobleck`` naming, but with **SnakeBeta**
|
||||
activations (alpha+beta, logscale) and ``weight_g``/``weight_v`` weight-norm —
|
||||
i.e. exactly FastVideo's existing native ``OobleckVAE`` decoder
|
||||
(``fastvideo/models/vaes/oobleck.py``). t2vs only needs DECODE (generate sound
|
||||
latents -> waveform), so this pins the decoder.
|
||||
|
||||
The framework decoder
|
||||
(``cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models.OobleckDecoder``,
|
||||
``nn.Sequential`` naming, ``output_padding=stride%2`` on the transpose convs) is
|
||||
the parity ORACLE. We build a tiny framework decoder, map its weights into the
|
||||
FastVideo decoder (Sequential -> conv1/block.N/res_unitM/snake1/conv2), and
|
||||
assert bit-exact decode. Strides include an ODD value (5, as in the real config
|
||||
``[2,4,5,6,8]``) to exercise the ``output_padding`` path that diverged before.
|
||||
|
||||
CPU / float32. Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_avae_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
_fw_models = pytest.importorskip(
|
||||
"cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
from cosmos_framework.model.vfm.tokenizers.audio.avae_utils.env import ( # noqa: E402
|
||||
AttrDict,
|
||||
)
|
||||
|
||||
from fastvideo.models.vaes.oobleck import OobleckDecoder as FvOobleckDecoder # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
FwOobleckDecoder = _fw_models.OobleckDecoder
|
||||
|
||||
|
||||
def _framework_decoder(dec_dim, vocoder_input_dim, dec_c_mults, dec_strides):
|
||||
"""Framework OobleckDecoder (the parity oracle), non-causal / no-antialias."""
|
||||
h = AttrDict({
|
||||
"vocoder_input_dim": vocoder_input_dim,
|
||||
"input_channels": 1,
|
||||
"stereo": True, # 2 audio channels
|
||||
"dec_dim": dec_dim,
|
||||
"dec_c_mults": dec_c_mults,
|
||||
"dec_strides": dec_strides,
|
||||
"dec_use_snake": True,
|
||||
"dec_use_nearest_upsample": False,
|
||||
"dec_anti_aliasing": False,
|
||||
"causal": False,
|
||||
"dec_use_tanh_at_final": False,
|
||||
"padding_mode": "zeros",
|
||||
})
|
||||
return FwOobleckDecoder(h).eval()
|
||||
|
||||
|
||||
def _framework_to_fastvideo_decoder_state(fw_decoder, num_blocks):
|
||||
"""Map framework Sequential decoder weights -> FastVideo decoder names.
|
||||
|
||||
framework: layers.0=first conv; layers.{1..K}=OobleckDecoderBlock
|
||||
(.layers.0 snake, .1 conv_t, .{2,3,4} ResidualUnit{.layers.0 snake,
|
||||
.1 conv, .2 snake, .3 conv}); layers.{1+K}=final snake; layers.{2+K}=final conv.
|
||||
FastVideo: conv1; block.{b}.{snake1,conv_t1,res_unit{1,2,3}.{snake1,conv1,snake2,conv2}};
|
||||
snake1; conv2. Snake alpha/beta: framework [C] -> FastVideo [1,C,1].
|
||||
"""
|
||||
out = {}
|
||||
for k, v in fw_decoder.state_dict().items():
|
||||
p = k.split(".")
|
||||
li = int(p[1])
|
||||
if li == 0:
|
||||
nk = "conv1." + ".".join(p[2:])
|
||||
elif li == 1 + num_blocks:
|
||||
nk = "snake1." + ".".join(p[2:])
|
||||
elif li == 2 + num_blocks:
|
||||
nk = "conv2." + ".".join(p[2:])
|
||||
else:
|
||||
b = li - 1
|
||||
sub = int(p[3])
|
||||
if sub == 0:
|
||||
nk = f"block.{b}.snake1." + ".".join(p[4:])
|
||||
elif sub == 1:
|
||||
nk = f"block.{b}.conv_t1." + ".".join(p[4:])
|
||||
else:
|
||||
r = sub - 2 # ResidualUnit index 0..2
|
||||
m = {0: "snake1", 1: "conv1", 2: "snake2", 3: "conv2"}[int(p[5])]
|
||||
nk = f"block.{b}.res_unit{r + 1}.{m}." + ".".join(p[6:])
|
||||
if nk.endswith(".alpha") or nk.endswith(".beta"):
|
||||
v = v.reshape(1, -1, 1)
|
||||
out[nk] = v
|
||||
return out
|
||||
|
||||
|
||||
# (dec_dim, vocoder_input_dim, dec_c_mults, dec_strides) — tiny; strides incl odd.
|
||||
_CASES = [
|
||||
pytest.param(4, 8, [1, 2], [5, 2], id="odd_stride5"),
|
||||
pytest.param(6, 8, [1, 2, 4], [2, 5, 6], id="real_stride_pattern_tiny"),
|
||||
pytest.param(4, 4, [1, 2], [4, 8], id="even_strides"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3AVAEParity:
|
||||
|
||||
@pytest.mark.parametrize(("dec_dim", "vin", "cmults", "strides"), _CASES)
|
||||
def test_decode_matches_framework(self, dec_dim, vin, cmults, strides):
|
||||
torch.manual_seed(0)
|
||||
fw = _framework_decoder(dec_dim, vin, cmults, strides)
|
||||
fv = FvOobleckDecoder(
|
||||
channels=dec_dim,
|
||||
input_channels=vin,
|
||||
audio_channels=2,
|
||||
upsampling_ratios=list(reversed(strides)), # framework reverses dec_strides
|
||||
channel_multiples=cmults,
|
||||
).eval()
|
||||
state = _framework_to_fastvideo_decoder_state(fw, num_blocks=len(strides))
|
||||
fv.load_state_dict(state, strict=True) # exact name + shape match
|
||||
|
||||
z = torch.randn(1, vin, 5)
|
||||
with torch.no_grad():
|
||||
a = fw(z)
|
||||
b = fv(z)
|
||||
assert a.shape == b.shape, f"shape: fw={a.shape} fv={b.shape}"
|
||||
max_abs = (a - b).abs().max().item()
|
||||
print(f"\n[avae_decode dim={dec_dim} strides={strides}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(b, a, atol=1e-6, rtol=1e-5)
|
||||
@@ -0,0 +1,342 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 denoise/CFG glue vs the framework.
|
||||
|
||||
The DiT forward and the sequence-packing are already framework-parity-verified
|
||||
(``test_cosmos3_dit_parity*`` / ``test_cosmos3_packing_parity``). This test pins
|
||||
the remaining glue that the native pipeline adds — the SEQUENTIAL classifier-free
|
||||
guidance velocity and one UniPC scheduler step — against the framework math
|
||||
(``diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline.get_cfg_velocity`` /
|
||||
``__call__``):
|
||||
|
||||
* for one denoise step, replicate the framework's ``get_cfg_velocity`` exactly
|
||||
on top of the OFFICIAL ``Cosmos3VFMNetwork`` forward (oracle): a conditional
|
||||
pass (prompt tokens) and an unconditional pass (negative-prompt tokens),
|
||||
each masking the prediction on conditioning frames
|
||||
(``pred * (1 - condition_mask)``), then ``v = uncond + g*(cond - uncond)``;
|
||||
* run FastVideo's :func:`cosmos3_get_cfg_velocity` with the native DiT (the
|
||||
framework weights copied in) + the native packer, and assert the velocity
|
||||
matches the oracle;
|
||||
* take one ``UniPCMultistepScheduler.step`` on each (the actual checkpoint
|
||||
scheduler) and assert the stepped latent matches;
|
||||
* drive :meth:`Cosmos3DenoiseEngine.denoise` for >= 2 steps and assert it
|
||||
equals the manual framework step-by-step loop.
|
||||
|
||||
CPU / float32, via the reference SDPA monkey-patch. The official model is the
|
||||
parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_denoise_cfg_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from .test_cosmos3_dit_parity import _copy_weights # noqa: E402
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
from .test_cosmos3_scheduler_parity import ( # noqa: E402
|
||||
_fastvideo_scheduler,
|
||||
_framework_scheduler,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Tiny special tokens (< tiny vocab_size=64), video path appends eos + sog.
|
||||
_SPECIAL_TOKENS = {
|
||||
"start_of_generation": 60,
|
||||
"end_of_generation": 61,
|
||||
"eos_token_id": 62,
|
||||
}
|
||||
|
||||
# Cosmos3 video flow_shift; framework scheduler is the parity oracle, FastVideo's
|
||||
# vendored UniPC (flow config) is the unit under test.
|
||||
_FLOW_SHIFT = 10.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Framework-oracle CFG velocity (replicates pipeline.get_cfg_velocity math).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _framework_pack(*, text_ids, vision_latent, cond_frames, timestep):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
# vision_latent is [1, C, T, H, W]; temporal dim is axis 2.
|
||||
gen_data_clean = GenerationDataClean(
|
||||
batch_size=1,
|
||||
is_image_batch=(vision_latent.shape[2] == 1),
|
||||
x0_tokens_vision=[vision_latent],
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1],
|
||||
)
|
||||
plans = [SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cond_frames))]
|
||||
return pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(text_ids)],
|
||||
gen_data_clean=gen_data_clean,
|
||||
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _framework_inputs(ps):
|
||||
"""Framework PackedSequence -> framework Cosmos3VFMNetwork forward kwargs."""
|
||||
return dict(packed_seq=ps)
|
||||
|
||||
|
||||
def _framework_cfg_velocity(
|
||||
*,
|
||||
vfm,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
vision_shape: tuple[int, int, int, int],
|
||||
cond_frames: list[int],
|
||||
cond_ids: list[int],
|
||||
uncond_ids: list[int],
|
||||
) -> torch.Tensor:
|
||||
"""Replicate the framework ``get_cfg_velocity`` on the oracle model.
|
||||
|
||||
Single vision item; sequential cond then uncond pass; mask condition
|
||||
frames; ``v = uncond + g*(cond - uncond)``.
|
||||
"""
|
||||
timestep_value = float(timestep.reshape(()).item())
|
||||
vision_latent = flat_latent.reshape(vision_shape) # [C, T, H, W]
|
||||
|
||||
def _run(text_ids: list[int]) -> torch.Tensor:
|
||||
ps = _framework_pack(
|
||||
text_ids=text_ids,
|
||||
# The framework packer expects a 5D [1, C, T, H, W] latent.
|
||||
vision_latent=vision_latent.unsqueeze(0),
|
||||
cond_frames=cond_frames,
|
||||
timestep=timestep_value,
|
||||
)
|
||||
out = vfm(**_framework_inputs(ps))
|
||||
preds = out.get("preds_vision")
|
||||
cond_mask = ps.vision.condition_mask[0] # [T] or [T,1,1]
|
||||
if preds is None:
|
||||
return torch.zeros_like(flat_latent)
|
||||
pred = preds[0].squeeze(0) # [C, T, H, W]
|
||||
keep = (1.0 - cond_mask.reshape(-1, 1, 1)).to(dtype=pred.dtype, device=pred.device)
|
||||
velocity = pred * keep if keep.sum() > 0 else torch.zeros_like(pred)
|
||||
return velocity.reshape(-1)
|
||||
|
||||
cond_v = _run(cond_ids)
|
||||
uncond_v = _run(uncond_ids)
|
||||
return uncond_v + guidance * (cond_v - uncond_v)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cases: T2V (no cond), I2V (cond frame 0), single-frame T2I.
|
||||
# ---------------------------------------------------------------------------
|
||||
_CASES = [
|
||||
pytest.param(2, 4, 4, 6, [], id="t2v_2x2x2"),
|
||||
pytest.param(3, 8, 4, 5, [0], id="i2v_3x4x2_cond0"),
|
||||
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
|
||||
]
|
||||
|
||||
|
||||
def _build_models(num_layers: int = 2, seed_model: int = 42):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
return vfm, dit
|
||||
|
||||
|
||||
def _fastvideo_velocity(dit, *, flat_latent, timestep, guidance, vision_shape, cond_frames, cond_ids, uncond_ids):
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3VisionSpec,
|
||||
cosmos3_get_cfg_velocity,
|
||||
)
|
||||
|
||||
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=list(cond_frames))
|
||||
return cosmos3_get_cfg_velocity(
|
||||
transformer=dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
specs=[spec],
|
||||
cond_token_ids=cond_ids,
|
||||
uncond_token_ids=uncond_ids,
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
class TestCosmos3DenoiseCFGParity:
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_cfg_velocity_matches_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
vfm, dit = _build_models()
|
||||
torch.manual_seed(0)
|
||||
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
timestep = torch.tensor([[500.0]]) # framework expects [1,1]; we reshape to scalar
|
||||
guidance = 6.0
|
||||
|
||||
fw_v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=cond,
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fv_v = _fastvideo_velocity(
|
||||
dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=cond,
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
assert fw_v.shape == fv_v.shape, f"shape: fw={fw_v.shape} fv={fv_v.shape}"
|
||||
max_abs = (fw_v - fv_v).abs().max().item()
|
||||
print(f"\n[cfg_velocity {grid_t}x{latent_h}x{latent_w} cond={cond}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_one_unipc_step_matches_framework(self):
|
||||
"""CFG velocity + one UniPC step: FastVideo == framework math."""
|
||||
vfm, dit = _build_models()
|
||||
grid_t, latent_h, latent_w = 2, 4, 4
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
torch.manual_seed(3)
|
||||
cond_ids = torch.randint(0, 60, (5,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (4,)).tolist()
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
guidance = 6.0
|
||||
|
||||
fw_sched = _framework_scheduler(4, _FLOW_SHIFT)
|
||||
fv_sched = _fastvideo_scheduler(4, _FLOW_SHIFT)
|
||||
t = fw_sched.timesteps[0]
|
||||
|
||||
fw_v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=flat_latent,
|
||||
timestep=t.reshape(1, 1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fw_stepped = fw_sched.step(model_output=fw_v, timestep=t, sample=flat_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
fv_v = _fastvideo_velocity(
|
||||
dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=t.reshape(1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fv_stepped = fv_sched.step(model_output=fv_v, timestep=t, sample=flat_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
max_abs = (fw_stepped - fv_stepped).abs().max().item()
|
||||
print(f"\n[unipc_step] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_stepped, fw_stepped, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_full_denoise_loop_matches_framework(self):
|
||||
"""Cosmos3DenoiseEngine.denoise (>= 2 steps) == framework step-by-step."""
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3DenoiseEngine,
|
||||
Cosmos3VisionSpec,
|
||||
)
|
||||
|
||||
vfm, dit = _build_models()
|
||||
grid_t, latent_h, latent_w = 2, 4, 4
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
torch.manual_seed(5)
|
||||
cond_ids = torch.randint(0, 60, (5,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (4,)).tolist()
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
guidance = 6.0
|
||||
num_steps = 3
|
||||
|
||||
# Manual framework loop (oracle).
|
||||
fw_sched = _framework_scheduler(num_steps, _FLOW_SHIFT)
|
||||
fw_latent = flat_latent.clone()
|
||||
for t in fw_sched.timesteps:
|
||||
v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=fw_latent,
|
||||
timestep=t.reshape(1, 1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fw_latent = fw_sched.step(model_output=v, timestep=t, sample=fw_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
# FastVideo engine loop.
|
||||
fv_sched = _fastvideo_scheduler(num_steps, _FLOW_SHIFT)
|
||||
engine = Cosmos3DenoiseEngine(
|
||||
transformer=dit,
|
||||
scheduler=fv_sched,
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=[])
|
||||
fv_latent = engine.denoise(
|
||||
flat_latent=flat_latent.clone(),
|
||||
timesteps=fv_sched.timesteps,
|
||||
guidance=guidance,
|
||||
specs=[spec],
|
||||
cond_token_ids=cond_ids,
|
||||
uncond_token_ids=uncond_ids,
|
||||
)
|
||||
|
||||
assert fv_latent.shape == fw_latent.shape
|
||||
max_abs = (fw_latent - fv_latent).abs().max().item()
|
||||
print(f"\n[full_denoise {num_steps} steps] final latent max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_latent, fw_latent, atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,256 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 DiT vs official ``Cosmos3VFMNetwork``.
|
||||
|
||||
Builds a tiny official-framework ``Cosmos3VFMNetwork`` AND a tiny FastVideo
|
||||
``Cosmos3VFMTransformer`` from the SAME tiny config, copies the framework
|
||||
weights into the FastVideo DiT via an explicit framework->fastvideo name map,
|
||||
runs BOTH forwards on identical deterministic inputs (CPU / float32), and
|
||||
asserts ``torch.allclose`` on the vision prediction output (``preds_vision``)
|
||||
and the per-token ``last_hidden_state``.
|
||||
|
||||
The official model is the parity ORACLE. It runs on CPU / float32 via the SDPA
|
||||
monkey-patch in ``test_cosmos3_reference_forward`` (flash2/flash3/natten are
|
||||
CUDA-only). The FastVideo DiT runs natively on CPU with plain SDPA.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_dit_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the reference harness's tiny-model builder + SDPA monkey-patch.
|
||||
from .test_cosmos3_reference_forward import ( # noqa: E402
|
||||
_apply_sdpa_patches,
|
||||
_build_tiny_cosmos3,
|
||||
_build_tiny_packed_seq,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tiny config shared by both models (must match _build_tiny_cosmos3).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _build_tiny_fastvideo_dit() -> "Cosmos3VFMTransformer": # noqa: F821
|
||||
from fastvideo.configs.models.dits.cosmos3 import (
|
||||
Cosmos3ArchConfig,
|
||||
Cosmos3VideoConfig,
|
||||
)
|
||||
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
|
||||
|
||||
arch = Cosmos3ArchConfig(
|
||||
hidden_size=16,
|
||||
num_hidden_layers=1,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
intermediate_size=32,
|
||||
vocab_size=64,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
latent_patch_size=2,
|
||||
latent_channel=16,
|
||||
rope_theta=5_000_000.0,
|
||||
mrope_section=[24, 20, 20],
|
||||
position_embedding_type="3d_rope",
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=4,
|
||||
enable_fps_modulation=False,
|
||||
# Dormant heads present in the checkpoint surface (constructed for
|
||||
# strict-load parity; not exercised by this video-path forward).
|
||||
action_gen=True,
|
||||
action_dim=64,
|
||||
max_action_dim=64,
|
||||
num_embodiment_domains=32,
|
||||
sound_gen=True,
|
||||
sound_dim=64,
|
||||
)
|
||||
cfg = Cosmos3VideoConfig(arch_config=arch)
|
||||
model = Cosmos3VFMTransformer(cfg, hf_config={})
|
||||
return model.to(torch.float32).eval()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Framework -> FastVideo weight name map.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _framework_to_fastvideo_state_dict(vfm, num_layers: int) -> dict[str, torch.Tensor]:
|
||||
"""Translate framework param names into the FastVideo DiT param names.
|
||||
|
||||
Framework (Cosmos3VFMNetwork):
|
||||
language_model.model.{embed_tokens,norm,norm_moe_gen}
|
||||
language_model.lm_head
|
||||
language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)
|
||||
language_model.model.layers.{i}.self_attn.{q,k}_norm(+ _moe_gen)
|
||||
language_model.model.layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
|
||||
language_model.model.layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
|
||||
vae2llm / llm2vae / time_embedder.mlp.{0,2}
|
||||
|
||||
FastVideo (Cosmos3VFMTransformer):
|
||||
embed_tokens / norm / norm_moe_gen / lm_head
|
||||
layers.{i}.self_attn.{to_q,to_k,to_v,to_out} (und)
|
||||
layers.{i}.self_attn.{add_q,add_k,add_v}_proj / to_add_out (gen)
|
||||
layers.{i}.self_attn.{norm_q,norm_k,norm_added_q,norm_added_k}
|
||||
layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
|
||||
layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
|
||||
proj_in / proj_out / time_embedder.linear_{1,2}
|
||||
"""
|
||||
src = dict(vfm.named_parameters())
|
||||
out: dict[str, torch.Tensor] = {}
|
||||
|
||||
def take(name: str) -> torch.Tensor:
|
||||
return src[name].detach().clone()
|
||||
|
||||
# ---- Top-level backbone ----
|
||||
out["embed_tokens.weight"] = take("language_model.model.embed_tokens.weight")
|
||||
out["norm.weight"] = take("language_model.model.norm.weight")
|
||||
out["norm_moe_gen.weight"] = take("language_model.model.norm_moe_gen.weight")
|
||||
out["lm_head.weight"] = take("language_model.lm_head.weight")
|
||||
|
||||
# ---- Vision adapters ----
|
||||
out["proj_in.weight"] = take("vae2llm.weight")
|
||||
out["proj_in.bias"] = take("vae2llm.bias")
|
||||
out["proj_out.weight"] = take("llm2vae.weight")
|
||||
out["proj_out.bias"] = take("llm2vae.bias")
|
||||
|
||||
# ---- Timestep embedder (mlp.0/mlp.2 -> linear_1/linear_2) ----
|
||||
out["time_embedder.linear_1.weight"] = take("time_embedder.mlp.0.weight")
|
||||
out["time_embedder.linear_1.bias"] = take("time_embedder.mlp.0.bias")
|
||||
out["time_embedder.linear_2.weight"] = take("time_embedder.mlp.2.weight")
|
||||
out["time_embedder.linear_2.bias"] = take("time_embedder.mlp.2.bias")
|
||||
|
||||
# ---- Per layer ----
|
||||
und_attn = {"q_proj": "to_q", "k_proj": "to_k", "v_proj": "to_v", "o_proj": "to_out"}
|
||||
gen_attn = {
|
||||
"q_proj_moe_gen": "add_q_proj",
|
||||
"k_proj_moe_gen": "add_k_proj",
|
||||
"v_proj_moe_gen": "add_v_proj",
|
||||
"o_proj_moe_gen": "to_add_out",
|
||||
}
|
||||
und_norm = {"q_norm": "norm_q", "k_norm": "norm_k"}
|
||||
gen_norm = {"q_norm_moe_gen": "norm_added_q", "k_norm_moe_gen": "norm_added_k"}
|
||||
|
||||
for i in range(num_layers):
|
||||
fw = f"language_model.model.layers.{i}"
|
||||
fv = f"layers.{i}"
|
||||
for s, d in und_attn.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in gen_attn.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in und_norm.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in gen_norm.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for mlp in ("mlp", "mlp_moe_gen"):
|
||||
for proj in ("gate_proj", "up_proj", "down_proj"):
|
||||
out[f"{fv}.{mlp}.{proj}.weight"] = take(f"{fw}.{mlp}.{proj}.weight")
|
||||
for ln in ("input_layernorm", "input_layernorm_moe_gen", "post_attention_layernorm",
|
||||
"post_attention_layernorm_moe_gen"):
|
||||
out[f"{fv}.{ln}.weight"] = take(f"{fw}.{ln}.weight")
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def _copy_weights(vfm, dit) -> None:
|
||||
"""Copy framework weights into the FastVideo DiT (shape-checked)."""
|
||||
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
|
||||
dst = dict(dit.named_parameters())
|
||||
# Every mapped tensor must land on an existing FastVideo param with a matching shape.
|
||||
for name, tensor in mapped.items():
|
||||
assert name in dst, f"FastVideo DiT missing param for mapped key {name!r}"
|
||||
assert dst[name].shape == tensor.shape, (f"shape mismatch for {name}: "
|
||||
f"dit={tuple(dst[name].shape)} fw={tuple(tensor.shape)}")
|
||||
with torch.no_grad():
|
||||
for name, tensor in mapped.items():
|
||||
dst[name].copy_(tensor.to(dst[name].dtype))
|
||||
|
||||
|
||||
def _fastvideo_inputs_from_packed_seq(ps) -> dict:
|
||||
"""Build the FastVideo DiT forward kwargs from a framework PackedSequence."""
|
||||
v = ps.vision
|
||||
return dict(
|
||||
text_ids=ps.text_ids,
|
||||
text_indexes=ps.text_indexes,
|
||||
position_ids=ps.position_ids,
|
||||
sequence_length=int(ps.sequence_length),
|
||||
split_lens=list(ps.split_lens),
|
||||
attn_modes=list(ps.attn_modes),
|
||||
vision_tokens=list(v.tokens),
|
||||
vision_token_shapes=list(v.token_shapes),
|
||||
vision_sequence_indexes=v.sequence_indexes,
|
||||
vision_timesteps=v.timesteps,
|
||||
vision_mse_loss_indexes=v.mse_loss_indexes,
|
||||
vision_noisy_frame_indexes=list(v.noisy_frame_indexes),
|
||||
fps_vision=None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestCosmos3DiTParity:
|
||||
|
||||
def _run_both(self, seed_model: int = 42, seed_data: int = 7):
|
||||
vfm = _build_tiny_cosmos3(seed=seed_model)
|
||||
dit = _build_tiny_fastvideo_dit()
|
||||
_copy_weights(vfm, dit)
|
||||
ps = _build_tiny_packed_seq(n_text=4, seed=seed_data)
|
||||
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=ps)
|
||||
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
|
||||
return fw_out, fv_out
|
||||
|
||||
def test_weight_map_is_complete(self):
|
||||
"""The framework->fastvideo map must cover EVERY FastVideo DiT parameter
|
||||
that is exercised by the video path (i.e. all non-dormant params).
|
||||
|
||||
Dormant action/audio heads have no framework counterpart in this tiny
|
||||
vision-only setup, so they are excluded from the copy; everything else
|
||||
must be covered.
|
||||
"""
|
||||
vfm = _build_tiny_cosmos3(seed=42)
|
||||
dit = _build_tiny_fastvideo_dit()
|
||||
mapped = set(_framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers))
|
||||
dit_params = set(n for n, _ in dit.named_parameters())
|
||||
dormant = {
|
||||
n
|
||||
for n in dit_params
|
||||
if n.startswith(("action_", "audio_"))
|
||||
}
|
||||
uncovered = dit_params - mapped - dormant
|
||||
assert not uncovered, f"FastVideo DiT params not covered by weight map: {sorted(uncovered)}"
|
||||
|
||||
def test_preds_vision_parity(self):
|
||||
fw_out, fv_out = self._run_both()
|
||||
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
|
||||
fv_pv = fv_out["preds_vision"][0]
|
||||
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
|
||||
max_abs = (fw_pv - fv_pv).abs().max().item()
|
||||
print(f"\n[preds_vision] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_last_hidden_state_parity(self):
|
||||
fw_out, fv_out = self._run_both()
|
||||
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
|
||||
fv_lhs = fv_out["last_hidden_state"]
|
||||
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
|
||||
max_abs = (fw_lhs - fv_lhs).abs().max().item()
|
||||
print(f"\n[last_hidden_state] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_parity_holds_across_seeds(self):
|
||||
"""Re-running with a different random init still matches (not a fluke)."""
|
||||
fw_out, fv_out = self._run_both(seed_model=99, seed_data=13)
|
||||
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,371 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 DiT vs ``Cosmos3VFMNetwork`` (mRoPE).
|
||||
|
||||
Companion to ``test_cosmos3_dit_parity.py`` (which covers ``3d_rope``). This
|
||||
module exercises the rotary mode the REAL ``nvidia/Cosmos3-Nano`` checkpoint
|
||||
uses: ``position_embedding_type="unified_3d_mrope"`` with the real-checkpoint
|
||||
settings (``mrope_section=[24,20,20]``, ``mrope_interleaved=True``,
|
||||
``rope_theta=5e6``, ``unified_3d_mrope_reset_spatial_ids=True``,
|
||||
``temporal_modality_margin=15000``).
|
||||
|
||||
Under unified 3D mRoPE there is NO additive latent position embedding
|
||||
(``latent_pos_embed is None``); all positional information rides on the
|
||||
per-token 3D (T, H, W) rotary embedding. The packed-sequence ``position_ids``
|
||||
are therefore shape ``[3, seq_len]``, built exactly like the framework data
|
||||
packer (``cosmos_framework.data.vfm.sequence_packing``):
|
||||
|
||||
* text tokens broadcast one monotone id across all three axes
|
||||
(``get_3d_mrope_ids_text_tokens``),
|
||||
* the temporal offset is bumped by ``temporal_modality_margin`` at the
|
||||
text->vision boundary,
|
||||
* vision tokens lay out a (T, H, W) grid with spatial ids reset per segment
|
||||
(``get_3d_mrope_ids_vae_tokens`` with ``reset_spatial_indices=True``).
|
||||
|
||||
Both models are built tiny from the SAME config, framework weights are copied
|
||||
into the FastVideo DiT (reusing the ``3d_rope`` test's weight map — the
|
||||
transformer key surface is identical across rotary modes), and BOTH forwards
|
||||
run on identical deterministic CPU / float32 inputs. The official model is the
|
||||
parity ORACLE (run on CPU via the SDPA monkey-patch in
|
||||
``test_cosmos3_reference_forward``).
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_dit_parity_mrope.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the reference harness's SDPA monkey-patch and the 3d_rope parity
|
||||
# test's weight-copy + input-builder helpers (key surface is rotary-agnostic).
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_copy_weights,
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
_framework_to_fastvideo_state_dict,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Real-checkpoint unified_3d_mrope settings (tiny model, real rope constants).
|
||||
_ROPE_THETA = 5_000_000.0
|
||||
_MROPE_SECTION = [24, 20, 20]
|
||||
_MROPE_INTERLEAVED = True
|
||||
_RESET_SPATIAL_IDS = True
|
||||
_TEMPORAL_MODALITY_MARGIN = 15_000
|
||||
_LATENT_PATCH_SIZE = 2
|
||||
_LATENT_CHANNEL = 16
|
||||
_TCF = 4 # temporal compression factor
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tiny model builders (framework + FastVideo) with unified_3d_mrope.
|
||||
# ---------------------------------------------------------------------------
|
||||
_SOUND_DIM = 64
|
||||
_SOUND_LATENT_FPS = 25
|
||||
_ACTION_DIM = 64
|
||||
_NUM_EMBODIMENT_DOMAINS = 32
|
||||
|
||||
|
||||
def _build_tiny_cosmos3_mrope(seed: int = 42, num_layers: int = 2, sound_gen: bool = False,
|
||||
action_gen: bool = False):
|
||||
"""Tiny framework ``Cosmos3VFMNetwork`` with ``unified_3d_mrope``.
|
||||
|
||||
``rope_theta`` / ``rope_scaling`` (carrying ``mrope_section`` +
|
||||
``mrope_interleaved``) are threaded through the materialized text config;
|
||||
``position_embedding_type="unified_3d_mrope"`` leaves ``latent_pos_embed``
|
||||
as ``None`` so positions ride solely on the 3D rotary embedding.
|
||||
|
||||
``sound_gen=True`` additionally builds the sound MoT heads (``sound2llm`` /
|
||||
``llm2sound`` / ``sound_modality_embed``) for the t2vs parity test.
|
||||
"""
|
||||
from cosmos_framework.model.vfm.mot.cosmos3_vfm_network import (
|
||||
Cosmos3VFMNetwork,
|
||||
Cosmos3VFMNetworkConfig,
|
||||
)
|
||||
from cosmos_framework.model.vfm.mot.unified_mot import (
|
||||
Qwen3MoTConfig,
|
||||
Qwen3VLTextForCausalLM,
|
||||
)
|
||||
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import Qwen3VLConfig
|
||||
|
||||
tiny_text_dict = dict(
|
||||
model_type="qwen3_vl_text",
|
||||
vocab_size=64,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=num_layers,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
attention_dropout=0.0,
|
||||
rope_theta=_ROPE_THETA,
|
||||
rope_scaling={
|
||||
"rope_type": "default",
|
||||
"mrope_section": _MROPE_SECTION,
|
||||
"mrope_interleaved": _MROPE_INTERLEAVED,
|
||||
},
|
||||
max_position_embeddings=262144,
|
||||
)
|
||||
mot_cfg = Qwen3MoTConfig(
|
||||
config_dict=tiny_text_dict,
|
||||
qk_norm_for_text=True,
|
||||
qk_norm_for_diffusion=True,
|
||||
include_visual=False,
|
||||
)
|
||||
tiny_vlm_cfg = Qwen3VLConfig(text_config=tiny_text_dict)
|
||||
sound_kwargs = dict(
|
||||
sound_gen=True,
|
||||
sound_dim=_SOUND_DIM,
|
||||
temporal_compression_factor_sound=1,
|
||||
sound_latent_fps=_SOUND_LATENT_FPS,
|
||||
) if sound_gen else {}
|
||||
action_kwargs = dict(
|
||||
action_gen=True,
|
||||
action_dim=_ACTION_DIM,
|
||||
num_embodiment_domains=_NUM_EMBODIMENT_DOMAINS,
|
||||
) if action_gen else {}
|
||||
vfm_cfg = Cosmos3VFMNetworkConfig(
|
||||
vision_gen=True,
|
||||
vlm_config=tiny_vlm_cfg,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
latent_downsample_factor=8,
|
||||
latent_channel_size=_LATENT_CHANNEL,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
max_latent_h=16,
|
||||
max_latent_w=16,
|
||||
max_latent_t=8,
|
||||
temporal_compression_factor_vision=_TCF,
|
||||
**sound_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
torch.manual_seed(seed)
|
||||
lm = Qwen3VLTextForCausalLM(config=mot_cfg)
|
||||
vfm = Cosmos3VFMNetwork(language_model=lm, config=vfm_cfg)
|
||||
# inv_freq is a non-persistent buffer; init it on CPU (mirrors from_pretrained).
|
||||
vfm.language_model.model.rotary_emb.init_weights(buffer_device=None)
|
||||
vfm.eval()
|
||||
return vfm
|
||||
|
||||
|
||||
def _build_tiny_fastvideo_dit_mrope(num_layers: int = 2):
|
||||
from fastvideo.configs.models.dits.cosmos3 import (
|
||||
Cosmos3ArchConfig,
|
||||
Cosmos3VideoConfig,
|
||||
)
|
||||
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
|
||||
|
||||
arch = Cosmos3ArchConfig(
|
||||
hidden_size=16,
|
||||
num_hidden_layers=num_layers,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
intermediate_size=32,
|
||||
vocab_size=64,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
latent_channel=_LATENT_CHANNEL,
|
||||
rope_theta=_ROPE_THETA,
|
||||
mrope_section=_MROPE_SECTION,
|
||||
mrope_interleaved=_MROPE_INTERLEAVED,
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
enable_fps_modulation=False,
|
||||
# Dormant heads present in the checkpoint surface (constructed for
|
||||
# strict-load parity; not exercised by this video-path forward).
|
||||
action_gen=True,
|
||||
action_dim=64,
|
||||
max_action_dim=64,
|
||||
num_embodiment_domains=32,
|
||||
sound_gen=True,
|
||||
sound_dim=64,
|
||||
)
|
||||
cfg = Cosmos3VideoConfig(arch_config=arch)
|
||||
model = Cosmos3VFMTransformer(cfg, hf_config={})
|
||||
return model.to(torch.float32).eval()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# [3, seq_len] mRoPE position-id builder (mirrors the framework data packer).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _build_mrope_position_ids(n_text: int, grid_t: int, patch_h: int, patch_w: int) -> torch.Tensor:
|
||||
"""Build ``[3, seq_len]`` (T, H, W) mRoPE ids for one text+vision sample.
|
||||
|
||||
Reproduces ``pack_input_sequence`` for a single causal-text + full-vision
|
||||
sample: monotone text ids on all axes, ``+temporal_modality_margin`` at the
|
||||
text->vision boundary, then a reset-spatial (T, H, W) vision grid.
|
||||
"""
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
get_3d_mrope_ids_text_tokens,
|
||||
get_3d_mrope_ids_vae_tokens,
|
||||
)
|
||||
|
||||
offset: int | float = 0
|
||||
text_ids, offset = get_3d_mrope_ids_text_tokens(num_tokens=n_text, temporal_offset=offset)
|
||||
# End of text modality: add the boundary margin before vision.
|
||||
offset += _TEMPORAL_MODALITY_MARGIN
|
||||
vision_ids, offset = get_3d_mrope_ids_vae_tokens(
|
||||
grid_t=grid_t,
|
||||
grid_h=patch_h,
|
||||
grid_w=patch_w,
|
||||
temporal_offset=offset,
|
||||
reset_spatial_indices=_RESET_SPATIAL_IDS,
|
||||
fps=None, # integer positions (enable_fps_modulation=False)
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
return torch.cat([text_ids, vision_ids], dim=1) # [3, seq_len]
|
||||
|
||||
|
||||
def _build_tiny_packed_seq_mrope(
|
||||
*,
|
||||
n_text: int = 6,
|
||||
grid_t: int = 2,
|
||||
latent_h: int = 4,
|
||||
latent_w: int = 4,
|
||||
seed: int = 7,
|
||||
):
|
||||
"""Minimal PackedSequence with ``[3, seq]`` mRoPE position ids.
|
||||
|
||||
Vision latent ``[C, grid_t, latent_h, latent_w]`` patchifies (patch=2) to a
|
||||
``(grid_t, latent_h/2, latent_w/2)`` token grid; all frames are noisy.
|
||||
"""
|
||||
from cosmos_framework.data.vfm.sequence_packing import ModalityData, PackedSequence
|
||||
|
||||
patch_h = latent_h // _LATENT_PATCH_SIZE
|
||||
patch_w = latent_w // _LATENT_PATCH_SIZE
|
||||
n_vision = grid_t * patch_h * patch_w
|
||||
total_len = n_text + n_vision
|
||||
|
||||
torch.manual_seed(seed)
|
||||
vision_tensor = torch.randn(_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
text_ids = torch.randint(0, 64, (n_text,))
|
||||
position_ids = _build_mrope_position_ids(n_text, grid_t, patch_h, patch_w) # [3, total_len]
|
||||
|
||||
noisy_frame_indexes = torch.arange(grid_t, dtype=torch.long) # all frames noisy
|
||||
vision_mod = ModalityData(
|
||||
sequence_indexes=torch.arange(n_text, total_len, dtype=torch.long),
|
||||
timesteps=torch.full((n_vision,), 500.0),
|
||||
mse_loss_indexes=torch.arange(n_text, total_len, dtype=torch.long),
|
||||
token_shapes=[(grid_t, patch_h, patch_w)],
|
||||
tokens=[vision_tensor],
|
||||
condition_mask=[torch.zeros(grid_t, dtype=torch.long)], # 0 = noisy
|
||||
noisy_frame_indexes=[noisy_frame_indexes],
|
||||
)
|
||||
packed_seq = PackedSequence(
|
||||
sample_lens=[total_len],
|
||||
split_lens=[n_text, n_vision],
|
||||
attn_modes=["causal", "full"],
|
||||
is_image_batch=(grid_t == 1),
|
||||
sequence_length=total_len,
|
||||
text_ids=text_ids,
|
||||
text_indexes=torch.arange(n_text, dtype=torch.long),
|
||||
position_ids=position_ids,
|
||||
vision=vision_mod,
|
||||
)
|
||||
return packed_seq
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
# (grid_t, latent_h, latent_w): a single image, a small video, and a taller
|
||||
# video, to exercise the spatial mRoPE overwrite + gen<->gen full attention.
|
||||
_GRIDS = [
|
||||
pytest.param(1, 8, 8, id="image_1x4x4"),
|
||||
pytest.param(2, 4, 4, id="video_2x2x2"),
|
||||
pytest.param(3, 8, 4, id="video_3x4x2"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3DiTParityMRoPE:
|
||||
|
||||
def _run_both(
|
||||
self,
|
||||
*,
|
||||
grid_t: int,
|
||||
latent_h: int,
|
||||
latent_w: int,
|
||||
seed_model: int = 42,
|
||||
seed_data: int = 7,
|
||||
num_layers: int = 2,
|
||||
):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
ps = _build_tiny_packed_seq_mrope(
|
||||
n_text=6, grid_t=grid_t, latent_h=latent_h, latent_w=latent_w, seed=seed_data
|
||||
)
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=ps)
|
||||
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
|
||||
return fw_out, fv_out
|
||||
|
||||
def test_position_ids_are_3xN_mrope(self):
|
||||
"""The packed mRoPE ids are ``[3, seq_len]`` with the text->vision margin."""
|
||||
ps = _build_tiny_packed_seq_mrope(n_text=6, grid_t=2, latent_h=4, latent_w=4)
|
||||
pos = ps.position_ids
|
||||
assert pos.ndim == 2 and pos.shape[0] == 3, f"expected [3, N], got {tuple(pos.shape)}"
|
||||
assert pos.shape[1] == int(ps.sequence_length)
|
||||
# Text axis is monotone 0..5 on all 3 rows; vision temporal jumps by the margin.
|
||||
assert pos[0, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
assert pos[1, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
assert pos[2, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
# First vision token temporal id == last_text_id (5) + margin + 1.
|
||||
assert pos[0, 6].item() == 5 + _TEMPORAL_MODALITY_MARGIN + 1
|
||||
# Reset spatial: first vision token H/W ids are 0.
|
||||
assert pos[1, 6].item() == 0 and pos[2, 6].item() == 0
|
||||
|
||||
def test_no_additive_latent_pos_embed(self):
|
||||
"""unified_3d_mrope must NOT build an additive latent position embedding."""
|
||||
dit = _build_tiny_fastvideo_dit_mrope()
|
||||
assert dit.position_embedding_type == "unified_3d_mrope"
|
||||
assert dit.latent_pos_embed is None
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
|
||||
def test_preds_vision_parity(self, grid_t, latent_h, latent_w):
|
||||
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
|
||||
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
|
||||
fv_pv = fv_out["preds_vision"][0]
|
||||
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
|
||||
max_abs = (fw_pv - fv_pv).abs().max().item()
|
||||
print(f"\n[preds_vision mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
|
||||
def test_last_hidden_state_parity(self, grid_t, latent_h, latent_w):
|
||||
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
|
||||
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
|
||||
fv_lhs = fv_out["last_hidden_state"]
|
||||
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
|
||||
max_abs = (fw_lhs - fv_lhs).abs().max().item()
|
||||
print(f"\n[last_hidden_state mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_parity_holds_across_seeds(self):
|
||||
"""A different random init still matches bit-for-bit (not a fluke)."""
|
||||
fw_out, fv_out = self._run_both(
|
||||
grid_t=2, latent_h=4, latent_w=4, seed_model=99, seed_data=13
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
fv_out["last_hidden_state"], fw_out["last_hidden_state"], atol=1e-4, rtol=1e-3
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 UniPC flow_shift vs the framework.
|
||||
|
||||
The framework selects the UniPC ``shift`` purely from the named resolution
|
||||
bucket the (H, W) belongs to, via ``OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS``
|
||||
(keyed by the VLM model size — Cosmos3-Nano uses the 8B backbone — and the
|
||||
resolution string), NOT from the task (T2V/I2V/T2I share a shift at a given
|
||||
resolution). FastVideo gets raw pixel ``height``/``width`` and must map back to
|
||||
the same shift.
|
||||
|
||||
This pins ``Cosmos3DenoisingStage._flow_shift_for_resolution`` against the
|
||||
framework's own tables: for every (resolution, aspect) entry in
|
||||
``VIDEO_RES_SIZE_INFO`` whose resolution has an 8B shift default, the FastVideo
|
||||
shift for that exact pixel size must equal the framework default.
|
||||
|
||||
The framework tables are the parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_flow_shift_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
# The official framework provides the parity oracle for the resolution->pixel
|
||||
# tables. (``cosmos_framework.inference.args`` — which holds the shift constant —
|
||||
# can't be imported here: it transitively requires ``multistorageclient``. The
|
||||
# small shift table is mirrored verbatim below with its source location.)
|
||||
_utils = pytest.importorskip(
|
||||
"cosmos_framework.data.vfm.utils",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
|
||||
Cosmos3DenoisingStage,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Cosmos3-Nano's VLM backbone is Qwen3-VL-8B (checkpoint config.json).
|
||||
_MODEL_SIZE = "8B"
|
||||
# Verbatim from cosmos_framework.inference.args.OmniSampleArgs
|
||||
# ._RESOLUTION_SHIFT_DEFAULTS (args.py:770), restricted to the 8B rows.
|
||||
_SHIFT_DEFAULTS = {
|
||||
("8B", "256"): 3.0,
|
||||
("8B", "480"): 5.0,
|
||||
("8B", "720"): 10.0,
|
||||
("8B", "768"): 10.0,
|
||||
("32B", "256"): 5.0,
|
||||
("32B", "480"): 5.0,
|
||||
("32B", "720"): 5.0,
|
||||
("32B", "768"): 5.0,
|
||||
}
|
||||
_VIDEO_RES = _utils.VIDEO_RES_SIZE_INFO
|
||||
_IMAGE_RES = _utils.IMAGE_RES_SIZE_INFO
|
||||
|
||||
|
||||
def _cases():
|
||||
seen = set()
|
||||
for resolution, by_aspect in {**_VIDEO_RES, **_IMAGE_RES}.items():
|
||||
key = (_MODEL_SIZE, resolution)
|
||||
if key not in _SHIFT_DEFAULTS:
|
||||
continue
|
||||
expected = _SHIFT_DEFAULTS[key]
|
||||
for aspect, (a, b) in by_aspect.items():
|
||||
cid = f"{resolution}_{aspect.replace(',', '-')}_{a}x{b}"
|
||||
if cid in seen:
|
||||
continue
|
||||
seen.add(cid)
|
||||
yield pytest.param(a, b, expected, id=cid)
|
||||
|
||||
|
||||
class TestCosmos3FlowShiftParity:
|
||||
|
||||
@pytest.mark.parametrize(("dim_a", "dim_b", "expected_shift"), list(_cases()))
|
||||
def test_flow_shift_matches_framework(self, dim_a, dim_b, expected_shift):
|
||||
got = Cosmos3DenoisingStage._flow_shift_for_resolution(dim_a, dim_b)
|
||||
assert got == expected_shift, (
|
||||
f"shift for {dim_a}x{dim_b}: got {got}, framework default {expected_shift}")
|
||||
@@ -0,0 +1,91 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo I2V conditioning pixel video vs the framework.
|
||||
|
||||
The Cosmos3 I2V path conditions on a *static repeat* of the input image. The
|
||||
framework (``cosmos_framework.inference.vision``):
|
||||
|
||||
* ``load_conditioning_image``: aspect-preserving resize + center crop + uint8
|
||||
quantization, then ``/127.5 - 1`` -> ``[3, 1, h, w]`` in [-1, 1];
|
||||
* ``build_conditioned_video_batch``: frame 0 = the image, and every remaining
|
||||
frame **repeats the last conditioning frame** (a static video) -> the clip
|
||||
is then VAE-encoded and only the latent condition frame(s) are kept clean.
|
||||
|
||||
Because the VAE is temporal, zero-filling the non-condition frames (the earlier
|
||||
FastVideo behavior) changes the encoded condition latent, so the repeat-fill is
|
||||
correctness-critical. This pins FastVideo's
|
||||
``Cosmos3DenoisingStage._image_to_video_tensor`` against the framework's
|
||||
image preprocessing + repeat-fill.
|
||||
|
||||
CPU / float32. The framework is the parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_i2v_conditioning_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
vision = pytest.importorskip(
|
||||
"cosmos_framework.inference.vision",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
|
||||
Cosmos3DenoisingStage,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
|
||||
def _make_image(path, h_in: int, w_in: int, seed: int = 0) -> None:
|
||||
rng = np.random.default_rng(seed)
|
||||
arr = rng.integers(0, 256, size=(h_in, w_in, 3), dtype=np.uint8)
|
||||
Image.fromarray(arr, "RGB").save(path)
|
||||
|
||||
|
||||
# (input H, input W, target H, target W, num_frames)
|
||||
_CASES = [
|
||||
pytest.param(120, 200, 256, 256, 9, id="square_from_landscape"),
|
||||
pytest.param(200, 120, 704, 1280, 13, id="wide_from_portrait"),
|
||||
pytest.param(256, 256, 256, 256, 5, id="same_size"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3I2VConditioningParity:
|
||||
|
||||
@pytest.mark.parametrize(("h_in", "w_in", "h", "w", "num_frames"), _CASES)
|
||||
def test_conditioning_video_matches_framework(self, tmp_path, h_in, w_in, h, w, num_frames):
|
||||
img_path = tmp_path / "cond.png"
|
||||
_make_image(img_path, h_in, w_in)
|
||||
|
||||
# ---- Framework oracle ----
|
||||
# load_conditioning_image -> [3, 1, h, w] in [-1, 1].
|
||||
cond = vision.load_conditioning_image(img_path, target_h=h, target_w=w).float()
|
||||
# Mirror build_conditioned_video_batch (vision.py lines 117-123) in fp32/CPU:
|
||||
# frame 0 = image; remaining frames repeat the last conditioning frame.
|
||||
t_cond = cond.shape[1]
|
||||
expected = torch.zeros(1, 3, num_frames, h, w, dtype=torch.float32)
|
||||
t_fill = min(t_cond, num_frames)
|
||||
expected[0, :, :t_fill] = cond[:, :t_fill]
|
||||
if t_fill < num_frames:
|
||||
expected[0, :, t_fill:] = expected[0, :, t_fill - 1:t_fill].expand(-1, num_frames - t_fill, -1, -1)
|
||||
|
||||
# ---- FastVideo: same PIL image through the stage helper ----
|
||||
pil = Image.open(img_path).convert("RGB")
|
||||
got = Cosmos3DenoisingStage._image_to_video_tensor(
|
||||
pil, num_frames, h, w, torch.device("cpu"), torch.float32)
|
||||
|
||||
assert got.shape == expected.shape, f"shape: got={got.shape} expected={expected.shape}"
|
||||
max_abs = (got - expected).abs().max().item()
|
||||
print(f"\n[i2v_cond {h}x{w} nf={num_frames}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(got, expected)
|
||||
|
||||
# Static repeat (not zero-fill): every frame equals frame 0, and the
|
||||
# frames past frame 0 are non-zero.
|
||||
assert torch.equal(got[0, :, 0], got[0, :, -1]), "non-condition frames must repeat the image"
|
||||
assert got[0, :, 1:].abs().sum() > 0, "non-condition frames must not be zero-filled"
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 unified 3D mRoPE position-ID parity (Tier A scaffold).
|
||||
|
||||
Reference: ``vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py``
|
||||
lines 113-177 (``compute_mrope_position_ids_text`` /
|
||||
``compute_mrope_position_ids_vision``). The reference test asserting
|
||||
these invariants lives at
|
||||
``tests/diffusion/models/cosmos3/test_cosmos3_transformer.py:32-57``.
|
||||
|
||||
The three invariants under test:
|
||||
|
||||
1. Text tokens broadcast the same monotonically-increasing positions
|
||||
across all three (t, h, w) axes. With ``num_tokens=3`` and
|
||||
``temporal_offset=5`` the result is ``[[5,6,7], [5,6,7], [5,6,7]]``
|
||||
and the next-offset is ``8``.
|
||||
|
||||
2. Vision tokens (no FPS modulation) flatten a ``(grid_t, grid_h, grid_w)``
|
||||
position grid in t-major order. With ``(2, 2, 3)`` and offset ``10``
|
||||
the resulting shape is ``(3, 12)`` and the temporal row begins
|
||||
``[10]*6 + [11]*6``; next-offset is ``12``.
|
||||
|
||||
3. FPS-modulated vision tokens scale the temporal axis by
|
||||
``base_fps / temporal_compression_factor / (fps / tcf)``. With
|
||||
``fps=12``, ``base_fps=24``, ``tcf=4``, ``grid_t=2`` the first row is
|
||||
``[10.0, 12.0]``.
|
||||
|
||||
The FastVideo side currently does NOT exist; the test is wrapped in
|
||||
``try/except ImportError`` and skips. Phase 2b replaces the skip with
|
||||
the real import + assertion path.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
|
||||
def test_compute_mrope_position_ids_text_and_vision() -> None:
|
||||
"""Asserts the 3 invariants of unified 3D mRoPE position-ID generation.
|
||||
|
||||
Once FastVideo's ``fastvideo.models.dits.cosmos3`` exports
|
||||
``compute_mrope_position_ids_text`` and
|
||||
``compute_mrope_position_ids_vision``, this test verifies they produce
|
||||
output tensors identical to the vllm-omni reference at
|
||||
transformer_cosmos3.py:113-177.
|
||||
"""
|
||||
try:
|
||||
from fastvideo.models.dits.cosmos3 import ( # type: ignore
|
||||
compute_mrope_position_ids_text,
|
||||
compute_mrope_position_ids_vision,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("FastVideo Cosmos3 not yet implemented (Phase 2b)")
|
||||
|
||||
text_ids, text_offset = compute_mrope_position_ids_text(num_tokens=3, temporal_offset=5)
|
||||
assert text_ids.tolist() == [[5, 6, 7], [5, 6, 7], [5, 6, 7]]
|
||||
assert text_offset == 8
|
||||
|
||||
vision_ids, vision_offset = compute_mrope_position_ids_vision(
|
||||
2, 2, 3, temporal_offset=10, fps=None
|
||||
)
|
||||
assert tuple(vision_ids.shape) == (3, 12)
|
||||
assert vision_ids[0].tolist() == [10] * 6 + [11] * 6
|
||||
assert vision_offset == 12
|
||||
|
||||
modulated_ids, modulated_offset = compute_mrope_position_ids_vision(
|
||||
2,
|
||||
1,
|
||||
1,
|
||||
temporal_offset=10,
|
||||
fps=12.0,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=4,
|
||||
)
|
||||
torch.testing.assert_close(modulated_ids[0], torch.tensor([10.0, 12.0]))
|
||||
assert modulated_offset == 13
|
||||
@@ -0,0 +1,312 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 sequence packing vs the OFFICIAL framework.
|
||||
|
||||
FastVideo's native packer
|
||||
(``fastvideo.pipelines.basic.cosmos3.sequence_packing.pack_cosmos3_video_sequence``)
|
||||
builds the packed-sequence inputs the ``Cosmos3VFMTransformer`` consumes. This
|
||||
test asserts, for the SAME logical inputs (prompt token ids, vision latents,
|
||||
condition-frame indices, diffusion timestep, fps), that FastVideo's packing
|
||||
matches the official ``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``
|
||||
oracle field-by-field:
|
||||
|
||||
* ``position_ids`` (exact, ``[3, seq]``),
|
||||
* ``text_ids`` / ``text_indexes``,
|
||||
* ``split_lens`` / ``attn_modes`` / ``sample_lens`` / ``sequence_length``,
|
||||
* vision ``sequence_indexes`` / ``token_shapes`` / ``timesteps`` /
|
||||
``mse_loss_indexes`` / ``noisy_frame_indexes`` / ``condition_mask``.
|
||||
|
||||
Coverage spans T2V (no condition frames), I2V (condition frame 0), and T2I
|
||||
(single conditioned frame), across multiple grids, plus a multi-sample batch.
|
||||
|
||||
Then BOTH the framework-packed and FastVideo-packed inputs are fed through the
|
||||
SAME tiny FastVideo DiT (framework weights copied in as in the existing DiT
|
||||
parity tests). Asserting bit-identical DiT output confirms FastVideo's own
|
||||
packing drives the DiT to the same result as the framework's packing.
|
||||
|
||||
The official framework is the parity ORACLE; it runs on CPU / float32 via the
|
||||
SDPA monkey-patch from ``test_cosmos3_reference_forward``.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_packing_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the DiT parity helpers (weight copy + framework->DiT kwarg builder) and
|
||||
# the mRoPE tiny-model builders (real-checkpoint rope constants).
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_copy_weights,
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
)
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_MROPE_SECTION,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_ROPE_THETA,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Tiny special-token ids (kept < tiny vocab_size=64). The video path appends
|
||||
# eos + start_of_generation after the prompt tokens.
|
||||
_SPECIAL_TOKENS = {
|
||||
"start_of_generation": 60,
|
||||
"end_of_generation": 61,
|
||||
"eos_token_id": 62,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Builders for the two packers from the SAME logical sample inputs.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _make_vision(grid_t: int, latent_h: int, latent_w: int, seed: int) -> torch.Tensor:
|
||||
"""Deterministic VAE latent ``[1, C, T, H, W]``."""
|
||||
torch.manual_seed(seed)
|
||||
return torch.randn(1, _LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
|
||||
|
||||
def _framework_pack(
|
||||
*,
|
||||
text_ids_per_sample: list[list[int]],
|
||||
visions: list[torch.Tensor],
|
||||
cond_frames_per_sample: list[list[int]],
|
||||
timesteps: list[float],
|
||||
is_image_batch: bool,
|
||||
):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
gen_data_clean = GenerationDataClean(
|
||||
batch_size=len(visions),
|
||||
is_image_batch=is_image_batch,
|
||||
x0_tokens_vision=list(visions),
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1] * len(visions),
|
||||
)
|
||||
plans = [
|
||||
SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cf))
|
||||
for cf in cond_frames_per_sample
|
||||
]
|
||||
return pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(t) for t in text_ids_per_sample],
|
||||
gen_data_clean=gen_data_clean,
|
||||
input_timesteps=torch.tensor(timesteps, dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _fastvideo_pack(
|
||||
*,
|
||||
text_ids_per_sample: list[list[int]],
|
||||
visions: list[torch.Tensor],
|
||||
cond_frames_per_sample: list[list[int]],
|
||||
timesteps: list[float],
|
||||
):
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
|
||||
samples = [
|
||||
Cosmos3SampleInputs(
|
||||
text_ids=list(t),
|
||||
vision=Cosmos3VisionItem(latent=v, condition_frame_indexes=list(cf)),
|
||||
timestep=float(ts),
|
||||
)
|
||||
for t, v, cf, ts in zip(text_ids_per_sample, visions, cond_frames_per_sample, timesteps)
|
||||
]
|
||||
return pack_cosmos3_video_sequence(
|
||||
samples,
|
||||
_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Field-by-field comparison.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _assert_packs_match(fw, fv) -> None:
|
||||
"""Assert the framework PackedSequence and FastVideo pack agree field-by-field."""
|
||||
# Structure.
|
||||
assert fv.split_lens == list(fw.split_lens), f"split_lens: fv={fv.split_lens} fw={list(fw.split_lens)}"
|
||||
assert fv.attn_modes == list(fw.attn_modes), f"attn_modes: fv={fv.attn_modes} fw={list(fw.attn_modes)}"
|
||||
assert fv.sample_lens == list(fw.sample_lens), f"sample_lens: fv={fv.sample_lens} fw={list(fw.sample_lens)}"
|
||||
assert int(fv.sequence_length) == int(fw.sequence_length)
|
||||
|
||||
# Text.
|
||||
torch.testing.assert_close(fv.text_ids, fw.text_ids.to(torch.long), rtol=0, atol=0)
|
||||
torch.testing.assert_close(fv.text_indexes, fw.text_indexes.to(torch.long), rtol=0, atol=0)
|
||||
|
||||
# position_ids: exact, [3, seq], same dtype.
|
||||
assert fv.position_ids.shape == fw.position_ids.shape, (
|
||||
f"position_ids shape: fv={tuple(fv.position_ids.shape)} fw={tuple(fw.position_ids.shape)}")
|
||||
assert fv.position_ids.dtype == fw.position_ids.dtype, (
|
||||
f"position_ids dtype: fv={fv.position_ids.dtype} fw={fw.position_ids.dtype}")
|
||||
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
|
||||
|
||||
# Vision.
|
||||
fwv = fw.vision
|
||||
torch.testing.assert_close(fv.vision_sequence_indexes, fwv.sequence_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert fv.vision_token_shapes == [tuple(s) for s in fwv.token_shapes], (
|
||||
f"token_shapes: fv={fv.vision_token_shapes} fw={[tuple(s) for s in fwv.token_shapes]}")
|
||||
torch.testing.assert_close(fv.vision_timesteps.to(torch.float32), fwv.timesteps.to(torch.float32))
|
||||
torch.testing.assert_close(fv.vision_mse_loss_indexes, fwv.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert len(fv.vision_noisy_frame_indexes) == len(fwv.noisy_frame_indexes)
|
||||
for a, b in zip(fv.vision_noisy_frame_indexes, fwv.noisy_frame_indexes):
|
||||
torch.testing.assert_close(a.to(torch.long), b.to(torch.long), rtol=0, atol=0)
|
||||
assert len(fv.vision_condition_mask) == len(fwv.condition_mask)
|
||||
for a, b in zip(fv.vision_condition_mask, fwv.condition_mask):
|
||||
torch.testing.assert_close(a.flatten().to(torch.float32), b.flatten().to(torch.float32))
|
||||
|
||||
|
||||
# (grid_t, latent_h, latent_w, n_text, cond_frames, id) — single-sample cases.
|
||||
_CASES = [
|
||||
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
|
||||
pytest.param(1, 4, 4, 5, [0], id="t2i_cond_1x2x2"),
|
||||
pytest.param(2, 4, 4, 4, [], id="t2v_2x2x2"),
|
||||
pytest.param(3, 8, 4, 6, [], id="t2v_3x4x2"),
|
||||
pytest.param(2, 4, 4, 5, [0], id="i2v_2x2x2"),
|
||||
pytest.param(3, 4, 4, 4, [0], id="i2v_3x2x2"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3PackingParity:
|
||||
|
||||
# -- Field-by-field packing parity -------------------------------------
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_packing_fields_match_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
torch.manual_seed(0)
|
||||
text_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
|
||||
timestep = 500.0
|
||||
|
||||
fw = _framework_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
is_image_batch=(grid_t == 1),
|
||||
)
|
||||
fv = _fastvideo_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
)
|
||||
_assert_packs_match(fw, fv)
|
||||
|
||||
def test_packing_fields_match_framework_multi_sample(self):
|
||||
"""A batch of two samples (T2V + I2V) packs identically to the framework."""
|
||||
torch.manual_seed(1)
|
||||
t0 = torch.randint(0, 60, (3,)).tolist()
|
||||
t1 = torch.randint(0, 60, (5,)).tolist()
|
||||
v0 = _make_vision(2, 4, 4, seed=11)
|
||||
v1 = _make_vision(2, 4, 4, seed=22)
|
||||
kwargs = dict(
|
||||
text_ids_per_sample=[t0, t1],
|
||||
visions=[v0, v1],
|
||||
cond_frames_per_sample=[[], [0]],
|
||||
timesteps=[500.0, 250.0],
|
||||
)
|
||||
fw = _framework_pack(is_image_batch=False, **kwargs)
|
||||
fv = _fastvideo_pack(**kwargs)
|
||||
_assert_packs_match(fw, fv)
|
||||
|
||||
# -- End-to-end: FastVideo packing drives the DiT identically ----------
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_fastvideo_packing_drives_dit_like_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
"""Feed BOTH the framework-packed and FastVideo-packed inputs through the
|
||||
SAME FastVideo DiT (framework weights copied in); assert identical output.
|
||||
"""
|
||||
num_layers = 2
|
||||
torch.manual_seed(0)
|
||||
text_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
|
||||
timestep = 500.0
|
||||
|
||||
fw_pack = _framework_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
is_image_batch=(grid_t == 1),
|
||||
)
|
||||
fv_pack = _fastvideo_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
)
|
||||
# Guard: the two packs must agree before we trust the DiT comparison.
|
||||
_assert_packs_match(fw_pack, fv_pack)
|
||||
|
||||
# One DiT instance, framework weights copied in (parity oracle weights).
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=42, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
|
||||
with torch.no_grad():
|
||||
out_fw = dit(**_fastvideo_inputs_from_packed_seq(fw_pack))
|
||||
out_fv = dit(**fv_pack.to_dit_kwargs())
|
||||
|
||||
# last_hidden_state must be bit-identical.
|
||||
lhs_fw = out_fw["last_hidden_state"]
|
||||
lhs_fv = out_fv["last_hidden_state"]
|
||||
assert lhs_fw.shape == lhs_fv.shape
|
||||
max_abs_lhs = (lhs_fw - lhs_fv).abs().max().item()
|
||||
print(f"\n[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
|
||||
f"last_hidden_state max abs diff = {max_abs_lhs:.3e}")
|
||||
torch.testing.assert_close(lhs_fv, lhs_fw, rtol=0, atol=0)
|
||||
|
||||
# preds_vision must be bit-identical when there are noisy frames to
|
||||
# predict. (A fully-conditioned clip has no noisy patches, so the DiT
|
||||
# emits no "preds_vision" — both packs agree the mse-loss set is empty,
|
||||
# already asserted by the field-parity guard above.)
|
||||
has_preds = fv_pack.vision_mse_loss_indexes.numel() > 0
|
||||
assert ("preds_vision" in out_fw) == has_preds
|
||||
assert ("preds_vision" in out_fv) == has_preds
|
||||
if has_preds:
|
||||
pv_fw = out_fw["preds_vision"][0]
|
||||
pv_fv = out_fv["preds_vision"][0]
|
||||
assert pv_fw.shape == pv_fv.shape
|
||||
max_abs_pv = (pv_fw - pv_fv).abs().max().item()
|
||||
print(f"[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
|
||||
f"preds_vision max abs diff = {max_abs_pv:.3e}")
|
||||
torch.testing.assert_close(pv_fv, pv_fw, rtol=0, atol=0)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user