Compare commits

...
Author SHA1 Message Date
leffffandClaude Fable 5 35a283531e [misc]: refresh Kandinsky5 e2e oracle calibration comment with current reference numbers
The calibration table/comment cited stats measured against the reference
video that shipped in b6ba7003, deleted in that same commit as part of
the generator_update_interval fix -- the currently-committed reference
(from the corrected run) has slightly different numbers. Recompute and
update the comment and both assertion messages; verified the current
reference plus its self-calibrating known-bad regression tests
(fastvideo/tests/workflow/test_kandinsky5_e2e_media_oracle.py) still
pass 4/4 against the new numbers with the same comfortable margin.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:31:15 -07:00
leffff 3c9e604199 [misc]: add Kandinsky5 QAD e2e golden reference video (recorded bootstrap run) 2026-07-26 02:31:15 -07:00
leffffandClaude Fable 5 2cb832290a [bugfix]: fit Kandinsky5 e2e stage 2 in one 80GiB GPU; make multi-GPU sharding viable
The generator-update fix costs a second full Adam state (~16GiB for the
2B student) on top of the critic's: a recorded stage-2 run peaked at
~76GiB and OOMed allocating the last MiBs with ~1.9GiB lost to
fragmentation. Set PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True for
the training stages to reclaim that headroom.

Also make KANDINSKY5_E2E_NUM_GPUS=2+ actually work as the fallback:
DP_SP_BatchSampler floor-divides batch count by the number of DP groups
with drop_last=True, so the previous 1-row dataset produced an EMPTY
dataloader on every rank at 2+ GPUs. Synthesize max(2, NUM_GPUS)
manifest rows pointing at the same clip -- identical latents/caption
keep the single-sample-overfit semantics.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:31:15 -07:00
leffffandClaude Fable 5 5aae1c093a [bugfix]: address third-round Kandinsky5 QAD review findings
- Never touch the user's documented dataset: preprocess roots are now
  env-overridable (KANDINSKY5_OVERFIT_DATA_DIR/_OUTPUT_DIR) and the
  nightly keeps every artifact under a single test-owned
  data/kandinsky5_e2e/ root, cleaning only that
- Override method.generator_update_interval to 1 in the E2E stage-2 run:
  with the recipe's 5, a 3-step run performed zero student updates and the
  exported "student" was stage-1 weights
- Harden the media oracle: exact frame count/geometry (closes the SSIM
  shorter-clip truncation hole), per-frame luma spatial std (solid clip at
  reference mean scored 15.4 on the old global RGB std), consecutive-frame
  temporal diff (frozen clips), and MS-SSIM floor recalibrated to 0.95
  against measured known-bads (solid 0.9189, frozen 0.8597, good 1.0000);
  known-bad clips are permanent self-calibrating regressions in
  fastvideo/tests/workflow/test_kandinsky5_e2e_media_oracle.py
- Capability-gate ATTN_QAT_INFER: is_attn_qat_infer_available() now
  requires an sm_120/sm_121 device in addition to the extension import
  (CUDA 13 wheels can bundle the extension on H100/GB200, where the first
  kernel call failed instead of falling back); the real platform resolver
  is covered in fastvideo/tests/api/test_attn_qat_infer_capability_gate.py
- Remove the now-stale golden reference: the generator-update fix changes
  what stage 2 trains, so the reference must be re-recorded from the
  corrected run (KANDINSKY5_E2E_WRITE_REFERENCE=1 bootstrap + review)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:31:15 -07:00
leffffandClaude Fable 5 068c7ba57e [misc]: tighten Kandinsky5 e2e MS-SSIM floor to 0.90 from recorded runs
An independent retrain-from-scratch validation run scored mean/min/max
MS-SSIM 1.0000 against the committed reference (recorded 2026-07-18),
proving the chain deterministic on fixed hardware. Replace the 0.5
bootstrap placeholder with 0.90: headroom for GPU/library drift, loud
failure on real remap/schedule regressions.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
leffff 9b95acadd8 [misc]: add Kandinsky5 QAD e2e golden reference video (recorded bootstrap run) 2026-07-26 02:29:52 -07:00
leffffandClaude Fable 5 5ee7980fcb [misc]: un-ignore nightly reference videos in .gitignore
The catch-all *.mp4 rule blocked committing the new Kandinsky5 QAD e2e
golden reference (existing nightly references were force-added). Add the
negation under the existing "Reference videos" section so current and
future reference_video_*.mp4 files stage normally.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
leffffandClaude Fable 5 3e1201c455 [misc]: make Kandinsky5 e2e reference video reviewable and SSIM-stable
Bootstrap run on a GPU box produced uniform grey static: the "random
noise" caption doubled as the generation prompt, and 3 stage-1 steps at
the recipe's full 5e-5 LR toward a pure-noise clip collapsed 4-step
no-CFG DMD sampling (50-step T2V sampling of the same export still
produced a good video, and the base checkpoint under the 4-step sampler
produced structured output -- isolating the damage to the finetune).
Grey-static references are indistinguishable from pipeline failure by
eye and decorrelate across runs, making the SSIM oracle flaky.

Use a semantic caption/prompt (clip content stays noise; irrelevant for
plumbing) and a 1e-6 stage-1 LR override so the deterministic reference
is blurry-but-structured: human-reviewable and stable under SSIM.
Document the review bar in the module docstring.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
leffffandClaude Fable 5 64f60b3d8a [bugfix]: address second-round Kandinsky5 QAD review findings
- Keep only ATTN_QAT_TRAIN strict in the stage backend guard; preserve the
  documented ATTN_QAT_INFER->FlashAttention fallback (+ regression tests)
- Resolve FSDP compute dtype from the recorded MixedPrecisionPolicy instead
  of scanning parameter storage dtypes (+ FSDP regression tests)
- E2E: export/strict-reload the final stage-2 student, instantiate it via
  the documented Kandinsky5DMDPipeline/Kandinsky5DMDConfig override,
  generate deterministically, and compare MS-SSIM against a committed
  reference (fail, not skip, when the reference is missing;
  KANDINSKY5_E2E_WRITE_REFERENCE=1 bootstraps it)
- E2E: delete all input/output roots up front and explicitly disable
  resume_from_checkpoint for both stages
- preprocess: path-specific ValueError on zero decoded frames, non-empty
  manifest validation, fixture VideoWriter.isOpened() check (+ tests)
- Stage-boundary regression test driving Kandinsky5DmdDenoisingStage with a
  backend spy; fails if the set_forward_context wrapper is removed
- distill_dmd_qat.sh now accepts a stage-1 DCP checkpoint, runs
  dcp_to_diffusers --verify, and injects the export into all three
  models.*.init_from overrides; README/YAML/finetune docs updated

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
leffffandClaude Sonnet 5 21d751fbb4 [bugfix]: address Kandinsky5 QAD PR review findings
Fixes all P1/P2 items from the review of 7c7704da:

- E2E test's data_path pointed at a nonexistent combined_parquet_dataset/
  subdir; preprocess_kandinsky5_overfit.py writes directly to its output
  root.
- Stage 2 loaded the raw DCP checkpoint-N/ directly, which isn't a
  diffusers model dir. Added dcp_to_diffusers --verify (strict-reload
  the export immediately) and wired the conversion into the E2E test,
  launcher scripts, and README between stage 1 and stage 2.
- Kandinsky5DmdDenoisingStage called the transformer without
  set_forward_context, so LocalAttention's missing-context guard caught
  the AssertionError and silently fell back to plain SDPA -- Stage 2
  validation never actually exercised ATTN_QAT_TRAIN. Wrapped the call
  to match the parent stage's contract and added a scoped backend
  assertion (ATTN_QAT_TRAIN/ATTN_QAT_INFER only, since other backends
  can legitimately fall back to auto-selection).
- Exported DMD checkpoints kept the base T2V model_index.json, so
  VideoGenerator.from_pretrained resolved them to Kandinsky5T2VPipeline
  instead of the DMD sampler. Added Kandinsky5DMDConfig and documented
  the required override_pipeline_cls_name + pipeline_config combo.
- The FSDP-training dtype-detection special case (scan params, default
  to bf16) applied unconditionally, so an explicit fp32 pipeline was
  silently cast to bf16 for plain inference. Extracted
  _resolve_target_dtype and gated the scan on isinstance(transformer,
  FSDPModule).
- Preprocessing tokenized captions with a 512-token cap that doesn't
  account for the 129-token system template runtime strips off; now
  reads text_encoder_max_lengths from pipeline_config to match runtime.
  Also replaced unconditional item["cap"][0] indexing (silently takes
  the first character of a string caption) with a get_caption() helper
  that accepts/validates both the documented string schema and the
  list-of-variants schema some producers use.
- test_kandinsky5_qat_attention_engages.py set FASTVIDEO_ATTENTION_BACKEND
  and cleared the attention backend selector cache without teardown,
  leaking ATTN_QAT_TRAIN into later tests in the same pytest process.
  Now uses monkeypatch + try/finally.

Added regression tests: test_kandinsky5_target_dtype.py,
test_kandinsky5_dmd_pipeline_resolution.py,
test_kandinsky5_overfit_caption.py.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
leffffandClaude Sonnet 5 2dbb0c94a9 [feat] Add Kandinsky5 QAD training pipeline: data preprocessing, QAT finetune, QAT-aware DMD distillation
Adds the training-side counterpart to the existing Kandinsky-5 T2V/I2V
inference pipeline (#1471): a full quantization-aware distillation (QAD)
recipe for Kandinsky-5.0 Lite T2V at 512x768/121 frames on the new
fastvideo/train/ stack.

- Data preprocessing: preprocess_kandinsky5_overfit.py encodes raw
  video + captions (Kandinsky5's shared HunyuanVideo VAE + dual
  Qwen/Reason1 + CLIP text encoders) into the t2v parquet schema.
- Stage 1 (Attn-QAT finetune): fastvideo.train.models.kandinsky5.Kandinsky5Model
  wires Kandinsky5 into the modular trainer; FASTVIDEO_ATTENTION_BACKEND=
  ATTN_QAT_TRAIN routes dense/local attention through a fake-quantized
  straight-through-estimator kernel during finetuning so the DiT learns
  to absorb quantization error. No weight quantization during training.
- Stage 2 (QAT-aware DMD2 distillation): Kandinsky5DmdDenoisingStage +
  Kandinsky5DMDPipeline distill the stage-1 checkpoint from many-step
  teacher sampling down to a 4-step student, with teacher/critic masked
  back to full-precision dense attention regardless of the student's
  QAT env var (component_loader.py's _loading_teacher_critic_model gate).

See examples/train/configs/fine_tuning/kandinsky5/README.md for the full
data -> stage-1 -> stage-2 -> inference (NVFP4/FP8) walkthrough.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-26 02:29:52 -07:00
36 changed files with 3729 additions and 25 deletions
+1
View File
@@ -75,6 +75,7 @@ docs/distillation/examples/
*.pkl
# Reference videos (negations must come after the catch-all on line below)
!fastvideo/tests/nightly/reference_video_*.mp4
# Static images
!docs/assets/images/**/*.png
+3 -1
View File
@@ -77,7 +77,9 @@ If forcing a backend fails, verify optional dependencies are installed:
- `SAGE_ATTN`: SageAttention package
- `SAGE_ATTN_THREE`: upstream `sageattn3` package
- `ATTN_QAT_INFER`: `fastvideo-kernel` checkout/source install that exposes
`attn_qat_infer`
`attn_qat_infer`, AND a consumer-Blackwell sm_120 GPU -- on any
other device the backend reports unavailable (even if a CUDA 13 wheel
bundles the extension) and selection falls back to FlashAttention
- `ATTN_QAT_TRAIN`: `fastvideo-kernel`; its runtime-JIT Triton implementation
selects an optimized route on SM100, joins the quantized and STE P@V paths on
SM120, and retains the previous route for unsupported configurations. See
@@ -0,0 +1,70 @@
#!/bin/bash
# Stage-2 QAT-aware DMD2 distillation for Kandinsky5 T2V (480p).
#
# Usage:
# bash distill_dmd_qat.sh <stage1-checkpoint> [--dotted.key value ...]
#
# <stage1-checkpoint> is either:
# - a raw stage-1 DCP checkpoint written by
# examples/train/configs/fine_tuning/kandinsky5/finetune_qat.sh
# (a checkpoint-<N>/ dir with a dcp/ subdir, or the stage-1 output_dir
# itself -- dcp_to_diffusers auto-picks the latest checkpoint). It is
# converted to a diffusers export (<checkpoint>-diffusers) with
# `dcp_to_diffusers --role student --verify` before launching --
# models.*.init_from need a diffusers model directory (model_index.json
# + component subfolders), not the raw DCP layout; or
# - an already-exported diffusers model dir (model_index.json present),
# used as-is with no conversion.
#
# The export is passed to all three models.*.init_from overrides: student,
# teacher, and critic all load the same weights. Teacher/critic are
# automatically masked back to dense attention and full-precision weights
# by the _loading_teacher_critic_model gate in
# fastvideo/models/loader/component_loader.py (family-agnostic, no
# Kandinsky5-specific handling needed) -- NOT by pointing them at a
# different export.
#
# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN keeps the student's dense/local
# attention fake-quantized during distillation.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../../.." && pwd)"
CHECKPOINT="${1:?Usage: $0 <stage1-checkpoint (DCP dir or diffusers export)> [extra --dotted.key overrides...]}"
shift
if [[ "${CHECKPOINT}" != /* ]]; then
CHECKPOINT="$(pwd)/${CHECKPOINT}"
fi
CHECKPOINT="${CHECKPOINT%/}"
if [[ -f "${CHECKPOINT}/model_index.json" ]]; then
# Already a diffusers export -- use as-is.
INIT_FROM="${CHECKPOINT}"
else
# Convert the raw DCP checkpoint. Runs BEFORE ATTN_QAT_TRAIN is
# exported below: backend selection for that env var refuses to start
# (ImportError) when the training kernel isn't built, and the
# conversion itself must not depend on it. --overwrite keeps relaunches
# from silently reusing a stale export for a newer checkpoint;
# --verify strictly reloads the exported transformer so a key-mapping
# bug fails here, not deep inside the training launch below.
INIT_FROM="${CHECKPOINT}-diffusers"
echo "Converting stage-1 DCP checkpoint ${CHECKPOINT} -> ${INIT_FROM}"
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
--checkpoint "${CHECKPOINT}" \
--output-dir "${INIT_FROM}" \
--role student \
--overwrite \
--verify
fi
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/dmd2_t2v_480p_qat.yaml" \
--models.student.init_from "${INIT_FROM}" \
--models.teacher.init_from "${INIT_FROM}" \
--models.critic.init_from "${INIT_FROM}" \
"$@"
@@ -0,0 +1,113 @@
# DMD2 QAT-aware distillation: Kandinsky5 T2V Lite, 480p (teacher many-step
# -> student 4-step), continuing from a stage-1 Attn-QAT finetune checkpoint.
#
# init_from below must be a diffusers model directory (model_index.json +
# component subfolders), not the raw checkpoint-N/ stage 1 writes (dcp/ +
# metadata/RNG state only). Launching via distill_dmd_qat.sh handles this:
# it runs fastvideo.train.entrypoint.dcp_to_diffusers (--verify) on the
# stage-1 checkpoint you pass it and overrides all three init_from values
# with the export -- the placeholder paths below only apply when launching
# the YAML directly.
#
# - Teacher: frozen, full-precision Kandinsky5 (loaded from the stage-1 QAT
# checkpoint's base weights, or the original pretrained checkpoint --
# the _loading_teacher_critic_model gate in component_loader.py masks
# quant_config and FASTVIDEO_ATTENTION_BACKEND for teacher/critic
# regardless, so they always run full precision / dense attention).
# - Student: trainable, initialized from the stage-1 QAT finetune checkpoint.
# - Critic: trainable, initialized from the same checkpoint as the student.
# - Validation: 4-step sampling via Kandinsky5DMDPipeline.
#
# Launch with FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN set (student only;
# teacher/critic are masked to dense automatically) so the student continues
# learning to absorb fake-quantized attention error during distillation.
#
# Scope: T2V, 480p only, dense/local attention only.
models:
student:
_target_: fastvideo.train.models.kandinsky5.Kandinsky5Model
init_from: outputs/kandinsky5_t2v_qat_finetune/checkpoint-1000-diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.kandinsky5.Kandinsky5Model
init_from: outputs/kandinsky5_t2v_qat_finetune/checkpoint-1000-diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.kandinsky5.Kandinsky5Model
init_from: outputs/kandinsky5_t2v_qat_finetune/checkpoint-1000-diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 4.5
dmd_denoising_steps: [1000, 750, 500, 250]
# Critic optimizer (required -- no fallback to training.optimizer)
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/kandinsky5_overfit_preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 31
num_height: 512
num_width: 768
num_frames: 121
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/kandinsky5_t2v_dmd2_4steps_qat
training_state_checkpointing_steps: 20
checkpoints_total_limit: 3
tracker:
project_name: distillation_kandinsky5
run_name: kandinsky5_t2v_dmd2_4steps_qat
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.kandinsky5.kandinsky5_dmd_pipeline.Kandinsky5DMDPipeline
dataset_file: data/kandinsky5_overfit_preprocessed/validation_prompts.json
every_steps: 50
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
guidance_scale: 6.0
ema:
_target_: fastvideo.train.callbacks.ema.EMACallback
decay: 0.98
start_iter: 0
pipeline:
flow_shift: 5
@@ -0,0 +1,198 @@
# Kandinsky5 QAD recipe (T2V, 480p)
Quantization-aware distillation for Kandinsky-5.0 Lite T2V at
**512×768, 121 frames, 24 fps** (the checkpoint's native `KANDINSKY5_T2V_LITE_5S`
preset -- 121/24 ~= 5.04s, matching the "5s" in the checkpoint name),
mirroring the Wan2.1 QAD recipe
(`examples/training/finetune/wan_t2v_1.3B/mixkit/`) on the new
`fastvideo/train/` stack. Scope is deliberately narrow: T2V only, 480p only,
dense/local attention only. Kandinsky5's NABLA sparse attention path is
never engaged at this resolution and is unsupported by this recipe --
attempting a larger resolution will raise loudly from `Kandinsky5Model`
rather than silently mis-scale the visual RoPE.
## Data
Kandinsky5 uses two text encoders (Qwen/Reason1 + CLIP), unlike Wan's single
encoder. The CLIP pooled projection is zero-padded into a row and prepended
to the Qwen sequence embeddings, then stored as a single `[seq+1, dim]`
tensor in the existing `text_embedding` Parquet field -- the same trick
`preprocess_hunyuan_overfit.py` uses for LLaMA+CLIP. No new Parquet schema
is needed.
Build a small Parquet dataset from raw videos + captions
(`videos2caption.json` + `videos/*.mp4`, matching the Hunyuan overfit script's
layout) with:
```bash
# from the repo root; reads data/kandinsky5_overfit and writes
# data/kandinsky5_overfit_preprocessed by default -- override with
# KANDINSKY5_OVERFIT_DATA_DIR / KANDINSKY5_OVERFIT_OUTPUT_DIR env vars
python -m fastvideo.pipelines.preprocess.preprocess_kandinsky5_overfit
```
This loads Kandinsky5's VAE (shared with HunyuanVideo) and both text
encoders via FastVideo's own loaders (`VAELoader`/`TextEncoderLoader`), and
writes `data_00000.parquet` + `validation_prompts.json` into the configured
output directory. Point `training.data.data_path` at that output directory
in the YAML configs below.
There is currently no scalable/production preprocessing pipeline (the kind
built on `BasePreprocessPipeline` that Wan/LTX2 use) for Kandinsky5 -- like
Hunyuan, its dual text encoders don't fit that base class's single-encoder
assumption, and `BasePreprocessPipeline.preprocess_video_and_text` isn't
cleanly overridable in pieces. If you need to scale past a handful of
overfit samples, extend `preprocess_kandinsky5_overfit.py`'s loop rather
than forking the production pipeline base class.
## Train stage 1 (Attn-QAT finetune)
```bash
bash examples/train/configs/fine_tuning/kandinsky5/finetune_qat.sh \
--training.data.data_path data/kandinsky5_overfit_preprocessed
```
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` (set by the script) routes
Kandinsky5's dense/local attention through the fake-quantized
straight-through-estimator Triton kernel, so the DiT learns to absorb
quantization error instead of fighting it. This is purely env-var driven,
model-agnostic, and does not touch weights during training -- weight-level
FP4/FP8 quantization is applied post-hoc at inference time (see below), not
during either training stage.
## Train stage 2 (QAT-aware DMD distillation)
Stage 1 writes a raw DCP checkpoint (`checkpoint-<N>/dcp` + metadata/RNG
state) under `training.checkpoint.output_dir`
(`outputs/kandinsky5_t2v_qat_finetune/` by default) -- `models.*.init_from`
needs a diffusers model directory instead (`model_index.json` + component
subfolders). Pass the stage-1 checkpoint to the launcher, which performs
the conversion and injects the export into all three `init_from` overrides:
```bash
bash examples/train/configs/distribution_matching/kandinsky5/distill_dmd_qat.sh \
outputs/kandinsky5_t2v_qat_finetune/checkpoint-<N>
```
Under the hood it runs exactly:
```bash
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
--checkpoint outputs/kandinsky5_t2v_qat_finetune/checkpoint-<N> \
--output-dir outputs/kandinsky5_t2v_qat_finetune/checkpoint-<N>-diffusers \
--role student --overwrite --verify
```
then launches `dmd2_t2v_480p_qat.yaml` with
`--models.student/teacher/critic.init_from` all pointing at the export.
(Passing an already-exported diffusers directory to the launcher skips the
conversion and uses it directly.)
`--role student` exports the trained transformer weights (the only role
that matters here: student/teacher/critic in stage 2 all load from the same
checkpoint, and `_loading_teacher_critic_model` handles the full-precision
masking for teacher/critic at load time, not at export time). `--verify`
strictly reloads the exported transformer immediately, so a key-mapping bug
fails here instead of deep inside the stage-2 launch.
Only the student is quantized (Attn-QAT); the teacher and critic stay full
precision / dense attention. This is enforced in the loader
(`fastvideo/models/loader/component_loader.py`, via the
`_loading_teacher_critic_model` flag), which masks `quant_config` and clears
`FASTVIDEO_ATTENTION_BACKEND` for teacher/critic -- the same global env var
reaches only the student, with no per-model flags or Kandinsky5-specific
handling. Validation runs the distilled student through
`Kandinsky5DMDPipeline` at the step counts configured in
`callbacks.validation.sampling_steps`.
## Export the stage-2 student for inference
Stage 2 writes raw DCP checkpoints too, so the distilled student must be
exported the same way stage 1 was before anything can load it -- the
`path/to/kandinsky5_dmd_checkpoint` used below is exactly this export:
```bash
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
--checkpoint outputs/kandinsky5_t2v_dmd2_4steps_qat/checkpoint-<N> \
--output-dir outputs/kandinsky5_t2v_dmd2_4steps_qat/checkpoint-<N>-diffusers \
--role student --verify
```
## Inference (NVFP4 / FP8 weight quantization)
Kandinsky5's attention out-projection (`self_attention.out_layer` /
`cross_attention.out_layer`) and FFN (`feed_forward.mlp.fc_in` /
`feed_forward.mlp.fc_out`) names differ from Wan's (`to_out`, `ffn.fc_in` /
`ffn.fc_out`); `to_query`/`to_key`/`to_value` already substring-match the
`to_q`/`to_k`/`to_v` entries the quant configs use for Wan. Both
`nvfp4_qat` and `FP8` now include Kandinsky5's out-projection/FFN names in
their default layer lists (`fastvideo/layers/quantization/nvfp4_qat_config.py`,
`fastvideo/layers/quantization/fp8_config.py`), so no `target_layers`
override is needed:
`dcp_to_diffusers` copies the base T2V checkpoint's `model_index.json`
unchanged into every export (see "Train stage 2" above), so `_class_name`
still says the base T2V pipeline and the registry
(`fastvideo/registry.py`) has no way to auto-detect that a given directory
is actually a DMD (four-step re-noise sampler) export -- it will resolve to
`Kandinsky5T2VPipeline` and run the full-length sampler on DMD-distilled
weights. Pass `override_pipeline_cls_name` and a `Kandinsky5DMDConfig`
explicitly to select the right pipeline/sampler and get
`dmd_denoising_steps` set (`Kandinsky5T2VConfig`'s default of `None` makes
`Kandinsky5DmdDenoisingStage` raise immediately):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig
from fastvideo.layers.quantization import get_quantization_config
gen = VideoGenerator.from_pretrained(
"path/to/kandinsky5_dmd_checkpoint", num_gpus=1,
override_pipeline_cls_name="Kandinsky5DMDPipeline",
pipeline_config=Kandinsky5DMDConfig(),
transformer_quant=get_quantization_config("nvfp4_qat")(),
use_fsdp_inference=False,
)
# Kandinsky5DmdDenoisingStage ignores request.sampling.num_inference_steps
# and guidance_scale -- it's a fixed 4-step, no-CFG sampler driven entirely
# by pipeline_config.dmd_denoising_steps (set above). Adjust
# Kandinsky5DMDConfig(dmd_denoising_steps=[...]) instead if the checkpoint
# was distilled/validated with a different schedule than the default
# [1000, 750, 500, 250].
gen.generate(request={"prompt": "...", "output": {"save_video": True}})
```
As with Wan, combine with `FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER` on an
sm_120 GPU for the full 4-bit path; on other GPUs the attention falls back
to a supported dense backend while the FP4 linear layers still run.
`flashinfer` is required for the FP4 path.
## Tests
- `fastvideo/tests/train/models/test_load_kandinsky5.py` -- loads the real
checkpoint and runs one transformer forward pass.
- `fastvideo/tests/train/models/test_kandinsky5_qat_attention_engages.py` --
confirms `ATTN_QAT_TRAIN` actually engages (not a silent SDPA fallback)
and that gradients flow through a forward+backward pass.
- `fastvideo/tests/api/test_kandinsky5_dmd_pipeline_resolution.py` --
confirms an unmodified export resolves to `Kandinsky5T2VPipeline` (the bug)
and that `override_pipeline_cls_name="Kandinsky5DMDPipeline"` fixes it (the
documented workaround above), plus `Kandinsky5DMDConfig`'s
`dmd_denoising_steps` default.
- `fastvideo/tests/nightly/test_e2e_kandinsky5_dmd_t2v_overfit.py` -- a few
steps of both training stages on a single synthetic sample, exercising the
whole recipe end to end: the `dcp_to_diffusers --verify` conversion
between the stages AND of the final stage-2 student, then reloading that
export through the documented `Kandinsky5DMDPipeline` /
`Kandinsky5DMDConfig` override, generating deterministically (fixed
seed), and comparing MS-SSIM against a committed reference video. The
test fails (does not skip) if the reference is missing; record it once on
a sanctioned GPU box with `KANDINSKY5_E2E_WRITE_REFERENCE=1`, review the
written video, and commit it (see the test's module docstring). Nightly
tests are not collected per-PR (`fastvideo/tests/contract/
test_ci_test_collection.py` allowlists the directory), so run it
explicitly:
```bash
pytest fastvideo/tests/nightly/test_e2e_kandinsky5_dmd_t2v_overfit.py -vs
```
@@ -0,0 +1,33 @@
#!/bin/bash
# QAD recipe -- quantization-aware finetune of Kandinsky5 T2V (480p) with
# fake-quant (Attn-QAT) attention.
#
# The fake-quantized attention path is selected purely by env var
# (config-driven, no monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN
# routes Kandinsky5's dense/local attention through the fake-quantized
# straight-through-estimator kernel, so the DiT learns to absorb the
# quantization error instead of fighting it. No weight quantization happens
# during this stage; that's applied post-hoc at inference time (see the
# stage-2 config and README).
#
# Scope: T2V, 480p only, dense/local attention only -- NABLA sparse attention
# is never engaged by Kandinsky5Model at this resolution.
#
# Data: preprocess with fastvideo/pipelines/preprocess/preprocess_kandinsky5_overfit.py
# (or an equivalent parquet dataset matching pyarrow_schema_t2v) first.
#
# Next: this writes a raw DCP checkpoint (checkpoint-N/dcp + metadata/RNG
# state), not something stage 2 can load directly. Pass that checkpoint to
# ../../distribution_matching/kandinsky5/distill_dmd_qat.sh, which converts
# it with fastvideo.train.entrypoint.dcp_to_diffusers (--verify) and feeds
# the export to all three models.*.init_from overrides automatically.
set -euo pipefail
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../../.." && pwd)"
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/t2v_480p_qat.yaml" \
"$@"
@@ -0,0 +1,85 @@
# Kandinsky5 T2V 480p stage-1 QAT finetune config.
#
# QAT here means attention-only fake-quant during training: launch with
# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN set (see finetune_qat.sh), which
# routes Kandinsky5's dense/local attention through the fake-quantized
# straight-through-estimator kernel. No weight quantization happens during
# training; that's applied post-hoc at inference time via the nvfp4_qat/FP8
# quant configs (see the stage-2 config and README for that step).
#
# Scope: T2V, 480p only, dense/local attention only -- NABLA sparse attention
# is never engaged by Kandinsky5Model at this resolution.
#
# Data must be preprocessed with Kandinsky5's VAE (shared with HunyuanVideo)
# + dual text encoders (Qwen/Reason1 + CLIP) into parquet format before
# training. See fastvideo/pipelines/preprocess/preprocess_kandinsky5_overfit.py.
models:
student:
_target_: fastvideo.train.models.kandinsky5.Kandinsky5Model
init_from: kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/kandinsky5_overfit_preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.1
seed: 1000
num_latent_t: 31
num_height: 512
num_width: 768
num_frames: 121
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 1.0e-4
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/kandinsky5_t2v_qat_finetune
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: fastvideo_kandinsky5
run_name: kandinsky5_t2v_qat_finetune
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.kandinsky5.kandinsky5_pipeline.Kandinsky5T2VPipeline
dataset_file: data/kandinsky5_overfit_preprocessed/validation_prompts.json
every_steps: 250
# Cut way down from 50 while shaking out validation-path bugs -- at
# 512x768x121 each step is ~30s, so 50 steps/sample x 6 samples/rank
# was ~2.5h per validation round. Bump back to 50 for real runs.
sampling_steps: [25]
guidance_scale: 6.0
pipeline:
flow_shift: 5
+26 -2
View File
@@ -50,9 +50,33 @@ def _get_attn_qat_infer() -> Callable[..., torch.Tensor] | None:
return _attn_qat_infer
# Consumer-Blackwell compute capability the modified SageAttention3 FP4
# kernel is compiled for (sm_120a -- see fastvideo-kernel/README.md).
_SUPPORTED_DEVICE_CAPABILITIES = frozenset({(12, 0)})
def _device_capability_supported() -> bool:
if not torch.cuda.is_available():
return False
try:
return tuple(torch.cuda.get_device_capability()) in _SUPPORTED_DEVICE_CAPABILITIES
except Exception: # pragma: no cover - defensive: never break backend selection
return False
def is_attn_qat_infer_available() -> bool:
return (torch.cuda.is_available() and torch.cuda.get_device_capability() == (12, 0)
and _get_attn_qat_infer() is not None)
"""True only when the extension imports AND the active device is a
consumer-Blackwell sm_120 GPU the kernel is compiled for.
The import check alone is not sufficient: CUDA 13 wheel builds can
carry the sm_120 extension on any host (e.g. H100, GB200),
where the import succeeds, backend selection picks this backend, and
the first kernel call then fails with an unsupported-capability error
instead of ever reaching the documented FlashAttention fallback in
fastvideo.platforms.cuda. Gating on the active device's capability
keeps that fallback working on every non-sm_120 GPU.
"""
return _device_capability_supported() and _get_attn_qat_infer() is not None
class AttnQatInferBackend(AttentionBackend):
+5 -2
View File
@@ -13,12 +13,15 @@ def _is_kandinsky5_transformer_block(n: str, m) -> bool:
class Kandinsky5ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
# the dense backends every DiT supports.
# NABLA block-sparse attention for attention_type="nabla" checkpoints, the
# dense backends every DiT supports, and the Attn-QAT backends needed for
# quantization-aware finetuning/distillation (fastvideo/train/models/kandinsky5/).
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.NABLA_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.ATTN_QAT_INFER,
AttentionBackendEnum.ATTN_QAT_TRAIN,
)
# Native FastVideo implementation uses the same parameter names as diffusers
@@ -7,6 +7,18 @@ from typing import Any
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and n.split(".")[-1].isdigit()
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass
class Reason1ArchConfig(TextEncoderArchConfig):
"""Architecture settings (defaults match Qwen2.5-VL-7B-Instruct)."""
@@ -61,6 +73,8 @@ class Reason1ArchConfig(TextEncoderArchConfig):
torch_dtype: str = "bfloat16"
_attn_implementation: str = "flash_attention_2"
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm])
@dataclass
+3 -3
View File
@@ -6,7 +6,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig, Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
@@ -21,6 +21,6 @@ __all__ = [
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
"Kandinsky5I2VConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig",
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+18
View File
@@ -120,3 +120,21 @@ class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
super().__post_init__()
# I2V needs the VAE encoder to encode the conditioning image.
self.vae_config.load_encoder = True
@dataclass
class Kandinsky5DMDConfig(Kandinsky5T2VConfig):
"""Kandinsky-5.0 DMD (few-step distilled) text-to-video pipeline configuration.
Checkpoints exported by ``fastvideo.train.entrypoint.dcp_to_diffusers``
copy their base T2V checkpoint's ``model_index.json`` unchanged, so
``_class_name`` still says the base T2V pipeline and the registry
(``fastvideo/registry.py``) cannot auto-detect a DMD export -- pass this
config together with ``override_pipeline_cls_name="Kandinsky5DMDPipeline"``
explicitly to ``VideoGenerator.from_pretrained``/``from_config`` (see
``examples/train/configs/fine_tuning/kandinsky5/README.md``). Without it,
``Kandinsky5T2VConfig``'s ``dmd_denoising_steps=None`` makes
``Kandinsky5DmdDenoisingStage`` raise immediately.
"""
dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 750, 500, 250])
@@ -26,6 +26,9 @@ FP8_DTYPE = torch.float8_e4m3fn
FP8_MAX = float(torch.finfo(FP8_DTYPE).max) # 448.0
FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0)
# Wan's "to_q"/"to_k"/"to_v" also substring-match Kandinsky5's "to_query"/
# "to_key"/"to_value", so only Kandinsky5's out-projection and FFN names need
# to be listed explicitly below.
_FP8_SUFFIXES = (
"ffn.fc_in",
"ffn.fc_out",
@@ -33,6 +36,11 @@ _FP8_SUFFIXES = (
"to_k",
"to_v",
"to_out",
# Kandinsky5
"self_attention.out_layer",
"cross_attention.out_layer",
"feed_forward.mlp.fc_in",
"feed_forward.mlp.fc_out",
)
@@ -46,7 +46,11 @@ from fastvideo.models.utils import set_weight_attrs
logger = logging.getLogger(__name__)
# Wan-style attention + FFN projection layers. Matched as substrings of the
# layer prefix (e.g. "blocks.0.attn1.to_q" contains "to_q").
# layer prefix (e.g. "blocks.0.attn1.to_q" contains "to_q"). Wan's "to_q"/
# "to_k"/"to_v" also substring-match Kandinsky5's "to_query"/"to_key"/
# "to_value", so only Kandinsky5's out-projection and FFN names (which don't
# share a substring with Wan's "to_out"/"ffn.fc_in"/"ffn.fc_out") need to be
# listed explicitly below.
DEFAULT_FP4_LAYERS = (
"ffn.fc_in",
"ffn.fc_out",
@@ -54,6 +58,11 @@ DEFAULT_FP4_LAYERS = (
"to_k",
"to_v",
"to_out",
# Kandinsky5
"self_attention.out_layer",
"cross_attention.out_layer",
"feed_forward.mlp.fc_in",
"feed_forward.mlp.fc_out",
)
+32 -11
View File
@@ -15,6 +15,7 @@ from fastvideo.configs.models.dits import Kandinsky5VideoConfig
from fastvideo.layers.layernorm import LayerNormScaleShift
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
@@ -285,6 +286,7 @@ class Kandinsky5Attention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
prefix: str = "",
use_nabla: bool = False,
quant_config: QuantizationConfig | None = None,
):
super().__init__()
assert num_channels % head_dim == 0
@@ -293,20 +295,24 @@ class Kandinsky5Attention(nn.Module):
self.to_query = ReplicatedLinear(num_channels,
num_channels,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_query")
self.to_key = ReplicatedLinear(num_channels,
num_channels,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_key")
self.to_value = ReplicatedLinear(num_channels,
num_channels,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_value")
self.query_norm = nn.RMSNorm(head_dim)
self.key_norm = nn.RMSNorm(head_dim)
self.out_layer = ReplicatedLinear(num_channels,
num_channels,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.out_layer")
self.local_attention = LocalAttention(
num_heads=self.num_heads,
@@ -413,9 +419,11 @@ class Kandinsky5Attention(nn.Module):
class Kandinsky5FeedForward(nn.Module):
def __init__(self, dim: int, ff_dim: int):
def __init__(self, dim: int, ff_dim: int, prefix: str = "",
quant_config: QuantizationConfig | None = None):
super().__init__()
self.mlp = MLP(dim, ff_dim, bias=False, act_type="gelu")
self.mlp = MLP(dim, ff_dim, bias=False, act_type="gelu",
quant_config=quant_config, prefix=f"{prefix}.mlp")
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.mlp(x)
@@ -467,7 +475,8 @@ class Kandinsky5TransformerEncoderBlock(nn.Module):
head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = ""):
prefix: str = "",
quant_config: QuantizationConfig | None = None):
super().__init__()
self.text_modulation = Kandinsky5Modulation(time_dim, model_dim, 6)
@@ -482,7 +491,8 @@ class Kandinsky5TransformerEncoderBlock(nn.Module):
model_dim,
head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.self_attention")
prefix=f"{prefix}.self_attention",
quant_config=quant_config)
self.feed_forward_norm = LayerNormScaleShift(
model_dim,
@@ -491,7 +501,9 @@ class Kandinsky5TransformerEncoderBlock(nn.Module):
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim)
self.feed_forward = Kandinsky5FeedForward(
model_dim, ff_dim, prefix=f"{prefix}.feed_forward",
quant_config=quant_config)
def forward(self, x: torch.Tensor, time_embed: torch.Tensor,
rope: torch.Tensor) -> torch.Tensor:
@@ -523,7 +535,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
use_nabla: bool = False):
use_nabla: bool = False,
quant_config: QuantizationConfig | None = None):
super().__init__()
self.visual_modulation = Kandinsky5Modulation(time_dim, model_dim, 9)
@@ -539,7 +552,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.self_attention",
use_nabla=use_nabla)
use_nabla=use_nabla,
quant_config=quant_config)
self.cross_attention_norm = LayerNormScaleShift(
model_dim,
@@ -552,7 +566,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
model_dim,
head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.cross_attention")
prefix=f"{prefix}.cross_attention",
quant_config=quant_config)
self.feed_forward_norm = LayerNormScaleShift(
model_dim,
@@ -561,7 +576,9 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim)
self.feed_forward = Kandinsky5FeedForward(
model_dim, ff_dim, prefix=f"{prefix}.feed_forward",
quant_config=quant_config)
def forward(self, visual_embed: torch.Tensor, text_embed: torch.Tensor,
time_embed: torch.Tensor, rope: torch.Tensor,
@@ -636,6 +653,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
quant_config = config.quant_config
self.quant_config = quant_config
head_dim = sum(arch.axes_dims)
self.in_visual_dim = arch.in_visual_dim
@@ -664,7 +683,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
arch.ff_dim,
head_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.text_transformer_blocks.{i}")
prefix=f"{config.prefix}.text_transformer_blocks.{i}",
quant_config=quant_config)
for i in range(arch.num_text_blocks)
])
self.visual_transformer_blocks = nn.ModuleList([
@@ -673,7 +693,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
head_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
use_nabla=arch.attention_type == "nabla")
use_nabla=arch.attention_type == "nabla",
quant_config=quant_config)
for i in range(arch.num_visual_blocks)
])
@@ -0,0 +1,82 @@
# SPDX-License-Identifier: Apache-2.0
"""Kandinsky5 DMD (few-step distilled) validation pipeline.
Composes the same stages as ``Kandinsky5T2VPipeline`` (text encoding, latent
prep, decoding) but swaps in ``Kandinsky5DmdDenoisingStage`` for the
denoising loop, mirroring ``wan_dmd_pipeline.py``'s split from the base T2V
pipeline. This exists solely to drive the stage-2 DMD training config's
validation callback.
"""
from __future__ import annotations
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.kandinsky5 import (
Kandinsky5DecodingStage,
Kandinsky5DmdDenoisingStage,
Kandinsky5LatentPreparationStage,
)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
class Kandinsky5DMDPipeline(ComposedPipelineBase):
"""Kandinsky-5.0 Lite DMD (few-step) text-to-video pipeline."""
_required_config_modules = [
"scheduler",
"text_encoder",
"text_encoder_2",
"tokenizer",
"tokenizer_2",
"transformer",
"vae",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
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="timestep_preparation_stage",
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=Kandinsky5LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
),
)
self.add_stage(
stage_name="denoising_stage",
stage=Kandinsky5DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=Kandinsky5DecodingStage(vae=self.get_module("vae"), pipeline=self),
)
EntryClass = Kandinsky5DMDPipeline
@@ -0,0 +1,309 @@
# SPDX-License-Identifier: Apache-2.0
"""Preprocess Kandinsky5 overfit data into parquet format.
Encodes videos with Kandinsky5's VAE (shared with HunyuanVideo) and
captions with dual text encoders (Qwen/Reason1 + CLIP) into the t2v
parquet schema expected by the training framework.
Unlike Hunyuan's version of this script, text encoders are loaded via
FastVideo's own ``TextEncoderLoader``/``VAELoader`` (matching
``Kandinsky5Model.ensure_negative_conditioning``) rather than raw
``transformers`` classes -- Kandinsky5's Qwen/Reason1 encoder is a
native FastVideo wrapper around Qwen2.5-VL's language model, not a
plain ``AutoModel.from_pretrained`` load.
The CLIP pooled projection (``[768]``) is zero-padded into a row and
prepended to the Qwen sequence embeddings, then stored as a single
``[seq+1, dim]`` tensor in the existing ``text_embedding`` parquet
field -- the same trick ``preprocess_hunyuan_overfit.py`` uses for
LLaMA+CLIP. No schema change and no separate attention-mask field are
needed: the standard collator derives ``text_attention_mask`` from
padding against the stored (unpadded) row count, so the prepended
pooled row is automatically counted as valid.
Usage:
CUDA_VISIBLE_DEVICES=0 python -m fastvideo.pipelines.preprocess.preprocess_kandinsky5_overfit
Input/output roots default to ``data/kandinsky5_overfit`` /
``data/kandinsky5_overfit_preprocessed`` and can be overridden with the
``KANDINSKY5_OVERFIT_DATA_DIR`` / ``KANDINSKY5_OVERFIT_OUTPUT_DIR`` env
vars (used by the nightly e2e test to keep its disposable roots separate
from a user's real dataset at the defaults).
"""
import json
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29520")
import cv2
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from transformers import AutoTokenizer
from fastvideo.configs.pipelines.base import preprocess_text
from fastvideo.configs.pipelines.kandinsky5 import (
Kandinsky5T2VConfig,
kandinsky5_clip_postprocess_text,
kandinsky5_qwen_postprocess_text,
kandinsky5_qwen_preprocess_text,
)
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.models.loader.component_loader import TextEncoderLoader, VAELoader
from fastvideo.utils import maybe_download_model
# --- Config ---
# Matches KANDINSKY5_T2V_LITE_5S's native preset (fastvideo/pipelines/basic/
# kandinsky5/presets.py): 121 frames @ 24fps ~= 5.04s, matching the "5s" in
# the checkpoint name. 512/768 fall in the 480p RoPE scale_factor band
# Kandinsky5Model asserts on, and satisfy the real divisor requirement
# (spatial_compression_ratio(8) * patch_size[1](2) = 16 -- not just 8).
NUM_FRAMES = 121 # 4k+1 for temporal compression ratio 4
MAX_HEIGHT = 512
MAX_WIDTH = 768
TRAIN_FPS = 24.0
MODEL_PATH = "kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers"
# Overridable so automation (e.g. the nightly e2e test) can point at its own
# test-owned directories instead of clobbering the documented default paths
# a user may have populated with their real dataset.
DATA_DIR = os.environ.get("KANDINSKY5_OVERFIT_DATA_DIR", "data/kandinsky5_overfit")
OUTPUT_DIR = os.environ.get("KANDINSKY5_OVERFIT_OUTPUT_DIR", "data/kandinsky5_overfit_preprocessed")
def load_video(path: str, num_frames: int, height: int, width: int) -> torch.Tensor:
"""Load video as [1, C, T, H, W] in [-1, 1], resized to (height, width).
Source clips are frequently at native resolution (e.g. 1920x1080);
encoding that directly through the VAE (instead of at the target
training resolution) uses far more memory than intended and can OOM.
"""
cap = cv2.VideoCapture(path)
frames: list[np.ndarray] = []
while len(frames) < num_frames:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
if frame.shape[0] != height or frame.shape[1] != width:
frame = cv2.resize(frame, (width, height), interpolation=cv2.INTER_AREA)
frames.append(frame)
cap.release()
if not frames:
# Without this, the repeat-last-frame fill below raises an opaque
# IndexError on frames[-1]; cv2.VideoCapture doesn't raise on a
# missing/corrupt file, it just decodes nothing.
raise ValueError(f"Could not decode any frames from video {path!r} -- "
"the file is missing, empty, or not readable by OpenCV.")
if len(frames) < num_frames:
# Repeat last frame to fill
while len(frames) < num_frames:
frames.append(frames[-1])
frames = frames[:num_frames]
video = np.stack(frames, axis=0)
video = torch.from_numpy(video).float()
video = video / 127.5 - 1.0 # [0,255] -> [-1,1]
video = video.permute(3, 0, 1, 2).unsqueeze(0) # [1,C,T,H,W]
return video
def get_caption(item: dict) -> str:
"""Extract a caption from one ``videos2caption.json`` entry.
The documented schema (``docs/training/data_preprocess.md``) stores
``cap`` as a plain string; some producers (e.g. this repo's own e2e
test fixtures, mirroring ``preprocess_hunyuan_overfit.py``) instead
store a non-empty list of caption variants and use the first one.
Indexing unconditionally with ``item["cap"][0]`` silently takes the
first *character* of a string caption instead of erroring, so accept
and validate both forms explicitly here.
"""
cap = item.get("cap")
if isinstance(cap, str):
if not cap:
raise ValueError(f"Empty 'cap' string for entry: {item}")
return cap
if isinstance(cap, list):
if not cap or not isinstance(cap[0], str) or not cap[0]:
raise ValueError(f"'cap' list must be non-empty with a non-empty first string, got: {item}")
return cap[0]
raise ValueError(f"'cap' must be a string or a list of strings, got {type(cap).__name__}: {item}")
def main() -> None:
# FastVideo's native text-encoder layers (e.g. CLIPAttention's
# QKVParallelLinear) are tensor-parallel-aware and assert the TP process
# group is initialized, even for a single-process/single-GPU run.
maybe_init_distributed_environment_and_model_parallel(1, 1)
device = torch.device("cuda:0")
model_path = maybe_download_model(MODEL_PATH)
os.makedirs(OUTPUT_DIR, exist_ok=True)
# Load captions. Validate up front (before spending minutes loading the
# VAE + two text encoders): an empty/malformed manifest would otherwise
# only surface as an IndexError on records[0] at parquet-write time.
manifest_path = os.path.join(DATA_DIR, "videos2caption.json")
with open(manifest_path) as f:
caption_data = json.load(f)
if not isinstance(caption_data, list) or not caption_data:
raise ValueError(f"Manifest {manifest_path} must be a non-empty JSON list of "
f"{{'path', 'cap'}} entries, got: {type(caption_data).__name__} "
f"with {len(caption_data) if isinstance(caption_data, list) else 'n/a'} entries")
pipeline_config = Kandinsky5T2VConfig()
# Kandinsky5T2VConfig.__post_init__ sets load_encoder=False by default --
# T2V inference only ever decodes generated latents, never encodes real
# video. Preprocessing needs the encoder to turn real clips into latents.
pipeline_config.vae_config.load_encoder = True
fastvideo_args = FastVideoArgs(
model_path=model_path,
dit_cpu_offload=False,
dit_layerwise_offload=False,
use_fsdp_inference=False,
text_encoder_cpu_offload=False,
vae_cpu_offload=False,
pipeline_config=pipeline_config,
)
fastvideo_args.device = device
# --- Load VAE (shared with HunyuanVideo) ---
print("Loading Kandinsky5 VAE...")
vae = VAELoader().load(os.path.join(model_path, "vae"), fastvideo_args)
vae = vae.to(device=device, dtype=torch.float16).eval()
print(f"VAE loaded ({sum(p.numel() for p in vae.parameters()) / 1e6:.0f}M)")
# --- Load text encoders ---
print("Loading Qwen/Reason1 text encoder...")
qwen_enc = TextEncoderLoader().load(
os.path.join(model_path, "text_encoder"),
fastvideo_args,
).to(device).eval()
qwen_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
qwen_tok_kwargs = dict(pipeline_config.text_encoder_configs[0].tokenizer_kwargs)
# TextEncodingStage overrides max_length from pipeline_config.text_encoder_max_lengths
# at runtime (fastvideo/pipelines/stages/text_encoding.py) -- Reason1Config's static
# tokenizer_kwargs default (text_len=512) doesn't account for the 129-token Kandinsky
# system template ENCODE_START_IDX strips off afterwards, so mirror the runtime value
# here or training conditions on fewer caption tokens than inference does.
qwen_tok_kwargs["max_length"] = pipeline_config.text_encoder_max_lengths[0]
print("Loading CLIP text encoder...")
clip_enc = TextEncoderLoader().load(
os.path.join(model_path, "text_encoder_2"),
fastvideo_args,
).to(device).eval()
clip_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer_2"))
clip_tok_kwargs = dict(pipeline_config.text_encoder_configs[1].tokenizer_kwargs)
clip_tok_kwargs["max_length"] = pipeline_config.text_encoder_max_lengths[1]
# --- Process each video ---
records = []
for item in caption_data:
video_name = item["path"]
caption = get_caption(item)
video_path = os.path.join(DATA_DIR, "videos", video_name)
print(f"\nProcessing: {video_name}")
print(f" Caption: {caption[:80]}...")
# Encode video. Kandinsky5's VAE encode returns channel-first
# [B, C, T, H, W], matching the storage convention Kandinsky5Model
# expects (it permutes to channel-last only right before calling
# the transformer).
video = load_video(video_path, NUM_FRAMES, MAX_HEIGHT, MAX_WIDTH).to(device=device, dtype=torch.float16)
print(f" Video shape: {video.shape}")
with torch.no_grad():
latent_dist = vae.encode(video)
latent = latent_dist.mean.squeeze(0).float().cpu() # [C, T, H, W]
print(f" Latent shape: {latent.shape}")
# Encode text with Qwen/Reason1.
qwen_text = kandinsky5_qwen_preprocess_text(caption)
with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None):
qwen_inputs = qwen_tok(qwen_text, **qwen_tok_kwargs).to(device)
qwen_out = qwen_enc(
input_ids=qwen_inputs.input_ids,
attention_mask=qwen_inputs.attention_mask,
output_hidden_states=True,
)
qwen_embeds, _qwen_mask = kandinsky5_qwen_postprocess_text(qwen_out, qwen_inputs.attention_mask)
qwen_embeds = qwen_embeds.squeeze(0) # [seq, dim]
# Encode text with CLIP.
clip_text = preprocess_text(caption)
with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None):
clip_inputs = clip_tok(clip_text, **clip_tok_kwargs).to(device)
clip_out = clip_enc(
input_ids=clip_inputs.input_ids,
attention_mask=clip_inputs.attention_mask,
)
clip_pooled = kandinsky5_clip_postprocess_text(clip_out).squeeze(0) # [768]
# Combine: [pooled_clip_row, qwen_embeds]
qwen_dim = qwen_embeds.shape[-1]
pooled_row = torch.zeros(qwen_dim, device=device, dtype=torch.float16)
pooled_row[:clip_pooled.shape[-1]] = clip_pooled
text_embedding = torch.cat(
[pooled_row.unsqueeze(0), qwen_embeds],
dim=0,
).float().cpu() # [seq+1, dim]
print(f" Text embedding shape: {text_embedding.shape}")
record = {
"id": video_name,
"vae_latent_bytes": latent.numpy().tobytes(),
"vae_latent_shape": list(latent.shape),
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
"text_embedding_bytes": text_embedding.numpy().tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
"file_name": video_name,
"caption": caption,
"media_type": "video",
"width": MAX_WIDTH,
"height": MAX_HEIGHT,
"num_frames": NUM_FRAMES,
"duration_sec": NUM_FRAMES / TRAIN_FPS,
"fps": TRAIN_FPS,
}
records.append(record)
# Clean up encoders
del qwen_enc, qwen_tok, clip_enc, clip_tok, vae
# Write parquet
table = pa.table(
{k: [r[k] for r in records]
for k in records[0]},
schema=pyarrow_schema_t2v,
)
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
pq.write_table(table, output_path)
print(f"\nWrote {len(records)} records to {output_path}")
# Write validation prompts for callback.
# Wrap in "data" key -- ValidationDataset expects field="data".
# Use "caption" field -- ValidationDataset aliases it to "prompt".
val_prompts = {"data": [{"caption": get_caption(item)} for item in caption_data]}
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
with open(val_path, "w") as f:
json.dump(val_prompts, f, indent=2)
print(f"Wrote validation prompts to {val_path}")
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
if __name__ == "__main__":
main()
+269 -2
View File
@@ -10,12 +10,18 @@ import torch
from diffusers.utils.torch_utils import randn_tensor
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention import LocalAttention
from fastvideo.attention.backends.nabla import NablaAttentionMetadataBuilder
from fastvideo.attention.selector import backend_name_to_enum
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.models.loader.component_loader import TransformerLoader, VAELoader
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.models.vaes.common import ParallelTiledVAE
from fastvideo.models.vision_utils import normalize, numpy_to_pt, pil_to_numpy, resize
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
@@ -24,7 +30,8 @@ from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import PRECISION_TO_TYPE
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import PRECISION_TO_TYPE, get_mixed_precision_state
logger = init_logger(__name__)
@@ -166,6 +173,115 @@ class Kandinsky5DenoisingStage(PipelineStage):
seq_len = int(mask.sum(1).max().item())
return torch.arange(seq_len, device=device)
# Backends whose selection is all-or-nothing: fastvideo.platforms.cuda
# raises ImportError immediately if the ATTN_QAT_TRAIN kernel isn't
# available (training must never silently de-quantize), rather than the
# generic "requested backend unsupported by this layer, falling back to
# automatic selection" path _cached_get_attn_backend takes for most
# other backend names (see fastvideo/attention/selector.py).
# ATTN_QAT_TRAIN is also always in Kandinsky5ArchConfig's
# _supported_attention_backends, so that generic fallback never applies
# to it here. That combination makes an exact backend match a safe,
# false-positive-free signal *only* for ATTN_QAT_TRAIN.
# ATTN_QAT_INFER must NOT be in this set: fastvideo.platforms.cuda
# deliberately resolves it to FlashAttention when the sm_120 kernel
# isn't built, and the QAD recipe README documents that fallback ("on
# other GPUs the attention falls back to a supported dense backend
# while the FP4 linear layers still run") -- asserting an exact match
# would abort inference on every non-sm_120 machine relying on it.
# Any other FASTVIDEO_ATTENTION_BACKEND value (e.g. a backend meant for
# a different model family sharing the same process/env var) can
# likewise legitimately resolve to something other than what was
# requested.
_STRICT_BACKENDS = frozenset({
AttentionBackendEnum.ATTN_QAT_TRAIN,
})
def _assert_local_attention_backend_engaged(self) -> None:
"""Guard against ``LocalAttention`` silently falling back to SDPA.
``Kandinsky5Attention.forward`` catches the ``AssertionError``
``LocalAttention`` raises when no pipeline forward context is set and
falls back to plain ``F.scaled_dot_product_attention`` -- silently
skipping whichever kernel ``FASTVIDEO_ATTENTION_BACKEND`` requested
(the fake-quantized ``ATTN_QAT_TRAIN`` kernel).
``LocalAttention.backend`` is resolved once at module construction,
independent of whether a forward context is later set, so this only
catches a backend that failed to resolve to what the env var
requested -- it does not prove forward context was present for a
given forward call. Pair it with always wrapping the actual
transformer call in ``set_forward_context`` (see the denoising loops
below), which is what prevents the runtime fallback.
Scoped to ``_STRICT_BACKENDS``: other backend names -- including
``ATTN_QAT_INFER``, whose FlashAttention fallback on non-sm_120
GPUs is documented, intentional behavior -- can legitimately differ
from what ``LocalAttention.backend`` resolved to (see
``_STRICT_BACKENDS``'s comment), so asserting on those would flag
expected behavior as a bug.
"""
backend_env = envs.FASTVIDEO_ATTENTION_BACKEND
if not backend_env:
return
expected = backend_name_to_enum(backend_env)
if expected is None or expected not in self._STRICT_BACKENDS:
return
for module in self.transformer.modules():
if isinstance(module, LocalAttention):
assert module.backend == expected, (
f"Kandinsky5 local attention resolved to backend {module.backend}, expected "
f"{expected} from FASTVIDEO_ATTENTION_BACKEND={backend_env}. This likely means "
"LocalAttention's missing-forward-context guard silently fell back to SDPA.")
return
def _resolve_target_dtype(self, fastvideo_args: FastVideoArgs) -> torch.dtype:
"""Resolve the transformer's actual compute dtype.
Trust ``pipeline_config.dit_precision`` directly for a normal,
non-FSDP load -- ``TransformerLoader.load()`` asserts every
parameter matches it exactly (``fastvideo/models/loader/
component_loader.py``), so this holds for standalone T2V/I2V
inference regardless of which precision (including fp32) was
requested.
FSDP2-wrapped modules are the one case where that assertion doesn't
carry forward: ``maybe_load_fsdp_model`` hardcodes its
``MixedPrecisionPolicy`` to ``param_dtype=torch.bfloat16`` (with
``cast_forward_inputs=False``) for every FSDP-wrapped load,
independent of ``dit_precision`` -- this covers both the live
transformer ``ValidationCallback`` reuses from training (whose
``pipeline_config.dit_precision`` reflects the fp32 master-weight
load dtype, not the actual bf16 compute dtype) and multi-GPU
``use_fsdp_inference=True`` runs.
For that FSDP case, read the policy itself
(``set_mixed_precision_policy`` records it right before
``maybe_load_fsdp_model`` shards -- every FSDP wrap in this repo
goes through that path), NOT the parameter storage dtypes:
FSDP2 computes in the policy's ``param_dtype`` regardless of what
dtype the parameters are stored/loaded in, so e.g. an fp16-loaded
FSDP model still runs its forward in bf16 -- scanning parameters
would resolve fp16 and, with ``cast_forward_inputs=False`` and
autocast keyed off the wrong dtype, hand bf16-computing modules
fp16 inputs.
"""
declared_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
try:
from torch.distributed.fsdp import FSDPModule
except Exception: # pragma: no cover - FSDP not always available
return declared_dtype
if not isinstance(self.transformer, FSDPModule):
return declared_dtype
try:
policy_dtype = get_mixed_precision_state().param_dtype
except ValueError:
policy_dtype = None
# bf16 mirrors maybe_load_fsdp_model's hardcoded param_dtype if the
# (thread-local) policy state is somehow unset here.
return policy_dtype if policy_dtype is not None else torch.bfloat16
@staticmethod
def fast_sta_nabla(
T: int,
@@ -267,9 +383,10 @@ class Kandinsky5DenoisingStage(PipelineStage):
loader = TransformerLoader()
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
fastvideo_args.model_loaded["transformer"] = True
self._assert_local_attention_backend_engaged()
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
target_dtype = self._resolve_target_dtype(fastvideo_args)
autocast_enabled = target_dtype != torch.float32 and not fastvideo_args.disable_autocast
latents = batch.latents
num_channels = getattr(
@@ -391,6 +508,156 @@ class Kandinsky5DenoisingStage(PipelineStage):
return result
class Kandinsky5DmdDenoisingStage(Kandinsky5DenoisingStage):
"""DMD (few fixed steps, no CFG) variant of Kandinsky5DenoisingStage.
Reuses the parent's RoPE/scale_factor/sparse-params helpers; only the
denoising loop differs: a fixed short timestep schedule
(``pipeline_config.dmd_denoising_steps``) with a single forward pass per
step and no classifier-free-guidance branch, matching DMD's distilled
few-step generator.
``dmd_denoising_steps`` values (e.g. ``[1000, 750, 500, 250]``) are
literal *final* target timesteps -- that's how ``DMD2Method`` on the
training side resolves them: nearest-sigma lookup against a scheduler
that has never had ``set_timesteps`` called on it, so e.g. "750" maps to
sigma 0.75. This stage therefore keeps a private
``FlowMatchEulerDiscreteScheduler`` (same ``shift`` as the pipeline
scheduler) and drives it directly via predict-x0 + re-noise, exactly
mirroring ``DMD2Method._student_rollout``'s "simulate" branch and Wan's
own ``DmdDenoisingStage``. It must NOT reuse the pipeline's shared
``scheduler`` object through ``scheduler.step()``: by the time this
stage runs, ``TimestepPreparationStage`` has already called
``scheduler.set_timesteps(timesteps=dmd_denoising_steps)`` on it, which
re-applies the flow-match ``shift`` warp on top of values that are
already final (e.g. sigma 0.75 -> 0.9375 at shift=5) -- a double shift
that leaves every step far noisier than the student was trained for, so
the sampled video is still mostly noise after the last step instead of
converged.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
self._sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=scheduler.shift)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None:
raise ValueError("latents must be prepared before Kandinsky5 DMD denoising.")
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
fastvideo_args.model_loaded["transformer"] = True
self._assert_local_attention_backend_engaged()
device = get_local_torch_device()
target_dtype = self._resolve_target_dtype(fastvideo_args)
autocast_enabled = target_dtype != torch.float32 and not fastvideo_args.disable_autocast
latents = batch.latents
num_channels = getattr(
self.transformer,
"in_visual_dim",
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
)
prompt_embeds = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
pooled = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
if batch.prompt_attention_mask is None or not batch.prompt_attention_mask:
raise ValueError("Kandinsky5 DMD requires Qwen prompt attention masks.")
text_rope_pos = self._text_rope_pos(batch.prompt_attention_mask[0].to(device), device)
height = int(batch.height)
width = int(batch.width)
temporal_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
num_latent_frames = (int(batch.num_frames) - 1) // temporal_ratio + 1
visual_rope_pos = [
torch.arange(num_latent_frames, device=device),
torch.arange(height // spatial_ratio // 2, device=device),
torch.arange(width // spatial_ratio // 2, device=device),
]
scale_factor = self._scale_factor(height, width)
sparse_params = self.get_sparse_params(latents, device)
dmd_steps = fastvideo_args.pipeline_config.dmd_denoising_steps
if not dmd_steps:
raise ValueError("Kandinsky5 DMD denoising requires "
"pipeline_config.dmd_denoising_steps to be set.")
timesteps = torch.tensor(dmd_steps, dtype=torch.long, device=device)
# I2V keeps the first (conditioning) frame fixed during denoising.
# Key off the actual image conditioning, not transformer.visual_cond:
# official T2V checkpoints also ship visual_cond=True, and skipping
# frame 0 for them leaves it as undenoised noise (see the same fix
# in Kandinsky5DenoisingStage.forward() above).
cond_frames = 1 if batch.image_latent is not None else 0
with tqdm(total=len(timesteps), desc="Kandinsky5 DMD Denoising") as progress_bar:
for i, timestep in enumerate(timesteps):
if hasattr(self, "interrupt") and self.interrupt:
continue
t_expand = timestep.unsqueeze(0).repeat(latents.shape[0]).to(device=device, dtype=target_dtype)
attn_metadata = None
if sparse_params is not None:
attn_metadata = NablaAttentionMetadataBuilder().build(
current_timestep=i,
sta_mask=sparse_params["sta_mask"],
P=sparse_params["P"],
visual_shape=sparse_params["visual_shape"],
)
autocast_ctx = (torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled)
if device.type == "cuda" else contextlib.nullcontext())
with set_forward_context(current_timestep=i, attn_metadata=attn_metadata,
forward_batch=batch), autocast_ctx:
pred_velocity = self.transformer(
hidden_states=latents.to(dtype=target_dtype),
encoder_hidden_states=prompt_embeds,
pooled_projections=pooled,
timestep=t_expand,
visual_rope_pos=visual_rope_pos,
text_rope_pos=text_rope_pos,
scale_factor=scale_factor,
sparse_params=sparse_params,
return_dict=True,
).sample
# Channel-last [B, T', H, W, C] -> channel-first [B, T', C, H, W]
# to match pred_noise_to_pred_video/add_noise's expected layout
# (same convention as Kandinsky5Model.add_noise/predict_x0).
sample_cf = latents[:, cond_frames:, :, :, :num_channels].permute(0, 1, 4, 2, 3)
velocity_cf = pred_velocity[:, cond_frames:].permute(0, 1, 4, 2, 3)
b, t = sample_cf.shape[:2]
step_timestep = timestep.reshape(1).to(device=device)
pred_x0_cf = pred_noise_to_pred_video(
pred_noise=velocity_cf.flatten(0, 1),
noise_input_latent=sample_cf.flatten(0, 1),
timestep=step_timestep,
scheduler=self._sample_scheduler,
).unflatten(0, (b, t))
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1].reshape(1).to(device=device)
noise = randn_tensor(
sample_cf.shape,
generator=batch.generator,
device=device,
dtype=pred_x0_cf.dtype,
)
next_cf = self._sample_scheduler.add_noise(
pred_x0_cf.flatten(0, 1),
noise.flatten(0, 1),
next_timestep,
).unflatten(0, (b, t))
else:
next_cf = pred_x0_cf
latents[:, cond_frames:, :, :, :num_channels] = next_cf.permute(0, 1, 3, 4, 2).to(latents.dtype)
progress_bar.update()
batch.latents = latents[:, :, :, :, :num_channels]
return batch
class Kandinsky5DecodingStage(DecodingStage):
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
@@ -0,0 +1,109 @@
# SPDX-License-Identifier: Apache-2.0
"""ATTN_QAT_INFER selection must be gated on device capability, through
the real platform resolver.
``is_attn_qat_infer_available()`` used to test only whether the kernel
extension imports. CUDA 13 wheel builds can carry the sm_120
extension on any host (e.g. H100 sm_90, GB200 sm_100): the import
succeeds, ``CudaPlatformBase.get_attn_backend_cls`` selects the
consumer-Blackwell backend, and the first kernel call fails with an
unsupported-capability error -- instead of the FlashAttention fallback
the QAD README documents for non-sm_120 GPUs.
These tests drive the REAL resolver (``fastvideo.platforms.cuda``) and
the REAL availability function; only the two physical facts are faked --
"does the extension import" (``_get_attn_qat_infer``) and "what GPU is
active" (``torch.cuda``). The stage-guard test
(fastvideo/tests/stages/test_kandinsky5_attention_backend_guard.py)
injects an already-resolved backend and by design cannot see this bug.
CPU-only: the ATTN_QAT_INFER branch and its fallback never require a
physical GPU to *resolve* (only to run).
"""
from __future__ import annotations
import torch
import fastvideo.attention.backends.attn_qat_infer as attn_qat_infer_module
from fastvideo.attention.backends.attn_qat_infer import is_attn_qat_infer_available
# The concrete resolver whose device queries go straight through torch.cuda
# (the NVML variant only differs in device-info plumbing, not selection
# logic) -- the abstract CudaPlatformBase can't run the fallthrough's
# has_device_capability() check.
from fastvideo.platforms.cuda import NonNvmlCudaPlatform
from fastvideo.platforms.interface import AttentionBackendEnum
ATTN_QAT_INFER_CLS = "fastvideo.attention.backends.attn_qat_infer.AttnQatInferBackend"
# What the resolver's fallthrough legitimately returns when ATTN_QAT_INFER
# is unavailable: FlashAttention, or SDPA when flash_attn isn't installed
# in the running environment (e.g. CPU-only CI).
FALLBACK_CLASSES = {
"fastvideo.attention.backends.flash_attn.FlashAttentionBackend",
"fastvideo.attention.backends.sdpa.SDPABackend",
}
def _fake_gpu(monkeypatch, *, capability: tuple[int, int], extension_imports: bool) -> None:
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device=None: capability)
monkeypatch.setattr(
attn_qat_infer_module,
"_get_attn_qat_infer",
lambda: (lambda *a, **k: None) if extension_imports else None,
)
def _resolve() -> str:
return NonNvmlCudaPlatform.get_attn_backend_cls(
AttentionBackendEnum.ATTN_QAT_INFER,
head_size=128,
dtype=torch.bfloat16,
)
def test_sm90_host_with_bundled_extension_falls_back(monkeypatch):
"""The reviewed failure: H100 + CUDA 13 wheel that bundles the sm_120
extension. Import succeeds; selection must still fall back."""
_fake_gpu(monkeypatch, capability=(9, 0), extension_imports=True)
assert not is_attn_qat_infer_available()
assert _resolve() in FALLBACK_CLASSES
def test_sm100_host_with_bundled_extension_falls_back(monkeypatch):
_fake_gpu(monkeypatch, capability=(10, 0), extension_imports=True)
assert not is_attn_qat_infer_available()
assert _resolve() in FALLBACK_CLASSES
def test_sm120_with_extension_selects_backend(monkeypatch):
_fake_gpu(monkeypatch, capability=(12, 0), extension_imports=True)
assert is_attn_qat_infer_available()
assert _resolve() == ATTN_QAT_INFER_CLS
def test_sm121_with_bundled_extension_falls_back(monkeypatch):
_fake_gpu(monkeypatch, capability=(12, 1), extension_imports=True)
assert not is_attn_qat_infer_available()
assert _resolve() in FALLBACK_CLASSES
def test_consumer_blackwell_without_extension_falls_back(monkeypatch):
_fake_gpu(monkeypatch, capability=(12, 0), extension_imports=False)
assert not is_attn_qat_infer_available()
assert _resolve() in FALLBACK_CLASSES
def test_no_cuda_reports_unavailable(monkeypatch):
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
monkeypatch.setattr(
attn_qat_infer_module,
"_get_attn_qat_infer",
lambda: (lambda *a, **k: None),
)
assert not is_attn_qat_infer_available()
@@ -0,0 +1,63 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression test for exported Kandinsky5 DMD checkpoints resolving to the
wrong pipeline.
``fastvideo.train.entrypoint.dcp_to_diffusers`` copies the base T2V
checkpoint's ``model_index.json`` unchanged into every export (see
``examples/train/configs/fine_tuning/kandinsky5/README.md``), so
``_class_name`` still says the base T2V pipeline. Without an explicit
override, ``fastvideo.registry.get_model_info`` resolves such a directory to
``Kandinsky5T2VPipeline`` (via the path-based "t2v" fallback detector in
``fastvideo/registry.py``) instead of ``Kandinsky5DMDPipeline`` -- silently
running the full-length multi-step sampler on DMD-distilled weights instead
of the four-step re-noise sampler. ``override_pipeline_cls_name`` bypasses
that resolution entirely and must be paired with ``Kandinsky5DMDConfig`` (not
the default ``Kandinsky5T2VConfig``, whose ``dmd_denoising_steps=None`` makes
``Kandinsky5DmdDenoisingStage`` raise immediately).
"""
import json
from fastvideo.configs.pipelines.kandinsky5 import (
Kandinsky5DMDConfig,
Kandinsky5T2VConfig,
)
from fastvideo.pipelines.basic.kandinsky5.kandinsky5_dmd_pipeline import Kandinsky5DMDPipeline
from fastvideo.pipelines.basic.kandinsky5.kandinsky5_pipeline import Kandinsky5T2VPipeline
from fastvideo.pipelines.pipeline_registry import PipelineType
from fastvideo.registry import get_model_info
def _write_fake_export(tmp_path, name: str):
"""Minimal on-disk stand-in for a dcp_to_diffusers export: just enough
for verify_model_config_and_directory to accept it."""
model_dir = tmp_path / name
(model_dir / "transformer").mkdir(parents=True)
(model_dir / "model_index.json").write_text(
json.dumps({
"_class_name": "Kandinsky5T2VPipeline",
"_diffusers_version": "0.30.0",
}))
return model_dir
def test_kandinsky5_dmd_config_sets_denoising_steps():
assert Kandinsky5T2VConfig().dmd_denoising_steps is None
assert Kandinsky5DMDConfig().dmd_denoising_steps == [1000, 750, 500, 250]
def test_exported_dmd_checkpoint_resolves_to_t2v_pipeline_without_override(tmp_path):
"""Documents the bug: an unmodified export resolves to the T2V pipeline."""
model_dir = _write_fake_export(tmp_path, "kandinsky5_t2v_dmd2_4steps_qat")
info = get_model_info(model_path=str(model_dir), pipeline_type=PipelineType.BASIC)
assert info.pipeline_cls is Kandinsky5T2VPipeline
def test_override_pipeline_cls_name_resolves_dmd_pipeline(tmp_path):
model_dir = _write_fake_export(tmp_path, "kandinsky5_t2v_dmd2_4steps_qat_override")
info = get_model_info(
model_path=str(model_dir),
pipeline_type=PipelineType.BASIC,
override_pipeline_cls_name="Kandinsky5DMDPipeline",
)
assert info.pipeline_cls is Kandinsky5DMDPipeline
@@ -0,0 +1,523 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-sample overfit / small end-to-end run for Kandinsky5 QAD:
stage-1 Attn-QAT finetune -> stage-2 QAT-aware DMD distillation, on the
new ``fastvideo/train/`` stack -- validated all the way to the public
inference artifact this recipe exists to produce.
Unlike ``test_e2e_dmd_t2v_crush_smol.py`` (which drives the legacy
``fastvideo/training/wan_distillation_pipeline.py`` CLI), this drives the
new-stack entrypoint (``fastvideo.train.entrypoint.train``) via the same
YAML configs used for real training
(``examples/train/configs/fine_tuning/kandinsky5/t2v_480p_qat.yaml`` and
``examples/train/configs/distribution_matching/kandinsky5/dmd2_t2v_480p_qat.yaml``),
with a handful of steps and dotted-key overrides pointing at a tiny local
dataset -- mirroring how a real run would be launched, just shrunk down.
The full validated chain, each arrow a hard assertion:
synthetic clip -> preprocess -> stage-1 train -> stage-1 DCP
-> dcp_to_diffusers --verify (strict reload)
-> stage-2 train (student/teacher/critic init_from that export)
-> stage-2 student DCP -> dcp_to_diffusers --verify (strict reload)
-> VideoGenerator.from_pretrained(export,
override_pipeline_cls_name="Kandinsky5DMDPipeline",
pipeline_config=Kandinsky5DMDConfig()) # the documented recipe
-> deterministic (fixed-seed) 4-step generation
-> degeneracy checks + MS-SSIM against the committed reference video.
Every artifact (raw clip, parquet, both stage output dirs, both diffusers
exports, generated video dir) lives under the single test-owned
``data/kandinsky5_e2e/`` root, which is deleted up front -- the
documented user dataset paths (``data/kandinsky5_overfit{,_preprocessed}``)
are never read, written, or removed. ``training.checkpoint.
resume_from_checkpoint`` is explicitly disabled for both stages: the
stage-1 YAML defaults to ``resume_from_checkpoint: latest``, so a rerun
over a dirty ``data/`` dir could otherwise no-op from an old checkpoint
and go green on stale artifacts.
Reference bootstrap: the committed reference
(``reference_video_kandinsky5_dmd_v0.mp4`` next to this file) is produced
by running this test once on a sanctioned GPU box with
``KANDINSKY5_E2E_WRITE_REFERENCE=1``, reviewing the written video by eye,
and committing it. Review bar: expect blurry-but-structured content
loosely matching the prompt (4-step no-CFG sampling of near-base weights
cannot look good) -- soft shapes, warm sunflower-ish colors, motion.
Uniform grey static, solid black, or hard garbage blocks mean a broken
pipeline: do NOT commit such a reference (probe the export with the plain
50-step ``Kandinsky5T2VPipeline`` to isolate weights-path vs sampler
problems). Until a reference exists the test FAILS (it does not skip) --
an artifact-existence-only pass is exactly the false-green this oracle
exists to prevent. Run it with::
pytest fastvideo/tests/nightly/test_e2e_kandinsky5_dmd_t2v_overfit.py -vs
The single training sample is a synthesized noise clip (not a downloaded
dataset): the exact directory/manifest layout of external raw-video HF
datasets (e.g. the ``crush-smol`` one Wan's e2e test uses) isn't something
this change can verify without a runnable environment, so depending on it
here would risk a silently-wrong assumption. ``preprocess_kandinsky5_overfit.py``
only needs a ``videos2caption.json`` + ``videos/*.mp4`` pair, both of which
this test fully controls.
"""
from __future__ import annotations
import json
import os
import shutil
import subprocess
import sys
from pathlib import Path
import cv2
import numpy as np
import pytest
REPO_ROOT = Path(__file__).resolve().parents[3]
THIS_FILE = Path(__file__).resolve()
NUM_GPUS = os.environ.get("KANDINSKY5_E2E_NUM_GPUS", "1")
# Every artifact this test reads or writes lives under this single
# test-owned root. It must NEVER point at (or contain) the documented
# default dataset paths -- data/kandinsky5_overfit and
# data/kandinsky5_overfit_preprocessed -- which a user may have populated
# with their real videos/captions for the actual recipe:
# _clean_previous_artifacts() deletes this root wholesale on every run,
# and the preprocess subprocess is pointed here via the
# KANDINSKY5_OVERFIT_DATA_DIR / KANDINSKY5_OVERFIT_OUTPUT_DIR env
# overrides instead of the script's user-facing defaults.
E2E_ROOT = Path("data") / "kandinsky5_e2e"
RAW_DATA_DIR = E2E_ROOT / "raw"
PREPROCESSED_DIR = E2E_ROOT / "preprocessed"
STAGE1_OUTPUT_DIR = E2E_ROOT / "stage1"
STAGE2_OUTPUT_DIR = E2E_ROOT / "stage2"
# Kept out of STAGE{1,2}_OUTPUT_DIR: those directories are scanned with
# glob("checkpoint-*") below to find the latest raw DCP checkpoint, and a
# "checkpoint-<N>-diffusers" export dir sitting alongside "checkpoint-<N>"
# would also match that pattern -- on a rerun that doesn't start from a
# clean data/ dir, it would sort after "checkpoint-<N>" and get picked as
# the checkpoint instead, feeding an already-exported diffusers directory
# back into dcp_to_diffusers as if it were a DCP checkpoint.
STAGE1_DIFFUSERS_DIR = E2E_ROOT / "stage1_diffusers"
STAGE2_DIFFUSERS_DIR = E2E_ROOT / "stage2_diffusers"
GENERATED_VIDEO_DIR = E2E_ROOT / "generated"
STAGE1_CONFIG = (REPO_ROOT / "examples" / "train" / "configs" / "fine_tuning" / "kandinsky5" /
"t2v_480p_qat.yaml")
STAGE2_CONFIG = (REPO_ROOT / "examples" / "train" / "configs" / "distribution_matching" / "kandinsky5" /
"dmd2_t2v_480p_qat.yaml")
# Deterministic-generation oracle. The reference is generated once (see the
# module docstring's bootstrap procedure), reviewed by a human, and
# committed next to this file -- mirroring reference_video_1_sample_v0.mp4
# in the Wan e2e test.
REFERENCE_VIDEO = THIS_FILE.parent / "reference_video_kandinsky5_dmd_v0.mp4"
WRITE_REFERENCE_ENV = "KANDINSKY5_E2E_WRITE_REFERENCE"
# Deliberately semantic even though the training clip is random noise: the
# caption's relation to the clip content is irrelevant for plumbing, but
# the caption doubles as the generation prompt, and a semantic prompt makes
# the reference video human-reviewable (blurry-but-structured output from
# 4-step sampling of near-base weights). An earlier "random noise" caption
# produced a grey-static reference that was visually indistinguishable
# from a broken pipeline AND decorrelates across runs, making the SSIM
# oracle both unreviewable and potentially flaky.
GENERATION_PROMPT = ("A curious raccoon peers through a vibrant field of yellow sunflowers, "
"soft natural light filtering through the petals, mid-shot")
GENERATION_SEED = 42
# Media-oracle calibration, measured against the committed reference
# (2026-07-20, from the generator_update_interval=1 / 2-GPU-sharded
# bootstrap run; regression-tested in
# fastvideo/tests/workflow/test_kandinsky5_e2e_media_oracle.py, which
# recomputes these known-bad clips from whatever reference is currently
# committed -- so it stays valid across reference swaps even though the
# numbers below are a point-in-time snapshot):
#
# clip spatial-std temporal-diff mean MS-SSIM
# good reference >= 11.42 >= 1.12 1.0000
# solid @ reference mean RGB 0.0 0.0 0.9162
# frozen (reference frame x121) 12.12 0.0 0.8534
#
# Two consequences drive the design below:
# - MS-SSIM alone CANNOT separate collapsed clips from good ones with
# safe margin (a solid clip scores ~0.92!), so the structural checks
# in _assert_video_not_degenerate are the primary defense against
# solid/frozen/truncated output; the SSIM floor's job is catching a
# *different structured video* (wrong remap / re-noise schedule).
# - The floors are set >= 5x below the good reference's observed values
# and >= 10x above every known-bad clip's.
# An independent retrain-from-scratch validation run against this exact
# reference is the last open verification step (see the module
# docstring's bootstrap procedure) -- 0.95 is chosen to sit comfortably
# above the ~0.92 solid-clip ceiling while leaving headroom for ordinary
# run-to-run drift once that number is in.
EXPECTED_FRAME_COUNT = 121
EXPECTED_FRAME_HEIGHT = 512
EXPECTED_FRAME_WIDTH = 768
MIN_FRAME_SPATIAL_STD = 2.0
MIN_FRAME_TEMPORAL_DIFF = 0.2
MIN_MEAN_MS_SSIM = 0.95
def _clean_previous_artifacts() -> None:
"""Delete the single test-owned root a previous run could have left.
Outputs matter as much as inputs: stage 1's YAML defaults to
``resume_from_checkpoint: latest``, and the final assertions glob for
checkpoints/videos -- stale files from an earlier run could satisfy
them without this run doing any work.
Scoped strictly to ``E2E_ROOT``: an earlier version of this cleanup
also removed ``data/kandinsky5_overfit{,_preprocessed}``, the
documented default dataset paths of the real recipe -- silently and
irreversibly deleting whatever real videos/captions a user had
prepared there. Nothing outside the test-owned root may be touched.
"""
if E2E_ROOT.exists():
shutil.rmtree(E2E_ROOT)
def _synthesize_single_sample() -> None:
"""Write one tiny synthetic clip + caption manifest in the layout
preprocess_kandinsky5_overfit.py expects.
Matches KANDINSKY5_T2V_LITE_5S's native preset (512x768, 121 frames,
24fps -- see preprocess_kandinsky5_overfit.py's own constants).
"""
videos_dir = RAW_DATA_DIR / "videos"
videos_dir.mkdir(parents=True, exist_ok=True)
video_path = videos_dir / "sample_0.mp4"
writer = cv2.VideoWriter(
str(video_path),
cv2.VideoWriter_fourcc(*"mp4v"),
24.0,
(768, 512),
)
if not writer.isOpened():
raise RuntimeError(
f"cv2.VideoWriter failed to open {video_path} (mp4v codec unavailable?) -- "
"the fixture clip would be empty and preprocessing would fail on it.")
rng = np.random.default_rng(0)
for _ in range(121):
frame = rng.integers(0, 256, size=(512, 768, 3), dtype=np.uint8)
writer.write(frame)
writer.release()
# One physical clip, but max(2, NUM_GPUS) manifest rows all pointing at
# it: DP_SP_BatchSampler (fastvideo/dataset/parquet_dataset_map_style.py)
# floor-divides the number of batches by the number of data-parallel
# groups with drop_last=True, so a dataset with fewer rows than GPUs
# yields an EMPTY dataloader on every rank. Duplicate rows keep the
# single-sample-overfit semantics (identical latents/caption) while
# making KANDINSKY5_E2E_NUM_GPUS=2+ a working escape hatch when one GPU
# doesn't have the memory headroom for stage 2 (see _run_stage).
num_rows = max(2, int(NUM_GPUS))
with open(RAW_DATA_DIR / "videos2caption.json", "w") as f:
json.dump([{
"path": "sample_0.mp4",
"cap": [GENERATION_PROMPT],
}] * num_rows, f)
def _run_preprocessing() -> None:
# Point the preprocessor at the test-owned roots. Without these env
# overrides it would read/write its user-facing defaults
# (data/kandinsky5_overfit{,_preprocessed}) -- the very directories a
# user populates for the real recipe, which this test must never touch.
env = dict(os.environ)
env["KANDINSKY5_OVERFIT_DATA_DIR"] = str(RAW_DATA_DIR)
env["KANDINSKY5_OVERFIT_OUTPUT_DIR"] = str(PREPROCESSED_DIR)
cmd = [
sys.executable, "-m",
"fastvideo.pipelines.preprocess.preprocess_kandinsky5_overfit",
]
subprocess.run(cmd, cwd=str(REPO_ROOT), env=env, check=True)
def _export_dcp_to_diffusers(checkpoint_dir: Path, output_dir: Path) -> None:
"""Convert a stage's DCP checkpoint to a diffusers-style directory.
``models.*.init_from`` / ``VideoGenerator.from_pretrained`` (like real
training YAML configs, see ``dmd2_t2v_480p_qat.yaml``) expect a
diffusers model directory (``model_index.json`` + component
subfolders), not the raw ``checkpoint-N`` DCP layout
``CheckpointManager`` writes (``dcp/`` + metadata/RNG state only).
``--verify`` strictly reloads the exported transformer immediately so a
key-mapping bug fails here, at the export boundary, rather than deep
inside the next launch that loads the directory.
"""
cmd = [
sys.executable, "-m",
"fastvideo.train.entrypoint.dcp_to_diffusers",
"--checkpoint", str(checkpoint_dir),
"--output-dir", str(output_dir),
"--role", "student",
"--overwrite",
"--verify",
]
subprocess.run(cmd, cwd=str(REPO_ROOT), check=True)
def _latest_checkpoint(output_dir: Path) -> Path:
checkpoints = sorted(output_dir.glob("checkpoint-*"))
assert checkpoints, f"no checkpoint produced under {output_dir}"
return checkpoints[-1]
def _run_stage(config: Path, output_dir: Path, *, max_train_steps: int, extra_overrides: list[str],
env_overrides: dict[str, str]) -> None:
cmd = [
sys.executable, "-m", "torch.distributed.run",
"--nnodes", "1",
"--nproc_per_node", NUM_GPUS,
"-m", "fastvideo.train.entrypoint.train",
"--config", str(config),
"--training.data.data_path", str(PREPROCESSED_DIR),
"--training.distributed.num_gpus", NUM_GPUS,
"--training.distributed.hsdp_shard_dim", NUM_GPUS,
"--training.loop.max_train_steps", str(max_train_steps),
"--training.checkpoint.output_dir", str(output_dir),
"--training.checkpoint.training_state_checkpointing_steps", str(max_train_steps),
# Explicitly disable resume for BOTH stages, even though
# _clean_previous_artifacts() already removed the output dirs: the
# stage-1 YAML defaults to resume_from_checkpoint: latest, and this
# test's correctness must not silently depend on that default (or a
# partially-failed cleanup) ever changing.
"--training.checkpoint.resume_from_checkpoint", "none",
"--callbacks.validation.every_steps", str(max_train_steps),
"--callbacks.validation.dataset_file", str(PREPROCESSED_DIR / "validation_prompts.json"),
*extra_overrides,
]
env = dict(os.environ)
# Stage 2 holds three 2B-param models plus TWO full Adam states (the
# student's only exists because generator_update_interval is overridden
# to 1 above -- a recorded run peaked at ~76 GiB and died allocating
# the last MiBs on an 80 GiB device with ~1.9 GiB lost to
# fragmentation). Expandable segments reclaims that headroom; if a
# single device still can't fit, shard with KANDINSKY5_E2E_NUM_GPUS=2+
# (the synthesized manifest guarantees >= one row per rank).
env.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
env.update(env_overrides)
subprocess.run(cmd, cwd=str(REPO_ROOT), env=env, check=True)
def _generate_from_export(export_dir: Path, output_dir: Path) -> Path:
"""Deterministically generate one video from an exported student via
the documented DMD inference path, in a fresh subprocess.
A subprocess (re-running this file with ``--generate``, see
``__main__``) keeps VideoGenerator's distributed/CUDA init isolated
from the two torchrun stages that already ran from this test process.
FASTVIDEO_ATTENTION_BACKEND is removed from the child env: generation
must exercise the default dense-attention inference path the README
documents for non-sm_120 GPUs, not inherit a QAT training backend.
"""
env = dict(os.environ)
env.pop("FASTVIDEO_ATTENTION_BACKEND", None)
cmd = [
sys.executable,
str(THIS_FILE), "--generate",
str(export_dir),
str(output_dir),
]
subprocess.run(cmd, cwd=str(REPO_ROOT), env=env, check=True)
videos = sorted(Path(output_dir).glob("**/*.mp4"))
assert len(videos) == 1, (
f"expected exactly one generated video under {output_dir}, found {videos}")
return videos[0]
def _generate_main(export_dir: str, output_dir: str) -> None:
"""Subprocess body for _generate_from_export.
This is the exact inference recipe the QAD README documents for a
stage-2 export (minus the optional NVFP4/FP8 weight quantization,
which needs flashinfer + sm_120 hardware): the export's
model_index.json still names the base T2V pipeline, so
override_pipeline_cls_name + Kandinsky5DMDConfig are required to get
the 4-step re-noise sampler -- see
fastvideo/tests/api/test_kandinsky5_dmd_pipeline_resolution.py.
"""
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig
generator = VideoGenerator.from_pretrained(
export_dir,
num_gpus=1,
override_pipeline_cls_name="Kandinsky5DMDPipeline",
pipeline_config=Kandinsky5DMDConfig(),
use_fsdp_inference=False,
)
generator.generate_video(
GENERATION_PROMPT,
output_path=output_dir,
save_video=True,
seed=GENERATION_SEED,
height=512,
width=768,
num_frames=121,
)
def _decode_video(video_path: Path) -> np.ndarray:
cap = cv2.VideoCapture(str(video_path))
frames = []
while True:
ret, frame = cap.read()
if not ret:
break
frames.append(frame)
cap.release()
assert frames, f"video {video_path} decoded to zero frames"
return np.stack(frames)
def _assert_video_not_degenerate(video_path: Path) -> None:
"""Structural referee that needs no reference: rejects the collapsed
outputs (solid, frozen, truncated, unreadable) a NaN run or broken
export/decode path produces.
This is the primary defense against collapsed clips -- the MS-SSIM
floor cannot be (see the calibration table above: a solid clip at the
reference's mean color scores 0.9189 against it). Every check here is
designed against a measured false-green:
- exact frame count/geometry: compute_video_ssim_torchvision
truncates both clips to the shorter length, so a 1-frame video
could otherwise sail through the comparison;
- per-frame *luma* spatial std: the previous global RGB std counted
differences between channel means as "variance", scoring a solid
[106, 98, 70] clip at 15.4;
- consecutive-frame temporal diff: a frozen (single repeated frame)
clip has real spatial content and passes the spatial check.
"""
frames = _decode_video(video_path).astype(np.float64)
expected_shape = (EXPECTED_FRAME_COUNT, EXPECTED_FRAME_HEIGHT, EXPECTED_FRAME_WIDTH, 3)
assert frames.shape == expected_shape, (
f"generated video {video_path} has shape {frames.shape}, expected {expected_shape} -- "
"wrong frame count or geometry (and the SSIM helper would silently truncate to the "
"shorter clip)")
luma = frames.mean(axis=-1)
spatial_std = luma.reshape(len(luma), -1).std(axis=1)
assert float(spatial_std.min()) >= MIN_FRAME_SPATIAL_STD, (
f"generated video {video_path} has a frame with luma spatial std "
f"{spatial_std.min():.3f} < {MIN_FRAME_SPATIAL_STD} (reference min: ~11.4) -- "
"solid/near-solid output from a NaN or collapsed run")
temporal_diff = np.abs(np.diff(luma, axis=0)).mean(axis=(1, 2))
assert float(temporal_diff.min()) >= MIN_FRAME_TEMPORAL_DIFF, (
f"generated video {video_path} has consecutive frames with mean |luma diff| "
f"{temporal_diff.min():.4f} < {MIN_FRAME_TEMPORAL_DIFF} (reference min: ~1.1) -- "
"frozen output")
def _assert_matches_reference(video_path: Path) -> None:
if not REFERENCE_VIDEO.exists():
if os.environ.get(WRITE_REFERENCE_ENV) == "1":
shutil.copy2(video_path, REFERENCE_VIDEO)
print(f"\nWrote new reference video to {REFERENCE_VIDEO} -- "
"review it by eye and commit it. This bootstrap run "
"compares the video against itself (SSIM=1) below.")
else:
pytest.fail(
f"reference video missing at {REFERENCE_VIDEO}. This test "
"must not pass on artifact existence alone -- run once on a "
f"sanctioned GPU box with {WRITE_REFERENCE_ENV}=1, review "
"the written video, and commit it (see module docstring).")
from fastvideo.tests.utils import compute_video_ssim_torchvision
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
str(REFERENCE_VIDEO),
str(video_path),
use_ms_ssim=True,
)
print("\n===== MS-SSIM vs committed Kandinsky5 DMD reference =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert mean_ssim >= MIN_MEAN_MS_SSIM, (
f"mean MS-SSIM {mean_ssim:.4f} below {MIN_MEAN_MS_SSIM} -- the exported/reloaded "
"stage-2 student no longer generates what the reviewed reference recorded "
"(wrong weight remap, re-noise schedule, or attention path?)")
@pytest.mark.nightly
def test_e2e_kandinsky5_dmd_overfit_single_sample():
if not STAGE1_CONFIG.exists() or not STAGE2_CONFIG.exists():
pytest.skip("Kandinsky5 QAT configs not found -- see examples/train/configs/{fine_tuning,distribution_matching}/kandinsky5/")
os.environ.setdefault("WANDB_MODE", "offline")
_clean_previous_artifacts()
_synthesize_single_sample()
_run_preprocessing()
_run_stage(
STAGE1_CONFIG,
STAGE1_OUTPUT_DIR,
max_train_steps=3,
extra_overrides=[
# The recipe's full 5e-5 is calibrated for a real dataset; 3
# steps of it toward this test's random-noise clip measurably
# degrades the model -- empirically enough to collapse the
# fragile 4-step no-CFG DMD sampling below into grey static
# (50-step T2V sampling of the same export still recovered).
# A tiny LR keeps the full optimizer/backward path exercised
# while the final generation stays structured: reviewable by
# eye at bootstrap, and stable across runs for the SSIM oracle
# (noise output decorrelates under any training
# nondeterminism; structured output does not).
"--training.optimizer.learning_rate", "1e-6",
],
env_overrides={"FASTVIDEO_ATTENTION_BACKEND": "ATTN_QAT_TRAIN"},
)
stage1_ckpt = _latest_checkpoint(STAGE1_OUTPUT_DIR)
stage1_diffusers_dir = STAGE1_DIFFUSERS_DIR / stage1_ckpt.name
_export_dcp_to_diffusers(stage1_ckpt, stage1_diffusers_dir)
assert (stage1_diffusers_dir / "model_index.json").exists(), (
f"dcp_to_diffusers export did not produce a diffusers model dir at {stage1_diffusers_dir}")
_run_stage(
STAGE2_CONFIG,
STAGE2_OUTPUT_DIR,
max_train_steps=3,
extra_overrides=[
"--models.student.init_from", str(stage1_diffusers_dir),
"--models.teacher.init_from", str(stage1_diffusers_dir),
"--models.critic.init_from", str(stage1_diffusers_dir),
# The recipe's generator_update_interval of 5 would mean ZERO
# student updates in a 3-step run (the trainer iterates 1..3
# and DMD2Method updates the generator only when
# iteration % interval == 0) -- the exported "stage-2 student"
# would just be the stage-1 weights, and this test could not
# catch a broken student backward/optimizer path. Update the
# generator every step instead.
"--method.generator_update_interval", "1",
],
env_overrides={"FASTVIDEO_ATTENTION_BACKEND": "ATTN_QAT_TRAIN"},
)
assert any(STAGE2_OUTPUT_DIR.glob("*.mp4")), (
f"no stage-2 validation video produced under {STAGE2_OUTPUT_DIR}")
# The actual deliverable of this recipe is the exported stage-2
# student: export it, strict-reload it (--verify), then instantiate it
# through the documented Kandinsky5DMDPipeline override and generate.
stage2_ckpt = _latest_checkpoint(STAGE2_OUTPUT_DIR)
stage2_diffusers_dir = STAGE2_DIFFUSERS_DIR / stage2_ckpt.name
_export_dcp_to_diffusers(stage2_ckpt, stage2_diffusers_dir)
assert (stage2_diffusers_dir / "model_index.json").exists(), (
f"dcp_to_diffusers export did not produce a diffusers model dir at {stage2_diffusers_dir}")
generated_video = _generate_from_export(stage2_diffusers_dir, GENERATED_VIDEO_DIR)
_assert_video_not_degenerate(generated_video)
_assert_matches_reference(generated_video)
if __name__ == "__main__":
if len(sys.argv) >= 2 and sys.argv[1] == "--generate":
_generate_main(sys.argv[2], sys.argv[3])
else:
test_e2e_kandinsky5_dmd_overfit_single_sample()
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression tests for ``Kandinsky5DenoisingStage._assert_local_attention_backend_engaged``.
The stage guard must be strict for ``ATTN_QAT_TRAIN`` only:
``fastvideo.platforms.cuda`` raises ``ImportError`` at backend selection if
the training kernel isn't built (training must never silently
de-quantize), so an exact match is guaranteed whenever module construction
succeeded -- a mismatch there really does mean a silent SDPA fallback.
``ATTN_QAT_INFER`` is different: ``fastvideo.platforms.cuda`` deliberately
resolves it to FlashAttention when the sm_120 kernel isn't built, and the
QAD recipe README (examples/train/configs/fine_tuning/kandinsky5/README.md)
documents that fallback as the supported path on non-sm_120 GPUs ("the
attention falls back to a supported dense backend while the FP4 linear
layers still run"). Treating it as exact-match-only would let module
construction succeed with FlashAttention and then abort denoising on every
machine relying on the documented fallback.
Pure logic tests on stand-in modules -- no GPU, no optional kernels, no
model load, no distributed init needed.
"""
from __future__ import annotations
import pytest
import torch
from fastvideo.attention import LocalAttention
from fastvideo.pipelines.stages.kandinsky5 import Kandinsky5DenoisingStage
from fastvideo.platforms import AttentionBackendEnum
class _ResolvedLocalAttention(LocalAttention):
"""``LocalAttention`` stand-in with a pre-resolved backend.
Bypasses ``LocalAttention.__init__`` (which runs real backend selection
-- unavailable without the optional kernels these tests are about)
while still passing the stage guard's
``isinstance(module, LocalAttention)`` check.
"""
def __init__(self, backend: AttentionBackendEnum) -> None:
torch.nn.Module.__init__(self)
self.backend = backend
def _make_stage(resolved_backend: AttentionBackendEnum) -> Kandinsky5DenoisingStage:
"""Bypass __init__: only set the field the guard reads."""
stage = Kandinsky5DenoisingStage.__new__(Kandinsky5DenoisingStage)
transformer = torch.nn.Module()
transformer.attn = _ResolvedLocalAttention(resolved_backend)
stage.transformer = transformer
return stage
def test_attn_qat_infer_flash_fallback_is_allowed(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_INFER")
stage = _make_stage(AttentionBackendEnum.FLASH_ATTN)
# Must not raise: FlashAttention is the documented ATTN_QAT_INFER
# fallback on GPUs without the sm_120 kernel.
stage._assert_local_attention_backend_engaged()
def test_attn_qat_infer_exact_match_is_allowed(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_INFER")
stage = _make_stage(AttentionBackendEnum.ATTN_QAT_INFER)
stage._assert_local_attention_backend_engaged()
def test_attn_qat_train_mismatch_raises(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_TRAIN")
stage = _make_stage(AttentionBackendEnum.TORCH_SDPA)
with pytest.raises(AssertionError, match="ATTN_QAT_TRAIN"):
stage._assert_local_attention_backend_engaged()
def test_attn_qat_train_exact_match_is_allowed(monkeypatch):
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_TRAIN")
stage = _make_stage(AttentionBackendEnum.ATTN_QAT_TRAIN)
stage._assert_local_attention_backend_engaged()
@@ -0,0 +1,207 @@
# SPDX-License-Identifier: Apache-2.0
"""Stage-boundary regression: ``Kandinsky5DmdDenoisingStage.forward`` must
invoke the selected attention backend on every denoising step.
``Kandinsky5Attention.forward`` catches the ``AssertionError``
``LocalAttention`` raises when no pipeline forward context is set and
silently falls back to plain ``F.scaled_dot_product_attention`` -- so if
the ``set_forward_context`` wrapper inside the stage's denoising loop were
ever removed, generation would still "work" while skipping whichever
kernel ``FASTVIDEO_ATTENTION_BACKEND`` selected. The GPU test
``fastvideo/tests/train/models/test_kandinsky5_qat_attention_engages.py``
installs ``set_forward_context`` itself around a raw transformer call, so
it cannot catch that removal; this test drives the *production stage* and
spies on the resolved backend impl (``LocalAttention.attn_impl.forward``,
which is only reachable once ``get_forward_context()`` succeeds inside
``LocalAttention.forward``).
CPU-only: a tiny randomly-initialized Kandinsky5 transformer
(~127K params) driven through the real stage code. TORCH_SDPA is pinned
via env var so backend resolution is identical on CPU-only and CUDA
machines. Output *numerics* are deliberately not asserted -- with random
weights the 4-step predict-x0/re-noise feedback loop amplifies
magnitudes without bound; the contract under test is backend routing,
not sample quality (the nightly e2e owns that).
"""
from __future__ import annotations
import types
import pytest
import torch
from fastvideo.attention import LocalAttention
from fastvideo.attention.selector import _cached_get_attn_backend
from fastvideo.configs.models.dits.kandinsky5 import (
Kandinsky5ArchConfig,
Kandinsky5VideoConfig,
)
from fastvideo.forward_context import set_forward_context
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
DMD_STEPS = [1000, 750, 500, 250]
TEXT_SEQ_LEN = 6
def _tiny_arch() -> Kandinsky5ArchConfig:
# model_dim must be divisible by head_dim = sum(axes_dims) = 32.
return Kandinsky5ArchConfig(
in_visual_dim=4,
in_text_dim=32,
in_text_dim2=16,
time_dim=32,
out_visual_dim=4,
patch_size=(1, 2, 2),
model_dim=64,
ff_dim=128,
num_text_blocks=1,
num_visual_blocks=1,
axes_dims=(8, 12, 12),
visual_cond=False,
attention_type="regular",
)
def _build_stage_and_spy(monkeypatch):
"""Tiny transformer + real DMD stage on CPU, with every resolved
backend impl wrapped in a call counter."""
# Pin the backend so resolution is identical on CPU-only and CUDA
# machines, and clear the process-wide selector cache keyed on
# (head_size, dtype, supported_backends) -- NOT on the env var (same
# defensive pattern as test_kandinsky5_qat_attention_engages.py).
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
_cached_get_attn_backend.cache_clear()
from fastvideo.models.dits.kandinsky5 import Kandinsky5Transformer3DModel
import fastvideo.pipelines.stages.kandinsky5 as k5_stage_mod
torch.manual_seed(0)
arch = _tiny_arch()
transformer = Kandinsky5Transformer3DModel(Kandinsky5VideoConfig(arch_config=arch), hf_config={})
transformer.eval()
# get_local_torch_device() never returns "cpu" (cuda -> mps fallback);
# pin the stage to CPU so this test runs identically everywhere.
monkeypatch.setattr(k5_stage_mod, "get_local_torch_device", lambda: torch.device("cpu"))
stage = k5_stage_mod.Kandinsky5DmdDenoisingStage(transformer, FlowMatchEulerDiscreteScheduler(shift=5.0))
backend_calls: list[int] = []
for module in transformer.modules():
if isinstance(module, LocalAttention):
def _spy(*args, __orig=module.attn_impl.forward, **kwargs):
backend_calls.append(1)
return __orig(*args, **kwargs)
module.attn_impl.forward = _spy
return stage, transformer, arch, backend_calls
def _make_batch() -> types.SimpleNamespace:
# 5 frames / 64x64 with the real VAE ratios (4 temporal, 8 spatial)
# -> latents [1, 2, 8, 8, 4], patchified to a 2x4x4 grid.
return types.SimpleNamespace(
latents=torch.randn(1, 2, 8, 8, 4),
prompt_embeds=[torch.randn(1, TEXT_SEQ_LEN, 32), torch.randn(1, 16)],
prompt_attention_mask=[torch.ones(1, TEXT_SEQ_LEN, dtype=torch.long)],
height=64,
width=64,
num_frames=5,
image_latent=None,
generator=None,
)
def _make_fastvideo_args(arch: Kandinsky5ArchConfig) -> types.SimpleNamespace:
return types.SimpleNamespace(
model_loaded={"transformer": True},
disable_autocast=True,
pipeline_config=types.SimpleNamespace(
dit_precision="fp32",
dmd_denoising_steps=list(DMD_STEPS),
dit_config=types.SimpleNamespace(arch_config=arch),
vae_config=types.SimpleNamespace(arch_config=types.SimpleNamespace(
temporal_compression_ratio=4,
spatial_compression_ratio=8,
)),
),
)
def test_dmd_stage_invokes_selected_backend_every_step(monkeypatch):
stage, _, arch, backend_calls = _build_stage_and_spy(monkeypatch)
try:
batch = _make_batch()
result = stage.forward(batch, _make_fastvideo_args(arch))
# 3 LocalAttention modules (text self-attn, visual self-attn, visual
# cross-attn) x one forward per DMD step. Any silent
# missing-forward-context fallback to F.scaled_dot_product_attention
# bypasses attn_impl.forward entirely and shows up here as a shortfall.
expected = 3 * len(DMD_STEPS)
assert len(backend_calls) == expected, (
f"selected attention backend invoked {len(backend_calls)} times, expected {expected} -- "
"the stage's set_forward_context wrapper is no longer covering the transformer call, "
"so Kandinsky5Attention silently fell back to raw SDPA")
assert result.latents.shape == (1, 2, 8, 8, 4)
finally:
_cached_get_attn_backend.cache_clear()
def test_raw_transformer_without_context_bypasses_backend(monkeypatch):
"""Documents the trap the stage wrapper exists to prevent: without a
forward context the transformer still returns output, but the selected
backend is never invoked."""
stage, transformer, _, backend_calls = _build_stage_and_spy(monkeypatch)
del stage # only the spied transformer is needed here
try:
with torch.no_grad():
out = transformer(
hidden_states=torch.randn(1, 2, 8, 8, 4),
encoder_hidden_states=torch.randn(1, TEXT_SEQ_LEN, 32),
pooled_projections=torch.randn(1, 16),
timestep=torch.tensor([500.0]),
visual_rope_pos=[torch.arange(2), torch.arange(4), torch.arange(4)],
text_rope_pos=torch.arange(TEXT_SEQ_LEN),
scale_factor=(1.0, 2.0, 2.0),
sparse_params=None,
return_dict=True,
).sample
assert out.shape == (1, 2, 8, 8, 4)
assert len(backend_calls) == 0, (
"expected the missing-forward-context fallback to bypass the backend impl; "
"if this now fails, Kandinsky5Attention's fallback behavior changed and the "
"stage guard docs/tests should be revisited")
finally:
_cached_get_attn_backend.cache_clear()
def test_dmd_stage_with_context_and_raw_without_share_one_spy(monkeypatch):
"""Same spy, both paths, one process: proves the counter difference
between the two tests above is the forward-context wrapper itself,
not some environmental difference."""
stage, transformer, arch, backend_calls = _build_stage_and_spy(monkeypatch)
try:
with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None):
transformer(
hidden_states=torch.randn(1, 2, 8, 8, 4),
encoder_hidden_states=torch.randn(1, TEXT_SEQ_LEN, 32),
pooled_projections=torch.randn(1, 16),
timestep=torch.tensor([500.0]),
visual_rope_pos=[torch.arange(2), torch.arange(4), torch.arange(4)],
text_rope_pos=torch.arange(TEXT_SEQ_LEN),
scale_factor=(1.0, 2.0, 2.0),
sparse_params=None,
return_dict=True,
)
assert len(backend_calls) == 3, "context-wrapped raw call should hit all 3 attention layers"
stage.forward(_make_batch(), _make_fastvideo_args(arch))
assert len(backend_calls) == 3 + 3 * len(DMD_STEPS)
finally:
_cached_get_attn_backend.cache_clear()
@@ -0,0 +1,139 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression tests for ``Kandinsky5DenoisingStage._resolve_target_dtype``.
Plain (non-FSDP) transformer: the stage must honor
``pipeline_config.dit_precision`` exactly -- ``TransformerLoader.load()``
asserts every parameter matches it for a non-FSDP load, so an explicit
fp32 pipeline must not be silently cast to bf16 (an earlier
parameter-scanning version of this helper did exactly that: all-fp32
params fell through to a bf16 "safe default").
FSDP2-wrapped transformer: the stage must read the active
``MixedPrecisionPolicy`` recorded by ``set_mixed_precision_policy`` (which
``maybe_load_fsdp_model`` always calls before sharding -- hardcoded to
``param_dtype=torch.bfloat16``, ``cast_forward_inputs=False`` today), NOT
the parameter storage dtypes. FSDP2 computes in the policy's
``param_dtype`` regardless of the dtype parameters are stored in, so an
fp16-loaded FSDP model still runs its forward in bf16; a parameter scan
resolved fp16 there and, with autocast keyed off the wrong dtype and
input casting disabled, handed bf16-computing modules fp16 inputs.
These are pure logic tests -- the FSDP "wrap" only swaps ``__class__`` the
same way ``fully_shard`` does (so the ``isinstance(..., FSDPModule)``
check fires) without any process group, GPU, or model load.
"""
from __future__ import annotations
import types
import pytest
import torch
from fastvideo import utils as fastvideo_utils
from fastvideo.pipelines.stages.kandinsky5 import Kandinsky5DenoisingStage
from fastvideo.utils import set_mixed_precision_policy
def _make_stage(transformer: torch.nn.Module) -> Kandinsky5DenoisingStage:
"""Bypass __init__: only set the field _resolve_target_dtype reads."""
stage = Kandinsky5DenoisingStage.__new__(Kandinsky5DenoisingStage)
stage.transformer = transformer
return stage
def _fastvideo_args(dit_precision: str) -> types.SimpleNamespace:
return types.SimpleNamespace(pipeline_config=types.SimpleNamespace(dit_precision=dit_precision))
def _fsdp_wrap(module: torch.nn.Module) -> torch.nn.Module:
"""Make ``isinstance(module, FSDPModule)`` true the way ``fully_shard``
does -- by swapping ``__class__`` to a dynamic ``(FSDPModule, orig)``
subclass -- without any actual sharding/distributed setup."""
from torch.distributed.fsdp import FSDPModule
orig_cls = module.__class__
module.__class__ = type(f"FSDP{orig_cls.__name__}", (FSDPModule, orig_cls), {})
return module
@pytest.fixture()
def mixed_precision_state_reset():
"""Snapshot/restore the process's thread-local mixed-precision state so
these tests neither depend on nor leak state across the pytest run."""
state_holder = fastvideo_utils._mixed_precision_state
had_state = hasattr(state_holder, "state")
prev_state = getattr(state_holder, "state", None)
if had_state:
del state_holder.state
yield
if had_state:
state_holder.state = prev_state
elif hasattr(state_holder, "state"):
del state_holder.state
def test_resolve_target_dtype_honors_explicit_fp32_for_plain_transformer():
transformer = torch.nn.Linear(4, 4).to(torch.float32)
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("fp32"))
assert resolved == torch.float32, (
"an explicit fp32 pipeline_config must not be silently cast to bf16 for a "
"plain (non-FSDP) transformer")
def test_resolve_target_dtype_honors_explicit_fp16_for_plain_transformer():
transformer = torch.nn.Linear(4, 4).to(torch.float16)
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("fp16"))
assert resolved == torch.float16
def test_resolve_target_dtype_honors_explicit_bf16_for_plain_transformer():
transformer = torch.nn.Linear(4, 4).to(torch.bfloat16)
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("bf16"))
assert resolved == torch.bfloat16
def test_resolve_target_dtype_fsdp_reads_policy_not_parameter_storage(mixed_precision_state_reset):
"""An fp16-*stored* FSDP model still computes in the policy's bf16 --
the exact mismatch a parameter scan gets wrong."""
set_mixed_precision_policy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32)
transformer = _fsdp_wrap(torch.nn.Linear(4, 4).to(torch.float16))
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("fp16"))
assert resolved == torch.bfloat16, (
"FSDP compute dtype comes from the MixedPrecisionPolicy param_dtype, not from "
"the dtype the parameters happen to be stored in")
def test_resolve_target_dtype_fsdp_follows_a_non_default_policy(mixed_precision_state_reset):
"""Proves the policy is actually read (not a hardcoded bf16)."""
set_mixed_precision_policy(param_dtype=torch.float16, reduce_dtype=torch.float32)
transformer = _fsdp_wrap(torch.nn.Linear(4, 4).to(torch.float32))
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("fp32"))
assert resolved == torch.float16
def test_resolve_target_dtype_fsdp_defaults_to_bf16_without_policy_state(mixed_precision_state_reset):
"""maybe_load_fsdp_model always records the policy before sharding, but
if the thread-local state is somehow unset the stage must still match
the loader's hardcoded bf16 param_dtype rather than crash or resolve
the storage dtype."""
transformer = _fsdp_wrap(torch.nn.Linear(4, 4).to(torch.float16))
stage = _make_stage(transformer)
resolved = stage._resolve_target_dtype(_fastvideo_args("fp16"))
assert resolved == torch.bfloat16
@@ -0,0 +1,17 @@
# Minimum config to instantiate Kandinsky5Model for the loading + forward
# smoke test. Real Kandinsky-5.0 Lite checkpoint loaded via
# load_module_from_path.
models:
student:
_target_: fastvideo.train.models.kandinsky5.Kandinsky5Model
init_from: kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers
trainable: false
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
dit_precision: bf16
pipeline: {}
@@ -0,0 +1,162 @@
# SPDX-License-Identifier: Apache-2.0
"""Forward+backward smoke test confirming ATTN_QAT_TRAIN actually engages
for Kandinsky5's dense/local attention path, rather than silently falling
back to plain SDPA.
Two silent-fallback traps this guards against:
- ``Kandinsky5Attention.forward`` catches ``AssertionError`` from
``LocalAttention`` when no forward context is set and falls back to
``F.scaled_dot_product_attention`` (a real fallback for standalone
parity tests, not relevant here since we always set a forward context,
but worth asserting past explicitly).
- ``_cached_get_attn_backend`` is process-wide ``@cache``d on
``(head_size, dtype, supported_attention_backends)`` -- NOT on the env
var -- so a stale backend selection from an earlier test in the same
process can silently linger. Cleared here before selecting, matching
the defensive pattern in
``fastvideo/models/loader/component_loader.py``.
This test itself must not become that stale-selection source for whichever
model test runs next in the same pytest process: ``FASTVIDEO_ATTENTION_BACKEND``
is set via ``monkeypatch`` (auto-restored at teardown, including on
failure) and the cache is cleared again in a ``finally`` block, since
``monkeypatch`` only restores the env var, not the separate process-wide
cache keyed on it.
Unlike ``fastvideo.platforms.cuda``'s other backends, ATTN_QAT_TRAIN has no
silent-fallback path at all if the kernel isn't built -- backend selection
raises ``ImportError`` loudly. This test skips (does not fail) when the
kernel isn't available, since it's an optional build artifact from
fastvideo-kernel/.
Runs the *full* transformer forward (mirroring test_load_kandinsky5.py)
rather than invoking an attention submodule directly. Two earlier versions
of this test tried isolating a single Kandinsky5Attention module -- first a
freshly-random-initialized one (the fake-quant kernel produced NaN/Inf on
untrained weights), then a real submodule pulled out of the loaded model
(FSDP2's fully_shard() registers its unshard/reshard hooks on the top-level
transformer's __call__, not on submodules individually, so invoking a
submodule directly left its Linear layers' parameters in sharded DTensor
form -- incompatible with a plain-Tensor input regardless of the
`trainable` flag). Going through the properly-hooked top-level forward call
sidesteps both problems.
"""
from __future__ import annotations
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29519")
from pathlib import Path
import pytest
import torch
from fastvideo.attention import LocalAttention
from fastvideo.attention.selector import _cached_get_attn_backend
from fastvideo.forward_context import set_forward_context
from fastvideo.platforms import AttentionBackendEnum
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures" /
"kandinsky5_t2v_min.yaml")
@pytest.mark.usefixtures("distributed_setup")
def test_kandinsky5_attn_qat_train_engages_and_backprops(monkeypatch):
if not torch.cuda.is_available():
pytest.skip("requires CUDA")
from fastvideo.attention.backends.attn_qat_train import (
is_attn_qat_train_available, )
if not is_attn_qat_train_available():
pytest.skip("fastvideo_kernel ATTN_QAT_TRAIN kernel is not built")
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "ATTN_QAT_TRAIN")
_cached_get_attn_backend.cache_clear()
try:
from fastvideo.train.models.kandinsky5 import Kandinsky5Model
from fastvideo.train.utils.config import load_run_config
cfg = load_run_config(_FIXTURE)
model = Kandinsky5Model(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
trainable=True,
)
device = torch.device("cuda:0")
dtype = torch.bfloat16
transformer = model.transformer.to(device=device, dtype=dtype)
attn = transformer.visual_transformer_blocks[0].self_attention
assert isinstance(attn.local_attention, LocalAttention)
assert attn.local_attention.backend == AttentionBackendEnum.ATTN_QAT_TRAIN, (
f"expected ATTN_QAT_TRAIN, got {attn.local_attention.backend} -- "
"backend selection silently fell back")
arch = transformer.config.arch_config
patch_size = arch.patch_size
in_visual_dim = arch.in_visual_dim
in_text_dim = arch.in_text_dim
in_text_dim2 = arch.in_text_dim2
grid_t, grid_h, grid_w = 2, 4, 4
latent_t = grid_t * patch_size[0]
latent_h = grid_h * patch_size[1]
latent_w = grid_w * patch_size[2]
latents = torch.randn(
1, latent_t, latent_h, latent_w, in_visual_dim, device=device, dtype=dtype,
requires_grad=True)
if bool(getattr(transformer, "visual_cond", False)):
# See Kandinsky5Model._build_distill_input_kwargs /
# Kandinsky5LatentPreparationStage: visual_cond=True checkpoints
# always expect [real | zero_cond | zero_mask] concatenated on the
# channel dim.
cond = torch.zeros_like(latents)
mask = torch.zeros(*latents.shape[:-1], 1, device=device, dtype=dtype)
hidden_states = torch.cat([latents, cond, mask], dim=-1)
else:
hidden_states = latents
encoder_hidden_states = torch.randn(1, 8, in_text_dim, device=device, dtype=dtype)
pooled_projections = torch.randn(1, in_text_dim2, device=device, dtype=dtype)
timestep = torch.tensor([500], device=device, dtype=dtype)
visual_rope_pos = [
torch.arange(grid_t, device=device),
torch.arange(grid_h, device=device),
torch.arange(grid_w, device=device),
]
text_rope_pos = torch.arange(encoder_hidden_states.shape[1], device=device)
with set_forward_context(current_timestep=0, attn_metadata=None):
out = transformer(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
pooled_projections=pooled_projections,
timestep=timestep,
visual_rope_pos=visual_rope_pos,
text_rope_pos=text_rope_pos,
scale_factor=(1.0, 2.0, 2.0),
sparse_params=None,
return_dict=True,
).sample
assert torch.isfinite(out).all().item(), "output contains NaN/Inf"
out.sum().backward()
assert latents.grad is not None
assert torch.isfinite(latents.grad).all().item(), "input grad contains NaN/Inf"
assert attn.to_query.weight.grad is not None
assert torch.isfinite(attn.to_query.weight.grad.to_local()
if hasattr(attn.to_query.weight.grad, "to_local") else
attn.to_query.weight.grad).all().item(), "weight grad contains NaN/Inf"
finally:
# monkeypatch restores FASTVIDEO_ATTENTION_BACKEND itself, but the
# selector cache is a separate process-wide functools.cache keyed on
# (head_size, dtype, supported_backends) -- not on the env var -- so
# it must be cleared again here or a later model test in this same
# pytest process could reuse this test's ATTN_QAT_TRAIN selection.
_cached_get_attn_backend.cache_clear()
@@ -0,0 +1,112 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU loading + forward smoke test for ``Kandinsky5Model``.
Loads the real Kandinsky-5.0 Lite checkpoint via ``Kandinsky5Model.__init__``
and runs one transformer forward pass on synthetic inputs. Kandinsky5's
transformer takes a structurally different forward signature than Wan/Hunyuan
-- dual text conditioning (``encoder_hidden_states`` + ``pooled_projections``),
explicit RoPE position tensors, a ``scale_factor``, and channel-last
``[B, T, H, W, C]`` hidden_states with a dict return (``.sample``) -- this
test mirrors the kwargs in ``Kandinsky5Model._build_distill_input_kwargs``
and ``Kandinsky5DenoisingStage``.
"""
from __future__ import annotations
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29518")
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
from pathlib import Path
import pytest
import torch
from fastvideo.forward_context import set_forward_context
from fastvideo.train.models.kandinsky5 import Kandinsky5Model
from fastvideo.train.utils.config import load_run_config
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures" /
"kandinsky5_t2v_min.yaml")
@pytest.mark.usefixtures("distributed_setup")
def test_kandinsky5_model_loads_and_forwards():
if not torch.cuda.is_available():
pytest.skip("requires CUDA")
cfg = load_run_config(_FIXTURE)
model = Kandinsky5Model(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
trainable=False,
)
transformer = model.transformer
assert isinstance(transformer, torch.nn.Module)
assert sum(p.numel() for p in transformer.parameters()) > 0
device = torch.device("cuda:0")
dtype = torch.bfloat16
transformer = transformer.to(device=device, dtype=dtype).eval()
arch = transformer.config.arch_config
patch_size = arch.patch_size
in_visual_dim = arch.in_visual_dim
out_visual_dim = arch.out_visual_dim
in_text_dim = arch.in_text_dim
in_text_dim2 = arch.in_text_dim2
# Small patch grid (T, H, W) so this fits next to the real checkpoint on
# a single GPU. scale_factor matches the 480p band Kandinsky5Model
# asserts on -- see fastvideo/train/models/kandinsky5/kandinsky5.py.
grid_t, grid_h, grid_w = 2, 4, 4
latent_t = grid_t * patch_size[0]
latent_h = grid_h * patch_size[1]
latent_w = grid_w * patch_size[2]
latents = torch.randn(
1, latent_t, latent_h, latent_w, in_visual_dim, device=device, dtype=dtype)
if bool(getattr(transformer, "visual_cond", False)):
# The shipped checkpoint has visual_cond=True (unified T2V/I2V):
# Kandinsky5VisualEmbeddings expects [real | zero_cond | zero_mask]
# concatenated on the channel dim, even for pure T2V. Mirrors
# Kandinsky5Model._build_distill_input_kwargs / Kandinsky5LatentPreparationStage.
cond = torch.zeros_like(latents)
mask = torch.zeros(*latents.shape[:-1], 1, device=device, dtype=dtype)
hidden_states = torch.cat([latents, cond, mask], dim=-1)
else:
hidden_states = latents
encoder_hidden_states = torch.randn(1, 8, in_text_dim, device=device, dtype=dtype)
pooled_projections = torch.randn(1, in_text_dim2, device=device, dtype=dtype)
timestep = torch.tensor([500], device=device, dtype=dtype)
visual_rope_pos = [
torch.arange(grid_t, device=device),
torch.arange(grid_h, device=device),
torch.arange(grid_w, device=device),
]
text_rope_pos = torch.arange(encoder_hidden_states.shape[1], device=device)
with torch.no_grad(), set_forward_context(
current_timestep=0,
attn_metadata=None,
):
out = transformer(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
pooled_projections=pooled_projections,
timestep=timestep,
visual_rope_pos=visual_rope_pos,
text_rope_pos=text_rope_pos,
scale_factor=(1.0, 2.0, 2.0),
sparse_params=None,
return_dict=True,
).sample
expected_shape = (1, latent_t, latent_h, latent_w, out_visual_dim)
assert tuple(out.shape) == expected_shape, (
f"output shape {tuple(out.shape)} != expected {expected_shape}")
assert torch.isfinite(out).all().item(), "output contains NaN/Inf"
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
"""Known-bad calibration tests for the Kandinsky5 e2e media oracle.
A review pass demonstrated a concrete false-green: a 121-frame solid clip
at the committed reference's mean RGB passed the old degeneracy check
(global RGB std counts channel-mean differences as "variance": 15.4) AND
the old MS-SSIM floor (scoring a mean of ~0.92 against the low-contrast
reference, with the SSIM helper additionally truncating both clips to the
shorter frame count). These tests turn that exploit -- and its frozen and
truncated siblings -- into permanent regressions against the *currently
committed* reference video, so the oracle's thresholds can never silently
drift below a known-bad clip again: if a future reference changes the
similarity landscape, the floor assertions here fail and force
recalibration.
CPU-only; skips when the committed reference is absent (pre-bootstrap
checkouts).
"""
from __future__ import annotations
import cv2
import numpy as np
import pytest
from fastvideo.tests.nightly.test_e2e_kandinsky5_dmd_t2v_overfit import (
EXPECTED_FRAME_COUNT,
MIN_MEAN_MS_SSIM,
REFERENCE_VIDEO,
_assert_video_not_degenerate,
_decode_video,
)
pytestmark = pytest.mark.skipif(
not REFERENCE_VIDEO.exists(),
reason="committed Kandinsky5 e2e reference video not present",
)
def _write_clip(path, frames) -> None:
height, width = frames[0].shape[:2]
writer = cv2.VideoWriter(str(path), cv2.VideoWriter_fourcc(*"mp4v"), 24.0, (width, height))
assert writer.isOpened(), "mp4v VideoWriter unavailable in this OpenCV build"
for frame in frames:
writer.write(frame)
writer.release()
def _mean_ms_ssim_vs_reference(path) -> float:
pytest.importorskip("pytorch_msssim")
pytest.importorskip("av")
from fastvideo.tests.utils import compute_video_ssim_torchvision
mean_ssim, _, _ = compute_video_ssim_torchvision(str(REFERENCE_VIDEO), str(path), use_ms_ssim=True)
return float(mean_ssim)
def test_committed_reference_passes_the_oracle():
"""Self-consistency: thresholds must never be tightened past what the
reviewed reference itself exhibits."""
_assert_video_not_degenerate(REFERENCE_VIDEO)
def test_solid_clip_at_reference_mean_color_is_rejected(tmp_path):
"""The review's exact exploit: solid clip at the reference's own mean
RGB. Global-std passed it at 15.4; per-frame luma spatial std is 0."""
reference = _decode_video(REFERENCE_VIDEO)
mean_color = reference.astype(np.float64).mean(axis=(0, 1, 2)).round().astype(np.uint8)
solid_frame = np.full(reference.shape[1:], mean_color, dtype=np.uint8)
solid_path = tmp_path / "solid.mp4"
_write_clip(solid_path, [solid_frame] * EXPECTED_FRAME_COUNT)
with pytest.raises(AssertionError, match="solid"):
_assert_video_not_degenerate(solid_path)
# The similarity floor must ALSO sit above this clip's score (measured
# ~0.9189 against the 2026-07-18 reference): the structural check is
# the primary defense, but the floor must not regress into known-bad
# territory either.
assert _mean_ms_ssim_vs_reference(solid_path) < MIN_MEAN_MS_SSIM
def test_frozen_clip_is_rejected(tmp_path):
"""A single real reference frame repeated 121x has genuine spatial
content (luma std ~12) -- only the temporal check catches it."""
reference = _decode_video(REFERENCE_VIDEO)
frozen_path = tmp_path / "frozen.mp4"
_write_clip(frozen_path, [reference[0]] * EXPECTED_FRAME_COUNT)
with pytest.raises(AssertionError, match="frozen"):
_assert_video_not_degenerate(frozen_path)
# Measured ~0.8597 against the 2026-07-18 reference.
assert _mean_ms_ssim_vs_reference(frozen_path) < MIN_MEAN_MS_SSIM
def test_truncated_clip_is_rejected(tmp_path):
"""compute_video_ssim_torchvision truncates both clips to the shorter
frame count, so a partial video could score arbitrarily well -- the
oracle must reject on frame count before similarity is even
consulted."""
reference = _decode_video(REFERENCE_VIDEO)
truncated_path = tmp_path / "truncated.mp4"
_write_clip(truncated_path, list(reference[:EXPECTED_FRAME_COUNT // 2]))
with pytest.raises(AssertionError, match="frame count or geometry"):
_assert_video_not_degenerate(truncated_path)
@@ -0,0 +1,45 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression test: preprocess_kandinsky5_overfit.py must not silently
truncate a string caption to its first character.
The documented ``videos2caption.json`` schema
(``docs/training/data_preprocess.md``) stores ``cap`` as a plain string
(e.g. ``"Ocean waves at sunset..."``), while some producers (this repo's
own e2e test fixtures, mirroring ``preprocess_hunyuan_overfit.py``) store a
non-empty list of caption variants instead. Indexing unconditionally with
``item["cap"][0]`` silently takes the first *character* of a string caption
("O") instead of erroring -- ``get_caption`` accepts and validates both
forms once at ingestion.
Pure logic test -- no GPU, no model load.
"""
from __future__ import annotations
import pytest
from fastvideo.pipelines.preprocess.preprocess_kandinsky5_overfit import get_caption
def test_get_caption_accepts_plain_string():
assert get_caption({"path": "a.mp4", "cap": "Ocean waves at sunset..."}) == "Ocean waves at sunset..."
def test_get_caption_accepts_nonempty_list():
assert get_caption({"path": "a.mp4", "cap": ["a synthetic test clip"]}) == "a synthetic test clip"
def test_get_caption_uses_first_list_element():
assert get_caption({"path": "a.mp4", "cap": ["first", "second"]}) == "first"
@pytest.mark.parametrize("bad_cap", [
"",
[],
[""],
None,
123,
{"nested": "dict"},
])
def test_get_caption_rejects_invalid_values(bad_cap):
with pytest.raises(ValueError):
get_caption({"path": "a.mp4", "cap": bad_cap})
@@ -0,0 +1,49 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression test: preprocess_kandinsky5_overfit.load_video must fail with
a path-specific error on an unreadable input, not an opaque IndexError.
``cv2.VideoCapture`` doesn't raise on a missing or corrupt file -- it just
decodes zero frames, and the repeat-last-frame fill then crashed on
``frames[-1]`` with an ``IndexError`` that named neither the file nor the
cause.
CPU-only -- decodes a tiny real clip written with cv2, no GPU or model
load.
"""
from __future__ import annotations
import cv2
import numpy as np
import pytest
from fastvideo.pipelines.preprocess.preprocess_kandinsky5_overfit import load_video
def test_load_video_missing_file_raises_path_specific_error(tmp_path):
missing = tmp_path / "does_not_exist.mp4"
with pytest.raises(ValueError, match="does_not_exist.mp4"):
load_video(str(missing), num_frames=5, height=64, width=64)
def test_load_video_corrupt_file_raises_path_specific_error(tmp_path):
corrupt = tmp_path / "corrupt.mp4"
corrupt.write_bytes(b"this is not an mp4 container")
with pytest.raises(ValueError, match="corrupt.mp4"):
load_video(str(corrupt), num_frames=5, height=64, width=64)
def test_load_video_short_clip_still_fills_by_repeating_last_frame(tmp_path):
"""The zero-frame guard must not break the documented short-clip fill."""
clip = tmp_path / "short.mp4"
writer = cv2.VideoWriter(str(clip), cv2.VideoWriter_fourcc(*"mp4v"), 24.0, (64, 64))
assert writer.isOpened(), "mp4v VideoWriter unavailable in this OpenCV build"
rng = np.random.default_rng(0)
for _ in range(3):
writer.write(rng.integers(0, 256, size=(64, 64, 3), dtype=np.uint8))
writer.release()
video = load_video(str(clip), num_frames=8, height=64, width=64)
assert video.shape == (1, 3, 8, 64, 64)
+9
View File
@@ -1216,6 +1216,15 @@ class ValidationCallback(Callback):
VSA_sparsity=tc.vsa_sparsity,
timesteps=sampling_timesteps_tensor,
)
# shallow_asdict(sampling_param) copies list-typed fields by
# reference. sampling_param is cached and reused across every
# validation sample/step, so without this reset every ForwardBatch
# would share (and keep appending to) the same
# prompt_attention_mask/negative_attention_mask list forever --
# index [0] would then hold the *first-ever* validation sample's
# mask instead of the current one, mismatching prompt_embeds.
batch.prompt_attention_mask = []
batch.negative_attention_mask = []
batch._inference_args = inference_args # type: ignore[attr-defined]
# Conditionally set I2V fields.
+46 -3
View File
@@ -20,6 +20,10 @@ The checkpoint must contain ``metadata.json`` (written by
``CheckpointManager``). If the checkpoint predates metadata
support, pass ``--config`` explicitly to provide the training
YAML.
Pass ``--verify`` to strictly reload the exported transformer immediately
after writing it, so a key-mapping bug fails here instead of deep inside a
later training/inference launch that loads the exported directory.
"""
from __future__ import annotations
@@ -101,10 +105,15 @@ def _save_role_pretrained(
"Pass --overwrite to replace it.")
def _copy_or_link(src: str, dest: str) -> None:
# Resolve symlinks ourselves: os.link's follow_symlinks=True
# default isn't honored on all filesystems (e.g. some
# network/overlay mounts), which can silently hard-link to
# the symlink itself instead of its target.
real_src = os.path.realpath(src)
try:
os.link(src, dest)
os.link(real_src, dest)
except OSError:
shutil.copy2(src, dest)
shutil.copy2(real_src, dest)
logger.info(
"Creating pretrained export dir at %s "
@@ -115,7 +124,7 @@ def _save_role_pretrained(
shutil.copytree(
local_base,
dst,
symlinks=True,
symlinks=False,
copy_function=_copy_or_link,
)
@@ -195,6 +204,27 @@ def _save_role_pretrained(
return str(dst)
def _strict_reload_verify(*, output_dir: str, training_config: Any) -> None:
"""Reload the just-exported transformer from disk and fail loudly on
any key mismatch.
``TransformerLoader.load()`` already loads strictly (``strict=True``)
for every non-Cosmos25 model, so this doesn't add leniency -- it moves
the failure to right after export, with a clear "the export is broken"
error, instead of surfacing deep inside a later training/inference
launch that happens to load this directory.
"""
from fastvideo.train.utils.moduleloader import load_module_from_path
logger.info("Verifying export: strictly reloading transformer from %s", output_dir)
load_module_from_path(
model_path=output_dir,
module_type="transformer",
training_config=training_config,
)
logger.info("Strict reload verification passed.")
def convert(
*,
checkpoint_dir: str,
@@ -202,6 +232,7 @@ def convert(
config_path: str | None = None,
role: str = "student",
overwrite: bool = False,
verify: bool = False,
) -> str:
"""Load a DCP checkpoint and export as a diffusers model.
@@ -293,6 +324,10 @@ def convert(
model=model,
)
logger.info("Export complete: %s", result)
if verify:
_strict_reload_verify(output_dir=result, training_config=tc)
return result
@@ -401,6 +436,13 @@ def main() -> None:
action="store_true",
help="Overwrite output-dir if it exists.",
)
parser.add_argument(
"--verify",
action="store_true",
help=("After exporting, strictly reload the transformer from "
"the exported directory to catch key-mapping bugs "
"immediately."),
)
args = parser.parse_args(sys.argv[1:])
convert(
@@ -409,6 +451,7 @@ def main() -> None:
config_path=args.config,
role=args.role,
overwrite=args.overwrite,
verify=args.verify,
)
@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
"""Kandinsky5 model plugin package."""
from fastvideo.train.models.kandinsky5.kandinsky5 import (
Kandinsky5Model as Kandinsky5Model, )
@@ -0,0 +1,773 @@
# SPDX-License-Identifier: Apache-2.0
"""Kandinsky5 model plugin (per-role instance).
Written directly against ``ModelBase`` (not a ``WanModel`` subclass) because
Kandinsky5 differs structurally in ways that don't compose cleanly with
Wan's implementation:
- Dual text encoders: Qwen/Reason1 sequence embeddings *and* a CLIP pooled
projection, vs. Wan's single encoder. The pair is packed into the
existing single ``text_embedding`` parquet field by zero-padding the
CLIP pooled vector into a row and prepending it to the Qwen sequence
(the same trick already used by ``HunyuanModel``/
``preprocess_hunyuan_overfit.py`` for LLaMA+CLIP).
- The transformer forward signature needs RoPE position tensors
(``visual_rope_pos``, ``text_rope_pos``) and a ``scale_factor`` that Wan
never has.
- Kandinsky5's native hidden_states layout is channel-last
``[B, T, H, W, C]``; the common ModelBase convention used at the
predict_noise/decode_latents boundary is ``[B, T, C, H, W]``.
- VAE denorm uses only ``scaling_factor`` (Kandinsky5 shares Hunyuan's
VAE, which has no ``latents_mean``/``latents_std``/
``handles_latent_denorm``).
- flow_shift default is 5.0 (vs. Wan's 3.0).
Scope: T2V at 480p only, dense/local attention only (NABLA sparse attention
is never engaged -- ``sparse_params`` is always ``None``).
"""
from __future__ import annotations
import copy
import os
from typing import Any, Literal, TYPE_CHECKING
import torch
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.distributed import (
get_sp_group,
get_world_group,
)
from fastvideo.forward_context import set_forward_context
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing, )
from fastvideo.training.training_utils import (
compute_density_for_timestep_sampling,
get_sigmas,
normalize_dit_input,
shift_timestep,
)
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.module_state import (
apply_trainable, )
from fastvideo.train.utils.moduleloader import (
load_module_from_path,
make_inference_args,
)
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
from fastvideo.train.utils.lora import LoraConfig
# 480p is the only supported resolution band for this recipe. Matches
# Kandinsky5DenoisingStage._scale_factor's low-resolution branch exactly;
# outside this band Kandinsky5 uses a different visual RoPE scale_factor
# ((1.0, 3.16, 3.16)) that this wrapper does not implement.
_SCALE_FACTOR_480P = (1.0, 2.0, 2.0)
_MIN_480P_SIDE = 480
_MAX_480P_SIDE = 854
class Kandinsky5Model(ModelBase):
"""Kandinsky5 per-role model: owns transformer + noise_scheduler."""
_transformer_cls_name: str = "Kandinsky5Transformer3DModel"
def __init__(
self,
*,
init_from: str,
training_config: TrainingConfig,
trainable: bool = True,
disable_custom_init_weights: bool = False,
flow_shift: float = 5.0,
enable_gradient_checkpointing_type: str
| None = None,
transformer_override_safetensor: str
| None = None,
lora: LoraConfig | dict[str, Any] | None = None,
) -> None:
super().__init__(
trainable=trainable,
lora=lora,
)
self._init_from = str(init_from)
self.transformer = self._load_transformer(
init_from=self._init_from,
trainable=self._trainable,
disable_custom_init_weights=(disable_custom_init_weights),
enable_gradient_checkpointing_type=(enable_gradient_checkpointing_type),
training_config=training_config,
transformer_override_safetensor=(transformer_override_safetensor),
)
self.noise_scheduler = (FlowMatchEulerDiscreteScheduler(shift=float(flow_shift)))
# Filled by init_preprocessors (student only).
self.vae: Any = None
self.training_config: TrainingConfig = training_config
self.dataloader: Any = None
self.validator: Any = None
self.start_step: int = 0
self.world_group: Any = None
self.sp_group: Any = None
# Qwen sequence embeds + mask, and CLIP pooled projection.
self.negative_prompt_embeds: (torch.Tensor | None) = None
self.negative_prompt_attention_mask: (torch.Tensor | None) = None
self.negative_pooled_embeds: (torch.Tensor | None) = None
self._requires_negative_conditioning = True
# Timestep mechanics.
self.timestep_shift: float = float(flow_shift)
self.num_train_timestep: int = int(self.noise_scheduler.num_train_timesteps)
self.min_timestep: int = 0
self.max_timestep: int = self.num_train_timestep
def _load_transformer(
self,
*,
init_from: str,
trainable: bool,
disable_custom_init_weights: bool,
enable_gradient_checkpointing_type: str | None,
training_config: TrainingConfig,
transformer_override_safetensor: str | None = None,
) -> torch.nn.Module:
transformer = load_module_from_path(
model_path=init_from,
module_type="transformer",
training_config=training_config,
disable_custom_init_weights=(disable_custom_init_weights),
override_transformer_cls_name=(self._transformer_cls_name),
transformer_override_safetensor=(transformer_override_safetensor),
)
ckpt_type = (enable_gradient_checkpointing_type or getattr(
getattr(training_config, "model", None),
"enable_gradient_checkpointing_type",
None,
))
if trainable and ckpt_type:
transformer = apply_activation_checkpointing(
transformer,
checkpointing_type=ckpt_type,
)
if self._enable_lora_if_configured(transformer):
return transformer
transformer = apply_trainable(transformer, trainable=trainable)
return transformer
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
def init_preprocessors(self, training_config: TrainingConfig) -> None:
self.vae = load_module_from_path(
model_path=str(training_config.model_path),
module_type="vae",
training_config=training_config,
)
self.world_group = get_world_group()
self.sp_group = get_sp_group()
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.train.utils.dataloader import (
build_parquet_t2v_train_dataloader, )
preprocessed_data_type = str(getattr(
training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
if preprocessed_data_type != "t2v":
raise ValueError("Unsupported Kandinsky5 preprocessed_data_type: "
f"{preprocessed_data_type!r}")
# Qwen's usable (post-template-trim) embedding length, +1 for the
# prepended CLIP pooled row (see module docstring / prepare_batch).
qwen_text_len = int(training_config.pipeline_config.text_encoder_configs[ # type: ignore[union-attr]
0].arch_config.text_len)
self.dataloader = build_parquet_t2v_train_dataloader(
training_config.data,
text_len=qwen_text_len + 1,
parquet_schema=pyarrow_schema_t2v,
)
self.start_step = 0
@property
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def set_requires_negative_conditioning(self, requires: bool) -> None:
self._requires_negative_conditioning = bool(requires)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
self.timestep_shift,
self.num_train_timestep,
)
return timestep.clamp(self.min_timestep, self.max_timestep)
def on_train_start(self) -> None:
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@torch.no_grad()
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
if self.vae is None:
raise RuntimeError("Kandinsky5 VAE is not initialized")
latents = latents_b_t_c_h_w.permute(0, 2, 1, 3, 4).float()
# Kandinsky5 shares Hunyuan's VAE: only a scalar scaling_factor, no
# latents_mean/latents_std, no handles_latent_denorm. Inverse of the
# encode-side `normalize_dit_input("hunyuan", ...)` multiply.
denorm = latents / self.vae.scaling_factor
media = self.vae.to(latents.device).decode(denorm)
return (media / 2 + 0.5).clamp(0, 1)
# ------------------------------------------------------------------
# Runtime primitives
# ------------------------------------------------------------------
def prepare_batch(
self,
raw_batch: dict[str, Any],
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
assert self.training_config is not None
tc = self.training_config
dtype = self._get_training_dtype()
device = self.device
training_batch = TrainingBatch()
# Unpack the CLIP-pooled-row-prepended text_embedding field: row 0
# is the zero-padded CLIP pooled vector, rows 1: are the Qwen
# sequence embeddings. See module docstring.
packed_embeds = raw_batch["text_embedding"]
packed_mask = raw_batch["text_attention_mask"]
pooled_dim = int(tc.pipeline_config.dit_config.arch_config.in_text_dim2 # type: ignore[union-attr]
)
pooled_projections = packed_embeds[:, 0, :pooled_dim]
qwen_embeds = packed_embeds[:, 1:, :]
qwen_mask = packed_mask[:, 1:]
infos = raw_batch.get("info_list")
if latents_source == "zeros":
batch_size = packed_embeds.shape[0]
vae_config = (
tc.pipeline_config.vae_config.arch_config # type: ignore[union-attr]
)
# Read off the loaded transformer module, not
# tc.pipeline_config.dit_config.arch_config: the latter carries
# YAML/dataclass defaults and isn't synced with the actual
# checkpoint's transformer/config.json (e.g. the shipped
# Kandinsky5-Lite checkpoint has in_visual_dim=16, vs. the
# dataclass default of 4).
num_channels = int(self.transformer.in_visual_dim)
spatial_compression_ratio = (vae_config.spatial_compression_ratio)
latent_height = (tc.data.num_height // spatial_compression_ratio)
latent_width = (tc.data.num_width // spatial_compression_ratio)
latents = torch.zeros(
batch_size,
num_channels,
tc.data.num_latent_t,
latent_height,
latent_width,
device=device,
dtype=dtype,
)
elif latents_source == "data":
if "vae_latent" not in raw_batch:
raise ValueError("vae_latent not found in batch "
"and latents_source='data'")
latents = raw_batch["vae_latent"]
latents = latents[:, :, :tc.data.num_latent_t]
latents = latents.to(device, dtype=dtype)
else:
raise ValueError(f"Unknown latents_source: "
f"{latents_source!r}")
training_batch.latents = latents
training_batch.encoder_hidden_states = qwen_embeds.to(device, dtype=dtype)
training_batch.encoder_attention_mask = qwen_mask.to(device, dtype=dtype)
training_batch.infos = infos
pooled_projections = pooled_projections.to(device, dtype=dtype)
training_batch.latents = normalize_dit_input("hunyuan", training_batch.latents, self.vae)
training_batch = self._prepare_dit_inputs(training_batch, generator, pooled_projections)
training_batch = self._build_attention_metadata(training_batch)
training_batch.attn_metadata_vsa = copy.copy(training_batch.attn_metadata)
if training_batch.attn_metadata is not None:
training_batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore[attr-defined]
return training_batch
def add_noise(
self,
clean_latents: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
b, t = clean_latents.shape[:2]
noisy = self.noise_scheduler.add_noise(
clean_latents.flatten(0, 1),
noise.flatten(0, 1),
timestep,
).unflatten(0, (b, t))
return noisy
def predict_noise(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = self._get_training_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}")
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
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))
transformer = self._get_transformer(timestep)
out = transformer(**input_kwargs)
sample = out.sample if hasattr(out, "sample") else out
# Kandinsky5's native hidden_states layout is channel-last
# [B, T, H, W, C]; convert to the common ModelBase convention
# [B, T, C, H, W].
pred_noise = sample.permute(0, 1, 4, 2, 3)
return pred_noise
def backward(
self,
loss: torch.Tensor,
ctx: Any,
*,
grad_accum_rounds: int,
) -> None:
timesteps, attn_metadata = ctx
with set_forward_context(
current_timestep=timesteps,
attn_metadata=attn_metadata,
):
(loss / max(1, int(grad_accum_rounds))).backward()
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _get_training_dtype(self) -> torch.dtype:
return torch.bfloat16
def _init_timestep_mechanics(self) -> None:
assert self.training_config is not None
tc = self.training_config
self.timestep_shift = float(tc.pipeline_config.flow_shift # type: ignore[union-attr]
)
self.num_train_timestep = int(self.noise_scheduler.num_train_timesteps)
self.min_timestep = 0
self.max_timestep = self.num_train_timestep
def ensure_negative_conditioning(self) -> None:
"""Encode the negative prompt with dual text encoders (Qwen + CLIP).
Every rank encodes independently (same rationale as Hunyuan's
override: avoids NCCL deadlocks from asymmetric rank-0-only
encoding). Cannot reuse ``encode_negative_prompt`` unmodified for
the Qwen encoder: ``kandinsky5_qwen_postprocess_text`` requires the
attention mask as a second positional argument and returns a
``(embeds, trimmed_mask)`` tuple, unlike the single-arg
postprocess-func contract that helper assumes.
"""
if self.negative_prompt_embeds is not None:
return
assert self.training_config is not None
tc = self.training_config
device = self.device
dtype = self._get_training_dtype()
from transformers import AutoTokenizer
from fastvideo.configs.pipelines.base import preprocess_text
from fastvideo.configs.pipelines.kandinsky5 import (
kandinsky5_clip_postprocess_text,
kandinsky5_qwen_postprocess_text,
kandinsky5_qwen_preprocess_text,
)
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.utils import maybe_download_model
pipeline_config = tc.pipeline_config
assert pipeline_config is not None
model_path = maybe_download_model(tc.model_path)
inference_args = make_inference_args(tc, model_path=model_path)
inference_args.text_encoder_cpu_offload = False
sampling_param = SamplingParam.from_pretrained(tc.model_path)
negative_prompt = sampling_param.negative_prompt
qwen_cfg = pipeline_config.text_encoder_configs[0]
clip_cfg = pipeline_config.text_encoder_configs[1]
loader = TextEncoderLoader()
# --- Qwen / Reason1 ---
qwen_enc = loader.load(
os.path.join(model_path, "text_encoder"),
inference_args,
).to(device).eval()
qwen_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
qwen_tok_kwargs = dict(qwen_cfg.tokenizer_kwargs)
qwen_text = kandinsky5_qwen_preprocess_text(negative_prompt)
with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None):
qwen_inputs = qwen_tok(qwen_text, **qwen_tok_kwargs).to(device)
qwen_out = qwen_enc(
input_ids=qwen_inputs.input_ids,
attention_mask=qwen_inputs.attention_mask,
output_hidden_states=True,
)
qwen_embeds, qwen_mask = kandinsky5_qwen_postprocess_text(qwen_out, qwen_inputs.attention_mask)
del qwen_enc, qwen_tok
# --- CLIP ---
clip_enc = loader.load(
os.path.join(model_path, "text_encoder_2"),
inference_args,
).to(device).eval()
clip_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer_2"))
clip_tok_kwargs = dict(clip_cfg.tokenizer_kwargs)
clip_text = preprocess_text(negative_prompt)
with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None):
clip_inputs = clip_tok(clip_text, **clip_tok_kwargs).to(device)
clip_out = clip_enc(
input_ids=clip_inputs.input_ids,
attention_mask=clip_inputs.attention_mask,
)
clip_pooled = kandinsky5_clip_postprocess_text(clip_out)
del clip_enc, clip_tok
self.negative_prompt_embeds = qwen_embeds.to(device=device, dtype=dtype)
self.negative_prompt_attention_mask = qwen_mask.to(device=device, dtype=dtype)
self.negative_pooled_embeds = clip_pooled.to(device=device, dtype=dtype)
def _sample_timesteps(
self,
batch_size: int,
device: torch.device,
generator: torch.Generator,
) -> torch.Tensor:
assert self.training_config is not None
tc = self.training_config
u = compute_density_for_timestep_sampling(
weighting_scheme=tc.model.weighting_scheme,
batch_size=batch_size,
generator=generator,
device=device,
logit_mean=tc.model.logit_mean,
logit_std=tc.model.logit_std,
mode_scale=tc.model.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices.cpu()].to(device=device)
def _build_attention_metadata(self, training_batch: TrainingBatch) -> TrainingBatch:
# Scope is 480p T2V with dense/local attention only -- VSA/VMOBA
# sparse attention is never engaged for Kandinsky5 in this recipe.
training_batch.attn_metadata = None
return training_batch
def _prepare_dit_inputs(
self,
training_batch: TrainingBatch,
generator: torch.Generator,
pooled_projections: torch.Tensor,
) -> TrainingBatch:
assert self.training_config is not None
tc = self.training_config
latents = training_batch.latents
assert isinstance(latents, torch.Tensor)
batch_size = latents.shape[0]
device = latents.device
num_height = int(tc.data.num_height)
num_width = int(tc.data.num_width)
if not (_MIN_480P_SIDE <= num_height <= _MAX_480P_SIDE and _MIN_480P_SIDE <= num_width <= _MAX_480P_SIDE):
raise ValueError("Kandinsky5Model only supports 480p training "
f"(height/width in [{_MIN_480P_SIDE}, {_MAX_480P_SIDE}]); "
f"got num_height={num_height}, num_width={num_width}. "
"A larger resolution needs a different visual RoPE "
"scale_factor, which this wrapper does not implement.")
noise = torch.randn(
latents.shape,
generator=generator,
device=device,
dtype=latents.dtype,
)
timesteps = self._sample_timesteps(batch_size, device, generator)
if int(tc.distributed.sp_size or 1) > 1:
self.sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
self.noise_scheduler,
device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = ((1.0 - sigmas) * latents + sigmas * noise)
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = latents.shape
# Trim the Qwen sequence to the longest real (non-pad) sample in
# this batch, mirroring Kandinsky5DenoisingStage._text_rope_pos.
# Kandinsky5's cross-attention has no attention-mask plumbing at
# all -- it always attends to the full encoder_hidden_states
# sequence -- so any padding left in the tensor beyond this point
# would silently leak into cross-attention. Batched *inference*
# already relies on the same property (tokenizer padding="True"
# pads only to the batch's longest real sample).
qwen_embeds = training_batch.encoder_hidden_states
qwen_mask = training_batch.encoder_attention_mask
assert qwen_embeds is not None and qwen_mask is not None
# bf16 cannot exactly accumulate sums beyond ~256, so count valid
# tokens in float32 to avoid an off-by-a-few error on long prompts.
valid_len = max(1, int(qwen_mask.float().sum(dim=1).max().item()))
qwen_embeds = qwen_embeds[:, :valid_len]
qwen_mask = qwen_mask[:, :valid_len]
text_rope_pos = torch.arange(valid_len, device=device)
patch_size = (
tc.pipeline_config.dit_config.arch_config.patch_size # type: ignore[union-attr]
)
spatial_compression_ratio = (
tc.pipeline_config.vae_config.arch_config.spatial_compression_ratio # type: ignore[union-attr]
)
latent_h = num_height // spatial_compression_ratio
latent_w = num_width // spatial_compression_ratio
visual_rope_pos = [
torch.arange(int(tc.data.num_latent_t), device=device),
torch.arange(latent_h // patch_size[1], device=device),
torch.arange(latent_w // patch_size[2], device=device),
]
training_batch.conditional_dict = {
"encoder_hidden_states": qwen_embeds,
"encoder_attention_mask": qwen_mask,
"pooled_projections": pooled_projections,
"text_rope_pos": text_rope_pos,
"visual_rope_pos": visual_rope_pos,
"scale_factor": _SCALE_FACTOR_480P,
"sparse_params": None,
}
if (self.negative_prompt_embeds is not None and self.negative_prompt_attention_mask is not None
and self.negative_pooled_embeds is not None):
neg_embeds = self.negative_prompt_embeds
neg_mask = self.negative_prompt_attention_mask
neg_pooled = self.negative_pooled_embeds
if neg_embeds.shape[0] == 1 and batch_size > 1:
neg_embeds = neg_embeds.expand(batch_size, *neg_embeds.shape[1:]).contiguous()
if neg_mask.shape[0] == 1 and batch_size > 1:
neg_mask = neg_mask.expand(batch_size, *neg_mask.shape[1:]).contiguous()
if neg_pooled.shape[0] == 1 and batch_size > 1:
neg_pooled = neg_pooled.expand(batch_size, *neg_pooled.shape[1:]).contiguous()
neg_text_rope_pos = torch.arange(neg_embeds.shape[1], device=device)
training_batch.unconditional_dict = {
"encoder_hidden_states": neg_embeds,
"encoder_attention_mask": neg_mask,
"pooled_projections": neg_pooled,
"text_rope_pos": neg_text_rope_pos,
"visual_rope_pos": visual_rope_pos,
"scale_factor": _SCALE_FACTOR_480P,
"sparse_params": None,
}
training_batch.latents = (training_batch.latents.permute(0, 2, 1, 3, 4))
return training_batch
def _build_distill_input_kwargs(
self,
noise_input: torch.Tensor,
timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
) -> dict[str, Any]:
if text_dict is None:
raise ValueError("text_dict cannot be None for "
"Kandinsky5 distillation")
# noise_input is [B, T, C, H, W] (common ModelBase convention);
# Kandinsky5's transformer expects channel-last [B, T, H, W, C].
hidden_states = noise_input.permute(0, 1, 3, 4, 2)
if getattr(self.transformer, "visual_cond", False):
# The shipped Kandinsky5-Lite checkpoint has visual_cond=True
# (unified T2V/I2V conditioning): Kandinsky5VisualEmbeddings
# always expects [real_latent | zero_cond | zero_mask]
# concatenated on the channel dim (2*in_visual_dim+1 total),
# even for pure T2V with no actual image conditioning. Mirrors
# Kandinsky5LatentPreparationStage's inference-time padding
# (fastvideo/pipelines/stages/kandinsky5.py).
cond = torch.zeros_like(hidden_states)
mask = torch.zeros(*hidden_states.shape[:-1], 1, device=hidden_states.device, dtype=hidden_states.dtype)
hidden_states = torch.cat([hidden_states, cond, mask], dim=-1)
return {
"hidden_states": hidden_states,
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"pooled_projections": text_dict["pooled_projections"],
"timestep": timestep,
"visual_rope_pos": text_dict["visual_rope_pos"],
"text_rope_pos": text_dict["text_rope_pos"],
"scale_factor": text_dict["scale_factor"],
"sparse_params": text_dict["sparse_params"],
"return_dict": True,
}
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
return self.transformer
def _get_uncond_text_dict(
self,
batch: TrainingBatch,
*,
cfg_uncond: dict[str, Any] | None,
) -> dict[str, torch.Tensor]:
if cfg_uncond is None:
text_dict = getattr(batch, "unconditional_dict", None)
if text_dict is None:
raise RuntimeError("Missing unconditional_dict; "
"ensure_negative_conditioning() "
"may have failed")
return text_dict
on_missing_raw = cfg_uncond.get("on_missing", "error")
if not isinstance(on_missing_raw, str):
raise ValueError("method_config.cfg_uncond.on_missing "
"must be a string, got "
f"{type(on_missing_raw).__name__}")
on_missing = on_missing_raw.strip().lower()
if on_missing not in {"error", "ignore"}:
raise ValueError("method_config.cfg_uncond.on_missing "
"must be one of {error, ignore}, got "
f"{on_missing_raw!r}")
for channel, policy_raw in cfg_uncond.items():
if channel in {"on_missing", "text"}:
continue
if policy_raw is None:
continue
if not isinstance(policy_raw, str):
raise ValueError("method_config.cfg_uncond values "
"must be strings, got "
f"{channel}="
f"{type(policy_raw).__name__}")
policy = policy_raw.strip().lower()
if policy == "keep":
continue
if on_missing == "ignore":
continue
raise ValueError("Kandinsky5Model does not support "
"cfg_uncond channel "
f"{channel!r} (policy={policy!r}). "
"Set cfg_uncond.on_missing=ignore or "
"remove the channel.")
text_policy_raw = cfg_uncond.get("text", None)
if text_policy_raw is None:
text_policy = "negative_prompt"
elif not isinstance(text_policy_raw, str):
raise ValueError("method_config.cfg_uncond.text must be "
"a string, got "
f"{type(text_policy_raw).__name__}")
else:
text_policy = (text_policy_raw.strip().lower())
if text_policy in {"negative_prompt"}:
text_dict = getattr(batch, "unconditional_dict", None)
if text_dict is None:
raise RuntimeError("Missing unconditional_dict; "
"ensure_negative_conditioning() "
"may have failed")
return text_dict
if text_policy == "keep":
if batch.conditional_dict is None:
raise RuntimeError("Missing conditional_dict in "
"TrainingBatch")
return batch.conditional_dict
if text_policy == "zero":
if batch.conditional_dict is None:
raise RuntimeError("Missing conditional_dict in "
"TrainingBatch")
cond = batch.conditional_dict
enc = cond["encoder_hidden_states"]
mask = cond["encoder_attention_mask"]
pooled = cond["pooled_projections"]
if not torch.is_tensor(enc) or not torch.is_tensor(mask) or not torch.is_tensor(pooled):
raise TypeError("conditional_dict must contain "
"tensor text inputs")
return {
"encoder_hidden_states": torch.zeros_like(enc),
"encoder_attention_mask": torch.zeros_like(mask),
"pooled_projections": torch.zeros_like(pooled),
"text_rope_pos": cond["text_rope_pos"],
"visual_rope_pos": cond["visual_rope_pos"],
"scale_factor": cond["scale_factor"],
"sparse_params": cond["sparse_params"],
}
if text_policy == "drop":
raise ValueError("cfg_uncond.text=drop is not supported "
"for Kandinsky5. Use "
"{negative_prompt, keep, zero}.")
raise ValueError("cfg_uncond.text must be one of "
"{negative_prompt, keep, zero, drop}, got "
f"{text_policy_raw!r}")
@@ -12,6 +12,8 @@ TRANSFORMER_BLOCK_NAMES = [
"temporal_transformer_blocks",
"transformer_double_blocks",
"transformer_single_blocks",
"text_transformer_blocks",
"visual_transformer_blocks",
]