[feat] Add Kandinsky-5 T2V/I2V pipeline support (#1471)

Co-authored-by: Aryan Kumar <aryan5v@users.noreply.github.com>
Co-authored-by: leffff <levnovitskiy@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
This commit is contained in:
Aryan Kumar
2026-07-07 14:43:30 -07:00
committed by GitHub
co-authored by Aryan Kumar leffff SolitaryThinker
parent e2f4d1a7b5
commit 02e1143f22
35 changed files with 1910 additions and 82 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()
+5 -1
View File
@@ -338,7 +338,11 @@ class FlashAttentionImpl(AttentionImpl):
)
qkv = torch.stack([query, key, value], dim=2)
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
if key_padding_mask.shape[-1] > qkv.shape[1]:
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
elif self.nvfp4_fa4:
output = self._forward_nvfp4(query, key, value)
+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)
+45 -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,51 @@ 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, and the shared convention for
# padding masks is the tokenizer-style 2D [batch, key_len]. 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 +125,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 -2
View File
@@ -6,6 +6,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
@@ -17,6 +18,6 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "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
@@ -28,6 +28,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
@@ -118,6 +119,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
@@ -138,6 +143,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.
@@ -152,6 +158,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:
@@ -491,18 +498,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)
@@ -551,7 +714,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",
@@ -937,7 +1100,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):
@@ -946,6 +1112,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. "
@@ -998,6 +1170,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 (
@@ -1028,6 +1202,7 @@ def _register_presets() -> None:
HUNYUAN_PRESETS,
HUNYUAN15_PRESETS,
HYWORLD_PRESETS,
KANDINSKY5_PRESETS,
LINGBOTWORLD_PRESETS,
LONGCAT_PRESETS,
LTX2_PRESETS,
@@ -0,0 +1,43 @@
# 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, and the shared convention
for padding masks is the tokenizer-style 2D [batch, key_len]. The builder
must 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 equal-length branch padded ``attn_mask.shape[1]`` (== 1 for 4D)
into a garbage-length mask (CUDA device-side assert, SIGABRT).
"""
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
@@ -280,7 +280,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/ ./fastvideo/tests/ops/ ./fastvideo/tests/training/test_trackers.py --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/ops/ ./fastvideo/tests/training/test_trackers.py ./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
+37 -1
View File
@@ -6,11 +6,45 @@ from fastvideo.logger import init_logger
import numpy as np
import torch
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.
@@ -52,6 +86,8 @@ def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
video2_path: Path to the second video.
use_ms_ssim: Whether to use Multi-Scale Structural Similarity(MS-SSIM) instead of SSIM.
"""
from pytorch_msssim import ms_ssim, ssim
print(f"Computing SSIM between {video1_path} and {video2_path}...")
if not os.path.exists(video1_path):
raise FileNotFoundError(f"Video1 not found: {video1_path}")