Compare commits

...
37 Commits
Author SHA1 Message Date
SolitaryThinker a288460497 [test]: collect the sdpa mask contract test in the CI unit lane
fastvideo/tests/attention/ has no CI lane (dark-directory allowlist), so
wire the new CPU-only contract test file into run_unit_test explicitly.
2026-07-06 17:26:02 -07:00
SolitaryThinker 45fc714f88 [bugfix]: keep sdpa metadata masks 2D; lift to 4D only inside the torch-sdpa impl
SDPAMetadataBuilder.build() lifted 2D [batch, key_len] padding masks to 4D
[batch, 1, 1, key_len], but the metadata is cross-backend: HYWorld and
HunyuanVideo15 build SDPAMetadata while the layer selector picks FLASH_ATTN
when installed. flash_attn.py's equal-length branch pads
attn_mask.shape[1] assuming 2D, so a 4D mask produced a garbage-length pad
and an out-of-bounds gather (CUDA device-side assert, SIGABRT in CI
transformer tests).

Store the mask 2D in the builder (the convention every in-tree
producer/consumer assumes) and do the 2D -> [batch, 1, 1, key_len] lift
inside _normalize_attn_mask_for_sdpa, the torch-sdpa-only consumer, so
torch.sdpa still does not misread 2D as its [query_len, key_len]
broadcast. The int -> bool coercion is unchanged. Adds a CPU-only
regression test for the contract.
2026-07-06 17:15:29 -07:00
Will LinandClaude Fable 5 38d1d7bdc5 [misc]: scope the kandinsky ssim FA4 pin per-test and name the min-length validator
The module-level FASTVIDEO_FA4=0 executed at pytest collection and disabled
FA4 for every SSIM test in the same process, diverging them from their
FA4-seeded references; set and restore it inside the test instead. Replace
the lambda prompt-embeds check with a curried V.min_list_length(2) whose
__name__ shows up in verification failure messages, and drop the no-op
num_frames re-assignment in the full-quality params map.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-06 22:53:24 +00:00
Will LinandClaude Fable 5 3a44ae6316 [bugfix]: probe gated-repo access with auth_check and keep transient hub failures loud
Replace the three copy-pasted .gitattributes-download guards with one shared
skip_if_gated_repo_inaccessible helper: local weights bypass the probe (an
offline box with cached weights runs instead of ERRORing at collection via
the uncaught LocalEntryNotFoundError), HfApi.auth_check identifies real
authorization failures, and only those skip — transient hub errors proceed
and let the download serve from cache or fail loudly, instead of
green-skipping with a message that falsely blames gated access. Drop the
except-ValueError wrappers around maybe_download_model for the same reason.
Also skip the HYWorld transformer test before its multi-GB download when an
env-forced backend has no reference latent.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-06 22:53:24 +00:00
Will LinandClaude Fable 5 6918418525 [bugfix]: coerce integer sdpa masks and keep torch 2D mask semantics
_normalize_attn_mask_for_sdpa passed tokenizer-produced int64 0/1 masks
straight to F.scaled_dot_product_attention, which rejects long dtypes — the
HYWorld pipeline crashed on exactly the SDPA-fallback machines the fastcheck
stabilization targeted; coerce non-bool/non-float masks to bool. Lift the
[batch, key_len] -> [batch, 1, 1, key_len] reinterpretation out of the impl
into SDPAMetadataBuilder.build, where that semantic is known, so 2D masks
reaching the impl keep torch's documented [query_len, key_len] broadcast
meaning. Reuse the identical metadata for HYWorld's prope attention instead
of rebuilding it, and note that the metadata does not choose the executing
kernel.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-06 22:53:09 +00:00
Aryan Kumar 881400490c [bugfix]: stabilize fastcheck lanes for Kandinsky PR 2026-07-06 11:42:26 -07:00
SolitaryThinker fca6fb3aa1 [test]: add reduced-size Kandinsky-5 T2V SSIM test 2026-07-05 16:31:59 -07:00
SolitaryThinker fa6e50fc7c [misc]: drop the bundled Modal checkout rewrite (separate PR material) 2026-07-05 16:09:06 -07:00
SolitaryThinker 9f28d34d6d [bugfix]: keep text-encoder CPU offload on by default
Revert the default flips from 40c8b7c1 (OffloadConfig.text_encoder and
FastVideoArgs.text_encoder_cpu_offload False -> True). Flipping the
global default changes the baseline VRAM footprint for every model
family, not just Kandinsky, and breaks the unit lane:
fastvideo/tests/api/test_parser.py::test_load_run_config_supports_yaml_roundtrip
still golden-checks text_encoder offload as True.

Nothing in this PR needs the flip: the Kandinsky examples pass
text_encoder_cpu_offload explicitly, and the offload-aware
TextEncodingStage (kept) makes offload work for Kandinsky's CLIP
either way. Offload stays opt-out rather than opt-in.
2026-07-05 15:25:06 -07:00
Will LinandClaude Fable 5 a7470d2da3 [bugfix]: retry the non-pr ci checkout too
The blob-less clone defers all file downloads to checkout, so the direct
'git checkout <commit>' on main-branch/merge-queue shards is a large lazy
blob fetch with no retry — the same transient-disconnect flake class the PR
path's retry loop was added for. Share one retry wrapper across both paths;
all four command compositions verified with bash -n.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 821b4bc275 [misc]: document why i2v keeps diffusers cond-channel semantics
The official kandinsky-5 repo zeros the visual_cond block for I2V while the
diffusers port copies the image latent into it. A same-seed A/B on the
Pro-distilled Diffusers export shows the copy is required (zeroing smears
faces mid-video), so keep the diffusers semantics and record the evidence.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 09c42cc73a [bugfix]: raise on unknown rope types in custom qwen2.5-vl encoder
The transformers-5.x compat fallback rescued every rope_type missing from
ROPE_INIT_FUNCTIONS, silently computing unscaled default frequencies for
typo'd or future scaling types; only 'default' (removed from the table in
transformers>=5) falls back now, anything else raises KeyError.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 d6ed0e54da [bugfix]: restore pipeline-class override short-circuit and dedup kandinsky5 detectors
get_model_info again skips model_index.json resolution entirely when
override_pipeline_cls_name is set, so overrides keep working for local
checkpoints without a diffusers model_index (and trigger no downloads).
Replace the nine copy-pasted kandinsky5 detectors with a parameterized
factory (identical matching semantics) and add a dedicated I2V Lite distilled
entry + preset so those checkpoints get distilled sampling defaults
(guidance 1.0 / 16 steps) instead of the sft fallback's 5.0 / 50.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 6ce4993ab7 [bugfix]: kandinsky5 stages: official rng order, attn metadata, tensor resize, decode dedup
Match the official kandinskylab/kandinsky-5 RNG order for I2V seed parity:
image encoding now runs after latent preparation and places the conditioning
latent itself, so the initial noise is the seeded generator's first draw
(upstream leaves the image-latent sample unseeded; drawing it second from the
same generator keeps FastVideo deterministic). The denoising stage now sets
the forward context with NABLA attention metadata per step — which also
routes the dense path through LocalAttention instead of the no-context SDPA
fallback. Resize tensor conditioning images to the requested resolution like
the PIL path. Collapse Kandinsky5DecodingStage to a channels-first permute +
super().forward(), regaining the base class's MPS handling.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 b9dfb12502 [feat]: register NABLA sparse attention as an attention backend
Move nablaT_v2 + the flex-attention guard out of the Kandinsky5 DiT into
fastvideo/attention/backends/nabla.py with backend/metadata/builder/impl
classes, a NABLA_ATTN enum entry, and CUDA platform wiring that refuses to
silently fall back to dense attention when flex_attention is unavailable.
Add a generic default_backend parameter to the selector and LocalAttention —
a layer-level default that the global force and FASTVIDEO_ATTENTION_BACKEND
still override — so nabla checkpoints select the backend by default while
users keep env-var control. The DiT's sparse path now dispatches through a
dedicated LocalAttention (with a direct-flex fallback for standalone parity
tests without a forward context); verified bit-identical to the previous
inline flex call.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 375ba410c7 [misc]: default text_encoder_cpu_offload to False
Keep text encoders GPU-resident by default (both the legacy arg and the
engine-config schema); offload remains an explicit opt-in now that the
swap path makes it work for encoders without FSDP shard conditions.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 a43284ae1b [bugfix]: swap cpu-offloaded text encoders to GPU for the forward
text_encoder_cpu_offload means params rest on CPU between calls but compute
on GPU: FSDP-wrapped encoders (T5/llama/clip) already get this via
CPUOffloadPolicy streaming, but encoders without _fsdp_shard_conditions
(reason1/Qwen2.5-VL) were parked on CPU and, after the encoder-device input
routing added for kandinsky5, silently ran their 7B forward on CPU (crashing
outright with flash-attn installed). Move plain CPU-parked encoders to the
target device for the forward and back afterwards, mirroring the
image-encoder/VAE offload pattern; FSDP (DTensor) encoders keep the existing
input routing. Revert the examples to text_encoder_cpu_offload=True now that
the default works. Verified e2e: T2V/I2V outputs are metric-identical to
GPU-resident runs, with ~12GB lower VRAM during denoising.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 6ab2d6f98d [bugfix]: kandinsky5 stages: nabla validation, seeded i2v latent, trajectory support
Validate NABLA resolutions (divisible by 128) in latent preparation instead
of crashing mid-denoise with a cryptic reshape error after all encoding work.
Encode the I2V conditioning image with a generator-seeded sample instead of
mode() for diffusers seed parity. Restore base DecodingStage behavior the
override dropped (pipeline.add_module on lazy VAE reload,
return_trajectory_decoded) and record trajectory latents in the denoise loop.
Always normalize PIL conditioning images deterministically; keep the
diffusers-style range heuristic only for raw tensors and log the assumption.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 9472ed0aab [bugfix]: surface kandinsky5 flex-attention failures and preserve rotary dtype
Replace the bare except around the flex_attention import (which swallowed
every error, printed blank lines to stdout, and silently degraded NABLA
requests to dense attention over ~95k tokens) with except ImportError plus a
logger warning, and raise a RuntimeError when sparse attention is requested
without flex_attention, matching the reference implementation. Also restore
the input dtype in _apply_rotary instead of hard-casting to bf16, which
truncated fp32 q/k through a bf16 round-trip in fp32/parity runs.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 4eb4e252eb [bugfix]: preserve tokenizer_kwargs customizations across update_model_arch
TextEncoderArchConfig.__post_init__ rebuilt tokenizer_kwargs from scratch, and
TextEncoderLoader.load -> update_model_arch re-runs __post_init__, silently
wiping customizations applied by pipeline configs (kandinsky5/gen3c/longcat
set 'padding'; multi-prompt kandinsky5 runs then crash on ragged sequences).
Merge defaults under existing keys instead of rebuilding.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 0434e9dcf0 [bugfix]: fix kandinsky5 registry detectors and distilled repo name
The i2v detector matched any path containing 'i2v' without excluding 't2v',
misrouting T2V checkpoints stored under i2v-containing directories. The
variant detectors all require lite/pro tokens, so the model_index _class_name
fallback ('kandinsky5t2vpipeline') could never match and I2V-Lite-distilled
checkpoints had no detector at all; add base T2V/I2V fallback registrations
after the variants. Also fix the registered Lite distilled repo to the real
kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers (the old
name 404s) and correct the example comments.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 36cf64b15f [bugfix]: fix bash syntax error in modal ci pr checkout
The multi-line retry checkout command ended with a newline, so splicing it
into '{checkout_command} &&' left a lone '&&' on the line after 'done' — a
bash syntax error failing every PR-triggered shard before checkout. End the
f-string at 'done' so the composed script reads 'done &&'; verified all four
compositions (PR/direct x kernel/no-kernel) with bash -n.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 7a270e66be [bugfix]: keep kandinsky5 example text encoders on GPU
text_encoder_cpu_offload=True makes the Qwen encoder compute on CPU, where
flash-attn has no kernels; with flash_attn installed both examples crash in
the first text-encoding forward. The offloaded path only worked in envs
without flash_attn via the silent SDPA fallback.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 74de6f27a2 [bugfix]: fix kandinsky5 t2v frame-0 denoising and restore trained prompt template
Denoising keyed cond_frames off transformer.visual_cond, but official T2V
checkpoints also ship visual_cond=true, so latent frame 0 was never stepped
and decoded as pure noise in every T2V video; key off batch.image_latent
instead. Restore the byte-exact upstream prompt template (typos included):
the checkpoints were trained with it, and ENCODE_START_IDX=129 matches the
upstream template while the corrected wording moves user content to 127,
silently dropping the first user-prompt tokens. Verified against the
checkpoint tokenizer and by end-to-end T2V/I2V generation.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Will LinandClaude Fable 5 6979040225 [bugfix]: make custom qwen2.5-vl encoder compatible with transformers 5.x
transformers>=5 stops forwarding text-model attributes from the composite
Qwen2_5_VLConfig, drops the 'default' ROPE_INIT_FUNCTIONS entry, and requires
flash-attn functions to be preloaded before _flash_attention_forward. Flatten
text_config at model entry, replicate the 4.x default rope init, and preload
flash-attn (falling back to SDPA). Fixes Kandinsky5 and Cosmos 2.5 encoder
loading under the transformers>=4.57.3 pin.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-05 14:58:12 -07:00
Aryan Kumar 80972d07fa [bugfix]: harden modal ci pr checkout 2026-07-05 14:58:12 -07:00
Aryan Kumar 260b326d8b [bugfix]: address kandinsky5 i2v review feedback 2026-07-05 14:58:12 -07:00
Aryan Kumar c57fa1eab0 [misc]: polish kandinsky5 sparse attention updates 2026-07-05 14:58:12 -07:00
leffff 562b951208 [feat] add support for Kandinsky 5 Pro, Lite I2V 2026-07-05 14:58:12 -07:00
leffff aa893056e8 [misc] I2V debug 2026-07-05 14:58:12 -07:00
leffff 313c2985f0 [feat] add support for Kandinsky5 I2V 2026-07-05 14:58:12 -07:00
leffff b3b3858c02 [misc] add skeleton for Kandinsky5 I2V pipeline 2026-07-05 14:58:12 -07:00
leffff d39e2e9fd7 fix nabla attention 2026-07-05 14:58:12 -07:00
leffff a66fb26781 add nabla 2026-07-05 14:58:12 -07:00
leffff f303d94780 add exmaple for kandinsky5 T2V Lite 5s 2026-07-05 14:58:12 -07:00
coderabbitai[bot]andCodeRabbit dae7c0da89 fix: apply CodeRabbit auto-fixes
Fixed 2 file(s) based on 2 unresolved review comments.

Co-authored-by: CodeRabbit <noreply@coderabbit.ai>
2026-07-05 14:58:12 -07:00
Aryan Kumar 0ce6dc928f [feat]: add kandinsky5 pipeline support 2026-07-05 14:58:12 -07:00
34 changed files with 1901 additions and 79 deletions
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
IMAGE_PATH = "assets/girl.png"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A woman stands up and walks away"
)
_ = generator.generate_video(
prompt,
image_path=IMAGE_PATH,
output_path=OUTPUT_PATH,
save_video=True,
height=1024,
width=1024,
num_frames=121,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
if __name__ == "__main__":
main()
+147
View File
@@ -0,0 +1,147 @@
# SPDX-License-Identifier: Apache-2.0
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
precomputed spatio-temporal-window (STA) mask carried on the attention
metadata. The mask spans the full sequence, so this backend does not support
sequence parallelism — use it via LocalAttention only.
"""
import math
from dataclasses import dataclass
from typing import Any
import torch
try:
from torch.nn.attention.flex_attention import BlockMask, flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
CAN_USE_FLEX_ATTN = True
except ImportError:
CAN_USE_FLEX_ATTN = False
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
def nablaT_v2(
q: torch.Tensor,
k: torch.Tensor,
sta: torch.Tensor,
thr: float = 0.9,
) -> "BlockMask":
q = q.transpose(1, 2).contiguous()
k = k.transpose(1, 2).contiguous()
# Map estimation
B, h, S, D = q.shape
s1 = S // 64
qa = q.reshape(B, h, s1, 64, D).mean(-2)
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
map = qa @ ka
map = torch.softmax(map / math.sqrt(D), dim=-1)
# Map binarization
vals, inds = map.sort(-1)
cvals = vals.cumsum_(-1)
mask = (cvals >= 1 - thr).int()
mask = mask.gather(-1, inds.argsort(-1))
mask = torch.logical_or(mask, sta)
# BlockMask creation
kv_nb = mask.sum(-1).to(torch.int32)
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
class NablaAttentionBackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "NABLA_ATTN"
@staticmethod
def get_impl_cls() -> type["NablaAttentionImpl"]:
return NablaAttentionImpl
@staticmethod
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
return NablaAttentionMetadata
@staticmethod
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
return NablaAttentionMetadataBuilder
@dataclass
class NablaAttentionMetadata(AttentionMetadata):
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
sta_mask: torch.Tensor = None # type: ignore[assignment]
# Cumulative-probability threshold for block-map binarization.
P: float = 0.9
visual_shape: tuple[int, int, int] = (0, 0, 0)
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build(
self,
current_timestep: int,
sta_mask: torch.Tensor,
P: float,
visual_shape: tuple[int, int, int],
**kwargs: Any,
) -> NablaAttentionMetadata:
return NablaAttentionMetadata(
current_timestep=current_timestep,
sta_mask=sta_mask,
P=P,
visual_shape=visual_shape,
)
class NablaAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
if not CAN_USE_FLEX_ATTN:
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
"which is unavailable in this PyTorch build.")
if causal:
raise ValueError("NABLA attention does not support causal masking.")
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: NablaAttentionMetadata,
) -> torch.Tensor:
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
return flex_attention(
query=query.transpose(1, 2),
key=key.transpose(1, 2),
value=value.transpose(1, 2),
block_mask=block_mask,
).transpose(1, 2)
+44 -1
View File
@@ -2,6 +2,7 @@
import torch
from dataclasses import dataclass
from torch.nn import functional as F
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
from fastvideo.logger import init_logger
@@ -19,7 +20,7 @@ class SDPABackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "SDPA"
return "TORCH_SDPA"
@staticmethod
def get_impl_cls() -> type["SDPAImpl"]:
@@ -49,9 +50,50 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
current_timestep: int,
attn_mask: torch.Tensor,
) -> SDPAMetadata:
# Store the mask exactly as passed. The metadata is cross-backend:
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
# layer's selector may pick FLASH_ATTN, whose impl requires the
# tokenizer-style 2D [batch, key_len] padding mask. Any reshaping for
# torch.sdpa happens inside the SDPA impl (_normalize_attn_mask_for_sdpa).
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
def _normalize_attn_mask_for_sdpa(
attn_mask: torch.Tensor | None,
query: torch.Tensor,
key: torch.Tensor,
) -> torch.Tensor | None:
if attn_mask is None:
return None
attn_mask = attn_mask.to(device=query.device)
# F.scaled_dot_product_attention only accepts bool or float masks;
# tokenizers commonly produce int64 0/1 padding masks.
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
attn_mask = attn_mask != 0
key_len = key.shape[-2]
if attn_mask.shape[-1] > key_len:
raise ValueError("Invalid attention mask length for SDPA: "
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
if attn_mask.shape[-1] < key_len:
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
# non-text tokens the tokenizer mask does not cover.
valid_value = True if attn_mask.dtype == torch.bool else 0.0
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
if attn_mask.dim() == 2:
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
# reinterpret 2D as its documented [query_len, key_len] broadcast.
return attn_mask[:, None, None, :]
if attn_mask.dim() == 3:
return attn_mask[:, None, :, :]
if attn_mask.dim() == 4:
return attn_mask
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
class SDPAImpl(AttentionImpl):
def __init__(
@@ -82,6 +124,7 @@ class SDPAImpl(AttentionImpl):
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
and hasattr(attn_metadata, "attn_mask")) else None
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
attn_kwargs = {
"attn_mask": attn_mask,
"dropout_p": self.dropout,
+5 -1
View File
@@ -252,6 +252,7 @@ class LocalAttention(nn.Module):
causal: bool = False,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -262,7 +263,10 @@ class LocalAttention(nn.Module):
num_kv_heads = num_heads
dtype = get_compute_dtype()
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
attn_backend = get_attn_backend(head_size,
dtype,
supported_attention_backends=supported_attention_backends,
default_backend=default_backend)
impl_cls = attn_backend.get_impl_cls()
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
+9 -1
View File
@@ -84,8 +84,9 @@ def get_attn_backend(
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
@cache
@@ -94,6 +95,7 @@ def _cached_get_attn_backend(
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
@@ -112,6 +114,12 @@ def _cached_get_attn_backend(
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
# Layer-level default (e.g. a checkpoint that requires a specific sparse
# backend). Lower precedence than the global force and the env var, so
# users can still override it.
if selected_backend is None and default_backend is not None:
selected_backend = default_backend
# get device-specific attn_backend
from fastvideo.platforms import current_platform
+14 -4
View File
@@ -2,14 +2,24 @@
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
@dataclass
class Kandinsky5ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [
lambda n, m:
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
])
_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.
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.NABLA_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
# Native FastVideo implementation uses the same parameter names as diffusers
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
+6 -1
View File
@@ -43,11 +43,16 @@ class TextEncoderArchConfig(EncoderArchConfig):
require_processor: bool = False
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
# update_model_arch re-runs __post_init__ after pipeline configs may
# have customized tokenizer_kwargs (e.g. kandinsky5/gen3c/longcat set
# "padding"); rebuilding the dict here would silently wipe those
# customizations, so only fill in defaults for keys not already set.
defaults = {
"truncation": True,
"max_length": self.text_len,
"return_tensors": "pt",
}
self.tokenizer_kwargs = defaults | self.tokenizer_kwargs
@dataclass
+3 -1
View File
@@ -5,6 +5,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.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
@@ -16,5 +17,6 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"HYWorldConfig", "Kandinsky5T2VConfig", "Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+122
View File
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, CLIPTextConfig
from fastvideo.configs.models.encoders.reason1 import Reason1Config
from fastvideo.configs.models.vaes import HunyuanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
# Byte-exact copy of the upstream Kandinsky5/diffusers template, including the
# "promt"/"scren" typos: the checkpoints were trained with this exact system # codespell:ignore promt,scren
# prompt, and ENCODE_START_IDX below is the tokenized length of everything
# before the user prompt. Fixing the typos shifts user content to index 127
# and mis-conditions every generation.
KANDINSKY5_PROMPT_TEMPLATE = "\n".join([
"<|im_start|>system\nYou are a promt engineer. Describe the video in detail.", # codespell:ignore promt
"Describe how the camera moves or shakes, describe the zoom and view angle, whether it follows the objects.",
"Describe the location of the video, main characters or objects and their action.",
"Describe the dynamism of the video and presented actions.",
"Name the visual style of the video: whether it is a professional footage, user generated content, some kind of animation, video game or scren content.", # codespell:ignore scren
"Describe the visual effects, postprocessing and transitions if they are presented in the video.",
"Pay attention to the order of key actions shown in the scene.<|im_end|>",
"<|im_start|>user\n{}<|im_end|>",
])
KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX = 129
def kandinsky5_qwen_preprocess_text(prompt: str) -> str:
if not prompt.strip():
prompt = "."
return KANDINSKY5_PROMPT_TEMPLATE.format(prompt)
def kandinsky5_qwen_postprocess_text(outputs: BaseEncoderOutput,
mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if outputs.hidden_states is None:
raise RuntimeError("Kandinsky5 Qwen prompt embeddings require hidden_states.")
hidden_states = outputs.hidden_states[-1]
prompt_embeds = hidden_states[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
mask = mask[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
if prompt_embeds.shape[1] == 0:
prompt_embeds = hidden_states[:, -1:]
mask = torch.ones((mask.shape[0], 1), dtype=mask.dtype, device=mask.device)
return prompt_embeds, mask
def kandinsky5_clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
if outputs.pooler_output is None:
raise RuntimeError("Kandinsky5 CLIP pooled output is required.")
return outputs.pooler_output
@dataclass
class Kandinsky5T2VConfig(PipelineConfig):
"""Kandinsky-5.0 Lite text-to-video pipeline configuration."""
dit_config: DiTConfig = field(default_factory=Kandinsky5VideoConfig)
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Reason1Config(), CLIPTextConfig()))
preprocess_text_funcs: tuple[Callable[[str], Any], ...] = field(
default_factory=lambda: (kandinsky5_qwen_preprocess_text, preprocess_text))
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
default_factory=lambda: (kandinsky5_qwen_postprocess_text, kandinsky5_clip_postprocess_text))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
text_encoder_max_lengths: tuple[int, ...] = field(
default_factory=lambda: (KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX + 512, 77))
flow_shift: float | None = 5.0
vae_tiling: bool = True
def __post_init__(self) -> None:
if len(self.text_encoder_configs) != 2:
raise ValueError(f"Kandinsky5 pipeline requires exactly 2 text encoders (qwen and clip), "
f"but got {len(self.text_encoder_configs)} encoder(s).")
if len(self.text_encoder_precisions) != 2:
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder precisions, "
f"but got {len(self.text_encoder_precisions)}.")
if len(self.text_encoder_max_lengths) != 2:
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder max lengths, "
f"but got {len(self.text_encoder_max_lengths)}.")
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
qwen_cfg = self.text_encoder_configs[0]
qwen_cfg.arch_config.output_hidden_states = True
qwen_cfg.arch_config.tokenizer_kwargs.update({
"padding": True,
"truncation": True,
"return_tensors": "pt",
})
clip_cfg = self.text_encoder_configs[1]
clip_cfg.arch_config.tokenizer_kwargs.update({
"padding": "max_length",
"max_length": 77,
"truncation": True,
"add_special_tokens": True,
"return_tensors": "pt",
})
@dataclass
class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
"""Kandinsky-5.0 image-to-video pipeline configuration."""
def __post_init__(self) -> None:
super().__post_init__()
# I2V needs the VAE encoder to encode the conditioning image.
self.vae_config.load_encoder = True
+2 -2
View File
@@ -699,8 +699,8 @@ class IndividualTokenRefinerBlock(nn.Module):
if mask is None:
attn_output = self.attn(q, k, v)
else:
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
attn_metadata = SDPAMetadataBuilder().build(
current_timestep=0,
attn_mask=mask,
)
+7 -9
View File
@@ -173,8 +173,11 @@ class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
img_v_prope = img_v_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
# end hyworld
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
# The metadata only carries the text padding mask; the executing
# attention kernel is still whatever the layer's selector picked
# (flash-attn when installed), not necessarily torch SDPA.
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
attn_metadata = SDPAMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
@@ -192,14 +195,9 @@ class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
)
# begin hyworld
# attention with prope
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_prope = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# attention with prope (same text mask, so reuse the metadata)
# NOTE: Do NOT pass freqs_cis to prope attention - HY-WorldPlay does not apply RoPE to prope
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata_prope):
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn_prope, _ = self.attn(
img_q_prope,
img_k_prope,
+76 -25
View File
@@ -10,13 +10,21 @@ import torch.nn as nn
import torch.nn.functional as F
from fastvideo.attention import LocalAttention
from fastvideo.attention.backends.nabla import CAN_USE_FLEX_ATTN, flex_attention, nablaT_v2
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.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
if not CAN_USE_FLEX_ATTN:
logger.warning("torch.nn.attention.flex_attention is unavailable in this PyTorch build; "
"Kandinsky5 NABLA sparse attention (Pro checkpoints) cannot be used.")
FRACTAL_PIXEL_SIZE = 8
_ARCH_CONFIG_DEFAULTS = Kandinsky5VideoConfig().arch_config
@@ -263,10 +271,9 @@ class Kandinsky5Modulation(nn.Module):
def _apply_rotary(x: torch.Tensor, rope: torch.Tensor) -> torch.Tensor:
orig_dtype = x.dtype
x_ = x.reshape(*x.shape[:-1], -1, 1, 2).to(torch.float32)
x_out = (rope * x_).sum(dim=-1)
return x_out.reshape(*x.shape).to(orig_dtype)
return x_out.reshape(*x.shape).to(x.dtype)
class Kandinsky5Attention(nn.Module):
@@ -277,6 +284,7 @@ class Kandinsky5Attention(nn.Module):
head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
prefix: str = "",
use_nabla: bool = False,
):
super().__init__()
assert num_channels % head_dim == 0
@@ -306,6 +314,17 @@ class Kandinsky5Attention(nn.Module):
causal=False,
supported_attention_backends=supported_attention_backends,
)
# NABLA checkpoints get a second attention layer whose backend defaults
# to NABLA_ATTN; FASTVIDEO_ATTENTION_BACKEND still overrides it.
self.nabla_attention = None
if use_nabla:
self.nabla_attention = LocalAttention(
num_heads=self.num_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
default_backend=AttentionBackendEnum.NABLA_ATTN,
)
def forward(
self,
@@ -339,27 +358,54 @@ class Kandinsky5Attention(nn.Module):
key = _apply_rotary(key, rotary_emb).type_as(key)
if sparse_params is not None:
raise NotImplementedError(
"Sparse attention is not yet supported for Kandinsky5 in FastVideo."
)
if self.nabla_attention is None:
raise RuntimeError("sparse_params passed to an attention layer built without use_nabla; "
"this checkpoint/config combination is inconsistent.")
try:
# Backend impl reads sta_mask/P from the forward-context
# attention metadata built by the denoising stage.
hidden_states = self.nabla_attention(query, key, value)
except AssertionError as exc:
# Standalone parity tests call the model without a pipeline
# forward context; run the NABLA kernel directly.
if "Forward context is not set" not in str(exc):
raise
attn_mask = nablaT_v2(query, key, sparse_params["sta_mask"], thr=sparse_params["P"])
hidden_states = flex_attention(
query=query.transpose(1, 2),
key=key.transpose(1, 2),
value=value.transpose(1, 2),
block_mask=attn_mask,
).transpose(1, 2)
else:
try:
hidden_states = self.local_attention(query, key, value)
try:
hidden_states = self.local_attention(query, key, value)
except AssertionError as exc:
# LocalAttention requires pipeline forward context. Standalone
# parity tests call the model directly, so fallback to Torch SDPA.
if "Forward context is not set" not in str(exc):
raise
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
hidden_states = F.scaled_dot_product_attention(query,
key,
value,
attn_mask=None,
is_causal=False)
hidden_states = hidden_states.transpose(1, 2)
hidden_states = hidden_states.flatten(2)
except AssertionError as exc:
# LocalAttention requires pipeline forward context. Standalone
# parity tests call the model directly, so fallback to Torch SDPA.
if "Forward context is not set" not in str(exc):
raise
query_shape = query.shape[:-2]
key_shape = key.shape[:-2]
query = query.reshape(query_shape[0], -1, self.num_heads,
query.shape[-1]).transpose(1, 2)
key = key.reshape(key_shape[0], -1, self.num_heads,
key.shape[-1]).transpose(1, 2)
value = value.reshape(key_shape[0], -1, self.num_heads,
value.shape[-1]).transpose(1, 2)
hidden_states = F.scaled_dot_product_attention(
query,
key,
value,
attn_mask=None,
is_causal=False,
)
hidden_states = hidden_states.transpose(1, 2).reshape(
*query_shape, self.num_heads, -1)
hidden_states = hidden_states.flatten(-2, -1)
hidden_states, _ = self.out_layer(hidden_states)
return hidden_states
@@ -476,7 +522,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = ""):
prefix: str = "",
use_nabla: bool = False):
super().__init__()
self.visual_modulation = Kandinsky5Modulation(time_dim, model_dim, 9)
@@ -491,7 +538,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
model_dim,
head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.self_attention")
prefix=f"{prefix}.self_attention",
use_nabla=use_nabla)
self.cross_attention_norm = LayerNormScaleShift(
model_dim,
@@ -624,7 +672,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
arch.ff_dim,
head_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.visual_transformer_blocks.{i}")
prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
use_nabla=arch.attention_type == "nabla")
for i in range(arch.num_visual_blocks)
])
@@ -694,6 +743,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
scale_factor)
to_fractal = sparse_params[
"to_fractal"] if sparse_params is not None else False
visual_embed, visual_rope = fractal_flatten(visual_embed, visual_rope,
visual_shape,
block_mask=to_fractal)
@@ -724,6 +774,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
if return_dict:
return Kandinsky5TransformerOutput(sample=x)
return x
def materialize_non_persistent_buffers(self, device: torch.device,
+44 -2
View File
@@ -467,6 +467,18 @@ class Qwen2_5_VisionTransformerPretrainedModel(nn.Module):
return hidden_states
def _compute_default_rope_parameters(config, device=None, seq_len=None, **kwargs):
# transformers>=5 removes the "default" entry from ROPE_INIT_FUNCTIONS and
# moves rope_theta inside rope_parameters; replicate the 4.x default init.
rope_params = getattr(config, "rope_parameters", None) or getattr(config, "rope_scaling", None) or {}
base = rope_params.get("rope_theta", getattr(config, "rope_theta", 10000.0))
head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0)
dim = int(head_dim * partial_rotary_factor)
inv_freq = 1.0 / (base**(torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim))
return inv_freq, 1.0
class Qwen2_5_VLRotaryEmbedding(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, device=None):
super().__init__()
@@ -479,7 +491,14 @@ class Qwen2_5_VLRotaryEmbedding(nn.Module):
self.original_max_seq_len = config.max_position_embeddings
self.config = config
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
if self.rope_type in ROPE_INIT_FUNCTIONS:
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
elif self.rope_type == "default":
# transformers>=5 drops the "default" entry from ROPE_INIT_FUNCTIONS.
self.rope_init_fn = _compute_default_rope_parameters
else:
raise KeyError(f"Unsupported rope_type '{self.rope_type}'; available: "
f"{['default', *ROPE_INIT_FUNCTIONS]}")
inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
self.register_buffer("inv_freq", inv_freq, persistent=False)
@@ -948,6 +967,16 @@ QWEN2_5_VL_ATTENTION_CLASSES = {
# If FlashAttention2 is not available, transparently fall back to SDPA.
if not is_flash_attn_2_available():
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
else:
# transformers>=5 only resolves the flash-attn functions when the model
# preloads them via its attention interface; this module bypasses
# PreTrainedModel, so _flash_attention_forward(implementation=None) raises
# unless we preload here.
try:
from transformers.modeling_flash_attention_utils import lazy_import_flash_attention
lazy_import_flash_attention("flash_attention_2")
except (ImportError, ValueError):
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
class Qwen2_5_VLDecoderLayer(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int):
@@ -1035,7 +1064,7 @@ class Qwen2_5_VLModel(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig):
super().__init__()
self.config = config
self.padding_idx = config.pad_token_id
self.padding_idx = getattr(config, "pad_token_id", None)
self.vocab_size = config.vocab_size
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
@@ -1408,6 +1437,18 @@ class Qwen2_5_VLCausalLMOutputWithPast(ModelOutput):
rope_deltas: Optional[torch.LongTensor] = None
def _flatten_text_config(config):
# transformers>=5 stops forwarding text-model attributes (hidden_size,
# vocab_size, rope_scaling, ...) from the composite Qwen2_5_VLConfig to
# config.text_config; this module reads them from the top level.
text_config = getattr(config, "text_config", None)
if text_config is not None:
for key, value in text_config.to_dict().items():
if not hasattr(config, key):
setattr(config, key, value)
return config
class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
_tied_weights_keys = ["lm_head.weight"]
config_class = Qwen2_5_VLConfig
@@ -1415,6 +1456,7 @@ class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
def __init__(self, config):
super().__init__()
config = _flatten_text_config(config)
self.config = config
self.visual = Qwen2_5_VisionTransformerPretrainedModel(config.vision_config)
+16 -1
View File
@@ -320,6 +320,19 @@ class TextEncoderLoader(ComponentLoader):
f"text encoder index {idx} out of range for text_encoder_configs (len={len(encoder_configs)}), model_path={model_path}"
)
encoder_config = encoder_configs[idx]
if (
model_config.get("architectures") == ["CLIPModel"]
and isinstance(model_config.get("text_config"), dict)
):
valid_arch_fields = {
f.name for f in dataclasses.fields(encoder_config.arch_config)
}
model_config = {
key: value
for key, value in deepcopy(model_config["text_config"]).items()
if key in valid_arch_fields
}
model_config["architectures"] = ["CLIPTextModel"]
encoder_config.update_model_arch(model_config)
if idx < 0 or idx >= len(encoder_precisions):
raise IndexError(
@@ -404,7 +417,9 @@ class TextEncoderLoader(ComponentLoader):
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
model, model_path, to_cpu=use_cpu_offload
model,
model_path,
to_cpu=fastvideo_args.text_encoder_cpu_offload,
)
) # type: ignore
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,89 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.kandinsky5 import (
Kandinsky5DecodingStage,
Kandinsky5DenoisingStage,
Kandinsky5ImageEncodingStage,
Kandinsky5LatentPreparationStage,
Kandinsky5NormalizationStage,
)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
class Kandinsky5I2VPipeline(ComposedPipelineBase):
"""Kandinsky-5.0 image-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"),
),
)
# Runs AFTER latent preparation so the initial noise is the seeded
# generator's first draw (official kandinsky-5 RNG order); this stage
# then samples the image latent and places it into the latents.
self.add_stage(
stage_name="image_encoding_stage",
stage=Kandinsky5ImageEncodingStage(vae=self.get_module("vae")),
)
self.add_stage(
stage_name="denoising_stage",
stage=Kandinsky5DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
stage_name="normalization_stage",
stage=Kandinsky5NormalizationStage(),
)
self.add_stage(
stage_name="decoding_stage",
stage=Kandinsky5DecodingStage(vae=self.get_module("vae"), pipeline=self),
)
EntryClass = Kandinsky5I2VPipeline
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.kandinsky5 import (
Kandinsky5DecodingStage,
Kandinsky5DenoisingStage,
Kandinsky5LatentPreparationStage,
)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
class Kandinsky5T2VPipeline(ComposedPipelineBase):
"""Kandinsky-5.0 Lite 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=Kandinsky5DenoisingStage(
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 = Kandinsky5T2VPipeline
@@ -0,0 +1,165 @@
# SPDX-License-Identifier: Apache-2.0
"""Kandinsky-5 model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_NEGATIVE_PROMPT = ("Static, 2D cartoon, cartoon, 2d animation, paintings, images, worst quality, low quality, ugly, "
"deformed, walking backwards")
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Main Kandinsky-5 denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
KANDINSKY5_T2V_LITE_5S = InferencePreset(
name="kandinsky5_t2v_lite_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Lite T2V 5s",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 5.0,
"num_inference_steps": 50,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_T2V_LITE_DISTILLED_5S = InferencePreset(
name="kandinsky5_t2v_lite_distilled_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Lite T2V Distilled 5s",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 16,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_T2V_PRO_5S = InferencePreset(
name="kandinsky5_t2v_pro_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Pro T2V 5s",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 5.0,
"num_inference_steps": 50,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_T2V_PRO_DISTILLED_5S = InferencePreset(
name="kandinsky5_t2v_pro_distilled_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Pro T2V Distilled 5s",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 16,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_I2V_LITE_5S = InferencePreset(
name="kandinsky5_i2v_lite_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Lite I2V 5s",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 5.0,
"num_inference_steps": 50,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_I2V_PRO_5S = InferencePreset(
name="kandinsky5_i2v_pro_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Pro I2V 5s",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 5.0,
"num_inference_steps": 50,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_I2V_LITE_DISTILLED_5S = InferencePreset(
name="kandinsky5_i2v_lite_distilled_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Lite I2V Distilled 5s",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 16,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
KANDINSKY5_I2V_PRO_DISTILLED_5S = InferencePreset(
name="kandinsky5_i2v_pro_distilled_5s",
version=1,
model_family="kandinsky5",
description="Kandinsky-5.0 Pro I2V Distilled 5s",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 512,
"width": 768,
"num_frames": 121,
"fps": 24,
"guidance_scale": 1.0,
"num_inference_steps": 16,
"negative_prompt": _NEGATIVE_PROMPT,
},
)
ALL_PRESETS = (KANDINSKY5_T2V_LITE_5S, KANDINSKY5_T2V_LITE_DISTILLED_5S, KANDINSKY5_T2V_PRO_5S,
KANDINSKY5_T2V_PRO_DISTILLED_5S, KANDINSKY5_I2V_LITE_5S, KANDINSKY5_I2V_LITE_DISTILLED_5S,
KANDINSKY5_I2V_PRO_5S, KANDINSKY5_I2V_PRO_DISTILLED_5S)
+5
View File
@@ -33,6 +33,8 @@ from fastvideo.pipelines.basic.ltx2.stages import ( # noqa: F401
from fastvideo.pipelines.stages.matrixgame2_denoising import MatrixGame2CausalDenoisingStage
from fastvideo.pipelines.stages.matrixgame3_denoising import MatrixGame3DenoisingStage
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.kandinsky5 import (Kandinsky5DecodingStage, Kandinsky5DenoisingStage,
Kandinsky5LatentPreparationStage)
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
from fastvideo.pipelines.stages.gen3c_stages import (Gen3CCFGPolicyStage, Gen3CConditioningStage, Gen3CDenoisingStage,
Gen3CLatentPreparationStage)
@@ -65,6 +67,9 @@ __all__ = [
"MatrixGame2CausalDenoisingStage",
"MatrixGame3DenoisingStage",
"HYWorldDenoisingStage",
"Kandinsky5DecodingStage",
"Kandinsky5DenoisingStage",
"Kandinsky5LatentPreparationStage",
"GameCraftDenoisingStage",
"Gen3CCFGPolicyStage",
"Gen3CConditioningStage",
+545
View File
@@ -0,0 +1,545 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import contextlib
from typing import Any
import PIL
import torch
from diffusers.utils.torch_utils import randn_tensor
from tqdm.auto import tqdm
from fastvideo.attention.backends.nabla import NablaAttentionMetadataBuilder
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.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
from fastvideo.pipelines.stages.base import PipelineStage
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
logger = init_logger(__name__)
class Kandinsky5LatentPreparationStage(PipelineStage):
def __init__(self, scheduler, transformer) -> None:
super().__init__()
self.scheduler = scheduler
self.transformer = transformer
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.height is None or batch.width is None:
raise ValueError("height and width must be provided for Kandinsky5.")
height = int(batch.height)
width = int(batch.width)
num_frames = int(batch.num_frames)
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
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
if num_frames % temporal_ratio != 1:
num_frames = num_frames // temporal_ratio * temporal_ratio + 1
batch.num_frames = num_frames
required_divisor_h = spatial_ratio * patch_size[1]
required_divisor_w = spatial_ratio * patch_size[2]
if height % required_divisor_h != 0 or width % required_divisor_w != 0:
raise ValueError(f"Kandinsky5 height must be divisible by {required_divisor_h} and width by "
f"{required_divisor_w}; "
f"got height={height}, width={width}.")
# NABLA sparse attention (Pro checkpoints) reshapes the post-patch grid
# into 8x8 blocks; validate here instead of crashing mid-denoise after
# all the encoding work is done.
arch_cfg = getattr(self.transformer, "config", None) or fastvideo_args.pipeline_config.dit_config.arch_config
if getattr(arch_cfg, "attention_type", "regular") == "nabla":
nabla_divisor_h = required_divisor_h * 8
nabla_divisor_w = required_divisor_w * 8
if height % nabla_divisor_h != 0 or width % nabla_divisor_w != 0:
raise ValueError(f"Kandinsky5 NABLA checkpoints require height divisible by {nabla_divisor_h} and "
f"width divisible by {nabla_divisor_w}; "
f"got height={height}, width={width}.")
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
batch_size = batch.prompt_embeds[0].shape[0]
batch_size *= batch.num_videos_per_prompt
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
device = get_local_torch_device()
num_latent_frames = (num_frames - 1) // temporal_ratio + 1
num_channels = getattr(
self.transformer,
"in_visual_dim",
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
)
shape = (
batch_size,
num_latent_frames,
height // spatial_ratio,
width // spatial_ratio,
num_channels,
)
if isinstance(batch.generator, list) and len(batch.generator) != batch_size:
raise ValueError(f"generator list length {len(batch.generator)} does not match batch size {batch_size}.")
visual_cond = getattr(self.transformer, "visual_cond", False)
if batch.latents is None:
latents = randn_tensor(shape, generator=batch.generator, device=device, dtype=dtype)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
else:
valid_shapes = [shape]
if visual_cond:
valid_shapes.append((*shape[:-1], num_channels * 2 + 1))
if tuple(batch.latents.shape) not in valid_shapes:
raise ValueError(f"Provided latents shape {list(batch.latents.shape)} does not match expected "
f"Kandinsky5 latent shape(s): {[list(s) for s in valid_shapes]}.")
latents = batch.latents.to(device=device, dtype=dtype)
if visual_cond and latents.shape[-1] == num_channels:
cond = torch.zeros_like(latents)
cond_mask = torch.zeros(
(*latents.shape[:-1], 1),
device=latents.device,
dtype=latents.dtype,
)
latents = torch.cat([latents, cond, cond_mask], dim=-1)
# I2V image conditioning is placed by Kandinsky5ImageEncodingStage,
# which runs AFTER this stage so the initial noise is the generator's
# first draw (matching the official kandinskylab/kandinsky-5 order).
batch.latents = latents
batch.raw_latent_shape = (
batch_size,
num_channels,
num_latent_frames,
height // spatial_ratio,
width // spatial_ratio,
)
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("num_frames", batch.num_frames, V.positive_int)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("latents", batch.latents, V.none_or_tensor)
return result
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
return result
class Kandinsky5DenoisingStage(PipelineStage):
def __init__(self, transformer, scheduler) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
@staticmethod
def _scale_factor(height: int, width: int) -> tuple[float, float, float]:
if 480 <= height <= 854 and 480 <= width <= 854:
return (1.0, 2.0, 2.0)
return (1.0, 3.16, 3.16)
@staticmethod
def _text_rope_pos(mask: torch.Tensor, device: torch.device) -> torch.Tensor:
seq_len = int(mask.sum(1).max().item())
return torch.arange(seq_len, device=device)
@staticmethod
def fast_sta_nabla(
T: int,
H: int,
W: int,
wT: int = 3,
wH: int = 3,
wW: int = 3,
device: torch.device | str = "cuda",
) -> torch.Tensor:
"""
Create a sparse temporal attention (STA) mask for efficient video generation.
This method generates a mask that limits attention to nearby frames and spatial positions, reducing
computational complexity for video generation.
Args:
T (int): Number of temporal frames
H (int): Height in latent space
W (int): Width in latent space
wT (int): Temporal attention window size
wH (int): Height attention window size
wW (int): Width attention window size
device (str): Device to create tensor on
Returns:
torch.Tensor: Sparse attention mask of shape (T*H*W, T*H*W)
"""
max_extent = int(torch.tensor([T, H, W], device=device).amax().item())
r = torch.arange(0, max_extent, 1, dtype=torch.int16, device=device)
mat = (r.unsqueeze(1) - r.unsqueeze(0)).abs()
sta_t, sta_h, sta_w = (
mat[:T, :T].flatten(),
mat[:H, :H].flatten(),
mat[:W, :W].flatten(),
)
sta_t = sta_t <= wT // 2
sta_h = sta_h <= wH // 2
sta_w = sta_w <= wW // 2
sta_hw = (sta_h.unsqueeze(1) * sta_w.unsqueeze(0)).reshape(H, H, W, W).transpose(1, 2).flatten()
sta = (sta_t.unsqueeze(1) * sta_hw.unsqueeze(0)).reshape(T, T, H * W, H * W).transpose(1, 2)
return sta.reshape(T * H * W, T * H * W)
def get_sparse_params(self, sample: torch.Tensor, device: torch.device) -> dict[str, Any] | None:
"""
Generate sparse attention parameters for the transformer based on sample dimensions.
This method computes the sparse attention configuration needed for efficient video processing in the
transformer model.
Args:
sample (torch.Tensor): Input sample tensor
device (torch.device): Device to place tensors on
Returns:
Dict: Dictionary containing sparse attention parameters
"""
assert self.transformer.config.patch_size[0] == 1
_, T, H, W, _ = sample.shape
T, H, W = (
T // self.transformer.config.patch_size[0],
H // self.transformer.config.patch_size[1],
W // self.transformer.config.patch_size[2],
)
if self.transformer.config.attention_type == "nabla":
sta_mask = self.fast_sta_nabla(
T,
H // 8,
W // 8,
self.transformer.config.attention_wT,
self.transformer.config.attention_wH,
self.transformer.config.attention_wW,
device=device,
)
sparse_params = {
"sta_mask": sta_mask.unsqueeze_(0).unsqueeze_(0),
"attention_type": self.transformer.config.attention_type,
"to_fractal": True,
"P": self.transformer.config.attention_P,
"wT": self.transformer.config.attention_wT,
"wW": self.transformer.config.attention_wW,
"wH": self.transformer.config.attention_wH,
"add_sta": self.transformer.config.attention_add_sta,
"visual_shape": (T, H, W),
"method": self.transformer.config.attention_method,
}
else:
sparse_params = None
return sparse_params
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.timesteps is None:
raise ValueError("timesteps must be prepared before Kandinsky5 denoising.")
if batch.latents is None:
raise ValueError("latents must be prepared before Kandinsky5 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
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
autocast_enabled = target_dtype != torch.float32 and not fastvideo_args.disable_autocast
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 requires Qwen prompt attention masks.")
text_rope_pos = self._text_rope_pos(batch.prompt_attention_mask[0].to(device), device)
neg_prompt_embeds = None
neg_pooled = None
negative_text_rope_pos = None
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds:
neg_prompt_embeds = batch.negative_prompt_embeds[0].to(device=device, dtype=target_dtype)
neg_pooled = batch.negative_prompt_embeds[1].to(device=device, dtype=target_dtype)
if batch.negative_attention_mask is None or not batch.negative_attention_mask:
raise ValueError("Kandinsky5 requires Qwen negative attention masks for CFG.")
negative_text_rope_pos = self._text_rope_pos(batch.negative_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)
# 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.
cond_frames = 1 if batch.image_latent is not None else 0
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
with tqdm(total=batch.num_inference_steps, desc="Kandinsky5 Denoising") as progress_bar:
for i, timestep in enumerate(batch.timesteps):
if hasattr(self, "interrupt") and self.interrupt:
break
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
if neg_prompt_embeds is not None and neg_pooled is not None:
uncond_pred_velocity = self.transformer(
hidden_states=latents.to(dtype=target_dtype),
encoder_hidden_states=neg_prompt_embeds,
pooled_projections=neg_pooled,
timestep=t_expand,
visual_rope_pos=visual_rope_pos,
text_rope_pos=negative_text_rope_pos,
scale_factor=scale_factor,
sparse_params=sparse_params,
return_dict=True,
).sample
pred_velocity = uncond_pred_velocity + batch.guidance_scale * (pred_velocity -
uncond_pred_velocity)
latents[:, cond_frames:, :, :, :num_channels] = self.scheduler.step(
pred_velocity[:, cond_frames:],
timestep,
latents[:, cond_frames:, :, :, :num_channels],
return_dict=False,
)[0]
if batch.return_trajectory_latents:
trajectory_timesteps.append(timestep)
# latents is mutated in place, so snapshot a channels-first copy.
trajectory_latents.append(latents[..., :num_channels].permute(0, 4, 1, 2, 3).cpu())
if i == len(batch.timesteps) - 1 or (i + 1) % self.scheduler.order == 0:
progress_bar.update()
if trajectory_latents:
batch.trajectory_latents = torch.stack(trajectory_latents, dim=1)
batch.trajectory_timesteps = torch.stack(trajectory_timesteps, dim=0).cpu()
batch.latents = latents[:, :, :, :, :num_channels]
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.min_list_length(2))
return result
class Kandinsky5DecodingStage(DecodingStage):
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
super().__init__(vae=vae, pipeline=pipeline)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None:
raise ValueError("latents must be available before Kandinsky5 decoding.")
# Kandinsky5 latents are channels-last [B, T, H, W, C]; the base stage
# (and the trajectory latents recorded by the denoising stage) work
# channels-first.
batch.latents = batch.latents.permute(0, 4, 1, 2, 3).contiguous()
return super().forward(batch, fastvideo_args)
class Kandinsky5ImageEncodingStage(EncodingStage):
"""Encode the conditioning image into a VAE latent for I2V."""
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
super().__init__(vae=vae)
@staticmethod
def _preprocess(image, height: int, width: int) -> torch.Tensor:
if isinstance(image, PIL.Image.Image):
image = resize(image, height, width)
image = numpy_to_pt(pil_to_numpy(image)) # always lands in [0, 1]
return normalize(image) # [0, 1] -> [-1, 1]
# Tensor input: no reliable way to tell [0, 1] from an already
# normalized [-1, 1] tensor whose values happen to be non-negative,
# so mirror diffusers' heuristic and say what we assumed.
if image.min() >= 0:
logger.warning("Kandinsky5 conditioning image tensor has no negative values; "
"assuming range [0, 1] and normalizing to [-1, 1]. "
"Pass a [-1, 1] tensor with negative values to skip normalization.")
image = normalize(image) # [0, 1] -> [-1, 1]
if image.ndim == 3:
image = image.unsqueeze(0)
if image.shape[-2:] != (height, width):
image = torch.nn.functional.interpolate(image.float(),
size=(height, width),
mode="bilinear",
antialias=True)
return image
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.pil_image is None:
raise ValueError("Kandinsky5 I2V requires an input image.")
if not fastvideo_args.model_loaded["vae"]:
vae = getattr(self, "vae", None)
if vae is None:
loader = VAELoader()
vae = loader.load(fastvideo_args.model_paths["vae"], fastvideo_args)
self.vae = vae
fastvideo_args.model_loaded["vae"] = True
device = get_local_torch_device()
vae = self.vae.to(device)
self.vae = vae
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = vae_dtype != torch.float32 and not fastvideo_args.disable_autocast
# [B, C, H, W] -> [B, C, 1, H, W]
image = self._preprocess(batch.pil_image, int(batch.height), int(batch.width))
image = image.to(device=device, dtype=torch.float32).unsqueeze(2)
# Encode the single conditioning frame without tiling (matches diffusers).
# The untested causal-VAE spatial_tiled_encode path corrupts the latent.
prev_use_tiling = vae.use_tiling
vae.use_tiling = False
try:
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if not vae_autocast_enabled:
image = image.to(vae_dtype)
# Sample with the batch generator (diffusers parity); mode()
# would make seed-for-seed reproduction of the reference
# pipeline impossible.
generator = batch.generator
if isinstance(generator, list) and len(generator) != image.shape[0]:
generator = generator[0]
image_latent = vae.encode(image).sample(generator=generator)
finally:
vae.use_tiling = prev_use_tiling
image_latent = image_latent * vae.scaling_factor
# [B, C, 1, H, W] -> [B, 1, H, W, C] to match channels-last latents
batch.image_latent = image_latent.permute(0, 2, 3, 4, 1).contiguous()
# Place the conditioning latent into the prepared latents: frame 0 of
# the main channels, the visual_cond channel block, and the mask.
# NOTE: the official kandinsky-5 repo leaves the visual_cond block
# zeros (generation_utils.py generate()), while the diffusers port
# copies the image latent into it. A same-seed A/B on
# Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers showed the Diffusers
# export requires the copy: zeroing the block produces smeared faces
# mid-video. Keep the diffusers semantics for Diffusers-format
# checkpoints.
latents = batch.latents
image_latent = batch.image_latent.to(device=latents.device, dtype=latents.dtype)
num_channels = image_latent.shape[-1]
latents[:, 0:1, :, :, :num_channels] = image_latent
if latents.shape[-1] > num_channels:
latents[:, 0:1, :, :, num_channels:2 * num_channels] = image_latent
latents[:, 0:1, :, :, 2 * num_channels:] = 1.0
batch.latents = latents
if fastvideo_args.vae_cpu_offload:
vae.to("cpu")
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("pil_image", batch.pil_image, V.not_none)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
# This stage runs after latent preparation and writes into its output.
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
return result
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)])
return result
class Kandinsky5NormalizationStage(PipelineStage):
"""Normalize the first latent frames to reduce I2V conditioning artifacts."""
COND_FRAMES = 4
REFERENCE_FRAMES = 5
@staticmethod
def _adaptive_mean_std(source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor:
source_mean = source.mean(dim=(1, 2, 3, 4), keepdim=True)
source_std = source.std(dim=(1, 2, 3, 4), keepdim=True)
# Magic constants limit how far the first frames may drift.
ref_mean = torch.clamp(reference.mean(dim=(1, 2, 3, 4), keepdim=True), source_mean - 0.05, source_mean + 0.1)
ref_std = torch.clamp(reference.std(dim=(1, 2, 3, 4), keepdim=True), source_std - 0.1, source_std + 0.25)
normalized = (source - source_mean) / source_std
return normalized * ref_std + ref_mean
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
latents = batch.latents
n = self.COND_FRAMES
if latents is None or latents.shape[1] <= n:
return batch
reference = latents[:, n:n + min(self.REFERENCE_FRAMES, latents.shape[1] - 1)]
latents[:, :n] = self._adaptive_mean_std(latents[:, :n].clone(), reference)
batch.latents = latents
return batch
+25 -5
View File
@@ -8,6 +8,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
import torch
from typing import Any
from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
@@ -203,6 +205,21 @@ class TextEncodingStage(PipelineStage):
encoder_config = encoder_cfgs[i]
preprocess_func = preprocess_funcs[i]
postprocess_func = postprocess_funcs[i]
# cpu_offload semantics: params rest on CPU between calls but the
# forward computes on GPU. FSDP2-wrapped encoders (CPUOffloadPolicy;
# DTensor params) stream themselves per-layer — leave inputs on the
# param device and let FSDP's root pre-forward move them. A plain
# module parked on CPU by text_encoder_cpu_offload is swapped to the
# target device for the forward and back afterwards, mirroring the
# image-encoder/VAE offload pattern.
first_param = next(text_encoder.parameters(), None)
encoder_device = first_param.device if first_param is not None else torch.device(target_device)
moved_for_forward = False
if (first_param is not None and not isinstance(first_param, DTensor)
and encoder_device.type != torch.device(target_device).type):
text_encoder = text_encoder.to(target_device)
encoder_device = torch.device(target_device)
moved_for_forward = True
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
if max_length is not None:
@@ -247,7 +264,7 @@ class TextEncodingStage(PipelineStage):
# pre-format prompts into message lists upstream and rely on
# the inner tokenizer + full tokenizer_kwargs (which include
# add_generation_prompt). Preserve that original path exactly.
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(target_device)
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(encoder_device)
else:
# Two-step approach matching Diffusers: format with chat
# template first, then tokenize the resulting strings.
@@ -261,9 +278,9 @@ class TextEncodingStage(PipelineStage):
enable_thinking=False,
)
formatted_texts.append(formatted)
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(target_device)
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(encoder_device)
else:
text_inputs = tok(processed_texts, **tok_kwargs).to(target_device)
text_inputs = tok(processed_texts, **tok_kwargs).to(encoder_device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
@@ -283,15 +300,18 @@ class TextEncodingStage(PipelineStage):
except Exception:
prompt_embeds, attention_mask = postprocess_func(outputs, attention_mask)
if is_ltx2 and getattr(outputs, "hidden_states", None):
audio_embed = outputs.hidden_states[0]
audio_embed = outputs.hidden_states[0].to(device=target_device)
if dtype is not None:
audio_embed = audio_embed.to(dtype=dtype)
audio_embeds_list.append(audio_embed)
prompt_embeds = prompt_embeds.to(device=target_device)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
embeds_list.append(prompt_embeds)
if return_attention_mask:
attn_masks_list.append(attention_mask)
attn_masks_list.append(attention_mask.to(device=target_device))
if moved_for_forward and fastvideo_args.text_encoder_cpu_offload:
text_encoder.to("cpu")
self._last_audio_embeds = audio_embeds_list if is_ltx2 else None
return self.return_embeds(embeds_list, attn_masks_list, return_type, return_attention_mask, indices)
+10
View File
@@ -219,6 +219,16 @@ class StageValidators:
return validator
@staticmethod
def min_list_length(min_length: int) -> Callable[[Any], bool]:
"""Return a validator that checks if a list has at least min_length items."""
def validator(value: Any) -> bool:
return StageValidators.list_min_length(value, min_length)
validator.__name__ = f"list_min_length_{min_length}"
return validator
@staticmethod
def min_dims(min_dims: int) -> Callable[[Any], bool]:
"""Return a validator that checks if tensor has at least min_dims dimensions and no NaN values."""
+8
View File
@@ -157,6 +157,14 @@ class CudaPlatformBase(Platform):
"ATTN_QAT_TRAIN selected but fastvideo_kernel.triton_kernels.attn_qat_train is not built. "
"Silent fallback would produce a non-QAT training run; refusing to proceed. "
"Install the training kernel or pick a different FASTVIDEO_ATTENTION_BACKEND.")
elif selected_backend == AttentionBackendEnum.NABLA_ATTN:
from fastvideo.attention.backends.nabla import CAN_USE_FLEX_ATTN
if CAN_USE_FLEX_ATTN:
logger.info("Using NABLA block-sparse flex-attention backend.")
return "fastvideo.attention.backends.nabla.NablaAttentionBackend"
raise ImportError("NABLA_ATTN selected but torch.nn.attention.flex_attention is unavailable in this "
"PyTorch build. Silent fallback to dense attention would be orders of magnitude "
"slower and diverge from the reference; upgrade PyTorch or pick a different backend.")
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
try:
from fastvideo_kernel import video_sparse_attn # noqa: F401
+1
View File
@@ -22,6 +22,7 @@ class AttentionBackendEnum(enum.Enum):
VMOBA_ATTN = enum.auto()
SLA_ATTN = enum.auto()
SAGE_SLA_ATTN = enum.auto()
NABLA_ATTN = enum.auto()
NO_ATTENTION = enum.auto()
+179 -4
View File
@@ -27,6 +27,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyu
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig,
Hunyuan15SR1080PConfig)
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
@@ -117,6 +118,10 @@ class ConfigInfo:
workload_types: tuple[WorkloadType, ...]
model_family: str | None = None
default_preset: str | None = None
# When set, overrides the model_index `_class_name` for pipeline resolution.
# Lets a model family map to a specific pipeline class by path/detector
# (e.g. a T2V and I2V checkpoint that share a `_class_name`).
pipeline_cls_name: str | None = None
# The central registry mapping a model name to its configuration information
@@ -137,6 +142,7 @@ def register_configs(
model_detectors: list[Callable[[str], bool]] | None = None,
model_family: str | None = None,
default_preset: str | None = None,
pipeline_cls_name: str | None = None,
) -> None:
"""Register config classes for a model family.
@@ -151,6 +157,7 @@ def register_configs(
workload_types=workload_types,
model_family=model_family,
default_preset=default_preset,
pipeline_cls_name=pipeline_cls_name,
)
if hf_model_paths:
@@ -490,18 +497,174 @@ def _register_configs() -> None:
default_preset="lingbotworld_i2v",
)
def _kandinsky5_detector(require: tuple[str, ...] = (), exclude: tuple[str, ...] = ()) -> Callable[[str], bool]:
def detect(path: str) -> bool:
path_lower = path.lower()
if "kandinsky5" not in path_lower and "kandinsky-5" not in path_lower:
return False
return (all(token in path_lower for token in require) and not any(token in path_lower for token in exclude))
return detect
# t2v/i2v exclude each other so a checkpoint stored under a directory
# containing the other token (e.g. ~/i2v_experiments/kandinsky5-t2v-ft)
# falls through to the model_index _class_name fallback detectors below
# instead of being misrouted.
_is_kandinsky5_t2v = _kandinsky5_detector(require=("t2v", ), exclude=("i2v", ))
_is_kandinsky5_i2v = _kandinsky5_detector(require=("i2v", ), exclude=("t2v", ))
_is_kandinsky5_t2v_lite = _kandinsky5_detector(require=("t2v", "lite"), exclude=("i2v", "distilled"))
_is_kandinsky5_t2v_pro = _kandinsky5_detector(require=("t2v", "pro"), exclude=("i2v", "distilled"))
_is_kandinsky5_t2v_lite_distilled = _kandinsky5_detector(require=("t2v", "lite", "distilled"), exclude=("i2v", ))
_is_kandinsky5_t2v_pro_distilled = _kandinsky5_detector(require=("t2v", "pro", "distilled"), exclude=("i2v", ))
_is_kandinsky5_i2v_lite = _kandinsky5_detector(require=("i2v", "lite"), exclude=("t2v", "distilled"))
_is_kandinsky5_i2v_pro = _kandinsky5_detector(require=("i2v", "pro"), exclude=("t2v", "distilled"))
_is_kandinsky5_i2v_lite_distilled = _kandinsky5_detector(require=("i2v", "lite", "distilled"), exclude=("t2v", ))
_is_kandinsky5_i2v_pro_distilled = _kandinsky5_detector(require=("i2v", "pro", "distilled"), exclude=("t2v", ))
# Kandinsky5 Lite T2V
register_configs(
sampling_param_cls=None,
pipeline_config_cls=PipelineConfig,
pipeline_config_cls=Kandinsky5T2VConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
],
model_detectors=[
lambda path: any(token in path.lower() for token in ("kandinsky5", "kandinsky-5")),
_is_kandinsky5_t2v_lite,
],
model_family="kandinsky5",
default_preset="kandinsky5_t2v_lite_5s",
pipeline_cls_name="Kandinsky5T2VPipeline",
)
# Kandinsky5 Pro T2V
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5T2VConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers",
],
model_detectors=[
_is_kandinsky5_t2v_pro,
],
model_family="kandinsky5",
default_preset="kandinsky5_t2v_pro_5s",
)
# Kandinsky5 Lite T2V Distilled
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5T2VConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers",
],
model_detectors=[
_is_kandinsky5_t2v_lite_distilled,
],
model_family="kandinsky5",
default_preset="kandinsky5_t2v_lite_distilled_5s",
)
# Kandinsky5 Pro T2V Distilled
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5T2VConfig,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers",
],
model_detectors=[
_is_kandinsky5_t2v_pro_distilled,
],
model_family="kandinsky5",
default_preset="kandinsky5_t2v_pro_distilled_5s",
)
# Kandinsky5 Lite I2V
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5I2VConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"],
model_detectors=[
_is_kandinsky5_i2v_lite,
],
model_family="kandinsky5",
default_preset="kandinsky5_i2v_lite_5s",
pipeline_cls_name="Kandinsky5I2VPipeline",
)
# Kandinsky5 Pro I2V
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5I2VConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"],
model_detectors=[
_is_kandinsky5_i2v_pro,
],
model_family="kandinsky5",
default_preset="kandinsky5_i2v_pro_5s",
)
# Kandinsky5 Pro I2V Distilled
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5I2VConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers"],
model_detectors=[
_is_kandinsky5_i2v_pro_distilled,
],
model_family="kandinsky5",
default_preset="kandinsky5_i2v_pro_distilled_5s",
)
# Kandinsky5 Lite I2V Distilled (no official hub repo yet; local
# conversions get distilled sampling defaults instead of the sft ones the
# fallback would apply).
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5I2VConfig,
workload_types=(),
model_detectors=[
_is_kandinsky5_i2v_lite_distilled,
],
model_family="kandinsky5",
default_preset="kandinsky5_i2v_lite_distilled_5s",
pipeline_cls_name="Kandinsky5I2VPipeline",
)
# Kandinsky5 fallbacks — registered AFTER the variant detectors so those
# win first-match. Catch checkpoints the variant detectors cannot resolve:
# token-less local paths matched via the model_index _class_name
# ("kandinsky5t2vpipeline" carries no lite/pro marker), variant combos
# without a dedicated entry (e.g. I2V Lite distilled), and t2v+i2v
# ambiguous paths resolved by the checkpoint's _class_name.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5T2VConfig,
workload_types=(),
model_detectors=[
_is_kandinsky5_t2v,
],
model_family="kandinsky5",
default_preset="kandinsky5_t2v_lite_5s",
pipeline_cls_name="Kandinsky5T2VPipeline",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Kandinsky5I2VConfig,
workload_types=(),
model_detectors=[
_is_kandinsky5_i2v,
],
model_family="kandinsky5",
default_preset="kandinsky5_i2v_lite_5s",
pipeline_cls_name="Kandinsky5I2VPipeline",
)
# LongCat (T2V, I2V, VC use same config; workload varies by path)
@@ -550,7 +713,7 @@ def _register_configs() -> None:
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
# Legacy HF paths (kept for backward compat — pre-rename names):
# Legacy HF paths (kept for backward compat - pre-rename names):
"FastVideo/Matrix-Game-2.0-Base-Diffusers",
"FastVideo/Matrix-Game-2.0-GTA-Diffusers",
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
@@ -902,7 +1065,10 @@ def get_model_info(
assert config_info is not None, "config_info must be resolved"
if override_pipeline_cls_name:
pipeline_name = override_pipeline_cls_name
# Explicit override: skip config resolution entirely so checkpoints
# without a diffusers model_index.json keep working (and no download
# is triggered just to log the replaced name).
pipeline_name: str | None = override_pipeline_cls_name
logger.info("Using override pipeline class name %s", pipeline_name)
else:
if os.path.exists(model_path):
@@ -911,6 +1077,12 @@ def get_model_info(
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name")
if config_info.pipeline_cls_name is not None:
# The resolved (path/detector-based) config pins the pipeline class,
# e.g. an I2V checkpoint whose `_class_name` would otherwise resolve
# to the T2V pipeline.
logger.info("Pinning pipeline class name from %s to %s", pipeline_name, config_info.pipeline_cls_name)
pipeline_name = config_info.pipeline_cls_name
if pipeline_name is None:
raise ValueError("Model config does not contain a _class_name attribute. "
@@ -961,6 +1133,8 @@ def _register_presets() -> None:
ALL_PRESETS as HUNYUAN15_PRESETS, )
from fastvideo.pipelines.basic.hyworld.presets import (
ALL_PRESETS as HYWORLD_PRESETS, )
from fastvideo.pipelines.basic.kandinsky5.presets import (
ALL_PRESETS as KANDINSKY5_PRESETS, )
from fastvideo.pipelines.basic.lingbotworld.presets import (
ALL_PRESETS as LINGBOTWORLD_PRESETS, )
from fastvideo.pipelines.basic.longcat.presets import (
@@ -990,6 +1164,7 @@ def _register_presets() -> None:
HUNYUAN_PRESETS,
HUNYUAN15_PRESETS,
HYWORLD_PRESETS,
KANDINSKY5_PRESETS,
LINGBOTWORLD_PRESETS,
LONGCAT_PRESETS,
LTX2_PRESETS,
@@ -0,0 +1,42 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only regression test for the SDPA metadata mask contract.
SDPAMetadata is cross-backend: HYWorld and HunyuanVideo15 build it while the
layer's attention selector may pick FLASH_ATTN, whose impl requires the
tokenizer-style 2D [batch, key_len] padding mask (flash_attn.py pads
``attn_mask.shape[1]`` assuming 2D). The builder must therefore store 2D
masks unchanged; only the torch-sdpa impl lifts them to 4D internally.
Regression guard for the CI break where the builder lifted 2D -> 4D and the
FLASH_ATTN consumer produced a garbage-length pad (CUDA device-side assert).
"""
from __future__ import annotations
import torch
from fastvideo.attention.backends.sdpa import (SDPAMetadataBuilder, _normalize_attn_mask_for_sdpa)
def test_builder_keeps_2d_padding_mask_2d() -> None:
"""A tokenizer-style 2D int mask must be stored as-is (cross-backend contract)."""
mask = torch.tensor([[1, 1, 0], [1, 0, 0]], dtype=torch.int64)
md = SDPAMetadataBuilder().build(current_timestep=0, attn_mask=mask)
assert md.attn_mask is not None
assert md.attn_mask.dim() == 2
assert md.attn_mask.shape == (2, 3)
assert torch.equal(md.attn_mask, mask)
def test_normalize_lifts_2d_int_mask_to_4d_bool_for_torch_sdpa() -> None:
"""The torch-sdpa consumer coerces int -> bool and lifts 2D -> [B,1,1,K]."""
batch, q_len, k_len, heads, head_dim = 2, 4, 3, 1, 8
query = torch.zeros(batch, heads, q_len, head_dim)
key = torch.zeros(batch, heads, k_len, head_dim)
mask = torch.tensor([[1, 1, 0], [1, 0, 0]], dtype=torch.int64)
out = _normalize_attn_mask_for_sdpa(mask, query, key)
assert out is not None
assert out.dtype == torch.bool
assert out.shape == (batch, 1, 1, k_len)
assert torch.equal(out[:, 0, 0, :], mask != 0)
+6 -3
View File
@@ -12,6 +12,7 @@ from fastvideo.configs.pipelines import CosmosConfig, PipelineConfig, WanT2V480P
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.tests.utils import skip_if_gated_repo_inaccessible
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.models.encoders import T5Config, T5LargeConfig
@@ -36,9 +37,11 @@ def t5_model_paths_and_config():
@pytest.fixture
def t5_large_model_paths_and_config():
base_model_path = "nvidia/Cosmos-Predict2-2B-Video2World"
model_path = maybe_download_model(base_model_path,
local_dir=os.path.join(
'data', base_model_path))
local_dir = os.path.join('data', base_model_path)
skip_if_gated_repo_inaccessible(base_model_path,
local_path=local_dir,
test_name="Cosmos T5-large encoder test")
model_path = maybe_download_model(base_model_path, local_dir=local_dir)
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path, CosmosConfig()
+1 -1
View File
@@ -276,7 +276,7 @@ def run_self_forcing_tests():
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test(
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ ./fastvideo/tests/attention/test_sdpa_metadata_mask_contract.py --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
)
@@ -0,0 +1,109 @@
# SPDX-License-Identifier: Apache-2.0
"""SSIM regression test for Kandinsky-5.0 Lite text-to-video.
Runs a reduced-size generation (256x384, 21 frames) with a fixed seed and
compares against a device-specific reference video via MS-SSIM. The preset
default is 512x768 x 121 frames; the reduced size keeps CI cheap while the
patch (1, 2, 2) and VAE temporal-4 / spatial-8 constraints stay satisfied
(dims divisible by 16, num_frames % 4 == 1).
"""
import os
import pytest
from fastvideo.logger import init_logger
from fastvideo.pipelines.basic.kandinsky5.presets import KANDINSKY5_T2V_LITE_5S
from fastvideo.tests.ssim.inference_similarity_utils import (
resolve_inference_device_reference_folder,
run_text_to_video_similarity_test,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 1
device_reference_folder = resolve_inference_device_reference_folder(logger)
_PRESET_DEFAULTS = KANDINSKY5_T2V_LITE_5S.defaults
KANDINSKY5_T2V_PARAMS = {
"num_gpus": 1,
"model_path": "kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
"height": 256,
"width": 384,
"num_frames": 21,
"num_inference_steps": _PRESET_DEFAULTS["num_inference_steps"],
"guidance_scale": _PRESET_DEFAULTS["guidance_scale"],
"seed": 1024,
"sp_size": 1,
"tp_size": 1,
"fps": _PRESET_DEFAULTS["fps"],
"neg_prompt": _PRESET_DEFAULTS["negative_prompt"],
}
KANDINSKY5_T2V_FULL_QUALITY_PARAMS = {
**KANDINSKY5_T2V_PARAMS,
"height": _PRESET_DEFAULTS["height"],
"width": _PRESET_DEFAULTS["width"],
# num_frames stays at the inherited reduced 21 even at full quality
# (preset default: 121).
}
KANDINSKY5_T2V_MODEL_TO_PARAMS = {
"Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers": KANDINSKY5_T2V_PARAMS,
}
FULL_QUALITY_KANDINSKY5_T2V_MODEL_TO_PARAMS = {
"Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers": KANDINSKY5_T2V_FULL_QUALITY_PARAMS,
}
KANDINSKY5_T2V_TEST_PROMPTS = [
"A curious raccoon peers through a vibrant field of yellow sunflowers, its "
"eyes wide with interest. The playful yet serene atmosphere is complemented "
"by soft natural light filtering through the petals. Mid-shot, warm and "
"cheerful tones.",
]
@pytest.mark.parametrize("prompt", KANDINSKY5_T2V_TEST_PROMPTS)
@pytest.mark.parametrize("attention_backend_name", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(KANDINSKY5_T2V_MODEL_TO_PARAMS.keys()))
def test_kandinsky5_t2v_inference_similarity(
prompt: str,
attention_backend_name: str,
model_id: str,
) -> None:
# The SSIM lane image bakes FASTVIDEO_FA4=1, but on L40S (sm89) the CLIP
# encoder's dense MHA inference call routes into the FA4 CuTeDSL JIT,
# which fails to compile there (nvvm.fmax signature mismatch vs the
# image's nvidia_cutlass_dsl). Pin FA4 off so reference seeding and CI
# runs both use the FA2 path with identical numerics. Scoped to this test
# and restored afterwards: the other SSIM references are seeded with FA4
# on, so a module-level pin would corrupt every test collected in the
# same pytest process.
saved_fa4 = os.environ.get("FASTVIDEO_FA4")
os.environ["FASTVIDEO_FA4"] = "0"
try:
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=KANDINSKY5_T2V_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_KANDINSKY5_T2V_MODEL_TO_PARAMS,
min_acceptable_ssim=0.93,
# Match examples/inference/basic/basic_kandinsky5_t2v.py: FSDP
# inference is not used for Kandinsky-5, and the Qwen2.5-VL text
# encoder stays on CPU between encodes.
init_kwargs_override={
"use_fsdp_inference": False,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": True,
},
)
finally:
if saved_fa4 is None:
os.environ.pop("FASTVIDEO_FA4", None)
else:
os.environ["FASTVIDEO_FA4"] = saved_fa4
+7 -12
View File
@@ -11,6 +11,7 @@ from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.tests.utils import skip_if_gated_repo_inaccessible
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import CosmosVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
@@ -25,18 +26,12 @@ BASE_MODEL_PATH = "nvidia/Cosmos-Predict2-2B-Video2World"
def _resolve_model_path() -> str:
try:
return maybe_download_model(
BASE_MODEL_PATH,
local_dir=os.path.join("data", BASE_MODEL_PATH),
)
except ValueError as exc:
pytest.skip(
"Skipping Cosmos transformer test because the configured "
"HuggingFace token cannot access the gated Cosmos weights: "
f"{exc}",
allow_module_level=True,
)
local_dir = os.path.join("data", BASE_MODEL_PATH)
skip_if_gated_repo_inaccessible(BASE_MODEL_PATH,
local_path=local_dir,
test_name="Cosmos transformer test",
allow_module_level=True)
return maybe_download_model(BASE_MODEL_PATH, local_dir=local_dir)
MODEL_PATH = _resolve_model_path()
@@ -19,6 +19,7 @@ if os.path.exists(COSMOS_PREDICT2_5_PATH) and COSMOS_PREDICT2_5_PATH not in sys.
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.tests.utils import skip_if_gated_repo_inaccessible
from fastvideo.utils import maybe_download_model
# Use Cosmos 2.5 specific config
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
@@ -43,7 +44,14 @@ CHECKPOINT_SUBDIR = "base/post-trained"
CHECKPOINT_FILENAME = "81edfebe-bd6a-4039-8c1d-737df1a790bf_ema_bf16.pt"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH, local_dir=None)
def _resolve_model_path() -> str:
skip_if_gated_repo_inaccessible(BASE_MODEL_PATH,
test_name="Cosmos 2.5 transformer test",
allow_module_level=True)
return maybe_download_model(BASE_MODEL_PATH, local_dir=None)
MODEL_PATH = _resolve_model_path()
TRANSFORMER_PATH = os.path.join(MODEL_PATH, CHECKPOINT_SUBDIR, "transformer")
if not os.path.exists(TRANSFORMER_PATH):
# Try without subdirectory
@@ -592,4 +600,3 @@ if __name__ == "__main__":
# Run tests directly
test_cosmos25_transformer()
test_cosmos25_transformer_video()
+21 -4
View File
@@ -15,6 +15,7 @@ from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -24,13 +25,25 @@ os.environ["MASTER_PORT"] = "29503"
MODEL_PATH = maybe_download_model("FastVideo/HY-WorldPlay-Bidirectional-Diffusers")
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
REFERENCE_LATENT = -197132.85557549074 # Pre-computed reference value
REFERENCE_LATENTS = {
AttentionBackendEnum.FLASH_ATTN: -197132.85557549074,
AttentionBackendEnum.TORCH_SDPA: -211007.95158730447,
}
@pytest.mark.usefixtures("distributed_setup")
def test_hyworld_transformer():
transformer_path = TRANSFORMER_PATH
# An env-forced backend is knowable before the multi-GB download and GPU
# load; the authoritative check on the layer's resolved backend still runs
# after construction below.
forced_backend_name = os.environ.get("FASTVIDEO_ATTENTION_BACKEND")
if forced_backend_name:
forced_backend = getattr(AttentionBackendEnum, forced_backend_name, None)
if forced_backend is not None and forced_backend not in REFERENCE_LATENTS:
pytest.skip(f"HYWorld transformer test has no reference latent for {forced_backend_name}.")
sp_rank = get_sp_parallel_rank()
sp_world_size = get_sp_world_size()
@@ -51,6 +64,9 @@ def test_hyworld_transformer():
loader = TransformerLoader()
model = loader.load(transformer_path, args).to(device, dtype=precision)
model.eval()
attention_backend = model.double_blocks[0].attn.backend
if attention_backend not in REFERENCE_LATENTS:
pytest.skip(f"HYWorld transformer test has no reference latent for {attention_backend}.")
total_params = sum(p.numel() for p in model.parameters())
weight_sum = sum(p.to(torch.float64).sum().item() for p in model.parameters())
@@ -147,9 +163,10 @@ def test_hyworld_transformer():
latent = output.double().sum().item()
diff = abs(REFERENCE_LATENT - latent)
relative_diff = diff / abs(REFERENCE_LATENT)
logger.info(f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}")
reference_latent = REFERENCE_LATENTS[attention_backend]
diff = abs(reference_latent - latent)
relative_diff = diff / abs(reference_latent)
logger.info(f"Reference latent: {reference_latent}, Current latent: {latent}")
logger.info(f"Absolute diff: {diff}, Relative diff: {relative_diff * 100:.4f}%")
# Allow 0.5% relative difference
+35
View File
@@ -11,6 +11,41 @@ from pytorch_msssim import ms_ssim, ssim
logger = init_logger(__name__)
def skip_if_gated_repo_inaccessible(repo_id: str,
*,
local_path: str | None = None,
test_name: str = "test",
allow_module_level: bool = False) -> None:
"""Skip a test when ``repo_id`` is gated/private and the token lacks access.
Local weights win: when ``local_path`` already exists the check is skipped
entirely, so cached machines keep running offline. Transient hub failures
(offline, DNS, 429/5xx) do NOT skip either — the subsequent download will
serve from cache or fail loudly, preserving the CI failure signal. Only a
positively-identified authorization problem produces a skip.
"""
import pytest
if local_path is not None and os.path.exists(local_path):
return
from huggingface_hub import HfApi
from huggingface_hub.errors import GatedRepoError, RepositoryNotFoundError
try:
HfApi().auth_check(repo_id)
except (GatedRepoError, RepositoryNotFoundError) as exc:
pytest.skip(
f"Skipping {test_name}: the configured HuggingFace token cannot "
f"access the gated repo {repo_id}: {exc}",
allow_module_level=allow_module_level,
)
except Exception as exc: # noqa: BLE001 - hub unreachable, proxies, 5xx
logger.warning(
"Gated-repo access probe for %s failed (%s); proceeding — the "
"download will serve from cache or fail loudly.", repo_id, exc)
def _read_video_frames(path: str) -> torch.Tensor:
"""Read video frames as a ``(T, C, H, W)`` uint8 tensor.