Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c7b3b40484 |
@@ -16,6 +16,8 @@ cache_dir/
|
||||
wandb/
|
||||
venv/
|
||||
.venv/
|
||||
/venv-port-magi-2/
|
||||
/archived/
|
||||
runs/
|
||||
samples/
|
||||
Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
@@ -7,3 +7,12 @@
|
||||
[submodule "fastvideo/third_party/eval/vbench"]
|
||||
path = fastvideo/third_party/eval/vbench
|
||||
url = https://github.com/Vchitect/VBench.git
|
||||
[submodule "fastvideo/third_party/magi_attention"]
|
||||
path = fastvideo/third_party/magi_attention
|
||||
url = https://github.com/SandAI-org/MagiAttention.git
|
||||
[submodule "fastvideo/third_party/magi_compiler"]
|
||||
path = fastvideo/third_party/magi_compiler
|
||||
url = https://github.com/SandAI-org/MagiCompiler.git
|
||||
[submodule "fastvideo/third_party/flash_attention"]
|
||||
path = fastvideo/third_party/flash_attention
|
||||
url = https://github.com/Dao-AILab/flash-attention.git
|
||||
|
||||
@@ -157,7 +157,7 @@ def legacy_from_pretrained_to_config(
|
||||
preset_refine["num_inference_steps"] = value
|
||||
elif key == "ltx2_refine_guidance_scale":
|
||||
preset_refine["guidance_scale"] = value
|
||||
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
|
||||
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast", "deterministic"}:
|
||||
engine[key] = value
|
||||
elif key == "override_text_encoder_quant":
|
||||
quantization["text_encoder_quant"] = value
|
||||
@@ -244,6 +244,7 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
|
||||
"enable_stage_verification": engine.enable_stage_verification,
|
||||
"use_fsdp_inference": engine.use_fsdp_inference,
|
||||
"disable_autocast": engine.disable_autocast,
|
||||
"deterministic": engine.deterministic,
|
||||
}
|
||||
if normalized.pipeline.workload_type is not None:
|
||||
kwargs["workload_type"] = normalized.pipeline.workload_type
|
||||
|
||||
@@ -81,6 +81,7 @@ class EngineConfig:
|
||||
enable_stage_verification: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
disable_autocast: bool = False
|
||||
deterministic: bool = False
|
||||
quantization: QuantizationConfig | None = None
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2PreviewVideoConfig, Magi2RefinerVideoConfig
|
||||
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
@@ -19,5 +20,6 @@ __all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"Magi2PreviewVideoConfig", "Magi2RefinerVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Architecture configurations for the MAGI-2 preview and refiner DiTs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_transformer_layer(name: str, module) -> bool:
|
||||
"""Select a complete MAGI-2 transformer layer as an FSDP shard unit."""
|
||||
del module
|
||||
parts = name.split(".")
|
||||
return len(parts) >= 3 and parts[0] == "block" and parts[1] == "layers" and parts[2].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2AttentionGatingConfig:
|
||||
"""Control the learned per-head attention output gate."""
|
||||
|
||||
enable: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2AttentionSinksConfig:
|
||||
"""Configure the learned Flash Attention sink logits."""
|
||||
|
||||
enable: bool = True
|
||||
sink_token_num: int = 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2MHCConfig:
|
||||
"""Configure multi-stream hyper-connections (MHC)."""
|
||||
|
||||
enable: bool = True
|
||||
num_stream: int = 4
|
||||
alpha_init: float = 0.01
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2MoEConfig:
|
||||
"""Configure the preview model's multi-head mixture of experts."""
|
||||
|
||||
num_experts: int = 256
|
||||
top_k: int = 6
|
||||
score_func: str = "sigmoid"
|
||||
route_norm: bool = True
|
||||
route_scale: float = 4.9
|
||||
moe_layers: list[int] = field(default_factory=lambda: list(range(2, 38)))
|
||||
expert_intermediate_size: int = 1280
|
||||
shared_expert_intermediate_size: int = 1280
|
||||
modality_specific_expert_intermediate_size: int = 1280
|
||||
num_heads: int = 12
|
||||
split_merge_intermediate_size: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2PreviewArchConfig(DiTArchConfig):
|
||||
"""Define the published 114B MAGI-2 preview transformer architecture."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.FLASH_ATTN, )
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
num_layers: int = 40
|
||||
hidden_size: int = 3072
|
||||
head_dim: int = 128
|
||||
num_query_groups: int = 24
|
||||
video_in_channels: int = 48
|
||||
audio_in_channels: int = 64
|
||||
text_in_channels: int = 5120
|
||||
params_dtype: str = "bfloat16"
|
||||
intermediate_factor: int = 4
|
||||
mm_layers: list[int] = field(default_factory=lambda: [0, 1, 38, 39])
|
||||
layer_activation_types: dict[int, str] = field(default_factory=dict)
|
||||
activation_type: str = "swiglu7"
|
||||
attn_softcap: float = -1.0
|
||||
attn_gating: Magi2AttentionGatingConfig = field(default_factory=Magi2AttentionGatingConfig)
|
||||
attn_sinks: Magi2AttentionSinksConfig = field(default_factory=Magi2AttentionSinksConfig)
|
||||
mhc_config: Magi2MHCConfig = field(default_factory=Magi2MHCConfig)
|
||||
moe_config: Magi2MoEConfig = field(default_factory=Magi2MoEConfig)
|
||||
|
||||
num_heads_q: int = 0
|
||||
num_heads_kv: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 48
|
||||
in_channels: int = 48
|
||||
out_channels: int = 48
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Derive attention dimensions and normalize nested checkpoint dictionaries."""
|
||||
super().__post_init__()
|
||||
if isinstance(self.attn_gating, dict):
|
||||
self.attn_gating = Magi2AttentionGatingConfig(**self.attn_gating)
|
||||
if isinstance(self.attn_sinks, dict):
|
||||
self.attn_sinks = Magi2AttentionSinksConfig(**self.attn_sinks)
|
||||
if isinstance(self.mhc_config, dict):
|
||||
self.mhc_config = Magi2MHCConfig(**self.mhc_config)
|
||||
if isinstance(self.moe_config, dict):
|
||||
self.moe_config = Magi2MoEConfig(**self.moe_config)
|
||||
self.num_heads_q = self.hidden_size // self.head_dim
|
||||
self.num_heads_kv = self.num_query_groups
|
||||
self.num_attention_heads = self.num_heads_q
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2RefinerArchConfig(DiTArchConfig):
|
||||
"""Define the published MAGI-2 1080p refiner transformer architecture."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.FLASH_ATTN, )
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
num_layers: int = 30
|
||||
hidden_size: int = 4096
|
||||
head_dim: int = 128
|
||||
num_query_groups: int = 8
|
||||
video_in_channels: int = 48
|
||||
audio_in_channels: int = 64
|
||||
text_in_channels: int = 5120
|
||||
checkpoint_qk_layernorm_rope: bool = False
|
||||
params_dtype: str = "bfloat16"
|
||||
mm_layers: list[int] = field(default_factory=lambda: [0, 1, 28, 29])
|
||||
local_attn_layers: list[int] = field(default_factory=lambda: list(range(30)))
|
||||
enable_attn_gating: bool = True
|
||||
activation_type: str = "swiglu7"
|
||||
layer_activation_types: dict[int, str] = field(default_factory=dict)
|
||||
post_norm_layers: list[int] = field(default_factory=list)
|
||||
|
||||
num_heads_q: int = 0
|
||||
num_heads_kv: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 48
|
||||
in_channels: int = 48
|
||||
out_channels: int = 48
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Derive the published grouped-query attention dimensions."""
|
||||
super().__post_init__()
|
||||
self.num_heads_q = self.hidden_size // self.head_dim
|
||||
self.num_heads_kv = self.num_query_groups
|
||||
self.num_attention_heads = self.num_heads_q
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2PreviewVideoConfig(DiTConfig):
|
||||
"""Load the MAGI-2 preview transformer through FastVideo's DiT loader."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=Magi2PreviewArchConfig)
|
||||
prefix: str = "magi2_preview"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2RefinerVideoConfig(DiTConfig):
|
||||
"""Load the MAGI-2 refiner transformer through FastVideo's DiT loader."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=Magi2RefinerArchConfig)
|
||||
prefix: str = "magi2_refiner"
|
||||
@@ -9,6 +9,7 @@ from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.encoders.qwen3_5 import Magi2Qwen35ArchConfig, Magi2Qwen35Config
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
@@ -17,5 +18,6 @@ __all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
|
||||
"Magi2Qwen35ArchConfig", "Magi2Qwen35Config"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Configuration for the MAGI-2 Qwen3.5-27B prompt encoder."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2Qwen35ArchConfig(TextEncoderArchConfig):
|
||||
"""Select Qwen3.5 hidden state -3 and the official token limit."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["Magi2Qwen35TextEncoder"])
|
||||
vocab_size: int = 248320
|
||||
hidden_size: int = 5120
|
||||
intermediate_size: int = 17408
|
||||
num_hidden_layers: int = 64
|
||||
num_attention_heads: int = 24
|
||||
num_key_value_heads: int = 4
|
||||
head_dim: int = 256
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 262144
|
||||
rms_norm_eps: float = 1e-6
|
||||
attention_bias: bool = False
|
||||
attention_dropout: float = 0.0
|
||||
attn_output_gate: bool = True
|
||||
linear_conv_kernel_dim: int = 4
|
||||
linear_key_head_dim: int = 128
|
||||
linear_value_head_dim: int = 128
|
||||
linear_num_key_heads: int = 16
|
||||
linear_num_value_heads: int = 48
|
||||
full_attention_interval: int = 4
|
||||
rope_theta: float = 10000000.0
|
||||
partial_rotary_factor: float = 0.25
|
||||
mrope_section: tuple[int, int, int] = (11, 11, 10)
|
||||
pad_token_id: int | None = None
|
||||
eos_token_id: int = 248044
|
||||
text_len: int = 7000
|
||||
hidden_state_skip_layer: int = 2
|
||||
output_hidden_states: bool = True
|
||||
tokenizer_kwargs: dict = field(default_factory=lambda: {"padding_side": "right"})
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2Qwen35Config(TextEncoderConfig):
|
||||
"""Load the published Qwen3.5 checkpoint through the FastVideo loader."""
|
||||
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=Magi2Qwen35ArchConfig)
|
||||
prefix: str = "magi2_qwen35"
|
||||
@@ -6,6 +6,8 @@ from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.magi2_turbo_vae import Magi2TurboVAEArchConfig, Magi2TurboVAEConfig
|
||||
from fastvideo.configs.models.vaes.magi2_wanvae import Magi2WanVAEArchConfig, Magi2WanVAEConfig
|
||||
from fastvideo.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
|
||||
from fastvideo.configs.models.vaes.flux2vae import Flux2VAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
@@ -19,6 +21,10 @@ __all__ = [
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
"Magi2TurboVAEArchConfig",
|
||||
"Magi2TurboVAEConfig",
|
||||
"Magi2WanVAEArchConfig",
|
||||
"Magi2WanVAEConfig",
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Configuration for the distilled MAGI-2 Turbo VAE decoder."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2TurboVAEArchConfig(VAEArchConfig):
|
||||
"""Describe the latent and output geometry used by Turbo VAE."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["Magi2TurboVAEModel"])
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 48
|
||||
decoder_block_out_channels: tuple[int, ...] = (64, 128, 256, 512)
|
||||
decoder_causal: bool = False
|
||||
decoder_layers_per_block: tuple[int, ...] = (2, 2, 2, 3, 3)
|
||||
patch_size: int = 2
|
||||
patch_size_t: int = 1
|
||||
resnet_norm_eps: float = 1e-6
|
||||
scaling_factor: float = 1.0
|
||||
decoder_spatio_temporal_scaling: tuple[bool, ...] = (False, True, True, True)
|
||||
decoder_spatio_only: tuple[bool, ...] = (False, True, False, False)
|
||||
decoder_is_dw_conv: tuple[bool, ...] = (False, False, False, False, False)
|
||||
decoder_dw_kernel_size: int = 5
|
||||
temporal_compression_ratio: int = 4
|
||||
spatial_compression_ratio: int = 16
|
||||
first_chunk_size: int = 7
|
||||
step_size: int = 7
|
||||
use_unpatchify: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2TurboVAEConfig(VAEConfig):
|
||||
"""Provide the published Turbo VAE JSON and checkpoint paths."""
|
||||
|
||||
arch_config: VAEArchConfig = field(default_factory=Magi2TurboVAEArchConfig)
|
||||
config_path: str = ""
|
||||
checkpoint_path: str = ""
|
||||
pretrained_dtype: str = "bfloat16"
|
||||
load_encoder: bool = False
|
||||
load_decoder: bool = True
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Wan 2.2 VAE configuration for MAGI-2 I2V reference images."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig, WanVAEConfig
|
||||
|
||||
|
||||
MAGI2_WAN_LATENTS_MEAN = (
|
||||
-0.2289, -0.0052, -0.1323, -0.2339, -0.2799, 0.0174, 0.1838, 0.1557,
|
||||
-0.1382, 0.0542, 0.2813, 0.0891, 0.1570, -0.0098, 0.0375, -0.1825,
|
||||
-0.2246, -0.1207, -0.0698, 0.5109, 0.2665, -0.2108, -0.2158, 0.2502,
|
||||
-0.2055, -0.0322, 0.1109, 0.1567, -0.0729, 0.0899, -0.2799, -0.1230,
|
||||
-0.0313, -0.1649, 0.0117, 0.0723, -0.2839, -0.2083, -0.0520, 0.3748,
|
||||
0.0152, 0.1957, 0.1433, -0.2944, 0.3573, -0.0548, -0.1681, -0.0667,
|
||||
)
|
||||
|
||||
MAGI2_WAN_LATENTS_STD = (
|
||||
0.4765, 1.0364, 0.4514, 1.1677, 0.5313, 0.4990, 0.4818, 0.5013,
|
||||
0.8158, 1.0344, 0.5894, 1.0901, 0.6885, 0.6165, 0.8454, 0.4978,
|
||||
0.5759, 0.3523, 0.7135, 0.6804, 0.5833, 1.4146, 0.8986, 0.5659,
|
||||
0.7069, 0.5338, 0.4889, 0.4917, 0.4069, 0.4999, 0.6866, 0.4093,
|
||||
0.5709, 0.6065, 0.6415, 0.4944, 0.5726, 1.2042, 0.5458, 1.6887,
|
||||
0.3971, 1.0600, 0.3943, 0.5537, 0.5444, 0.4089, 0.7468, 0.7744,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2WanVAEArchConfig(WanVAEArchConfig):
|
||||
"""Match the published Wan2.2 VAE encoder architecture and normalization."""
|
||||
|
||||
in_channels: int = 12
|
||||
out_channels: int = 12
|
||||
base_dim: int = 160
|
||||
decoder_base_dim: int | None = 256
|
||||
z_dim: int = 48
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
is_residual: bool = True
|
||||
clip_output: bool = False
|
||||
latents_mean: tuple[float, ...] = MAGI2_WAN_LATENTS_MEAN
|
||||
latents_std: tuple[float, ...] = MAGI2_WAN_LATENTS_STD
|
||||
patch_size: int | None = 2
|
||||
scale_factor_temporal: int = 4
|
||||
scale_factor_spatial: int = 16
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2WanVAEConfig(WanVAEConfig):
|
||||
"""Load only the Wan encoder used by MAGI-2 image conditioning."""
|
||||
|
||||
arch_config: WanVAEArchConfig = field(default_factory=Magi2WanVAEArchConfig)
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = False
|
||||
@@ -9,6 +9,7 @@ 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.configs.pipelines.magi2 import Magi2PreviewPipelineConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.wan import (LucyEditDevConfig, SelfForcingWanT2V480PConfig, WanI2V480PConfig,
|
||||
@@ -19,5 +20,6 @@ __all__ = [
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"Kandinsky5I2VConfig", "Magi2PreviewPipelineConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline configuration for the MAGI-2 Preview 1080p release profile."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import (
|
||||
Magi2PreviewVideoConfig,
|
||||
Magi2RefinerVideoConfig,
|
||||
)
|
||||
from fastvideo.configs.models.vaes import Magi2TurboVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Magi2PreviewPipelineConfig(PipelineConfig):
|
||||
"""Configure the published 10-second MAGI-2 Preview pipeline."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Magi2PreviewVideoConfig)
|
||||
refiner_dit_config: DiTConfig = field(default_factory=Magi2RefinerVideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=Magi2TurboVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple = field(default_factory=tuple)
|
||||
preprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
postprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
|
||||
|
||||
preview_width: int = 896
|
||||
preview_height: int = 512
|
||||
output_width: int = 1920
|
||||
output_height: int = 1088
|
||||
output_frames: int = 249
|
||||
output_fps: int = 25
|
||||
|
||||
preview_video_channels: int = 48
|
||||
preview_video_length: int = 32
|
||||
preview_latent_height: int = 32
|
||||
preview_latent_width: int = 56
|
||||
audio_channels: int = 64
|
||||
audio_latent_length: int = 250
|
||||
refiner_latent_height: int = 68
|
||||
refiner_latent_width: int = 120
|
||||
|
||||
preview_flow_shift: float = 7.0
|
||||
preview_video_guidance_scale: float = 5.0
|
||||
preview_audio_guidance_scale: float = 7.0
|
||||
refiner_flow_shift: float = 5.0
|
||||
refiner_video_guidance_scale: float = 2.0
|
||||
refiner_audio_guidance_scale: float = 5.0
|
||||
refiner_noise_index: int = 220
|
||||
|
||||
|
||||
__all__ = ["Magi2PreviewPipelineConfig"]
|
||||
@@ -46,6 +46,9 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_CFG_GATE_STEP: float = 1.0
|
||||
FASTVIDEO_HOST_IP: str = ""
|
||||
FASTVIDEO_LOOPBACK_IP: str = ""
|
||||
MAGI2_DETERMINISTIC: bool = False
|
||||
MAGI2_SAVE_LATENT_PATH: str | None = None
|
||||
MAGI_ATTENTION_DETERMINISTIC_MODE: bool = False
|
||||
|
||||
|
||||
def get_default_cache_root() -> str:
|
||||
@@ -320,6 +323,14 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - Wan2.2 high/low-noise expert switch invalidates the cache.
|
||||
"FASTVIDEO_CFG_GATE_STEP":
|
||||
lambda: float(os.getenv("FASTVIDEO_CFG_GATE_STEP", "1.0")),
|
||||
|
||||
# MAGI-2 reproducibility and post-refiner latent capture.
|
||||
"MAGI2_DETERMINISTIC":
|
||||
lambda: os.getenv("MAGI2_DETERMINISTIC", "0") == "1",
|
||||
"MAGI2_SAVE_LATENT_PATH":
|
||||
lambda: os.getenv("MAGI2_SAVE_LATENT_PATH"),
|
||||
"MAGI_ATTENTION_DETERMINISTIC_MODE":
|
||||
lambda: os.getenv("MAGI_ATTENTION_DETERMINISTIC_MODE", "0") == "1",
|
||||
}
|
||||
|
||||
# end-env-vars-definition
|
||||
|
||||
@@ -152,6 +152,7 @@ class FastVideoArgs:
|
||||
torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
disable_autocast: bool = False
|
||||
deterministic: bool = False
|
||||
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
@@ -596,6 +597,12 @@ class FastVideoArgs:
|
||||
action=StoreBoolean,
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--deterministic",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.deterministic,
|
||||
help="Enable deterministic model kernels when the pipeline supports them",
|
||||
)
|
||||
|
||||
# VSA parameters
|
||||
parser.add_argument(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,178 @@
|
||||
# Copyright (c) 2026 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Checkpoint loading for MAGI-2.
|
||||
|
||||
Weights ship as a single HuggingFace-style safetensors directory holding the
|
||||
full, unsharded expert stack. Every expert-parallel rank reads that same
|
||||
directory but pulls only the slice of each MoE tensor it owns, so one copy on
|
||||
disk serves any ``ep_size``.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from collections.abc import Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
|
||||
_INDEX_NAME = "model.safetensors.index.json"
|
||||
_SINGLE_NAME = "model.safetensors"
|
||||
|
||||
# MoE tensors are concatenated over (head, expert) along dim 0 and split across
|
||||
# expert-parallel ranks; everything else is replicated.
|
||||
_EP_SHARDED_SUFFIXES = (
|
||||
"moe_mlp.gate",
|
||||
"moe_mlp.W_gate",
|
||||
"moe_mlp.W_up",
|
||||
"moe_mlp.W_down",
|
||||
"moe_mlp.router.expert_bias",
|
||||
"moe_mlp.router.expert_bias_ema",
|
||||
)
|
||||
|
||||
|
||||
def print_rank_0(message: str) -> None:
|
||||
"""Print one checkpoint message from global rank zero."""
|
||||
if psm.get_global_rank() == 0:
|
||||
print(message)
|
||||
|
||||
|
||||
def read_safetensors_weight_map(checkpoint_dir: str) -> dict[str, str]:
|
||||
"""Map every tensor name in a HuggingFace-style directory to its file."""
|
||||
index_path = os.path.join(checkpoint_dir, _INDEX_NAME)
|
||||
if os.path.exists(index_path):
|
||||
with open(index_path) as f:
|
||||
weight_map = json.load(f)["weight_map"]
|
||||
return {key: os.path.join(checkpoint_dir, shard) for key, shard in weight_map.items()}
|
||||
|
||||
single_path = os.path.join(checkpoint_dir, _SINGLE_NAME)
|
||||
if not os.path.exists(single_path):
|
||||
raise FileNotFoundError(f"No {_INDEX_NAME} or {_SINGLE_NAME} under {checkpoint_dir}")
|
||||
with safe_open(single_path, framework="pt") as f:
|
||||
return {key: single_path for key in f}
|
||||
|
||||
|
||||
def load_safetensors_dir(checkpoint_dir: str, desc: str = "Loading shards") -> dict[str, torch.Tensor]:
|
||||
"""Load every tensor from a HuggingFace-style safetensors directory."""
|
||||
return _load_by_shard(read_safetensors_weight_map(checkpoint_dir), reader=_read_whole, desc=desc)
|
||||
|
||||
|
||||
def _is_ep_sharded_key(key: str) -> bool:
|
||||
return ".moe_mlp." in key and key.endswith(_EP_SHARDED_SUFFIXES)
|
||||
|
||||
|
||||
def _read_whole(handle: Any, key: str, _target: torch.Tensor | None) -> torch.Tensor:
|
||||
return handle.get_tensor(key)
|
||||
|
||||
|
||||
def _read_ep_slice(handle: Any, key: str, target: torch.Tensor | None) -> torch.Tensor:
|
||||
"""Read only the (head, expert) rows this expert-parallel rank owns.
|
||||
|
||||
``num_heads`` is padded up to a multiple of ``ep_size``, so the trailing
|
||||
ranks may fall wholly or partly past the real experts and get zeros.
|
||||
"""
|
||||
if target is None or not _is_ep_sharded_key(key):
|
||||
return _read_whole(handle, key, target)
|
||||
|
||||
ep_size = psm.get_world_size("ep")
|
||||
view = handle.get_slice(key)
|
||||
real_rows = view.get_shape()[0]
|
||||
shard_rows = target.shape[0]
|
||||
if ep_size <= 1 or shard_rows == real_rows:
|
||||
return handle.get_tensor(key)
|
||||
|
||||
start = psm.get_local_rank("ep") * shard_rows
|
||||
if start >= real_rows:
|
||||
return torch.zeros_like(target, device="cpu")
|
||||
|
||||
chunk = view[start : min(start + shard_rows, real_rows)]
|
||||
if chunk.shape[0] == shard_rows:
|
||||
return chunk
|
||||
padding = torch.zeros((shard_rows - chunk.shape[0], *chunk.shape[1:]), dtype=chunk.dtype)
|
||||
return torch.cat([chunk, padding], dim=0)
|
||||
|
||||
|
||||
def _load_by_shard(
|
||||
weight_map: Mapping[str, str],
|
||||
reader,
|
||||
desc: str,
|
||||
targets: Mapping[str, torch.Tensor] | None = None,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Read tensors grouped by file so each shard is opened exactly once."""
|
||||
by_shard: dict[str, list[str]] = defaultdict(list)
|
||||
for key, path in weight_map.items():
|
||||
by_shard[path].append(key)
|
||||
|
||||
def read_shard(item: tuple[str, list[str]]) -> dict[str, torch.Tensor]:
|
||||
path, keys = item
|
||||
with safe_open(path, framework="pt") as handle:
|
||||
return {key: reader(handle, key, None if targets is None else targets.get(key)) for key in keys}
|
||||
|
||||
state_dict: dict[str, torch.Tensor] = {}
|
||||
with ThreadPoolExecutor() as executor:
|
||||
futures = list(executor.map(read_shard, by_shard.items()))
|
||||
for shard_state in tqdm(futures, desc=desc, total=len(futures), disable=psm.get_global_rank() != 0):
|
||||
state_dict.update(shard_state)
|
||||
return state_dict
|
||||
|
||||
|
||||
def _apply_router_bias_ema(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
|
||||
"""Pair the router bias with exponential moving average (EMA) model weights.
|
||||
|
||||
The checkpoint stores both ``expert_bias`` (unaveraged training buffer) and
|
||||
``expert_bias_ema`` (EMA-smoothed buffer). Inference uses the EMA model
|
||||
weights, so the router must use the paired EMA bias; otherwise routing
|
||||
biases mismatch the weights and quality degrades. Set the env var
|
||||
``MAGI2_ROUTER_BIAS_SOURCE=main`` to keep the raw ``expert_bias`` instead.
|
||||
"""
|
||||
if (os.environ.get("MAGI2_ROUTER_BIAS_SOURCE") or "ema").strip().lower() == "main":
|
||||
print_rank_0("[magi2] Router bias source: expert_bias (raw, MAGI2_ROUTER_BIAS_SOURCE=main)")
|
||||
return state_dict
|
||||
|
||||
copied = 0
|
||||
for key, value in list(state_dict.items()):
|
||||
if not key.endswith("moe_mlp.router.expert_bias_ema"):
|
||||
continue
|
||||
target_key = key[: -len("_ema")]
|
||||
target = state_dict.get(target_key)
|
||||
if target is None or tuple(target.shape) != tuple(value.shape):
|
||||
continue
|
||||
state_dict[target_key] = value.clone()
|
||||
copied += 1
|
||||
|
||||
print_rank_0(f"[magi2] Router bias: expert_bias_ema -> expert_bias ({copied} tensor(s))")
|
||||
return state_dict
|
||||
|
||||
|
||||
def load_magi2_model_state_dict(model: torch.nn.Module, checkpoint_dir: str) -> Mapping[str, Any]:
|
||||
"""Load the rank-local preview state and select EMA router biases."""
|
||||
if not checkpoint_dir:
|
||||
raise ValueError("checkpoint_dir must point at a safetensors checkpoint directory")
|
||||
|
||||
print_rank_0(f"[magi2] Loading preview weights from {checkpoint_dir}")
|
||||
targets = model.state_dict()
|
||||
weight_map = {key: path for key, path in read_safetensors_weight_map(checkpoint_dir).items() if key in targets}
|
||||
|
||||
state_dict = _load_by_shard(weight_map, reader=_read_ep_slice, desc="Loading shards", targets=targets)
|
||||
for key, value in state_dict.items():
|
||||
target = targets[key]
|
||||
if value.dtype != target.dtype and value.is_floating_point() and target.is_floating_point():
|
||||
state_dict[key] = value.to(dtype=target.dtype)
|
||||
return _apply_router_bias_ema(state_dict)
|
||||
@@ -0,0 +1,55 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Production checkpoint loaders for the MAGI-2 preview and refiner DiTs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2PreviewVideoConfig, Magi2RefinerVideoConfig
|
||||
from fastvideo.models.dits.magi2 import Magi2PreviewDiT
|
||||
from fastvideo.models.dits.magi2_checkpointing import load_magi2_model_state_dict, load_safetensors_dir
|
||||
from fastvideo.models.dits.magi2_refiner import Magi2RefinerDiT
|
||||
|
||||
|
||||
def load_magi2_preview_model(
|
||||
checkpoint_dir: str,
|
||||
config: Magi2PreviewVideoConfig,
|
||||
device: torch.device | str,
|
||||
) -> Magi2PreviewDiT:
|
||||
"""Build the EP-local preview model and strictly load its official state."""
|
||||
model = Magi2PreviewDiT(config=config)
|
||||
state_dict = load_magi2_model_state_dict(model, checkpoint_dir)
|
||||
incompatible = model.load_state_dict(state_dict, strict=True)
|
||||
if incompatible.missing_keys or incompatible.unexpected_keys:
|
||||
raise RuntimeError(
|
||||
"Strict MAGI-2 preview loading reported incompatible keys: "
|
||||
f"missing={incompatible.missing_keys}, unexpected={incompatible.unexpected_keys}"
|
||||
)
|
||||
del state_dict
|
||||
gc.collect()
|
||||
model = model.to(device=device).eval()
|
||||
torch.cuda.empty_cache()
|
||||
return model
|
||||
|
||||
|
||||
def load_magi2_refiner_model(
|
||||
checkpoint_dir: str,
|
||||
config: Magi2RefinerVideoConfig,
|
||||
device: torch.device | str,
|
||||
) -> Magi2RefinerDiT:
|
||||
"""Build a replicated refiner and strictly load every official tensor."""
|
||||
model = Magi2RefinerDiT(config=config)
|
||||
state_dict = load_safetensors_dir(checkpoint_dir, desc="Loading MAGI-2 refiner shards")
|
||||
incompatible = model.load_state_dict(state_dict, strict=True)
|
||||
if incompatible.missing_keys or incompatible.unexpected_keys:
|
||||
raise RuntimeError(
|
||||
"Strict MAGI-2 refiner loading reported incompatible keys: "
|
||||
f"missing={incompatible.missing_keys}, unexpected={incompatible.unexpected_keys}"
|
||||
)
|
||||
del state_dict
|
||||
gc.collect()
|
||||
model = model.to(device=device).eval()
|
||||
torch.cuda.empty_cache()
|
||||
return model
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Exact distributed primitives used by the MAGI-2 inference kernels."""
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime.psm import initialize_expert_parallel, initialize_model_parallel, psm
|
||||
|
||||
__all__ = ["initialize_expert_parallel", "initialize_model_parallel", "psm"]
|
||||
@@ -0,0 +1,142 @@
|
||||
# Copyright (c) 2025 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import List, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
|
||||
def _divide(numerator: int, denominator: int) -> int:
|
||||
assert numerator % denominator == 0, (
|
||||
f"{numerator} is not divisible by {denominator}"
|
||||
)
|
||||
return numerator // denominator
|
||||
|
||||
|
||||
class FakeHandle:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def wait(self):
|
||||
pass
|
||||
|
||||
|
||||
|
||||
def scatter_seqlen_gather_head(
|
||||
tensor: torch.Tensor,
|
||||
split_sizes: List[int] = None,
|
||||
group: dist.ProcessGroup = None,
|
||||
async_op: bool = True,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, Union[dist.Work, FakeHandle]]]:
|
||||
"""
|
||||
Scatter seq_len and gather head_number, for example:
|
||||
input: (seq_len * cp, hn, hd)
|
||||
output: (seq_len, cp * hn, hd)
|
||||
NOTE: seq_len of output maybe not equal, which depends on split_sizes[rank]
|
||||
NOTE: rearrange the tensor after communication: (cp, seq, hn, hd) -> (seq, cp * hn, hd)
|
||||
"""
|
||||
if group is None or dist.get_world_size(group) == 1:
|
||||
return tensor, FakeHandle() if async_op else tensor
|
||||
group_world_size = dist.get_world_size(group)
|
||||
if split_sizes is None:
|
||||
assert tensor.shape[0] % group_world_size == 0, (
|
||||
f"tensor.shape[0] {tensor.shape[0]} % group_world_size {group_world_size} != 0"
|
||||
)
|
||||
split_sizes = [tensor.shape[0] // group_world_size] * group_world_size
|
||||
assert tensor.is_contiguous()
|
||||
assert tensor.dim() == 3, f"tensor must be 3D, but got {tensor.dim()}D"
|
||||
output = torch.empty(
|
||||
[group_world_size * split_sizes[dist.get_rank(group)], *tensor.shape[1:]],
|
||||
device=tensor.device,
|
||||
dtype=tensor.dtype,
|
||||
)
|
||||
output_split_sizes = [split_sizes[dist.get_rank(group)]] * group_world_size
|
||||
if async_op:
|
||||
handle = dist.all_to_all_single(
|
||||
output,
|
||||
tensor,
|
||||
output_split_sizes=output_split_sizes,
|
||||
input_split_sizes=split_sizes,
|
||||
group=group,
|
||||
async_op=True,
|
||||
)
|
||||
return output, handle
|
||||
else:
|
||||
dist.all_to_all_single(
|
||||
output,
|
||||
tensor,
|
||||
output_split_sizes=output_split_sizes,
|
||||
input_split_sizes=split_sizes,
|
||||
group=group,
|
||||
async_op=False,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def batch_scatter_head_gather_seqlen(
|
||||
inputs: List[torch.Tensor],
|
||||
split_sizes: List[int] = None,
|
||||
group: dist.ProcessGroup = None,
|
||||
) -> List[torch.Tensor]:
|
||||
"""
|
||||
Batch scatter head_number and gather seq_len, for example:
|
||||
inputs[i] input: (seq_len_i, cp * hn_i, hd)
|
||||
outputs[i] output: (seq_len_i * cp, hn_i, hd)
|
||||
NOTE: seq_len of inputs maybe not equal across ranks, which depends on split_sizes[rank]
|
||||
NOTE: fuse along head dim before communication, and split back after
|
||||
"""
|
||||
if group is None or dist.get_world_size(group) == 1:
|
||||
return inputs
|
||||
rank = dist.get_rank(group)
|
||||
group_world_size = dist.get_world_size(group)
|
||||
if split_sizes is None:
|
||||
split_sizes = [inputs[0].shape[0]] * group_world_size
|
||||
assert all(input.shape[0] == split_sizes[rank] for input in inputs), (
|
||||
f"inputs[0].shape[0] {inputs[0].shape[0]} != split_sizes[rank] {split_sizes[rank]}"
|
||||
)
|
||||
assert all(input.dim() == 3 for input in inputs), (
|
||||
f"inputs[0].dim() {inputs[0].dim()} != 3"
|
||||
)
|
||||
for idx in range(len(inputs)):
|
||||
_, hn, _ = inputs[idx].shape
|
||||
if group_world_size % hn == 0 and group_world_size != hn:
|
||||
inputs[idx] = torch.repeat_interleave(
|
||||
inputs[idx], repeats=_divide(group_world_size, hn), dim=1
|
||||
)
|
||||
inputs[idx] = rearrange(
|
||||
inputs[idx], "seq (cp hn) hd -> (cp seq) hn hd", cp=group_world_size
|
||||
).contiguous()
|
||||
|
||||
head_split_number = [input.shape[1] for input in inputs]
|
||||
fused_input = torch.cat(inputs, dim=1).contiguous()
|
||||
input_split_sizes = [fused_input.shape[0] // group_world_size] * group_world_size
|
||||
|
||||
fused_output = torch.empty(
|
||||
[sum(split_sizes), *fused_input.shape[1:]],
|
||||
device=fused_input.device,
|
||||
dtype=fused_input.dtype,
|
||||
)
|
||||
dist.all_to_all_single(
|
||||
fused_output,
|
||||
fused_input,
|
||||
output_split_sizes=split_sizes,
|
||||
input_split_sizes=input_split_sizes,
|
||||
group=group,
|
||||
async_op=False,
|
||||
)
|
||||
outputs = torch.split(fused_output, head_split_number, dim=1)
|
||||
return outputs
|
||||
@@ -0,0 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Context-parallel scheduling for MAGI-2."""
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime.context_parallel.ulysses_scheduler import ulysses_scheduler
|
||||
|
||||
__all__ = ["ulysses_scheduler"]
|
||||
@@ -0,0 +1,146 @@
|
||||
# Copyright (c) 2025 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Generic, List, Optional, TypeVar
|
||||
|
||||
import torch
|
||||
from torch.utils._pytree import tree_map
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime.gather_scatter_primitive import (
|
||||
gather_from_context_parallel_region,
|
||||
scatter_to_context_parallel_region,
|
||||
)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class UlyssesScheduler(Generic[T]):
|
||||
"""
|
||||
A naive implementation of Ulysses scheduler for context parallel processing.
|
||||
|
||||
This scheduler handles tensor dispatching and undispatching operations when tensors
|
||||
enter and exit the context parallel region. It supports arbitrary nested data structures
|
||||
containing tensors and applies the same balanced sequence partition to every tensor.
|
||||
|
||||
The scheduler splits input tensors along the sequence dimension across multiple GPUs
|
||||
in the context parallel group, enabling parallel processing of long sequences.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the naive Ulysses scheduler."""
|
||||
self._cp_split_sizes: Optional[List[int]] = None
|
||||
|
||||
@property
|
||||
def cp_split_sizes(self):
|
||||
"""Get the current context parallel split sizes."""
|
||||
return self._cp_split_sizes
|
||||
|
||||
def _dispatch(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Dispatch a tensor to the context parallel region.
|
||||
|
||||
This method splits the tensor along the sequence dimension across the context
|
||||
parallel group. The balanced split sizes differ by at most one token.
|
||||
|
||||
Args:
|
||||
x: Input tensor with shape [seq_len, ...] where seq_len is the sequence length.
|
||||
|
||||
Returns:
|
||||
Dispatched tensor that has been split and distributed across the context parallel group.
|
||||
|
||||
Raises:
|
||||
AssertionError: If the split sizes change between calls, indicating inconsistent
|
||||
sequence lengths or context parallel group size.
|
||||
"""
|
||||
seq_len = x.shape[0]
|
||||
cp_world_size = psm.get_world_size("cp")
|
||||
if seq_len % cp_world_size == 0:
|
||||
cp_split_sizes = [seq_len // cp_world_size] * cp_world_size
|
||||
else:
|
||||
num_ranks_with_one_extra = seq_len % cp_world_size
|
||||
min_tokens_per_rank = (seq_len - num_ranks_with_one_extra) // cp_world_size
|
||||
cp_split_sizes = [min_tokens_per_rank + 1] * num_ranks_with_one_extra + [min_tokens_per_rank] * (
|
||||
cp_world_size - num_ranks_with_one_extra
|
||||
)
|
||||
if self._cp_split_sizes is not None:
|
||||
assert (
|
||||
self._cp_split_sizes == cp_split_sizes
|
||||
), f"cp_split_sizes changed from {self._cp_split_sizes} to {cp_split_sizes}"
|
||||
self._cp_split_sizes = cp_split_sizes
|
||||
x = scatter_to_context_parallel_region(x, cp_split_sizes, group=psm.get_parallel_group("cp"))
|
||||
return x
|
||||
|
||||
def _undispatch(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Undispatch a tensor from the context parallel region.
|
||||
|
||||
This method gathers the balanced tensor slices from all ranks in the context
|
||||
parallel group and concatenates them into the original sequence.
|
||||
|
||||
Args:
|
||||
x: Dispatched tensor from the context parallel region.
|
||||
|
||||
Returns:
|
||||
Reconstructed tensor with the original sequence length.
|
||||
"""
|
||||
x = gather_from_context_parallel_region(x, self._cp_split_sizes, group=psm.get_parallel_group("cp"))
|
||||
return x
|
||||
|
||||
def dispatch(self, tensors: T) -> T:
|
||||
"""
|
||||
Apply dispatch operation to all tensor leaf nodes in a nested data structure.
|
||||
|
||||
This method recursively applies the _dispatch operation to all tensors in the
|
||||
input data structure, preparing them for context parallel computation. The
|
||||
structure of the input is preserved in the output.
|
||||
|
||||
Args:
|
||||
tensors: Arbitrary nested data structure containing tensors (single tensor,
|
||||
tuple, list, dict, etc.). All tensors should have the same sequence
|
||||
length in their first dimension.
|
||||
|
||||
Returns:
|
||||
A new data structure with the same structure as input, where all tensors
|
||||
have been dispatched to the context parallel region.
|
||||
"""
|
||||
return tree_map(self._dispatch, tensors)
|
||||
|
||||
def undispatch(self, tensors: T) -> T:
|
||||
"""
|
||||
Apply undispatch operation to all tensor leaf nodes in a nested data structure.
|
||||
|
||||
This method recursively applies the _undispatch operation to all tensors in the
|
||||
input data structure, reconstructing them from the context parallel region. The
|
||||
structure of the input is preserved in the output.
|
||||
|
||||
Args:
|
||||
tensors: Arbitrary nested data structure containing dispatched tensors.
|
||||
|
||||
Returns:
|
||||
A new data structure with the same structure as input, where all tensors
|
||||
have been reconstructed from the context parallel region.
|
||||
"""
|
||||
output = tree_map(self._undispatch, tensors)
|
||||
self._cp_split_sizes = None
|
||||
return output
|
||||
|
||||
|
||||
_ULYSSES_SCHEDULER = UlyssesScheduler()
|
||||
|
||||
|
||||
def ulysses_scheduler() -> UlyssesScheduler:
|
||||
assert _ULYSSES_SCHEDULER is not None, "ulysses scheduler is not initialized"
|
||||
return _ULYSSES_SCHEDULER
|
||||
@@ -0,0 +1,33 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Connect MAGI-2 collectives to FastVideo-owned process groups."""
|
||||
|
||||
from fastvideo.distributed import get_dp_group, get_sp_group
|
||||
from fastvideo.models.dits.magi2_runtime import psm as psm_manager
|
||||
from fastvideo.models.dits.magi2_runtime.psm import bind_process_groups
|
||||
|
||||
|
||||
MAGI2_PARALLEL_SIZE = 8
|
||||
|
||||
|
||||
def bind_fastvideo_parallel_state() -> None:
|
||||
"""Use FastVideo sequence parallelism for MAGI-2 context and experts."""
|
||||
sequence_group = get_sp_group()
|
||||
data_group = get_dp_group()
|
||||
if sequence_group.world_size != MAGI2_PARALLEL_SIZE:
|
||||
raise ValueError(
|
||||
"MAGI-2 Preview strict inference requires sp_size=8; "
|
||||
f"received sp_size={sequence_group.world_size}"
|
||||
)
|
||||
bind_process_groups(
|
||||
cp_group=sequence_group.device_group,
|
||||
cp_ranks=sequence_group.ranks,
|
||||
dp_group=data_group.device_group,
|
||||
dp_ranks=data_group.ranks,
|
||||
ep_group=sequence_group.device_group,
|
||||
ep_ranks=sequence_group.ranks,
|
||||
)
|
||||
if psm_manager.get_world_size("cp") != MAGI2_PARALLEL_SIZE:
|
||||
raise RuntimeError("MAGI-2 context-parallel binding did not preserve eight ranks")
|
||||
|
||||
|
||||
__all__ = ["MAGI2_PARALLEL_SIZE", "bind_fastvideo_parallel_state"]
|
||||
@@ -0,0 +1,11 @@
|
||||
# Vendored flash_mh_moe inference subset.
|
||||
# Routing/sorting: pure PyTorch.
|
||||
# Forward GEMM: fused Triton kernel with deterministic scatter support.
|
||||
from .fwd import flash_mh_moe_fwd
|
||||
from .route import compute_topk_probs_and_indices, flash_mh_moe_global_sort
|
||||
|
||||
__all__ = [
|
||||
"compute_topk_probs_and_indices",
|
||||
"flash_mh_moe_global_sort",
|
||||
"flash_mh_moe_fwd",
|
||||
]
|
||||
@@ -0,0 +1,40 @@
|
||||
# Copyright (c) 2025-2026 SandAI. All Rights Reserved.
|
||||
# Apache-2.0.
|
||||
"""Flash-MH-MoE forward — fused Triton kernel with deterministic scatter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _is_deterministic() -> bool:
|
||||
return os.environ.get("MAGI2_DETERMINISTIC", "0") == "1"
|
||||
|
||||
|
||||
def flash_mh_moe_fwd(
|
||||
x: torch.Tensor,
|
||||
gather_ids: torch.Tensor,
|
||||
probs: torch.Tensor,
|
||||
expert_offsets: torch.Tensor,
|
||||
W_gate: torch.Tensor,
|
||||
W_up: torch.Tensor,
|
||||
W_down: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Flash-MH-MoE forward: gather → gate/up GEMM → SwiGLU7 → down GEMM → scatter.
|
||||
|
||||
When ``MAGI2_DETERMINISTIC=1``, uses sequential scatter for bit-exact reproducibility.
|
||||
"""
|
||||
from .triton import mh_moe_fwd_func
|
||||
|
||||
return mh_moe_fwd_func(
|
||||
x=x,
|
||||
gather_ids=gather_ids,
|
||||
probs=probs,
|
||||
expert_offsets=expert_offsets,
|
||||
W_gate=W_gate,
|
||||
W_up=W_up,
|
||||
W_down=W_down,
|
||||
deterministic=_is_deterministic(),
|
||||
)
|
||||
@@ -0,0 +1,249 @@
|
||||
# Copyright (c) 2025-2026 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Literal, TypeAlias
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
TopKScoreFunc: TypeAlias = Literal["softmax", "sigmoid"]
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Fused top-k probabilities and indices computation
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _compute_topk_probs_and_indices_torch_impl(
|
||||
router_logits: torch.Tensor,
|
||||
top_k: int,
|
||||
score_func: TopKScoreFunc = "softmax",
|
||||
expert_bias: torch.Tensor | None = None,
|
||||
num_virtual_group: int = 1,
|
||||
topk_virtual_group: int = 0,
|
||||
route_norm: bool = True,
|
||||
norm_eps: float = 1e-12,
|
||||
use_score_matrix_for_topk_probs: bool = False,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Select and normalize each token's expert-routing probabilities."""
|
||||
H, S, E = router_logits.shape
|
||||
K = top_k
|
||||
|
||||
# ── Gating scores ──
|
||||
|
||||
match score_func:
|
||||
case "softmax":
|
||||
router_scores = F.softmax(router_logits, dim=-1)
|
||||
case "sigmoid":
|
||||
router_scores = F.sigmoid(router_logits)
|
||||
case _:
|
||||
raise ValueError(f"Unknown score_func: {score_func}")
|
||||
|
||||
# ── Aux-free bias for topk selection (not for actual probs) ──
|
||||
|
||||
topk_scores = router_scores
|
||||
if expert_bias is not None:
|
||||
bias = expert_bias.view(H, 1, E) # [H, E] -> [H, 1, E]
|
||||
topk_scores = topk_scores + bias
|
||||
|
||||
# ── TopK per head per token (with optional virtual group filtering) ──
|
||||
|
||||
if num_virtual_group > 1:
|
||||
assert (
|
||||
topk_virtual_group > 0 and topk_virtual_group <= num_virtual_group
|
||||
), f"topk_virtual_group ({topk_virtual_group}) must be in (0, num_virtual_group ({num_virtual_group})]"
|
||||
assert (
|
||||
E % num_virtual_group == 0
|
||||
), f"num_experts ({E}) must be divisible by num_virtual_group ({num_virtual_group})"
|
||||
experts_per_group = E // num_virtual_group
|
||||
|
||||
# Score each group by sum of its top-2 expert scores
|
||||
# [H, S, E] -> [H, S, num_virtual_group, experts_per_group]
|
||||
# -> topk(2, dim=-1) -> [H, S, num_virtual_group, 2]
|
||||
# -> sum(dim=-1) -> [H, S, num_virtual_group]
|
||||
group_scores = (
|
||||
topk_scores.view(H, S, num_virtual_group, experts_per_group)
|
||||
.topk(min(2, experts_per_group), dim=-1)[0]
|
||||
.sum(dim=-1)
|
||||
)
|
||||
|
||||
# Select top-k virtual groups
|
||||
# [H, S, num_virtual_group] -> topk(k, dim=-1) -> [H, S, topk_virtual_group]
|
||||
_, group_idx = torch.topk(
|
||||
group_scores,
|
||||
k=topk_virtual_group,
|
||||
dim=-1,
|
||||
sorted=False,
|
||||
)
|
||||
|
||||
# Build a mask to select experts from the top-k virtual groups
|
||||
# group_mask: [H, S, num_virtual_group]
|
||||
# score_mask: [H, S, num_virtual_group, 1] -> expand(dim=-1) -> [H, S, E]
|
||||
group_mask = torch.zeros_like(group_scores)
|
||||
group_mask.scatter_(-1, group_idx, 1)
|
||||
score_mask = (
|
||||
group_mask.unsqueeze(-1)
|
||||
.expand(H, S, num_virtual_group, experts_per_group)
|
||||
.reshape(H, S, E)
|
||||
)
|
||||
|
||||
# Mask out scores not in the top-k virtual groups before selecting top-k experts
|
||||
# topk_scores: [H, S, E] -> masked_fill(~score_mask, 0.0) -> [H, S, E]
|
||||
# -> topk(k, dim=-1) -> topk_probs: [H, S, K]
|
||||
scores_for_choice = topk_scores.masked_fill(~score_mask.bool(), 0.0)
|
||||
_, topk_indices = torch.topk(scores_for_choice, k=K, dim=-1, sorted=False)
|
||||
else:
|
||||
_, topk_indices = torch.topk(topk_scores, K, dim=-1) # [H, S, K]
|
||||
|
||||
if use_score_matrix_for_topk_probs:
|
||||
# The score-matrix path masks unselected experts, normalizes across all
|
||||
# experts, and then gathers the selected probabilities. Gathering before
|
||||
# normalization is algebraically equivalent but can differ numerically
|
||||
# because the two paths normalize different tensor representations.
|
||||
global_routing_matrix = torch.zeros(
|
||||
H,
|
||||
S,
|
||||
E,
|
||||
device=topk_indices.device,
|
||||
dtype=torch.int32,
|
||||
).scatter_(-1, topk_indices, 1)
|
||||
|
||||
global_score_matrix = global_routing_matrix * router_scores
|
||||
|
||||
# ── Normalize ──
|
||||
if route_norm:
|
||||
global_score_matrix = F.normalize(
|
||||
global_score_matrix, p=1, dim=-1, eps=norm_eps
|
||||
)
|
||||
|
||||
topk_probs = global_score_matrix.gather(-1, topk_indices)
|
||||
else:
|
||||
# Expert selection uses biased or virtual-group-masked topk_scores, while
|
||||
# probability magnitudes come from the unbiased router_scores.
|
||||
topk_probs = router_scores.gather(-1, topk_indices) # [H, S, K]
|
||||
|
||||
# ── Normalize ──
|
||||
|
||||
if route_norm:
|
||||
topk_probs = F.normalize(topk_probs, p=1, dim=-1, eps=norm_eps)
|
||||
|
||||
return (
|
||||
topk_probs,
|
||||
topk_indices,
|
||||
)
|
||||
|
||||
|
||||
score_func_name2id_map = {
|
||||
"softmax": 0,
|
||||
"sigmoid": 1,
|
||||
}
|
||||
|
||||
|
||||
def compute_topk_probs_and_indices(
|
||||
router_logits: torch.Tensor,
|
||||
top_k: int,
|
||||
score_func: TopKScoreFunc = "softmax",
|
||||
expert_bias: torch.Tensor | None = None,
|
||||
num_virtual_group: int = 1,
|
||||
topk_virtual_group: int = 0,
|
||||
route_norm: bool = True,
|
||||
norm_eps: float = 1e-12,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compute the top-k probabilities and indices from router logits.
|
||||
|
||||
Args:
|
||||
router_logits: Tensor of shape [H, S, E], the router logits for each head, token, and expert.
|
||||
top_k: The number of top probabilities and indices to compute.
|
||||
score_func: The function to convert logits to scores, either "softmax" or "sigmoid".
|
||||
expert_bias: Optional tensor of shape [H, E] to be added to the router scores before top-k selection,
|
||||
for biasing certain experts.
|
||||
num_virtual_group: If >1, the experts are divided into this many virtual groups
|
||||
for a two-stage top-k selection (first select top-k virtual groups,
|
||||
then select top-k experts within those groups).
|
||||
topk_virtual_group: If num_virtual_group > 1,
|
||||
the number of top virtual groups to select before selecting top-k experts.
|
||||
route_norm: Whether to normalize the top-k probabilities to sum to 1. Defaults to ``True``.
|
||||
norm_eps: the epsilon used for route_norm. Defaults to ``1e-12``.
|
||||
|
||||
|
||||
Returns:
|
||||
A tuple of (topk_probs, topk_indices):
|
||||
- topk_probs: Tensor of shape [H, S, top_k], the top-k probabilities for each head and token.
|
||||
- topk_indices: Tensor of shape [H, S, top_k],
|
||||
the indices of the top-k experts for each head and token.
|
||||
"""
|
||||
return _compute_topk_probs_and_indices_torch_impl(
|
||||
router_logits=router_logits,
|
||||
top_k=top_k,
|
||||
score_func=score_func,
|
||||
expert_bias=expert_bias,
|
||||
num_virtual_group=num_virtual_group,
|
||||
topk_virtual_group=topk_virtual_group,
|
||||
route_norm=route_norm,
|
||||
norm_eps=norm_eps,
|
||||
)
|
||||
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Mapping topk probs/indices to sorted gather ids and probs and offsets
|
||||
# prepared as meta inputs for flash_mh_moe kernels
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def flash_mh_moe_global_sort(
|
||||
topk_probs: torch.Tensor,
|
||||
topk_indices: torch.Tensor,
|
||||
num_experts: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Map per-head expert ids to global [0, H*E) and argsort into CSR format.
|
||||
|
||||
Args:
|
||||
topk_probs (torch.Tensor): Routing probabilities with shape ``(H, S, K)``.
|
||||
topk_indices (torch.Tensor): Per-head expert indices in ``[0, E)`` with shape ``(H, S, K)``.
|
||||
num_experts (int): E, number of experts per head.
|
||||
|
||||
Returns:
|
||||
gather_ids (torch.Tensor): Token indices to gather from ``[S]`` with shape ``(T,)`` (int32).
|
||||
probs_sorted (torch.Tensor): Routing probabilities sorted by expert with shape ``(T,)``.
|
||||
expert_offsets (torch.Tensor): CSR offsets with shape ``(H*E+1,)`` (long).
|
||||
"""
|
||||
H, S, K = topk_indices.shape
|
||||
E = num_experts
|
||||
device = topk_indices.device
|
||||
|
||||
head_offset = torch.arange(H, device=device).view(H, 1, 1) * E
|
||||
global_indices = topk_indices + head_offset
|
||||
|
||||
flat_indices = global_indices.reshape(-1)
|
||||
flat_probs = topk_probs.reshape(-1)
|
||||
flat_token_ids = (
|
||||
torch.arange(S, device=device).view(1, S, 1).expand(H, S, K).reshape(-1)
|
||||
)
|
||||
|
||||
sorted_order = flat_indices.argsort(stable=True)
|
||||
gather_ids = flat_token_ids[sorted_order].to(torch.int32)
|
||||
probs_sorted = flat_probs[sorted_order].float()
|
||||
|
||||
expert_counts = torch.zeros(H * E, device=device, dtype=torch.long)
|
||||
expert_counts.scatter_add_(
|
||||
0,
|
||||
flat_indices[sorted_order],
|
||||
torch.ones_like(sorted_order, dtype=torch.long),
|
||||
)
|
||||
expert_offsets = torch.zeros(H * E + 1, device=device, dtype=torch.long)
|
||||
expert_offsets[1:] = expert_counts.cumsum(0)
|
||||
|
||||
return gather_ids, probs_sorted, expert_offsets
|
||||
@@ -0,0 +1,3 @@
|
||||
from .mh_moe_fwd import mh_moe_fwd_func
|
||||
|
||||
__all__ = ["mh_moe_fwd_func"]
|
||||
@@ -0,0 +1,335 @@
|
||||
# Copyright (c) 2025-2026 SandAI. All Rights Reserved.
|
||||
# Apache-2.0.
|
||||
"""Triton kernel for multi-head MoE forward (inference).
|
||||
|
||||
Multi-head MoE forward: gather → GEMM(gate,up) → SwiGLU7 → GEMM(down) → scatter.
|
||||
Supports deterministic mode (sequential scatter_back) for bit-exact reproducibility.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
__all__ = ["mh_moe_fwd_func"]
|
||||
|
||||
SWIGLU7_ALPHA = tl.constexpr(1.702)
|
||||
SWIGLU7_LIMIT = tl.constexpr(7.0)
|
||||
SWIGLU7_BIAS = tl.constexpr(1.0)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _swiglu7_fwd(gate, up, out_dtype: tl.constexpr):
|
||||
gate_clamped = tl.minimum(gate, SWIGLU7_LIMIT)
|
||||
up_clamped = tl.maximum(tl.minimum(up, SWIGLU7_LIMIT), -SWIGLU7_LIMIT)
|
||||
sig = tl.sigmoid(SWIGLU7_ALPHA * gate_clamped)
|
||||
swish = gate_clamped * sig
|
||||
return (swish * (up_clamped + SWIGLU7_BIAS)).to(out_dtype)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _binary_search_expert_id(
|
||||
cu_expert_tile_counts_ptr,
|
||||
tile_id,
|
||||
NUM_EXPERTS: tl.constexpr,
|
||||
LOG2_NUM_EXPERTS: tl.constexpr,
|
||||
):
|
||||
"""Map a flattened output tile to the expert that owns the tile."""
|
||||
lo = 0
|
||||
hi = NUM_EXPERTS
|
||||
for _ in tl.static_range(0, LOG2_NUM_EXPERTS + 1):
|
||||
mid = (lo + hi + 1) // 2
|
||||
below = tl.load(cu_expert_tile_counts_ptr + mid) <= tile_id
|
||||
lo = tl.where(below, mid, lo)
|
||||
hi = tl.where(below, hi, mid - 1)
|
||||
return lo
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _mh_moe_fwd_kernel(
|
||||
x_ptr,
|
||||
wg_ptr,
|
||||
wu_ptr,
|
||||
wd_ptr,
|
||||
y_ptr,
|
||||
gather_ids_ptr,
|
||||
probs_ptr,
|
||||
expert_offsets_ptr,
|
||||
cu_expert_tile_counts_ptr,
|
||||
stride_x_s,
|
||||
stride_x_h,
|
||||
stride_x_dh,
|
||||
stride_wg_ne,
|
||||
stride_wg_dh,
|
||||
stride_wg_de,
|
||||
stride_wu_ne,
|
||||
stride_wu_dh,
|
||||
stride_wu_de,
|
||||
stride_wd_ne,
|
||||
stride_wd_de,
|
||||
stride_wd_dh,
|
||||
stride_y_s,
|
||||
stride_y_h,
|
||||
stride_y_dh,
|
||||
D_HEAD: tl.constexpr,
|
||||
D_EXPERT: tl.constexpr,
|
||||
NUM_HEADS: tl.constexpr,
|
||||
NUM_EXPERTS: tl.constexpr,
|
||||
LOG2_NUM_EXPERTS: tl.constexpr,
|
||||
BLOCK_T: tl.constexpr,
|
||||
BLOCK_DH: tl.constexpr,
|
||||
BLOCK_DE: tl.constexpr,
|
||||
ACC_DTYPE: tl.constexpr = tl.float32,
|
||||
DETERMINISTIC: tl.constexpr = False,
|
||||
):
|
||||
"""Run the fused gather, expert SwiGLU7 projections, and output scatter."""
|
||||
tile_id = tl.program_id(0)
|
||||
|
||||
total_tiles = tl.load(cu_expert_tile_counts_ptr + NUM_EXPERTS)
|
||||
if tile_id >= total_tiles:
|
||||
return
|
||||
|
||||
expert_id = _binary_search_expert_id(
|
||||
cu_expert_tile_counts_ptr, tile_id, NUM_EXPERTS, LOG2_NUM_EXPERTS
|
||||
)
|
||||
expert_id_i64 = expert_id.to(tl.int64)
|
||||
head_id = expert_id // (NUM_EXPERTS // NUM_HEADS)
|
||||
|
||||
tile_in_expert = tile_id - tl.load(cu_expert_tile_counts_ptr + expert_id)
|
||||
token_start = tl.load(expert_offsets_ptr + expert_id) + tile_in_expert * BLOCK_T
|
||||
expert_end = tl.load(expert_offsets_ptr + expert_id + 1)
|
||||
n = tl.minimum(token_start + BLOCK_T, expert_end) - token_start
|
||||
|
||||
block_dh_offs = tl.arange(0, BLOCK_DH)
|
||||
block_de_offs = tl.arange(0, BLOCK_DE)
|
||||
block_t_offs = tl.arange(0, BLOCK_T)
|
||||
dh_offs = tl.arange(0, D_HEAD)
|
||||
|
||||
block_token_offs = token_start + block_t_offs
|
||||
block_token_mask = block_t_offs < n
|
||||
|
||||
gather_ids = tl.load(
|
||||
gather_ids_ptr + block_token_offs, mask=block_token_mask, other=0
|
||||
)
|
||||
probs = tl.load(probs_ptr + block_token_offs, mask=block_token_mask, other=0.0)
|
||||
x_offs = gather_ids * stride_x_s + head_id * stride_x_h
|
||||
|
||||
y_acc = tl.zeros([BLOCK_T, D_HEAD], dtype=ACC_DTYPE)
|
||||
for dexpert_start in tl.range(0, D_EXPERT, BLOCK_DE):
|
||||
dexpert_block_offs = dexpert_start + block_de_offs
|
||||
|
||||
gate_acc = tl.zeros([BLOCK_T, BLOCK_DE], dtype=ACC_DTYPE)
|
||||
up_acc = tl.zeros([BLOCK_T, BLOCK_DE], dtype=ACC_DTYPE)
|
||||
for dhead_start in tl.static_range(0, D_HEAD, BLOCK_DH):
|
||||
dhead_block_offs = dhead_start + block_dh_offs
|
||||
|
||||
x_block = tl.load(
|
||||
x_ptr + x_offs[:, None] + dhead_block_offs[None, :] * stride_x_dh,
|
||||
mask=block_token_mask[:, None],
|
||||
other=0.0,
|
||||
)
|
||||
w_gate = tl.load(
|
||||
wg_ptr
|
||||
+ expert_id_i64 * stride_wg_ne
|
||||
+ dhead_block_offs[:, None] * stride_wg_dh
|
||||
+ dexpert_block_offs[None, :] * stride_wg_de
|
||||
)
|
||||
w_up = tl.load(
|
||||
wu_ptr
|
||||
+ expert_id_i64 * stride_wu_ne
|
||||
+ dhead_block_offs[:, None] * stride_wu_dh
|
||||
+ dexpert_block_offs[None, :] * stride_wu_de
|
||||
)
|
||||
|
||||
gate_acc += tl.dot(x_block, w_gate)
|
||||
up_acc += tl.dot(x_block, w_up)
|
||||
|
||||
hidden = _swiglu7_fwd(gate_acc, up_acc, wd_ptr.dtype.element_ty)
|
||||
|
||||
w_down = tl.load(
|
||||
wd_ptr
|
||||
+ expert_id_i64 * stride_wd_ne
|
||||
+ dexpert_block_offs[:, None] * stride_wd_de
|
||||
+ dh_offs[None, :] * stride_wd_dh
|
||||
)
|
||||
|
||||
y_acc += tl.dot(hidden, w_down)
|
||||
|
||||
y_acc = y_acc * probs[:, None]
|
||||
|
||||
if DETERMINISTIC:
|
||||
# Store to flat (T, D_HEAD) buffer; host scatter_back accumulates deterministically.
|
||||
y_ptrs = (
|
||||
y_ptr
|
||||
+ block_token_offs[:, None] * stride_y_s
|
||||
+ dh_offs[None, :] * stride_y_dh
|
||||
)
|
||||
tl.store(
|
||||
y_ptrs,
|
||||
y_acc.to(y_ptr.dtype.element_ty),
|
||||
mask=block_token_mask[:, None],
|
||||
)
|
||||
else:
|
||||
# Scatter-reduce by atomic add into (S, H, D_HEAD).
|
||||
y_offs = gather_ids * stride_y_s + head_id * stride_y_h
|
||||
y_ptrs = y_ptr + y_offs[:, None] + dh_offs[None, :] * stride_y_dh
|
||||
tl.atomic_add(
|
||||
y_ptrs,
|
||||
y_acc.to(y_ptr.dtype.element_ty),
|
||||
mask=block_token_mask[:, None],
|
||||
)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Host-side scatter_back (deterministic path)
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _scatter_back(
|
||||
y_sorted: torch.Tensor,
|
||||
ref_y: torch.Tensor,
|
||||
gather_ids: torch.Tensor,
|
||||
expert_offsets: torch.Tensor,
|
||||
num_flatten_experts: int,
|
||||
num_heads: int,
|
||||
) -> torch.Tensor:
|
||||
"""Scatter prob-weighted expert outputs back to (S, H, D) deterministically."""
|
||||
D = ref_y.size(2)
|
||||
num_experts_per_head = num_flatten_experts // num_heads
|
||||
|
||||
y = torch.zeros_like(ref_y).view(-1, D) # (S*H, D)
|
||||
|
||||
expert_seqlens = expert_offsets.diff()
|
||||
head_id_values = torch.arange(num_flatten_experts, device=gather_ids.device) // num_experts_per_head
|
||||
head_ids = torch.repeat_interleave(head_id_values, expert_seqlens)
|
||||
|
||||
if num_heads == 1:
|
||||
scatter_idx = gather_ids.long()
|
||||
else:
|
||||
scatter_idx = gather_ids.long() * num_heads + head_ids.long()
|
||||
|
||||
scatter_idx = scatter_idx.unsqueeze(-1).expand_as(y_sorted)
|
||||
y.scatter_add_(0, scatter_idx, y_sorted.to(y.dtype))
|
||||
|
||||
return y.view_as(ref_y)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Public forward function
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _select_block_config() -> tuple[int, int, int, int, int]:
|
||||
"""(BLOCK_T, BLOCK_DH, BLOCK_DE, NUM_STAGES, NUM_WARPS) for d_head=256."""
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major >= 10: # Blackwell
|
||||
return (128, 64, 32, 2, 8)
|
||||
return (128, 64, 32, 2, 8) # Hopper
|
||||
|
||||
|
||||
def mh_moe_fwd_func(
|
||||
x: torch.Tensor,
|
||||
gather_ids: torch.Tensor,
|
||||
probs: torch.Tensor,
|
||||
expert_offsets: torch.Tensor,
|
||||
W_gate: torch.Tensor,
|
||||
W_up: torch.Tensor,
|
||||
W_down: torch.Tensor,
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Triton multi-head MoE forward: gather → gate/up GEMM → SwiGLU7 → down GEMM → scatter.
|
||||
|
||||
When ``deterministic=True``, the kernel writes per-expert outputs to a flat buffer
|
||||
and the host performs a sequential scatter_back for bit-exact reproducibility.
|
||||
|
||||
Args:
|
||||
x: (S, H, d_head)
|
||||
gather_ids: (T,) int32
|
||||
probs: (T,) float32
|
||||
expert_offsets: (H*E+1,) int32/int64
|
||||
W_gate: (H*E, d_head, d_expert)
|
||||
W_up: (H*E, d_head, d_expert)
|
||||
W_down: (H*E, d_expert, d_head)
|
||||
deterministic: use sequential scatter for reproducibility
|
||||
"""
|
||||
T = gather_ids.size(0)
|
||||
if T == 0:
|
||||
return torch.zeros_like(x)
|
||||
|
||||
num_heads = x.size(1)
|
||||
d_head, d_expert = x.size(2), W_down.size(1)
|
||||
|
||||
BLOCK_T, BLOCK_DH, BLOCK_DE, NUM_STAGES, NUM_WARPS = _select_block_config()
|
||||
assert d_head % BLOCK_DH == 0, f"d_head={d_head} must be divisible by {BLOCK_DH}"
|
||||
assert d_expert % BLOCK_DE == 0, f"d_expert={d_expert} must be divisible by {BLOCK_DE}"
|
||||
|
||||
if deterministic:
|
||||
y = torch.empty(T, 1, d_head, device=x.device, dtype=x.dtype)
|
||||
else:
|
||||
y = torch.zeros_like(x)
|
||||
|
||||
num_flatten_experts = expert_offsets.size(0) - 1
|
||||
num_flatten_experts_log2 = max(1, math.ceil(math.log2(max(num_flatten_experts, 1) + 1)))
|
||||
|
||||
expert_token_counts = expert_offsets.diff()
|
||||
expert_tile_counts = (expert_token_counts + BLOCK_T - 1) // BLOCK_T
|
||||
cu_expert_tile_counts = torch.cat([
|
||||
torch.zeros(1, dtype=torch.int32, device=expert_offsets.device),
|
||||
torch.cumsum(expert_tile_counts, dim=0, dtype=torch.int32),
|
||||
])
|
||||
|
||||
grid_size = (T + BLOCK_T - 1) // BLOCK_T + num_flatten_experts
|
||||
|
||||
_mh_moe_fwd_kernel[(grid_size,)](
|
||||
x_ptr=x,
|
||||
wg_ptr=W_gate,
|
||||
wu_ptr=W_up,
|
||||
wd_ptr=W_down,
|
||||
y_ptr=y,
|
||||
gather_ids_ptr=gather_ids,
|
||||
probs_ptr=probs,
|
||||
expert_offsets_ptr=expert_offsets,
|
||||
cu_expert_tile_counts_ptr=cu_expert_tile_counts,
|
||||
stride_x_s=x.stride(0),
|
||||
stride_x_h=x.stride(1),
|
||||
stride_x_dh=x.stride(2),
|
||||
stride_wg_ne=W_gate.stride(0),
|
||||
stride_wg_dh=W_gate.stride(1),
|
||||
stride_wg_de=W_gate.stride(2),
|
||||
stride_wu_ne=W_up.stride(0),
|
||||
stride_wu_dh=W_up.stride(1),
|
||||
stride_wu_de=W_up.stride(2),
|
||||
stride_wd_ne=W_down.stride(0),
|
||||
stride_wd_de=W_down.stride(1),
|
||||
stride_wd_dh=W_down.stride(2),
|
||||
stride_y_s=y.stride(0),
|
||||
stride_y_h=y.stride(1),
|
||||
stride_y_dh=y.stride(2),
|
||||
D_HEAD=d_head,
|
||||
D_EXPERT=d_expert,
|
||||
NUM_HEADS=num_heads,
|
||||
NUM_EXPERTS=num_flatten_experts,
|
||||
LOG2_NUM_EXPERTS=num_flatten_experts_log2,
|
||||
BLOCK_T=BLOCK_T,
|
||||
BLOCK_DH=BLOCK_DH,
|
||||
BLOCK_DE=BLOCK_DE,
|
||||
ACC_DTYPE=tl.float32,
|
||||
DETERMINISTIC=deterministic,
|
||||
num_stages=NUM_STAGES,
|
||||
num_warps=NUM_WARPS,
|
||||
)
|
||||
|
||||
if deterministic:
|
||||
y = _scatter_back(
|
||||
y_sorted=y.view(T, d_head),
|
||||
ref_y=x,
|
||||
gather_ids=gather_ids,
|
||||
expert_offsets=expert_offsets,
|
||||
num_flatten_experts=num_flatten_experts,
|
||||
num_heads=num_heads,
|
||||
)
|
||||
|
||||
return y
|
||||
@@ -0,0 +1,221 @@
|
||||
# Copyright (c) 2025 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from functools import partial
|
||||
from typing import List, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.utils._pytree import tree_map
|
||||
|
||||
|
||||
|
||||
class Metadata:
|
||||
def __init__(self, dtype: torch.dtype, numel: int, ndim: int, shape: List[int]):
|
||||
self.dtype = dtype
|
||||
self.numel = numel
|
||||
self.ndim = ndim
|
||||
self.shape = shape
|
||||
|
||||
def __repr__(self):
|
||||
return f"Metadata(dtype={self.dtype}, numel={self.numel}, ndim={self.ndim}, shape={self.shape})"
|
||||
|
||||
|
||||
def _gather_metadata(tensor_list: List[torch.Tensor], group: dist.ProcessGroup) -> List[List[Metadata]]:
|
||||
"""Exchange tensor shapes and dtypes for an arbitrary distributed gather."""
|
||||
dist.get_rank(group)
|
||||
world_size = dist.get_world_size(group)
|
||||
|
||||
local_rank = torch.distributed.get_rank() % torch.cuda.device_count()
|
||||
assert (
|
||||
local_rank == torch.cuda.current_device()
|
||||
), f"local_rank {local_rank} != current_device {torch.cuda.current_device()}"
|
||||
device = tensor_list[0].device if len(tensor_list) > 0 else torch.device("cuda")
|
||||
|
||||
# ========== Step 1: flatten local tensor list ==========
|
||||
|
||||
# Metadata: [dtype_code, numel, ndim, *shape]
|
||||
local_metadata = []
|
||||
|
||||
dtype_map = {torch.float32: 0, torch.float16: 1, torch.bfloat16: 2, torch.int32: 3, torch.int64: 4, torch.uint8: 5}
|
||||
reverse_dtype_map = {v: k for k, v in dtype_map.items()}
|
||||
|
||||
for t in tensor_list:
|
||||
dtype_code = dtype_map[t.dtype]
|
||||
shape = list(t.shape)
|
||||
numel = t.numel()
|
||||
local_metadata.append(torch.tensor([dtype_code, numel, len(shape)] + shape, dtype=torch.int32, device=device))
|
||||
|
||||
if local_metadata:
|
||||
local_metadata_tensor = torch.cat(local_metadata)
|
||||
else:
|
||||
local_metadata_tensor = torch.empty(0, dtype=torch.int32, device=device)
|
||||
local_metadata_tensor = local_metadata_tensor.contiguous()
|
||||
local_metadata_len = torch.tensor([local_metadata_tensor.numel()], dtype=torch.int32, device=device)
|
||||
|
||||
# ========== Step 2: all_gather metadata lengths ==========
|
||||
metadata_lens = [torch.empty_like(local_metadata_len) for _ in range(world_size)]
|
||||
dist.all_gather(metadata_lens, local_metadata_len, group)
|
||||
|
||||
# ========== Step 3: all_gather metadata payloads (with cpu tensor) ==========
|
||||
metadata_lists = [torch.empty(m.item(), dtype=torch.int32, device=device) for m in metadata_lens]
|
||||
dist.all_gather(metadata_lists, local_metadata_tensor, group)
|
||||
|
||||
# ========== Step 4: decode metadata and reconstruct tensor list ==========
|
||||
result = []
|
||||
for metadata_list in metadata_lists:
|
||||
offset = 0
|
||||
local_metadata = []
|
||||
while offset < metadata_list.numel():
|
||||
dtype_code = metadata_list[offset].item()
|
||||
numel = metadata_list[offset + 1].item()
|
||||
ndim = metadata_list[offset + 2].item()
|
||||
shape = metadata_list[offset + 3 : offset + 3 + ndim].tolist()
|
||||
offset += 3 + ndim
|
||||
|
||||
local_metadata.append(Metadata(reverse_dtype_map[dtype_code], numel, ndim, shape))
|
||||
result.append(local_metadata)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _get_dtype_and_assert_consistency(metadata_lists: List[List[Metadata]]):
|
||||
dtype_set = set()
|
||||
for metadata_list in metadata_lists:
|
||||
for metadata in metadata_list:
|
||||
dtype_set.add(metadata.dtype)
|
||||
assert len(dtype_set) == 1, f"Metadata lists are not consistent: {dtype_set}"
|
||||
return dtype_set.pop()
|
||||
|
||||
|
||||
def _get_numel_for_each_rank(metadata_lists: List[List[Metadata]]) -> List[int]:
|
||||
return [sum(meta.numel for meta in metadata_list) for metadata_list in metadata_lists]
|
||||
|
||||
|
||||
def gather_arbitrary_tensor_list(tensor_list: List[torch.Tensor], group: dist.ProcessGroup) -> List[torch.Tensor]:
|
||||
"""
|
||||
Magic gather primitive. Provide the following features:
|
||||
1. Support tensor list with different length for each rank.
|
||||
2. Support arbitrary Tensor, which means the Tensor can have different shapes but same dtype.
|
||||
3. Support empty tensor_list in some ranks without padding.
|
||||
|
||||
Args:
|
||||
tensor_list: A list of tensors to gather.
|
||||
group: The process group to use.
|
||||
|
||||
Returns:
|
||||
A list of tensors gathered from all ranks.
|
||||
"""
|
||||
|
||||
dist.get_rank(group)
|
||||
world_size = dist.get_world_size(group)
|
||||
|
||||
local_rank = torch.distributed.get_rank() % torch.cuda.device_count()
|
||||
assert (
|
||||
local_rank == torch.cuda.current_device()
|
||||
), f"local_rank {local_rank} != current_device {torch.cuda.current_device()}"
|
||||
device = tensor_list[0].device if len(tensor_list) > 0 else torch.device("cuda")
|
||||
|
||||
# Step 1: Gather metadata
|
||||
metadata_lists = _gather_metadata(tensor_list, group)
|
||||
tensor_dtype = _get_dtype_and_assert_consistency(metadata_lists)
|
||||
|
||||
# Step 2: Flatten local tensors into a single 1D buffer
|
||||
if tensor_list:
|
||||
flat_tensor = torch.cat([t.flatten() for t in tensor_list], dim=0).contiguous()
|
||||
else:
|
||||
flat_tensor = torch.empty(0, dtype=tensor_dtype, device=device) # dummy, will be ignored
|
||||
|
||||
# Step 3: Gather lengths from metadata
|
||||
all_numels_int = _get_numel_for_each_rank(metadata_lists)
|
||||
|
||||
# Step 4: Allocate buffers and gather flat tensor data
|
||||
output_flat_tensors = []
|
||||
for numel in all_numels_int:
|
||||
output_flat_tensors.append(torch.empty(numel, dtype=tensor_dtype, device=device))
|
||||
dist.all_gather(output_flat_tensors, flat_tensor, group)
|
||||
|
||||
# Step 5: Reconstruct individual tensors using metadata
|
||||
gathered_tensor_lists = []
|
||||
for i in range(world_size):
|
||||
flat = output_flat_tensors[i]
|
||||
if flat.numel() == 0:
|
||||
continue
|
||||
metadata_list = metadata_lists[i]
|
||||
offset = 0
|
||||
for meta in metadata_list:
|
||||
numel = meta.numel
|
||||
t = flat[offset : offset + numel].view(meta.shape).to(meta.dtype)
|
||||
offset += numel
|
||||
gathered_tensor_lists.append(t)
|
||||
|
||||
return gathered_tensor_lists
|
||||
|
||||
|
||||
def _scatter_to_context_parallel_region(input: torch.Tensor, split_sizes: List[int], group: dist.ProcessGroup = None):
|
||||
"""Split the tensor along its first dimension and keep the
|
||||
corresponding slice."""
|
||||
# Split along first dimension with padding.
|
||||
rank = dist.get_rank(group)
|
||||
dim_offset = sum(split_sizes[:rank])
|
||||
output = input[dim_offset : dim_offset + split_sizes[rank]].contiguous()
|
||||
return output
|
||||
|
||||
|
||||
def scatter_to_context_parallel_region(
|
||||
inputs: Union[torch.Tensor, List[torch.Tensor]], split_sizes: List[int] = None, group: dist.ProcessGroup = None
|
||||
):
|
||||
"""Split the tensor along its first dimension and keep the
|
||||
corresponding slice."""
|
||||
if group is None or torch.distributed.get_world_size(group) == 1:
|
||||
return inputs
|
||||
|
||||
if split_sizes is None:
|
||||
assert (
|
||||
inputs.shape[0] % dist.get_world_size(group) == 0
|
||||
), f"inputs.shape[0] {inputs.shape[0]} % dist.get_world_size(group) {dist.get_world_size(group)} != 0"
|
||||
split_sizes = [inputs.shape[0] // dist.get_world_size(group)] * dist.get_world_size(group)
|
||||
|
||||
partial_func = partial(_scatter_to_context_parallel_region, split_sizes=split_sizes, group=group)
|
||||
return tree_map(partial_func, inputs)
|
||||
|
||||
|
||||
def _gather_from_context_parallel_region(
|
||||
input: Union[torch.Tensor, List[torch.Tensor]], split_sizes: List[int], group: dist.ProcessGroup = None
|
||||
):
|
||||
"""Gather context-parallel slices and concatenate them in rank order."""
|
||||
input = input.contiguous()
|
||||
dim_size = list(input.size())
|
||||
dim_size[0] = sum(split_sizes)
|
||||
|
||||
output = torch.empty(dim_size, dtype=input.dtype, device=input.device)
|
||||
outputs = list(torch.split(output, split_sizes, dim=0))
|
||||
torch.distributed.all_gather(outputs, input, group=group)
|
||||
output = torch.concat(outputs, dim=0)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def gather_from_context_parallel_region(
|
||||
inputs: Union[torch.Tensor, List[torch.Tensor]], split_sizes: List[int] = None, group: dist.ProcessGroup = None
|
||||
):
|
||||
"""Gather tensors and concatinate along the first dimension."""
|
||||
|
||||
if group is None or torch.distributed.get_world_size(group) == 1:
|
||||
return inputs
|
||||
|
||||
if split_sizes is None:
|
||||
split_sizes = [inputs.shape[0] * dist.get_world_size(group)]
|
||||
partial_func = partial(_gather_from_context_parallel_region, split_sizes=split_sizes, group=group)
|
||||
return tree_map(partial_func, inputs)
|
||||
@@ -0,0 +1,174 @@
|
||||
# Copyright (c) 2025-2026 SandAI. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Process-group management for MAGI-2 inference.
|
||||
|
||||
Inference uses three parallelism axes:
|
||||
|
||||
* **Context parallel (cp)** splits one sample's token sequence across ranks.
|
||||
* **Data parallel (dp)** gives different samples to different ranks.
|
||||
* **Expert parallel (ep)** shards MoE experts across ranks.
|
||||
|
||||
Ranks are numbered cp-major: a cp group is a contiguous block while a dp
|
||||
group strides across blocks. With ``world_size=8, cp_size=4`` the cp
|
||||
groups are ``[0,1,2,3]`` and ``[4,5,6,7]``; dp groups are ``[0,4]``,
|
||||
``[1,5]``, ``[2,6]``, ``[3,7]``.
|
||||
|
||||
All inference code accesses groups through the ``psm`` singleton below.
|
||||
"""
|
||||
|
||||
from datetime import timedelta
|
||||
from typing import List, Optional
|
||||
|
||||
import torch.distributed as dist
|
||||
|
||||
_CP_GROUP = None
|
||||
_CP_RANKS: Optional[List[int]] = None
|
||||
_DP_GROUP = None
|
||||
_DP_RANKS: Optional[List[int]] = None
|
||||
_EP_GROUP = None
|
||||
_EP_RANKS: Optional[List[int]] = None
|
||||
|
||||
|
||||
def bind_process_groups(
|
||||
cp_group: "dist.ProcessGroup",
|
||||
cp_ranks: list[int],
|
||||
dp_group: "dist.ProcessGroup",
|
||||
dp_ranks: list[int],
|
||||
ep_group: "dist.ProcessGroup",
|
||||
ep_ranks: list[int],
|
||||
) -> None:
|
||||
"""Bind MAGI-2 parallel axes to process groups owned by FastVideo."""
|
||||
global _CP_GROUP, _CP_RANKS, _DP_GROUP, _DP_RANKS, _EP_GROUP, _EP_RANKS
|
||||
|
||||
if not dist.is_initialized():
|
||||
raise RuntimeError("torch.distributed must be initialized before binding MAGI-2 groups")
|
||||
rank = dist.get_rank()
|
||||
for group_name, group, ranks in (
|
||||
("cp", cp_group, cp_ranks),
|
||||
("dp", dp_group, dp_ranks),
|
||||
("ep", ep_group, ep_ranks),
|
||||
):
|
||||
if group is None or rank not in ranks:
|
||||
raise ValueError(f"Rank {rank} does not belong to the MAGI-2 {group_name} group {ranks}")
|
||||
|
||||
_CP_GROUP, _CP_RANKS = cp_group, list(cp_ranks)
|
||||
_DP_GROUP, _DP_RANKS = dp_group, list(dp_ranks)
|
||||
_EP_GROUP, _EP_RANKS = ep_group, list(ep_ranks)
|
||||
|
||||
|
||||
def initialize_model_parallel(cp_size: int = 1, distributed_timeout_minutes: int = 30) -> None:
|
||||
"""Create cp and dp groups. Must be called after ``dist.init_process_group``."""
|
||||
global _CP_GROUP, _CP_RANKS, _DP_GROUP, _DP_RANKS
|
||||
|
||||
assert dist.is_initialized(), "torch.distributed must be initialized first"
|
||||
assert _CP_GROUP is None, "parallel state is already initialized"
|
||||
|
||||
world_size = dist.get_world_size()
|
||||
assert world_size % cp_size == 0, f"world_size {world_size} not divisible by cp_size {cp_size}"
|
||||
dp_size = world_size // cp_size
|
||||
rank = dist.get_rank()
|
||||
timeout = timedelta(minutes=distributed_timeout_minutes)
|
||||
|
||||
for i in range(dp_size):
|
||||
ranks = list(range(i * cp_size, (i + 1) * cp_size))
|
||||
group = dist.new_group(ranks, timeout=timeout)
|
||||
if rank in ranks:
|
||||
_CP_GROUP, _CP_RANKS = group, ranks
|
||||
|
||||
for i in range(cp_size):
|
||||
ranks = list(range(i, world_size, cp_size))
|
||||
group = dist.new_group(ranks, timeout=timeout)
|
||||
if rank in ranks:
|
||||
_DP_GROUP, _DP_RANKS = group, ranks
|
||||
|
||||
|
||||
def initialize_expert_parallel(ep_size: int) -> None:
|
||||
"""Create ep groups. Call after ``initialize_model_parallel``."""
|
||||
global _EP_GROUP, _EP_RANKS
|
||||
if ep_size <= 1:
|
||||
return
|
||||
world_size = dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
assert world_size % ep_size == 0
|
||||
for i in range(world_size // ep_size):
|
||||
ranks = list(range(i * ep_size, (i + 1) * ep_size))
|
||||
group = dist.new_group(ranks)
|
||||
if rank in ranks:
|
||||
_EP_GROUP, _EP_RANKS = group, ranks
|
||||
|
||||
|
||||
class ParallelStateManager:
|
||||
"""Minimal psm interface used by MAGI-2 inference."""
|
||||
|
||||
@staticmethod
|
||||
def is_initialized() -> bool:
|
||||
return _CP_GROUP is not None
|
||||
|
||||
@staticmethod
|
||||
def get_global_rank() -> int:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
return dist.get_rank()
|
||||
return 0
|
||||
|
||||
def get_world_size(self, dim: str = "") -> int:
|
||||
if not dist.is_available() or not dist.is_initialized():
|
||||
return 1
|
||||
if dim == "cp":
|
||||
return dist.get_world_size(group=_CP_GROUP)
|
||||
elif dim == "dp":
|
||||
return dist.get_world_size(group=_DP_GROUP)
|
||||
elif dim == "ep":
|
||||
return dist.get_world_size(group=_EP_GROUP) if _EP_GROUP else dist.get_world_size()
|
||||
return dist.get_world_size()
|
||||
|
||||
def get_local_rank(self, dim: str) -> int:
|
||||
if not dist.is_available() or not dist.is_initialized():
|
||||
return 0
|
||||
if dim == "cp":
|
||||
return dist.get_rank(group=_CP_GROUP)
|
||||
elif dim == "dp":
|
||||
return dist.get_rank(group=_DP_GROUP)
|
||||
elif dim == "ep":
|
||||
return dist.get_rank(group=_EP_GROUP) if _EP_GROUP else dist.get_rank()
|
||||
return 0
|
||||
|
||||
def is_group_first_rank(self, dim: str) -> bool:
|
||||
return self.get_local_rank(dim) == 0
|
||||
|
||||
def is_group_last_rank(self, dim: str) -> bool:
|
||||
return self.get_local_rank(dim) == (self.get_world_size(dim) - 1)
|
||||
|
||||
def get_parallel_group(self, dim: str) -> Optional["dist.ProcessGroup"]:
|
||||
if not dist.is_available() or not dist.is_initialized():
|
||||
return None
|
||||
if dim == "cp":
|
||||
return _CP_GROUP
|
||||
elif dim == "dp":
|
||||
return _DP_GROUP
|
||||
elif dim == "ep":
|
||||
return _EP_GROUP
|
||||
return None
|
||||
|
||||
def get_global_ranks(self, dim: str) -> List[int]:
|
||||
if dim == "cp":
|
||||
return _CP_RANKS or [self.get_global_rank()]
|
||||
elif dim == "dp":
|
||||
return _DP_RANKS or [self.get_global_rank()]
|
||||
elif dim == "ep":
|
||||
return _EP_RANKS or list(range(dist.get_world_size()))
|
||||
return [self.get_global_rank()]
|
||||
|
||||
|
||||
psm = ParallelStateManager()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MAGI-2 Stable Audio Open decoder loading and source-key conversion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.models.vaes.oobleck import OobleckDecoder
|
||||
|
||||
_SOURCE_PREFIX = "pretransform.model."
|
||||
|
||||
|
||||
class Magi2AudioVAE(nn.Module):
|
||||
"""Decode MAGI-2 audio latents with the Stable Audio Oobleck decoder."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
decoder: OobleckDecoder,
|
||||
sampling_rate: int,
|
||||
downsampling_ratio: int,
|
||||
) -> None:
|
||||
"""Store the decoder and the published audio sampling geometry."""
|
||||
super().__init__()
|
||||
self.decoder = decoder
|
||||
self.sampling_rate = sampling_rate
|
||||
self.hop_length = downsampling_ratio
|
||||
|
||||
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode channel-major audio latents into channel-major waveforms."""
|
||||
if latent.ndim == 2:
|
||||
latent = latent.unsqueeze(0)
|
||||
return self.decoder(latent)
|
||||
|
||||
|
||||
def _source_decoder_to_fastvideo_key(
|
||||
source_name: str,
|
||||
decoder_block_count: int,
|
||||
) -> str:
|
||||
"""Map one Stable Audio sequential decoder key to an Oobleck decoder key."""
|
||||
parts = source_name.split(".")
|
||||
if parts[:2] != ["decoder", "layers"] or len(parts) < 4:
|
||||
raise ValueError(f"Unsupported Stable Audio decoder key: {source_name}")
|
||||
|
||||
top_level_index = int(parts[2])
|
||||
parameter_parts = parts[3:]
|
||||
if top_level_index == 0:
|
||||
return ".".join(["decoder", "conv1", *parameter_parts])
|
||||
if top_level_index == decoder_block_count + 1:
|
||||
return ".".join(["decoder", "snake1", *parameter_parts])
|
||||
if top_level_index == decoder_block_count + 2:
|
||||
return ".".join(["decoder", "conv2", *parameter_parts])
|
||||
if not 1 <= top_level_index <= decoder_block_count:
|
||||
raise ValueError(f"Unsupported Stable Audio decoder key: {source_name}")
|
||||
|
||||
if parameter_parts[:1] != ["layers"] or len(parameter_parts) < 3:
|
||||
raise ValueError(f"Unsupported Stable Audio decoder block key: {source_name}")
|
||||
block_index = top_level_index - 1
|
||||
block_layer_index = int(parameter_parts[1])
|
||||
block_parameter_parts = parameter_parts[2:]
|
||||
if block_layer_index == 0:
|
||||
target_parts = ["decoder", "block", str(block_index), "snake1"]
|
||||
elif block_layer_index == 1:
|
||||
target_parts = ["decoder", "block", str(block_index), "conv_t1"]
|
||||
elif 2 <= block_layer_index <= 4:
|
||||
if block_parameter_parts[:1] != ["layers"] or len(block_parameter_parts) < 3:
|
||||
raise ValueError(f"Unsupported Stable Audio residual key: {source_name}")
|
||||
residual_layer_index = int(block_parameter_parts[1])
|
||||
residual_names = ("snake1", "conv1", "snake2", "conv2")
|
||||
if residual_layer_index >= len(residual_names):
|
||||
raise ValueError(f"Unsupported Stable Audio residual key: {source_name}")
|
||||
target_parts = [
|
||||
"decoder",
|
||||
"block",
|
||||
str(block_index),
|
||||
f"res_unit{block_layer_index - 1}",
|
||||
residual_names[residual_layer_index],
|
||||
]
|
||||
block_parameter_parts = block_parameter_parts[2:]
|
||||
else:
|
||||
raise ValueError(f"Unsupported Stable Audio decoder block key: {source_name}")
|
||||
return ".".join([*target_parts, *block_parameter_parts])
|
||||
|
||||
|
||||
def _load_decoder_state(
|
||||
decoder_model: Magi2AudioVAE,
|
||||
checkpoint_path: Path,
|
||||
decoder_block_count: int,
|
||||
) -> None:
|
||||
"""Read only decoder tensors and load the mapped state strictly."""
|
||||
target_state = decoder_model.state_dict()
|
||||
mapped_state: dict[str, torch.Tensor] = {}
|
||||
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint:
|
||||
for checkpoint_name in list(checkpoint.keys()):
|
||||
if not checkpoint_name.startswith(f"{_SOURCE_PREFIX}decoder."):
|
||||
continue
|
||||
source_name = checkpoint_name.removeprefix(_SOURCE_PREFIX)
|
||||
target_name = _source_decoder_to_fastvideo_key(
|
||||
source_name,
|
||||
decoder_block_count,
|
||||
)
|
||||
if target_name not in target_state:
|
||||
raise KeyError(
|
||||
f"Stable Audio key {checkpoint_name} maps to unknown key {target_name}"
|
||||
)
|
||||
source_tensor = checkpoint.get_tensor(checkpoint_name)
|
||||
target_tensor = target_state[target_name]
|
||||
if source_tensor.numel() != target_tensor.numel():
|
||||
raise ValueError(
|
||||
f"Stable Audio tensor size differs for {checkpoint_name}: "
|
||||
f"source={tuple(source_tensor.shape)}, target={tuple(target_tensor.shape)}"
|
||||
)
|
||||
mapped_state[target_name] = source_tensor.reshape(target_tensor.shape)
|
||||
incompatible = decoder_model.load_state_dict(mapped_state, strict=True)
|
||||
if incompatible.missing_keys or incompatible.unexpected_keys:
|
||||
raise RuntimeError(
|
||||
"Stable Audio strict load failed: "
|
||||
f"missing={incompatible.missing_keys}, "
|
||||
f"unexpected={incompatible.unexpected_keys}"
|
||||
)
|
||||
|
||||
|
||||
def load_magi2_audio_vae(
|
||||
component_dir: str | Path,
|
||||
device: torch.device | str,
|
||||
) -> Magi2AudioVAE:
|
||||
"""Build and strictly load MAGI-2's published FP32 audio decoder."""
|
||||
component_path = Path(component_dir)
|
||||
config_path = component_path / "model_config.json"
|
||||
checkpoint_path = component_path / "model.safetensors"
|
||||
with config_path.open(encoding="utf-8") as config_file:
|
||||
full_config = json.load(config_file)
|
||||
sampling_rate = int(full_config["sample_rate"])
|
||||
vae_config = full_config["model"]["pretransform"]["config"]
|
||||
decoder_spec = vae_config["decoder"]
|
||||
decoder_config = decoder_spec["config"]
|
||||
if decoder_spec["type"] != "oobleck" or not decoder_config["use_snake"]:
|
||||
raise ValueError("MAGI-2 audio decoding requires the published Oobleck Snake decoder")
|
||||
if decoder_config.get("use_nearest_upsample", False):
|
||||
raise ValueError("MAGI-2 audio decoding requires transposed-convolution upsampling")
|
||||
if decoder_config.get("final_tanh", True):
|
||||
raise ValueError("MAGI-2 audio decoding requires final_tanh=false")
|
||||
|
||||
strides = [int(stride) for stride in decoder_config["strides"]]
|
||||
downsampling_ratio = int(vae_config["downsampling_ratio"])
|
||||
if math.prod(strides) != downsampling_ratio:
|
||||
raise ValueError("Stable Audio strides do not match the declared downsampling ratio")
|
||||
decoder = OobleckDecoder(
|
||||
channels=int(decoder_config["channels"]),
|
||||
input_channels=int(decoder_config["latent_dim"]),
|
||||
audio_channels=int(decoder_config["out_channels"]),
|
||||
upsampling_ratios=list(reversed(strides)),
|
||||
channel_multiples=[int(value) for value in decoder_config["c_mults"]],
|
||||
)
|
||||
audio_vae = Magi2AudioVAE(
|
||||
decoder=decoder,
|
||||
sampling_rate=sampling_rate,
|
||||
downsampling_ratio=downsampling_ratio,
|
||||
)
|
||||
_load_decoder_state(audio_vae, checkpoint_path, len(strides))
|
||||
audio_vae.to(device=device, dtype=torch.float32)
|
||||
audio_vae.requires_grad_(False)
|
||||
return audio_vae.eval()
|
||||
|
||||
|
||||
__all__ = ["Magi2AudioVAE", "load_magi2_audio_vae"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,529 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Copyright (c) 2025 SandAI. All Rights Reserved.
|
||||
"""Strict source-checkpoint loading for the MAGI-2 Wan image encoder."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.configs.models.vaes.magi2_wanvae import Magi2WanVAEConfig
|
||||
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
|
||||
class CausalConv3d(nn.Conv3d):
|
||||
"""Apply left-padded temporal convolution with an optional feature cache."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
"""Store causal padding separately from the underlying convolution."""
|
||||
super().__init__(*args, **kwargs)
|
||||
self._padding = (
|
||||
self.padding[2],
|
||||
self.padding[2],
|
||||
self.padding[1],
|
||||
self.padding[1],
|
||||
2 * self.padding[0],
|
||||
0,
|
||||
)
|
||||
self.padding = (0, 0, 0)
|
||||
|
||||
@torch.compile
|
||||
def forward(self, x: torch.Tensor, cache_x: torch.Tensor | None = None) -> torch.Tensor:
|
||||
"""Convolve one temporal chunk after prepending the prior feature cache."""
|
||||
padding = list(self._padding)
|
||||
if cache_x is not None and self._padding[4] > 0:
|
||||
cache_x = cache_x.to(x.device)
|
||||
x = torch.cat([cache_x, x], dim=2)
|
||||
padding[4] -= cache_x.shape[2]
|
||||
x = F.pad(x, padding)
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
"""Normalize channel vectors with the Wan encoder's root-mean-square rule."""
|
||||
|
||||
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False):
|
||||
"""Create the broadcastable scale and optional bias parameters."""
|
||||
super().__init__()
|
||||
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
||||
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
||||
|
||||
self.channel_first = channel_first
|
||||
self.scale = dim**0.5
|
||||
self.gamma = nn.Parameter(torch.ones(shape))
|
||||
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
|
||||
|
||||
@torch.compile
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize the channel dimension and apply the learned scale."""
|
||||
dim = 1 if self.channel_first else -1
|
||||
return F.normalize(x, dim=dim) * self.scale * self.gamma + self.bias
|
||||
|
||||
|
||||
class Resample(nn.Module):
|
||||
"""Downsample spatial dimensions and selected temporal transitions."""
|
||||
|
||||
def __init__(self, dim: int, mode: str):
|
||||
"""Build the spatial convolution and optional temporal convolution."""
|
||||
if mode not in ("downsample2d", "downsample3d"):
|
||||
raise ValueError(f"Unsupported MAGI-2 Wan encoder resampling mode: {mode}")
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mode = mode
|
||||
self.resample = nn.Sequential(
|
||||
nn.ZeroPad2d((0, 1, 0, 1)),
|
||||
nn.Conv2d(dim, dim, 3, stride=(2, 2)),
|
||||
)
|
||||
if mode == "downsample3d":
|
||||
self.time_conv = CausalConv3d(
|
||||
dim,
|
||||
dim,
|
||||
(3, 1, 1),
|
||||
stride=(2, 1, 1),
|
||||
padding=(0, 0, 0),
|
||||
)
|
||||
|
||||
@torch.compile
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
feat_cache: list[torch.Tensor | None] | None = None,
|
||||
feat_idx: list[int] = [0], # noqa: B006 - compiled nested calls share this cache cursor.
|
||||
) -> torch.Tensor:
|
||||
"""Downsample one chunk and update its temporal feature cache."""
|
||||
batch_size, channels, time, height, width = x.size()
|
||||
x = rearrange(x, "b c t h w -> (b t) c h w")
|
||||
x = self.resample(x)
|
||||
x = rearrange(x, "(b t) c h w -> b c t h w", b=batch_size, t=time)
|
||||
|
||||
if self.mode == "downsample3d" and feat_cache is not None:
|
||||
index = feat_idx[0]
|
||||
if feat_cache[index] is None:
|
||||
feat_cache[index] = x.clone()
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
cache_x = x[:, :, -1:, :, :].clone()
|
||||
x = self.time_conv(torch.cat([feat_cache[index][:, :, -1:, :, :], x], 2))
|
||||
feat_cache[index] = cache_x
|
||||
feat_idx[0] += 1
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
"""Apply two cached causal convolutions and an additive shortcut."""
|
||||
|
||||
def __init__(self, in_dim: int, out_dim: int, dropout: float = 0.0):
|
||||
"""Build the normalization, convolution, and shortcut modules."""
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
self.residual = nn.Sequential(
|
||||
RMSNorm(in_dim, images=False),
|
||||
nn.SiLU(),
|
||||
CausalConv3d(in_dim, out_dim, 3, padding=1),
|
||||
RMSNorm(out_dim, images=False),
|
||||
nn.SiLU(),
|
||||
nn.Dropout(dropout),
|
||||
CausalConv3d(out_dim, out_dim, 3, padding=1),
|
||||
)
|
||||
self.shortcut = CausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity()
|
||||
|
||||
@torch.compile
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
feat_cache: list[torch.Tensor | None] | None = None,
|
||||
feat_idx: list[int] = [0], # noqa: B006 - compiled nested calls share this cache cursor.
|
||||
) -> torch.Tensor:
|
||||
"""Run the residual path while preserving the causal convolution caches."""
|
||||
shortcut = self.shortcut(x)
|
||||
for layer in self.residual:
|
||||
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
||||
index = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[index] is not None:
|
||||
cache_x = torch.cat(
|
||||
[feat_cache[index][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
|
||||
dim=2,
|
||||
)
|
||||
x = layer(x, feat_cache[index])
|
||||
feat_cache[index] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x + shortcut
|
||||
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
"""Apply single-head spatial self-attention independently per frame."""
|
||||
|
||||
def __init__(self, dim: int):
|
||||
"""Build normalization, fused query-key-value projection, and output projection."""
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.norm = RMSNorm(dim)
|
||||
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
||||
self.proj = nn.Conv2d(dim, dim, 1)
|
||||
nn.init.zeros_(self.proj.weight)
|
||||
|
||||
@torch.compile
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Attend over the spatial positions of each frame and add the residual."""
|
||||
identity = x
|
||||
batch_size, channels, time, height, width = x.size()
|
||||
x = rearrange(x, "b c t h w -> (b t) c h w")
|
||||
x = self.norm(x)
|
||||
query, key, value = (
|
||||
self.to_qkv(x)
|
||||
.reshape(batch_size * time, 1, channels * 3, -1)
|
||||
.permute(0, 1, 3, 2)
|
||||
.contiguous()
|
||||
.chunk(3, dim=-1)
|
||||
)
|
||||
x = F.scaled_dot_product_attention(query, key, value)
|
||||
x = x.squeeze(1).permute(0, 2, 1).reshape(batch_size * time, channels, height, width)
|
||||
x = self.proj(x)
|
||||
x = rearrange(x, "(b t) c h w -> b c t h w", t=time)
|
||||
return x + identity
|
||||
|
||||
|
||||
def patchify(x: torch.Tensor, patch_size: int) -> torch.Tensor:
|
||||
"""Move each spatial patch into the channel dimension."""
|
||||
if patch_size == 1:
|
||||
return x
|
||||
if x.dim() == 4:
|
||||
return rearrange(x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
|
||||
if x.dim() == 5:
|
||||
return rearrange(x, "b c f (h q) (w r) -> b (c r q) f h w", q=patch_size, r=patch_size)
|
||||
raise ValueError(f"Invalid MAGI-2 Wan image shape: {tuple(x.shape)}")
|
||||
|
||||
|
||||
class AvgDown3D(nn.Module):
|
||||
"""Average channel groups after spatial and temporal rearrangement."""
|
||||
|
||||
def __init__(self, in_channels: int, out_channels: int, factor_t: int, factor_s: int = 1):
|
||||
"""Record the factors and channel-group size for the shortcut path."""
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.factor_t = factor_t
|
||||
self.factor_s = factor_s
|
||||
self.factor = self.factor_t * self.factor_s * self.factor_s
|
||||
|
||||
if in_channels * self.factor % out_channels != 0:
|
||||
raise ValueError("MAGI-2 Wan shortcut channels must divide evenly")
|
||||
self.group_size = in_channels * self.factor // out_channels
|
||||
|
||||
@torch.compile
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Rearrange a video into channel groups and average each group."""
|
||||
pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
|
||||
x = F.pad(x, (0, 0, 0, 0, pad_t, 0))
|
||||
batch_size, channels, time, height, width = x.shape
|
||||
x = x.view(
|
||||
batch_size,
|
||||
channels,
|
||||
time // self.factor_t,
|
||||
self.factor_t,
|
||||
height // self.factor_s,
|
||||
self.factor_s,
|
||||
width // self.factor_s,
|
||||
self.factor_s,
|
||||
)
|
||||
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
|
||||
x = x.view(
|
||||
batch_size,
|
||||
channels * self.factor,
|
||||
time // self.factor_t,
|
||||
height // self.factor_s,
|
||||
width // self.factor_s,
|
||||
)
|
||||
x = x.view(
|
||||
batch_size,
|
||||
self.out_channels,
|
||||
self.group_size,
|
||||
time // self.factor_t,
|
||||
height // self.factor_s,
|
||||
width // self.factor_s,
|
||||
)
|
||||
return x.mean(dim=2)
|
||||
|
||||
|
||||
class DownResidualBlock(nn.Module):
|
||||
"""Combine residual encoder blocks with an averaged downsample shortcut."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_dim: int,
|
||||
out_dim: int,
|
||||
dropout: float,
|
||||
mult: int,
|
||||
temporal_downsample: bool = False,
|
||||
down_flag: bool = False,
|
||||
):
|
||||
"""Build the main and shortcut paths for one encoder resolution."""
|
||||
super().__init__()
|
||||
self.avg_shortcut = AvgDown3D(
|
||||
in_dim,
|
||||
out_dim,
|
||||
factor_t=2 if temporal_downsample else 1,
|
||||
factor_s=2 if down_flag else 1,
|
||||
)
|
||||
downsamples: list[nn.Module] = []
|
||||
for _ in range(mult):
|
||||
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
||||
in_dim = out_dim
|
||||
if down_flag:
|
||||
mode = "downsample3d" if temporal_downsample else "downsample2d"
|
||||
downsamples.append(Resample(out_dim, mode=mode))
|
||||
self.downsamples = nn.Sequential(*downsamples)
|
||||
|
||||
@torch.compile
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
feat_cache: list[torch.Tensor | None] | None = None,
|
||||
feat_idx: list[int] = [0], # noqa: B006 - compiled nested calls share this cache cursor.
|
||||
) -> torch.Tensor:
|
||||
"""Run one resolution block and combine its two paths."""
|
||||
shortcut_input = x.clone()
|
||||
for module in self.downsamples:
|
||||
x = module(x, feat_cache, feat_idx)
|
||||
return x + self.avg_shortcut(shortcut_input)
|
||||
|
||||
|
||||
class Encoder3d(nn.Module):
|
||||
"""Encode patchified conditioning images into Gaussian moments."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
z_dim: int,
|
||||
dim_mult: tuple[int, ...],
|
||||
num_res_blocks: int,
|
||||
attn_scales: tuple[float, ...],
|
||||
temporal_downsample: tuple[bool, ...],
|
||||
dropout: float,
|
||||
):
|
||||
"""Build the published MAGI-2 Wan image-encoder hierarchy."""
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temporal_downsample = temporal_downsample
|
||||
|
||||
dims = [dim * multiplier for multiplier in [1, *dim_mult]]
|
||||
self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
|
||||
downsamples: list[nn.Module] = []
|
||||
for index, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)):
|
||||
temporal_downsample_flag = temporal_downsample[index] if index < len(temporal_downsample) else False
|
||||
downsamples.append(
|
||||
DownResidualBlock(
|
||||
in_dim=in_dim,
|
||||
out_dim=out_dim,
|
||||
dropout=dropout,
|
||||
mult=num_res_blocks,
|
||||
temporal_downsample=temporal_downsample_flag,
|
||||
down_flag=index != len(dim_mult) - 1,
|
||||
)
|
||||
)
|
||||
self.downsamples = nn.Sequential(*downsamples)
|
||||
self.middle = nn.Sequential(
|
||||
ResidualBlock(out_dim, out_dim, dropout),
|
||||
AttentionBlock(out_dim),
|
||||
ResidualBlock(out_dim, out_dim, dropout),
|
||||
)
|
||||
self.head = nn.Sequential(
|
||||
RMSNorm(out_dim, images=False),
|
||||
nn.SiLU(),
|
||||
CausalConv3d(out_dim, z_dim, 3, padding=1),
|
||||
)
|
||||
|
||||
@torch.compile
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
feat_cache: list[torch.Tensor | None] | None = None,
|
||||
feat_idx: list[int] = [0], # noqa: B006 - compiled nested calls share this cache cursor.
|
||||
) -> torch.Tensor:
|
||||
"""Encode one temporal chunk while updating every causal cache slot."""
|
||||
if feat_cache is not None:
|
||||
index = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[index] is not None:
|
||||
cache_x = torch.cat(
|
||||
[feat_cache[index][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
|
||||
dim=2,
|
||||
)
|
||||
x = self.conv1(x, feat_cache[index])
|
||||
feat_cache[index] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = self.conv1(x)
|
||||
|
||||
for layer in self.downsamples:
|
||||
x = layer(x, feat_cache, feat_idx) if feat_cache is not None else layer(x)
|
||||
for layer in self.middle:
|
||||
if isinstance(layer, ResidualBlock) and feat_cache is not None:
|
||||
x = layer(x, feat_cache, feat_idx)
|
||||
else:
|
||||
x = layer(x)
|
||||
for layer in self.head:
|
||||
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
||||
index = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[index] is not None:
|
||||
cache_x = torch.cat(
|
||||
[feat_cache[index][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
|
||||
dim=2,
|
||||
)
|
||||
x = layer(x, feat_cache[index])
|
||||
feat_cache[index] = cache_x
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
|
||||
def count_causal_convolutions(model: nn.Module) -> int:
|
||||
"""Count cache slots required by the encoder's causal convolutions."""
|
||||
return sum(isinstance(module, CausalConv3d) for module in model.modules())
|
||||
|
||||
|
||||
class Magi2WanImageEncoder(nn.Module):
|
||||
"""Encode MAGI-2 I2V reference images into normalized 48-channel latents."""
|
||||
|
||||
def __init__(self, config: Magi2WanVAEConfig):
|
||||
"""Build the encoder and quantization convolution with source checkpoint names."""
|
||||
super().__init__()
|
||||
self.z_dim = config.z_dim
|
||||
self.patch_size = config.patch_size
|
||||
self.encoder = Encoder3d(
|
||||
dim=config.base_dim,
|
||||
z_dim=config.z_dim * 2,
|
||||
dim_mult=tuple(config.dim_mult),
|
||||
num_res_blocks=config.num_res_blocks,
|
||||
attn_scales=tuple(config.attn_scales),
|
||||
temporal_downsample=tuple(config.temperal_downsample),
|
||||
dropout=config.dropout,
|
||||
)
|
||||
self.conv1 = CausalConv3d(config.z_dim * 2, config.z_dim * 2, 1)
|
||||
self.register_buffer("mean", torch.tensor(config.latents_mean, dtype=torch.float32), persistent=False)
|
||||
self.register_buffer("std", torch.tensor(config.latents_std, dtype=torch.float32), persistent=False)
|
||||
self.register_buffer("inverse_std", 1.0 / self.std, persistent=False)
|
||||
self.clear_cache()
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""Reset every causal feature cache before and after an image batch."""
|
||||
self._enc_conv_num = count_causal_convolutions(self.encoder)
|
||||
self._enc_conv_idx = [0]
|
||||
self._enc_feat_map: list[torch.Tensor | None] = [None] * self._enc_conv_num
|
||||
|
||||
def set_normalization(self, mean: tuple[float, ...], std: tuple[float, ...]) -> None:
|
||||
"""Create normalization tensors on the encoder device with FP32 arithmetic."""
|
||||
device = self.conv1.weight.device
|
||||
self.mean = torch.tensor(mean, dtype=torch.float32, device=device)
|
||||
self.std = torch.tensor(std, dtype=torch.float32, device=device)
|
||||
self.inverse_std = 1.0 / self.std
|
||||
|
||||
def encode(self, video: torch.Tensor) -> torch.Tensor:
|
||||
"""Encode FP32 video produced by the pipeline's BF16 image round trip."""
|
||||
if video.dtype != torch.float32:
|
||||
raise TypeError(f"MAGI-2 Wan encoder input must be torch.float32, received {video.dtype}")
|
||||
self.clear_cache()
|
||||
video = patchify(video, patch_size=self.patch_size)
|
||||
time = video.shape[2]
|
||||
chunk_count = 1 + (time - 1) // 4
|
||||
for chunk_index in range(chunk_count):
|
||||
self._enc_conv_idx = [0]
|
||||
if chunk_index == 0:
|
||||
encoded = self.encoder(
|
||||
video[:, :, :1, :, :],
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx,
|
||||
)
|
||||
else:
|
||||
encoded_chunk = self.encoder(
|
||||
video[:, :, 1 + 4 * (chunk_index - 1) : 1 + 4 * chunk_index, :, :],
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx,
|
||||
)
|
||||
encoded = torch.cat([encoded, encoded_chunk], 2)
|
||||
latent_mean, _ = self.conv1(encoded).chunk(2, dim=1)
|
||||
latent = (latent_mean - self.mean.view(1, self.z_dim, 1, 1, 1)) * self.inverse_std.view(
|
||||
1,
|
||||
self.z_dim,
|
||||
1,
|
||||
1,
|
||||
1,
|
||||
)
|
||||
self.clear_cache()
|
||||
return latent.float()
|
||||
|
||||
def forward(self, video: torch.Tensor) -> torch.Tensor:
|
||||
"""Encode one pipeline-ready reference-image tensor."""
|
||||
return self.encode(video)
|
||||
|
||||
|
||||
def _is_decoder_tensor(source_name: str) -> bool:
|
||||
"""Identify checkpoint tensors that belong only to video decoding."""
|
||||
return source_name.startswith("decoder.") or source_name.startswith("conv2.")
|
||||
|
||||
|
||||
def load_magi2_wan_image_encoder(
|
||||
checkpoint_path: str | Path,
|
||||
device: torch.device | str,
|
||||
) -> Magi2WanImageEncoder:
|
||||
"""Build and strictly load every MAGI-2 Wan encoder and quant tensor."""
|
||||
config = Magi2WanVAEConfig()
|
||||
with torch.device("meta"):
|
||||
model = Magi2WanImageEncoder(config)
|
||||
target_state = model.state_dict()
|
||||
source_state = torch.load(
|
||||
Path(checkpoint_path),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
mmap=True,
|
||||
)
|
||||
encoder_state: dict[str, torch.Tensor] = {}
|
||||
for source_name, source_tensor in source_state.items():
|
||||
if _is_decoder_tensor(source_name):
|
||||
continue
|
||||
if source_name not in target_state:
|
||||
raise KeyError(f"Unexpected MAGI-2 Wan encoder checkpoint tensor: {source_name}")
|
||||
target_tensor = target_state[source_name]
|
||||
if source_tensor.shape != target_tensor.shape or source_tensor.dtype != target_tensor.dtype:
|
||||
raise ValueError(
|
||||
f"MAGI-2 Wan encoder tensor metadata differs for {source_name}: "
|
||||
f"source={tuple(source_tensor.shape)}/{source_tensor.dtype}, "
|
||||
f"target={tuple(target_tensor.shape)}/{target_tensor.dtype}"
|
||||
)
|
||||
encoder_state[source_name] = source_tensor
|
||||
missing_names = sorted(set(target_state) - set(encoder_state))
|
||||
if missing_names:
|
||||
raise RuntimeError(f"MAGI-2 Wan encoder checkpoint is missing tensors: {missing_names}")
|
||||
incompatible = model.load_state_dict(encoder_state, strict=True, assign=True)
|
||||
if incompatible.missing_keys or incompatible.unexpected_keys:
|
||||
raise RuntimeError(
|
||||
"MAGI-2 Wan encoder strict load failed: "
|
||||
f"missing={incompatible.missing_keys}, unexpected={incompatible.unexpected_keys}"
|
||||
)
|
||||
model.mean = torch.empty(len(config.latents_mean), dtype=torch.float32)
|
||||
model.std = torch.empty(len(config.latents_std), dtype=torch.float32)
|
||||
model.inverse_std = torch.empty(len(config.latents_std), dtype=torch.float32)
|
||||
model.to(device=device, dtype=torch.float32)
|
||||
model.set_normalization(config.latents_mean, config.latents_std)
|
||||
model.requires_grad_(False)
|
||||
return model.eval()
|
||||
|
||||
|
||||
__all__ = ["Magi2WanImageEncoder", "load_magi2_wan_image_encoder"]
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MAGI-2 Preview pipeline package."""
|
||||
@@ -0,0 +1,265 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo pipeline for MAGI-2 Preview text-to-video and image-to-video."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages import (
|
||||
Magi2AudioDecodingStage,
|
||||
Magi2DataProxyConfig,
|
||||
Magi2InputValidationStage,
|
||||
Magi2LatentPreparationStage,
|
||||
Magi2LatentSavingStage,
|
||||
Magi2PreviewDenoisingStage,
|
||||
Magi2ReferenceImageStage,
|
||||
Magi2RefinerDataProxyConfig,
|
||||
Magi2RefinerStage,
|
||||
Magi2TextEncodingStage,
|
||||
Magi2VideoDecodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
|
||||
|
||||
def _configure_deterministic_kernels(seed: int) -> None:
|
||||
"""Enable the official deterministic attention, MoE, and PyTorch paths."""
|
||||
os.environ["MAGI2_DETERMINISTIC"] = "1"
|
||||
os.environ["MAGI_ATTENTION_DETERMINISTIC_MODE"] = "1"
|
||||
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
torch.use_deterministic_algorithms(True)
|
||||
|
||||
|
||||
class Magi2Pipeline(ComposedPipelineBase):
|
||||
"""Generate a 10-second video and stereo audio with MAGI-2 Preview."""
|
||||
|
||||
is_video_pipeline = True
|
||||
_required_config_modules = [
|
||||
"transformer",
|
||||
"transformer_2",
|
||||
"text_encoder",
|
||||
"image_encoder",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_path: str,
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs,
|
||||
required_config_modules: list[str] | None = None,
|
||||
loaded_modules: dict[str, torch.nn.Module] | None = None,
|
||||
) -> None:
|
||||
"""Configure deterministic kernels before distributed model imports."""
|
||||
deterministic_environment = os.environ.get(
|
||||
"MAGI2_DETERMINISTIC",
|
||||
"0",
|
||||
) == "1"
|
||||
if fastvideo_args.deterministic or deterministic_environment:
|
||||
_configure_deterministic_kernels(seed=42)
|
||||
fastvideo_args.dit_layerwise_offload = False
|
||||
fastvideo_args.dit_cpu_offload = True
|
||||
super().__init__(
|
||||
model_path=model_path,
|
||||
fastvideo_args=fastvideo_args,
|
||||
required_config_modules=required_config_modules,
|
||||
loaded_modules=loaded_modules,
|
||||
)
|
||||
|
||||
def load_modules(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
loaded_modules: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Strictly load official components while keeping decoder copies on one rank."""
|
||||
self._load_config(self.model_path)
|
||||
from fastvideo.configs.models.encoders import Magi2Qwen35Config
|
||||
from fastvideo.configs.models.vaes import Magi2TurboVAEConfig
|
||||
from fastvideo.models.dits.magi2_loader import (
|
||||
load_magi2_preview_model,
|
||||
load_magi2_refiner_model,
|
||||
)
|
||||
from fastvideo.models.dits.magi2_runtime.fastvideo_parallel import (
|
||||
bind_fastvideo_parallel_state,
|
||||
)
|
||||
from fastvideo.models.encoders.qwen3_5 import Magi2Qwen35TextEncoder
|
||||
from fastvideo.models.vaes.magi2_audio_vae import load_magi2_audio_vae
|
||||
from fastvideo.models.vaes.magi2_turbo_vae import Magi2TurboVAEModel
|
||||
from fastvideo.models.vaes.magi2_wan_loader import (
|
||||
load_magi2_wan_image_encoder,
|
||||
)
|
||||
|
||||
bind_fastvideo_parallel_state()
|
||||
checkpoint_root = Path(self.model_path)
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
provided_modules = loaded_modules or {}
|
||||
modules: dict[str, Any] = {}
|
||||
|
||||
modules["transformer"] = provided_modules.get("transformer")
|
||||
if modules["transformer"] is None:
|
||||
modules["transformer"] = load_magi2_preview_model(
|
||||
str(checkpoint_root / "transformer"),
|
||||
pipeline_config.dit_config,
|
||||
"cpu",
|
||||
)
|
||||
modules["transformer_2"] = provided_modules.get("transformer_2")
|
||||
if modules["transformer_2"] is None:
|
||||
modules["transformer_2"] = load_magi2_refiner_model(
|
||||
str(checkpoint_root / "transformer_2"),
|
||||
pipeline_config.refiner_dit_config,
|
||||
"cpu",
|
||||
)
|
||||
|
||||
is_decode_rank = psm.is_group_first_rank("cp")
|
||||
modules["text_encoder"] = provided_modules.get("text_encoder")
|
||||
modules["image_encoder"] = provided_modules.get("image_encoder")
|
||||
modules["vae"] = provided_modules.get("vae")
|
||||
modules["audio_vae"] = provided_modules.get("audio_vae")
|
||||
if is_decode_rank and modules["text_encoder"] is None:
|
||||
modules["text_encoder"] = Magi2Qwen35TextEncoder.from_pretrained_local(
|
||||
str(checkpoint_root / "text_encoder"),
|
||||
Magi2Qwen35Config(),
|
||||
torch.bfloat16,
|
||||
torch.device("cpu"),
|
||||
)
|
||||
if is_decode_rank and modules["image_encoder"] is None:
|
||||
modules["image_encoder"] = load_magi2_wan_image_encoder(
|
||||
checkpoint_root / "image_encoder" / "Wan2.2_VAE.pth",
|
||||
"cpu",
|
||||
)
|
||||
if is_decode_rank and modules["vae"] is None:
|
||||
turbo_config = Magi2TurboVAEConfig(
|
||||
config_path=str(
|
||||
checkpoint_root
|
||||
/ "vae"
|
||||
/ "TurboV3-Wan22-TinyShallow_7_7.json"
|
||||
),
|
||||
checkpoint_path=str(checkpoint_root / "vae" / "checkpoint.ckpt"),
|
||||
pretrained_dtype="bfloat16",
|
||||
)
|
||||
modules["vae"] = Magi2TurboVAEModel(turbo_config).to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
if is_decode_rank and modules["audio_vae"] is None:
|
||||
modules["audio_vae"] = load_magi2_audio_vae(
|
||||
checkpoint_root / "audio_vae",
|
||||
"cpu",
|
||||
)
|
||||
modules["scheduler"] = provided_modules.get(
|
||||
"scheduler",
|
||||
FlowUniPCMultistepScheduler(),
|
||||
)
|
||||
return modules
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Compose conditioning, denoising, refinement, and decoding stages."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
self.add_stage(
|
||||
"input_validation_stage",
|
||||
Magi2InputValidationStage(
|
||||
output_frames=pipeline_config.output_frames,
|
||||
output_height=pipeline_config.output_height,
|
||||
output_width=pipeline_config.output_width,
|
||||
output_fps=pipeline_config.output_fps,
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"reference_image_stage",
|
||||
Magi2ReferenceImageStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
preview_height=pipeline_config.preview_height,
|
||||
preview_width=pipeline_config.preview_width,
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"text_encoding_stage",
|
||||
Magi2TextEncodingStage(self.get_module("text_encoder")),
|
||||
)
|
||||
self.add_stage(
|
||||
"latent_preparation_stage",
|
||||
Magi2LatentPreparationStage(
|
||||
video_channels=pipeline_config.preview_video_channels,
|
||||
video_length=pipeline_config.preview_video_length,
|
||||
video_height=pipeline_config.preview_latent_height,
|
||||
video_width=pipeline_config.preview_latent_width,
|
||||
audio_length=pipeline_config.audio_latent_length,
|
||||
audio_channels=pipeline_config.audio_channels,
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"preview_denoising_stage",
|
||||
Magi2PreviewDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
data_proxy_config=Magi2DataProxyConfig(),
|
||||
flow_shift=pipeline_config.preview_flow_shift,
|
||||
video_guidance_scale=pipeline_config.preview_video_guidance_scale,
|
||||
audio_guidance_scale=pipeline_config.preview_audio_guidance_scale,
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
"refiner_stage",
|
||||
Magi2RefinerStage(
|
||||
transformer=self.get_module("transformer_2"),
|
||||
data_proxy_config=Magi2RefinerDataProxyConfig(
|
||||
t_patch_size=1,
|
||||
patch_size=1,
|
||||
frame_receptive_field=11,
|
||||
spatial_rope_interpolation="extra",
|
||||
text_offset=0,
|
||||
coords_style="v1",
|
||||
attn_config={
|
||||
"mode": "window",
|
||||
"block_t_size": 8,
|
||||
"block_size": 4,
|
||||
"window": {
|
||||
"level": "block",
|
||||
"block_mode": "grid",
|
||||
"block_t_radius": 2,
|
||||
"block_h_radius": 2,
|
||||
"block_w_radius": 2,
|
||||
"win_size": 384,
|
||||
"frame_receptive_field": -1,
|
||||
"auto_range_merge": True,
|
||||
"sparse_load": False,
|
||||
"full_attn_layers": [],
|
||||
},
|
||||
},
|
||||
magi2_refiner_condition_input="none",
|
||||
),
|
||||
latent_height=pipeline_config.refiner_latent_height,
|
||||
latent_width=pipeline_config.refiner_latent_width,
|
||||
noise_index=pipeline_config.refiner_noise_index,
|
||||
flow_shift=pipeline_config.refiner_flow_shift,
|
||||
video_guidance_scale=pipeline_config.refiner_video_guidance_scale,
|
||||
audio_guidance_scale=pipeline_config.refiner_audio_guidance_scale,
|
||||
audio_channels=pipeline_config.audio_channels,
|
||||
),
|
||||
)
|
||||
self.add_stage("latent_saving_stage", Magi2LatentSavingStage())
|
||||
self.add_stage(
|
||||
"video_decoding_stage",
|
||||
Magi2VideoDecodingStage(self.get_module("vae")),
|
||||
)
|
||||
self.add_stage(
|
||||
"audio_decoding_stage",
|
||||
Magi2AudioDecodingStage(self.get_module("audio_vae")),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = Magi2Pipeline
|
||||
|
||||
__all__ = ["Magi2Pipeline"]
|
||||
@@ -0,0 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline-local export for the MAGI-2 configuration class."""
|
||||
|
||||
from fastvideo.configs.pipelines.magi2 import Magi2PreviewPipelineConfig
|
||||
|
||||
__all__ = ["Magi2PreviewPipelineConfig"]
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Inference preset for the published MAGI-2 Preview 1080p profile."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
MAGI2_NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, "
|
||||
"overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly "
|
||||
"drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy "
|
||||
"background, three legs, many people in the background, walking backwards, low quality, worst quality, poor "
|
||||
"quality, noise, background noise, hiss, hum, buzz, crackle, static, compression artifacts, MP3 artifacts, "
|
||||
"digital clipping, distortion, muffled, muddy, unclear, echo, reverb, room echo, over-reverberated, hollow "
|
||||
"sound, distant, washed out, harsh, shrill, piercing, grating, tinny, thin sound, boomy, bass-heavy, flat EQ, "
|
||||
"over-compressed, abrupt cut, jarring transition, sudden silence, looping artifact, music, instrumental, "
|
||||
"sirens, alarms, crowd noise, unrelated sound effects, chaotic, disorganized, messy, cheap sound, emotionless, "
|
||||
"flat delivery, deadpan, lifeless, apathetic, robotic, mechanical, monotone, flat intonation, undynamic, boring, "
|
||||
"reading from a script, AI voice, synthetic, text-to-speech, TTS, insincere, fake emotion, exaggerated, overly "
|
||||
"dramatic, melodramatic, cheesy, cringey, hesitant, unconfident, tired, weak voice, stuttering, stammering, "
|
||||
"mumbling, slurred speech, mispronounced, bad articulation, lisp, vocal fry, creaky voice, mouth clicks, lip "
|
||||
"smacks, wet mouth sounds, heavy breathing, audible inhales, plosives, p-pops, coughing, clearing throat, "
|
||||
"sneezing, speaking too fast, rushed, speaking too slow, dragged out, unnatural pauses, awkward silence, choppy, "
|
||||
"disconnected, multiple speakers, two voices, background talking, out of tune, off-key, autotune artifacts")
|
||||
|
||||
_PREVIEW_STAGE = PresetStageSpec(
|
||||
name="preview",
|
||||
kind="denoising",
|
||||
description="Joint video and audio preview denoising pass.",
|
||||
allowed_overrides=frozenset({"num_inference_steps"}),
|
||||
)
|
||||
|
||||
_REFINER_STAGE = PresetStageSpec(
|
||||
name="refiner",
|
||||
kind="refinement",
|
||||
description="Spatial-temporal 1080p latent refinement pass.",
|
||||
allowed_overrides=frozenset({"num_inference_steps_sr"}),
|
||||
)
|
||||
|
||||
MAGI2_PREVIEW_1080P = InferencePreset(
|
||||
name="magi2_preview_1080p",
|
||||
version=1,
|
||||
model_family="magi2",
|
||||
description="MAGI-2 Preview text- or image-conditioned 10-second 1080p video with stereo audio.",
|
||||
workload_type=None,
|
||||
stage_schemas=(_PREVIEW_STAGE, _REFINER_STAGE),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 1088,
|
||||
"width": 1920,
|
||||
"num_frames": 249,
|
||||
"fps": 25,
|
||||
"num_inference_steps": 100,
|
||||
"num_inference_steps_sr": 5,
|
||||
"guidance_scale": 5.0,
|
||||
"negative_prompt": MAGI2_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (MAGI2_PREVIEW_1080P, )
|
||||
|
||||
__all__ = ["ALL_PRESETS", "MAGI2_NEGATIVE_PROMPT", "MAGI2_PREVIEW_1080P"]
|
||||
@@ -0,0 +1,62 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MAGI-2 pipeline components."""
|
||||
|
||||
from fastvideo.pipelines.basic.magi2.stages.audio_decoding import (
|
||||
MAGI2_AUDIO_TIME_STRETCH,
|
||||
Magi2AudioDecodingStage,
|
||||
decode_magi2_audio,
|
||||
resample_magi2_audio,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.conditioning import (
|
||||
Magi2ReferenceImageStage,
|
||||
Magi2TextEncodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.output import (
|
||||
Magi2LatentSavingStage,
|
||||
Magi2VideoDecodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.preview_data_proxy import (
|
||||
Magi2DataProxy,
|
||||
Magi2DataProxyConfig,
|
||||
ModelInput,
|
||||
Modality,
|
||||
VarlenHandler,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.refiner_data_proxy import (
|
||||
Magi2RefinerDataProxy,
|
||||
Magi2RefinerDataProxyConfig,
|
||||
RefinerModelInput,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.runtime import (
|
||||
Magi2InputValidationStage,
|
||||
Magi2LatentPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.sampling import (
|
||||
Magi2PreviewDenoisingStage,
|
||||
Magi2RefinerStage,
|
||||
ZeroSNRDDPMDiscretization,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"MAGI2_AUDIO_TIME_STRETCH",
|
||||
"Magi2AudioDecodingStage",
|
||||
"Magi2DataProxy",
|
||||
"Magi2DataProxyConfig",
|
||||
"Magi2InputValidationStage",
|
||||
"Magi2LatentPreparationStage",
|
||||
"Magi2LatentSavingStage",
|
||||
"Magi2PreviewDenoisingStage",
|
||||
"Magi2ReferenceImageStage",
|
||||
"Magi2RefinerStage",
|
||||
"Magi2RefinerDataProxy",
|
||||
"Magi2RefinerDataProxyConfig",
|
||||
"Magi2TextEncodingStage",
|
||||
"Magi2VideoDecodingStage",
|
||||
"ModelInput",
|
||||
"Modality",
|
||||
"RefinerModelInput",
|
||||
"VarlenHandler",
|
||||
"ZeroSNRDDPMDiscretization",
|
||||
"decode_magi2_audio",
|
||||
"resample_magi2_audio",
|
||||
]
|
||||
@@ -0,0 +1,91 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio decoding for the MAGI-2 joint video-audio pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scipy.signal import resample
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
MAGI2_AUDIO_TIME_STRETCH = 441.0 / 512.0
|
||||
|
||||
|
||||
def resample_magi2_audio(
|
||||
sample_major_audio: np.ndarray,
|
||||
time_stretching: float = MAGI2_AUDIO_TIME_STRETCH,
|
||||
) -> np.ndarray:
|
||||
"""Resample sample-major stereo audio with MAGI-2's FFT interpolation."""
|
||||
output_length = int(sample_major_audio.shape[0] * time_stretching)
|
||||
return resample(sample_major_audio, output_length)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def decode_magi2_audio(audio_vae, sample_major_latent: torch.Tensor) -> np.ndarray:
|
||||
"""Decode a ``[latent samples, channels]`` tensor into stereo samples."""
|
||||
waveform = audio_vae.decode(sample_major_latent.T)
|
||||
sample_major_audio = waveform.squeeze(0).T.cpu().numpy()
|
||||
return resample_magi2_audio(sample_major_audio)
|
||||
|
||||
|
||||
class Magi2AudioDecodingStage(PipelineStage):
|
||||
"""Decode preview audio latents on the context-parallel leader rank."""
|
||||
|
||||
def __init__(self, audio_vae) -> None:
|
||||
"""Store the published Stable Audio decoder."""
|
||||
super().__init__()
|
||||
self.audio_vae = audio_vae
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Accept the audio latent that the joint denoising stage produced."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default stage verification record."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Decode and store sample-major audio on the output-producing rank."""
|
||||
if batch.audio_latents is None:
|
||||
raise ValueError("MAGI-2 audio decoding requires audio_latents")
|
||||
if not psm.is_group_first_rank("cp"):
|
||||
return batch
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
self.audio_vae.to(device=device, dtype=torch.float32)
|
||||
sample_major_latent = batch.audio_latents.squeeze(0).to(device)
|
||||
batch.extra["audio"] = decode_magi2_audio(
|
||||
self.audio_vae,
|
||||
sample_major_latent,
|
||||
)
|
||||
batch.extra["audio_sample_rate"] = int(self.audio_vae.sampling_rate)
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.audio_vae.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MAGI2_AUDIO_TIME_STRETCH",
|
||||
"Magi2AudioDecodingStage",
|
||||
"decode_magi2_audio",
|
||||
"resample_magi2_audio",
|
||||
]
|
||||
@@ -0,0 +1,238 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Reference-image and Qwen3.5 conditioning stages for MAGI-2."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils import load_image
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.distributed import get_sp_group
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
def _resizepad(image: Image.Image, target_height: int, target_width: int) -> Image.Image:
|
||||
"""Fit an RGB image inside a white canvas with the official letterboxing rule."""
|
||||
width, height = image.size
|
||||
if width <= 0 or height <= 0:
|
||||
raise ValueError(
|
||||
f"MAGI-2 received an invalid image size: width={width}, height={height}"
|
||||
)
|
||||
scale = min(target_width / width, target_height / height)
|
||||
resized_width = max(1, int(round(width * scale)))
|
||||
resized_height = max(1, int(round(height * scale)))
|
||||
resized = image.convert("RGB").resize(
|
||||
(resized_width, resized_height),
|
||||
resample=Image.Resampling.LANCZOS,
|
||||
)
|
||||
canvas = Image.new("RGB", (target_width, target_height), (255, 255, 255))
|
||||
canvas.paste(
|
||||
resized,
|
||||
(
|
||||
(target_width - resized_width) // 2,
|
||||
(target_height - resized_height) // 2,
|
||||
),
|
||||
)
|
||||
return canvas
|
||||
|
||||
|
||||
def _ensure_figure_reference(prompt: str) -> str:
|
||||
"""Address the single I2V image as ``<Figure 1>`` in a plain or JSON prompt."""
|
||||
try:
|
||||
prompt_object = json.loads(prompt)
|
||||
if not isinstance(prompt_object, dict):
|
||||
raise ValueError("prompt JSON is not an object")
|
||||
prompt_object["reference_layer"] = [
|
||||
"The first frame refers to <Figure 1>"
|
||||
]
|
||||
return json.dumps(prompt_object, ensure_ascii=False)
|
||||
except (json.JSONDecodeError, TypeError, ValueError):
|
||||
return prompt + "reference_layer:The first frame refers to <Figure 1>"
|
||||
|
||||
|
||||
class Magi2ReferenceImageStage(PipelineStage):
|
||||
"""Encode one I2V reference image on the context-parallel leader."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_encoder: Any | None,
|
||||
preview_height: int,
|
||||
preview_width: int,
|
||||
) -> None:
|
||||
"""Store the Wan encoder and the published preview resolution."""
|
||||
super().__init__()
|
||||
self.image_encoder = image_encoder
|
||||
self.preview_height = preview_height
|
||||
self.preview_width = preview_width
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=32)
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default validation record; ``forward`` gives precise errors."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default validation record for T2V and I2V outputs."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Encode and broadcast the official FP32 Wan posterior mean for I2V."""
|
||||
if fastvideo_args.workload_type != WorkloadType.I2V:
|
||||
batch.magi2_ref_image_feat = None
|
||||
batch.magi2_ref_image_feat_len = None
|
||||
batch.magi2_ref_image_special_tokens = None
|
||||
return batch
|
||||
|
||||
is_leader = psm.is_group_first_rank("cp")
|
||||
payload: dict[str, torch.Tensor] | None = None
|
||||
if is_leader:
|
||||
if self.image_encoder is None:
|
||||
raise RuntimeError("MAGI-2 I2V requires the Wan image encoder")
|
||||
image = batch.pil_image
|
||||
if image is None and batch.image_path is not None:
|
||||
image = load_image(batch.image_path)
|
||||
if not isinstance(image, Image.Image):
|
||||
raise TypeError("MAGI-2 I2V requires one PIL image or image path")
|
||||
|
||||
maximum_length = max(self.preview_height, self.preview_width)
|
||||
if image.width > image.height:
|
||||
target_width = maximum_length
|
||||
target_height = int(image.height * maximum_length / image.width)
|
||||
else:
|
||||
target_height = maximum_length
|
||||
target_width = int(image.width * maximum_length / image.height)
|
||||
image = _resizepad(image, target_height, target_width)
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
image_tensor = self.video_processor.preprocess(
|
||||
image,
|
||||
height=target_height,
|
||||
width=target_width,
|
||||
)
|
||||
image_tensor = image_tensor.to(device=device, dtype=torch.bfloat16)
|
||||
image_tensor = image_tensor.unsqueeze(2)[:, :3]
|
||||
self.image_encoder.to(device=device, dtype=torch.float32)
|
||||
reference_latent = self.image_encoder.encode(image_tensor.float())
|
||||
latent_height = int(reference_latent.shape[-2])
|
||||
latent_width = int(reference_latent.shape[-1])
|
||||
payload = {
|
||||
"reference_latent": reference_latent.unsqueeze(1),
|
||||
"reference_length": torch.tensor(
|
||||
[[[latent_height, latent_width]]],
|
||||
device=device,
|
||||
dtype=torch.long,
|
||||
),
|
||||
}
|
||||
if fastvideo_args.image_encoder_cpu_offload:
|
||||
self.image_encoder.to("cpu")
|
||||
|
||||
broadcast_payload = get_sp_group().broadcast_tensor_dict(payload, src=0)
|
||||
if broadcast_payload is None:
|
||||
raise RuntimeError("MAGI-2 reference-image broadcast returned no payload")
|
||||
batch.magi2_ref_image_feat = broadcast_payload["reference_latent"]
|
||||
batch.magi2_ref_image_feat_len = broadcast_payload["reference_length"]
|
||||
if not isinstance(batch.prompt, str):
|
||||
raise TypeError("MAGI-2 accepts one prompt string per request")
|
||||
batch.prompt = _ensure_figure_reference(batch.prompt)
|
||||
if is_leader and fastvideo_args.image_encoder_cpu_offload:
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
class Magi2TextEncodingStage(PipelineStage):
|
||||
"""Encode positive and negative prompts on the context-parallel leader."""
|
||||
|
||||
def __init__(self, text_encoder: Any | None) -> None:
|
||||
"""Store the native Qwen3.5 encoder loaded from the release checkpoint."""
|
||||
super().__init__()
|
||||
self.text_encoder = text_encoder
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record; prompt validation occurs in ``forward``."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for the broadcast conditioning tensors."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Produce Qwen3.5 hidden states and the optional figure-token embedding."""
|
||||
if not isinstance(batch.prompt, str):
|
||||
raise TypeError("MAGI-2 accepts one prompt string per request")
|
||||
if not isinstance(batch.negative_prompt, str):
|
||||
raise TypeError("MAGI-2 requires one negative prompt string")
|
||||
|
||||
is_leader = psm.is_group_first_rank("cp")
|
||||
payload: dict[str, torch.Tensor | None] | None = None
|
||||
if is_leader:
|
||||
if self.text_encoder is None:
|
||||
raise RuntimeError("MAGI-2 requires the Qwen3.5 text encoder")
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
self.text_encoder.to(device=device)
|
||||
text_context = self.text_encoder.encode(batch.prompt).to(torch.bfloat16)
|
||||
negative_context = self.text_encoder.encode(batch.negative_prompt).to(
|
||||
torch.bfloat16
|
||||
)
|
||||
special_tokens = None
|
||||
if batch.magi2_ref_image_feat is not None:
|
||||
special_tokens = self.text_encoder.get_special_token(
|
||||
batch.prompt,
|
||||
["<Figure 1>"],
|
||||
text_context,
|
||||
).unsqueeze(0)
|
||||
payload = {
|
||||
"text_context": text_context,
|
||||
"negative_context": negative_context,
|
||||
"special_tokens": special_tokens,
|
||||
}
|
||||
if fastvideo_args.text_encoder_cpu_offload:
|
||||
self.text_encoder.to("cpu")
|
||||
|
||||
broadcast_payload = get_sp_group().broadcast_tensor_dict(payload, src=0)
|
||||
if broadcast_payload is None:
|
||||
raise RuntimeError("MAGI-2 text-conditioning broadcast returned no payload")
|
||||
batch.magi2_text_context = broadcast_payload["text_context"]
|
||||
batch.magi2_negative_context = broadcast_payload["negative_context"]
|
||||
batch.magi2_ref_image_special_tokens = broadcast_payload["special_tokens"]
|
||||
if is_leader and fastvideo_args.text_encoder_cpu_offload:
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = ["Magi2ReferenceImageStage", "Magi2TextEncodingStage"]
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Post-refiner latent capture and Turbo VAE decoding for MAGI-2."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.dits.magi2_runtime import psm
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class Magi2LatentSavingStage(PipelineStage):
|
||||
"""Write each post-refiner video latent on the output-producing rank."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the per-process sample counter used in latent filenames."""
|
||||
super().__init__()
|
||||
self.sample_index = 0
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for the post-refiner latent."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record after the optional filesystem write."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Save ``latent_N.pt`` when ``MAGI2_SAVE_LATENT_PATH`` is configured."""
|
||||
del fastvideo_args
|
||||
latent_directory = envs.MAGI2_SAVE_LATENT_PATH
|
||||
if not latent_directory or not psm.is_group_first_rank("cp"):
|
||||
return batch
|
||||
if batch.latents is None:
|
||||
raise ValueError("MAGI-2 latent saving requires a post-refiner latent")
|
||||
output_directory = Path(latent_directory)
|
||||
output_directory.mkdir(parents=True, exist_ok=True)
|
||||
output_path = output_directory / f"latent_{self.sample_index}.pt"
|
||||
torch.save(batch.latents.detach().cpu(), output_path)
|
||||
self.sample_index += 1
|
||||
return batch
|
||||
|
||||
|
||||
class Magi2VideoDecodingStage(PipelineStage):
|
||||
"""Decode the 1080p latent with Turbo VAE on one rank per video."""
|
||||
|
||||
performance_component_metric = "vae_decode_time_s"
|
||||
|
||||
def __init__(self, turbo_vae: Any | None) -> None:
|
||||
"""Store the distilled Turbo VAE decoder loaded from ``ckpt/turbo_vae``."""
|
||||
super().__init__()
|
||||
self.turbo_vae = turbo_vae
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record; ``forward`` validates the video latent."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record because non-leader ranks intentionally skip decode."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Decode and store a CPU float video tensor in ``[B,C,T,H,W]`` layout."""
|
||||
if not psm.is_group_first_rank("cp"):
|
||||
batch.output = None
|
||||
return batch
|
||||
if batch.latents is None:
|
||||
raise ValueError("MAGI-2 Turbo VAE decoding requires a video latent")
|
||||
if self.turbo_vae is None:
|
||||
raise RuntimeError("MAGI-2 requires the distilled Turbo VAE decoder")
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
self.turbo_vae.to(device=device, dtype=torch.bfloat16)
|
||||
decoder_input = batch.latents.squeeze(0).to(torch.bfloat16)
|
||||
if decoder_input.dim() == 4:
|
||||
decoder_input = decoder_input.unsqueeze(0)
|
||||
decoded_video = self.turbo_vae.decode(decoder_input).float()
|
||||
batch.output = decoded_video.mul(0.5).add(0.5).clamp(0, 1).cpu()
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.turbo_vae.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = ["Magi2LatentSavingStage", "Magi2VideoDecodingStage"]
|
||||
@@ -0,0 +1,892 @@
|
||||
# Copyright (c) 2026 SandAI. All Rights Reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pack MAGI-2 preview video, audio, text, and image-conditioning tokens."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from itertools import chain
|
||||
from typing import Any, Literal
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from einops import rearrange
|
||||
from torch.nn import functional as F
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime.psm import psm
|
||||
|
||||
|
||||
class Modality(IntEnum):
|
||||
"""Modality identifiers consumed by the MAGI-2 preview transformer."""
|
||||
|
||||
VIDEO = 0
|
||||
AUDIO = 1
|
||||
TEXT = 2
|
||||
TIME = 3
|
||||
|
||||
|
||||
@dataclass
|
||||
class VarlenHandler:
|
||||
"""Variable-length attention boundaries for a packed token sequence."""
|
||||
|
||||
cu_seqlens_q: torch.Tensor
|
||||
cu_seqlens_k: torch.Tensor
|
||||
max_seqlen_q: int
|
||||
max_seqlen_k: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Magi2DataProxyConfig:
|
||||
"""Configuration that controls MAGI-2 preview token packing."""
|
||||
|
||||
t_patch_size: int = 1
|
||||
patch_size: int = 1
|
||||
spatial_rope_interpolation: Literal["inter", "extra"] = "extra"
|
||||
add_time_token: bool = False
|
||||
time_channel_dim: int = 64
|
||||
time_aligned_rope: bool = False
|
||||
audio_latent_fps: float = 25.0
|
||||
time_pos_fps: float = 3.125
|
||||
vae_first_latent_is_image: bool = True
|
||||
video_fps: float = 25.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelInput:
|
||||
"""Multimodal tensors and lengths consumed by ``Magi2DataProxy``."""
|
||||
|
||||
x_t: torch.Tensor
|
||||
audio_x_t: torch.Tensor
|
||||
audio_feat_len: torch.Tensor | list[int]
|
||||
txt_feat: torch.Tensor
|
||||
txt_feat_len: torch.Tensor | list[int]
|
||||
t: torch.Tensor
|
||||
ref_audio_feat: torch.Tensor | None = None
|
||||
ref_audio_feat_len: torch.Tensor | list[int] | None = None
|
||||
ref_video_feat: torch.Tensor | None = None
|
||||
ref_video_feat_len: torch.Tensor | list[int] | None = None
|
||||
per_token_video_t: torch.Tensor | None = None
|
||||
per_token_audio_t: torch.Tensor | None = None
|
||||
ref_image_feat: torch.Tensor | None = None
|
||||
ref_image_feat_len: torch.Tensor | None = None
|
||||
ref_image_special_token_embedding: torch.Tensor | None = None
|
||||
|
||||
|
||||
def _to_int(value: int | torch.Tensor) -> int:
|
||||
"""Convert a scalar tensor or Python integer into an integer."""
|
||||
if isinstance(value, torch.Tensor):
|
||||
return int(value.detach().reshape(-1)[0].item())
|
||||
return int(value)
|
||||
|
||||
|
||||
def _pad_cat(
|
||||
tensors: list[torch.Tensor],
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Right-pad feature channels and concatenate token segments."""
|
||||
if not tensors:
|
||||
return torch.empty(0, 0, device=device, dtype=dtype)
|
||||
max_channel = max(tensor.shape[-1] for tensor in tensors)
|
||||
return torch.cat(
|
||||
[
|
||||
F.pad(
|
||||
tensor.to(device=device, dtype=dtype),
|
||||
(0, max_channel - tensor.shape[-1]),
|
||||
)
|
||||
for tensor in tensors
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
|
||||
def _segment_paint(
|
||||
values: list[int],
|
||||
seqlens: list[int],
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
output_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""Expand one modality identifier over each packed token segment."""
|
||||
mapping = torch.empty(output_size, dtype=dtype, device=device)
|
||||
offset = 0
|
||||
for value, seqlen in zip(values, seqlens, strict=True):
|
||||
mapping[offset:offset + seqlen] = int(value)
|
||||
offset += seqlen
|
||||
return mapping
|
||||
|
||||
|
||||
def _seqlens2cu_seqlens(
|
||||
seqlens: list[int],
|
||||
device: torch.device | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Convert sequence lengths into int32 cumulative boundaries."""
|
||||
seqlens_tensor = torch.tensor(seqlens, dtype=torch.int32, device=device)
|
||||
return F.pad(torch.cumsum(seqlens_tensor, dim=0), (1, 0))
|
||||
|
||||
|
||||
def _ceil_div(dividend: int, divisor: int) -> int:
|
||||
"""Return integer ceiling division."""
|
||||
return (dividend + divisor - 1) // divisor
|
||||
|
||||
|
||||
def _len_to_list(value: torch.Tensor) -> list[int]:
|
||||
"""Convert one packed reference-image grid shape into Python integers."""
|
||||
return [int(entry) for entry in value.detach().to(torch.long).reshape(-1).tolist()]
|
||||
|
||||
|
||||
def _get_coords(
|
||||
shape: tuple[int, int, int],
|
||||
ref_feat_shape: tuple[int, int, int],
|
||||
offset_thw: tuple[int, int, int] = (0, 0, 0),
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
time_positions: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Build token coordinates with source-grid and reference-grid metadata."""
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
ori_t, ori_h, ori_w = shape
|
||||
ref_t, ref_h, ref_w = ref_feat_shape
|
||||
offset_t, offset_h, offset_w = offset_thw
|
||||
if time_positions is None:
|
||||
time_range = torch.arange(ori_t, device=device, dtype=dtype) + offset_t
|
||||
else:
|
||||
time_range = time_positions.to(device=device, dtype=dtype) + offset_t
|
||||
height_range = torch.arange(ori_h, device=device, dtype=dtype) + offset_h
|
||||
width_range = torch.arange(ori_w, device=device, dtype=dtype) + offset_w
|
||||
time_grid, height_grid, width_grid = torch.meshgrid(
|
||||
time_range,
|
||||
height_range,
|
||||
width_range,
|
||||
indexing="ij",
|
||||
)
|
||||
coords_grid = torch.stack([time_grid, height_grid, width_grid], dim=-1)
|
||||
coords_flat = coords_grid.reshape(-1, 3)
|
||||
metadata = torch.tensor(
|
||||
[ori_t, ori_h, ori_w, ref_t, ref_h, ref_w],
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
return torch.cat([coords_flat, metadata.expand(coords_flat.size(0), -1)], dim=-1)
|
||||
|
||||
|
||||
def _sinusoidal_embedding_1d(
|
||||
dim: int,
|
||||
position: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Encode normalized diffusion times into cosine-sine feature channels."""
|
||||
position = position.to(torch.float32) * 1000.0
|
||||
half = dim // 2
|
||||
frequencies = torch.exp(
|
||||
-math.log(10000)
|
||||
* torch.arange(
|
||||
start=0,
|
||||
end=half,
|
||||
dtype=torch.float32,
|
||||
device=position.device,
|
||||
)
|
||||
/ half
|
||||
)
|
||||
arguments = position[:, None].float() * frequencies[None]
|
||||
embedding = torch.cat([torch.cos(arguments), torch.sin(arguments)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])],
|
||||
dim=-1,
|
||||
)
|
||||
return embedding
|
||||
|
||||
|
||||
@dataclass
|
||||
class SingleData:
|
||||
"""One sample's token segments and packing metadata."""
|
||||
|
||||
video_x_t: torch.Tensor
|
||||
audio_x_t: torch.Tensor
|
||||
audio_feat_len: int
|
||||
txt_feat: torch.Tensor
|
||||
txt_feat_len: int
|
||||
t: int
|
||||
h: int
|
||||
w: int
|
||||
patch_size: int
|
||||
t_patch_size: int
|
||||
spatial_rope_interpolation: Literal["inter", "extra"]
|
||||
diffusion_t: torch.Tensor | None = None
|
||||
per_token_video_t: torch.Tensor | None = None
|
||||
per_token_audio_t: torch.Tensor | None = None
|
||||
time_channel_dim: int = 0
|
||||
vae_first_latent_is_image: bool = True
|
||||
video_fps: float = 25.0
|
||||
time_pos_fps: float = 3.125
|
||||
ref_image_feats: list[torch.Tensor] | None = None
|
||||
ref_image_feat_lens: list[list[int]] | None = None
|
||||
ref_image_special_tokens: list[torch.Tensor] | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Trim variable-length segments and derive immutable packing sizes."""
|
||||
self.video_token_num = self.video_x_t.shape[0]
|
||||
self.origin_audio_feat_len = self.audio_x_t.shape[0]
|
||||
self.audio_x_t = self.audio_x_t[:self.audio_feat_len]
|
||||
self.txt_feat = self.txt_feat[:self.txt_feat_len]
|
||||
if self.per_token_audio_t is not None:
|
||||
self.per_token_audio_t = self.per_token_audio_t[:self.audio_feat_len]
|
||||
|
||||
self.ref_image_feats = self.ref_image_feats or []
|
||||
self.ref_image_feat_lens = self.ref_image_feat_lens or []
|
||||
self.ref_image_special_tokens = self.ref_image_special_tokens or []
|
||||
self.ref_image_token_nums = [
|
||||
int(math.prod(feat_len)) for feat_len in self.ref_image_feat_lens
|
||||
]
|
||||
self.ref_image_feats = [
|
||||
feat[:token_num]
|
||||
for feat, token_num in zip(
|
||||
self.ref_image_feats,
|
||||
self.ref_image_token_nums,
|
||||
strict=True,
|
||||
)
|
||||
]
|
||||
self.num_ref_images = len(self.ref_image_feats)
|
||||
self.total_ref_image_feat_len = sum(self.ref_image_token_nums)
|
||||
self.video_channel = self.video_x_t.shape[-1]
|
||||
self.audio_channel = self.audio_x_t.shape[-1]
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Return the device that owns this sample's packed tokens."""
|
||||
return self.video_x_t.device
|
||||
|
||||
@property
|
||||
def default_dtype(self) -> torch.dtype:
|
||||
"""Return the video-token dtype used for the packed sequence."""
|
||||
return self.video_x_t.dtype
|
||||
|
||||
@property
|
||||
def add_time_token(self) -> bool:
|
||||
"""Return whether this sample appends a standalone time token."""
|
||||
return self.diffusion_t is not None
|
||||
|
||||
@property
|
||||
def total_token_num(self) -> int:
|
||||
"""Return the complete token count for one sample."""
|
||||
total = self.video_token_num + self.audio_feat_len + self.txt_feat_len
|
||||
total += self.total_ref_image_feat_len + self.num_ref_images
|
||||
return total + (1 if self.add_time_token else 0)
|
||||
|
||||
@property
|
||||
def feat_to_cat(self) -> list[torch.Tensor]:
|
||||
"""Return token segments in transformer consumption order."""
|
||||
tensors = [self.video_x_t, self.audio_x_t, self.txt_feat]
|
||||
for image_index in range(self.num_ref_images):
|
||||
tensors.append(self.ref_image_special_tokens[image_index])
|
||||
tensors.append(self.ref_image_feats[image_index])
|
||||
if self.add_time_token:
|
||||
assert self.diffusion_t is not None
|
||||
tensors.append(
|
||||
self.diffusion_t.to(
|
||||
device=self.device,
|
||||
dtype=self.default_dtype,
|
||||
).reshape(1, 1)
|
||||
)
|
||||
return tensors
|
||||
|
||||
@property
|
||||
def token_sequence(self) -> torch.Tensor:
|
||||
"""Return channel-padded tokens for one sample."""
|
||||
return _pad_cat(self.feat_to_cat, self.device, self.default_dtype)
|
||||
|
||||
@property
|
||||
def modality_map_seqlens(self) -> tuple[list[int], list[int]]:
|
||||
"""Return segment lengths and modality identifiers in token order."""
|
||||
seqlens = [self.video_token_num, self.audio_feat_len, self.txt_feat_len]
|
||||
modalities: list[int] = [
|
||||
int(Modality.VIDEO),
|
||||
int(Modality.AUDIO),
|
||||
int(Modality.TEXT),
|
||||
]
|
||||
for image_index in range(self.num_ref_images):
|
||||
seqlens.append(1)
|
||||
modalities.append(int(Modality.TEXT))
|
||||
seqlens.append(self.ref_image_token_nums[image_index])
|
||||
modalities.append(int(Modality.VIDEO))
|
||||
if self.add_time_token:
|
||||
seqlens.append(1)
|
||||
modalities.append(int(Modality.TIME))
|
||||
return seqlens, modalities
|
||||
|
||||
@property
|
||||
def modality_mapping(self) -> torch.Tensor:
|
||||
"""Expand segment modality identifiers to one value per token."""
|
||||
seqlens, modalities = self.modality_map_seqlens
|
||||
return _segment_paint(
|
||||
modalities,
|
||||
seqlens,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
output_size=self.total_token_num,
|
||||
)
|
||||
|
||||
def _default_coords(
|
||||
self,
|
||||
shape: tuple[int, int, int],
|
||||
ref_feat_shape: tuple[int, int, int],
|
||||
offset_thw: tuple[int, int, int] = (0, 0, 0),
|
||||
time_positions: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Build coordinate rows on this sample's device and dtype."""
|
||||
return _get_coords(
|
||||
shape,
|
||||
ref_feat_shape,
|
||||
offset_thw=offset_thw,
|
||||
device=self.device,
|
||||
dtype=self.default_dtype,
|
||||
time_positions=time_positions,
|
||||
)
|
||||
|
||||
@property
|
||||
def coords_to_cat(self) -> list[torch.Tensor]:
|
||||
"""Return coordinate segments in the same order as feature segments."""
|
||||
t_steps = self.t // self.t_patch_size
|
||||
h_steps = self.h // self.patch_size
|
||||
w_steps = self.w // self.patch_size
|
||||
if self.spatial_rope_interpolation == "inter":
|
||||
video_h_ref, video_w_ref = 32, 32
|
||||
else:
|
||||
video_h_ref, video_w_ref = h_steps, w_steps
|
||||
|
||||
video_coords = self._default_coords(
|
||||
(t_steps, h_steps, w_steps),
|
||||
(t_steps, video_h_ref, video_w_ref),
|
||||
)
|
||||
magic_audio_ref_t = (self.audio_feat_len - 1) // 8 + 1
|
||||
audio_coords = self._default_coords(
|
||||
(self.audio_feat_len, 1, 1),
|
||||
(magic_audio_ref_t // self.t_patch_size, 1, 1),
|
||||
)
|
||||
coords = [
|
||||
video_coords,
|
||||
audio_coords,
|
||||
self._default_coords(
|
||||
(self.txt_feat_len, 1, 1),
|
||||
(1, 1, 1),
|
||||
offset_thw=(-self.txt_feat_len, 0, 0),
|
||||
),
|
||||
]
|
||||
|
||||
for image_index in range(self.num_ref_images):
|
||||
token_len = self.ref_image_token_nums[image_index]
|
||||
if len(self.ref_image_feat_lens[image_index]) >= 2:
|
||||
image_h = int(self.ref_image_feat_lens[image_index][0])
|
||||
image_w = int(self.ref_image_feat_lens[image_index][1])
|
||||
else:
|
||||
image_h = image_w = int(math.ceil(math.sqrt(token_len)))
|
||||
time_offset = t_steps + 2 + image_index
|
||||
coords.append(
|
||||
torch.tensor(
|
||||
[
|
||||
[
|
||||
time_offset,
|
||||
-1,
|
||||
-1,
|
||||
1,
|
||||
image_h,
|
||||
image_w,
|
||||
1,
|
||||
image_h,
|
||||
image_w,
|
||||
]
|
||||
],
|
||||
device=self.device,
|
||||
dtype=self.default_dtype,
|
||||
)
|
||||
)
|
||||
coords.append(
|
||||
self._default_coords(
|
||||
(1, image_h, image_w),
|
||||
(1, image_h, image_w),
|
||||
offset_thw=(time_offset, 0, 0),
|
||||
)[:token_len]
|
||||
)
|
||||
|
||||
if self.add_time_token:
|
||||
coords.append(self._default_coords((1, 1, 1), (1, 1, 1))[:1])
|
||||
return coords
|
||||
|
||||
@property
|
||||
def coords_mapping(self) -> torch.Tensor:
|
||||
"""Return one nine-value coordinate row per token."""
|
||||
return torch.cat(self.coords_to_cat, dim=0)
|
||||
|
||||
@property
|
||||
def time_token_sequence(self) -> torch.Tensor:
|
||||
"""Return per-token diffusion-time features in feature-segment order."""
|
||||
if self.time_channel_dim == 0:
|
||||
return torch.empty(self.total_token_num, 0, device=self.device)
|
||||
assert self.per_token_video_t is not None
|
||||
assert self.per_token_audio_t is not None
|
||||
time_parts = [
|
||||
self.per_token_video_t.squeeze(-1),
|
||||
self.per_token_audio_t.squeeze(-1),
|
||||
torch.zeros(self.txt_feat_len, device=self.device),
|
||||
]
|
||||
for image_index in range(self.num_ref_images):
|
||||
time_parts.append(torch.zeros(1, device=self.device))
|
||||
time_parts.append(
|
||||
torch.zeros(
|
||||
self.ref_image_token_nums[image_index],
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
if self.add_time_token:
|
||||
assert self.diffusion_t is not None
|
||||
time_parts.append(self.diffusion_t.reshape(1).to(self.device))
|
||||
raw_time = torch.cat(time_parts, dim=0)
|
||||
if self.time_channel_dim == 1:
|
||||
return raw_time.unsqueeze(-1)
|
||||
return _sinusoidal_embedding_1d(self.time_channel_dim, raw_time)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SimplePackedData:
|
||||
"""A batch represented as adjacent variable-length sample sequences."""
|
||||
|
||||
items: list[SingleData]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Cache token-length summaries shared by attention and depacking."""
|
||||
if len(self.items) == 0:
|
||||
raise ValueError("SimplePackedData must contain at least one item.")
|
||||
self._total_token_num_list = [item.total_token_num for item in self.items]
|
||||
self._total_token_num_sum = sum(self._total_token_num_list)
|
||||
self._total_token_num_max = max(self._total_token_num_list)
|
||||
self._total_token_cu_seqlens = _seqlens2cu_seqlens(
|
||||
self._total_token_num_list
|
||||
)
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Return the device used by every packed sample."""
|
||||
return self.items[0].device
|
||||
|
||||
@property
|
||||
def default_dtype(self) -> torch.dtype:
|
||||
"""Return the token dtype used by every packed sample."""
|
||||
return self.items[0].default_dtype
|
||||
|
||||
@property
|
||||
def token_sequence(self) -> torch.Tensor:
|
||||
"""Return every sample's feature segments as one packed tensor."""
|
||||
feature_segments = list(
|
||||
chain.from_iterable(item.feat_to_cat for item in self.items)
|
||||
)
|
||||
return _pad_cat(feature_segments, self.device, self.default_dtype)
|
||||
|
||||
@property
|
||||
def modality_mapping(self) -> torch.Tensor:
|
||||
"""Return one modality identifier per token across all samples."""
|
||||
seqlens: list[int] = []
|
||||
modalities: list[int] = []
|
||||
for item in self.items:
|
||||
item_seqlens, item_modalities = item.modality_map_seqlens
|
||||
seqlens.extend(item_seqlens)
|
||||
modalities.extend(item_modalities)
|
||||
return _segment_paint(
|
||||
modalities,
|
||||
seqlens,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
output_size=self.total_token_num,
|
||||
)
|
||||
|
||||
@property
|
||||
def coords_mapping(self) -> torch.Tensor:
|
||||
"""Return coordinate rows across all packed samples."""
|
||||
coordinate_segments = list(
|
||||
chain.from_iterable(item.coords_to_cat for item in self.items)
|
||||
)
|
||||
return torch.cat(coordinate_segments, dim=0)
|
||||
|
||||
@property
|
||||
def time_token_sequence(self) -> torch.Tensor:
|
||||
"""Return diffusion-time features across all packed samples."""
|
||||
return torch.cat(
|
||||
[item.time_token_sequence for item in self.items],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
@property
|
||||
def total_token_num(self) -> int:
|
||||
"""Return the total number of unpadded tokens."""
|
||||
return self._total_token_num_sum
|
||||
|
||||
@property
|
||||
def cu_seqlen(self) -> torch.Tensor:
|
||||
"""Return cumulative token boundaries for each packed sample."""
|
||||
return self._total_token_cu_seqlens.clone()
|
||||
|
||||
@property
|
||||
def max_seqlen(self) -> int:
|
||||
"""Return the longest unpadded sample sequence."""
|
||||
return self._total_token_num_max
|
||||
|
||||
def __getitem__(self, index: int) -> SingleData:
|
||||
"""Return one packed sample descriptor."""
|
||||
return self.items[index]
|
||||
|
||||
def depack_token_sequence(
|
||||
self,
|
||||
token_sequence: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Restore video and audio tensors from model-output token segments."""
|
||||
videos: list[torch.Tensor] = []
|
||||
audios: list[torch.Tensor] = []
|
||||
token_counts = [item.total_token_num for item in self.items]
|
||||
for item, token_slice in zip(
|
||||
self.items,
|
||||
torch.split(token_sequence, token_counts, dim=0),
|
||||
strict=True,
|
||||
):
|
||||
video_flat = token_slice[:item.video_token_num, :item.video_channel]
|
||||
output_channel = item.video_channel // (
|
||||
item.t_patch_size * item.patch_size * item.patch_size
|
||||
)
|
||||
video = rearrange(
|
||||
video_flat,
|
||||
"(T H W) (pT pH pW C) -> C (T pT) (H pH) (W pW)",
|
||||
T=item.t // item.t_patch_size,
|
||||
H=item.h // item.patch_size,
|
||||
W=item.w // item.patch_size,
|
||||
pT=item.t_patch_size,
|
||||
pH=item.patch_size,
|
||||
pW=item.patch_size,
|
||||
C=output_channel,
|
||||
).contiguous()
|
||||
audio = torch.zeros(
|
||||
item.origin_audio_feat_len,
|
||||
item.audio_channel,
|
||||
device=token_sequence.device,
|
||||
dtype=token_sequence.dtype,
|
||||
)
|
||||
audio[:item.audio_feat_len] = token_slice[
|
||||
item.video_token_num:item.video_token_num + item.audio_feat_len,
|
||||
:item.audio_channel,
|
||||
]
|
||||
videos.append(video)
|
||||
audios.append(audio)
|
||||
return torch.stack(videos, dim=0), torch.stack(audios, dim=0)
|
||||
|
||||
|
||||
class Magi2DataProxy:
|
||||
"""Pack MAGI-2 preview inputs and restore preview-model outputs."""
|
||||
|
||||
def __init__(self, config: Magi2DataProxyConfig) -> None:
|
||||
self.config = config
|
||||
self.patch_size = config.patch_size
|
||||
self.t_patch_size = config.t_patch_size
|
||||
self._saved_data: dict[str, Any] = {}
|
||||
|
||||
def saved_for_output(self, **kwargs: Any) -> None:
|
||||
"""Save packing metadata that is required to depack model outputs."""
|
||||
self._saved_data.update(kwargs)
|
||||
|
||||
def get_saved_data(self, key: str) -> Any:
|
||||
"""Return saved packing metadata by name."""
|
||||
return self._saved_data[key]
|
||||
|
||||
def _reduce_max_token_num_for_ep_cp(self, total_token_num: int) -> int:
|
||||
"""Find the maximum token count across active expert and context groups."""
|
||||
if not (dist.is_available() and dist.is_initialized()):
|
||||
return total_token_num
|
||||
device = (
|
||||
torch.device("cuda", torch.cuda.current_device())
|
||||
if torch.cuda.is_available()
|
||||
else torch.device("cpu")
|
||||
)
|
||||
local_token_count = torch.tensor(
|
||||
[total_token_num],
|
||||
device=device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
maximum_counts = [local_token_count]
|
||||
ep_size = psm.get_world_size("ep")
|
||||
if ep_size > 1:
|
||||
ep_maximum = local_token_count.clone()
|
||||
dist.all_reduce(
|
||||
ep_maximum,
|
||||
op=dist.ReduceOp.MAX,
|
||||
group=psm.get_parallel_group("ep"),
|
||||
)
|
||||
maximum_counts.append(ep_maximum)
|
||||
cp_size = psm.get_world_size("cp")
|
||||
if cp_size > 1:
|
||||
cp_maximum = local_token_count.clone()
|
||||
dist.all_reduce(
|
||||
cp_maximum,
|
||||
op=dist.ReduceOp.MAX,
|
||||
group=psm.get_parallel_group("cp"),
|
||||
)
|
||||
maximum_counts.append(cp_maximum)
|
||||
return max(int(maximum.item()) for maximum in maximum_counts)
|
||||
|
||||
def _pad_for_ep_cp(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
coords_mapping: torch.Tensor,
|
||||
modality_mapping: torch.Tensor,
|
||||
time_token_sequence: torch.Tensor,
|
||||
varlen_handler: VarlenHandler,
|
||||
align_to: int = 48,
|
||||
) -> tuple[
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
VarlenHandler,
|
||||
int,
|
||||
]:
|
||||
"""Pad distributed packed inputs to a shared multiple-of-48 length."""
|
||||
cp_size = max(1, psm.get_world_size("cp"))
|
||||
ep_size = max(1, psm.get_world_size("ep"))
|
||||
if cp_size <= 1 and ep_size <= 1:
|
||||
return (
|
||||
x,
|
||||
coords_mapping,
|
||||
modality_mapping,
|
||||
time_token_sequence,
|
||||
varlen_handler,
|
||||
0,
|
||||
)
|
||||
|
||||
total_token_num = x.shape[0]
|
||||
target_size = self._reduce_max_token_num_for_ep_cp(total_token_num)
|
||||
padded_size = _ceil_div(target_size, align_to) * align_to
|
||||
pad_size = padded_size - total_token_num
|
||||
if pad_size <= 0:
|
||||
return (
|
||||
x,
|
||||
coords_mapping,
|
||||
modality_mapping,
|
||||
time_token_sequence,
|
||||
varlen_handler,
|
||||
0,
|
||||
)
|
||||
|
||||
x = F.pad(x, (0, 0, 0, pad_size), value=0)
|
||||
pad_coords = _get_coords(
|
||||
shape=(pad_size, 1, 1),
|
||||
ref_feat_shape=(pad_size, 1, 1),
|
||||
offset_thw=(0, 0, 0),
|
||||
device=x.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
coords_mapping = torch.cat([coords_mapping, pad_coords], dim=0)
|
||||
modality_mapping = F.pad(
|
||||
modality_mapping,
|
||||
(0, pad_size),
|
||||
value=int(Modality.TEXT),
|
||||
)
|
||||
if time_token_sequence.shape[-1] == 0:
|
||||
time_token_sequence = time_token_sequence.new_zeros(padded_size, 0)
|
||||
else:
|
||||
time_token_sequence = F.pad(
|
||||
time_token_sequence,
|
||||
(0, 0, 0, pad_size),
|
||||
value=0,
|
||||
)
|
||||
padded_boundary = torch.tensor(
|
||||
[padded_size],
|
||||
device=x.device,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
varlen_handler = VarlenHandler(
|
||||
cu_seqlens_q=torch.cat(
|
||||
[varlen_handler.cu_seqlens_q, padded_boundary]
|
||||
),
|
||||
cu_seqlens_k=torch.cat(
|
||||
[varlen_handler.cu_seqlens_k, padded_boundary]
|
||||
),
|
||||
max_seqlen_q=max(varlen_handler.max_seqlen_q, pad_size),
|
||||
max_seqlen_k=max(varlen_handler.max_seqlen_k, pad_size),
|
||||
)
|
||||
return (
|
||||
x,
|
||||
coords_mapping,
|
||||
modality_mapping,
|
||||
time_token_sequence,
|
||||
varlen_handler,
|
||||
pad_size,
|
||||
)
|
||||
|
||||
def img2tokens(self, x_t: torch.Tensor) -> torch.Tensor:
|
||||
"""Extract non-overlapping 3D patches with channel-major features."""
|
||||
kernel_size = (self.t_patch_size, self.patch_size, self.patch_size)
|
||||
if not all(
|
||||
size >= kernel
|
||||
for size, kernel in zip(x_t.shape[2:], kernel_size, strict=True)
|
||||
):
|
||||
return torch.empty(
|
||||
x_t.shape[0],
|
||||
0,
|
||||
x_t.shape[1] * math.prod(kernel_size),
|
||||
device=x_t.device,
|
||||
dtype=x_t.dtype,
|
||||
)
|
||||
patches = (
|
||||
x_t.unfold(2, self.t_patch_size, self.t_patch_size)
|
||||
.unfold(3, self.patch_size, self.patch_size)
|
||||
.unfold(4, self.patch_size, self.patch_size)
|
||||
)
|
||||
batch_size, channels, t_steps, h_steps, w_steps, _, _, _ = patches.shape
|
||||
return (
|
||||
patches.permute(0, 2, 3, 4, 1, 5, 6, 7)
|
||||
.reshape(batch_size, t_steps * h_steps * w_steps, -1)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
def process_input(
|
||||
self,
|
||||
data: ModelInput,
|
||||
) -> tuple[
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
VarlenHandler,
|
||||
torch.Tensor,
|
||||
]:
|
||||
"""Pack a batch without allowing conditional samples to attend each other."""
|
||||
batch_size, input_video_channel, video_t, video_h, video_w = data.x_t.shape
|
||||
video_tokens = self.img2tokens(data.x_t)
|
||||
audio_tokens = data.audio_x_t.contiguous()
|
||||
text_tokens = data.txt_feat.contiguous()
|
||||
|
||||
per_token_video_tokens = None
|
||||
per_token_audio = None
|
||||
if self.config.time_channel_dim > 0 and data.per_token_video_t is not None:
|
||||
per_token_video_tokens = self.img2tokens(data.per_token_video_t)[:, :, :1]
|
||||
per_token_audio = data.per_token_audio_t
|
||||
|
||||
ref_image_data = data.ref_image_feat
|
||||
ref_image_feat_len = data.ref_image_feat_len
|
||||
ref_image_special_tokens = data.ref_image_special_token_embedding
|
||||
has_ref_image = (
|
||||
ref_image_data is not None
|
||||
and ref_image_data.ndim >= 5
|
||||
and ref_image_feat_len is not None
|
||||
and ref_image_feat_len.ndim >= 2
|
||||
and ref_image_special_tokens is not None
|
||||
and ref_image_special_tokens.ndim >= 3
|
||||
)
|
||||
num_ref_images = ref_image_data.shape[1] if has_ref_image else 0
|
||||
ref_device = text_tokens.device
|
||||
ref_dtype = text_tokens.dtype
|
||||
|
||||
items: list[SingleData] = []
|
||||
for batch_index in range(batch_size):
|
||||
image_features: list[torch.Tensor] = []
|
||||
image_feature_lengths: list[list[int]] = []
|
||||
image_special_tokens: list[torch.Tensor] = []
|
||||
for image_index in range(num_ref_images):
|
||||
assert ref_image_data is not None
|
||||
assert ref_image_feat_len is not None
|
||||
assert ref_image_special_tokens is not None
|
||||
image_features.append(
|
||||
self.img2tokens(
|
||||
ref_image_data[batch_index, image_index].unsqueeze(0)
|
||||
).squeeze(0)
|
||||
)
|
||||
image_feature_lengths.append(
|
||||
_len_to_list(ref_image_feat_len[batch_index, image_index])
|
||||
)
|
||||
image_special_tokens.append(
|
||||
ref_image_special_tokens[batch_index, image_index]
|
||||
.to(device=ref_device, dtype=ref_dtype)
|
||||
.unsqueeze(0)
|
||||
)
|
||||
|
||||
items.append(
|
||||
SingleData(
|
||||
video_x_t=video_tokens[batch_index],
|
||||
audio_x_t=audio_tokens[batch_index],
|
||||
audio_feat_len=_to_int(data.audio_feat_len[batch_index]),
|
||||
txt_feat=text_tokens[batch_index],
|
||||
txt_feat_len=_to_int(data.txt_feat_len[batch_index]),
|
||||
ref_image_feats=image_features,
|
||||
ref_image_feat_lens=image_feature_lengths,
|
||||
ref_image_special_tokens=image_special_tokens,
|
||||
t=video_t,
|
||||
h=video_h,
|
||||
w=video_w,
|
||||
patch_size=self.patch_size,
|
||||
t_patch_size=self.t_patch_size,
|
||||
spatial_rope_interpolation=self.config.spatial_rope_interpolation,
|
||||
diffusion_t=(
|
||||
data.t[batch_index] if self.config.add_time_token else None
|
||||
),
|
||||
per_token_video_t=(
|
||||
per_token_video_tokens[batch_index]
|
||||
if per_token_video_tokens is not None
|
||||
else None
|
||||
),
|
||||
per_token_audio_t=(
|
||||
per_token_audio[batch_index]
|
||||
if per_token_audio is not None
|
||||
else None
|
||||
),
|
||||
time_channel_dim=self.config.time_channel_dim,
|
||||
time_pos_fps=self.config.time_pos_fps,
|
||||
vae_first_latent_is_image=self.config.vae_first_latent_is_image,
|
||||
video_fps=self.config.video_fps,
|
||||
)
|
||||
)
|
||||
|
||||
packed = SimplePackedData(items)
|
||||
packed_cu_seqlens = packed.cu_seqlen.to(
|
||||
device=data.x_t.device,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
varlen_handler = VarlenHandler(
|
||||
cu_seqlens_q=packed_cu_seqlens,
|
||||
cu_seqlens_k=packed_cu_seqlens.clone(),
|
||||
max_seqlen_q=packed.max_seqlen,
|
||||
max_seqlen_k=packed.max_seqlen,
|
||||
)
|
||||
(
|
||||
token_sequence,
|
||||
coords_mapping,
|
||||
modality_mapping,
|
||||
time_token_sequence,
|
||||
varlen_handler,
|
||||
pad_size,
|
||||
) = self._pad_for_ep_cp(
|
||||
packed.token_sequence,
|
||||
packed.coords_mapping,
|
||||
packed.modality_mapping,
|
||||
packed.time_token_sequence,
|
||||
varlen_handler,
|
||||
)
|
||||
self.saved_for_output(
|
||||
simple_packed_data=packed,
|
||||
input_video_channel=input_video_channel,
|
||||
pad_size=pad_size,
|
||||
)
|
||||
return (
|
||||
token_sequence,
|
||||
coords_mapping,
|
||||
modality_mapping,
|
||||
varlen_handler,
|
||||
time_token_sequence,
|
||||
)
|
||||
|
||||
def process_output(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Remove distributed padding and restore video and audio tensors."""
|
||||
packed: SimplePackedData = self.get_saved_data("simple_packed_data")
|
||||
pad_size = self.get_saved_data("pad_size")
|
||||
if pad_size > 0:
|
||||
x = x[:-pad_size]
|
||||
return packed.depack_token_sequence(x)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Request validation and latent initialization for MAGI-2 inference."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class Magi2InputValidationStage(PipelineStage):
|
||||
"""Validate the published ten-second 1080p T2V and I2V request profile."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
output_frames: int,
|
||||
output_height: int,
|
||||
output_width: int,
|
||||
output_fps: int,
|
||||
) -> None:
|
||||
"""Store the fixed output geometry supported by the release checkpoint."""
|
||||
super().__init__()
|
||||
self.output_frames = output_frames
|
||||
self.output_height = output_height
|
||||
self.output_width = output_width
|
||||
self.output_fps = output_fps
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record; ``forward`` emits model-specific errors."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record after request normalization and seeding."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Require release geometry and reset every request's random-number stream."""
|
||||
if fastvideo_args.workload_type not in (WorkloadType.T2V, WorkloadType.I2V):
|
||||
raise ValueError("MAGI-2 supports only T2V and I2V workloads")
|
||||
if not isinstance(batch.prompt, str) or not batch.prompt.strip():
|
||||
raise ValueError("MAGI-2 requires one non-empty prompt string")
|
||||
expected_geometry = (
|
||||
self.output_frames,
|
||||
self.output_height,
|
||||
self.output_width,
|
||||
self.output_fps,
|
||||
)
|
||||
received_geometry = (
|
||||
batch.num_frames,
|
||||
batch.height,
|
||||
batch.width,
|
||||
batch.fps,
|
||||
)
|
||||
if received_geometry != expected_geometry:
|
||||
raise ValueError(
|
||||
"MAGI-2 Preview supports 249 frames at 1088x1920 and 25 fps; "
|
||||
f"received frames/height/width/fps={received_geometry}"
|
||||
)
|
||||
if batch.num_inference_steps <= 0 or batch.num_inference_steps_sr <= 0:
|
||||
raise ValueError("MAGI-2 denoising step counts must be positive")
|
||||
seed = 42 if batch.seed is None else int(batch.seed)
|
||||
batch.seed = seed
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
if fastvideo_args.deterministic or envs.MAGI2_DETERMINISTIC:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
return batch
|
||||
|
||||
|
||||
class Magi2LatentPreparationStage(PipelineStage):
|
||||
"""Draw video noise before audio noise with the official release shapes."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
video_channels: int,
|
||||
video_length: int,
|
||||
video_height: int,
|
||||
video_width: int,
|
||||
audio_length: int,
|
||||
audio_channels: int,
|
||||
) -> None:
|
||||
"""Store the preview video and audio latent geometry."""
|
||||
super().__init__()
|
||||
self.video_shape = (
|
||||
1,
|
||||
video_channels,
|
||||
video_length,
|
||||
video_height,
|
||||
video_width,
|
||||
)
|
||||
self.audio_shape = (1, audio_length, audio_channels)
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for a validated and seeded request."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for joint video and audio noise tensors."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Draw FP32 video and audio noise in the release model's RNG order."""
|
||||
del fastvideo_args
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
batch.latents = torch.randn(
|
||||
self.video_shape,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
batch.audio_latents = torch.randn(
|
||||
self.audio_shape,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = ["Magi2InputValidationStage", "Magi2LatentPreparationStage"]
|
||||
@@ -0,0 +1,500 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Preview and refiner denoising stages for MAGI-2."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.preview_data_proxy import (
|
||||
Magi2DataProxy,
|
||||
Magi2DataProxyConfig,
|
||||
ModelInput,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages.refiner_data_proxy import (
|
||||
Magi2RefinerDataProxy,
|
||||
Magi2RefinerDataProxyConfig,
|
||||
RefinerModelInput,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
def _pad_or_trim(
|
||||
tensor: torch.Tensor,
|
||||
target_size: int,
|
||||
dimension: int,
|
||||
pad_value: float = 0.0,
|
||||
) -> torch.Tensor:
|
||||
"""Match one tensor dimension to the shared classifier-free-guidance length."""
|
||||
current_size = tensor.size(dimension)
|
||||
if current_size < target_size:
|
||||
padding_amount = target_size - current_size
|
||||
padding = [0] * (2 * tensor.dim())
|
||||
padding_dimension = tensor.dim() - 1 - dimension
|
||||
padding[2 * padding_dimension + 1] = padding_amount
|
||||
return F.pad(tensor, tuple(padding), "constant", pad_value)
|
||||
slices = [slice(None)] * tensor.dim()
|
||||
slices[dimension] = slice(0, target_size)
|
||||
return tensor[tuple(slices)]
|
||||
|
||||
|
||||
class Magi2PreviewDenoisingStage(PipelineStage):
|
||||
"""Run the joint video-audio preview transformer with two-way CFG."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: Any,
|
||||
data_proxy_config: Magi2DataProxyConfig,
|
||||
flow_shift: float,
|
||||
video_guidance_scale: float,
|
||||
audio_guidance_scale: float,
|
||||
) -> None:
|
||||
"""Store the preview transformer and published sampling constants."""
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.data_proxy = Magi2DataProxy(data_proxy_config)
|
||||
self.flow_shift = flow_shift
|
||||
self.video_guidance_scale = video_guidance_scale
|
||||
self.audio_guidance_scale = audio_guidance_scale
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record; ``forward`` validates required tensors."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for denoised video and audio latents."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
@staticmethod
|
||||
def _prepare_reference_cfg(
|
||||
reference_latent: torch.Tensor | None,
|
||||
reference_length: torch.Tensor | None,
|
||||
special_tokens: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]:
|
||||
"""Duplicate I2V conditioning across both classifier-free-guidance halves."""
|
||||
if reference_latent is None:
|
||||
return None, None, None
|
||||
if reference_length is None or special_tokens is None:
|
||||
raise ValueError("MAGI-2 I2V conditioning is incomplete")
|
||||
return (
|
||||
torch.cat([reference_latent, reference_latent], dim=0),
|
||||
torch.cat([reference_length, reference_length], dim=0),
|
||||
torch.cat([special_tokens, special_tokens], dim=0),
|
||||
)
|
||||
|
||||
def _prepare_model_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
video_latent: torch.Tensor,
|
||||
audio_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> ModelInput:
|
||||
"""Build the exact two-sample conditional/unconditional preview input."""
|
||||
text_context = batch.magi2_text_context
|
||||
negative_context = batch.magi2_negative_context
|
||||
if text_context is None or negative_context is None:
|
||||
raise ValueError("MAGI-2 preview denoising requires text conditioning")
|
||||
text_length = text_context.shape[1]
|
||||
negative_length = negative_context.shape[1]
|
||||
shared_length = max(text_length, negative_length)
|
||||
text_context = _pad_or_trim(text_context, shared_length, 1)
|
||||
negative_context = _pad_or_trim(negative_context, shared_length, 1)
|
||||
|
||||
reference_video = torch.empty_like(video_latent)
|
||||
reference_video = torch.cat(
|
||||
[reference_video, torch.zeros_like(reference_video)],
|
||||
dim=0,
|
||||
)
|
||||
reference_video_length = torch.tensor(
|
||||
[0, 0],
|
||||
device=video_latent.device,
|
||||
)
|
||||
reference_image, reference_image_length, special_tokens = (
|
||||
self._prepare_reference_cfg(
|
||||
batch.magi2_ref_image_feat,
|
||||
batch.magi2_ref_image_feat_len,
|
||||
batch.magi2_ref_image_special_tokens,
|
||||
)
|
||||
)
|
||||
|
||||
timestep_value = timestep.reshape(-1)[0].to(
|
||||
device=video_latent.device,
|
||||
dtype=video_latent.dtype,
|
||||
)
|
||||
normalized_timestep = torch.stack([timestep_value, timestep_value]) / 1000.0
|
||||
_, _, video_time, video_height, video_width = video_latent.shape
|
||||
audio_time = audio_latent.shape[1]
|
||||
per_token_video_t = normalized_timestep.view(-1, 1, 1, 1, 1).expand(
|
||||
2,
|
||||
1,
|
||||
video_time,
|
||||
video_height,
|
||||
video_width,
|
||||
).clone()
|
||||
per_token_audio_t = normalized_timestep.view(-1, 1, 1).expand(
|
||||
2,
|
||||
audio_time,
|
||||
1,
|
||||
).clone()
|
||||
reference_audio = torch.zeros(
|
||||
1,
|
||||
0,
|
||||
audio_latent.shape[-1],
|
||||
device=video_latent.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
return ModelInput(
|
||||
x_t=torch.cat([video_latent, video_latent], dim=0),
|
||||
audio_x_t=torch.cat([audio_latent, audio_latent], dim=0),
|
||||
audio_feat_len=torch.tensor(
|
||||
[audio_time, audio_time],
|
||||
device=video_latent.device,
|
||||
),
|
||||
txt_feat=torch.cat([text_context, negative_context], dim=0),
|
||||
txt_feat_len=torch.tensor(
|
||||
[text_length, negative_length],
|
||||
device=video_latent.device,
|
||||
),
|
||||
t=normalized_timestep,
|
||||
per_token_video_t=per_token_video_t,
|
||||
per_token_audio_t=per_token_audio_t,
|
||||
ref_audio_feat=torch.cat(
|
||||
[reference_audio, torch.zeros_like(reference_audio)],
|
||||
dim=0,
|
||||
),
|
||||
ref_audio_feat_len=torch.tensor(
|
||||
[0, 0],
|
||||
device=video_latent.device,
|
||||
),
|
||||
ref_video_feat=reference_video,
|
||||
ref_video_feat_len=reference_video_length,
|
||||
ref_image_feat=reference_image,
|
||||
ref_image_feat_len=reference_image_length,
|
||||
ref_image_special_token_embedding=special_tokens,
|
||||
)
|
||||
|
||||
def _predict_velocity(
|
||||
self,
|
||||
model_input: ModelInput,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Pack tokens, execute the preview transformer, and restore both modalities."""
|
||||
transformer_input = self.data_proxy.process_input(model_input)
|
||||
transformer_output = self.transformer(*transformer_input)
|
||||
return self.data_proxy.process_output(transformer_output)
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Advance the video and audio latents through the published preview schedule."""
|
||||
if batch.latents is None or batch.audio_latents is None:
|
||||
raise ValueError("MAGI-2 preview denoising requires joint noise latents")
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
self.transformer.to(device)
|
||||
video_latent = batch.latents.clone()
|
||||
audio_latent = batch.audio_latents.clone()
|
||||
video_scheduler = FlowUniPCMultistepScheduler()
|
||||
audio_scheduler = FlowUniPCMultistepScheduler()
|
||||
video_scheduler.set_timesteps(
|
||||
batch.num_inference_steps,
|
||||
device=device,
|
||||
shift=self.flow_shift,
|
||||
)
|
||||
audio_scheduler.set_timesteps(
|
||||
batch.num_inference_steps,
|
||||
device=device,
|
||||
shift=self.flow_shift,
|
||||
)
|
||||
video_guidance = torch.tensor(
|
||||
[self.video_guidance_scale],
|
||||
device=device,
|
||||
)
|
||||
for timestep in video_scheduler.timesteps:
|
||||
model_input = self._prepare_model_input(
|
||||
batch,
|
||||
video_latent,
|
||||
audio_latent,
|
||||
timestep,
|
||||
)
|
||||
video_prediction, audio_prediction = self._predict_velocity(model_input)
|
||||
conditional_video = video_prediction[0:1]
|
||||
unconditional_video = video_prediction[1:2]
|
||||
video_velocity = unconditional_video + video_guidance * (
|
||||
conditional_video - unconditional_video
|
||||
)
|
||||
conditional_audio = audio_prediction[0:1]
|
||||
unconditional_audio = audio_prediction[1:2]
|
||||
audio_velocity = unconditional_audio + self.audio_guidance_scale * (
|
||||
conditional_audio - unconditional_audio
|
||||
)
|
||||
video_latent = video_scheduler.step(
|
||||
video_velocity,
|
||||
timestep,
|
||||
video_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
audio_latent = audio_scheduler.step(
|
||||
audio_velocity,
|
||||
timestep,
|
||||
audio_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
batch.latents = video_latent
|
||||
batch.audio_latents = audio_latent
|
||||
if fastvideo_args.dit_cpu_offload:
|
||||
self.transformer.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
class ZeroSNRDDPMDiscretization:
|
||||
"""Construct the refiner's published zero-terminal-SNR noise schedule."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
linear_start: float = 0.00085,
|
||||
linear_end: float = 0.0120,
|
||||
num_timesteps: int = 1000,
|
||||
) -> None:
|
||||
"""Precompute the cumulative alpha schedule in FP64 NumPy arithmetic."""
|
||||
self.num_timesteps = num_timesteps
|
||||
betas = torch.linspace(
|
||||
linear_start**0.5,
|
||||
linear_end**0.5,
|
||||
num_timesteps,
|
||||
dtype=torch.float64,
|
||||
) ** 2
|
||||
self.alphas_cumprod = np.cumprod(1.0 - betas.numpy(), axis=0)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
count: int,
|
||||
do_append_zero: bool = True,
|
||||
device: str = "cpu",
|
||||
flip: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Return sampled sigma values with optional terminal zero and reversal."""
|
||||
sigmas = self.get_sigmas(count, device=device)
|
||||
if do_append_zero:
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
return torch.flip(sigmas, (0,)) if flip else sigmas
|
||||
|
||||
def get_sigmas(self, count: int, device: str = "cpu") -> torch.Tensor:
|
||||
"""Convert cumulative alphas into the refiner's sigma parameterization."""
|
||||
if count < self.num_timesteps:
|
||||
timesteps = np.linspace(
|
||||
self.num_timesteps - 1,
|
||||
0,
|
||||
count,
|
||||
endpoint=False,
|
||||
).astype(int)[::-1]
|
||||
alphas_cumprod = self.alphas_cumprod[timesteps]
|
||||
elif count == self.num_timesteps:
|
||||
alphas_cumprod = self.alphas_cumprod
|
||||
else:
|
||||
raise ValueError(
|
||||
f"sigma count must be <= {self.num_timesteps}, received {count}"
|
||||
)
|
||||
alphas = torch.tensor(
|
||||
alphas_cumprod,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).sqrt()
|
||||
first_alpha = alphas[0].clone()
|
||||
terminal_alpha = alphas[-1].clone()
|
||||
alphas -= terminal_alpha
|
||||
alphas *= first_alpha / (first_alpha - terminal_alpha)
|
||||
return torch.flip(alphas, (0,))
|
||||
|
||||
|
||||
class Magi2RefinerStage(PipelineStage):
|
||||
"""Upsample preview latents and run the five-step 1080p refiner."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: Any,
|
||||
data_proxy_config: Magi2RefinerDataProxyConfig,
|
||||
latent_height: int,
|
||||
latent_width: int,
|
||||
noise_index: int,
|
||||
flow_shift: float,
|
||||
video_guidance_scale: float,
|
||||
audio_guidance_scale: float,
|
||||
audio_channels: int,
|
||||
) -> None:
|
||||
"""Store the refiner transformer, token proxy, and release constants."""
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.data_proxy = Magi2RefinerDataProxy(data_proxy_config)
|
||||
self.latent_height = latent_height
|
||||
self.latent_width = latent_width
|
||||
self.noise_index = noise_index
|
||||
self.flow_shift = flow_shift
|
||||
self.video_guidance_scale = video_guidance_scale
|
||||
self.audio_guidance_scale = audio_guidance_scale
|
||||
self.audio_channels = audio_channels
|
||||
self.sigmas = ZeroSNRDDPMDiscretization()(
|
||||
1000,
|
||||
do_append_zero=False,
|
||||
flip=True,
|
||||
)
|
||||
|
||||
def verify_input(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record; ``forward`` validates latent conditioning."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
"""Return the default record for the post-refiner video latent."""
|
||||
del batch, fastvideo_args
|
||||
return VerificationResult()
|
||||
|
||||
def _predict_velocity(
|
||||
self,
|
||||
video_latent: torch.Tensor,
|
||||
audio_latent: torch.Tensor,
|
||||
text_context: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Pack one text branch, execute the refiner, and restore both modalities."""
|
||||
reference_audio = torch.zeros(
|
||||
1,
|
||||
0,
|
||||
self.audio_channels,
|
||||
device=video_latent.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
model_input = RefinerModelInput(
|
||||
x_t=video_latent,
|
||||
audio_x_t=audio_latent,
|
||||
audio_feat_len=torch.tensor([audio_latent.shape[1]]),
|
||||
txt_feat=text_context.to(torch.float32),
|
||||
txt_feat_len=torch.tensor([text_context.shape[1]]),
|
||||
ref_audio_feat=reference_audio,
|
||||
ref_audio_feat_len=torch.tensor([0]),
|
||||
ref_video_feat=torch.empty_like(video_latent),
|
||||
ref_video_feat_len=torch.tensor([0]),
|
||||
)
|
||||
transformer_input = self.data_proxy.process_input(model_input)
|
||||
transformer_output = self.transformer(*transformer_input)
|
||||
return self.data_proxy.process_output(transformer_output)
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""Inject refiner noise and advance the upsampled latent for five steps."""
|
||||
if batch.latents is None:
|
||||
raise ValueError("MAGI-2 refiner requires the preview video latent")
|
||||
if batch.magi2_text_context is None or batch.magi2_negative_context is None:
|
||||
raise ValueError("MAGI-2 refiner requires positive and negative text context")
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
video_latent = F.interpolate(
|
||||
batch.latents,
|
||||
size=(
|
||||
batch.latents.shape[2] * 2 - 1,
|
||||
self.latent_height,
|
||||
self.latent_width,
|
||||
),
|
||||
mode="trilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
noise = torch.randn_like(video_latent, device=video_latent.device)
|
||||
sigma = self.sigmas.to(video_latent.device)[self.noise_index]
|
||||
video_latent = video_latent * sigma + noise * (1 - sigma**2) ** 0.5
|
||||
refiner_audio = torch.zeros(
|
||||
1,
|
||||
0,
|
||||
self.audio_channels,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
video_scheduler = FlowUniPCMultistepScheduler()
|
||||
audio_scheduler = FlowUniPCMultistepScheduler()
|
||||
video_scheduler.set_timesteps(
|
||||
batch.num_inference_steps_sr,
|
||||
device=device,
|
||||
shift=self.flow_shift,
|
||||
)
|
||||
audio_scheduler.set_timesteps(
|
||||
batch.num_inference_steps_sr,
|
||||
device=device,
|
||||
shift=self.flow_shift,
|
||||
)
|
||||
video_guidance = torch.tensor(
|
||||
self.video_guidance_scale,
|
||||
device=device,
|
||||
).expand(1, 1, video_latent.shape[2], 1, 1).clone()
|
||||
self.transformer.to(device)
|
||||
for timestep in video_scheduler.timesteps:
|
||||
conditional_video, conditional_audio = self._predict_velocity(
|
||||
video_latent,
|
||||
refiner_audio,
|
||||
batch.magi2_text_context,
|
||||
)
|
||||
unconditional_video, unconditional_audio = self._predict_velocity(
|
||||
video_latent,
|
||||
refiner_audio,
|
||||
batch.magi2_negative_context,
|
||||
)
|
||||
video_velocity = unconditional_video + video_guidance * (
|
||||
conditional_video - unconditional_video
|
||||
)
|
||||
audio_velocity = unconditional_audio + self.audio_guidance_scale * (
|
||||
conditional_audio - unconditional_audio
|
||||
)
|
||||
video_latent = video_scheduler.step(
|
||||
video_velocity,
|
||||
timestep,
|
||||
video_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if audio_velocity.numel() > 0:
|
||||
refiner_audio = audio_scheduler.step(
|
||||
audio_velocity,
|
||||
timestep,
|
||||
refiner_audio,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
batch.latents = video_latent
|
||||
if fastvideo_args.dit_cpu_offload:
|
||||
self.transformer.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Magi2PreviewDenoisingStage",
|
||||
"Magi2RefinerStage",
|
||||
"ZeroSNRDDPMDiscretization",
|
||||
]
|
||||
@@ -104,6 +104,11 @@ class ForwardBatch:
|
||||
negative_attention_mask: list[torch.Tensor] | None = None
|
||||
clip_embedding_pos: list[torch.Tensor] | None = None
|
||||
clip_embedding_neg: list[torch.Tensor] | None = None
|
||||
magi2_text_context: torch.Tensor | None = None
|
||||
magi2_negative_context: torch.Tensor | None = None
|
||||
magi2_ref_image_feat: torch.Tensor | None = None
|
||||
magi2_ref_image_feat_len: torch.Tensor | None = None
|
||||
magi2_ref_image_special_tokens: torch.Tensor | None = None
|
||||
|
||||
# Additional text-related parameters
|
||||
max_sequence_length: int | None = None
|
||||
@@ -125,6 +130,7 @@ class ForwardBatch:
|
||||
|
||||
# Latent tensors
|
||||
latents: torch.Tensor | None = None
|
||||
audio_latents: torch.Tensor | None = None
|
||||
lq_latents: torch.Tensor | None = None
|
||||
raw_latent_shape: tuple[int, ...] | None = None
|
||||
noise_pred: torch.Tensor | None = None
|
||||
|
||||
@@ -32,6 +32,7 @@ from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsk
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.magi2 import Magi2PreviewPipelineConfig
|
||||
from fastvideo.configs.pipelines.flux_2 import (
|
||||
Flux2KleinPipelineConfig,
|
||||
Flux2PipelineConfig,
|
||||
@@ -299,6 +300,17 @@ def _register_configs() -> None:
|
||||
default_preset="ltx2_base",
|
||||
)
|
||||
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Magi2PreviewPipelineConfig,
|
||||
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
|
||||
model_detectors=[
|
||||
lambda path: "magi2pipeline" in path.lower(),
|
||||
],
|
||||
model_family="magi2",
|
||||
default_preset="magi2_preview_1080p",
|
||||
)
|
||||
|
||||
# Stable Audio Open (text-to-audio). Both variants must be loaded
|
||||
# from the FastVideo-curated converted Diffusers-format repos —
|
||||
# the upstream `stabilityai/stable-audio-open-{1.0,small}` repos
|
||||
@@ -1214,6 +1226,8 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as LONGCAT_PRESETS, )
|
||||
from fastvideo.pipelines.basic.ltx2.presets import (
|
||||
ALL_PRESETS as LTX2_PRESETS, )
|
||||
from fastvideo.pipelines.basic.magi2.presets import (
|
||||
ALL_PRESETS as MAGI2_PRESETS, )
|
||||
from fastvideo.pipelines.basic.matrixgame2.presets import (
|
||||
ALL_PRESETS as MATRIXGAME2_PRESETS, )
|
||||
from fastvideo.pipelines.basic.matrixgame3.presets import (
|
||||
@@ -1242,6 +1256,7 @@ def _register_presets() -> None:
|
||||
LINGBOTWORLD_PRESETS,
|
||||
LONGCAT_PRESETS,
|
||||
LTX2_PRESETS,
|
||||
MAGI2_PRESETS,
|
||||
MATRIXGAME2_PRESETS,
|
||||
MATRIXGAME3_PRESETS,
|
||||
SD35_PRESETS,
|
||||
|
||||
@@ -129,6 +129,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
|
||||
"enable_stage_verification": True,
|
||||
"use_fsdp_inference": False,
|
||||
"disable_autocast": False,
|
||||
"deterministic": False,
|
||||
"quantization": None,
|
||||
},
|
||||
"pipeline": {
|
||||
|
||||
+1
Submodule fastvideo/third_party/flash_attention added at b613d9e2c8
+1
Submodule fastvideo/third_party/magi_attention added at 2c6413571c
+1
Submodule fastvideo/third_party/magi_compiler added at 5950612ddf
@@ -0,0 +1,223 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Create a FastVideo MAGI-2 repository from the official checkpoint layout.
|
||||
|
||||
The published ``sand-ai/MAGI-2-preview`` snapshot already stores tensors in the
|
||||
names and dtypes that the FastVideo-native loaders consume. This converter keeps
|
||||
the tensor files byte-for-byte and changes only the component directory layout.
|
||||
Hard links avoid a second 286 GiB allocation while producing regular files that
|
||||
Hugging Face upload tools can publish as a self-contained repository.
|
||||
|
||||
Example:
|
||||
python scripts/checkpoint_conversion/convert_magi2_to_fastvideo.py \
|
||||
--source official_weights/magi2 \
|
||||
--output converted_weights/magi2
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
SOURCE_REPOSITORY = "sand-ai/MAGI-2-preview"
|
||||
SOURCE_REVISION = "2dea51b64db47ee5b4402d36fd90829a0c58913b"
|
||||
|
||||
COMPONENT_DIRECTORY_MAPPING: dict[str, str] = {
|
||||
"preview": "transformer",
|
||||
"refiner": "transformer_2",
|
||||
"text_encoder": "text_encoder",
|
||||
"vae": "image_encoder",
|
||||
"turbo_vae": "vae",
|
||||
"stable-audio-open-1.0": "audio_vae",
|
||||
}
|
||||
|
||||
REQUIRED_COMPONENT_FILES: dict[str, tuple[str, ...]] = {
|
||||
"preview": ("model.safetensors.index.json",),
|
||||
"refiner": ("model.safetensors.index.json",),
|
||||
"text_encoder": ("config.json", "model.safetensors.index.json", "tokenizer.json"),
|
||||
"vae": ("Wan2.2_VAE.pth",),
|
||||
"turbo_vae": ("TurboV3-Wan22-TinyShallow_7_7.json", "checkpoint.ckpt"),
|
||||
"stable-audio-open-1.0": ("model_config.json", "model.safetensors"),
|
||||
}
|
||||
|
||||
# The distilled decoder checkpoint contains these feature-matching heads for
|
||||
# training. MAGI-2 inference strictly loads the 88 ``decoder.*`` tensors.
|
||||
TURBO_VAE_SKIPPED_KEYS: tuple[str, ...] = (
|
||||
"aligned_feature_projection_heads.0.0.conv.weight",
|
||||
"aligned_feature_projection_heads.0.0.conv.bias",
|
||||
"aligned_feature_projection_heads.0.1.conv.weight",
|
||||
"aligned_feature_projection_heads.0.1.conv.bias",
|
||||
"aligned_feature_projection_heads.1.0.conv.weight",
|
||||
"aligned_feature_projection_heads.1.0.conv.bias",
|
||||
"aligned_feature_projection_heads.1.1.conv.weight",
|
||||
"aligned_feature_projection_heads.1.1.conv.bias",
|
||||
)
|
||||
|
||||
MODEL_INDEX: dict[str, object] = {
|
||||
"_class_name": "Magi2Pipeline",
|
||||
"_diffusers_version": "0.37.0",
|
||||
"transformer": ["fastvideo.models.dits.magi2", "Magi2PreviewDiT"],
|
||||
"transformer_2": ["fastvideo.models.dits.magi2_refiner", "Magi2RefinerDiT"],
|
||||
"text_encoder": ["fastvideo.models.encoders.qwen3_5", "Magi2Qwen35TextEncoder"],
|
||||
"image_encoder": ["fastvideo.models.vaes.magi2_wan_loader", "Magi2WanImageEncoder"],
|
||||
"vae": ["fastvideo.models.vaes.magi2_turbo_vae", "Magi2TurboVAEModel"],
|
||||
"audio_vae": ["fastvideo.models.vaes.magi2_audio_vae", "Magi2AudioVAE"],
|
||||
"scheduler": [
|
||||
"fastvideo.models.schedulers.scheduling_flow_unipc_multistep",
|
||||
"FlowUniPCMultistepScheduler",
|
||||
],
|
||||
}
|
||||
|
||||
SCHEDULER_CONFIG: dict[str, object] = {
|
||||
"_class_name": "FlowUniPCMultistepScheduler",
|
||||
"_diffusers_version": "0.37.0",
|
||||
"disable_corrector": [],
|
||||
"lower_order_final": True,
|
||||
"num_train_timesteps": 1000,
|
||||
"predict_x0": True,
|
||||
"prediction_type": "flow_prediction",
|
||||
"shift": 1.0,
|
||||
"solver_order": 2,
|
||||
"solver_type": "bh2",
|
||||
}
|
||||
|
||||
TRANSFORMER_CONFIGS: dict[str, dict[str, object]] = {
|
||||
"transformer": {
|
||||
"_class_name": "Magi2PreviewDiT",
|
||||
"_diffusers_version": "0.37.0",
|
||||
"source_subfolder": "preview",
|
||||
},
|
||||
"transformer_2": {
|
||||
"_class_name": "Magi2RefinerDiT",
|
||||
"_diffusers_version": "0.37.0",
|
||||
"source_subfolder": "refiner",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _validate_source(source: Path) -> None:
|
||||
"""Require component metadata and every shard referenced by an index."""
|
||||
missing_files = [
|
||||
str(source / component / relative_path)
|
||||
for component, relative_paths in REQUIRED_COMPONENT_FILES.items()
|
||||
for relative_path in relative_paths
|
||||
if not (source / component / relative_path).is_file()
|
||||
]
|
||||
if missing_files:
|
||||
raise FileNotFoundError(
|
||||
"The MAGI-2 source snapshot is incomplete; missing: "
|
||||
+ ", ".join(missing_files)
|
||||
)
|
||||
|
||||
missing_shards: list[str] = []
|
||||
for component, relative_paths in REQUIRED_COMPONENT_FILES.items():
|
||||
for relative_path in relative_paths:
|
||||
if not relative_path.endswith(".index.json"):
|
||||
continue
|
||||
index_path = source / component / relative_path
|
||||
index_payload = json.loads(index_path.read_text(encoding="utf-8"))
|
||||
weight_map = index_payload.get("weight_map")
|
||||
if not isinstance(weight_map, dict):
|
||||
raise ValueError(f"Checkpoint index has no weight_map object: {index_path}")
|
||||
if not weight_map or not all(isinstance(shard_name, str) for shard_name in weight_map.values()):
|
||||
raise ValueError(f"Checkpoint index has an invalid weight_map object: {index_path}")
|
||||
shard_names = set(weight_map.values())
|
||||
missing_shards.extend(
|
||||
str(index_path.parent / shard_name)
|
||||
for shard_name in sorted(shard_names)
|
||||
if not (index_path.parent / shard_name).is_file()
|
||||
)
|
||||
if missing_shards:
|
||||
raise FileNotFoundError(
|
||||
"The MAGI-2 source snapshot is missing indexed checkpoint shards: "
|
||||
+ ", ".join(missing_shards)
|
||||
)
|
||||
|
||||
|
||||
def _hardlink_component(source: Path, destination: Path) -> None:
|
||||
"""Replicate one component tree with regular-file hard links."""
|
||||
destination.mkdir(parents=True)
|
||||
for source_path in sorted(source.rglob("*")):
|
||||
relative_path = source_path.relative_to(source)
|
||||
destination_path = destination / relative_path
|
||||
if source_path.is_dir():
|
||||
destination_path.mkdir()
|
||||
continue
|
||||
if not source_path.is_file():
|
||||
raise ValueError(f"Unsupported checkpoint entry: {source_path}")
|
||||
os.link(source_path, destination_path)
|
||||
|
||||
|
||||
def _write_json(path: Path, payload: object) -> None:
|
||||
"""Write stable, reviewable JSON with a trailing newline."""
|
||||
path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8")
|
||||
|
||||
|
||||
def _conversion_manifest() -> dict[str, object]:
|
||||
"""Describe directory and tensor-name transformations for every component."""
|
||||
return {
|
||||
"format_version": 1,
|
||||
"source": {
|
||||
"repository": SOURCE_REPOSITORY,
|
||||
"revision": SOURCE_REVISION,
|
||||
},
|
||||
"component_directory_mapping": COMPONENT_DIRECTORY_MAPPING,
|
||||
"tensor_mapping": {
|
||||
"transformer": "identity; expert_bias_ema replaces expert_bias during inference loading",
|
||||
"transformer_2": "identity",
|
||||
"text_encoder": "identity; Transformers Qwen3.5 loads the language-model tensors",
|
||||
"image_encoder": "identity; ignore decoder.* and conv2.* tensors",
|
||||
"vae": "strip module.; load decoder.* tensors",
|
||||
"audio_vae": "strip pretransform.model.; map the published sequential decoder to OobleckDecoder",
|
||||
},
|
||||
"skipped_checkpoint_keys": {
|
||||
"vae/checkpoint.ckpt": list(TURBO_VAE_SKIPPED_KEYS),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def convert_checkpoint_layout(source: Path, output: Path) -> None:
|
||||
"""Hard-link official components and add FastVideo repository metadata."""
|
||||
source = source.resolve()
|
||||
output = output.resolve()
|
||||
_validate_source(source)
|
||||
if output.exists():
|
||||
raise FileExistsError(f"Output path already exists: {output}")
|
||||
output.mkdir(parents=True)
|
||||
|
||||
for source_name, destination_name in COMPONENT_DIRECTORY_MAPPING.items():
|
||||
_hardlink_component(source / source_name, output / destination_name)
|
||||
|
||||
scheduler_dir = output / "scheduler"
|
||||
scheduler_dir.mkdir()
|
||||
_write_json(scheduler_dir / "scheduler_config.json", SCHEDULER_CONFIG)
|
||||
for component_name, config in TRANSFORMER_CONFIGS.items():
|
||||
_write_json(output / component_name / "config.json", config)
|
||||
_write_json(output / "model_index.json", MODEL_INDEX)
|
||||
_write_json(output / "magi2_conversion_manifest.json", _conversion_manifest())
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Parse source and destination paths and create the converted repository."""
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--source",
|
||||
type=Path,
|
||||
required=True,
|
||||
help="Local snapshot of sand-ai/MAGI-2-preview.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=Path,
|
||||
required=True,
|
||||
help="Destination FastVideo model repository; the path must not exist.",
|
||||
)
|
||||
arguments = parser.parse_args()
|
||||
convert_checkpoint_layout(arguments.source, arguments.output)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,168 @@
|
||||
# MAGI-2 Preview Port Status
|
||||
|
||||
## Completion
|
||||
|
||||
The FastVideo MAGI-2 Preview port supports text-to-video (T2V), image-to-video
|
||||
(I2V), and joint audio generation. Strict release-profile parity passes against
|
||||
the official SandAI implementation for both modalities.
|
||||
|
||||
| Field | Value |
|
||||
| ---------------------- | ------------------------------------------------------------- |
|
||||
| Branch | `model/port-magi-2` |
|
||||
| Base commit | `1c04ace57351f7340e8d1a2e8b8f62180856ed16` |
|
||||
| Official code revision | `073c84f2102ec3c9287623113a103c14402770ad` |
|
||||
| Checkpoint revision | `2dea51b64db47ee5b4402d36fd90829a0c58913b` |
|
||||
| Workloads | T2V and I2V with generated audio |
|
||||
| Release profile | 100 preview steps, 5 refiner steps, seed 42 |
|
||||
| Distributed topology | 8 Hopper GPUs; sequence, context, and expert parallel width 8 |
|
||||
| Numerical requirement | Exact shape, dtype, stride, and tensor bytes |
|
||||
| Full pipeline result | Exact for every captured stage and both decoded outputs |
|
||||
| Video output | `uint8 [249, 1088, 1920, 3]` at 25 frames per second |
|
||||
| Audio output | `float32 [441000, 2]` at 44.1 kHz |
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | FastVideo implementation | Strict parity result |
|
||||
| ------------------------ | ---------------------------------------------------------------- | -------------------------------------------- |
|
||||
| Qwen3.5 prompt encoder | `fastvideo/models/encoders/qwen3_5.py` | Exact |
|
||||
| Wan 2.2 image encoder | `fastvideo/models/vaes/magi2_wan_loader.py` | Exact |
|
||||
| Preview input packing | `fastvideo/pipelines/basic/magi2/stages/preview_data_proxy.py` | Exact |
|
||||
| Preview transformer | `fastvideo/models/dits/magi2.py` | Exact across 40 layers |
|
||||
| Refiner input packing | `fastvideo/pipelines/basic/magi2/stages/refiner_data_proxy.py` | Exact |
|
||||
| Refiner transformer | `fastvideo/models/dits/magi2_refiner.py` | Exact across 30 layers |
|
||||
| Flow UniPC scheduler | `fastvideo/models/schedulers/scheduling_flow_unipc_multistep.py` | Exact |
|
||||
| Turbo VAE decoder | `fastvideo/models/vaes/magi2_turbo_vae.py` | Exact across first, middle, and last windows |
|
||||
| Stable Audio VAE decoder | `fastvideo/models/vaes/magi2_audio_vae.py` | Exact before and after resampling |
|
||||
| T2V pipeline | `fastvideo/pipelines/basic/magi2/magi2_pipeline.py` | Exact, 100 + 5 steps |
|
||||
| I2V pipeline | `fastvideo/pipelines/basic/magi2/magi2_pipeline.py` | Exact, 100 + 5 steps |
|
||||
|
||||
## Architecture Fidelity
|
||||
|
||||
The preview `Transformer` preserves the 40-layer hierarchical head-parallel
|
||||
design described in sections 3.1 and 3.2 of the SandAI MAGI-2 report. The
|
||||
implementation exchanges fixed head activations between ranks, shards each
|
||||
layer's experts across the local high-bandwidth interconnect, and uses the
|
||||
12-head, 256-expert, top-6 MagiMoE routing layout. Routing runs in FP32, expert
|
||||
computation runs in BF16, and each expert uses the SwiGLU7 feed-forward layout.
|
||||
The strict preview parity test compares every layer boundary across eight ranks.
|
||||
|
||||
The refiner preserves the 30-layer local-attention design, including the exact
|
||||
temporal and spatial attention ranges for the 1080p latent grid. The strict
|
||||
refiner parity test compares every layer boundary across eight ranks.
|
||||
|
||||
## Full-Profile Output Digests
|
||||
|
||||
The official and FastVideo manifests differ only in the `implementation`
|
||||
label. Every stage-boundary digest and output digest is identical.
|
||||
|
||||
| Workload | Tensor | Shape | SHA-256 |
|
||||
| -------- | ------------- | ---------------------- | ------------------------------------------------------------------ |
|
||||
| T2V | Decoded video | `[249, 1088, 1920, 3]` | `0326570a07353cc78d488117f265b881ef6c681a3d660dc83acce2222b62e9a3` |
|
||||
| T2V | Decoded audio | `[441000, 2]` | `fa62b45da7cfd055829ba79d1fe24d09736ed38e1075098aca71c2d0baad1937` |
|
||||
| I2V | Decoded video | `[249, 1088, 1920, 3]` | `fd211e0f647e2b3304f636202940b19ec138b520abdfd4bfabba74678097edc9` |
|
||||
| I2V | Decoded audio | `[441000, 2]` | `32d578e901d9443d21b3659b068992e112e6430c97e1a56231a99499f6fa3db0` |
|
||||
|
||||
The full-profile manifests are stored at:
|
||||
|
||||
- `archived/magi2_parity/validation/pipeline/official/capture.json`
|
||||
- `archived/magi2_parity/validation/pipeline/fastvideo/capture.json`
|
||||
|
||||
## Compiled-Path Verification
|
||||
|
||||
The official and FastVideo pipelines also ran from separate empty node-local
|
||||
MagiCompiler caches with one preview step and one refiner step. Their T2V and
|
||||
I2V manifests match exactly at every captured stage and decoded output. This
|
||||
reduced schedule verifies compilation and production wiring; the 100 + 5 step
|
||||
comparison provides the full release-profile numerical result.
|
||||
|
||||
| Workload | Tensor | SHA-256 |
|
||||
| -------- | ------------- | ------------------------------------------------------------------ |
|
||||
| T2V | Decoded video | `ca379741f7af306cdc2d2fbbd38dfabfa35f215b613b8c7fe8b9193dcef154f5` |
|
||||
| T2V | Decoded audio | `2e44652e3c4047560255fac37b606ecc26e44d140ce500ee0a7f6328cf30c22e` |
|
||||
| I2V | Decoded video | `354a56ace8504a61487f299867c4a31beee2fd60fcce7f5bc1eb6acefdf9f532` |
|
||||
| I2V | Decoded audio | `622420e1a0777b16804ec1a8e1b53ff5a2ac2dad7edae7823b8d39a8377f67aa` |
|
||||
|
||||
The compiled manifests are stored under
|
||||
`archived/magi2_parity/validation/pipeline_compiled/`.
|
||||
|
||||
## Checkpoint Conversion
|
||||
|
||||
`scripts/checkpoint_conversion/convert_magi2_to_fastvideo.py` converts the
|
||||
official mixed checkpoint layout into the component layout consumed by the
|
||||
FastVideo registry. The converter validates every indexed shard and uses hard
|
||||
links for tensor files so that both layouts share the same file data.
|
||||
|
||||
| Official component | FastVideo component | Role |
|
||||
| ------------------------ | ------------------- | ---------------------------------------- |
|
||||
| `preview/` | `transformer/` | Joint video-audio preview transformer |
|
||||
| `refiner/` | `transformer_2/` | 1080p video refiner |
|
||||
| `text_encoder/` | `text_encoder/` | Qwen3.5 prompt encoder and tokenizer |
|
||||
| `vae/Wan2.2_VAE.pth` | `image_encoder/` | Wan reference-image encoder |
|
||||
| `turbo_vae/` | `vae/` | Default distilled sliding-window decoder |
|
||||
| `stable-audio-open-1.0/` | `audio_vae/` | Stable Audio VAE decoder |
|
||||
|
||||
The preview, refiner, Turbo VAE, Wan VAE, and audio VAE production loaders
|
||||
reject missing tensors, unexpected tensors, and shape mismatches. The Qwen3.5
|
||||
loader selects the text backbone from the multimodal checkpoint and leaves the
|
||||
visual tower and language-model head unused, matching the release pipeline.
|
||||
|
||||
## Runtime Contract
|
||||
|
||||
- Turbo VAE decoding is the default video-decoding path. The decoder uses a
|
||||
seven-latent first window and seven-latent subsequent windows on one context-
|
||||
parallel leader rank per video.
|
||||
- `--deterministic` and `MAGI2_DETERMINISTIC=1` set fixed Python, NumPy,
|
||||
PyTorch, and CUDA seeds and activate deterministic MAGI-2 and MagiAttention
|
||||
kernels.
|
||||
- `MAGI2_SAVE_LATENT_PATH` writes each sample's post-refiner latent from the
|
||||
leader rank. An empty environment value disables latent saving.
|
||||
- I2V conditioning preserves the official BF16 image round trip before the Wan
|
||||
encoder performs FP32 encoding.
|
||||
- Random tensors are drawn in the official order: preview video noise, preview
|
||||
audio noise, and refiner video noise.
|
||||
- The pipeline returns decoded video and generated audio tensors before media
|
||||
container encoding.
|
||||
|
||||
## Pinned Dependencies
|
||||
|
||||
| Dependency | Repository | Revision |
|
||||
| ---------------------- | ---------------------------------------------- | ------------------------------------------ |
|
||||
| MagiAttention | `https://github.com/SandAI-org/MagiAttention` | `2c6413571c2cac6a80d1f85a434c6713fe0f5286` |
|
||||
| MagiCompiler | `https://github.com/SandAI-org/MagiCompiler` | `5950612ddf1205f9ba9c3238a8f02a078023e15c` |
|
||||
| Flash Attention Hopper | `https://github.com/Dao-AILab/flash-attention` | `b613d9e2c8475945baff3fd68f2030af1b890acf` |
|
||||
|
||||
The revisions match the build arguments in the published MAGI-2 Dockerfile.
|
||||
The dependencies live under `fastvideo/third_party/` and are declared in the
|
||||
repository's `.gitmodules` file.
|
||||
|
||||
## Validation Coverage
|
||||
|
||||
| Validation file | Coverage |
|
||||
| ------------------------------------------ | ---------------------------------------------------------- |
|
||||
| `test_magi2_text_encoder_parity.py` | Prompt formatting, CJK splitting, tokenization, embeddings |
|
||||
| `test_magi2_image_encoder_parity.py` | I2V resize, BF16 round trip, Wan posterior mean |
|
||||
| `test_magi2_preview_data_proxy_parity.py` | Preview packing, coordinates, and context padding |
|
||||
| `test_magi2_preview_transformer_parity.py` | T2V and I2V boundaries across all 40 layers |
|
||||
| `test_magi2_refiner_data_proxy_parity.py` | Local-attention ranges and refiner packing |
|
||||
| `test_magi2_refiner_transformer_parity.py` | Boundaries across all 30 refiner layers |
|
||||
| `test_magi2_scheduler_parity.py` | Preview and refiner scheduler arrays and updates |
|
||||
| `test_magi2_turbo_vae_parity.py` | First, middle, and last temporal decode windows |
|
||||
| `test_magi2_audio_vae_parity.py` | Stable Audio VAE decoding and 44.1 kHz resampling |
|
||||
| `test_magi2_runtime_controls.py` | Determinism and leader-only latent saving |
|
||||
| `test_magi2_checkpoint_conversion.py` | Component mapping and indexed-shard validation |
|
||||
| `test_magi2_registry_and_metadata.py` | T2V and I2V registry routing |
|
||||
| `test_magi2_pipeline_parity.py` | Full T2V and I2V stage and decoded-output parity |
|
||||
|
||||
The validation suite and environment setup records are described in
|
||||
`tests/local_tests/magi2/README.md`.
|
||||
|
||||
## Decisions
|
||||
|
||||
| Decision | Reason |
|
||||
| ---------------------------------------- | ----------------------------------------------------------- |
|
||||
| Branch from local `main` | The requested base is local `main` commit `1c04ace...`. |
|
||||
| Preserve the official dependency pins | Kernel and compiler revisions affect strict parity. |
|
||||
| Use Turbo VAE by default | The official release and port contract select this decoder. |
|
||||
| Compare tensors before MP4/AAC encoding | Media codecs operate outside model inference. |
|
||||
| Use eager mode for full byte-parity runs | Eager execution isolates model math from compiler caching. |
|
||||
| Keep MagiCompiler enabled in production | The pinned compiler provides the official optimized path. |
|
||||
@@ -0,0 +1,167 @@
|
||||
# MAGI-2 Preview Local Port Validation
|
||||
|
||||
This directory contains the reproducible local validation suite for the MAGI-2
|
||||
Preview text-to-video (T2V) and image-to-video (I2V) port. The validation target
|
||||
is exact tensor equality with the official SandAI implementation for model
|
||||
components, pipeline stages, terminal video tensors, and terminal audio tensors.
|
||||
|
||||
## Sources
|
||||
|
||||
| Source | Location | Revision |
|
||||
| ------------------------- | -------------------------------------------------------------------------------------------- | ------------------------------------------ |
|
||||
| Official inference code | `https://github.com/SandAI-org/MAGI-2-preview` | `073c84f2102ec3c9287623113a103c14402770ad` |
|
||||
| Official local checkout | `/mnt/weka/shrd/wm/junda/fv-hub/MAGI-2-preview` | `073c84f2102ec3c9287623113a103c14402770ad` |
|
||||
| Published checkpoints | `https://huggingface.co/sand-ai/MAGI-2-preview` | `2dea51b64db47ee5b4402d36fd90829a0c58913b` |
|
||||
| Local checkpoint snapshot | `official_weights/magi2` | `2dea51b64db47ee5b4402d36fd90829a0c58913b` |
|
||||
| MagiAttention | `fastvideo/third_party/magi_attention` | `2c6413571c2cac6a80d1f85a434c6713fe0f5286` |
|
||||
| MagiCompiler | `fastvideo/third_party/magi_compiler` | `5950612ddf1205f9ba9c3238a8f02a078023e15c` |
|
||||
| Flash Attention Hopper | `https://github.com/Dao-AILab/flash-attention/tree/b613d9e2c8475945baff3fd68f2030af1b890acf` | `b613d9e2c8475945baff3fd68f2030af1b890acf` |
|
||||
|
||||
The Hugging Face repository is public and does not require an access token. No
|
||||
Hugging Face token environment variable is used.
|
||||
|
||||
## Checkpoint Layout
|
||||
|
||||
The checkpoint repository has a custom mixed layout without a root
|
||||
`model_index.json` file. The converter maps the official component directories
|
||||
into a FastVideo model repository and writes the component metadata. The tensor
|
||||
files remain byte-identical hard links.
|
||||
|
||||
| Checkpoint path | Pipeline role |
|
||||
| ------------------------ | ------------------------------------------------------------------ |
|
||||
| `preview/` | Joint video-audio preview transformer. |
|
||||
| `refiner/` | 1080p video refiner transformer. |
|
||||
| `text_encoder/` | Qwen3.5-27B prompt encoder and tokenizer. |
|
||||
| `vae/Wan2.2_VAE.pth` | Wan 2.2 reference-image encoder for I2V. |
|
||||
| `turbo_vae/` | Default distilled decoder; first window 7 latents, step 7 latents. |
|
||||
| `stable-audio-open-1.0/` | Stable Audio variational autoencoder (VAE) audio decoder. |
|
||||
|
||||
The checkpoint snapshot contains 306,721,986,476 bytes. The converted repository
|
||||
uses hard links, so the source and converted layouts share the same file data on
|
||||
the local filesystem.
|
||||
|
||||
Run the converter with:
|
||||
|
||||
```bash
|
||||
venv-port-magi-2/bin/python \
|
||||
scripts/checkpoint_conversion/convert_magi2_to_fastvideo.py \
|
||||
--source official_weights/magi2 \
|
||||
--output converted_weights/magi2
|
||||
```
|
||||
|
||||
## Environment
|
||||
|
||||
The Python environment is `venv-port-magi-2/` at the FastVideo worktree root.
|
||||
It uses Python 3.12, PyTorch 2.11 with CUDA 12.8, and one numeric software stack
|
||||
for the official implementation and the FastVideo implementation.
|
||||
|
||||
Setup commands:
|
||||
|
||||
```bash
|
||||
uv venv --python /mnt/weka/home/junda.su/.local/bin/python3.12 venv-port-magi-2
|
||||
UV_TORCH_BACKEND=cu128 uv pip install \
|
||||
--python venv-port-magi-2/bin/python \
|
||||
pip setuptools wheel ninja torch==2.11.0 torchvision torchaudio
|
||||
```
|
||||
|
||||
The official code imports Flash Attention 3 through `flash_attn_interface`.
|
||||
There is no wheel for the official pinned revision and this Python, CUDA, and
|
||||
PyTorch combination. Build revision
|
||||
`b613d9e2c8475945baff3fd68f2030af1b890acf` with `MAX_JOBS=16`. Do not increase
|
||||
the build job count.
|
||||
|
||||
Initialize the pinned MagiAttention dependency before installation:
|
||||
|
||||
```bash
|
||||
git -C fastvideo/third_party/magi_attention submodule update --init --recursive
|
||||
```
|
||||
|
||||
The environment setup record and command logs are stored under
|
||||
`archived/magi2_parity/env_setup/`.
|
||||
|
||||
## Official Execution Contract
|
||||
|
||||
Use eight Hopper GPUs. Set both deterministic environment variables for strict
|
||||
official parity because the official `--deterministic` implementation does not
|
||||
set the environment variable that controls deterministic mixture-of-experts
|
||||
(MoE) scatter:
|
||||
|
||||
```bash
|
||||
MAGI2_DETERMINISTIC=1 \
|
||||
MAGI_ATTENTION_DETERMINISTIC_MODE=1 \
|
||||
MAGI2_SAVE_LATENT_PATH=archived/magi2_parity/validation/reference_latents \
|
||||
torchrun --nproc_per_node=8 inference/pipeline/entry.py \
|
||||
--resolution 1080p \
|
||||
--seconds 10 \
|
||||
--seed 42 \
|
||||
--prompt-file assets/sample_000.txt \
|
||||
--output archived/magi2_parity/validation/reference_t2v
|
||||
```
|
||||
|
||||
Add `--image assets/sample_000.jpeg` for I2V. The FastVideo deterministic option
|
||||
sets fixed Python, NumPy, PyTorch, and CUDA seeds and activates both MAGI-2
|
||||
kernel determinism controls.
|
||||
|
||||
## Required Parity Coverage
|
||||
|
||||
Each component test loads the official component and the production FastVideo
|
||||
component. A reused FastVideo component requires a non-skipped parity pass.
|
||||
Strict parity compares shape, dtype, stride, and every tensor value.
|
||||
|
||||
| Scope | Test |
|
||||
| --------------------------- | ------------------------------------------ |
|
||||
| Prompt and Qwen3.5 encoding | `test_magi2_text_encoder_parity.py` |
|
||||
| I2V Wan VAE encoding | `test_magi2_image_encoder_parity.py` |
|
||||
| Preview token packing | `test_magi2_preview_data_proxy_parity.py` |
|
||||
| Preview transformer | `test_magi2_preview_transformer_parity.py` |
|
||||
| Refiner token packing | `test_magi2_refiner_data_proxy_parity.py` |
|
||||
| Refiner transformer | `test_magi2_refiner_transformer_parity.py` |
|
||||
| Flow UniPC scheduler | `test_magi2_scheduler_parity.py` |
|
||||
| Turbo VAE decoding | `test_magi2_turbo_vae_parity.py` |
|
||||
| Stable Audio VAE decoding | `test_magi2_audio_vae_parity.py` |
|
||||
| Runtime controls | `test_magi2_runtime_controls.py` |
|
||||
| Checkpoint conversion | `test_magi2_checkpoint_conversion.py` |
|
||||
| Registry and metadata | `test_magi2_registry_and_metadata.py` |
|
||||
| T2V and I2V pipeline stages | `test_magi2_pipeline_parity.py` |
|
||||
| Decoded video and audio | `test_magi2_pipeline_parity.py` |
|
||||
|
||||
Run the local suite with:
|
||||
|
||||
```bash
|
||||
venv-port-magi-2/bin/python -m pytest tests/local_tests/magi2 -sv
|
||||
```
|
||||
|
||||
The component tests compare prompt normalization, tokenization, image encoding,
|
||||
scheduler arrays, packed model inputs, every preview and refiner transformer
|
||||
layer boundary, all three Turbo VAE temporal-window roles, audio decoding before
|
||||
resampling, and resampled audio. The pipeline test compares stage-boundary
|
||||
digests for conditioned prompts, reference-image latents, text embeddings,
|
||||
initial video and audio noise, preview outputs, the noise-injected refiner input,
|
||||
the post-refiner latent, and decoded outputs. End-to-end parity compares the
|
||||
decoded `uint8 [T, H, W, 3]` video array and the sample-major stereo audio array
|
||||
before media encoding.
|
||||
|
||||
## Release-Profile Result
|
||||
|
||||
The official and FastVideo workers ran both T2V and I2V with seed 42, 100
|
||||
preview steps, 5 refiner steps, and deterministic MAGI-2 and MagiAttention
|
||||
kernels. The manifest comparison found one expected metadata difference: the
|
||||
`implementation` label. All captured tensor metadata and SHA-256 digests match.
|
||||
|
||||
| Artifact | Path |
|
||||
| ------------------- | ----------------------------------------------------------------------- |
|
||||
| Official manifest | `archived/magi2_parity/validation/pipeline/official/capture.json` |
|
||||
| FastVideo manifest | `archived/magi2_parity/validation/pipeline/fastvideo/capture.json` |
|
||||
| Completion report | `tests/local_tests/magi2/PORT_STATUS.md` |
|
||||
| Environment record | `archived/magi2_parity/env_setup/SETUP_SUMMARY.md` |
|
||||
|
||||
## Review Requirements
|
||||
|
||||
The port is ready for review when all required component tests and the combined
|
||||
T2V/I2V pipeline test report non-skipped strict parity passes. Review also
|
||||
verifies strict checkpoint loading, Turbo VAE default selection, one-rank-per-
|
||||
video decoding, deterministic controls, post-refiner latent saving through
|
||||
`MAGI2_SAVE_LATENT_PATH`, T2V behavior, I2V behavior, and audio output.
|
||||
|
||||
Generated tensors, checkpoints, converted weights, compiler caches, and run
|
||||
logs remain outside version control.
|
||||
@@ -0,0 +1,96 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared utilities for MAGI-2 component parity scaffolds."""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(
|
||||
os.getenv(
|
||||
"MAGI2_OFFICIAL_REF_DIR",
|
||||
"/mnt/weka/shrd/wm/junda/fv-hub/MAGI-2-preview",
|
||||
)
|
||||
)
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
|
||||
|
||||
def require_path(path: Path, purpose: str) -> Path:
|
||||
"""Return a required local path or skip with its concrete purpose."""
|
||||
if not path.exists():
|
||||
pytest.skip(f"Missing {purpose}: {path}")
|
||||
return path
|
||||
|
||||
|
||||
def require_complete_safetensor_index(component_dir: Path) -> Path:
|
||||
"""Require every shard named by a component's safetensors index."""
|
||||
index_path = require_path(
|
||||
component_dir / "model.safetensors.index.json",
|
||||
"safetensors index",
|
||||
)
|
||||
with index_path.open(encoding="utf-8") as index_file:
|
||||
weight_index = json.load(index_file)
|
||||
shard_names = sorted(set(weight_index["weight_map"].values()))
|
||||
missing_shards = [name for name in shard_names if not (component_dir / name).is_file()]
|
||||
if missing_shards:
|
||||
pytest.skip(
|
||||
f"Incomplete weights in {component_dir}; missing shards: "
|
||||
f"{', '.join(missing_shards)}"
|
||||
)
|
||||
return component_dir
|
||||
|
||||
|
||||
def import_official_module(module_name: str):
|
||||
"""Import one module from the user-specified official MAGI-2 checkout."""
|
||||
require_path(OFFICIAL_REF_DIR, "official MAGI-2 checkout")
|
||||
reference_path = str(OFFICIAL_REF_DIR)
|
||||
if reference_path not in sys.path:
|
||||
sys.path.insert(0, reference_path)
|
||||
try:
|
||||
official_module = importlib.import_module(module_name)
|
||||
except ModuleNotFoundError as exc:
|
||||
if exc.name == module_name or module_name.startswith(f"{exc.name}."):
|
||||
raise
|
||||
pytest.skip(
|
||||
f"Official module {module_name} requires missing dependency {exc.name!r}"
|
||||
)
|
||||
module_path = Path(official_module.__file__).resolve()
|
||||
if not module_path.is_relative_to(OFFICIAL_REF_DIR.resolve()):
|
||||
raise RuntimeError(
|
||||
f"Imported {module_name} from {module_path}, expected {OFFICIAL_REF_DIR}"
|
||||
)
|
||||
return official_module
|
||||
|
||||
|
||||
def assert_tensor_exact(
|
||||
fastvideo_tensor: torch.Tensor,
|
||||
official_tensor: torch.Tensor,
|
||||
tensor_name: str,
|
||||
) -> None:
|
||||
"""Require identical tensor metadata and identical values."""
|
||||
assert fastvideo_tensor.shape == official_tensor.shape, tensor_name
|
||||
assert fastvideo_tensor.dtype == official_tensor.dtype, tensor_name
|
||||
assert fastvideo_tensor.stride() == official_tensor.stride(), tensor_name
|
||||
assert torch.equal(fastvideo_tensor, official_tensor), tensor_name
|
||||
|
||||
|
||||
def assert_array_exact(
|
||||
fastvideo_array: np.ndarray,
|
||||
official_array: np.ndarray,
|
||||
array_name: str,
|
||||
) -> None:
|
||||
"""Require identical NumPy array metadata and identical values."""
|
||||
assert fastvideo_array.shape == official_array.shape, array_name
|
||||
assert fastvideo_array.dtype == official_array.dtype, array_name
|
||||
assert fastvideo_array.strides == official_array.strides, array_name
|
||||
assert np.array_equal(fastvideo_array, official_array), array_name
|
||||
@@ -0,0 +1,502 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Capture strict MAGI-2 pipeline boundaries from one implementation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import random
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
CONVERTED_ROOT = Path(
|
||||
os.environ.get("MAGI2_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "magi2")
|
||||
)
|
||||
WORLD_SIZE = 8
|
||||
SEED = 42
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
"""Parse the implementation, denoising counts, and requested modalities."""
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--implementation",
|
||||
choices=("official", "fastvideo"),
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument("--output-dir", type=Path, required=True)
|
||||
parser.add_argument("--preview-steps", type=int, default=100)
|
||||
parser.add_argument("--refiner-steps", type=int, default=5)
|
||||
parser.add_argument(
|
||||
"--cases",
|
||||
nargs="+",
|
||||
choices=("t2v", "i2v"),
|
||||
default=("t2v", "i2v"),
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _seed_all() -> None:
|
||||
"""Reset every random-number source used by the release pipeline."""
|
||||
random.seed(SEED)
|
||||
np.random.seed(SEED)
|
||||
torch.manual_seed(SEED)
|
||||
torch.cuda.manual_seed_all(SEED)
|
||||
|
||||
|
||||
def _enable_determinism() -> None:
|
||||
"""Apply the deterministic controls exposed by the official entry point."""
|
||||
os.environ["MAGI2_DETERMINISTIC"] = "1"
|
||||
os.environ["MAGI_ATTENTION_DETERMINISTIC_MODE"] = "1"
|
||||
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
||||
_seed_all()
|
||||
torch.use_deterministic_algorithms(True)
|
||||
|
||||
|
||||
def _hash_array(array: np.ndarray) -> str:
|
||||
"""Hash the exact bytes of one contiguous NumPy array without a byte copy."""
|
||||
contiguous_array = np.ascontiguousarray(array)
|
||||
return hashlib.sha256(memoryview(contiguous_array).cast("B")).hexdigest()
|
||||
|
||||
|
||||
def _tensor_digest(tensor: torch.Tensor | None) -> dict[str, Any] | None:
|
||||
"""Describe and hash a tensor while preserving its source metadata."""
|
||||
if tensor is None:
|
||||
return None
|
||||
source_tensor = tensor.detach()
|
||||
byte_array = source_tensor.to(device="cpu").contiguous().view(torch.uint8).numpy()
|
||||
return {
|
||||
"shape": list(source_tensor.shape),
|
||||
"dtype": str(source_tensor.dtype),
|
||||
"stride": list(source_tensor.stride()),
|
||||
"sha256": _hash_array(byte_array),
|
||||
}
|
||||
|
||||
|
||||
def _array_digest(array: np.ndarray | None) -> dict[str, Any] | None:
|
||||
"""Describe and hash a NumPy array while preserving its source metadata."""
|
||||
if array is None:
|
||||
return None
|
||||
source_array = np.asarray(array)
|
||||
return {
|
||||
"shape": list(source_array.shape),
|
||||
"dtype": str(source_array.dtype),
|
||||
"stride": list(source_array.strides),
|
||||
"sha256": _hash_array(source_array),
|
||||
}
|
||||
|
||||
|
||||
def _case_inputs(case_name: str) -> tuple[str, str | None]:
|
||||
"""Return the official release prompt and optional reference image path."""
|
||||
if case_name == "t2v":
|
||||
prompt_path = OFFICIAL_ROOT / "assets" / "sample_enhanced_t2v.json"
|
||||
return prompt_path.read_text(encoding="utf-8").strip(), None
|
||||
if case_name == "i2v":
|
||||
prompt_path = OFFICIAL_ROOT / "assets" / "sample_000.txt"
|
||||
image_path = OFFICIAL_ROOT / "assets" / "sample_000.jpeg"
|
||||
return prompt_path.read_text(encoding="utf-8").strip(), str(image_path)
|
||||
raise ValueError(f"Unknown MAGI-2 parity case: {case_name}")
|
||||
|
||||
|
||||
def _capture_official_case(
|
||||
engine: Any,
|
||||
case_name: str,
|
||||
preview_steps: int,
|
||||
refiner_steps: int,
|
||||
is_capture_rank: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Run the official engine and capture every externally visible stage boundary."""
|
||||
prompt, image_path = _case_inputs(case_name)
|
||||
capture: dict[str, Any] = {}
|
||||
text_call_count = 0
|
||||
|
||||
original_encode_images = engine._encode_images
|
||||
original_get_text_embedding = engine.get_text_embedding
|
||||
original_get_special_token = engine.get_special_token
|
||||
original_preview_sample = engine.sampler.sample
|
||||
original_refiner_forward = engine._forward_magi2_refiner
|
||||
original_refiner_sample = engine.evaluate_magi2_refiner_with_latent
|
||||
|
||||
def capture_images(image: Any, height: int, width: int):
|
||||
"""Record the Wan image-encoder output and its figure identity."""
|
||||
encoded_images = original_encode_images(image, height, width)
|
||||
if is_capture_rank:
|
||||
reference_latent, reference_length, reference_ids = encoded_images
|
||||
capture["reference_latent"] = _tensor_digest(reference_latent)
|
||||
capture["reference_length"] = _tensor_digest(reference_length)
|
||||
capture["reference_ids"] = _tensor_digest(reference_ids)
|
||||
return encoded_images
|
||||
|
||||
def capture_text(prompt_text: str) -> torch.Tensor:
|
||||
"""Record positive and negative Qwen3.5 conditioning in call order."""
|
||||
nonlocal text_call_count
|
||||
text_embedding = original_get_text_embedding(prompt_text)
|
||||
if is_capture_rank:
|
||||
if text_call_count == 0:
|
||||
capture["conditioned_prompt"] = prompt_text
|
||||
capture["positive_text"] = _tensor_digest(text_embedding)
|
||||
elif text_call_count == 1:
|
||||
capture["negative_prompt"] = prompt_text
|
||||
capture["negative_text"] = _tensor_digest(text_embedding)
|
||||
else:
|
||||
raise RuntimeError("MAGI-2 encoded more than two prompts in one request")
|
||||
text_call_count += 1
|
||||
return text_embedding
|
||||
|
||||
def capture_special_token(
|
||||
prompt_text: str,
|
||||
figure_ids: torch.Tensor,
|
||||
text_feature: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Record the pooled text embedding that prefixes image patches."""
|
||||
special_token = original_get_special_token(
|
||||
prompt_text,
|
||||
figure_ids,
|
||||
text_feature,
|
||||
)
|
||||
if is_capture_rank:
|
||||
capture["special_tokens"] = _tensor_digest(special_token.unsqueeze(0))
|
||||
return special_token
|
||||
|
||||
def capture_preview(sampler_input: Any) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Record the initial noise and the completed preview latents."""
|
||||
if is_capture_rank:
|
||||
capture["initial_video_noise"] = _tensor_digest(sampler_input.latent)
|
||||
capture["initial_audio_noise"] = _tensor_digest(
|
||||
sampler_input.audio_latent
|
||||
)
|
||||
preview_video, preview_audio = original_preview_sample(sampler_input)
|
||||
if is_capture_rank:
|
||||
capture["preview_video"] = _tensor_digest(preview_video)
|
||||
capture["preview_audio"] = _tensor_digest(preview_audio)
|
||||
return preview_video, preview_audio
|
||||
|
||||
def capture_refiner_forward(
|
||||
latent_video: torch.Tensor,
|
||||
latent_audio: torch.Tensor,
|
||||
txt_feat: torch.Tensor,
|
||||
ref_audio_feat: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Record the noise-injected 1080p latent entering the refiner."""
|
||||
if is_capture_rank and "refiner_input" not in capture:
|
||||
capture["refiner_input"] = _tensor_digest(latent_video)
|
||||
return original_refiner_forward(
|
||||
latent_video,
|
||||
latent_audio,
|
||||
txt_feat,
|
||||
ref_audio_feat,
|
||||
)
|
||||
|
||||
def capture_refiner_sample(**kwargs: Any) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Record the video latent that the refiner passes to the decoder."""
|
||||
refined_video, refined_audio = original_refiner_sample(**kwargs)
|
||||
if is_capture_rank:
|
||||
capture["refined_video"] = _tensor_digest(refined_video)
|
||||
return refined_video, refined_audio
|
||||
|
||||
engine._encode_images = capture_images
|
||||
engine.get_text_embedding = capture_text
|
||||
engine.get_special_token = capture_special_token
|
||||
engine.sampler.sample = capture_preview
|
||||
engine._forward_magi2_refiner = capture_refiner_forward
|
||||
engine.evaluate_magi2_refiner_with_latent = capture_refiner_sample
|
||||
_seed_all()
|
||||
try:
|
||||
video, audio = engine.evaluate(
|
||||
prompt=prompt,
|
||||
image=image_path,
|
||||
eval_task_type="text2video" if image_path is None else "image2video",
|
||||
seconds=10.0,
|
||||
preview_width=896,
|
||||
preview_height=512,
|
||||
br_num_inference_steps=preview_steps,
|
||||
refiner_width=1920,
|
||||
refiner_height=1088,
|
||||
magi2_refiner_num_inference_steps=refiner_steps,
|
||||
)
|
||||
finally:
|
||||
engine._encode_images = original_encode_images
|
||||
engine.get_text_embedding = original_get_text_embedding
|
||||
engine.get_special_token = original_get_special_token
|
||||
engine.sampler.sample = original_preview_sample
|
||||
engine._forward_magi2_refiner = original_refiner_forward
|
||||
engine.evaluate_magi2_refiner_with_latent = original_refiner_sample
|
||||
|
||||
if is_capture_rank:
|
||||
capture.setdefault("special_tokens", None)
|
||||
capture["decoded_video"] = _array_digest(video)
|
||||
capture["decoded_audio"] = _array_digest(audio)
|
||||
return capture
|
||||
|
||||
|
||||
def _load_official_engine() -> Any:
|
||||
"""Initialize the official eight-rank runtime and load every release component."""
|
||||
sys.path.insert(0, str(OFFICIAL_ROOT))
|
||||
from inference.common.magi2_config import load_config
|
||||
from inference.infra.checkpoint.load_checkpoint import (
|
||||
load_magi2_model,
|
||||
load_magi2_refiner,
|
||||
)
|
||||
from inference.infra.distributed import (
|
||||
initialize_expert_parallel,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
|
||||
config = load_config(str(OFFICIAL_ROOT / "configs" / "magi2_refiner.json"))
|
||||
dist.init_process_group(backend="nccl")
|
||||
rank = dist.get_rank()
|
||||
if dist.get_world_size() != WORLD_SIZE:
|
||||
raise RuntimeError(
|
||||
f"MAGI-2 pipeline parity requires {WORLD_SIZE} ranks, "
|
||||
f"received {dist.get_world_size()}"
|
||||
)
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
torch.cuda.set_device(local_rank)
|
||||
initialize_model_parallel(cp_size=WORLD_SIZE)
|
||||
initialize_expert_parallel(ep_size=WORLD_SIZE)
|
||||
|
||||
from inference.pipeline.inference_engine import Magi2InferenceEngine
|
||||
|
||||
preview_model = load_magi2_model(config)
|
||||
refiner_model = load_magi2_refiner(config)
|
||||
return Magi2InferenceEngine(
|
||||
model=preview_model,
|
||||
config=config.evaluation_config,
|
||||
device=f"cuda:{local_rank}",
|
||||
weight_dtype=torch.bfloat16,
|
||||
magi2_refiner=refiner_model,
|
||||
)
|
||||
|
||||
|
||||
def _canonical_fastvideo_video(video: torch.Tensor | None) -> np.ndarray | None:
|
||||
"""Convert FastVideo's BCHWT float output into the release's THWC bytes."""
|
||||
if video is None:
|
||||
return None
|
||||
if video.shape[0] != 1:
|
||||
raise ValueError(f"MAGI-2 parity expects one decoded video, received {video.shape}")
|
||||
return video[0].permute(1, 2, 3, 0).mul(255).numpy().astype(np.uint8)
|
||||
|
||||
|
||||
def _capture_fastvideo_case(
|
||||
pipeline: Any,
|
||||
fastvideo_args: Any,
|
||||
case_name: str,
|
||||
preview_steps: int,
|
||||
refiner_steps: int,
|
||||
is_capture_rank: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Run FastVideo stages and capture the boundaries used by the official worker."""
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
prompt, image_path = _case_inputs(case_name)
|
||||
fastvideo_args.workload_type = (
|
||||
WorkloadType.T2V if case_name == "t2v" else WorkloadType.I2V
|
||||
)
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=prompt,
|
||||
negative_prompt=pipeline.negative_prompt,
|
||||
image_path=image_path,
|
||||
seed=SEED,
|
||||
num_frames=249,
|
||||
height=1088,
|
||||
width=1920,
|
||||
fps=25,
|
||||
num_inference_steps=preview_steps,
|
||||
num_inference_steps_sr=refiner_steps,
|
||||
)
|
||||
capture: dict[str, Any] = {}
|
||||
refiner_stage = pipeline.refiner_stage
|
||||
original_refiner_predict = refiner_stage._predict_velocity
|
||||
|
||||
def capture_refiner_predict(
|
||||
video_latent: torch.Tensor,
|
||||
audio_latent: torch.Tensor,
|
||||
text_context: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Record the noise-injected latent before the first refiner forward."""
|
||||
if is_capture_rank and "refiner_input" not in capture:
|
||||
capture["refiner_input"] = _tensor_digest(video_latent)
|
||||
return original_refiner_predict(video_latent, audio_latent, text_context)
|
||||
|
||||
refiner_stage._predict_velocity = capture_refiner_predict
|
||||
_seed_all()
|
||||
try:
|
||||
for stage_name, stage in pipeline._stage_name_mapping.items():
|
||||
batch = stage(batch, fastvideo_args)
|
||||
if not is_capture_rank:
|
||||
continue
|
||||
if stage_name == "reference_image_stage":
|
||||
capture["conditioned_prompt"] = batch.prompt
|
||||
capture["reference_latent"] = _tensor_digest(
|
||||
batch.magi2_ref_image_feat
|
||||
)
|
||||
capture["reference_length"] = _tensor_digest(
|
||||
batch.magi2_ref_image_feat_len
|
||||
)
|
||||
reference_ids = (
|
||||
None
|
||||
if batch.magi2_ref_image_feat is None
|
||||
else torch.tensor([[1]], dtype=torch.long)
|
||||
)
|
||||
capture["reference_ids"] = _tensor_digest(reference_ids)
|
||||
elif stage_name == "text_encoding_stage":
|
||||
capture["negative_prompt"] = batch.negative_prompt
|
||||
capture["positive_text"] = _tensor_digest(batch.magi2_text_context)
|
||||
capture["negative_text"] = _tensor_digest(
|
||||
batch.magi2_negative_context
|
||||
)
|
||||
capture["special_tokens"] = _tensor_digest(
|
||||
batch.magi2_ref_image_special_tokens
|
||||
)
|
||||
elif stage_name == "latent_preparation_stage":
|
||||
capture["initial_video_noise"] = _tensor_digest(batch.latents)
|
||||
capture["initial_audio_noise"] = _tensor_digest(
|
||||
batch.audio_latents
|
||||
)
|
||||
elif stage_name == "preview_denoising_stage":
|
||||
capture["preview_video"] = _tensor_digest(batch.latents)
|
||||
capture["preview_audio"] = _tensor_digest(batch.audio_latents)
|
||||
elif stage_name == "refiner_stage":
|
||||
capture["refined_video"] = _tensor_digest(batch.latents)
|
||||
elif stage_name == "video_decoding_stage":
|
||||
canonical_video = _canonical_fastvideo_video(batch.output)
|
||||
capture["decoded_video"] = _array_digest(canonical_video)
|
||||
elif stage_name == "audio_decoding_stage":
|
||||
capture["decoded_audio"] = _array_digest(batch.extra.get("audio"))
|
||||
finally:
|
||||
refiner_stage._predict_velocity = original_refiner_predict
|
||||
return capture
|
||||
|
||||
|
||||
def _load_fastvideo_pipeline() -> tuple[Any, Any]:
|
||||
"""Initialize FastVideo's eight-rank runtime and load the converted checkpoint."""
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
|
||||
from fastvideo.pipelines.basic.magi2.magi2_pipeline import Magi2Pipeline
|
||||
from fastvideo.pipelines.basic.magi2.pipeline_configs import (
|
||||
Magi2PreviewPipelineConfig,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.presets import MAGI2_NEGATIVE_PROMPT
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=str(CONVERTED_ROOT),
|
||||
workload_type=WorkloadType.T2V,
|
||||
num_gpus=WORLD_SIZE,
|
||||
tp_size=1,
|
||||
sp_size=WORLD_SIZE,
|
||||
pipeline_config=Magi2PreviewPipelineConfig(),
|
||||
deterministic=True,
|
||||
dit_cpu_offload=True,
|
||||
dit_layerwise_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
image_encoder_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
enable_stage_verification=True,
|
||||
)
|
||||
pipeline = Magi2Pipeline(str(CONVERTED_ROOT), fastvideo_args)
|
||||
pipeline.negative_prompt = MAGI2_NEGATIVE_PROMPT
|
||||
pipeline.post_init()
|
||||
if dist.get_world_size() != WORLD_SIZE:
|
||||
raise RuntimeError(
|
||||
f"MAGI-2 pipeline parity requires {WORLD_SIZE} ranks, "
|
||||
f"received {dist.get_world_size()}"
|
||||
)
|
||||
return pipeline, fastvideo_args
|
||||
|
||||
|
||||
def _write_capture(
|
||||
output_dir: Path,
|
||||
implementation: str,
|
||||
preview_steps: int,
|
||||
refiner_steps: int,
|
||||
captures: dict[str, Any],
|
||||
) -> None:
|
||||
"""Atomically write the rank-zero digest manifest for one implementation."""
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
artifact = {
|
||||
"schema_version": 1,
|
||||
"implementation": implementation,
|
||||
"world_size": WORLD_SIZE,
|
||||
"preview_steps": preview_steps,
|
||||
"refiner_steps": refiner_steps,
|
||||
"cases": captures,
|
||||
}
|
||||
artifact_path = output_dir / "capture.json"
|
||||
temporary_path = output_dir / "capture.json.tmp"
|
||||
temporary_path.write_text(
|
||||
json.dumps(artifact, indent=2, sort_keys=True),
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.replace(temporary_path, artifact_path)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Load one implementation, capture all cases, and persist rank-zero digests."""
|
||||
args = _parse_args()
|
||||
if args.preview_steps <= 0 or args.refiner_steps <= 0:
|
||||
raise ValueError("MAGI-2 parity step counts must be positive")
|
||||
os.environ["MAGI2_CKPT_ROOT"] = str(WEIGHTS_ROOT)
|
||||
os.environ.pop("MAGI2_SAVE_LATENT_PATH", None)
|
||||
os.environ.pop("NEGATIVE_PROMPT", None)
|
||||
os.environ.pop("SKIP_LOAD_MODEL", None)
|
||||
_enable_determinism()
|
||||
|
||||
if args.implementation == "official":
|
||||
engine = _load_official_engine()
|
||||
is_capture_rank = dist.get_rank() == 0
|
||||
captures = {
|
||||
case_name: _capture_official_case(
|
||||
engine,
|
||||
case_name,
|
||||
args.preview_steps,
|
||||
args.refiner_steps,
|
||||
is_capture_rank,
|
||||
)
|
||||
for case_name in args.cases
|
||||
}
|
||||
else:
|
||||
pipeline, fastvideo_args = _load_fastvideo_pipeline()
|
||||
is_capture_rank = dist.get_rank() == 0
|
||||
captures = {
|
||||
case_name: _capture_fastvideo_case(
|
||||
pipeline,
|
||||
fastvideo_args,
|
||||
case_name,
|
||||
args.preview_steps,
|
||||
args.refiner_steps,
|
||||
is_capture_rank,
|
||||
)
|
||||
for case_name in args.cases
|
||||
}
|
||||
|
||||
if is_capture_rank:
|
||||
_write_capture(
|
||||
args.output_dir,
|
||||
args.implementation,
|
||||
args.preview_steps,
|
||||
args.refiner_steps,
|
||||
captures,
|
||||
)
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,364 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run one isolated side of MAGI-2 preview transformer numerical parity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
import random
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
WORLD_SIZE = 8
|
||||
SEED = 42
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
"""Parse the implementation and artifact directory for one torchrun job."""
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--implementation",
|
||||
choices=("official", "fastvideo"),
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument("--output-dir", type=Path, required=True)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _enable_determinism() -> None:
|
||||
"""Apply the deterministic controls used by the official entry point."""
|
||||
os.environ["MAGI2_DETERMINISTIC"] = "1"
|
||||
os.environ["MAGI_ATTENTION_DETERMINISTIC_MODE"] = "1"
|
||||
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
||||
random.seed(SEED)
|
||||
np.random.seed(SEED)
|
||||
torch.manual_seed(SEED)
|
||||
torch.cuda.manual_seed_all(SEED)
|
||||
torch.use_deterministic_algorithms(True)
|
||||
|
||||
|
||||
def _initialize_distributed(implementation: str) -> tuple[int, torch.device]:
|
||||
"""Initialize the official eight-rank CP and EP process-group topology."""
|
||||
if dist.is_initialized():
|
||||
raise RuntimeError("The parity worker requires a fresh distributed process")
|
||||
dist.init_process_group(backend="nccl")
|
||||
rank = dist.get_rank()
|
||||
if dist.get_world_size() != WORLD_SIZE:
|
||||
raise RuntimeError(
|
||||
f"MAGI-2 preview parity requires {WORLD_SIZE} ranks, "
|
||||
f"received {dist.get_world_size()}"
|
||||
)
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
if implementation == "official":
|
||||
sys.path.insert(0, str(OFFICIAL_ROOT))
|
||||
from inference.infra.distributed import (
|
||||
initialize_expert_parallel,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
else:
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
from fastvideo.models.dits.magi2_runtime import (
|
||||
initialize_expert_parallel,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
|
||||
initialize_model_parallel(cp_size=WORLD_SIZE)
|
||||
initialize_expert_parallel(ep_size=WORLD_SIZE)
|
||||
return rank, torch.device("cuda", local_rank)
|
||||
|
||||
|
||||
def _load_runtime(implementation: str, device: torch.device) -> tuple[Any, Any, type]:
|
||||
"""Load one implementation's production transformer, proxy, and input type."""
|
||||
if implementation == "official":
|
||||
from inference.common.magi2_config import load_config
|
||||
from inference.infra.checkpoint.load_checkpoint import load_magi2_model
|
||||
from inference.pipeline.preview_data_proxy import Magi2DataProxy, ModelInput
|
||||
|
||||
config = load_config(str(OFFICIAL_ROOT / "configs" / "magi2_preview.json"))
|
||||
config.engine_config.load = str(WEIGHTS_ROOT / "preview")
|
||||
config.engine_config.cp_size = WORLD_SIZE
|
||||
config.engine_config.ep_size = WORLD_SIZE
|
||||
model = load_magi2_model(config)
|
||||
proxy = Magi2DataProxy(config.evaluation_config.data_proxy_config)
|
||||
return model, proxy, ModelInput
|
||||
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2PreviewVideoConfig
|
||||
from fastvideo.models.dits.magi2_loader import load_magi2_preview_model
|
||||
from fastvideo.pipelines.basic.magi2.stages.preview_data_proxy import (
|
||||
Magi2DataProxy,
|
||||
Magi2DataProxyConfig,
|
||||
ModelInput,
|
||||
)
|
||||
|
||||
with (OFFICIAL_ROOT / "configs" / "magi2_preview.json").open(
|
||||
encoding="utf-8"
|
||||
) as config_file:
|
||||
proxy_values = json.load(config_file)["evaluation_config"]["data_proxy_config"]
|
||||
model = load_magi2_preview_model(
|
||||
checkpoint_dir=str(WEIGHTS_ROOT / "preview"),
|
||||
config=Magi2PreviewVideoConfig(),
|
||||
device=device,
|
||||
)
|
||||
proxy = Magi2DataProxy(Magi2DataProxyConfig(**proxy_values))
|
||||
return model, proxy, ModelInput
|
||||
|
||||
|
||||
def _pattern(
|
||||
shape: tuple[int, ...],
|
||||
*,
|
||||
offset: int,
|
||||
scale: float,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
"""Create a stable, nonuniform tensor without consuming random numbers."""
|
||||
element_count = math.prod(shape)
|
||||
values = torch.arange(element_count, dtype=torch.float32)
|
||||
values = ((values + offset).remainder(127) - 63) * scale
|
||||
return values.reshape(shape).to(dtype=dtype)
|
||||
|
||||
|
||||
def _build_model_input(
|
||||
model_input_type: type,
|
||||
mode: str,
|
||||
device: torch.device,
|
||||
) -> Any:
|
||||
"""Build a compact CFG input that exercises every preview modality."""
|
||||
video = _pattern((1, 48, 2, 2, 2), offset=3, scale=1 / 64).repeat(2, 1, 1, 1, 1)
|
||||
audio = _pattern((1, 4, 64), offset=11, scale=1 / 32).repeat(2, 1, 1)
|
||||
text = _pattern(
|
||||
(2, 3, 5120),
|
||||
offset=29,
|
||||
scale=1 / 128,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
normalized_time = torch.tensor([0.75, 0.75], dtype=torch.float32)
|
||||
input_values: dict[str, Any] = {
|
||||
"x_t": video.to(device),
|
||||
"audio_x_t": audio.to(device),
|
||||
"audio_feat_len": torch.tensor([4, 4], dtype=torch.int32, device=device),
|
||||
"txt_feat": text.to(device),
|
||||
"txt_feat_len": torch.tensor([3, 2], dtype=torch.int32, device=device),
|
||||
"t": normalized_time.to(device),
|
||||
"per_token_video_t": normalized_time.view(2, 1, 1, 1, 1)
|
||||
.expand(2, 1, 2, 2, 2)
|
||||
.clone()
|
||||
.to(device),
|
||||
"per_token_audio_t": normalized_time.view(2, 1, 1)
|
||||
.expand(2, 4, 1)
|
||||
.clone()
|
||||
.to(device),
|
||||
}
|
||||
if mode == "i2v":
|
||||
conditioned_image = _pattern(
|
||||
(1, 1, 48, 1, 2, 2),
|
||||
offset=47,
|
||||
scale=1 / 64,
|
||||
)
|
||||
conditioned_token = _pattern(
|
||||
(1, 1, 5120),
|
||||
offset=71,
|
||||
scale=1 / 128,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
input_values.update(
|
||||
{
|
||||
"ref_image_feat": torch.cat(
|
||||
[conditioned_image, torch.zeros_like(conditioned_image)],
|
||||
dim=0,
|
||||
).to(device),
|
||||
"ref_image_feat_len": torch.tensor(
|
||||
[[[2, 2]], [[2, 2]]],
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
),
|
||||
"ref_image_special_token_embedding": torch.cat(
|
||||
[conditioned_token, torch.zeros_like(conditioned_token)],
|
||||
dim=0,
|
||||
).to(device),
|
||||
}
|
||||
)
|
||||
elif mode != "t2v":
|
||||
raise ValueError(f"Unknown preview parity mode: {mode}")
|
||||
return model_input_type(**input_values)
|
||||
|
||||
|
||||
def _tensor_record(tensor: torch.Tensor) -> dict[str, Any]:
|
||||
"""Copy tensor values to CPU and retain the source tensor metadata."""
|
||||
return {
|
||||
"shape": tuple(tensor.shape),
|
||||
"dtype": str(tensor.dtype),
|
||||
"stride": tuple(tensor.stride()),
|
||||
"layout": str(tensor.layout),
|
||||
"device_type": tensor.device.type,
|
||||
"device_index": tensor.device.index,
|
||||
"requires_grad": tensor.requires_grad,
|
||||
"is_contiguous": tensor.is_contiguous(),
|
||||
"value": tensor.detach().to(device="cpu").contiguous(),
|
||||
}
|
||||
|
||||
|
||||
def _packed_input_record(packed_input: tuple[Any, ...]) -> dict[str, Any]:
|
||||
"""Record every tensor and sequence boundary passed into the transformer."""
|
||||
tokens, coords, modalities, varlen_handler, time_tokens = packed_input
|
||||
return {
|
||||
"tokens": _tensor_record(tokens),
|
||||
"coords": _tensor_record(coords),
|
||||
"modalities": _tensor_record(modalities),
|
||||
"time_tokens": _tensor_record(time_tokens),
|
||||
"cu_seqlens_q": _tensor_record(varlen_handler.cu_seqlens_q),
|
||||
"cu_seqlens_k": _tensor_record(varlen_handler.cu_seqlens_k),
|
||||
"max_seqlen_q": varlen_handler.max_seqlen_q,
|
||||
"max_seqlen_k": varlen_handler.max_seqlen_k,
|
||||
}
|
||||
|
||||
|
||||
def _capture_model_boundaries(
|
||||
model: torch.nn.Module,
|
||||
) -> tuple[dict[str, Any], list[torch.utils.hooks.RemovableHandle]]:
|
||||
"""Attach eager hooks at the adapters and every transformer layer boundary."""
|
||||
capture: dict[str, Any] = {
|
||||
"pre_adapter": {},
|
||||
"layer_boundaries": [
|
||||
{"layer_index": layer_index} for layer_index in range(len(model.block.layers))
|
||||
],
|
||||
}
|
||||
handles: list[torch.utils.hooks.RemovableHandle] = []
|
||||
|
||||
def capture_pre_adapter(
|
||||
_module: torch.nn.Module,
|
||||
_inputs: tuple[Any, ...],
|
||||
output: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> None:
|
||||
"""Record projected multimodal tokens and rotary embeddings."""
|
||||
hidden_states, rope = output
|
||||
capture["pre_adapter"] = {
|
||||
"hidden_states": _tensor_record(hidden_states),
|
||||
"rope": _tensor_record(rope),
|
||||
}
|
||||
|
||||
def make_layer_input_hook(layer_index: int):
|
||||
"""Create a hook that records one layer's input hidden states."""
|
||||
|
||||
def capture_layer_input(
|
||||
_module: torch.nn.Module,
|
||||
inputs: tuple[Any, ...],
|
||||
) -> None:
|
||||
"""Record the hidden states entering one transformer layer."""
|
||||
capture["layer_boundaries"][layer_index]["input"] = _tensor_record(
|
||||
inputs[0]
|
||||
)
|
||||
|
||||
return capture_layer_input
|
||||
|
||||
def make_layer_output_hook(layer_index: int):
|
||||
"""Create a hook that records one layer's output hidden states."""
|
||||
|
||||
def capture_layer_output(
|
||||
_module: torch.nn.Module,
|
||||
_inputs: tuple[Any, ...],
|
||||
output: torch.Tensor,
|
||||
) -> None:
|
||||
"""Record the hidden states leaving one transformer layer."""
|
||||
capture["layer_boundaries"][layer_index]["output"] = _tensor_record(output)
|
||||
|
||||
return capture_layer_output
|
||||
|
||||
def capture_post_adapter(
|
||||
_module: torch.nn.Module,
|
||||
_inputs: tuple[Any, ...],
|
||||
output: torch.Tensor,
|
||||
) -> None:
|
||||
"""Record the local video-and-audio projection before CP gathering."""
|
||||
capture["post_adapter"] = _tensor_record(output)
|
||||
|
||||
handles.append(model.pre_adapter.register_forward_hook(capture_pre_adapter))
|
||||
for layer_index, layer in enumerate(model.block.layers):
|
||||
handles.append(
|
||||
layer.register_forward_pre_hook(make_layer_input_hook(layer_index))
|
||||
)
|
||||
handles.append(
|
||||
layer.register_forward_hook(make_layer_output_hook(layer_index))
|
||||
)
|
||||
handles.append(model.post_adapter.register_forward_hook(capture_post_adapter))
|
||||
return capture, handles
|
||||
|
||||
|
||||
def _run_case(
|
||||
model: torch.nn.Module,
|
||||
proxy: Any,
|
||||
model_input_type: type,
|
||||
mode: str,
|
||||
device: torch.device,
|
||||
) -> dict[str, Any]:
|
||||
"""Run one T2V or I2V forward pass and capture all parity boundaries."""
|
||||
model_input = _build_model_input(model_input_type, mode, device)
|
||||
packed_input = proxy.process_input(model_input)
|
||||
capture, handles = _capture_model_boundaries(model)
|
||||
capture["packed_input"] = _packed_input_record(packed_input)
|
||||
try:
|
||||
with torch.inference_mode():
|
||||
model_output = model(*packed_input)
|
||||
finally:
|
||||
for handle in handles:
|
||||
handle.remove()
|
||||
video_output, audio_output = proxy.process_output(model_output)
|
||||
capture["model_output"] = _tensor_record(model_output)
|
||||
capture["depacked_output"] = {
|
||||
"video": _tensor_record(video_output),
|
||||
"audio": _tensor_record(audio_output),
|
||||
}
|
||||
if any(
|
||||
set(layer_capture) != {"layer_index", "input", "output"}
|
||||
for layer_capture in capture["layer_boundaries"]
|
||||
):
|
||||
raise RuntimeError("A preview transformer layer hook did not execute")
|
||||
return capture
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Load one implementation, run both modalities, and save rank-local captures."""
|
||||
args = _parse_args()
|
||||
os.environ.setdefault("MAGI_COMPILE_COMPILE_MODE", "NONE")
|
||||
os.environ.pop("SKIP_LOAD_MODEL", None)
|
||||
os.environ["MAGI2_CKPT_ROOT"] = str(WEIGHTS_ROOT)
|
||||
_enable_determinism()
|
||||
rank, device = _initialize_distributed(args.implementation)
|
||||
model, proxy, model_input_type = _load_runtime(args.implementation, device)
|
||||
artifact = {
|
||||
"schema_version": 1,
|
||||
"implementation": args.implementation,
|
||||
"rank": rank,
|
||||
"world_size": dist.get_world_size(),
|
||||
"cases": {
|
||||
mode: _run_case(model, proxy, model_input_type, mode, device)
|
||||
for mode in ("t2v", "i2v")
|
||||
},
|
||||
}
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
artifact_path = args.output_dir / f"rank_{rank}.pt"
|
||||
temporary_path = args.output_dir / f"rank_{rank}.pt.tmp"
|
||||
torch.save(artifact, temporary_path)
|
||||
os.replace(temporary_path, artifact_path)
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,388 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run one isolated side of MAGI-2 refiner transformer numerical parity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
import random
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
WORLD_SIZE = 8
|
||||
SEED = 42
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
"""Parse the implementation and artifact directory for one torchrun job."""
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--implementation",
|
||||
choices=("official", "fastvideo"),
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument("--output-dir", type=Path, required=True)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _enable_determinism() -> None:
|
||||
"""Apply the deterministic controls used by the official entry point."""
|
||||
os.environ["MAGI2_DETERMINISTIC"] = "1"
|
||||
os.environ["MAGI_ATTENTION_DETERMINISTIC_MODE"] = "1"
|
||||
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
||||
random.seed(SEED)
|
||||
np.random.seed(SEED)
|
||||
torch.manual_seed(SEED)
|
||||
torch.cuda.manual_seed_all(SEED)
|
||||
torch.use_deterministic_algorithms(True)
|
||||
|
||||
|
||||
def _initialize_distributed(implementation: str) -> tuple[int, torch.device]:
|
||||
"""Initialize the official eight-rank CP and EP process-group topology."""
|
||||
if dist.is_initialized():
|
||||
raise RuntimeError("The parity worker requires a fresh distributed process")
|
||||
dist.init_process_group(backend="nccl")
|
||||
rank = dist.get_rank()
|
||||
if dist.get_world_size() != WORLD_SIZE:
|
||||
raise RuntimeError(
|
||||
f"MAGI-2 refiner parity requires {WORLD_SIZE} ranks, "
|
||||
f"received {dist.get_world_size()}"
|
||||
)
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
if implementation == "official":
|
||||
sys.path.insert(0, str(OFFICIAL_ROOT))
|
||||
from inference.infra.distributed import (
|
||||
initialize_expert_parallel,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
else:
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
from fastvideo.models.dits.magi2_runtime import (
|
||||
initialize_expert_parallel,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
|
||||
initialize_model_parallel(cp_size=WORLD_SIZE)
|
||||
initialize_expert_parallel(ep_size=WORLD_SIZE)
|
||||
return rank, torch.device("cuda", local_rank)
|
||||
|
||||
|
||||
def _load_runtime(implementation: str, device: torch.device) -> tuple[Any, Any, type]:
|
||||
"""Load one implementation's production transformer, proxy, and input type."""
|
||||
if implementation == "official":
|
||||
from inference.common.magi2_config import load_config
|
||||
from inference.infra.checkpoint.load_checkpoint import load_magi2_refiner
|
||||
from inference.pipeline.inference_engine import EvalInput
|
||||
from inference.pipeline.refiner_data_proxy import Magi2RefinerDataProxy
|
||||
|
||||
config = load_config(str(OFFICIAL_ROOT / "configs" / "magi2_refiner.json"))
|
||||
config.engine_config.cp_size = WORLD_SIZE
|
||||
config.engine_config.ep_size = WORLD_SIZE
|
||||
config.evaluation_config.magi2_refiner_model_path = str(WEIGHTS_ROOT / "refiner")
|
||||
model = load_magi2_refiner(config).to(device=device).eval()
|
||||
proxy = Magi2RefinerDataProxy(
|
||||
config.evaluation_config.magi2_refiner_data_proxy_config
|
||||
)
|
||||
return model, proxy, EvalInput
|
||||
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2RefinerVideoConfig
|
||||
from fastvideo.models.dits.magi2_loader import load_magi2_refiner_model
|
||||
from fastvideo.pipelines.basic.magi2.stages.refiner_data_proxy import (
|
||||
Magi2RefinerDataProxy,
|
||||
Magi2RefinerDataProxyConfig,
|
||||
RefinerModelInput,
|
||||
)
|
||||
|
||||
with (OFFICIAL_ROOT / "configs" / "magi2_refiner.json").open(
|
||||
encoding="utf-8"
|
||||
) as config_file:
|
||||
proxy_values = json.load(config_file)["evaluation_config"][
|
||||
"magi2_refiner_data_proxy_config"
|
||||
]
|
||||
model = load_magi2_refiner_model(
|
||||
checkpoint_dir=str(WEIGHTS_ROOT / "refiner"),
|
||||
config=Magi2RefinerVideoConfig(),
|
||||
device=device,
|
||||
)
|
||||
proxy = Magi2RefinerDataProxy(Magi2RefinerDataProxyConfig(**proxy_values))
|
||||
return model, proxy, RefinerModelInput
|
||||
|
||||
|
||||
def _pattern(
|
||||
shape: tuple[int, ...],
|
||||
*,
|
||||
offset: int,
|
||||
scale: float,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
"""Create a stable, nonuniform tensor without consuming random numbers."""
|
||||
element_count = math.prod(shape)
|
||||
values = torch.arange(element_count, dtype=torch.float32)
|
||||
values = ((values + offset).remainder(127) - 63) * scale
|
||||
return values.reshape(shape).to(dtype=dtype)
|
||||
|
||||
|
||||
def _build_model_input(model_input_type: type, device: torch.device) -> Any:
|
||||
"""Build one packed input that exercises every refiner modality and CP split."""
|
||||
input_values = {
|
||||
"x_t": _pattern((1, 48, 9, 4, 24), offset=3, scale=1 / 64).to(device),
|
||||
"audio_x_t": _pattern((1, 3, 64), offset=11, scale=1 / 32).to(device),
|
||||
"audio_feat_len": torch.tensor([3], dtype=torch.int64),
|
||||
"txt_feat": _pattern((1, 5, 5120), offset=29, scale=1 / 128).to(device),
|
||||
"txt_feat_len": torch.tensor([5], dtype=torch.int64),
|
||||
"ref_audio_feat": _pattern((1, 2, 64), offset=47, scale=1 / 32).to(device),
|
||||
"ref_audio_feat_len": torch.tensor([2], dtype=torch.int64),
|
||||
"ref_video_feat": _pattern((1, 48, 9, 4, 24), offset=71, scale=1 / 64).to(device),
|
||||
"ref_video_feat_len": torch.tensor([4], dtype=torch.int64),
|
||||
}
|
||||
return model_input_type(**input_values)
|
||||
|
||||
|
||||
def _tensor_record(tensor: torch.Tensor) -> dict[str, Any]:
|
||||
"""Copy tensor values to CPU and retain the source tensor metadata."""
|
||||
return {
|
||||
"shape": tuple(tensor.shape),
|
||||
"dtype": str(tensor.dtype),
|
||||
"stride": tuple(tensor.stride()),
|
||||
"layout": str(tensor.layout),
|
||||
"device_type": tensor.device.type,
|
||||
"device_index": tensor.device.index,
|
||||
"requires_grad": tensor.requires_grad,
|
||||
"is_contiguous": tensor.is_contiguous(),
|
||||
"value": tensor.detach().to(device="cpu").contiguous(),
|
||||
}
|
||||
|
||||
|
||||
def _scalar_or_tensor_record(value: int | float | torch.Tensor | None) -> Any:
|
||||
"""Record scalar attention metadata while preserving tensor metadata."""
|
||||
if isinstance(value, torch.Tensor):
|
||||
return _tensor_record(value)
|
||||
return value
|
||||
|
||||
|
||||
def _local_attention_record(handler: Any) -> dict[str, Any] | None:
|
||||
"""Record every field that controls the refiner's local attention kernel."""
|
||||
if handler is None:
|
||||
return None
|
||||
field_names = (
|
||||
"q_ranges",
|
||||
"k_ranges",
|
||||
"max_seqlen_q",
|
||||
"max_seqlen_k",
|
||||
"attn_type_map",
|
||||
"softmax_scale",
|
||||
"bwd_q_ranges",
|
||||
"bwd_k_ranges",
|
||||
"bwd_attn_type_map",
|
||||
"auto_range_merge",
|
||||
"sparse_load",
|
||||
)
|
||||
return {
|
||||
field_name: _scalar_or_tensor_record(getattr(handler, field_name, None))
|
||||
for field_name in field_names
|
||||
}
|
||||
|
||||
|
||||
def _packed_input_record(packed_input: tuple[Any, ...]) -> dict[str, Any]:
|
||||
"""Record every tensor and sequence boundary passed into the transformer."""
|
||||
tokens, coords, modalities, varlen_handler, local_attn_handler = packed_input
|
||||
return {
|
||||
"tokens": _tensor_record(tokens),
|
||||
"coords": _tensor_record(coords),
|
||||
"modalities": _tensor_record(modalities),
|
||||
"cu_seqlens_q": _tensor_record(varlen_handler.cu_seqlens_q),
|
||||
"cu_seqlens_k": _tensor_record(varlen_handler.cu_seqlens_k),
|
||||
"max_seqlen_q": _scalar_or_tensor_record(varlen_handler.max_seqlen_q),
|
||||
"max_seqlen_k": _scalar_or_tensor_record(varlen_handler.max_seqlen_k),
|
||||
"local_attention": _local_attention_record(local_attn_handler),
|
||||
}
|
||||
|
||||
|
||||
def _capture_model_boundaries(
|
||||
model: torch.nn.Module,
|
||||
) -> tuple[dict[str, Any], list[torch.utils.hooks.RemovableHandle]]:
|
||||
"""Attach eager hooks at the adapters and all 30 transformer layers."""
|
||||
capture: dict[str, Any] = {
|
||||
"pre_adapter": {},
|
||||
"layer_boundaries": [
|
||||
{"layer_index": layer_index} for layer_index in range(len(model.block.layers))
|
||||
],
|
||||
"post_adapter": {},
|
||||
}
|
||||
if len(capture["layer_boundaries"]) != 30:
|
||||
raise RuntimeError(
|
||||
"MAGI-2 refiner parity requires 30 layer boundaries, "
|
||||
f"received {len(capture['layer_boundaries'])}"
|
||||
)
|
||||
handles: list[torch.utils.hooks.RemovableHandle] = []
|
||||
|
||||
def capture_pre_adapter(
|
||||
_module: torch.nn.Module,
|
||||
inputs: tuple[Any, ...],
|
||||
output: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> None:
|
||||
"""Record CP-local adapter inputs, hidden states, and rotary embeddings."""
|
||||
tokens, coords, video_mask, audio_mask, text_mask = inputs
|
||||
hidden_states, rope = output
|
||||
capture["pre_adapter"] = {
|
||||
"tokens": _tensor_record(tokens),
|
||||
"coords": _tensor_record(coords),
|
||||
"video_mask": _tensor_record(video_mask),
|
||||
"audio_mask": _tensor_record(audio_mask),
|
||||
"text_mask": _tensor_record(text_mask),
|
||||
"hidden_states": _tensor_record(hidden_states),
|
||||
"rope": _tensor_record(rope),
|
||||
}
|
||||
|
||||
def make_layer_input_hook(layer_index: int):
|
||||
"""Create a hook that records one layer's input hidden states."""
|
||||
|
||||
def capture_layer_input(
|
||||
_module: torch.nn.Module,
|
||||
inputs: tuple[Any, ...],
|
||||
) -> None:
|
||||
"""Record the hidden states entering one transformer layer."""
|
||||
capture["layer_boundaries"][layer_index]["input"] = _tensor_record(
|
||||
inputs[0]
|
||||
)
|
||||
|
||||
return capture_layer_input
|
||||
|
||||
def make_layer_output_hook(layer_index: int):
|
||||
"""Create a hook that records one layer's output hidden states."""
|
||||
|
||||
def capture_layer_output(
|
||||
_module: torch.nn.Module,
|
||||
_inputs: tuple[Any, ...],
|
||||
output: torch.Tensor,
|
||||
) -> None:
|
||||
"""Record the hidden states leaving one transformer layer."""
|
||||
capture["layer_boundaries"][layer_index]["output"] = _tensor_record(output)
|
||||
|
||||
return capture_layer_output
|
||||
|
||||
def capture_post_adapter(
|
||||
_module: torch.nn.Module,
|
||||
inputs: tuple[Any, ...],
|
||||
output: torch.Tensor,
|
||||
) -> None:
|
||||
"""Record the CP-local final projection inputs, masks, and output."""
|
||||
hidden_states, video_mask, audio_mask = inputs
|
||||
capture["post_adapter"] = {
|
||||
"hidden_states": _tensor_record(hidden_states),
|
||||
"video_mask": _tensor_record(video_mask),
|
||||
"audio_mask": _tensor_record(audio_mask),
|
||||
"output": _tensor_record(output),
|
||||
}
|
||||
|
||||
handles.append(model.pre_adapter.register_forward_hook(capture_pre_adapter))
|
||||
for layer_index, layer in enumerate(model.block.layers):
|
||||
handles.append(
|
||||
layer.register_forward_pre_hook(make_layer_input_hook(layer_index))
|
||||
)
|
||||
handles.append(layer.register_forward_hook(make_layer_output_hook(layer_index)))
|
||||
handles.append(model.post_adapter.register_forward_hook(capture_post_adapter))
|
||||
return capture, handles
|
||||
|
||||
|
||||
def _run_case(
|
||||
model: torch.nn.Module,
|
||||
proxy: Any,
|
||||
model_input_type: type,
|
||||
device: torch.device,
|
||||
) -> dict[str, Any]:
|
||||
"""Run one refiner forward pass and capture all parity boundaries."""
|
||||
model_input = _build_model_input(model_input_type, device)
|
||||
packed_input = proxy.process_input(model_input)
|
||||
capture, handles = _capture_model_boundaries(model)
|
||||
capture["packed_input"] = _packed_input_record(packed_input)
|
||||
try:
|
||||
with torch.inference_mode():
|
||||
model_output = model(*packed_input)
|
||||
finally:
|
||||
for handle in handles:
|
||||
handle.remove()
|
||||
video_output, audio_output = proxy.process_output(model_output)
|
||||
capture["model_output"] = _tensor_record(model_output)
|
||||
capture["depacked_output"] = {
|
||||
"video": _tensor_record(video_output),
|
||||
"audio": _tensor_record(audio_output),
|
||||
}
|
||||
if set(capture["pre_adapter"]) != {
|
||||
"tokens",
|
||||
"coords",
|
||||
"video_mask",
|
||||
"audio_mask",
|
||||
"text_mask",
|
||||
"hidden_states",
|
||||
"rope",
|
||||
}:
|
||||
raise RuntimeError("The refiner pre-adapter hook did not execute")
|
||||
if any(
|
||||
set(layer_capture) != {"layer_index", "input", "output"}
|
||||
for layer_capture in capture["layer_boundaries"]
|
||||
):
|
||||
raise RuntimeError("A refiner transformer layer hook did not execute")
|
||||
if set(capture["post_adapter"]) != {
|
||||
"hidden_states",
|
||||
"video_mask",
|
||||
"audio_mask",
|
||||
"output",
|
||||
}:
|
||||
raise RuntimeError("The refiner post-adapter hook did not execute")
|
||||
if set(capture) != {
|
||||
"packed_input",
|
||||
"pre_adapter",
|
||||
"layer_boundaries",
|
||||
"post_adapter",
|
||||
"model_output",
|
||||
"depacked_output",
|
||||
}:
|
||||
raise RuntimeError("The refiner parity capture schema is incomplete")
|
||||
return capture
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Load one implementation, run the refiner, and save rank-local captures."""
|
||||
args = _parse_args()
|
||||
os.environ.setdefault("MAGI_COMPILE_COMPILE_MODE", "NONE")
|
||||
os.environ.pop("SKIP_LOAD_MODEL", None)
|
||||
os.environ["MAGI2_CKPT_ROOT"] = str(WEIGHTS_ROOT)
|
||||
_enable_determinism()
|
||||
rank, device = _initialize_distributed(args.implementation)
|
||||
model, proxy, model_input_type = _load_runtime(args.implementation, device)
|
||||
artifact = {
|
||||
"schema_version": 1,
|
||||
"implementation": args.implementation,
|
||||
"rank": rank,
|
||||
"world_size": dist.get_world_size(),
|
||||
"case": _run_case(model, proxy, model_input_type, device),
|
||||
}
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
artifact_path = args.output_dir / f"rank_{rank}.pt"
|
||||
temporary_path = args.output_dir / f"rank_{rank}.pt.tmp"
|
||||
torch.save(artifact, temporary_path)
|
||||
os.replace(temporary_path, artifact_path)
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict Stable Audio VAE decode parity for MAGI-2 Preview.
|
||||
|
||||
Coverage scope: implementation_subcomponent. The official side instantiates
|
||||
``inference.pipeline.audio_decoder.SAAudioFeatureExtractor``, which builds the
|
||||
official ``AudioAutoencoder`` from ``inference.model.sa_audio_vae``. The
|
||||
FastVideo side uses ``fastvideo.models.vaes.magi2_audio_vae.Magi2AudioVAE``,
|
||||
which wraps ``fastvideo.models.vaes.oobleck.OobleckDecoder``. The test also
|
||||
checks MAGI-2's latent transpose, sample-major stereo layout, and 441/512
|
||||
resampling contract.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.models.vaes.magi2_audio_vae import load_magi2_audio_vae
|
||||
from fastvideo.pipelines.basic.magi2.stages.audio_decoding import (
|
||||
decode_magi2_audio,
|
||||
)
|
||||
from tests.local_tests.magi2._parity_utils import (
|
||||
LOCAL_WEIGHTS_DIR,
|
||||
assert_array_exact,
|
||||
assert_tensor_exact,
|
||||
import_official_module,
|
||||
require_path,
|
||||
)
|
||||
|
||||
|
||||
PARITY_COVERAGE = "both"
|
||||
STABLE_AUDIO_DIR = LOCAL_WEIGHTS_DIR / "stable-audio-open-1.0"
|
||||
STABLE_AUDIO_CONFIG_PATH = STABLE_AUDIO_DIR / "model_config.json"
|
||||
STABLE_AUDIO_CHECKPOINT_PATH = STABLE_AUDIO_DIR / "model.safetensors"
|
||||
|
||||
|
||||
def _deterministic_audio_latent(device: torch.device) -> torch.Tensor:
|
||||
"""Create a short sample-major latent with MAGI-2's 64 audio channels."""
|
||||
latent_values = torch.arange(8 * 64, dtype=torch.float32)
|
||||
latent_values = ((latent_values % 127) - 63) / 64
|
||||
return latent_values.reshape(8, 64).to(device=device)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="MAGI-2 Stable Audio VAE parity requires CUDA.",
|
||||
)
|
||||
def test_magi2_stable_audio_decode_and_resample_exact_parity() -> None:
|
||||
"""Require exact decoded waveform and sample-major resampled audio."""
|
||||
require_path(STABLE_AUDIO_CONFIG_PATH, "Stable Audio configuration")
|
||||
require_path(STABLE_AUDIO_CHECKPOINT_PATH, "Stable Audio checkpoint")
|
||||
official_audio_module = import_official_module(
|
||||
"inference.pipeline.audio_decoder"
|
||||
)
|
||||
device = torch.device("cuda:0")
|
||||
torch.cuda.set_device(device)
|
||||
official_audio_vae = official_audio_module.SAAudioFeatureExtractor(
|
||||
str(STABLE_AUDIO_DIR)
|
||||
)
|
||||
fastvideo_audio_vae = load_magi2_audio_vae(STABLE_AUDIO_DIR, device)
|
||||
assert official_audio_vae.sample_rate == fastvideo_audio_vae.sampling_rate == 44100
|
||||
assert official_audio_vae.downsampling_ratio == fastvideo_audio_vae.hop_length == 2048
|
||||
|
||||
sample_major_latent = _deterministic_audio_latent(device)
|
||||
official_channel_major_latent = sample_major_latent.T
|
||||
fastvideo_channel_major_latent = sample_major_latent.unsqueeze(0).permute(0, 2, 1).contiguous()
|
||||
with torch.inference_mode():
|
||||
official_waveform = official_audio_vae.decode(official_channel_major_latent).detach().cpu()
|
||||
fastvideo_waveform = fastvideo_audio_vae.decode(
|
||||
fastvideo_channel_major_latent
|
||||
).detach().cpu()
|
||||
|
||||
assert_tensor_exact(fastvideo_waveform, official_waveform, "decoded audio waveform")
|
||||
official_sample_major = official_waveform.squeeze(0).T.numpy()
|
||||
fastvideo_sample_major = fastvideo_waveform.squeeze(0).T.numpy()
|
||||
assert_array_exact(fastvideo_sample_major, official_sample_major, "sample-major stereo audio")
|
||||
|
||||
official_resampled = official_audio_module.resample_audio_sinc(
|
||||
official_sample_major,
|
||||
441 / 512,
|
||||
)
|
||||
fastvideo_resampled = decode_magi2_audio(
|
||||
fastvideo_audio_vae,
|
||||
sample_major_latent,
|
||||
)
|
||||
assert_array_exact(fastvideo_resampled, official_resampled, "resampled stereo audio")
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Checkpoint-layout conversion coverage for the MAGI-2 port."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.utils import verify_model_config_and_directory
|
||||
from scripts.checkpoint_conversion.convert_magi2_to_fastvideo import (
|
||||
COMPONENT_DIRECTORY_MAPPING,
|
||||
MODEL_INDEX,
|
||||
REQUIRED_COMPONENT_FILES,
|
||||
SOURCE_REVISION,
|
||||
TURBO_VAE_SKIPPED_KEYS,
|
||||
convert_checkpoint_layout,
|
||||
)
|
||||
|
||||
|
||||
def _create_minimal_source(source: Path) -> dict[str, Path]:
|
||||
"""Create identifying files for every published checkpoint component."""
|
||||
source_files: dict[str, Path] = {}
|
||||
for component, relative_paths in REQUIRED_COMPONENT_FILES.items():
|
||||
for relative_path in relative_paths:
|
||||
source_path = source / component / relative_path
|
||||
source_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if relative_path.endswith(".index.json"):
|
||||
shard_name = f"{component}-00001-of-00001.safetensors"
|
||||
source_path.write_text(
|
||||
json.dumps({"weight_map": {"weight": shard_name}}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
shard_path = source_path.parent / shard_name
|
||||
shard_path.write_bytes(shard_name.encode())
|
||||
source_files[f"{component}/{shard_name}"] = shard_path
|
||||
else:
|
||||
source_path.write_bytes(f"{component}/{relative_path}".encode())
|
||||
source_files[f"{component}/{relative_path}"] = source_path
|
||||
return source_files
|
||||
|
||||
|
||||
def test_magi2_converter_preserves_files_and_records_all_mappings(tmp_path: Path) -> None:
|
||||
"""Require hard-linked components and complete conversion metadata."""
|
||||
source = tmp_path / "official"
|
||||
output = tmp_path / "fastvideo"
|
||||
source_files = _create_minimal_source(source)
|
||||
|
||||
convert_checkpoint_layout(source, output)
|
||||
|
||||
assert verify_model_config_and_directory(str(output)) == MODEL_INDEX
|
||||
for source_relative_path, source_path in source_files.items():
|
||||
source_component, component_relative_path = source_relative_path.split("/", 1)
|
||||
destination_component = COMPONENT_DIRECTORY_MAPPING[source_component]
|
||||
destination_path = output / destination_component / component_relative_path
|
||||
assert destination_path.read_bytes() == source_path.read_bytes()
|
||||
assert destination_path.stat().st_ino == source_path.stat().st_ino
|
||||
|
||||
manifest = json.loads((output / "magi2_conversion_manifest.json").read_text())
|
||||
assert manifest["source"]["revision"] == SOURCE_REVISION
|
||||
assert manifest["component_directory_mapping"] == COMPONENT_DIRECTORY_MAPPING
|
||||
assert manifest["skipped_checkpoint_keys"]["vae/checkpoint.ckpt"] == list(
|
||||
TURBO_VAE_SKIPPED_KEYS
|
||||
)
|
||||
assert len(TURBO_VAE_SKIPPED_KEYS) == 8
|
||||
|
||||
|
||||
def test_magi2_converter_rejects_an_incomplete_snapshot(tmp_path: Path) -> None:
|
||||
"""Reject a source tree before creating a partial output repository."""
|
||||
source = tmp_path / "official"
|
||||
source.mkdir()
|
||||
output = tmp_path / "fastvideo"
|
||||
|
||||
try:
|
||||
convert_checkpoint_layout(source, output)
|
||||
except FileNotFoundError as error:
|
||||
assert "source snapshot is incomplete" in str(error)
|
||||
else:
|
||||
raise AssertionError("An incomplete MAGI-2 snapshot was accepted")
|
||||
assert not output.exists()
|
||||
@@ -0,0 +1,140 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict production-loader parity for the MAGI-2 Wan 2.2 image encoder.
|
||||
|
||||
Coverage scope: both. The test loads the published ``Wan2_2_VAE`` reference
|
||||
and the FastVideo ``Magi2WanImageEncoder`` through its production loader. A
|
||||
synthetic image passes through the I2V letterbox, ``VideoProcessor``, and BF16
|
||||
round trip before FP32 encoding.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.vaes.magi2_wanvae import Magi2WanVAEConfig
|
||||
from fastvideo.models.vaes.magi2_wan_loader import (
|
||||
Magi2WanImageEncoder,
|
||||
load_magi2_wan_image_encoder,
|
||||
)
|
||||
from tests.local_tests.magi2._parity_utils import (
|
||||
LOCAL_WEIGHTS_DIR,
|
||||
assert_tensor_exact,
|
||||
import_official_module,
|
||||
require_path,
|
||||
)
|
||||
|
||||
|
||||
PARITY_COVERAGE = "both"
|
||||
WAN_VAE_PATH = LOCAL_WEIGHTS_DIR / "vae" / "Wan2.2_VAE.pth"
|
||||
|
||||
|
||||
def _resizepad(image, target_height: int, target_width: int):
|
||||
"""Fit an RGB image inside a white canvas without cropping its edges."""
|
||||
image_module = pytest.importorskip("PIL.Image")
|
||||
width, height = image.size
|
||||
if width <= 0 or height <= 0:
|
||||
raise ValueError(f"Invalid image size, width: {width}, height: {height}")
|
||||
scale = min(target_width / width, target_height / height)
|
||||
resized_width = max(1, int(round(width * scale)))
|
||||
resized_height = max(1, int(round(height * scale)))
|
||||
resized_image = image.convert("RGB").resize(
|
||||
(resized_width, resized_height),
|
||||
resample=image_module.Resampling.LANCZOS,
|
||||
)
|
||||
canvas = image_module.new("RGB", (target_width, target_height), (255, 255, 255))
|
||||
canvas.paste(
|
||||
resized_image,
|
||||
((target_width - resized_width) // 2, (target_height - resized_height) // 2),
|
||||
)
|
||||
return canvas
|
||||
|
||||
|
||||
def _prepare_reference_image(device: torch.device) -> torch.Tensor:
|
||||
"""Apply the I2V resize, letterbox, and BF16 input conversion.
|
||||
|
||||
The source dimensions force the long-edge scaling branch and a non-square
|
||||
letterbox. The returned tensor has the ``[B, C, T, H, W]`` layout consumed
|
||||
by ``Wan2_2_VAE.encode``.
|
||||
"""
|
||||
pil_image_module = pytest.importorskip("PIL.Image")
|
||||
video_processor_module = pytest.importorskip("diffusers.video_processor")
|
||||
pixels = (
|
||||
np.arange(48 * 80 * 3, dtype=np.uint32).reshape(48, 80, 3) % 251
|
||||
).astype(np.uint8)
|
||||
source_image = pil_image_module.fromarray(pixels, mode="RGB")
|
||||
generation_height, generation_width = 64, 96
|
||||
max_length = max(generation_height, generation_width)
|
||||
target_width = max_length
|
||||
target_height = int(source_image.height * max_length / source_image.width)
|
||||
resized_image = _resizepad(
|
||||
source_image,
|
||||
target_height,
|
||||
target_width,
|
||||
)
|
||||
video_processor = video_processor_module.VideoProcessor(vae_scale_factor=32)
|
||||
image_tensor = video_processor.preprocess(
|
||||
resized_image,
|
||||
height=target_height,
|
||||
width=target_width,
|
||||
)
|
||||
return image_tensor.to(device=device, dtype=torch.bfloat16).unsqueeze(2)[:, :3].float()
|
||||
|
||||
|
||||
def _assert_magi2_wan_config(
|
||||
official_vae,
|
||||
fastvideo_encoder: Magi2WanImageEncoder,
|
||||
) -> None:
|
||||
"""Validate the production encoder configuration against published values."""
|
||||
config = Magi2WanVAEConfig()
|
||||
assert config.base_dim == 160
|
||||
assert config.decoder_base_dim == 256
|
||||
assert config.z_dim == 48
|
||||
assert config.in_channels == 12
|
||||
assert config.out_channels == 12
|
||||
assert config.temperal_downsample == (False, True, True)
|
||||
assert config.patch_size == 2
|
||||
assert config.is_residual is True
|
||||
assert config.clip_output is False
|
||||
assert torch.equal(
|
||||
torch.tensor(config.latents_mean, dtype=torch.float32),
|
||||
official_vae.mean.detach().cpu(),
|
||||
)
|
||||
assert torch.equal(
|
||||
torch.tensor(config.latents_std, dtype=torch.float32),
|
||||
official_vae.std.detach().cpu(),
|
||||
)
|
||||
assert torch.equal(fastvideo_encoder.mean.detach().cpu(), official_vae.mean.detach().cpu())
|
||||
assert torch.equal(fastvideo_encoder.std.detach().cpu(), official_vae.std.detach().cpu())
|
||||
assert torch.equal(
|
||||
fastvideo_encoder.inverse_std.detach().cpu(),
|
||||
official_vae.scale[1].detach().cpu(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="MAGI-2 I2V image-encoder parity requires CUDA.",
|
||||
)
|
||||
def test_load_magi2_wan_image_encoder_i2v_exact_parity() -> None:
|
||||
"""Require exact post-normalization image latents from both encoders."""
|
||||
vae_path = require_path(WAN_VAE_PATH, "Wan 2.2 VAE weights")
|
||||
official_vae_module = import_official_module("inference.model.vae2_2")
|
||||
device = torch.device("cuda:0")
|
||||
image_tensor = _prepare_reference_image(device)
|
||||
|
||||
official_vae = official_vae_module.get_vae2_2(
|
||||
str(vae_path),
|
||||
device=str(device),
|
||||
weight_dtype=torch.float32,
|
||||
)
|
||||
fastvideo_encoder = load_magi2_wan_image_encoder(vae_path, device)
|
||||
_assert_magi2_wan_config(official_vae, fastvideo_encoder)
|
||||
|
||||
with torch.inference_mode():
|
||||
official_latent = official_vae.encode(image_tensor).detach().cpu()
|
||||
fastvideo_latent = fastvideo_encoder.encode(image_tensor).detach().cpu()
|
||||
|
||||
assert image_tensor.dtype == torch.float32
|
||||
assert image_tensor.shape[:3] == (1, 3, 1)
|
||||
assert_tensor_exact(fastvideo_latent, official_latent, "I2V Wan latent")
|
||||
@@ -0,0 +1,63 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo process-group binding coverage for MAGI-2."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.models.dits.magi2_runtime import fastvideo_parallel
|
||||
from fastvideo.models.dits.magi2_runtime import psm as psm_manager
|
||||
|
||||
|
||||
def test_fastvideo_groups_are_reused_for_context_and_expert_parallelism(monkeypatch) -> None:
|
||||
"""Bind one eight-rank sequence group to both MAGI-2 parallel axes."""
|
||||
sequence_process_group = object()
|
||||
data_process_group = object()
|
||||
sequence_group = SimpleNamespace(
|
||||
world_size=8,
|
||||
device_group=sequence_process_group,
|
||||
ranks=list(range(8)),
|
||||
)
|
||||
data_group = SimpleNamespace(
|
||||
world_size=1,
|
||||
device_group=data_process_group,
|
||||
ranks=[3],
|
||||
)
|
||||
monkeypatch.setattr(fastvideo_parallel, "get_sp_group", lambda: sequence_group)
|
||||
monkeypatch.setattr(fastvideo_parallel, "get_dp_group", lambda: data_group)
|
||||
monkeypatch.setattr(
|
||||
psm_manager,
|
||||
"get_world_size",
|
||||
lambda dim="": 8 if dim in {"cp", "ep"} else 1,
|
||||
)
|
||||
|
||||
recorded: dict = {}
|
||||
monkeypatch.setattr(
|
||||
fastvideo_parallel,
|
||||
"bind_process_groups",
|
||||
lambda **kwargs: recorded.update(kwargs),
|
||||
)
|
||||
fastvideo_parallel.bind_fastvideo_parallel_state()
|
||||
|
||||
assert recorded["cp_group"] is sequence_process_group
|
||||
assert recorded["ep_group"] is sequence_process_group
|
||||
assert recorded["dp_group"] is data_process_group
|
||||
assert recorded["cp_ranks"] == list(range(8))
|
||||
|
||||
|
||||
def test_fastvideo_binding_requires_the_published_parallel_size(monkeypatch) -> None:
|
||||
"""Reject sequence groups that cannot match the published eight-rank layout."""
|
||||
monkeypatch.setattr(
|
||||
fastvideo_parallel,
|
||||
"get_sp_group",
|
||||
lambda: SimpleNamespace(world_size=4),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
fastvideo_parallel,
|
||||
"get_dp_group",
|
||||
lambda: SimpleNamespace(world_size=2),
|
||||
)
|
||||
with pytest.raises(ValueError, match="sp_size=8"):
|
||||
fastvideo_parallel.bind_fastvideo_parallel_state()
|
||||
@@ -0,0 +1,196 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict stage and end-to-end parity for the MAGI-2 release pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
CONVERTED_ROOT = Path(
|
||||
os.environ.get("MAGI2_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "magi2")
|
||||
)
|
||||
WORKER_PATH = Path(__file__).with_name("_pipeline_parity_worker.py")
|
||||
CAPTURE_ROOT = REPO_ROOT / "archived" / "magi2_parity" / "validation" / "pipeline"
|
||||
OFFICIAL_REVISION = "073c84f2102ec3c9287623113a103c14402770ad"
|
||||
WORLD_SIZE = 8
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
torch.cuda.device_count() < WORLD_SIZE,
|
||||
reason=f"MAGI-2 pipeline numerical parity requires {WORLD_SIZE} CUDA devices",
|
||||
)
|
||||
|
||||
|
||||
def _require_parity_sources() -> None:
|
||||
"""Require clean pinned source trees and complete converted components."""
|
||||
revision = subprocess.run(
|
||||
["git", "rev-parse", "HEAD"],
|
||||
cwd=OFFICIAL_ROOT,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
).stdout.strip()
|
||||
assert revision == OFFICIAL_REVISION
|
||||
official_changes = subprocess.run(
|
||||
["git", "status", "--porcelain=v1"],
|
||||
cwd=OFFICIAL_ROOT,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
).stdout
|
||||
assert not official_changes, (
|
||||
f"Official MAGI-2 checkout must be clean for parity:\n{official_changes}"
|
||||
)
|
||||
required_paths = (
|
||||
WEIGHTS_ROOT / "preview" / "model.safetensors.index.json",
|
||||
WEIGHTS_ROOT / "refiner" / "model.safetensors.index.json",
|
||||
WEIGHTS_ROOT / "turbo_vae" / "checkpoint.ckpt",
|
||||
CONVERTED_ROOT / "model_index.json",
|
||||
CONVERTED_ROOT / "transformer" / "model.safetensors.index.json",
|
||||
CONVERTED_ROOT / "transformer_2" / "model.safetensors.index.json",
|
||||
)
|
||||
for required_path in required_paths:
|
||||
assert required_path.is_file(), f"MAGI-2 parity input is missing: {required_path}"
|
||||
|
||||
|
||||
def _run_implementation(implementation: str) -> Path:
|
||||
"""Launch one full release-profile implementation in an isolated torchrun job."""
|
||||
output_dir = CAPTURE_ROOT / implementation
|
||||
environment = os.environ.copy()
|
||||
environment.update(
|
||||
{
|
||||
"CUBLAS_WORKSPACE_CONFIG": ":4096:8",
|
||||
"MAGI2_CKPT_ROOT": str(WEIGHTS_ROOT),
|
||||
"MAGI2_CONVERTED_WEIGHTS_DIR": str(CONVERTED_ROOT),
|
||||
"MAGI2_DETERMINISTIC": "1",
|
||||
"MAGI2_LOCAL_WEIGHTS_DIR": str(WEIGHTS_ROOT),
|
||||
"MAGI2_OFFICIAL_REF_DIR": str(OFFICIAL_ROOT),
|
||||
"MAGI2_TEXT_ENC_OFFLOAD_MODE": "roundtrip",
|
||||
"MAGI2_VAE_OFFLOAD_MODE": "roundtrip",
|
||||
"MAGI_ATTENTION_DETERMINISTIC_MODE": "1",
|
||||
"MAGI_COMPILE_COMPILE_MODE": "NONE",
|
||||
"OMP_NUM_THREADS": "1",
|
||||
"PYTHONHASHSEED": "42",
|
||||
}
|
||||
)
|
||||
environment.pop("MAGI2_SAVE_LATENT_PATH", None)
|
||||
environment.pop("NEGATIVE_PROMPT", None)
|
||||
environment.pop("SKIP_LOAD_MODEL", None)
|
||||
completed = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--standalone",
|
||||
"--nnodes=1",
|
||||
f"--nproc-per-node={WORLD_SIZE}",
|
||||
str(WORKER_PATH),
|
||||
"--implementation",
|
||||
implementation,
|
||||
"--output-dir",
|
||||
str(output_dir),
|
||||
],
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=21600,
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
raise AssertionError(
|
||||
f"{implementation} MAGI-2 pipeline torchrun failed.\n"
|
||||
f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}"
|
||||
)
|
||||
return output_dir
|
||||
|
||||
|
||||
def _load_capture(output_dir: Path) -> dict[str, Any]:
|
||||
"""Load and validate one stage-digest manifest."""
|
||||
artifact_path = output_dir / "capture.json"
|
||||
assert artifact_path.is_file(), f"MAGI-2 capture is missing: {artifact_path}"
|
||||
artifact = json.loads(artifact_path.read_text(encoding="utf-8"))
|
||||
assert artifact["schema_version"] == 1
|
||||
assert artifact["world_size"] == WORLD_SIZE
|
||||
assert artifact["preview_steps"] == 100
|
||||
assert artifact["refiner_steps"] == 5
|
||||
assert artifact["cases"].keys() == {"t2v", "i2v"}
|
||||
return artifact
|
||||
|
||||
|
||||
def _first_difference(
|
||||
fastvideo_value: Any,
|
||||
official_value: Any,
|
||||
path: str,
|
||||
) -> str | None:
|
||||
"""Return the first structural or numerical difference in two manifests."""
|
||||
if type(fastvideo_value) is not type(official_value):
|
||||
return (
|
||||
f"{path}: type differs: {type(fastvideo_value).__name__} versus "
|
||||
f"{type(official_value).__name__}"
|
||||
)
|
||||
if isinstance(official_value, dict):
|
||||
if fastvideo_value.keys() != official_value.keys():
|
||||
return (
|
||||
f"{path}: keys differ: {sorted(fastvideo_value)} versus "
|
||||
f"{sorted(official_value)}"
|
||||
)
|
||||
for key in official_value:
|
||||
difference = _first_difference(
|
||||
fastvideo_value[key],
|
||||
official_value[key],
|
||||
f"{path}.{key}",
|
||||
)
|
||||
if difference is not None:
|
||||
return difference
|
||||
return None
|
||||
if isinstance(official_value, list):
|
||||
if len(fastvideo_value) != len(official_value):
|
||||
return (
|
||||
f"{path}: list length differs: {len(fastvideo_value)} versus "
|
||||
f"{len(official_value)}"
|
||||
)
|
||||
for index, (fastvideo_item, official_item) in enumerate(
|
||||
zip(fastvideo_value, official_value, strict=True)
|
||||
):
|
||||
difference = _first_difference(
|
||||
fastvideo_item,
|
||||
official_item,
|
||||
f"{path}[{index}]",
|
||||
)
|
||||
if difference is not None:
|
||||
return difference
|
||||
return None
|
||||
if fastvideo_value != official_value:
|
||||
return f"{path}: {fastvideo_value!r} differs from {official_value!r}"
|
||||
return None
|
||||
|
||||
|
||||
def test_magi2_pipeline_matches_official_release_exactly() -> None:
|
||||
"""Match T2V and I2V stage tensors, decoded video, and decoded audio exactly."""
|
||||
_require_parity_sources()
|
||||
official_capture = _load_capture(_run_implementation("official"))
|
||||
fastvideo_capture = _load_capture(_run_implementation("fastvideo"))
|
||||
assert official_capture["implementation"] == "official"
|
||||
assert fastvideo_capture["implementation"] == "fastvideo"
|
||||
difference = _first_difference(
|
||||
fastvideo_capture["cases"],
|
||||
official_capture["cases"],
|
||||
"cases",
|
||||
)
|
||||
assert difference is None, difference
|
||||
@@ -0,0 +1,531 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict MAGI-2 preview data-proxy implementation parity.
|
||||
|
||||
Coverage scope: implementation_subcomponent. The tests load the pinned official
|
||||
``Magi2DataProxy`` source and compare it with the FastVideo-owned component on
|
||||
deterministic text-to-video (T2V) and image-to-video (I2V) tensors.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from dataclasses import asdict, dataclass
|
||||
from enum import IntEnum
|
||||
from functools import lru_cache
|
||||
import importlib.util
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
import typing
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
import fastvideo.pipelines.basic.magi2.stages.preview_data_proxy as fastvideo_proxy_module
|
||||
from fastvideo.pipelines.basic.magi2.stages.preview_data_proxy import (
|
||||
Magi2DataProxy,
|
||||
Magi2DataProxyConfig,
|
||||
ModelInput,
|
||||
Modality,
|
||||
VarlenHandler,
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REPO_ROOT = Path(
|
||||
os.environ.get(
|
||||
"MAGI2_OFFICIAL_REF_DIR",
|
||||
REPO_ROOT.parent / "MAGI-2-preview",
|
||||
)
|
||||
)
|
||||
OFFICIAL_PROXY_PATH = (
|
||||
OFFICIAL_REPO_ROOT / "inference" / "pipeline" / "preview_data_proxy.py"
|
||||
)
|
||||
OFFICIAL_MODEL_PATH = (
|
||||
OFFICIAL_REPO_ROOT / "inference" / "model" / "magi2_preview.py"
|
||||
)
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ParallelStateStub:
|
||||
"""Expose fixed parallel sizes and identifiable process-group handles."""
|
||||
|
||||
cp_size: int = 1
|
||||
ep_size: int = 1
|
||||
|
||||
def get_world_size(self, dimension: str = "") -> int:
|
||||
"""Return the configured size for one parallel dimension."""
|
||||
if dimension == "cp":
|
||||
return self.cp_size
|
||||
if dimension == "ep":
|
||||
return self.ep_size
|
||||
return max(self.cp_size, self.ep_size)
|
||||
|
||||
def get_parallel_group(self, dimension: str) -> str:
|
||||
"""Return a stable process-group marker for distributed call assertions."""
|
||||
return f"{dimension}-group"
|
||||
|
||||
|
||||
def _load_official_model_utility_module() -> types.ModuleType:
|
||||
"""Load official coordinate, modality, and time-embedding definitions."""
|
||||
if not OFFICIAL_MODEL_PATH.is_file():
|
||||
raise FileNotFoundError(
|
||||
f"Official MAGI-2 model source is missing: {OFFICIAL_MODEL_PATH}"
|
||||
)
|
||||
source = OFFICIAL_MODEL_PATH.read_text(encoding="utf-8")
|
||||
source_module = ast.parse(source, filename=str(OFFICIAL_MODEL_PATH))
|
||||
symbol_names = {
|
||||
"Modality",
|
||||
"VarlenHandler",
|
||||
"get_coords",
|
||||
"sinusoidal_embedding_1d",
|
||||
}
|
||||
symbol_nodes = [
|
||||
node
|
||||
for node in source_module.body
|
||||
if isinstance(node, (ast.ClassDef, ast.FunctionDef))
|
||||
and node.name in symbol_names
|
||||
]
|
||||
found_names = {node.name for node in symbol_nodes}
|
||||
if found_names != symbol_names:
|
||||
raise AssertionError(
|
||||
f"Official model utility definitions differ: expected {symbol_names}, "
|
||||
f"found {found_names}"
|
||||
)
|
||||
|
||||
module = types.ModuleType("inference.model.magi2_preview")
|
||||
module.__file__ = str(OFFICIAL_MODEL_PATH)
|
||||
module.__dict__.update(
|
||||
{
|
||||
"Coords": tuple[int, int, int],
|
||||
"IntEnum": IntEnum,
|
||||
"Optional": typing.Optional,
|
||||
"dataclass": dataclass,
|
||||
"math": math,
|
||||
"torch": torch,
|
||||
}
|
||||
)
|
||||
utility_module = ast.Module(body=symbol_nodes, type_ignores=[])
|
||||
module_name = module.__name__
|
||||
previous_module = sys.modules.get(module_name)
|
||||
sys.modules[module_name] = module
|
||||
try:
|
||||
exec(
|
||||
compile(utility_module, str(OFFICIAL_MODEL_PATH), "exec"),
|
||||
module.__dict__,
|
||||
)
|
||||
finally:
|
||||
if previous_module is None:
|
||||
sys.modules.pop(module_name, None)
|
||||
else:
|
||||
sys.modules[module_name] = previous_module
|
||||
return module
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _load_official_proxy_module() -> types.ModuleType:
|
||||
"""Import the official proxy while replacing unrelated model dependencies."""
|
||||
if not OFFICIAL_PROXY_PATH.is_file():
|
||||
raise FileNotFoundError(
|
||||
f"Official MAGI-2 data-proxy source is missing: {OFFICIAL_PROXY_PATH}"
|
||||
)
|
||||
if importlib.util.find_spec("unfoldNd") is None:
|
||||
pytest.skip(
|
||||
"Official MAGI-2 data-proxy parity requires unfoldNd==0.2.3."
|
||||
)
|
||||
|
||||
model_stub = _load_official_model_utility_module()
|
||||
distributed_stub = types.ModuleType("inference.infra.distributed")
|
||||
distributed_stub.psm = _ParallelStateStub()
|
||||
replacement_modules = {
|
||||
"inference.model.magi2_preview": model_stub,
|
||||
"inference.infra.distributed": distributed_stub,
|
||||
}
|
||||
missing_module = object()
|
||||
previous_modules = {
|
||||
name: sys.modules.get(name, missing_module)
|
||||
for name in replacement_modules
|
||||
}
|
||||
sys.modules.update(replacement_modules)
|
||||
|
||||
official_path_entry = str(OFFICIAL_REPO_ROOT)
|
||||
added_path = official_path_entry not in sys.path
|
||||
if added_path:
|
||||
sys.path.insert(0, official_path_entry)
|
||||
module_name = "magi2_official_preview_data_proxy"
|
||||
try:
|
||||
spec = importlib.util.spec_from_file_location(module_name, OFFICIAL_PROXY_PATH)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(
|
||||
f"Cannot create an import spec for {OFFICIAL_PROXY_PATH}"
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[module_name] = module
|
||||
spec.loader.exec_module(module)
|
||||
finally:
|
||||
if added_path:
|
||||
sys.path.remove(official_path_entry)
|
||||
for name, previous_module in previous_modules.items():
|
||||
if previous_module is missing_module:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = previous_module
|
||||
return module
|
||||
|
||||
|
||||
def _build_proxy_pair(
|
||||
official_module: types.ModuleType,
|
||||
config: Magi2DataProxyConfig,
|
||||
) -> tuple[Any, Magi2DataProxy]:
|
||||
"""Construct official and FastVideo proxies from identical config values."""
|
||||
official_config = official_module.DataProxyConfig(**asdict(config))
|
||||
return (
|
||||
official_module.Magi2DataProxy(official_config),
|
||||
Magi2DataProxy(config),
|
||||
)
|
||||
|
||||
|
||||
def _build_model_input_pair(
|
||||
official_module: types.ModuleType,
|
||||
input_tensors: dict[str, Any],
|
||||
) -> tuple[Any, ModelInput]:
|
||||
"""Construct both input dataclasses over the same immutable test tensors."""
|
||||
return (
|
||||
official_module.ModelInput(**input_tensors),
|
||||
ModelInput(**input_tensors),
|
||||
)
|
||||
|
||||
|
||||
def _assert_tensor_exact(
|
||||
fastvideo_tensor: torch.Tensor,
|
||||
official_tensor: torch.Tensor,
|
||||
) -> None:
|
||||
"""Require identical tensor metadata and values."""
|
||||
assert fastvideo_tensor.shape == official_tensor.shape
|
||||
assert fastvideo_tensor.dtype == official_tensor.dtype
|
||||
assert fastvideo_tensor.stride() == official_tensor.stride()
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=0, rtol=0)
|
||||
|
||||
|
||||
def _assert_packed_inputs_exact(
|
||||
fastvideo_output: tuple[torch.Tensor, torch.Tensor, torch.Tensor, VarlenHandler, torch.Tensor],
|
||||
official_output: tuple[torch.Tensor, torch.Tensor, torch.Tensor, Any, torch.Tensor],
|
||||
) -> None:
|
||||
"""Compare all packed tensors and variable-length attention boundaries."""
|
||||
fastvideo_tokens, fastvideo_coords, fastvideo_modalities, fastvideo_varlen, fastvideo_time = fastvideo_output
|
||||
official_tokens, official_coords, official_modalities, official_varlen, official_time = official_output
|
||||
_assert_tensor_exact(fastvideo_tokens, official_tokens)
|
||||
_assert_tensor_exact(fastvideo_coords, official_coords)
|
||||
_assert_tensor_exact(fastvideo_modalities, official_modalities)
|
||||
_assert_tensor_exact(fastvideo_time, official_time)
|
||||
_assert_tensor_exact(fastvideo_varlen.cu_seqlens_q, official_varlen.cu_seqlens_q)
|
||||
_assert_tensor_exact(fastvideo_varlen.cu_seqlens_k, official_varlen.cu_seqlens_k)
|
||||
assert fastvideo_varlen.max_seqlen_q == official_varlen.max_seqlen_q
|
||||
assert fastvideo_varlen.max_seqlen_k == official_varlen.max_seqlen_k
|
||||
|
||||
|
||||
def _make_t2v_input_tensors() -> dict[str, Any]:
|
||||
"""Create two distinguishable CFG samples with video, audio, and text."""
|
||||
video = torch.arange(2 * 3 * 2 * 4 * 4, dtype=torch.float32).reshape(
|
||||
2, 3, 2, 4, 4
|
||||
)
|
||||
video[1].add_(1000)
|
||||
audio = torch.arange(2 * 5 * 4, dtype=torch.float32).reshape(2, 5, 4)
|
||||
audio[1].add_(2000)
|
||||
text = torch.arange(2 * 4 * 6, dtype=torch.float32).reshape(2, 4, 6)
|
||||
text[1].add_(3000)
|
||||
per_token_video_time = torch.linspace(
|
||||
0.05,
|
||||
0.85,
|
||||
steps=2 * 2 * 4 * 4,
|
||||
dtype=torch.float32,
|
||||
).reshape(2, 1, 2, 4, 4)
|
||||
per_token_audio_time = torch.linspace(
|
||||
0.1,
|
||||
0.9,
|
||||
steps=2 * 5,
|
||||
dtype=torch.float32,
|
||||
).reshape(2, 5, 1)
|
||||
return {
|
||||
"x_t": video,
|
||||
"audio_x_t": audio,
|
||||
"audio_feat_len": torch.tensor([5, 3], dtype=torch.int32),
|
||||
"txt_feat": text,
|
||||
"txt_feat_len": torch.tensor([4, 2], dtype=torch.int32),
|
||||
"t": torch.tensor([750, 750], dtype=torch.int64),
|
||||
"per_token_video_t": per_token_video_time,
|
||||
"per_token_audio_t": per_token_audio_time,
|
||||
}
|
||||
|
||||
|
||||
def _make_i2v_input_tensors() -> dict[str, Any]:
|
||||
"""Create two CFG samples with two conditioning images per sample."""
|
||||
generator = torch.Generator(device="cpu").manual_seed(314159)
|
||||
return {
|
||||
"x_t": torch.randn(
|
||||
(2, 3, 2, 4, 4),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
"audio_x_t": torch.randn(
|
||||
(2, 4, 4),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
"audio_feat_len": torch.tensor([4, 2], dtype=torch.int32),
|
||||
"txt_feat": torch.randn(
|
||||
(2, 3, 6),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).to(torch.bfloat16),
|
||||
"txt_feat_len": torch.tensor([3, 2], dtype=torch.int32),
|
||||
"t": torch.tensor([0.6, 0.6], dtype=torch.float32),
|
||||
"per_token_video_t": torch.full(
|
||||
(2, 1, 2, 4, 4),
|
||||
0.6,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
"per_token_audio_t": torch.full(
|
||||
(2, 4, 1),
|
||||
0.6,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
"ref_image_feat": torch.randn(
|
||||
(2, 2, 3, 1, 4, 4),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
"ref_image_feat_len": torch.tensor(
|
||||
[[[2, 2], [2, 2]], [[2, 2], [2, 2]]],
|
||||
dtype=torch.int32,
|
||||
),
|
||||
"ref_image_special_token_embedding": torch.randn(
|
||||
(2, 2, 6),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).to(torch.bfloat16),
|
||||
}
|
||||
|
||||
|
||||
def _assert_process_output_exact(
|
||||
official_proxy: Any,
|
||||
fastvideo_proxy: Magi2DataProxy,
|
||||
packed_token_count: int,
|
||||
output_channel: int,
|
||||
) -> None:
|
||||
"""Compare depacked video and audio after deterministic model-like output."""
|
||||
model_output = torch.arange(
|
||||
packed_token_count * output_channel,
|
||||
dtype=torch.float32,
|
||||
).reshape(packed_token_count, output_channel)
|
||||
official_video, official_audio = official_proxy.process_output(model_output)
|
||||
fastvideo_video, fastvideo_audio = fastvideo_proxy.process_output(model_output)
|
||||
_assert_tensor_exact(fastvideo_video, official_video)
|
||||
_assert_tensor_exact(fastvideo_audio, official_audio)
|
||||
|
||||
|
||||
def test_magi2_data_proxy_config_defaults_match_official() -> None:
|
||||
"""Keep FastVideo proxy defaults aligned with the official configuration."""
|
||||
official_module = _load_official_proxy_module()
|
||||
official_defaults = official_module.DataProxyConfig().model_dump()
|
||||
|
||||
assert asdict(Magi2DataProxyConfig()) == official_defaults
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("parallel_dimension", "cp_size", "ep_size"),
|
||||
[("cp", 2, 1), ("ep", 1, 2)],
|
||||
ids=("context_parallel", "expert_parallel"),
|
||||
)
|
||||
def test_process_input_distributed_remote_max_padding_matches_official(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
parallel_dimension: str,
|
||||
cp_size: int,
|
||||
ep_size: int,
|
||||
) -> None:
|
||||
"""Match distributed maximum reduction and multiple-of-48 padding."""
|
||||
official_module = _load_official_proxy_module()
|
||||
parallel_state = _ParallelStateStub(cp_size=cp_size, ep_size=ep_size)
|
||||
monkeypatch.setattr(official_module, "psm", parallel_state)
|
||||
monkeypatch.setattr(fastvideo_proxy_module, "psm", parallel_state)
|
||||
reduced_groups: list[Any] = []
|
||||
|
||||
def fake_all_reduce(
|
||||
tensor: torch.Tensor,
|
||||
op: Any = None,
|
||||
group: Any = None,
|
||||
) -> None:
|
||||
"""Simulate a remote rank whose packed sequence contains 49 tokens."""
|
||||
del op
|
||||
reduced_groups.append(group)
|
||||
tensor.fill_(49)
|
||||
|
||||
monkeypatch.setattr(torch.distributed, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True)
|
||||
monkeypatch.setattr(torch.distributed, "all_reduce", fake_all_reduce)
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
config = Magi2DataProxyConfig(
|
||||
t_patch_size=2,
|
||||
patch_size=2,
|
||||
spatial_rope_interpolation="extra",
|
||||
add_time_token=False,
|
||||
time_channel_dim=1,
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
official_module,
|
||||
_make_t2v_input_tensors(),
|
||||
)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, _, modalities, varlen_handler, time_features = fastvideo_output
|
||||
assert tokens.shape == (96, 24)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 13, 22, 96]
|
||||
assert varlen_handler.cu_seqlens_k.tolist() == [0, 13, 22, 96]
|
||||
assert varlen_handler.max_seqlen_q == 74
|
||||
assert varlen_handler.max_seqlen_k == 74
|
||||
assert modalities[22:].tolist() == [int(Modality.TEXT)] * 74
|
||||
assert torch.count_nonzero(tokens[22:]) == 0
|
||||
assert torch.count_nonzero(time_features[22:]) == 0
|
||||
assert reduced_groups == [
|
||||
f"{parallel_dimension}-group",
|
||||
f"{parallel_dimension}-group",
|
||||
]
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=96,
|
||||
output_channel=24,
|
||||
)
|
||||
|
||||
|
||||
def test_process_input_t2v_cfg_and_cp_padding_match_official(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Match T2V packing, CFG boundaries, time features, and CP padding."""
|
||||
official_module = _load_official_proxy_module()
|
||||
parallel_state = _ParallelStateStub(cp_size=2, ep_size=1)
|
||||
monkeypatch.setattr(official_module, "psm", parallel_state)
|
||||
monkeypatch.setattr(fastvideo_proxy_module, "psm", parallel_state)
|
||||
config = Magi2DataProxyConfig(
|
||||
t_patch_size=2,
|
||||
patch_size=2,
|
||||
spatial_rope_interpolation="extra",
|
||||
add_time_token=False,
|
||||
time_channel_dim=7,
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
official_module,
|
||||
_make_t2v_input_tensors(),
|
||||
)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, _, modalities, varlen_handler, time_features = fastvideo_output
|
||||
assert tokens.shape == (48, 24)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 13, 22, 48]
|
||||
assert varlen_handler.cu_seqlens_k.tolist() == [0, 13, 22, 48]
|
||||
expected_modalities = (
|
||||
[int(Modality.VIDEO)] * 4
|
||||
+ [int(Modality.AUDIO)] * 5
|
||||
+ [int(Modality.TEXT)] * 4
|
||||
+ [int(Modality.VIDEO)] * 4
|
||||
+ [int(Modality.AUDIO)] * 3
|
||||
+ [int(Modality.TEXT)] * 2
|
||||
+ [int(Modality.TEXT)] * 26
|
||||
)
|
||||
assert modalities.tolist() == expected_modalities
|
||||
assert not torch.equal(tokens[:4], tokens[13:17])
|
||||
assert_close(
|
||||
time_features[9:13],
|
||||
torch.tensor([1, 1, 1, 0, 0, 0, 0], dtype=torch.float32).expand(4, -1),
|
||||
atol=0,
|
||||
rtol=0,
|
||||
)
|
||||
assert torch.count_nonzero(time_features[22:]) == 0
|
||||
assert official_proxy._saved_data.keys() == fastvideo_proxy._saved_data.keys()
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=48,
|
||||
output_channel=24,
|
||||
)
|
||||
|
||||
|
||||
def test_process_input_i2v_token_order_and_coordinates_match_official(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Match I2V special-token order, image coordinates, times, and depacking."""
|
||||
official_module = _load_official_proxy_module()
|
||||
parallel_state = _ParallelStateStub(cp_size=1, ep_size=1)
|
||||
monkeypatch.setattr(official_module, "psm", parallel_state)
|
||||
monkeypatch.setattr(fastvideo_proxy_module, "psm", parallel_state)
|
||||
config = Magi2DataProxyConfig(
|
||||
t_patch_size=1,
|
||||
patch_size=2,
|
||||
spatial_rope_interpolation="extra",
|
||||
add_time_token=False,
|
||||
time_channel_dim=1,
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
input_tensors = _make_i2v_input_tensors()
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
official_module,
|
||||
input_tensors,
|
||||
)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, coords, modalities, varlen_handler, time_features = fastvideo_output
|
||||
assert tokens.shape == (47, 12)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 25, 47]
|
||||
first_sample_modalities = (
|
||||
[int(Modality.VIDEO)] * 8
|
||||
+ [int(Modality.AUDIO)] * 4
|
||||
+ [int(Modality.TEXT)] * 3
|
||||
+ [int(Modality.TEXT)]
|
||||
+ [int(Modality.VIDEO)] * 4
|
||||
+ [int(Modality.TEXT)]
|
||||
+ [int(Modality.VIDEO)] * 4
|
||||
)
|
||||
second_sample_modalities = (
|
||||
[int(Modality.VIDEO)] * 8
|
||||
+ [int(Modality.AUDIO)] * 2
|
||||
+ [int(Modality.TEXT)] * 2
|
||||
+ [int(Modality.TEXT)]
|
||||
+ [int(Modality.VIDEO)] * 4
|
||||
+ [int(Modality.TEXT)]
|
||||
+ [int(Modality.VIDEO)] * 4
|
||||
)
|
||||
assert modalities.tolist() == first_sample_modalities + second_sample_modalities
|
||||
|
||||
first_special_offset = 8 + 4 + 3
|
||||
first_special = input_tensors["ref_image_special_token_embedding"][0, 0]
|
||||
_assert_tensor_exact(
|
||||
tokens[first_special_offset, :6],
|
||||
first_special.to(tokens.dtype),
|
||||
)
|
||||
first_image_patch = input_tensors["ref_image_feat"][0, 0, :, :, :2, :2].flatten()
|
||||
_assert_tensor_exact(tokens[first_special_offset + 1], first_image_patch)
|
||||
assert coords[first_special_offset].tolist() == [4, -1, -1, 1, 2, 2, 1, 2, 2]
|
||||
assert coords[first_special_offset + 5].tolist() == [5, -1, -1, 1, 2, 2, 1, 2, 2]
|
||||
assert torch.count_nonzero(time_features[12:25]) == 0
|
||||
assert official_proxy._saved_data.keys() == fastvideo_proxy._saved_data.keys()
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=47,
|
||||
output_channel=12,
|
||||
)
|
||||
@@ -0,0 +1,160 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict eight-GPU numerical parity for the MAGI-2 preview transformer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
WORKER_PATH = Path(__file__).with_name("_preview_transformer_parity_worker.py")
|
||||
CAPTURE_ROOT = REPO_ROOT / "archived" / "magi2_parity" / "validation" / "preview_transformer"
|
||||
OFFICIAL_REVISION = "073c84f2102ec3c9287623113a103c14402770ad"
|
||||
WORLD_SIZE = 8
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
torch.cuda.device_count() < WORLD_SIZE,
|
||||
reason=f"MAGI-2 preview numerical parity requires {WORLD_SIZE} CUDA devices",
|
||||
)
|
||||
|
||||
|
||||
def _require_parity_sources() -> None:
|
||||
"""Require the pinned official source and complete preview checkpoint."""
|
||||
model_path = OFFICIAL_ROOT / "inference" / "model" / "magi2_preview.py"
|
||||
if not model_path.is_file():
|
||||
raise AssertionError(f"Official MAGI-2 source is missing: {model_path}")
|
||||
index_path = WEIGHTS_ROOT / "preview" / "model.safetensors.index.json"
|
||||
if not index_path.is_file():
|
||||
raise AssertionError(f"MAGI-2 preview checkpoint index is missing: {index_path}")
|
||||
revision = subprocess.run(
|
||||
["git", "rev-parse", "HEAD"],
|
||||
cwd=OFFICIAL_ROOT,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
).stdout.strip()
|
||||
assert revision == OFFICIAL_REVISION
|
||||
|
||||
|
||||
def _run_implementation(implementation: str) -> Path:
|
||||
"""Launch one implementation in an isolated eight-rank torchrun job."""
|
||||
output_dir = CAPTURE_ROOT / implementation
|
||||
environment = os.environ.copy()
|
||||
environment.update(
|
||||
{
|
||||
"CUBLAS_WORKSPACE_CONFIG": ":4096:8",
|
||||
"MAGI2_CKPT_ROOT": str(WEIGHTS_ROOT),
|
||||
"MAGI2_DETERMINISTIC": "1",
|
||||
"MAGI2_LOCAL_WEIGHTS_DIR": str(WEIGHTS_ROOT),
|
||||
"MAGI2_OFFICIAL_REF_DIR": str(OFFICIAL_ROOT),
|
||||
"MAGI_ATTENTION_DETERMINISTIC_MODE": "1",
|
||||
"MAGI_COMPILE_COMPILE_MODE": "NONE",
|
||||
"OMP_NUM_THREADS": "1",
|
||||
"PYTHONHASHSEED": "42",
|
||||
}
|
||||
)
|
||||
environment.pop("SKIP_LOAD_MODEL", None)
|
||||
completed = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--standalone",
|
||||
"--nnodes=1",
|
||||
f"--nproc-per-node={WORLD_SIZE}",
|
||||
str(WORKER_PATH),
|
||||
"--implementation",
|
||||
implementation,
|
||||
"--output-dir",
|
||||
str(output_dir),
|
||||
],
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=7200,
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
raise AssertionError(
|
||||
f"{implementation} preview transformer torchrun failed.\n"
|
||||
f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}"
|
||||
)
|
||||
return output_dir
|
||||
|
||||
|
||||
def _assert_exact(actual: Any, expected: Any, path: str) -> None:
|
||||
"""Recursively require identical capture structure, metadata, and values."""
|
||||
assert type(actual) is type(expected), (
|
||||
f"{path}: type differs: {type(actual).__name__} versus "
|
||||
f"{type(expected).__name__}"
|
||||
)
|
||||
if isinstance(expected, torch.Tensor):
|
||||
assert actual.shape == expected.shape, path
|
||||
assert actual.dtype == expected.dtype, path
|
||||
assert actual.stride() == expected.stride(), path
|
||||
if not torch.equal(actual, expected):
|
||||
difference = (actual.float() - expected.float()).abs()
|
||||
raise AssertionError(
|
||||
f"{path}: tensor values differ; mismatched="
|
||||
f"{torch.count_nonzero(actual != expected).item()}, "
|
||||
f"max_abs={difference.max().item()}"
|
||||
)
|
||||
return
|
||||
if isinstance(expected, dict):
|
||||
assert actual.keys() == expected.keys(), path
|
||||
for key in expected:
|
||||
_assert_exact(actual[key], expected[key], f"{path}.{key}")
|
||||
return
|
||||
if isinstance(expected, (list, tuple)):
|
||||
assert len(actual) == len(expected), path
|
||||
for index, (actual_value, expected_value) in enumerate(
|
||||
zip(actual, expected, strict=True)
|
||||
):
|
||||
_assert_exact(actual_value, expected_value, f"{path}[{index}]")
|
||||
return
|
||||
assert actual == expected, path
|
||||
|
||||
|
||||
def _load_rank_capture(output_dir: Path, rank: int) -> dict[str, Any]:
|
||||
"""Load one rank artifact and validate its fixed capture envelope."""
|
||||
artifact_path = output_dir / f"rank_{rank}.pt"
|
||||
if not artifact_path.is_file():
|
||||
raise AssertionError(f"Preview parity capture is missing: {artifact_path}")
|
||||
artifact = torch.load(artifact_path, map_location="cpu", weights_only=True)
|
||||
assert artifact["schema_version"] == 1
|
||||
assert artifact["rank"] == rank
|
||||
assert artifact["world_size"] == WORLD_SIZE
|
||||
assert artifact["cases"].keys() == {"t2v", "i2v"}
|
||||
return artifact
|
||||
|
||||
|
||||
def test_magi2_preview_transformer_matches_official_exactly() -> None:
|
||||
"""Match T2V and I2V tensors at every preview transformer boundary."""
|
||||
_require_parity_sources()
|
||||
official_output_dir = _run_implementation("official")
|
||||
fastvideo_output_dir = _run_implementation("fastvideo")
|
||||
for rank in range(WORLD_SIZE):
|
||||
official_capture = _load_rank_capture(official_output_dir, rank)
|
||||
fastvideo_capture = _load_rank_capture(fastvideo_output_dir, rank)
|
||||
assert official_capture["implementation"] == "official"
|
||||
assert fastvideo_capture["implementation"] == "fastvideo"
|
||||
_assert_exact(
|
||||
fastvideo_capture["cases"],
|
||||
official_capture["cases"],
|
||||
f"rank_{rank}.cases",
|
||||
)
|
||||
@@ -0,0 +1,604 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict MAGI-2 refiner data-proxy implementation parity.
|
||||
|
||||
Coverage scope: implementation_subcomponent. The tests load the pinned official
|
||||
``Magi2RefinerDataProxy`` source and compare every packed or depacked tensor
|
||||
with the FastVideo-owned component. A scheduler check verifies that uneven
|
||||
context-parallel splitting remains owned by ``UlyssesScheduler``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from dataclasses import asdict, dataclass, fields
|
||||
from functools import lru_cache
|
||||
import importlib
|
||||
import importlib.util
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
import typing
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
from torch.utils._pytree import tree_map
|
||||
|
||||
from fastvideo.pipelines.basic.magi2.stages.refiner_data_proxy import (
|
||||
Magi2RefinerDataProxy,
|
||||
Magi2RefinerDataProxyConfig,
|
||||
Modality,
|
||||
RefinerModelInput,
|
||||
VarlenHandler,
|
||||
WindowLocalAttnHandler,
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REPO_ROOT = Path(
|
||||
os.environ.get(
|
||||
"MAGI2_OFFICIAL_REF_DIR",
|
||||
REPO_ROOT.parent / "MAGI-2-preview",
|
||||
)
|
||||
)
|
||||
OFFICIAL_PROXY_PATH = OFFICIAL_REPO_ROOT / "inference" / "pipeline" / "refiner_data_proxy.py"
|
||||
OFFICIAL_MODEL_PATH = OFFICIAL_REPO_ROOT / "inference" / "model" / "magi2_refiner.py"
|
||||
OFFICIAL_SCHEDULER_PATH = (
|
||||
OFFICIAL_REPO_ROOT
|
||||
/ "inference"
|
||||
/ "infra"
|
||||
/ "parallelism"
|
||||
/ "context_parallel"
|
||||
/ "ulysses_scheduler.py"
|
||||
)
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ParallelStateStub:
|
||||
"""Expose one fixed context-parallel size and process-group marker."""
|
||||
|
||||
cp_size: int
|
||||
|
||||
def get_world_size(self, dimension: str = "") -> int:
|
||||
"""Return the configured context-parallel world size."""
|
||||
return self.cp_size if dimension == "cp" else 1
|
||||
|
||||
def get_parallel_group(self, dimension: str) -> str:
|
||||
"""Return an identifiable process-group marker for call assertions."""
|
||||
return f"{dimension}-group"
|
||||
|
||||
|
||||
def _load_official_refiner_model_utilities() -> types.ModuleType:
|
||||
"""Load the official frame-local attention helper without model dependencies."""
|
||||
if not OFFICIAL_MODEL_PATH.is_file():
|
||||
raise FileNotFoundError(f"Official MAGI-2 refiner source is missing: {OFFICIAL_MODEL_PATH}")
|
||||
source_module = ast.parse(
|
||||
OFFICIAL_MODEL_PATH.read_text(encoding="utf-8"),
|
||||
filename=str(OFFICIAL_MODEL_PATH),
|
||||
)
|
||||
symbol_names = {
|
||||
"FFAHandler",
|
||||
"calc_local_qk_range",
|
||||
"calc_local_attn_ffa_handler",
|
||||
}
|
||||
symbol_nodes = [
|
||||
node
|
||||
for node in source_module.body
|
||||
if isinstance(node, (ast.ClassDef, ast.FunctionDef))
|
||||
and node.name in symbol_names
|
||||
]
|
||||
found_names = {node.name for node in symbol_nodes}
|
||||
if found_names != symbol_names:
|
||||
raise AssertionError(
|
||||
f"Official refiner utility definitions differ: expected {symbol_names}, found {found_names}"
|
||||
)
|
||||
|
||||
module = types.ModuleType("inference.model.magi2_refiner")
|
||||
module.__file__ = str(OFFICIAL_MODEL_PATH)
|
||||
module.__dict__.update(
|
||||
{
|
||||
"Optional": typing.Optional,
|
||||
"dataclass": dataclass,
|
||||
"torch": torch,
|
||||
}
|
||||
)
|
||||
module_name = module.__name__
|
||||
previous_module = sys.modules.get(module_name)
|
||||
sys.modules[module_name] = module
|
||||
try:
|
||||
exec(
|
||||
compile(
|
||||
ast.Module(body=symbol_nodes, type_ignores=[]),
|
||||
str(OFFICIAL_MODEL_PATH),
|
||||
"exec",
|
||||
),
|
||||
module.__dict__,
|
||||
)
|
||||
finally:
|
||||
if previous_module is None:
|
||||
sys.modules.pop(module_name, None)
|
||||
else:
|
||||
sys.modules[module_name] = previous_module
|
||||
return module
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _load_official_proxy_module() -> types.ModuleType:
|
||||
"""Import the official proxy while isolating unrelated refiner model code."""
|
||||
if not OFFICIAL_PROXY_PATH.is_file():
|
||||
raise FileNotFoundError(f"Official MAGI-2 refiner proxy is missing: {OFFICIAL_PROXY_PATH}")
|
||||
if importlib.util.find_spec("unfoldNd") is None:
|
||||
pytest.skip("Official MAGI-2 refiner data-proxy parity requires unfoldNd==0.2.3.")
|
||||
|
||||
model_module_name = "inference.model.magi2_refiner"
|
||||
missing_module = object()
|
||||
previous_model_module = sys.modules.get(model_module_name, missing_module)
|
||||
sys.modules[model_module_name] = _load_official_refiner_model_utilities()
|
||||
official_path_entry = str(OFFICIAL_REPO_ROOT)
|
||||
added_path = official_path_entry not in sys.path
|
||||
if added_path:
|
||||
sys.path.insert(0, official_path_entry)
|
||||
module_name = "magi2_official_refiner_data_proxy"
|
||||
try:
|
||||
spec = importlib.util.spec_from_file_location(module_name, OFFICIAL_PROXY_PATH)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(f"Cannot create an import spec for {OFFICIAL_PROXY_PATH}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[module_name] = module
|
||||
spec.loader.exec_module(module)
|
||||
finally:
|
||||
if added_path:
|
||||
sys.path.remove(official_path_entry)
|
||||
if previous_model_module is missing_module:
|
||||
sys.modules.pop(model_module_name, None)
|
||||
else:
|
||||
sys.modules[model_module_name] = previous_model_module
|
||||
return module
|
||||
|
||||
|
||||
def _load_official_scheduler_class() -> type[Any]:
|
||||
"""Load the official Ulysses scheduler class with injectable collectives."""
|
||||
if not OFFICIAL_SCHEDULER_PATH.is_file():
|
||||
raise FileNotFoundError(f"Official Ulysses scheduler is missing: {OFFICIAL_SCHEDULER_PATH}")
|
||||
source_module = ast.parse(
|
||||
OFFICIAL_SCHEDULER_PATH.read_text(encoding="utf-8"),
|
||||
filename=str(OFFICIAL_SCHEDULER_PATH),
|
||||
)
|
||||
class_nodes = [
|
||||
node
|
||||
for node in source_module.body
|
||||
if isinstance(node, ast.ClassDef) and node.name == "UlyssesScheduler"
|
||||
]
|
||||
if len(class_nodes) != 1:
|
||||
raise AssertionError("Official source must define exactly one UlyssesScheduler class")
|
||||
module = types.ModuleType("magi2_official_ulysses_scheduler")
|
||||
module.__dict__.update(
|
||||
{
|
||||
"Generic": typing.Generic,
|
||||
"List": list,
|
||||
"Optional": typing.Optional,
|
||||
"T": typing.TypeVar("T"),
|
||||
"torch": torch,
|
||||
"tree_map": tree_map,
|
||||
}
|
||||
)
|
||||
exec(
|
||||
compile(
|
||||
ast.Module(body=class_nodes, type_ignores=[]),
|
||||
str(OFFICIAL_SCHEDULER_PATH),
|
||||
"exec",
|
||||
),
|
||||
module.__dict__,
|
||||
)
|
||||
return module.UlyssesScheduler
|
||||
|
||||
|
||||
def _cuda_device() -> torch.device:
|
||||
"""Return the CUDA device required by the pinned official proxy."""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("Official MAGI-2 refiner data-proxy parity requires CUDA.")
|
||||
return torch.device("cuda", torch.cuda.current_device())
|
||||
|
||||
|
||||
def _build_proxy_pair(
|
||||
official_module: types.ModuleType,
|
||||
config: Magi2RefinerDataProxyConfig,
|
||||
) -> tuple[Any, Magi2RefinerDataProxy]:
|
||||
"""Construct official and FastVideo proxies from identical config values."""
|
||||
official_config = official_module.Magi2RefinerDataProxyConfig(**asdict(config))
|
||||
return (
|
||||
official_module.Magi2RefinerDataProxy(official_config),
|
||||
Magi2RefinerDataProxy(config),
|
||||
)
|
||||
|
||||
|
||||
def _build_model_input_pair(
|
||||
input_tensors: dict[str, torch.Tensor],
|
||||
) -> tuple[types.SimpleNamespace, RefinerModelInput]:
|
||||
"""Construct both proxy inputs over the same immutable test tensors."""
|
||||
return (
|
||||
types.SimpleNamespace(**input_tensors),
|
||||
RefinerModelInput(**input_tensors),
|
||||
)
|
||||
|
||||
|
||||
def _make_input_tensors(
|
||||
*,
|
||||
device: torch.device,
|
||||
batch_size: int,
|
||||
ref_video_token_limit: int,
|
||||
distinguish_cfg_samples: bool = False,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Create deterministic refiner video, audio, text, and reference features."""
|
||||
video = torch.arange(
|
||||
batch_size * 3 * 3 * 6 * 8,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(batch_size, 3, 3, 6, 8)
|
||||
audio = torch.arange(
|
||||
batch_size * 5 * 4,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(batch_size, 5, 4)
|
||||
text = torch.arange(
|
||||
batch_size * 4 * 6,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(batch_size, 4, 6)
|
||||
reference_audio = torch.arange(
|
||||
batch_size * 3 * 4,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(batch_size, 3, 4)
|
||||
reference_video = torch.arange(
|
||||
batch_size * 3 * 3 * 6 * 8,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(batch_size, 3, 3, 6, 8)
|
||||
if distinguish_cfg_samples and batch_size > 1:
|
||||
video[1].add_(1000)
|
||||
audio[1].add_(2000)
|
||||
text[1].add_(3000)
|
||||
reference_audio[1].add_(4000)
|
||||
reference_video[1].add_(5000)
|
||||
return {
|
||||
"x_t": video,
|
||||
"audio_x_t": audio,
|
||||
"audio_feat_len": torch.full((batch_size,), 4, dtype=torch.int64),
|
||||
"txt_feat": text,
|
||||
"txt_feat_len": torch.full((batch_size,), 3, dtype=torch.int64),
|
||||
"ref_audio_feat": reference_audio,
|
||||
"ref_audio_feat_len": torch.full((batch_size,), 2, dtype=torch.int64),
|
||||
"ref_video_feat": reference_video,
|
||||
"ref_video_feat_len": torch.full(
|
||||
(batch_size,),
|
||||
ref_video_token_limit,
|
||||
dtype=torch.int64,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _assert_tensor_exact(
|
||||
fastvideo_tensor: torch.Tensor,
|
||||
official_tensor: torch.Tensor,
|
||||
) -> None:
|
||||
"""Require identical tensor metadata and values."""
|
||||
assert fastvideo_tensor.shape == official_tensor.shape
|
||||
assert fastvideo_tensor.dtype == official_tensor.dtype
|
||||
assert fastvideo_tensor.stride() == official_tensor.stride()
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=0, rtol=0)
|
||||
|
||||
|
||||
def _assert_scalar_or_tensor_exact(
|
||||
fastvideo_value: int | torch.Tensor,
|
||||
official_value: int | torch.Tensor,
|
||||
) -> None:
|
||||
"""Compare attention-length metadata without discarding tensor dtype."""
|
||||
if isinstance(official_value, torch.Tensor):
|
||||
assert isinstance(fastvideo_value, torch.Tensor)
|
||||
_assert_tensor_exact(fastvideo_value, official_value)
|
||||
else:
|
||||
assert fastvideo_value == official_value
|
||||
|
||||
|
||||
def _assert_local_attn_handler_exact(
|
||||
fastvideo_handler: WindowLocalAttnHandler | None,
|
||||
official_handler: Any,
|
||||
) -> None:
|
||||
"""Compare every attention-range field exposed by the official handler."""
|
||||
if official_handler is None:
|
||||
assert fastvideo_handler is None
|
||||
return
|
||||
assert fastvideo_handler is not None
|
||||
_assert_tensor_exact(fastvideo_handler.q_ranges, official_handler.q_ranges)
|
||||
_assert_tensor_exact(fastvideo_handler.k_ranges, official_handler.k_ranges)
|
||||
_assert_tensor_exact(fastvideo_handler.attn_type_map, official_handler.attn_type_map)
|
||||
assert fastvideo_handler.max_seqlen_q == official_handler.max_seqlen_q
|
||||
assert fastvideo_handler.max_seqlen_k == official_handler.max_seqlen_k
|
||||
assert fastvideo_handler.softmax_scale == official_handler.softmax_scale
|
||||
for field_name in (
|
||||
"bwd_q_ranges",
|
||||
"bwd_k_ranges",
|
||||
"bwd_attn_type_map",
|
||||
):
|
||||
official_value = getattr(official_handler, field_name, None)
|
||||
fastvideo_value = getattr(fastvideo_handler, field_name)
|
||||
if official_value is None:
|
||||
assert fastvideo_value is None
|
||||
else:
|
||||
assert fastvideo_value is not None
|
||||
_assert_tensor_exact(fastvideo_value, official_value)
|
||||
assert fastvideo_handler.auto_range_merge == official_handler.auto_range_merge
|
||||
assert fastvideo_handler.sparse_load == official_handler.sparse_load
|
||||
|
||||
|
||||
def _assert_packed_inputs_exact(
|
||||
fastvideo_output: tuple[
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
VarlenHandler,
|
||||
WindowLocalAttnHandler | None,
|
||||
],
|
||||
official_output: tuple[torch.Tensor, torch.Tensor, torch.Tensor, Any, Any],
|
||||
) -> None:
|
||||
"""Compare all packed tensors and attention metadata."""
|
||||
(
|
||||
fastvideo_tokens,
|
||||
fastvideo_coords,
|
||||
fastvideo_modalities,
|
||||
fastvideo_varlen,
|
||||
fastvideo_local_attn,
|
||||
) = fastvideo_output
|
||||
(
|
||||
official_tokens,
|
||||
official_coords,
|
||||
official_modalities,
|
||||
official_varlen,
|
||||
official_local_attn,
|
||||
) = official_output
|
||||
_assert_tensor_exact(fastvideo_tokens, official_tokens)
|
||||
_assert_tensor_exact(fastvideo_coords, official_coords)
|
||||
_assert_tensor_exact(fastvideo_modalities, official_modalities)
|
||||
_assert_tensor_exact(fastvideo_varlen.cu_seqlens_q, official_varlen.cu_seqlens_q)
|
||||
_assert_tensor_exact(fastvideo_varlen.cu_seqlens_k, official_varlen.cu_seqlens_k)
|
||||
_assert_scalar_or_tensor_exact(fastvideo_varlen.max_seqlen_q, official_varlen.max_seqlen_q)
|
||||
_assert_scalar_or_tensor_exact(fastvideo_varlen.max_seqlen_k, official_varlen.max_seqlen_k)
|
||||
_assert_local_attn_handler_exact(fastvideo_local_attn, official_local_attn)
|
||||
|
||||
|
||||
def _assert_process_output_exact(
|
||||
official_proxy: Any,
|
||||
fastvideo_proxy: Magi2RefinerDataProxy,
|
||||
packed_token_count: int,
|
||||
output_channel: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Compare depacked video and audio after deterministic model-like output."""
|
||||
model_output = torch.arange(
|
||||
packed_token_count * output_channel,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
).reshape(packed_token_count, output_channel)
|
||||
official_video, official_audio = official_proxy.process_output(model_output)
|
||||
fastvideo_video, fastvideo_audio = fastvideo_proxy.process_output(model_output)
|
||||
_assert_tensor_exact(fastvideo_video, official_video)
|
||||
_assert_tensor_exact(fastvideo_audio, official_audio)
|
||||
|
||||
|
||||
def test_magi2_refiner_data_proxy_config_defaults_match_official() -> None:
|
||||
"""Keep FastVideo refiner proxy defaults aligned with the official config."""
|
||||
official_module = _load_official_proxy_module()
|
||||
official_defaults = official_module.Magi2RefinerDataProxyConfig().model_dump()
|
||||
|
||||
assert asdict(Magi2RefinerDataProxyConfig()) == official_defaults
|
||||
|
||||
|
||||
def test_process_input_default_frame_local_attention_matches_official() -> None:
|
||||
"""Match default patching, coordinates, varlen metadata, and frame ranges."""
|
||||
device = _cuda_device()
|
||||
official_module = _load_official_proxy_module()
|
||||
config = Magi2RefinerDataProxyConfig()
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
_make_input_tensors(
|
||||
device=device,
|
||||
batch_size=1,
|
||||
ref_video_token_limit=0,
|
||||
)
|
||||
)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, _, modalities, varlen_handler, local_attn_handler = fastvideo_output
|
||||
assert tokens.shape == (45, 12)
|
||||
assert modalities.tolist() == (
|
||||
[int(Modality.VIDEO)] * 36
|
||||
+ [int(Modality.AUDIO)] * 4
|
||||
+ [int(Modality.TEXT)] * 3
|
||||
+ [int(Modality.AUDIO)] * 2
|
||||
)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 45]
|
||||
assert local_attn_handler is not None
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=tokens.shape[0],
|
||||
output_channel=tokens.shape[1],
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def test_process_input_cfg_boundaries_and_no_time_features_match_official() -> None:
|
||||
"""Preserve adjacent CFG samples without adding timestep feature tokens."""
|
||||
device = _cuda_device()
|
||||
official_module = _load_official_proxy_module()
|
||||
config = Magi2RefinerDataProxyConfig(
|
||||
patch_size=2,
|
||||
frame_receptive_field=-1,
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
input_tensors = _make_input_tensors(
|
||||
device=device,
|
||||
batch_size=2,
|
||||
ref_video_token_limit=1,
|
||||
distinguish_cfg_samples=True,
|
||||
)
|
||||
official_input, fastvideo_input = _build_model_input_pair(input_tensors)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, _, _, varlen_handler, local_attn_handler = fastvideo_output
|
||||
assert tokens.shape == (92, 12)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 46, 92]
|
||||
assert local_attn_handler is None
|
||||
assert not torch.equal(tokens[:36], tokens[46:82])
|
||||
assert "t" not in {input_field.name for input_field in fields(RefinerModelInput)}
|
||||
assert len(fastvideo_output) == 5
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=tokens.shape[0],
|
||||
output_channel=tokens.shape[1],
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def test_process_input_shipping_block_window_and_depacking_match_official() -> None:
|
||||
"""Match the release block-grid token order, ranges, and inverse depacking."""
|
||||
device = _cuda_device()
|
||||
official_module = _load_official_proxy_module()
|
||||
config = Magi2RefinerDataProxyConfig(
|
||||
t_patch_size=1,
|
||||
patch_size=1,
|
||||
frame_receptive_field=11,
|
||||
spatial_rope_interpolation="extra",
|
||||
coords_style="v1",
|
||||
text_offset=0,
|
||||
attn_config={
|
||||
"mode": "window",
|
||||
"block_t_size": 8,
|
||||
"block_size": 4,
|
||||
"window": {
|
||||
"level": "block",
|
||||
"block_mode": "grid",
|
||||
"block_t_radius": 2,
|
||||
"block_h_radius": 2,
|
||||
"block_w_radius": 2,
|
||||
"win_size": 384,
|
||||
"frame_receptive_field": -1,
|
||||
"auto_range_merge": True,
|
||||
"sparse_load": False,
|
||||
"full_attn_layers": [],
|
||||
},
|
||||
},
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
_make_input_tensors(
|
||||
device=device,
|
||||
batch_size=1,
|
||||
ref_video_token_limit=4,
|
||||
)
|
||||
)
|
||||
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
|
||||
tokens, coords, modalities, varlen_handler, local_attn_handler = fastvideo_output
|
||||
assert tokens.shape == (157, 6)
|
||||
assert coords.shape == (157, 9)
|
||||
assert modalities.shape == (157,)
|
||||
assert varlen_handler.cu_seqlens_q.tolist() == [0, 157]
|
||||
assert local_attn_handler is not None
|
||||
assert local_attn_handler.auto_range_merge is True
|
||||
_assert_process_output_exact(
|
||||
official_proxy,
|
||||
fastvideo_proxy,
|
||||
packed_token_count=tokens.shape[0],
|
||||
output_channel=tokens.shape[1],
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def test_process_input_distributed_context_keeps_unpadded_official_sequence(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Keep distributed reduction and uneven splitting outside the refiner proxy."""
|
||||
device = _cuda_device()
|
||||
official_module = _load_official_proxy_module()
|
||||
config = Magi2RefinerDataProxyConfig(
|
||||
patch_size=2,
|
||||
frame_receptive_field=-1,
|
||||
)
|
||||
official_proxy, fastvideo_proxy = _build_proxy_pair(official_module, config)
|
||||
official_input, fastvideo_input = _build_model_input_pair(
|
||||
_make_input_tensors(
|
||||
device=device,
|
||||
batch_size=1,
|
||||
ref_video_token_limit=1,
|
||||
)
|
||||
)
|
||||
|
||||
def reject_all_reduce(*args: Any, **kwargs: Any) -> None:
|
||||
"""Fail if preview-style distributed maximum padding enters this proxy."""
|
||||
del args, kwargs
|
||||
raise AssertionError("The refiner proxy must not perform all_reduce")
|
||||
|
||||
monkeypatch.setattr(torch.distributed, "all_reduce", reject_all_reduce)
|
||||
official_output = official_proxy.process_input(official_input)
|
||||
fastvideo_output = fastvideo_proxy.process_input(fastvideo_input)
|
||||
_assert_packed_inputs_exact(fastvideo_output, official_output)
|
||||
packed_tokens = fastvideo_output[0]
|
||||
assert packed_tokens.shape[0] == 46
|
||||
assert packed_tokens.shape[0] % 48 != 0
|
||||
assert "pad_size" not in fastvideo_proxy._saved_data
|
||||
assert "pad_size" not in official_proxy._saved_data
|
||||
|
||||
official_scheduler_class = _load_official_scheduler_class()
|
||||
target_scheduler_module = importlib.import_module(
|
||||
"fastvideo.models.dits.magi2_runtime.context_parallel.ulysses_scheduler"
|
||||
)
|
||||
parallel_state = _ParallelStateStub(cp_size=3)
|
||||
official_capture: list[tuple[list[int], Any]] = []
|
||||
target_capture: list[tuple[list[int], Any]] = []
|
||||
|
||||
def official_scatter(
|
||||
tensor: torch.Tensor,
|
||||
split_sizes: list[int],
|
||||
group: Any,
|
||||
) -> torch.Tensor:
|
||||
"""Capture the official scheduler's uneven split without communication."""
|
||||
official_capture.append((split_sizes, group))
|
||||
return tensor
|
||||
|
||||
def target_scatter(
|
||||
tensor: torch.Tensor,
|
||||
split_sizes: list[int],
|
||||
group: Any,
|
||||
) -> torch.Tensor:
|
||||
"""Capture the FastVideo scheduler's uneven split without communication."""
|
||||
target_capture.append((split_sizes, group))
|
||||
return tensor
|
||||
|
||||
official_scheduler_class._dispatch.__globals__["psm"] = parallel_state
|
||||
official_scheduler_class._dispatch.__globals__["scatter_to_context_parallel_region"] = official_scatter
|
||||
monkeypatch.setattr(target_scheduler_module, "psm", parallel_state)
|
||||
monkeypatch.setattr(target_scheduler_module, "scatter_to_context_parallel_region", target_scatter)
|
||||
official_scheduler = official_scheduler_class()
|
||||
target_scheduler = target_scheduler_module.UlyssesScheduler()
|
||||
|
||||
official_dispatched = official_scheduler.dispatch(packed_tokens)
|
||||
target_dispatched = target_scheduler.dispatch(packed_tokens)
|
||||
_assert_tensor_exact(target_dispatched, official_dispatched)
|
||||
assert official_scheduler.cp_split_sizes == [16, 15, 15]
|
||||
assert target_scheduler.cp_split_sizes == [16, 15, 15]
|
||||
assert official_capture == [([16, 15, 15], "cp-group")]
|
||||
assert target_capture == official_capture
|
||||
@@ -0,0 +1,194 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict eight-GPU numerical parity for the MAGI-2 refiner transformer.
|
||||
|
||||
Coverage scope: both. The FastVideo side uses the production checkpoint loader,
|
||||
and the captures compare every implementation boundary that owns numerical
|
||||
behavior from the packed proxy input through the depacked video and audio.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(
|
||||
os.environ.get("MAGI2_OFFICIAL_REF_DIR", REPO_ROOT.parent / "MAGI-2-preview")
|
||||
)
|
||||
WEIGHTS_ROOT = Path(
|
||||
os.environ.get("MAGI2_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "magi2")
|
||||
)
|
||||
WORKER_PATH = Path(__file__).with_name("_refiner_transformer_parity_worker.py")
|
||||
CAPTURE_ROOT = REPO_ROOT / "archived" / "magi2_parity" / "validation" / "refiner_transformer"
|
||||
OFFICIAL_REVISION = "073c84f2102ec3c9287623113a103c14402770ad"
|
||||
WORLD_SIZE = 8
|
||||
|
||||
|
||||
def _require_parity_sources() -> None:
|
||||
"""Require eight GPUs, the pinned source, and every refiner weight shard."""
|
||||
if torch.cuda.device_count() < WORLD_SIZE:
|
||||
raise AssertionError(
|
||||
f"MAGI-2 refiner numerical parity requires {WORLD_SIZE} CUDA devices; "
|
||||
f"found {torch.cuda.device_count()}"
|
||||
)
|
||||
model_path = OFFICIAL_ROOT / "inference" / "model" / "magi2_refiner.py"
|
||||
if not model_path.is_file():
|
||||
raise AssertionError(f"Official MAGI-2 refiner source is missing: {model_path}")
|
||||
index_path = WEIGHTS_ROOT / "refiner" / "model.safetensors.index.json"
|
||||
if not index_path.is_file():
|
||||
raise AssertionError(f"MAGI-2 refiner checkpoint index is missing: {index_path}")
|
||||
weight_map = json.loads(index_path.read_text(encoding="utf-8"))["weight_map"]
|
||||
missing_shards = sorted(
|
||||
shard_name
|
||||
for shard_name in set(weight_map.values())
|
||||
if not (index_path.parent / shard_name).is_file()
|
||||
)
|
||||
if missing_shards:
|
||||
raise AssertionError(f"MAGI-2 refiner checkpoint shards are missing: {missing_shards}")
|
||||
revision = subprocess.run(
|
||||
["git", "rev-parse", "HEAD"],
|
||||
cwd=OFFICIAL_ROOT,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
).stdout.strip()
|
||||
assert revision == OFFICIAL_REVISION
|
||||
checkout_changes = subprocess.run(
|
||||
["git", "status", "--porcelain=v1", "--untracked-files=all"],
|
||||
cwd=OFFICIAL_ROOT,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
).stdout.strip()
|
||||
assert not checkout_changes, (
|
||||
"Official MAGI-2 checkout contains local changes:\n" + checkout_changes
|
||||
)
|
||||
|
||||
|
||||
def _run_implementation(implementation: str) -> Path:
|
||||
"""Launch one implementation in an isolated eight-rank torchrun job."""
|
||||
output_dir = CAPTURE_ROOT / implementation
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
environment = os.environ.copy()
|
||||
environment.update(
|
||||
{
|
||||
"CUBLAS_WORKSPACE_CONFIG": ":4096:8",
|
||||
"MAGI2_CKPT_ROOT": str(WEIGHTS_ROOT),
|
||||
"MAGI2_DETERMINISTIC": "1",
|
||||
"MAGI2_LOCAL_WEIGHTS_DIR": str(WEIGHTS_ROOT),
|
||||
"MAGI2_OFFICIAL_REF_DIR": str(OFFICIAL_ROOT),
|
||||
"MAGI_ATTENTION_DETERMINISTIC_MODE": "1",
|
||||
"MAGI_COMPILE_COMPILE_MODE": "NONE",
|
||||
"OMP_NUM_THREADS": "1",
|
||||
"PYTHONHASHSEED": "42",
|
||||
}
|
||||
)
|
||||
environment.pop("SKIP_LOAD_MODEL", None)
|
||||
completed = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--standalone",
|
||||
"--nnodes=1",
|
||||
f"--nproc-per-node={WORLD_SIZE}",
|
||||
str(WORKER_PATH),
|
||||
"--implementation",
|
||||
implementation,
|
||||
"--output-dir",
|
||||
str(output_dir),
|
||||
],
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=7200,
|
||||
)
|
||||
log_path = CAPTURE_ROOT / f"{implementation}.log"
|
||||
log_path.write_text(
|
||||
f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}",
|
||||
encoding="utf-8",
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
raise AssertionError(
|
||||
f"{implementation} refiner transformer torchrun failed; see {log_path}.\n"
|
||||
f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}"
|
||||
)
|
||||
return output_dir
|
||||
|
||||
|
||||
def _assert_exact(actual: Any, expected: Any, path: str) -> None:
|
||||
"""Recursively require identical capture structure, metadata, and values."""
|
||||
assert type(actual) is type(expected), (
|
||||
f"{path}: type differs: {type(actual).__name__} versus "
|
||||
f"{type(expected).__name__}"
|
||||
)
|
||||
if isinstance(expected, torch.Tensor):
|
||||
assert actual.shape == expected.shape, path
|
||||
assert actual.dtype == expected.dtype, path
|
||||
assert actual.stride() == expected.stride(), path
|
||||
if not torch.equal(actual, expected):
|
||||
difference = (actual.float() - expected.float()).abs()
|
||||
raise AssertionError(
|
||||
f"{path}: tensor values differ; mismatched="
|
||||
f"{torch.count_nonzero(actual != expected).item()}, "
|
||||
f"max_abs={difference.max().item()}"
|
||||
)
|
||||
return
|
||||
if isinstance(expected, dict):
|
||||
assert actual.keys() == expected.keys(), path
|
||||
for key in expected:
|
||||
_assert_exact(actual[key], expected[key], f"{path}.{key}")
|
||||
return
|
||||
if isinstance(expected, (list, tuple)):
|
||||
assert len(actual) == len(expected), path
|
||||
for index, (actual_value, expected_value) in enumerate(
|
||||
zip(actual, expected, strict=True)
|
||||
):
|
||||
_assert_exact(actual_value, expected_value, f"{path}[{index}]")
|
||||
return
|
||||
assert actual == expected, path
|
||||
|
||||
|
||||
def _load_rank_capture(output_dir: Path, rank: int) -> dict[str, Any]:
|
||||
"""Load one rank artifact and validate its fixed capture envelope."""
|
||||
artifact_path = output_dir / f"rank_{rank}.pt"
|
||||
if not artifact_path.is_file():
|
||||
raise AssertionError(f"Refiner parity capture is missing: {artifact_path}")
|
||||
artifact = torch.load(artifact_path, map_location="cpu", weights_only=True)
|
||||
assert artifact["schema_version"] == 1
|
||||
assert artifact["rank"] == rank
|
||||
assert artifact["world_size"] == WORLD_SIZE
|
||||
assert set(artifact) == {
|
||||
"schema_version",
|
||||
"implementation",
|
||||
"rank",
|
||||
"world_size",
|
||||
"case",
|
||||
}
|
||||
return artifact
|
||||
|
||||
|
||||
def test_magi2_refiner_transformer_matches_official_exactly() -> None:
|
||||
"""Match every refiner tensor boundary through video and audio depacking."""
|
||||
_require_parity_sources()
|
||||
official_output_dir = _run_implementation("official")
|
||||
fastvideo_output_dir = _run_implementation("fastvideo")
|
||||
for rank in range(WORLD_SIZE):
|
||||
official_capture = _load_rank_capture(official_output_dir, rank)
|
||||
fastvideo_capture = _load_rank_capture(fastvideo_output_dir, rank)
|
||||
assert official_capture["implementation"] == "official"
|
||||
assert fastvideo_capture["implementation"] == "fastvideo"
|
||||
_assert_exact(
|
||||
fastvideo_capture["case"],
|
||||
official_capture["case"],
|
||||
f"rank_{rank}.case",
|
||||
)
|
||||
@@ -0,0 +1,144 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU-only registry, preset, and model-index checks for MAGI-2 Preview."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.configs.pipelines import Magi2PreviewPipelineConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.registry import (
|
||||
_get_config_info,
|
||||
get_default_preset,
|
||||
get_model_family,
|
||||
get_model_info,
|
||||
get_pipeline_config_cls_from_name,
|
||||
get_registered_model_paths,
|
||||
)
|
||||
from scripts.checkpoint_conversion.convert_magi2_to_fastvideo import (
|
||||
MODEL_INDEX,
|
||||
REQUIRED_COMPONENT_FILES,
|
||||
SOURCE_REPOSITORY,
|
||||
convert_checkpoint_layout,
|
||||
)
|
||||
|
||||
|
||||
def _create_indexed_source(source: Path) -> Path:
|
||||
"""Create the smallest valid source layout with one shard per index."""
|
||||
indexed_shard = source / "preview" / "preview.safetensors"
|
||||
for component, relative_paths in REQUIRED_COMPONENT_FILES.items():
|
||||
for relative_path in relative_paths:
|
||||
source_path = source / component / relative_path
|
||||
source_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if relative_path.endswith(".index.json"):
|
||||
shard_name = f"{component}.safetensors"
|
||||
source_path.write_text(json.dumps({"weight_map": {"weight": shard_name}}), encoding="utf-8")
|
||||
shard_path = source_path.parent / shard_name
|
||||
shard_path.write_bytes(b"weights")
|
||||
if component == "preview":
|
||||
indexed_shard = shard_path
|
||||
else:
|
||||
source_path.write_bytes(b"component")
|
||||
return indexed_shard
|
||||
|
||||
|
||||
@pytest.mark.parametrize("workload_type", [WorkloadType.T2V, WorkloadType.I2V])
|
||||
def test_get_model_info_magi2_model_index_for_each_workload(
|
||||
tmp_path: Path,
|
||||
workload_type: WorkloadType,
|
||||
) -> None:
|
||||
"""Resolve the MAGI-2 pipeline class from local model-index metadata."""
|
||||
model_path = tmp_path / "checkpoint"
|
||||
(model_path / "transformer").mkdir(parents=True)
|
||||
(model_path / "model_index.json").write_text(
|
||||
json.dumps({
|
||||
"_class_name": "Magi2Pipeline",
|
||||
"_diffusers_version": "0.37.0",
|
||||
"transformer": [None, None],
|
||||
}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
model_info = get_model_info(str(model_path), workload_type=workload_type)
|
||||
config_info = _get_config_info(str(model_path))
|
||||
|
||||
assert model_info.pipeline_cls.__name__ == "Magi2Pipeline"
|
||||
assert model_info.pipeline_config_cls is Magi2PreviewPipelineConfig
|
||||
assert config_info is not None
|
||||
assert config_info.workload_types == (WorkloadType.T2V, WorkloadType.I2V)
|
||||
assert get_pipeline_config_cls_from_name(str(model_path)) is Magi2PreviewPipelineConfig
|
||||
assert get_model_family(str(model_path)) == "magi2"
|
||||
assert get_default_preset(str(model_path)) == "magi2_preview_1080p"
|
||||
|
||||
|
||||
def test_magi2_preview_pipeline_config_published_geometry() -> None:
|
||||
"""Expose the published preview and refiner latent geometry."""
|
||||
pipeline_config = Magi2PreviewPipelineConfig()
|
||||
|
||||
pipeline_config.check_pipeline_config()
|
||||
assert (pipeline_config.preview_width, pipeline_config.preview_height) == (896, 512)
|
||||
assert (pipeline_config.output_width, pipeline_config.output_height) == (1920, 1088)
|
||||
assert pipeline_config.output_frames == 249
|
||||
assert pipeline_config.output_fps == 25
|
||||
assert (pipeline_config.refiner_latent_width, pipeline_config.refiner_latent_height) == (120, 68)
|
||||
assert pipeline_config.text_encoder_configs == ()
|
||||
assert pipeline_config.text_encoder_precisions == ()
|
||||
|
||||
|
||||
def test_get_preset_magi2_matches_published_profile() -> None:
|
||||
"""Keep the registered preset aligned with the published inference profile."""
|
||||
preset = get_preset("magi2_preview_1080p", "magi2")
|
||||
|
||||
assert preset.defaults["seed"] == 42
|
||||
assert preset.defaults["height"] == 1088
|
||||
assert preset.defaults["width"] == 1920
|
||||
assert preset.defaults["num_frames"] == 249
|
||||
assert preset.defaults["fps"] == 25
|
||||
assert preset.defaults["num_inference_steps"] == 100
|
||||
assert preset.defaults["num_inference_steps_sr"] == 5
|
||||
negative_prompt = preset.defaults["negative_prompt"]
|
||||
assert len(negative_prompt) == 1721
|
||||
assert hashlib.sha256(negative_prompt.encode()).hexdigest() == (
|
||||
"5ac0746de6c7e0388a16122d9c7e751b1cb067242218bf6f41ec1e88271cdcb9"
|
||||
)
|
||||
|
||||
|
||||
def test_model_index_fastvideo_components_use_defining_module_paths() -> None:
|
||||
"""Record each FastVideo component with its defining Python module."""
|
||||
expected_libraries = {
|
||||
"audio_vae": "fastvideo.models.vaes.magi2_audio_vae",
|
||||
"image_encoder": "fastvideo.models.vaes.magi2_wan_loader",
|
||||
"scheduler": "fastvideo.models.schedulers.scheduling_flow_unipc_multistep",
|
||||
"text_encoder": "fastvideo.models.encoders.qwen3_5",
|
||||
"transformer": "fastvideo.models.dits.magi2",
|
||||
"transformer_2": "fastvideo.models.dits.magi2_refiner",
|
||||
"vae": "fastvideo.models.vaes.magi2_turbo_vae",
|
||||
}
|
||||
|
||||
assert {
|
||||
component: MODEL_INDEX[component][0]
|
||||
for component in expected_libraries
|
||||
} == expected_libraries
|
||||
|
||||
|
||||
def test_source_repository_remains_converter_provenance() -> None:
|
||||
"""Keep the official checkpoint ID as converter provenance metadata."""
|
||||
assert SOURCE_REPOSITORY == "sand-ai/MAGI-2-preview"
|
||||
assert SOURCE_REPOSITORY not in get_registered_model_paths()
|
||||
|
||||
|
||||
def test_convert_checkpoint_layout_missing_indexed_shard(tmp_path: Path) -> None:
|
||||
"""Reject a source snapshot when a weight index references a missing shard."""
|
||||
source = tmp_path / "official"
|
||||
indexed_shard = _create_indexed_source(source)
|
||||
indexed_shard.unlink()
|
||||
output = tmp_path / "fastvideo"
|
||||
|
||||
with pytest.raises(FileNotFoundError, match="missing indexed checkpoint shards"):
|
||||
convert_checkpoint_layout(source, output)
|
||||
assert not output.exists()
|
||||
@@ -0,0 +1,163 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Public runtime-control coverage for MAGI-2 inference."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.api.compat import (
|
||||
generator_config_to_fastvideo_args,
|
||||
legacy_from_pretrained_to_config,
|
||||
)
|
||||
from fastvideo.api.schema import EngineConfig, GeneratorConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.magi2.magi2_pipeline import (
|
||||
_configure_deterministic_kernels,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi2.stages import output as magi2_output_stage
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
def test_deterministic_cli_flag_is_enabled_without_a_value() -> None:
|
||||
"""Expose ``--deterministic`` through the legacy FastVideo CLI."""
|
||||
parser = FlexibleArgumentParser()
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
arguments = parser.parse_args(
|
||||
["--model-path", "/models/magi2", "--deterministic"]
|
||||
)
|
||||
assert arguments.deterministic is True
|
||||
|
||||
|
||||
def test_deterministic_control_round_trips_through_typed_config(monkeypatch) -> None:
|
||||
"""Preserve deterministic mode across legacy and typed engine adapters."""
|
||||
typed_config = legacy_from_pretrained_to_config(
|
||||
"/models/magi2",
|
||||
{"deterministic": True},
|
||||
)
|
||||
assert typed_config.engine.deterministic is True
|
||||
|
||||
captured_kwargs: dict = {}
|
||||
|
||||
def capture_fastvideo_args(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return kwargs
|
||||
|
||||
monkeypatch.setattr(FastVideoArgs, "from_kwargs", capture_fastvideo_args)
|
||||
generator_config_to_fastvideo_args(
|
||||
GeneratorConfig(
|
||||
model_path="/models/magi2",
|
||||
engine=EngineConfig(deterministic=True),
|
||||
)
|
||||
)
|
||||
assert captured_kwargs["deterministic"] is True
|
||||
|
||||
|
||||
def test_magi2_environment_controls_are_lazy(monkeypatch) -> None:
|
||||
"""Read deterministic and latent-capture values from each worker's environment."""
|
||||
monkeypatch.setenv("MAGI2_DETERMINISTIC", "1")
|
||||
monkeypatch.setenv("MAGI_ATTENTION_DETERMINISTIC_MODE", "1")
|
||||
monkeypatch.setenv("MAGI2_SAVE_LATENT_PATH", "/tmp/magi2-latents")
|
||||
|
||||
assert envs.MAGI2_DETERMINISTIC is True
|
||||
assert envs.MAGI_ATTENTION_DETERMINISTIC_MODE is True
|
||||
assert envs.MAGI2_SAVE_LATENT_PATH == "/tmp/magi2-latents"
|
||||
|
||||
|
||||
def test_deterministic_kernel_configuration_repeats_cpu_rng_sequences(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
"""Seed Python, NumPy, and CPU PyTorch with repeatable sequences."""
|
||||
python_rng_state = random.getstate()
|
||||
numpy_rng_state = np.random.get_state()
|
||||
torch_rng_state = torch.get_rng_state()
|
||||
deterministic_algorithms = torch.are_deterministic_algorithms_enabled()
|
||||
deterministic_warn_only = torch.is_deterministic_algorithms_warn_only_enabled()
|
||||
cuda_seed_calls: list[int] = []
|
||||
monkeypatch.setattr(torch.cuda, "manual_seed_all", cuda_seed_calls.append)
|
||||
monkeypatch.delenv("CUBLAS_WORKSPACE_CONFIG", raising=False)
|
||||
monkeypatch.delenv("MAGI2_DETERMINISTIC", raising=False)
|
||||
monkeypatch.delenv("MAGI_ATTENTION_DETERMINISTIC_MODE", raising=False)
|
||||
|
||||
try:
|
||||
_configure_deterministic_kernels(seed=1234)
|
||||
first_sequence = (random.random(), np.random.random(4), torch.rand(4))
|
||||
_configure_deterministic_kernels(seed=1234)
|
||||
second_sequence = (random.random(), np.random.random(4), torch.rand(4))
|
||||
|
||||
assert first_sequence[0] == second_sequence[0]
|
||||
assert np.array_equal(first_sequence[1], second_sequence[1])
|
||||
assert torch.equal(first_sequence[2], second_sequence[2])
|
||||
assert cuda_seed_calls and set(cuda_seed_calls) == {1234}
|
||||
assert torch.are_deterministic_algorithms_enabled()
|
||||
assert envs.MAGI2_DETERMINISTIC is True
|
||||
assert envs.MAGI_ATTENTION_DETERMINISTIC_MODE is True
|
||||
assert os.environ["CUBLAS_WORKSPACE_CONFIG"] == ":4096:8"
|
||||
finally:
|
||||
random.setstate(python_rng_state)
|
||||
np.random.set_state(numpy_rng_state)
|
||||
torch.set_rng_state(torch_rng_state)
|
||||
torch.use_deterministic_algorithms(
|
||||
deterministic_algorithms,
|
||||
warn_only=deterministic_warn_only,
|
||||
)
|
||||
|
||||
|
||||
def test_latent_saving_stage_writes_leader_latent(tmp_path, monkeypatch) -> None:
|
||||
"""Save the post-refiner latent on the context-parallel leader rank."""
|
||||
latent_directory = tmp_path / "latents"
|
||||
monkeypatch.setenv("MAGI2_SAVE_LATENT_PATH", str(latent_directory))
|
||||
monkeypatch.setattr(
|
||||
magi2_output_stage.psm,
|
||||
"is_group_first_rank",
|
||||
lambda dimension: dimension == "cp",
|
||||
)
|
||||
stage = magi2_output_stage.Magi2LatentSavingStage()
|
||||
latent = torch.arange(12, dtype=torch.float32).reshape(1, 3, 4)
|
||||
batch = ForwardBatch(data_type="video", latents=latent)
|
||||
|
||||
returned_batch = stage.forward(batch, cast(FastVideoArgs, object()))
|
||||
|
||||
assert returned_batch is batch
|
||||
assert torch.equal(
|
||||
torch.load(latent_directory / "latent_0.pt", weights_only=True),
|
||||
latent,
|
||||
)
|
||||
assert stage.sample_index == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("configured_directory", "is_leader"),
|
||||
[("", True), ("nonleader-latents", False)],
|
||||
ids=["empty-path", "non-leader"],
|
||||
)
|
||||
def test_latent_saving_stage_skips_disabled_ranks(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
configured_directory: str,
|
||||
is_leader: bool,
|
||||
) -> None:
|
||||
"""Skip latent writes for an empty path and for non-leader ranks."""
|
||||
monkeypatch.chdir(tmp_path)
|
||||
latent_directory = "" if configured_directory == "" else str(tmp_path / configured_directory)
|
||||
monkeypatch.setenv("MAGI2_SAVE_LATENT_PATH", latent_directory)
|
||||
monkeypatch.setattr(
|
||||
magi2_output_stage.psm,
|
||||
"is_group_first_rank",
|
||||
lambda dimension: is_leader,
|
||||
)
|
||||
stage = magi2_output_stage.Magi2LatentSavingStage()
|
||||
batch = ForwardBatch(data_type="video", latents=torch.ones(1))
|
||||
|
||||
returned_batch = stage.forward(batch, cast(FastVideoArgs, object()))
|
||||
|
||||
assert returned_batch is batch
|
||||
assert list(tmp_path.rglob("latent_*.pt")) == []
|
||||
assert stage.sample_index == 0
|
||||
@@ -0,0 +1,183 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict MAGI-2 Flow UniPC scheduler parity against the official source."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import math
|
||||
import typing
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
try:
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import (
|
||||
KarrasDiffusionSchedulers,
|
||||
SchedulerMixin,
|
||||
SchedulerOutput,
|
||||
)
|
||||
from diffusers.utils import deprecate
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler as FastVideoFlowUniPCMultistepScheduler,
|
||||
)
|
||||
except ModuleNotFoundError as exc:
|
||||
if exc.name != "diffusers":
|
||||
raise
|
||||
pytest.skip(
|
||||
"MAGI-2 scheduler parity requires the Diffusers package.",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_SCHEDULER_PATH = (
|
||||
REPO_ROOT.parent
|
||||
/ "MAGI-2-preview"
|
||||
/ "inference"
|
||||
/ "pipeline"
|
||||
/ "sampler.py"
|
||||
)
|
||||
OFFICIAL_SCHEDULER_CLASS_NAME = "FlowUniPCMultistepScheduler"
|
||||
|
||||
|
||||
def _load_official_scheduler_class() -> type:
|
||||
"""Load only the official scheduler class without model-runtime imports."""
|
||||
if not OFFICIAL_SCHEDULER_PATH.is_file():
|
||||
raise FileNotFoundError(
|
||||
f"Official MAGI-2 scheduler source is missing: {OFFICIAL_SCHEDULER_PATH}"
|
||||
)
|
||||
|
||||
source = OFFICIAL_SCHEDULER_PATH.read_text(encoding="utf-8")
|
||||
source_module = ast.parse(source, filename=str(OFFICIAL_SCHEDULER_PATH))
|
||||
scheduler_nodes = [
|
||||
node
|
||||
for node in source_module.body
|
||||
if isinstance(node, ast.ClassDef)
|
||||
and node.name == OFFICIAL_SCHEDULER_CLASS_NAME
|
||||
]
|
||||
if len(scheduler_nodes) != 1:
|
||||
raise AssertionError(
|
||||
f"Expected one {OFFICIAL_SCHEDULER_CLASS_NAME} definition in "
|
||||
f"{OFFICIAL_SCHEDULER_PATH}, found {len(scheduler_nodes)}"
|
||||
)
|
||||
|
||||
scheduler_namespace = {
|
||||
"__name__": "magi2_official_scheduler",
|
||||
"Any": typing.Any,
|
||||
"ConfigMixin": ConfigMixin,
|
||||
"KarrasDiffusionSchedulers": KarrasDiffusionSchedulers,
|
||||
"List": list,
|
||||
"Optional": typing.Optional,
|
||||
"SchedulerMixin": SchedulerMixin,
|
||||
"SchedulerOutput": SchedulerOutput,
|
||||
"Tuple": tuple,
|
||||
"Union": typing.Union,
|
||||
"deprecate": deprecate,
|
||||
"math": math,
|
||||
"np": np,
|
||||
"register_to_config": register_to_config,
|
||||
"torch": torch,
|
||||
}
|
||||
scheduler_module = ast.Module(body=scheduler_nodes, type_ignores=[])
|
||||
exec(
|
||||
compile(scheduler_module, str(OFFICIAL_SCHEDULER_PATH), "exec"),
|
||||
scheduler_namespace,
|
||||
)
|
||||
return scheduler_namespace[OFFICIAL_SCHEDULER_CLASS_NAME]
|
||||
|
||||
|
||||
OfficialFlowUniPCMultistepScheduler = _load_official_scheduler_class()
|
||||
|
||||
|
||||
def _build_scheduler_pair(
|
||||
num_inference_steps: int,
|
||||
shift: float,
|
||||
) -> tuple[SchedulerMixin, SchedulerMixin]:
|
||||
"""Construct official and FastVideo schedulers with a shipping schedule."""
|
||||
official_scheduler = OfficialFlowUniPCMultistepScheduler()
|
||||
fastvideo_scheduler = FastVideoFlowUniPCMultistepScheduler()
|
||||
official_scheduler.set_timesteps(num_inference_steps, device="cpu", shift=shift)
|
||||
fastvideo_scheduler.set_timesteps(num_inference_steps, device="cpu", shift=shift)
|
||||
return official_scheduler, fastvideo_scheduler
|
||||
|
||||
|
||||
def _assert_scheduler_tensor_exact(
|
||||
fastvideo_tensor: torch.Tensor,
|
||||
official_tensor: torch.Tensor,
|
||||
) -> None:
|
||||
"""Require identical scheduler tensor metadata and values."""
|
||||
assert fastvideo_tensor.shape == official_tensor.shape
|
||||
assert fastvideo_tensor.dtype == official_tensor.dtype
|
||||
assert fastvideo_tensor.stride() == official_tensor.stride()
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("num_inference_steps", "shift"),
|
||||
[(100, 7.0), (5, 5.0)],
|
||||
ids=("preview", "refiner"),
|
||||
)
|
||||
def test_set_timesteps_magi2_shipping_schedules_match_official(
|
||||
num_inference_steps: int,
|
||||
shift: float,
|
||||
) -> None:
|
||||
"""Match the preview and refiner timestep and sigma schedules exactly."""
|
||||
official_scheduler, fastvideo_scheduler = _build_scheduler_pair(
|
||||
num_inference_steps,
|
||||
shift,
|
||||
)
|
||||
|
||||
_assert_scheduler_tensor_exact(
|
||||
fastvideo_scheduler.timesteps,
|
||||
official_scheduler.timesteps,
|
||||
)
|
||||
_assert_scheduler_tensor_exact(
|
||||
fastvideo_scheduler.sigmas,
|
||||
official_scheduler.sigmas,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("num_inference_steps", "shift"),
|
||||
[(100, 7.0), (5, 5.0)],
|
||||
ids=("preview", "refiner"),
|
||||
)
|
||||
def test_step_magi2_shipping_denoise_trajectory_matches_official(
|
||||
num_inference_steps: int,
|
||||
shift: float,
|
||||
) -> None:
|
||||
"""Match every scheduler output across each shipping denoise schedule."""
|
||||
official_scheduler, fastvideo_scheduler = _build_scheduler_pair(
|
||||
num_inference_steps,
|
||||
shift,
|
||||
)
|
||||
generator = torch.Generator(device="cpu").manual_seed(20260805)
|
||||
tensor_shape = (1, 2, 3, 4, 4)
|
||||
official_sample = torch.randn(
|
||||
tensor_shape,
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
fastvideo_sample = official_sample.clone()
|
||||
model_outputs = [
|
||||
torch.randn(tensor_shape, generator=generator, dtype=torch.float32)
|
||||
for _ in range(num_inference_steps)
|
||||
]
|
||||
|
||||
for step_index, model_output in enumerate(model_outputs):
|
||||
official_sample = official_scheduler.step(
|
||||
model_output,
|
||||
official_scheduler.timesteps[step_index],
|
||||
official_sample,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
fastvideo_sample = fastvideo_scheduler.step(
|
||||
model_output,
|
||||
fastvideo_scheduler.timesteps[step_index],
|
||||
fastvideo_sample,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
_assert_scheduler_tensor_exact(fastvideo_sample, official_sample)
|
||||
@@ -0,0 +1,218 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict prompt and Qwen3.5 encoder parity for MAGI-2 Preview.
|
||||
|
||||
Coverage scope: both. The official side uses
|
||||
``inference.pipeline.inference_engine.initialize_text_encoder`` with the
|
||||
published Qwen3.5 checkpoint. The FastVideo side targets
|
||||
``fastvideo.models.encoders.qwen3_5.Magi2Qwen35TextEncoder`` with the same checkpoint.
|
||||
The comparison covers structured-prompt normalization, CJK token splitting,
|
||||
token IDs, attention masks, the skip-layer-2 hidden state, and ``<Figure 1>``
|
||||
token pooling.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import json
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.models.encoders.qwen3_5 import (
|
||||
Magi2Qwen35TextEncoder,
|
||||
json_to_compact_markdown,
|
||||
)
|
||||
from tests.local_tests.magi2._parity_utils import (
|
||||
LOCAL_WEIGHTS_DIR,
|
||||
OFFICIAL_REF_DIR,
|
||||
assert_tensor_exact,
|
||||
import_official_module,
|
||||
require_complete_safetensor_index,
|
||||
)
|
||||
|
||||
|
||||
PARITY_COVERAGE = "both"
|
||||
TEXT_ENCODER_DIR = LOCAL_WEIGHTS_DIR / "text_encoder"
|
||||
|
||||
|
||||
def _structured_i2v_prompt() -> str:
|
||||
"""Return one deterministic prompt that exercises JSON, CJK, and figure tokens."""
|
||||
prompt = {
|
||||
"global_layer": {
|
||||
"context": "夜晚的城市街道",
|
||||
"description": "A cyclist passes a quiet café.",
|
||||
"aesthetics": {
|
||||
"style": "documentary",
|
||||
"mood_atmosphere": "calm",
|
||||
"color_scheme": "blue and amber",
|
||||
},
|
||||
},
|
||||
"reference_layer": ["The first frame refers to <Figure 1>"],
|
||||
}
|
||||
return json.dumps(prompt, ensure_ascii=False)
|
||||
|
||||
|
||||
def _tokenize(
|
||||
tokenizer,
|
||||
max_length: int,
|
||||
normalized_prompt: str,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Tokenize with the exact keyword arguments used by the official encoder."""
|
||||
tokens = tokenizer(
|
||||
[normalized_prompt],
|
||||
return_tensors="pt",
|
||||
padding="longest",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
)
|
||||
return tokens["input_ids"], tokens["attention_mask"]
|
||||
|
||||
|
||||
def _load_official_encoder(model_path: str, device: torch.device):
|
||||
"""Instantiate Qwen3.5 through the official component implementation."""
|
||||
qwen35_module = import_official_module("inference.model.qwen35")
|
||||
return qwen35_module.Qwen35TextEncoder(
|
||||
model_path=model_path,
|
||||
device=str(device),
|
||||
precision=torch.bfloat16,
|
||||
skip_layer=2,
|
||||
)
|
||||
|
||||
|
||||
def test_iter_indexed_language_model_weights_reads_only_language_model_tensors(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Read only indexed language-model keys and ignore unrelated missing shards."""
|
||||
from fastvideo.models.encoders import qwen3_5 as qwen35_module
|
||||
|
||||
language_shards = ["language-a.safetensors", "language-b.safetensors"]
|
||||
for shard_name in language_shards:
|
||||
(tmp_path / shard_name).touch()
|
||||
weight_map = {
|
||||
"model.language_model.layers.0.weight": language_shards[1],
|
||||
"model.visual.weight": "missing-visual.safetensors",
|
||||
"model.language_model.embed_tokens.weight": language_shards[0],
|
||||
"lm_head.weight": "missing-output.safetensors",
|
||||
}
|
||||
tensors_by_shard = {
|
||||
language_shards[0]: {
|
||||
"model.language_model.embed_tokens.weight": torch.tensor([1.0]),
|
||||
},
|
||||
language_shards[1]: {
|
||||
"model.language_model.layers.0.weight": torch.tensor([2.0]),
|
||||
},
|
||||
}
|
||||
opened_shards: list[str] = []
|
||||
read_requests: list[tuple[str, str]] = []
|
||||
|
||||
@contextmanager
|
||||
def recording_safe_open(filename: str, *, framework: str, device: str):
|
||||
"""Record each shard and reject tensor reads from the wrong shard."""
|
||||
shard_name = Path(filename).name
|
||||
opened_shards.append(shard_name)
|
||||
assert framework == "pt"
|
||||
assert device == "cpu"
|
||||
|
||||
def read_tensor(checkpoint_name: str) -> torch.Tensor:
|
||||
read_requests.append((shard_name, checkpoint_name))
|
||||
return tensors_by_shard[shard_name][checkpoint_name]
|
||||
|
||||
yield SimpleNamespace(get_tensor=read_tensor)
|
||||
|
||||
monkeypatch.setattr(qwen35_module, "safe_open", recording_safe_open)
|
||||
loaded_weights = list(
|
||||
qwen35_module._iter_indexed_language_model_weights(
|
||||
tmp_path,
|
||||
weight_map,
|
||||
torch.device("cpu"),
|
||||
)
|
||||
)
|
||||
|
||||
assert opened_shards == language_shards
|
||||
assert read_requests == [
|
||||
(language_shards[0], "model.language_model.embed_tokens.weight"),
|
||||
(language_shards[1], "model.language_model.layers.0.weight"),
|
||||
]
|
||||
assert [name for name, _ in loaded_weights] == [name for _, name in read_requests]
|
||||
assert [tensor.item() for _, tensor in loaded_weights] == [1.0, 2.0]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="MAGI-2 Qwen3.5 parity requires CUDA.",
|
||||
)
|
||||
def test_magi2_text_encoder_structured_i2v_prompt_exact_parity() -> None:
|
||||
"""Require exact tokenizer, hidden-state, and figure-token outputs."""
|
||||
model_dir = require_complete_safetensor_index(TEXT_ENCODER_DIR)
|
||||
device = torch.device("cuda:0")
|
||||
prompt = _structured_i2v_prompt()
|
||||
|
||||
official_encoder = _load_official_encoder(
|
||||
str(model_dir),
|
||||
torch.device("cpu"),
|
||||
).to(device)
|
||||
official_normalized_prompt = official_encoder._normalize_prompt(prompt)
|
||||
official_input_ids, official_attention_mask = _tokenize(
|
||||
official_encoder.tokenizer,
|
||||
official_encoder.max_length,
|
||||
official_normalized_prompt,
|
||||
)
|
||||
official_embedding = official_encoder.encode(prompt).detach().cpu()
|
||||
release_prompt = (
|
||||
OFFICIAL_REF_DIR / "assets" / "sample_enhanced_t2v.json"
|
||||
).read_text(encoding="utf-8").strip()
|
||||
official_release_embedding = official_encoder.encode(release_prompt).detach().cpu()
|
||||
official_figure_embedding = official_encoder.get_special_token(
|
||||
prompt,
|
||||
["<Figure 1>"],
|
||||
official_embedding,
|
||||
)
|
||||
|
||||
del official_encoder
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
from fastvideo.configs.models.encoders.qwen3_5 import Magi2Qwen35Config
|
||||
fastvideo_encoder = Magi2Qwen35TextEncoder.from_pretrained_local(
|
||||
model_path=str(model_dir),
|
||||
model_config=Magi2Qwen35Config(),
|
||||
dtype=torch.bfloat16,
|
||||
device=torch.device("cpu"),
|
||||
).to(device)
|
||||
fastvideo_normalized_prompt = json_to_compact_markdown(prompt)
|
||||
fastvideo_input_ids, fastvideo_attention_mask = _tokenize(
|
||||
fastvideo_encoder.tokenizer,
|
||||
fastvideo_encoder.config.text_len,
|
||||
fastvideo_normalized_prompt,
|
||||
)
|
||||
fastvideo_embedding = fastvideo_encoder.encode(prompt).detach().cpu()
|
||||
fastvideo_release_embedding = fastvideo_encoder.encode(
|
||||
release_prompt
|
||||
).detach().cpu()
|
||||
fastvideo_figure_embedding = fastvideo_encoder.get_special_token(
|
||||
prompt,
|
||||
["<Figure 1>"],
|
||||
fastvideo_embedding,
|
||||
)
|
||||
|
||||
assert fastvideo_normalized_prompt == official_normalized_prompt
|
||||
assert_tensor_exact(fastvideo_input_ids, official_input_ids, "token IDs")
|
||||
assert_tensor_exact(
|
||||
fastvideo_attention_mask,
|
||||
official_attention_mask,
|
||||
"attention mask",
|
||||
)
|
||||
assert_tensor_exact(fastvideo_embedding, official_embedding, "skip-layer-2 hidden state")
|
||||
assert_tensor_exact(
|
||||
fastvideo_release_embedding,
|
||||
official_release_embedding,
|
||||
"release-prompt skip-layer-2 hidden state",
|
||||
)
|
||||
assert_tensor_exact(
|
||||
fastvideo_figure_embedding,
|
||||
official_figure_embedding,
|
||||
"<Figure 1> pooled embedding",
|
||||
)
|
||||
@@ -0,0 +1,121 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict MAGI-2 preview and refiner state-structure parity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_ROOT = Path(os.environ.get("MAGI2_OFFICIAL_ROOT", REPO_ROOT.parent / "MAGI-2-preview"))
|
||||
WEIGHTS_ROOT = Path(os.environ.get("MAGI2_WEIGHTS_ROOT", REPO_ROOT / "official_weights" / "magi2"))
|
||||
|
||||
if not (OFFICIAL_ROOT / "inference" / "model" / "magi2_preview.py").is_file():
|
||||
pytest.skip(f"Official MAGI-2 checkout is missing: {OFFICIAL_ROOT}", allow_module_level=True)
|
||||
if not (WEIGHTS_ROOT / "preview" / "model.safetensors.index.json").is_file():
|
||||
pytest.skip(f"Official MAGI-2 weights are missing: {WEIGHTS_ROOT}", allow_module_level=True)
|
||||
|
||||
|
||||
def _structure_dump_script(component: str, implementation: str) -> str:
|
||||
"""Create an isolated process script that instantiates one model on meta."""
|
||||
if implementation == "official":
|
||||
model_setup = {
|
||||
"preview": """
|
||||
from inference.common.magi2_config import load_config
|
||||
from inference.model.magi2_preview import Transformer
|
||||
config = load_config(str(official_root / "configs" / "magi2_preview.json"))
|
||||
with torch.device("meta"):
|
||||
model = Transformer(config.arch_config, ep_size=1)
|
||||
""",
|
||||
"refiner": """
|
||||
from inference.common.magi2_config import load_config
|
||||
from inference.model.magi2_refiner import Transformer
|
||||
config = load_config(str(official_root / "configs" / "magi2_refiner.json"))
|
||||
with torch.device("meta"):
|
||||
model = Transformer(config.magi2_refiner_arch_config)
|
||||
""",
|
||||
}[component]
|
||||
path_setup = "sys.path.insert(0, str(official_root))"
|
||||
else:
|
||||
model_setup = {
|
||||
"preview": """
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2PreviewVideoConfig
|
||||
from fastvideo.models.dits.magi2 import Magi2PreviewDiT
|
||||
with torch.device("meta"):
|
||||
model = Magi2PreviewDiT(config=Magi2PreviewVideoConfig())
|
||||
""",
|
||||
"refiner": """
|
||||
from fastvideo.configs.models.dits.magi2 import Magi2RefinerVideoConfig
|
||||
from fastvideo.models.dits.magi2_refiner import Magi2RefinerDiT
|
||||
with torch.device("meta"):
|
||||
model = Magi2RefinerDiT(config=Magi2RefinerVideoConfig())
|
||||
""",
|
||||
}[component]
|
||||
path_setup = "sys.path.insert(0, str(repo_root))"
|
||||
|
||||
setup_script = textwrap.dedent(
|
||||
f"""
|
||||
import json
|
||||
import os
|
||||
import pathlib
|
||||
import sys
|
||||
import torch
|
||||
|
||||
repo_root = pathlib.Path({str(REPO_ROOT)!r})
|
||||
official_root = pathlib.Path({str(OFFICIAL_ROOT)!r})
|
||||
{path_setup}
|
||||
os.environ["SKIP_LOAD_MODEL"] = "1"
|
||||
os.environ["MAGI2_DISABLE_MAGI_COMPILE"] = "1"
|
||||
os.environ["MAGI_COMPILE_COMPILE_MODE"] = "NONE"
|
||||
"""
|
||||
)
|
||||
dump_script = textwrap.dedent(
|
||||
"""
|
||||
structure = {
|
||||
name: {"shape": list(tensor.shape), "dtype": str(tensor.dtype)}
|
||||
for name, tensor in model.state_dict().items()
|
||||
}
|
||||
print("MAGI2_STRUCTURE=" + json.dumps(structure, sort_keys=True))
|
||||
"""
|
||||
)
|
||||
return setup_script + textwrap.dedent(model_setup) + dump_script
|
||||
|
||||
|
||||
def _dump_structure(component: str, implementation: str) -> dict[str, dict]:
|
||||
"""Run one model definition in isolation and parse its state metadata."""
|
||||
completed = subprocess.run(
|
||||
[sys.executable, "-c", _structure_dump_script(component, implementation)],
|
||||
cwd=REPO_ROOT,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300,
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
raise AssertionError(
|
||||
f"{implementation} {component} structure dump failed.\n"
|
||||
f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}"
|
||||
)
|
||||
lines = [line for line in completed.stdout.splitlines() if line.startswith("MAGI2_STRUCTURE=")]
|
||||
if len(lines) != 1:
|
||||
raise AssertionError(f"Expected one structure record, received stdout:\n{completed.stdout}")
|
||||
return json.loads(lines[0].removeprefix("MAGI2_STRUCTURE="))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("component", ["preview", "refiner"])
|
||||
def test_magi2_transformer_state_structure_matches_official(component: str) -> None:
|
||||
"""Match every state name, shape, and dtype to the official definition."""
|
||||
official_structure = _dump_structure(component, "official")
|
||||
fastvideo_structure = _dump_structure(component, "fastvideo")
|
||||
assert fastvideo_structure == official_structure
|
||||
|
||||
index_path = WEIGHTS_ROOT / component / "model.safetensors.index.json"
|
||||
checkpoint_keys = set(json.loads(index_path.read_text(encoding="utf-8"))["weight_map"])
|
||||
assert set(fastvideo_structure) == checkpoint_keys
|
||||
@@ -0,0 +1,115 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict temporal sliding-window parity for the MAGI-2 Turbo VAE decoder.
|
||||
|
||||
Coverage scope: both. The official side uses
|
||||
``inference.model.turbo_vaed.get_turbo_vaed`` with the published distilled
|
||||
decoder. The FastVideo side targets
|
||||
``fastvideo.models.vaes.magi2_turbo_vae.Magi2TurboVAEModel``. Fifteen latent
|
||||
frames exercise the published first, middle, and last temporal windows after
|
||||
padding from 15 to 21 frames.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.models.vaes.magi2_turbo_vae import Magi2TurboVAEModel
|
||||
from tests.local_tests.magi2._parity_utils import (
|
||||
LOCAL_WEIGHTS_DIR,
|
||||
assert_tensor_exact,
|
||||
import_official_module,
|
||||
require_path,
|
||||
)
|
||||
|
||||
|
||||
PARITY_COVERAGE = "both"
|
||||
TURBO_VAE_DIR = LOCAL_WEIGHTS_DIR / "turbo_vae"
|
||||
TURBO_VAE_CONFIG_PATH = TURBO_VAE_DIR / "TurboV3-Wan22-TinyShallow_7_7.json"
|
||||
TURBO_VAE_CHECKPOINT_PATH = TURBO_VAE_DIR / "checkpoint.ckpt"
|
||||
|
||||
|
||||
def _load_turbo_config() -> dict:
|
||||
"""Load the exact configuration paired with the distilled checkpoint."""
|
||||
config_path = require_path(TURBO_VAE_CONFIG_PATH, "Turbo VAE configuration")
|
||||
with config_path.open(encoding="utf-8") as config_file:
|
||||
return json.load(config_file)
|
||||
|
||||
|
||||
def _deterministic_video_latent(device: torch.device) -> torch.Tensor:
|
||||
"""Create a reproducible BF16 latent that crosses two window boundaries."""
|
||||
latent_values = torch.arange(1 * 48 * 15 * 4 * 4, dtype=torch.float32)
|
||||
latent_values = ((latent_values % 257) - 128) / 128
|
||||
return latent_values.reshape(1, 48, 15, 4, 4).to(
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
|
||||
def _temporal_window_roles(
|
||||
num_frames: int,
|
||||
first_chunk_size: int,
|
||||
step_size: int,
|
||||
) -> tuple[str, ...]:
|
||||
"""Return the Turbo VAE window roles selected after temporal padding."""
|
||||
padding_frames = 0
|
||||
if num_frames < first_chunk_size:
|
||||
padding_frames = first_chunk_size - num_frames
|
||||
elif (num_frames - first_chunk_size) % step_size != 0:
|
||||
padding_frames = step_size - (num_frames - first_chunk_size) % step_size
|
||||
padded_frames = num_frames + padding_frames
|
||||
if padded_frames == first_chunk_size:
|
||||
return ("single",)
|
||||
|
||||
window_roles = ["first"]
|
||||
for window_start in range(first_chunk_size, padded_frames, step_size):
|
||||
is_last_window = window_start + step_size == padded_frames
|
||||
window_roles.append("last" if is_last_window else "middle")
|
||||
return tuple(window_roles)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="MAGI-2 Turbo VAE parity requires CUDA.",
|
||||
)
|
||||
def test_magi2_turbo_vae_sliding_window_decode_exact_parity() -> None:
|
||||
"""Require exact decoded video values across first, middle, and last windows."""
|
||||
checkpoint_path = require_path(
|
||||
TURBO_VAE_CHECKPOINT_PATH,
|
||||
"Turbo VAE checkpoint",
|
||||
)
|
||||
config = _load_turbo_config()
|
||||
assert config["first_chunk_size"] == 7
|
||||
assert config["step_size"] == 7
|
||||
assert config["temporal_compression_ratio"] == 4
|
||||
|
||||
official_module = import_official_module("inference.model.turbo_vaed")
|
||||
device = torch.device("cuda:0")
|
||||
official_vae = official_module.get_turbo_vaed(
|
||||
str(TURBO_VAE_CONFIG_PATH),
|
||||
str(checkpoint_path),
|
||||
device=str(device),
|
||||
weight_dtype=torch.bfloat16,
|
||||
)
|
||||
from fastvideo.configs.models.vaes.magi2_turbo_vae import Magi2TurboVAEConfig
|
||||
|
||||
fastvideo_config = Magi2TurboVAEConfig(
|
||||
config_path=str(TURBO_VAE_CONFIG_PATH),
|
||||
checkpoint_path=str(checkpoint_path),
|
||||
pretrained_dtype="bfloat16",
|
||||
)
|
||||
fastvideo_vae = Magi2TurboVAEModel(fastvideo_config).eval()
|
||||
latent = _deterministic_video_latent(device)
|
||||
assert _temporal_window_roles(
|
||||
latent.shape[2],
|
||||
config["first_chunk_size"],
|
||||
config["step_size"],
|
||||
) == ("first", "middle", "last")
|
||||
|
||||
with torch.inference_mode():
|
||||
official_video = official_vae.decode(latent).float().detach().cpu()
|
||||
fastvideo_video = fastvideo_vae.decode(latent).float().detach().cpu()
|
||||
|
||||
assert_tensor_exact(fastvideo_video, official_video, "Turbo VAE decoded video")
|
||||
Reference in New Issue
Block a user