Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
35a283531e | ||
|
|
3c9e604199 | ||
|
|
2cb832290a | ||
|
|
5aae1c093a | ||
|
|
068c7ba57e | ||
|
|
9b95acadd8 | ||
|
|
5ee7980fcb | ||
|
|
3e1201c455 | ||
|
|
64f60b3d8a | ||
|
|
21d751fbb4 | ||
|
|
2dbb0c94a9 |
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
Binary file not shown.
@@ -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)
|
||||
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user