Compare commits

...
29 Commits
Author SHA1 Message Date
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
23 changed files with 1613 additions and 49 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)
+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"
]
+116
View File
@@ -0,0 +1,116 @@
# 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).")
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
+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",
+534
View File
@@ -0,0 +1,534 @@
# 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 = spatial_ratio * patch_size[1]
if height % required_divisor != 0 or width % required_divisor != 0:
raise ValueError(f"Kandinsky5 height/width must be divisible by {required_divisor}; "
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 = required_divisor * 8
if height % nabla_divisor != 0 or width % nabla_divisor != 0:
raise ValueError(f"Kandinsky5 NABLA checkpoints require height/width divisible by {nabla_divisor}; "
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}.")
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:
latents = batch.latents.to(device=device, dtype=dtype)
visual_cond = getattr(self.transformer, "visual_cond", False)
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:
continue
t_expand = timestep.unsqueeze(0).repeat(latents.shape[0]).to(device=device, dtype=target_dtype)
attn_metadata = None
if sparse_params is not None:
attn_metadata = NablaAttentionMetadataBuilder().build(
current_timestep=i,
sta_mask=sparse_params["sta_mask"],
P=sparse_params["P"],
visual_shape=sparse_params["visual_shape"],
)
autocast_ctx = (torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled)
if device.type == "cuda" else contextlib.nullcontext())
with set_forward_context(current_timestep=i, attn_metadata=attn_metadata,
forward_batch=batch), autocast_ctx:
pred_velocity = self.transformer(
hidden_states=latents.to(dtype=target_dtype),
encoder_hidden_states=prompt_embeds,
pooled_projections=pooled,
timestep=t_expand,
visual_rope_pos=visual_rope_pos,
text_rope_pos=text_rope_pos,
scale_factor=scale_factor,
sparse_params=sparse_params,
return_dict=True,
).sample
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.list_not_empty)
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)
+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,
+22 -4
View File
@@ -106,12 +106,30 @@ def run_test_command(test_command: str,
if pr_number:
print(f"PR number: {pr_number}")
# For PRs (including forks), use GitHub's PR refs to get the correct commit
# The blob-less clone (--filter=blob:none) defers all file-content
# downloads to the checkout, so BOTH paths below perform a large lazy blob
# fetch from GitHub. Retry transient GitHub/HTTP2 disconnects on each;
# otherwise Modal shards can fail before pytest starts.
def with_retries(inner_command: str) -> str:
return f"""
for attempt in 1 2 3; do
{inner_command} &&
break
status=$?
if [ "$attempt" -eq 3 ]; then
exit "$status"
fi
sleep $((attempt * 5))
done"""
# For PRs (including forks), use GitHub's PR refs to get the correct commit.
if pr_number and pr_number != "false":
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
checkout_command = with_retries(f"git fetch --prune --no-tags --depth=1 origin refs/pull/{pr_number}/head &&"
"\n git checkout --detach FETCH_HEAD")
print(f"Using PR ref for checkout: {checkout_command}")
else:
checkout_command = f"git checkout {git_commit}"
checkout_command = with_retries(f"git checkout {git_commit}")
print(f"Using direct commit checkout: {checkout_command}")
build_kernel_command = """
@@ -125,7 +143,7 @@ def run_test_command(test_command: str,
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
git clone --filter=blob:none --no-checkout {git_repo} /FastVideo &&
cd /FastVideo &&
{checkout_command} &&
git submodule update --init --recursive &&