Compare commits
37
Commits
main
...
sdpa-mask-2d-fix
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a288460497 | ||
|
|
45fc714f88 | ||
|
|
38d1d7bdc5 | ||
|
|
3a44ae6316 | ||
|
|
6918418525 | ||
|
|
881400490c | ||
|
|
fca6fb3aa1 | ||
|
|
fa6e50fc7c | ||
|
|
9f28d34d6d | ||
|
|
a7470d2da3 | ||
|
|
821b4bc275 | ||
|
|
09c42cc73a | ||
|
|
d6ed0e54da | ||
|
|
6ce4993ab7 | ||
|
|
b9dfb12502 | ||
|
|
375ba410c7 | ||
|
|
a43284ae1b | ||
|
|
6ab2d6f98d | ||
|
|
9472ed0aab | ||
|
|
4eb4e252eb | ||
|
|
0434e9dcf0 | ||
|
|
36cf64b15f | ||
|
|
7a270e66be | ||
|
|
74de6f27a2 | ||
|
|
6979040225 | ||
|
|
80972d07fa | ||
|
|
260b326d8b | ||
|
|
c57fa1eab0 | ||
|
|
562b951208 | ||
|
|
aa893056e8 | ||
|
|
313c2985f0 | ||
|
|
b3b3858c02 | ||
|
|
d39e2e9fd7 | ||
|
|
a66fb26781 | ||
|
|
f303d94780 | ||
|
|
dae7c0da89 | ||
|
|
0ce6dc928f |
@@ -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()
|
||||
@@ -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)
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
from torch.nn import functional as F
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -19,7 +20,7 @@ class SDPABackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SDPA"
|
||||
return "TORCH_SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
@@ -49,9 +50,50 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
# Store the mask exactly as passed. The metadata is cross-backend:
|
||||
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
|
||||
# layer's selector may pick FLASH_ATTN, whose impl requires the
|
||||
# tokenizer-style 2D [batch, key_len] padding mask. Any reshaping for
|
||||
# torch.sdpa happens inside the SDPA impl (_normalize_attn_mask_for_sdpa).
|
||||
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
|
||||
|
||||
def _normalize_attn_mask_for_sdpa(
|
||||
attn_mask: torch.Tensor | None,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> torch.Tensor | None:
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
attn_mask = attn_mask.to(device=query.device)
|
||||
# F.scaled_dot_product_attention only accepts bool or float masks;
|
||||
# tokenizers commonly produce int64 0/1 padding masks.
|
||||
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
|
||||
attn_mask = attn_mask != 0
|
||||
|
||||
key_len = key.shape[-2]
|
||||
if attn_mask.shape[-1] > key_len:
|
||||
raise ValueError("Invalid attention mask length for SDPA: "
|
||||
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
|
||||
if attn_mask.shape[-1] < key_len:
|
||||
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
|
||||
# non-text tokens the tokenizer mask does not cover.
|
||||
valid_value = True if attn_mask.dtype == torch.bool else 0.0
|
||||
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
|
||||
|
||||
if attn_mask.dim() == 2:
|
||||
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
|
||||
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
|
||||
# reinterpret 2D as its documented [query_len, key_len] broadcast.
|
||||
return attn_mask[:, None, None, :]
|
||||
if attn_mask.dim() == 3:
|
||||
return attn_mask[:, None, :, :]
|
||||
if attn_mask.dim() == 4:
|
||||
return attn_mask
|
||||
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -82,6 +124,7 @@ class SDPAImpl(AttentionImpl):
|
||||
|
||||
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -27,6 +27,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyu
|
||||
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig,
|
||||
Hunyuan15SR1080PConfig)
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -117,6 +118,10 @@ class ConfigInfo:
|
||||
workload_types: tuple[WorkloadType, ...]
|
||||
model_family: str | None = None
|
||||
default_preset: str | None = None
|
||||
# When set, overrides the model_index `_class_name` for pipeline resolution.
|
||||
# Lets a model family map to a specific pipeline class by path/detector
|
||||
# (e.g. a T2V and I2V checkpoint that share a `_class_name`).
|
||||
pipeline_cls_name: str | None = None
|
||||
|
||||
|
||||
# The central registry mapping a model name to its configuration information
|
||||
@@ -137,6 +142,7 @@ def register_configs(
|
||||
model_detectors: list[Callable[[str], bool]] | None = None,
|
||||
model_family: str | None = None,
|
||||
default_preset: str | None = None,
|
||||
pipeline_cls_name: str | None = None,
|
||||
) -> None:
|
||||
"""Register config classes for a model family.
|
||||
|
||||
@@ -151,6 +157,7 @@ def register_configs(
|
||||
workload_types=workload_types,
|
||||
model_family=model_family,
|
||||
default_preset=default_preset,
|
||||
pipeline_cls_name=pipeline_cls_name,
|
||||
)
|
||||
|
||||
if hf_model_paths:
|
||||
@@ -490,18 +497,174 @@ def _register_configs() -> None:
|
||||
default_preset="lingbotworld_i2v",
|
||||
)
|
||||
|
||||
def _kandinsky5_detector(require: tuple[str, ...] = (), exclude: tuple[str, ...] = ()) -> Callable[[str], bool]:
|
||||
|
||||
def detect(path: str) -> bool:
|
||||
path_lower = path.lower()
|
||||
if "kandinsky5" not in path_lower and "kandinsky-5" not in path_lower:
|
||||
return False
|
||||
return (all(token in path_lower for token in require) and not any(token in path_lower for token in exclude))
|
||||
|
||||
return detect
|
||||
|
||||
# t2v/i2v exclude each other so a checkpoint stored under a directory
|
||||
# containing the other token (e.g. ~/i2v_experiments/kandinsky5-t2v-ft)
|
||||
# falls through to the model_index _class_name fallback detectors below
|
||||
# instead of being misrouted.
|
||||
_is_kandinsky5_t2v = _kandinsky5_detector(require=("t2v", ), exclude=("i2v", ))
|
||||
_is_kandinsky5_i2v = _kandinsky5_detector(require=("i2v", ), exclude=("t2v", ))
|
||||
_is_kandinsky5_t2v_lite = _kandinsky5_detector(require=("t2v", "lite"), exclude=("i2v", "distilled"))
|
||||
_is_kandinsky5_t2v_pro = _kandinsky5_detector(require=("t2v", "pro"), exclude=("i2v", "distilled"))
|
||||
_is_kandinsky5_t2v_lite_distilled = _kandinsky5_detector(require=("t2v", "lite", "distilled"), exclude=("i2v", ))
|
||||
_is_kandinsky5_t2v_pro_distilled = _kandinsky5_detector(require=("t2v", "pro", "distilled"), exclude=("i2v", ))
|
||||
_is_kandinsky5_i2v_lite = _kandinsky5_detector(require=("i2v", "lite"), exclude=("t2v", "distilled"))
|
||||
_is_kandinsky5_i2v_pro = _kandinsky5_detector(require=("i2v", "pro"), exclude=("t2v", "distilled"))
|
||||
_is_kandinsky5_i2v_lite_distilled = _kandinsky5_detector(require=("i2v", "lite", "distilled"), exclude=("t2v", ))
|
||||
_is_kandinsky5_i2v_pro_distilled = _kandinsky5_detector(require=("i2v", "pro", "distilled"), exclude=("t2v", ))
|
||||
|
||||
# Kandinsky5 Lite T2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=PipelineConfig,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: any(token in path.lower() for token in ("kandinsky5", "kandinsky-5")),
|
||||
_is_kandinsky5_t2v_lite,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5T2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro T2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_pro,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_pro_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite T2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_lite_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro T2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_pro_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_pro_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite I2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_lite,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro I2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_pro,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_pro_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro I2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_pro_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_pro_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite I2V Distilled (no official hub repo yet; local
|
||||
# conversions get distilled sampling defaults instead of the sft ones the
|
||||
# fallback would apply).
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_lite_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_distilled_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 fallbacks — registered AFTER the variant detectors so those
|
||||
# win first-match. Catch checkpoints the variant detectors cannot resolve:
|
||||
# token-less local paths matched via the model_index _class_name
|
||||
# ("kandinsky5t2vpipeline" carries no lite/pro marker), variant combos
|
||||
# without a dedicated entry (e.g. I2V Lite distilled), and t2v+i2v
|
||||
# ambiguous paths resolved by the checkpoint's _class_name.
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5T2VPipeline",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# LongCat (T2V, I2V, VC use same config; workload varies by path)
|
||||
@@ -550,7 +713,7 @@ def _register_configs() -> None:
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
# Legacy HF paths (kept for backward compat — pre-rename names):
|
||||
# Legacy HF paths (kept for backward compat - pre-rename names):
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
@@ -902,7 +1065,10 @@ def get_model_info(
|
||||
assert config_info is not None, "config_info must be resolved"
|
||||
|
||||
if override_pipeline_cls_name:
|
||||
pipeline_name = override_pipeline_cls_name
|
||||
# Explicit override: skip config resolution entirely so checkpoints
|
||||
# without a diffusers model_index.json keep working (and no download
|
||||
# is triggered just to log the replaced name).
|
||||
pipeline_name: str | None = override_pipeline_cls_name
|
||||
logger.info("Using override pipeline class name %s", pipeline_name)
|
||||
else:
|
||||
if os.path.exists(model_path):
|
||||
@@ -911,6 +1077,12 @@ def get_model_info(
|
||||
config = maybe_download_model_index(model_path)
|
||||
|
||||
pipeline_name = config.get("_class_name")
|
||||
if config_info.pipeline_cls_name is not None:
|
||||
# The resolved (path/detector-based) config pins the pipeline class,
|
||||
# e.g. an I2V checkpoint whose `_class_name` would otherwise resolve
|
||||
# to the T2V pipeline.
|
||||
logger.info("Pinning pipeline class name from %s to %s", pipeline_name, config_info.pipeline_cls_name)
|
||||
pipeline_name = config_info.pipeline_cls_name
|
||||
|
||||
if pipeline_name is None:
|
||||
raise ValueError("Model config does not contain a _class_name attribute. "
|
||||
@@ -961,6 +1133,8 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as HUNYUAN15_PRESETS, )
|
||||
from fastvideo.pipelines.basic.hyworld.presets import (
|
||||
ALL_PRESETS as HYWORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.kandinsky5.presets import (
|
||||
ALL_PRESETS as KANDINSKY5_PRESETS, )
|
||||
from fastvideo.pipelines.basic.lingbotworld.presets import (
|
||||
ALL_PRESETS as LINGBOTWORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.longcat.presets import (
|
||||
@@ -990,6 +1164,7 @@ def _register_presets() -> None:
|
||||
HUNYUAN_PRESETS,
|
||||
HUNYUAN15_PRESETS,
|
||||
HYWORLD_PRESETS,
|
||||
KANDINSKY5_PRESETS,
|
||||
LINGBOTWORLD_PRESETS,
|
||||
LONGCAT_PRESETS,
|
||||
LTX2_PRESETS,
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only regression test for the SDPA metadata mask contract.
|
||||
|
||||
SDPAMetadata is cross-backend: HYWorld and HunyuanVideo15 build it while the
|
||||
layer's attention selector may pick FLASH_ATTN, whose impl requires the
|
||||
tokenizer-style 2D [batch, key_len] padding mask (flash_attn.py pads
|
||||
``attn_mask.shape[1]`` assuming 2D). The builder must therefore store 2D
|
||||
masks unchanged; only the torch-sdpa impl lifts them to 4D internally.
|
||||
|
||||
Regression guard for the CI break where the builder lifted 2D -> 4D and the
|
||||
FLASH_ATTN consumer produced a garbage-length pad (CUDA device-side assert).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.attention.backends.sdpa import (SDPAMetadataBuilder, _normalize_attn_mask_for_sdpa)
|
||||
|
||||
|
||||
def test_builder_keeps_2d_padding_mask_2d() -> None:
|
||||
"""A tokenizer-style 2D int mask must be stored as-is (cross-backend contract)."""
|
||||
mask = torch.tensor([[1, 1, 0], [1, 0, 0]], dtype=torch.int64)
|
||||
md = SDPAMetadataBuilder().build(current_timestep=0, attn_mask=mask)
|
||||
assert md.attn_mask is not None
|
||||
assert md.attn_mask.dim() == 2
|
||||
assert md.attn_mask.shape == (2, 3)
|
||||
assert torch.equal(md.attn_mask, mask)
|
||||
|
||||
|
||||
def test_normalize_lifts_2d_int_mask_to_4d_bool_for_torch_sdpa() -> None:
|
||||
"""The torch-sdpa consumer coerces int -> bool and lifts 2D -> [B,1,1,K]."""
|
||||
batch, q_len, k_len, heads, head_dim = 2, 4, 3, 1, 8
|
||||
query = torch.zeros(batch, heads, q_len, head_dim)
|
||||
key = torch.zeros(batch, heads, k_len, head_dim)
|
||||
mask = torch.tensor([[1, 1, 0], [1, 0, 0]], dtype=torch.int64)
|
||||
|
||||
out = _normalize_attn_mask_for_sdpa(mask, query, key)
|
||||
assert out is not None
|
||||
assert out.dtype == torch.bool
|
||||
assert out.shape == (batch, 1, 1, k_len)
|
||||
assert torch.equal(out[:, 0, 0, :], mask != 0)
|
||||
@@ -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()
|
||||
|
||||
@@ -276,7 +276,7 @@ def run_self_forcing_tests():
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test(
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ ./fastvideo/tests/attention/test_sdpa_metadata_mask_contract.py --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""SSIM regression test for Kandinsky-5.0 Lite text-to-video.
|
||||
|
||||
Runs a reduced-size generation (256x384, 21 frames) with a fixed seed and
|
||||
compares against a device-specific reference video via MS-SSIM. The preset
|
||||
default is 512x768 x 121 frames; the reduced size keeps CI cheap while the
|
||||
patch (1, 2, 2) and VAE temporal-4 / spatial-8 constraints stay satisfied
|
||||
(dims divisible by 16, num_frames % 4 == 1).
|
||||
"""
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.kandinsky5.presets import KANDINSKY5_T2V_LITE_5S
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
device_reference_folder = resolve_inference_device_reference_folder(logger)
|
||||
|
||||
_PRESET_DEFAULTS = KANDINSKY5_T2V_LITE_5S.defaults
|
||||
|
||||
KANDINSKY5_T2V_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
"height": 256,
|
||||
"width": 384,
|
||||
"num_frames": 21,
|
||||
"num_inference_steps": _PRESET_DEFAULTS["num_inference_steps"],
|
||||
"guidance_scale": _PRESET_DEFAULTS["guidance_scale"],
|
||||
"seed": 1024,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"fps": _PRESET_DEFAULTS["fps"],
|
||||
"neg_prompt": _PRESET_DEFAULTS["negative_prompt"],
|
||||
}
|
||||
|
||||
KANDINSKY5_T2V_FULL_QUALITY_PARAMS = {
|
||||
**KANDINSKY5_T2V_PARAMS,
|
||||
"height": _PRESET_DEFAULTS["height"],
|
||||
"width": _PRESET_DEFAULTS["width"],
|
||||
# num_frames stays at the inherited reduced 21 even at full quality
|
||||
# (preset default: 121).
|
||||
}
|
||||
|
||||
KANDINSKY5_T2V_MODEL_TO_PARAMS = {
|
||||
"Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers": KANDINSKY5_T2V_PARAMS,
|
||||
}
|
||||
FULL_QUALITY_KANDINSKY5_T2V_MODEL_TO_PARAMS = {
|
||||
"Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers": KANDINSKY5_T2V_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
KANDINSKY5_T2V_TEST_PROMPTS = [
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its "
|
||||
"eyes wide with interest. The playful yet serene atmosphere is complemented "
|
||||
"by soft natural light filtering through the petals. Mid-shot, warm and "
|
||||
"cheerful tones.",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", KANDINSKY5_T2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["FLASH_ATTN"])
|
||||
@pytest.mark.parametrize("model_id", list(KANDINSKY5_T2V_MODEL_TO_PARAMS.keys()))
|
||||
def test_kandinsky5_t2v_inference_similarity(
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
# The SSIM lane image bakes FASTVIDEO_FA4=1, but on L40S (sm89) the CLIP
|
||||
# encoder's dense MHA inference call routes into the FA4 CuTeDSL JIT,
|
||||
# which fails to compile there (nvvm.fmax signature mismatch vs the
|
||||
# image's nvidia_cutlass_dsl). Pin FA4 off so reference seeding and CI
|
||||
# runs both use the FA2 path with identical numerics. Scoped to this test
|
||||
# and restored afterwards: the other SSIM references are seeded with FA4
|
||||
# on, so a module-level pin would corrupt every test collected in the
|
||||
# same pytest process.
|
||||
saved_fa4 = os.environ.get("FASTVIDEO_FA4")
|
||||
os.environ["FASTVIDEO_FA4"] = "0"
|
||||
try:
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=KANDINSKY5_T2V_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_KANDINSKY5_T2V_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.93,
|
||||
# Match examples/inference/basic/basic_kandinsky5_t2v.py: FSDP
|
||||
# inference is not used for Kandinsky-5, and the Qwen2.5-VL text
|
||||
# encoder stays on CPU between encodes.
|
||||
init_kwargs_override={
|
||||
"use_fsdp_inference": False,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"pin_cpu_memory": True,
|
||||
},
|
||||
)
|
||||
finally:
|
||||
if saved_fa4 is None:
|
||||
os.environ.pop("FASTVIDEO_FA4", None)
|
||||
else:
|
||||
os.environ["FASTVIDEO_FA4"] = saved_fa4
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -11,6 +11,41 @@ from pytorch_msssim import ms_ssim, ssim
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def skip_if_gated_repo_inaccessible(repo_id: str,
|
||||
*,
|
||||
local_path: str | None = None,
|
||||
test_name: str = "test",
|
||||
allow_module_level: bool = False) -> None:
|
||||
"""Skip a test when ``repo_id`` is gated/private and the token lacks access.
|
||||
|
||||
Local weights win: when ``local_path`` already exists the check is skipped
|
||||
entirely, so cached machines keep running offline. Transient hub failures
|
||||
(offline, DNS, 429/5xx) do NOT skip either — the subsequent download will
|
||||
serve from cache or fail loudly, preserving the CI failure signal. Only a
|
||||
positively-identified authorization problem produces a skip.
|
||||
"""
|
||||
import pytest
|
||||
|
||||
if local_path is not None and os.path.exists(local_path):
|
||||
return
|
||||
|
||||
from huggingface_hub import HfApi
|
||||
from huggingface_hub.errors import GatedRepoError, RepositoryNotFoundError
|
||||
|
||||
try:
|
||||
HfApi().auth_check(repo_id)
|
||||
except (GatedRepoError, RepositoryNotFoundError) as exc:
|
||||
pytest.skip(
|
||||
f"Skipping {test_name}: the configured HuggingFace token cannot "
|
||||
f"access the gated repo {repo_id}: {exc}",
|
||||
allow_module_level=allow_module_level,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - hub unreachable, proxies, 5xx
|
||||
logger.warning(
|
||||
"Gated-repo access probe for %s failed (%s); proceeding — the "
|
||||
"download will serve from cache or fail loudly.", repo_id, exc)
|
||||
|
||||
|
||||
def _read_video_frames(path: str) -> torch.Tensor:
|
||||
"""Read video frames as a ``(T, C, H, W)`` uint8 tensor.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user