Compare commits
29
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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)
|
||||
@@ -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,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, CLIPTextConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1Config
|
||||
from fastvideo.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
# Byte-exact copy of the upstream Kandinsky5/diffusers template, including the
|
||||
# "promt"/"scren" typos: the checkpoints were trained with this exact system # codespell:ignore promt,scren
|
||||
# prompt, and ENCODE_START_IDX below is the tokenized length of everything
|
||||
# before the user prompt. Fixing the typos shifts user content to index 127
|
||||
# and mis-conditions every generation.
|
||||
KANDINSKY5_PROMPT_TEMPLATE = "\n".join([
|
||||
"<|im_start|>system\nYou are a promt engineer. Describe the video in detail.", # codespell:ignore promt
|
||||
"Describe how the camera moves or shakes, describe the zoom and view angle, whether it follows the objects.",
|
||||
"Describe the location of the video, main characters or objects and their action.",
|
||||
"Describe the dynamism of the video and presented actions.",
|
||||
"Name the visual style of the video: whether it is a professional footage, user generated content, some kind of animation, video game or scren content.", # codespell:ignore scren
|
||||
"Describe the visual effects, postprocessing and transitions if they are presented in the video.",
|
||||
"Pay attention to the order of key actions shown in the scene.<|im_end|>",
|
||||
"<|im_start|>user\n{}<|im_end|>",
|
||||
])
|
||||
KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX = 129
|
||||
|
||||
|
||||
def kandinsky5_qwen_preprocess_text(prompt: str) -> str:
|
||||
if not prompt.strip():
|
||||
prompt = "."
|
||||
return KANDINSKY5_PROMPT_TEMPLATE.format(prompt)
|
||||
|
||||
|
||||
def kandinsky5_qwen_postprocess_text(outputs: BaseEncoderOutput,
|
||||
mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if outputs.hidden_states is None:
|
||||
raise RuntimeError("Kandinsky5 Qwen prompt embeddings require hidden_states.")
|
||||
hidden_states = outputs.hidden_states[-1]
|
||||
prompt_embeds = hidden_states[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
mask = mask[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
if prompt_embeds.shape[1] == 0:
|
||||
prompt_embeds = hidden_states[:, -1:]
|
||||
mask = torch.ones((mask.shape[0], 1), dtype=mask.dtype, device=mask.device)
|
||||
return prompt_embeds, mask
|
||||
|
||||
|
||||
def kandinsky5_clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError("Kandinsky5 CLIP pooled output is required.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5T2VConfig(PipelineConfig):
|
||||
"""Kandinsky-5.0 Lite text-to-video pipeline configuration."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Kandinsky5VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Reason1Config(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_postprocess_text, kandinsky5_clip_postprocess_text))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX + 512, 77))
|
||||
|
||||
flow_shift: float | None = 5.0
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if len(self.text_encoder_configs) != 2:
|
||||
raise ValueError(f"Kandinsky5 pipeline requires exactly 2 text encoders (qwen and clip), "
|
||||
f"but got {len(self.text_encoder_configs)} encoder(s).")
|
||||
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
qwen_cfg = self.text_encoder_configs[0]
|
||||
qwen_cfg.arch_config.output_hidden_states = True
|
||||
qwen_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": True,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
clip_cfg = self.text_encoder_configs[1]
|
||||
clip_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": "max_length",
|
||||
"max_length": 77,
|
||||
"truncation": True,
|
||||
"add_special_tokens": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
|
||||
"""Kandinsky-5.0 image-to-video pipeline configuration."""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# I2V needs the VAE encoder to encode the conditioning image.
|
||||
self.vae_config.load_encoder = True
|
||||
@@ -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,534 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from typing import Any
|
||||
|
||||
import PIL
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.attention.backends.nabla import NablaAttentionMetadataBuilder
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader, VAELoader
|
||||
from fastvideo.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.models.vision_utils import normalize, numpy_to_pt, pil_to_numpy, resize
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Kandinsky5LatentPreparationStage(PipelineStage):
|
||||
|
||||
def __init__(self, scheduler, transformer) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("height and width must be provided for Kandinsky5.")
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
num_frames = int(batch.num_frames)
|
||||
|
||||
temporal_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
|
||||
if num_frames % temporal_ratio != 1:
|
||||
num_frames = num_frames // temporal_ratio * temporal_ratio + 1
|
||||
batch.num_frames = num_frames
|
||||
|
||||
required_divisor = spatial_ratio * patch_size[1]
|
||||
if height % required_divisor != 0 or width % required_divisor != 0:
|
||||
raise ValueError(f"Kandinsky5 height/width must be divisible by {required_divisor}; "
|
||||
f"got height={height}, width={width}.")
|
||||
|
||||
# NABLA sparse attention (Pro checkpoints) reshapes the post-patch grid
|
||||
# into 8x8 blocks; validate here instead of crashing mid-denoise after
|
||||
# all the encoding work is done.
|
||||
arch_cfg = getattr(self.transformer, "config", None) or fastvideo_args.pipeline_config.dit_config.arch_config
|
||||
if getattr(arch_cfg, "attention_type", "regular") == "nabla":
|
||||
nabla_divisor = required_divisor * 8
|
||||
if height % nabla_divisor != 0 or width % nabla_divisor != 0:
|
||||
raise ValueError(f"Kandinsky5 NABLA checkpoints require height/width divisible by {nabla_divisor}; "
|
||||
f"got height={height}, width={width}.")
|
||||
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
device = get_local_torch_device()
|
||||
num_latent_frames = (num_frames - 1) // temporal_ratio + 1
|
||||
num_channels = getattr(
|
||||
self.transformer,
|
||||
"in_visual_dim",
|
||||
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
|
||||
)
|
||||
shape = (
|
||||
batch_size,
|
||||
num_latent_frames,
|
||||
height // spatial_ratio,
|
||||
width // spatial_ratio,
|
||||
num_channels,
|
||||
)
|
||||
|
||||
if isinstance(batch.generator, list) and len(batch.generator) != batch_size:
|
||||
raise ValueError(f"generator list length {len(batch.generator)} does not match batch size {batch_size}.")
|
||||
|
||||
if batch.latents is None:
|
||||
latents = randn_tensor(shape, generator=batch.generator, device=device, dtype=dtype)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
latents = batch.latents.to(device=device, dtype=dtype)
|
||||
|
||||
visual_cond = getattr(self.transformer, "visual_cond", False)
|
||||
if visual_cond and latents.shape[-1] == num_channels:
|
||||
cond = torch.zeros_like(latents)
|
||||
cond_mask = torch.zeros(
|
||||
(*latents.shape[:-1], 1),
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
latents = torch.cat([latents, cond, cond_mask], dim=-1)
|
||||
|
||||
# I2V image conditioning is placed by Kandinsky5ImageEncodingStage,
|
||||
# which runs AFTER this stage so the initial noise is the generator's
|
||||
# first draw (matching the official kandinskylab/kandinsky-5 order).
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = (
|
||||
batch_size,
|
||||
num_channels,
|
||||
num_latent_frames,
|
||||
height // spatial_ratio,
|
||||
width // spatial_ratio,
|
||||
)
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("latents", batch.latents, V.none_or_tensor)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5DenoisingStage(PipelineStage):
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _scale_factor(height: int, width: int) -> tuple[float, float, float]:
|
||||
if 480 <= height <= 854 and 480 <= width <= 854:
|
||||
return (1.0, 2.0, 2.0)
|
||||
return (1.0, 3.16, 3.16)
|
||||
|
||||
@staticmethod
|
||||
def _text_rope_pos(mask: torch.Tensor, device: torch.device) -> torch.Tensor:
|
||||
seq_len = int(mask.sum(1).max().item())
|
||||
return torch.arange(seq_len, device=device)
|
||||
|
||||
@staticmethod
|
||||
def fast_sta_nabla(
|
||||
T: int,
|
||||
H: int,
|
||||
W: int,
|
||||
wT: int = 3,
|
||||
wH: int = 3,
|
||||
wW: int = 3,
|
||||
device: torch.device | str = "cuda",
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Create a sparse temporal attention (STA) mask for efficient video generation.
|
||||
|
||||
This method generates a mask that limits attention to nearby frames and spatial positions, reducing
|
||||
computational complexity for video generation.
|
||||
|
||||
Args:
|
||||
T (int): Number of temporal frames
|
||||
H (int): Height in latent space
|
||||
W (int): Width in latent space
|
||||
wT (int): Temporal attention window size
|
||||
wH (int): Height attention window size
|
||||
wW (int): Width attention window size
|
||||
device (str): Device to create tensor on
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Sparse attention mask of shape (T*H*W, T*H*W)
|
||||
"""
|
||||
max_extent = int(torch.tensor([T, H, W], device=device).amax().item())
|
||||
r = torch.arange(0, max_extent, 1, dtype=torch.int16, device=device)
|
||||
mat = (r.unsqueeze(1) - r.unsqueeze(0)).abs()
|
||||
sta_t, sta_h, sta_w = (
|
||||
mat[:T, :T].flatten(),
|
||||
mat[:H, :H].flatten(),
|
||||
mat[:W, :W].flatten(),
|
||||
)
|
||||
sta_t = sta_t <= wT // 2
|
||||
sta_h = sta_h <= wH // 2
|
||||
sta_w = sta_w <= wW // 2
|
||||
sta_hw = (sta_h.unsqueeze(1) * sta_w.unsqueeze(0)).reshape(H, H, W, W).transpose(1, 2).flatten()
|
||||
sta = (sta_t.unsqueeze(1) * sta_hw.unsqueeze(0)).reshape(T, T, H * W, H * W).transpose(1, 2)
|
||||
return sta.reshape(T * H * W, T * H * W)
|
||||
|
||||
def get_sparse_params(self, sample: torch.Tensor, device: torch.device) -> dict[str, Any] | None:
|
||||
"""
|
||||
Generate sparse attention parameters for the transformer based on sample dimensions.
|
||||
|
||||
This method computes the sparse attention configuration needed for efficient video processing in the
|
||||
transformer model.
|
||||
|
||||
Args:
|
||||
sample (torch.Tensor): Input sample tensor
|
||||
device (torch.device): Device to place tensors on
|
||||
|
||||
Returns:
|
||||
Dict: Dictionary containing sparse attention parameters
|
||||
"""
|
||||
assert self.transformer.config.patch_size[0] == 1
|
||||
_, T, H, W, _ = sample.shape
|
||||
T, H, W = (
|
||||
T // self.transformer.config.patch_size[0],
|
||||
H // self.transformer.config.patch_size[1],
|
||||
W // self.transformer.config.patch_size[2],
|
||||
)
|
||||
if self.transformer.config.attention_type == "nabla":
|
||||
sta_mask = self.fast_sta_nabla(
|
||||
T,
|
||||
H // 8,
|
||||
W // 8,
|
||||
self.transformer.config.attention_wT,
|
||||
self.transformer.config.attention_wH,
|
||||
self.transformer.config.attention_wW,
|
||||
device=device,
|
||||
)
|
||||
|
||||
sparse_params = {
|
||||
"sta_mask": sta_mask.unsqueeze_(0).unsqueeze_(0),
|
||||
"attention_type": self.transformer.config.attention_type,
|
||||
"to_fractal": True,
|
||||
"P": self.transformer.config.attention_P,
|
||||
"wT": self.transformer.config.attention_wT,
|
||||
"wW": self.transformer.config.attention_wW,
|
||||
"wH": self.transformer.config.attention_wH,
|
||||
"add_sta": self.transformer.config.attention_add_sta,
|
||||
"visual_shape": (T, H, W),
|
||||
"method": self.transformer.config.attention_method,
|
||||
}
|
||||
else:
|
||||
sparse_params = None
|
||||
|
||||
return sparse_params
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("timesteps must be prepared before Kandinsky5 denoising.")
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be prepared before Kandinsky5 denoising.")
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
autocast_enabled = target_dtype != torch.float32 and not fastvideo_args.disable_autocast
|
||||
latents = batch.latents
|
||||
num_channels = getattr(
|
||||
self.transformer,
|
||||
"in_visual_dim",
|
||||
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
|
||||
)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
pooled = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
if batch.prompt_attention_mask is None or not batch.prompt_attention_mask:
|
||||
raise ValueError("Kandinsky5 requires Qwen prompt attention masks.")
|
||||
text_rope_pos = self._text_rope_pos(batch.prompt_attention_mask[0].to(device), device)
|
||||
|
||||
neg_prompt_embeds = None
|
||||
neg_pooled = None
|
||||
negative_text_rope_pos = None
|
||||
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
neg_pooled = batch.negative_prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
if batch.negative_attention_mask is None or not batch.negative_attention_mask:
|
||||
raise ValueError("Kandinsky5 requires Qwen negative attention masks for CFG.")
|
||||
negative_text_rope_pos = self._text_rope_pos(batch.negative_attention_mask[0].to(device), device)
|
||||
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
temporal_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
num_latent_frames = (int(batch.num_frames) - 1) // temporal_ratio + 1
|
||||
visual_rope_pos = [
|
||||
torch.arange(num_latent_frames, device=device),
|
||||
torch.arange(height // spatial_ratio // 2, device=device),
|
||||
torch.arange(width // spatial_ratio // 2, device=device),
|
||||
]
|
||||
scale_factor = self._scale_factor(height, width)
|
||||
|
||||
sparse_params = self.get_sparse_params(latents, device)
|
||||
|
||||
# I2V keeps the first (conditioning) frame fixed during denoising.
|
||||
# Key off the actual image conditioning, not transformer.visual_cond:
|
||||
# official T2V checkpoints also ship visual_cond=True, and skipping
|
||||
# frame 0 for them leaves it as undenoised noise.
|
||||
cond_frames = 1 if batch.image_latent is not None else 0
|
||||
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
with tqdm(total=batch.num_inference_steps, desc="Kandinsky5 Denoising") as progress_bar:
|
||||
for i, timestep in enumerate(batch.timesteps):
|
||||
if hasattr(self, "interrupt") and self.interrupt:
|
||||
continue
|
||||
|
||||
t_expand = timestep.unsqueeze(0).repeat(latents.shape[0]).to(device=device, dtype=target_dtype)
|
||||
attn_metadata = None
|
||||
if sparse_params is not None:
|
||||
attn_metadata = NablaAttentionMetadataBuilder().build(
|
||||
current_timestep=i,
|
||||
sta_mask=sparse_params["sta_mask"],
|
||||
P=sparse_params["P"],
|
||||
visual_shape=sparse_params["visual_shape"],
|
||||
)
|
||||
autocast_ctx = (torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled)
|
||||
if device.type == "cuda" else contextlib.nullcontext())
|
||||
with set_forward_context(current_timestep=i, attn_metadata=attn_metadata,
|
||||
forward_batch=batch), autocast_ctx:
|
||||
pred_velocity = self.transformer(
|
||||
hidden_states=latents.to(dtype=target_dtype),
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
pooled_projections=pooled,
|
||||
timestep=t_expand,
|
||||
visual_rope_pos=visual_rope_pos,
|
||||
text_rope_pos=text_rope_pos,
|
||||
scale_factor=scale_factor,
|
||||
sparse_params=sparse_params,
|
||||
return_dict=True,
|
||||
).sample
|
||||
|
||||
if neg_prompt_embeds is not None and neg_pooled is not None:
|
||||
uncond_pred_velocity = self.transformer(
|
||||
hidden_states=latents.to(dtype=target_dtype),
|
||||
encoder_hidden_states=neg_prompt_embeds,
|
||||
pooled_projections=neg_pooled,
|
||||
timestep=t_expand,
|
||||
visual_rope_pos=visual_rope_pos,
|
||||
text_rope_pos=negative_text_rope_pos,
|
||||
scale_factor=scale_factor,
|
||||
sparse_params=sparse_params,
|
||||
return_dict=True,
|
||||
).sample
|
||||
pred_velocity = uncond_pred_velocity + batch.guidance_scale * (pred_velocity -
|
||||
uncond_pred_velocity)
|
||||
|
||||
latents[:, cond_frames:, :, :, :num_channels] = self.scheduler.step(
|
||||
pred_velocity[:, cond_frames:],
|
||||
timestep,
|
||||
latents[:, cond_frames:, :, :, :num_channels],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(timestep)
|
||||
# latents is mutated in place, so snapshot a channels-first copy.
|
||||
trajectory_latents.append(latents[..., :num_channels].permute(0, 4, 1, 2, 3).cpu())
|
||||
|
||||
if i == len(batch.timesteps) - 1 or (i + 1) % self.scheduler.order == 0:
|
||||
progress_bar.update()
|
||||
|
||||
if trajectory_latents:
|
||||
batch.trajectory_latents = torch.stack(trajectory_latents, dim=1)
|
||||
batch.trajectory_timesteps = torch.stack(trajectory_timesteps, dim=0).cpu()
|
||||
|
||||
batch.latents = latents[:, :, :, :, :num_channels]
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5DecodingStage(DecodingStage):
|
||||
|
||||
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
|
||||
super().__init__(vae=vae, pipeline=pipeline)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be available before Kandinsky5 decoding.")
|
||||
# Kandinsky5 latents are channels-last [B, T, H, W, C]; the base stage
|
||||
# (and the trajectory latents recorded by the denoising stage) work
|
||||
# channels-first.
|
||||
batch.latents = batch.latents.permute(0, 4, 1, 2, 3).contiguous()
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class Kandinsky5ImageEncodingStage(EncodingStage):
|
||||
"""Encode the conditioning image into a VAE latent for I2V."""
|
||||
|
||||
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
|
||||
super().__init__(vae=vae)
|
||||
|
||||
@staticmethod
|
||||
def _preprocess(image, height: int, width: int) -> torch.Tensor:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
image = resize(image, height, width)
|
||||
image = numpy_to_pt(pil_to_numpy(image)) # always lands in [0, 1]
|
||||
return normalize(image) # [0, 1] -> [-1, 1]
|
||||
# Tensor input: no reliable way to tell [0, 1] from an already
|
||||
# normalized [-1, 1] tensor whose values happen to be non-negative,
|
||||
# so mirror diffusers' heuristic and say what we assumed.
|
||||
if image.min() >= 0:
|
||||
logger.warning("Kandinsky5 conditioning image tensor has no negative values; "
|
||||
"assuming range [0, 1] and normalizing to [-1, 1]. "
|
||||
"Pass a [-1, 1] tensor with negative values to skip normalization.")
|
||||
image = normalize(image) # [0, 1] -> [-1, 1]
|
||||
if image.ndim == 3:
|
||||
image = image.unsqueeze(0)
|
||||
if image.shape[-2:] != (height, width):
|
||||
image = torch.nn.functional.interpolate(image.float(),
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
antialias=True)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.pil_image is None:
|
||||
raise ValueError("Kandinsky5 I2V requires an input image.")
|
||||
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
vae = getattr(self, "vae", None)
|
||||
if vae is None:
|
||||
loader = VAELoader()
|
||||
vae = loader.load(fastvideo_args.model_paths["vae"], fastvideo_args)
|
||||
self.vae = vae
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
device = get_local_torch_device()
|
||||
vae = self.vae.to(device)
|
||||
self.vae = vae
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = vae_dtype != torch.float32 and not fastvideo_args.disable_autocast
|
||||
|
||||
# [B, C, H, W] -> [B, C, 1, H, W]
|
||||
image = self._preprocess(batch.pil_image, int(batch.height), int(batch.width))
|
||||
image = image.to(device=device, dtype=torch.float32).unsqueeze(2)
|
||||
|
||||
# Encode the single conditioning frame without tiling (matches diffusers).
|
||||
# The untested causal-VAE spatial_tiled_encode path corrupts the latent.
|
||||
prev_use_tiling = vae.use_tiling
|
||||
vae.use_tiling = False
|
||||
try:
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
if not vae_autocast_enabled:
|
||||
image = image.to(vae_dtype)
|
||||
# Sample with the batch generator (diffusers parity); mode()
|
||||
# would make seed-for-seed reproduction of the reference
|
||||
# pipeline impossible.
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list) and len(generator) != image.shape[0]:
|
||||
generator = generator[0]
|
||||
image_latent = vae.encode(image).sample(generator=generator)
|
||||
finally:
|
||||
vae.use_tiling = prev_use_tiling
|
||||
|
||||
image_latent = image_latent * vae.scaling_factor
|
||||
# [B, C, 1, H, W] -> [B, 1, H, W, C] to match channels-last latents
|
||||
batch.image_latent = image_latent.permute(0, 2, 3, 4, 1).contiguous()
|
||||
|
||||
# Place the conditioning latent into the prepared latents: frame 0 of
|
||||
# the main channels, the visual_cond channel block, and the mask.
|
||||
# NOTE: the official kandinsky-5 repo leaves the visual_cond block
|
||||
# zeros (generation_utils.py generate()), while the diffusers port
|
||||
# copies the image latent into it. A same-seed A/B on
|
||||
# Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers showed the Diffusers
|
||||
# export requires the copy: zeroing the block produces smeared faces
|
||||
# mid-video. Keep the diffusers semantics for Diffusers-format
|
||||
# checkpoints.
|
||||
latents = batch.latents
|
||||
image_latent = batch.image_latent.to(device=latents.device, dtype=latents.dtype)
|
||||
num_channels = image_latent.shape[-1]
|
||||
latents[:, 0:1, :, :, :num_channels] = image_latent
|
||||
if latents.shape[-1] > num_channels:
|
||||
latents[:, 0:1, :, :, num_channels:2 * num_channels] = image_latent
|
||||
latents[:, 0:1, :, :, 2 * num_channels:] = 1.0
|
||||
batch.latents = latents
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
vae.to("cpu")
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("pil_image", batch.pil_image, V.not_none)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
# This stage runs after latent preparation and writes into its output.
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5NormalizationStage(PipelineStage):
|
||||
"""Normalize the first latent frames to reduce I2V conditioning artifacts."""
|
||||
|
||||
COND_FRAMES = 4
|
||||
REFERENCE_FRAMES = 5
|
||||
|
||||
@staticmethod
|
||||
def _adaptive_mean_std(source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor:
|
||||
source_mean = source.mean(dim=(1, 2, 3, 4), keepdim=True)
|
||||
source_std = source.std(dim=(1, 2, 3, 4), keepdim=True)
|
||||
# Magic constants limit how far the first frames may drift.
|
||||
ref_mean = torch.clamp(reference.mean(dim=(1, 2, 3, 4), keepdim=True), source_mean - 0.05, source_mean + 0.1)
|
||||
ref_std = torch.clamp(reference.std(dim=(1, 2, 3, 4), keepdim=True), source_std - 0.1, source_std + 0.25)
|
||||
normalized = (source - source_mean) / source_std
|
||||
return normalized * ref_std + ref_mean
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
latents = batch.latents
|
||||
n = self.COND_FRAMES
|
||||
if latents is None or latents.shape[1] <= n:
|
||||
return batch
|
||||
|
||||
reference = latents[:, n:n + min(self.REFERENCE_FRAMES, latents.shape[1] - 1)]
|
||||
latents[:, :n] = self._adaptive_mean_std(latents[:, :n].clone(), reference)
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -106,12 +106,30 @@ def run_test_command(test_command: str,
|
||||
if pr_number:
|
||||
print(f"PR number: {pr_number}")
|
||||
|
||||
# For PRs (including forks), use GitHub's PR refs to get the correct commit
|
||||
# The blob-less clone (--filter=blob:none) defers all file-content
|
||||
# downloads to the checkout, so BOTH paths below perform a large lazy blob
|
||||
# fetch from GitHub. Retry transient GitHub/HTTP2 disconnects on each;
|
||||
# otherwise Modal shards can fail before pytest starts.
|
||||
def with_retries(inner_command: str) -> str:
|
||||
return f"""
|
||||
for attempt in 1 2 3; do
|
||||
{inner_command} &&
|
||||
break
|
||||
|
||||
status=$?
|
||||
if [ "$attempt" -eq 3 ]; then
|
||||
exit "$status"
|
||||
fi
|
||||
sleep $((attempt * 5))
|
||||
done"""
|
||||
|
||||
# For PRs (including forks), use GitHub's PR refs to get the correct commit.
|
||||
if pr_number and pr_number != "false":
|
||||
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
|
||||
checkout_command = with_retries(f"git fetch --prune --no-tags --depth=1 origin refs/pull/{pr_number}/head &&"
|
||||
"\n git checkout --detach FETCH_HEAD")
|
||||
print(f"Using PR ref for checkout: {checkout_command}")
|
||||
else:
|
||||
checkout_command = f"git checkout {git_commit}"
|
||||
checkout_command = with_retries(f"git checkout {git_commit}")
|
||||
print(f"Using direct commit checkout: {checkout_command}")
|
||||
|
||||
build_kernel_command = """
|
||||
@@ -125,7 +143,7 @@ def run_test_command(test_command: str,
|
||||
command = f"""
|
||||
source $HOME/.local/bin/env &&
|
||||
source /opt/venv/bin/activate &&
|
||||
git clone {git_repo} /FastVideo &&
|
||||
git clone --filter=blob:none --no-checkout {git_repo} /FastVideo &&
|
||||
cd /FastVideo &&
|
||||
{checkout_command} &&
|
||||
git submodule update --init --recursive &&
|
||||
|
||||
Reference in New Issue
Block a user