Compare commits

...
59 changed files with 6266 additions and 1124 deletions
+1
View File
@@ -23,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
outputs_audio/
outputs_video
checkpoints/
sbatch.sh
+5
View File
@@ -64,6 +64,7 @@ column links a runnable script in `examples/inference/basic/` where one exists.
| ltx2 | `FastVideo/LTX2-Distilled-Diffusers`<br>`FastVideo/LTX2.3-Distilled-Diffusers`<br>`FastVideo/LTX-2.3-Distilled-Diffusers` | T2V | [basic_ltx2_distilled.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2_distilled.py) |
| ltx2 | `Lightricks/LTX-2.3`<br>`FastVideo/LTX2.3-base`<br>`FastVideo/LTX2.3-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
| ltx2 | `Lightricks/LTX-2`<br>`FastVideo/LTX2-base`<br>`FastVideo/LTX2-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
| mmaudio | `FastVideo/MMAudio-large-44k-v2-Diffusers` | V2A, T2A | [basic_mmaudio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mmaudio.py) |
| matrixgame | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-Base-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Diffusers`<br>`mignonjia/mg_longtuning_distilled_zelda`<br>`mignonjia/mg_sf_distilled_zelda_1k_steps`<br>`mignonjia/mg_sf_distilled_zelda`<br>`mignonjia/mg_causal_zelda`<br>`mignonjia/mg_bidirectional_zelda` | I2V | [basic_matrixgame2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame2.py) |
| matrixgame | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | I2V | [basic_matrixgame3.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame3.py) |
| minimax_h3 | `MiniMaxAI/MiniMax-H3` | T2V, I2V | [T2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_t2v.py)<br>[FL2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_fl2va.py)<br>[Ref2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_ref2va.py) |
@@ -94,6 +95,10 @@ column links a runnable script in `examples/inference/basic/` where one exists.
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
under the generic T2V workload option in the registry.
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
audio. Use the Ref2VA example when passing ordered image, video, or audio
references.
+44
View File
@@ -0,0 +1,44 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio large-44k-v2 video-to-audio example."""
import argparse
import os
from fastvideo import VideoGenerator
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--video-path", required=True)
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
parser.add_argument("--duration-seconds", type=float, default=8.0)
parser.add_argument("--prompt", default="")
parser.add_argument("--negative-prompt", default="music")
return parser.parse_args()
def main() -> None:
args = parse_args()
generator = VideoGenerator.from_pretrained(
os.environ.get(
"MMAUDIO_MODEL_PATH",
"converted_weights/mmaudio/large_44k_v2",
),
workload_type="v2a",
num_gpus=1,
)
result = generator.generate_video(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
video_path=args.video_path,
audio_end_in_s=args.duration_seconds,
output_path=args.output_path,
save_video=True,
return_frames=False,
)
print(result["video_path"])
generator.shutdown()
if __name__ == "__main__":
main()
+1 -1
View File
@@ -100,7 +100,7 @@ class ComponentConfig:
@dataclass
class PipelineSelection:
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
workload_type: Literal["t2v", "i2v", "t2i", "i2i", "v2a", "t2a"] | None = None
preset: str | None = None
preset_version: int | None = None
components: ComponentConfig = field(default_factory=ComponentConfig)
+13 -1
View File
@@ -168,8 +168,12 @@ class FlashAttentionBackend(AttentionBackend):
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor, key_len: int) -> torch.Tensor:
# Normalize attn_mask to [B, key_len] where True means valid token.
if attn_mask.dim() == 4:
if attn_mask.shape[1] != 1 or attn_mask.shape[-2] != 1:
raise ValueError("FLASH_ATTN only supports 4D key-padding masks with shape [B, 1, 1, K]")
attn_mask = attn_mask[:, 0, 0, :]
elif attn_mask.dim() == 3:
if attn_mask.shape[-2] != 1:
raise ValueError("FLASH_ATTN only supports 3D key-padding masks with shape [B, 1, K]")
attn_mask = attn_mask[:, 0, :]
elif attn_mask.dim() != 2:
raise ValueError(f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}")
@@ -274,6 +278,14 @@ class FlashAttentionImpl(AttentionImpl):
)
attn_mask = attn_metadata.attn_mask
if getattr(attn_metadata, "is_causal", False):
return flash_attn_func_compilable(
query,
key,
value,
softmax_scale=self.softmax_scale,
causal=True,
)
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
# sequence lengths. Cross-attention can violate this.
@@ -302,7 +314,7 @@ class FlashAttentionImpl(AttentionImpl):
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=self.causal, dropout_p=0, softmax_scale=None)
elif self.nvfp4_fa4:
output = self._forward_nvfp4(query, key, value)
+2
View File
@@ -41,6 +41,8 @@ class SDPABackend(AttentionBackend):
class SDPAMetadata(AttentionMetadata):
current_timestep: int
attn_mask: torch.Tensor | None = None
# The mask is exactly native causal attention, with no additional padding.
is_causal: bool = False
class SDPAMetadataBuilder(AttentionMetadataBuilder):
+9 -1
View File
@@ -2,7 +2,13 @@ from fastvideo.configs.models.base import ModelConfig
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.models.encoders.base import EncoderConfig
from fastvideo.configs.models.vaes.base import VAEConfig
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig, LTX2AudioEncoderConfig, LTX2VocoderConfig)
from fastvideo.configs.models.audio import (
BigVGANV2Config,
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
MMAudioVAEConfig,
)
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
__all__ = [
@@ -13,5 +19,7 @@ __all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
"MMAudioVAEConfig",
"BigVGANV2Config",
"UpsamplerConfig",
]
@@ -5,9 +5,21 @@ from fastvideo.configs.models.audio.ltx2_audio_vae import (
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
from fastvideo.configs.models.audio.mmaudio_vae import (
MMAudioVAEArchConfig,
MMAudioVAEConfig,
)
from fastvideo.configs.models.audio.bigvgan import (
BigVGANV2ArchConfig,
BigVGANV2Config,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
"MMAudioVAEArchConfig",
"MMAudioVAEConfig",
"BigVGANV2ArchConfig",
"BigVGANV2Config",
]
+18
View File
@@ -0,0 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable BigVGAN-v2 vocoder configuration."""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class BigVGANV2ArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=lambda: ["BigVGANV2"])
sample_rate: int = 44100
num_mels: int = 128
@dataclass
class BigVGANV2Config(ModelConfig):
arch_config: ArchConfig = field(default_factory=BigVGANV2ArchConfig)
@@ -0,0 +1,21 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio audio VAE configuration."""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class MMAudioVAEArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=lambda: ["MMAudioVAE"])
mode: str = "44k"
data_dim: int = 128
embed_dim: int = 40
hidden_dim: int = 512
need_encoder: bool = False
@dataclass
class MMAudioVAEConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=MMAudioVAEArchConfig)
+2 -1
View File
@@ -11,6 +11,7 @@ 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.minimax_h3 import MiniMaxH3Config
from fastvideo.configs.models.dits.mmaudio import MMAudioArchConfig, MMAudioTransformerConfig
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
@@ -24,5 +25,5 @@ __all__ = [
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
"MiniMaxH3Config", "ZImageDiTConfig"
"MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
]
+49
View File
@@ -0,0 +1,49 @@
# SPDX-License-Identifier: Apache-2.0
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_mmaudio_transformer_block(name: str, module) -> bool:
del module
parts = name.split(".")
return len(parts) >= 2 and parts[-1].isdigit() and parts[-2] in {"joint_blocks", "fused_blocks"}
@dataclass
class MMAudioArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_mmaudio_transformer_block])
param_names_mapping: dict = field(default_factory=lambda: {r"^(.*)$": r"\1"})
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.TORCH_SDPA, )
latent_dim: int = 40
clip_dim: int = 1024
sync_dim: int = 768
text_dim: int = 1024
hidden_dim: int = 896
depth: int = 21
fused_depth: int = 14
num_heads: int = 14
mlp_ratio: float = 4.0
latent_seq_len: int = 345
clip_seq_len: int = 64
sync_seq_len: int = 192
text_seq_len: int = 77
v2: bool = True
def __post_init__(self) -> None:
super().__post_init__()
self.hidden_size = self.hidden_dim
self.num_attention_heads = self.num_heads
self.num_channels_latents = self.latent_dim
self.in_channels = self.latent_dim
self.out_channels = self.latent_dim
@dataclass
class MMAudioTransformerConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=MMAudioArchConfig)
prefix: str = "MMAudio"
+49 -11
View File
@@ -1,6 +1,10 @@
from fastvideo.configs.models.encoders.base import (BaseEncoderOutput, EncoderConfig, ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.encoders.base import (
BaseEncoderOutput,
EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig,
)
from fastvideo.configs.models.encoders.clip import CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5ArchConfig, LingBotWorld2UMT5Config
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
@@ -12,15 +16,49 @@ from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import (MiniMaxH3Qwen3VLArchConfig, MiniMaxH3Qwen3VLConfig)
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
StableAudioConditionerConfig)
from fastvideo.configs.models.encoders.stable_audio_conditioner import (
StableAudioConditionerArchConfig,
StableAudioConditionerConfig,
)
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
from fastvideo.configs.models.encoders.mmaudio_synchformer import MMAudioSynchformerArchConfig, MMAudioSynchformerConfig
from fastvideo.configs.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextArchConfig,
MMAudioDFNCLIPTextConfig,
MMAudioDFNCLIPVisionArchConfig,
MMAudioDFNCLIPVisionConfig,
)
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig",
"MiniMaxH3Qwen3VLArchConfig", "MiniMaxH3Qwen3VLConfig"
"EncoderConfig",
"TextEncoderConfig",
"ImageEncoderConfig",
"BaseEncoderOutput",
"CLIPTextConfig",
"CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig",
"LlamaConfig",
"T5Config",
"T5LargeConfig",
"Qwen2_5_VLConfig",
"Reason1ArchConfig",
"Reason1Config",
"LTX2GemmaConfig",
"SiglipVisionConfig",
"StableAudioConditionerArchConfig",
"StableAudioConditionerConfig",
"T5GemmaEncoderConfig",
"Qwen3TextConfig",
"Mistral3TextConfig",
"LingBotWorld2UMT5ArchConfig",
"LingBotWorld2UMT5Config",
"LingBotVideoQwen3VLTextConfig",
"MiniMaxH3Qwen3VLArchConfig",
"MiniMaxH3Qwen3VLConfig",
"MMAudioSynchformerArchConfig",
"MMAudioSynchformerConfig",
"MMAudioDFNCLIPTextArchConfig",
"MMAudioDFNCLIPTextConfig",
"MMAudioDFNCLIPVisionArchConfig",
"MMAudioDFNCLIPVisionConfig",
]
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
"""Native DFN5B CLIP conditioner configurations for MMAudio."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.clip import (
CLIPTextArchConfig,
CLIPTextConfig,
CLIPVisionArchConfig,
CLIPVisionConfig,
)
@dataclass
class MMAudioDFNCLIPTextArchConfig(CLIPTextArchConfig):
architectures: list[str] = field(default_factory=lambda: ["MMAudioDFNCLIPTextEncoder"])
vocab_size: int = 49408
hidden_size: int = 1024
intermediate_size: int = 4096
projection_dim: int = 1024
num_hidden_layers: int = 24
num_attention_heads: int = 16
max_position_embeddings: int = 77
text_len: int = 77
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
pad_token_id: int = 0
bos_token_id: int = 49406
eos_token_id: int = 49407
# MMAudio moves this encoder as one unit between CPU and GPU. Keeping it a
# plain module avoids nesting FSDP CPU-offload semantics inside the custom
# OpenCLIP causal-mask forward used by this pipeline.
_fsdp_shard_conditions: list = field(default_factory=list)
@dataclass
class MMAudioDFNCLIPTextConfig(CLIPTextConfig):
arch_config: CLIPTextArchConfig = field(default_factory=MMAudioDFNCLIPTextArchConfig)
# OpenCLIP supplies an explicit additive triangular mask to
# nn.MultiheadAttention. The MMAudio adapter reproduces that path instead
# of using SDPA's is_causal shortcut, which rounds differently in bf16.
is_causal: bool = False
prefix: str = "mmaudio_dfn_clip_text"
@dataclass
class MMAudioDFNCLIPVisionArchConfig(CLIPVisionArchConfig):
architectures: list[str] = field(default_factory=lambda: ["MMAudioDFNCLIPVisionEncoder"])
hidden_size: int = 1280
intermediate_size: int = 5120
projection_dim: int = 1024
num_hidden_layers: int = 32
num_attention_heads: int = 16
image_size: int = 378
patch_size: int = 14
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
@dataclass
class MMAudioDFNCLIPVisionConfig(CLIPVisionConfig):
arch_config: CLIPVisionArchConfig = field(default_factory=MMAudioDFNCLIPVisionArchConfig)
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = True
enable_scale: bool = True
is_causal: bool = False
prefix: str = "mmaudio_dfn_clip_vision"
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
"""Configuration for MMAudio's Synchformer visual conditioner."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
ImageEncoderArchConfig,
ImageEncoderConfig,
)
@dataclass
class MMAudioSynchformerArchConfig(ImageEncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["MMAudioSynchformerVisualEncoder"])
image_size: int = 224
num_channels: int = 3
segment_size: int = 16
segment_stride: int = 8
hidden_size: int = 768
tokens_per_segment: int = 8
@dataclass
class MMAudioSynchformerConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(default_factory=MMAudioSynchformerArchConfig)
prefix: str = "synchformer"
+2 -1
View File
@@ -11,6 +11,7 @@ from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
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,
@@ -22,5 +23,5 @@ __all__ = [
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig",
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "MMAudioV2AConfig", "get_pipeline_config_cls_from_name"
]
+5
View File
@@ -63,6 +63,11 @@ class PipelineConfig:
# Image encoder configuration
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
image_encoder_precision: str = "fp32"
# Optional multi-encoder contract. Existing pipelines continue to use the
# singular fields above; V2A and other multimodal pipelines can opt into
# indexed ``image_encoder``, ``image_encoder_2``, ... components.
image_encoder_configs: tuple[EncoderConfig, ...] | None = None
image_encoder_precisions: tuple[str, ...] | None = None
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp32", )
+66
View File
@@ -0,0 +1,66 @@
# SPDX-License-Identifier: Apache-2.0
"""Pipeline configuration for the native MMAudio video-to-audio port."""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, EncoderConfig, ModelConfig
from fastvideo.configs.models.audio import BigVGANV2Config, MMAudioVAEConfig
from fastvideo.configs.models.dits import MMAudioTransformerConfig
from fastvideo.configs.models.encoders import (
MMAudioDFNCLIPTextConfig,
MMAudioDFNCLIPVisionConfig,
MMAudioSynchformerConfig,
)
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class MMAudioV2AConfig(PipelineConfig):
"""MMAudio large-44k-v2 inference defaults.
The published demo moves every module to bfloat16. Keeping the same
per-component precision here is important: condition features seed the
complete flow trajectory, so silently encoding them in fp32 changes the
generated waveform even when the transformer weights are identical.
"""
dit_config: DiTConfig = field(default_factory=MMAudioTransformerConfig)
dit_precision: str = "bf16"
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (MMAudioDFNCLIPTextConfig(), ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
image_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (
MMAudioDFNCLIPVisionConfig(),
MMAudioSynchformerConfig(),
))
image_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
audio_decoder_config: ModelConfig = field(default_factory=MMAudioVAEConfig)
audio_decoder_precision: str = "bf16"
vocoder_config: ModelConfig = field(default_factory=BigVGANV2Config)
vocoder_precision: str = "bf16"
# Published large_44k_v2 default sequence contract. The official demo
# supports other durations, although quality can drop far away from the
# eight-second training duration.
duration_s: float = 8.0
max_audio_duration_s: float | None = None
sampling_rate: int = 44100
spectrogram_frame_rate: int = 512
latent_downsample_rate: int = 2
clip_frame_rate: int = 8
sync_frame_rate: int = 25
sync_segment_size: int = 16
sync_segment_stride: int = 8
sync_downsample_rate: int = 2
clip_image_size: int = 384
sync_image_size: int = 224
clip_batch_size_multiplier: int = 40
sync_batch_size_multiplier: int = 40
num_inference_steps: int = 25
guidance_scale: float = 4.5
vae_tiling: bool = False
vae_sp: bool = False
+21 -3
View File
@@ -50,7 +50,7 @@ from fastvideo.api.schema import (
SamplingConfig,
)
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.logger import init_logger
from fastvideo.pipelines import ForwardBatch
from fastvideo.utils import align_to, shallow_asdict
@@ -605,7 +605,13 @@ class VideoGenerator:
# Single prompt generation (original behavior)
if prompt is None:
raise ValueError("Either prompt or prompt_txt must be provided")
if fastvideo_args.workload_type is WorkloadType.V2A:
# Video semantics are sufficient conditioning for V2A models;
# model-specific text stages interpret the empty string using
# their native tokenizer/empty-prompt contract.
prompt = ""
else:
raise ValueError("Either prompt or prompt_txt must be provided")
output_path = self._prepare_output_path(sampling_param.output_path, prompt)
kwargs["output_path"] = output_path
if prompt_embeds is not None:
@@ -624,6 +630,13 @@ class VideoGenerator:
return False
return args.workload_type.value.endswith("2i")
def _is_audio_workload(self) -> bool:
"""Return True when the workload produces standalone audio."""
args = getattr(self, "fastvideo_args", None)
if args is None:
return False
return args.workload_type.value.endswith("2a")
def _prepare_output_path(
self,
output_path: str,
@@ -643,7 +656,12 @@ class VideoGenerator:
warning is logged.
- If the target path already exists, a numeric suffix is appended.
"""
target_ext = ".png" if self._is_image_workload() else ".mp4"
if self._is_image_workload():
target_ext = ".png"
elif self._is_audio_workload():
target_ext = ".wav"
else:
target_ext = ".mp4"
def _sanitize_filename_component(name: str) -> str:
# Remove characters invalid on common filesystems, strip spaces/dots
+2
View File
@@ -61,6 +61,8 @@ class WorkloadType(str, Enum):
T2V = "t2v" # Text to Video
T2I = "t2i" # Text to Image
I2I = "i2i" # Image to Image
V2A = "v2a" # Video to Audio
T2A = "t2a" # Text to Audio
@classmethod
def from_string(cls, value: str) -> "WorkloadType":
+376
View File
@@ -0,0 +1,376 @@
# SPDX-License-Identifier: MIT
"""Reusable native BigVGAN-v2 vocoder.
Adapted from NVIDIA BigVGAN-v2 and its alias-free activation implementation.
The CUDA activation kernel is intentionally excluded; FastVideo uses the
portable PyTorch path for deterministic loading and parity.
"""
from __future__ import annotations
import math
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.nn import Conv1d, ConvTranspose1d
from torch.nn.utils.parametrizations import weight_norm
from torch.nn.utils.parametrize import remove_parametrizations
class AttrDict(dict):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.__dict__ = self
def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
def init_weights(module: nn.Module, mean: float = 0.0, std: float = 0.01) -> None:
if "Conv" in module.__class__.__name__:
module.weight.data.normal_(mean, std)
def kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> torch.Tensor:
even = kernel_size % 2 == 0
half_size = kernel_size // 2
delta_f = 4 * half_width
amplitude = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
if amplitude > 50.0:
beta = 0.1102 * (amplitude - 8.7)
elif amplitude >= 21.0:
beta = 0.5842 * (amplitude - 21) ** 0.4 + 0.07886 * (amplitude - 21.0)
else:
beta = 0.0
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
time = torch.arange(-half_size, half_size) + 0.5 if even else torch.arange(kernel_size) - half_size
if cutoff == 0:
kernel = torch.zeros_like(time)
else:
kernel = 2 * cutoff * window * torch.sinc(2 * cutoff * time)
kernel /= kernel.sum()
return kernel.view(1, 1, kernel_size)
class LowPassFilter1d(nn.Module):
def __init__(
self,
cutoff: float = 0.5,
half_width: float = 0.6,
stride: int = 1,
padding: bool = True,
padding_mode: str = "replicate",
kernel_size: int = 12,
) -> None:
super().__init__()
self.kernel_size = kernel_size
self.even = kernel_size % 2 == 0
self.pad_left = kernel_size // 2 - int(self.even)
self.pad_right = kernel_size // 2
self.stride = stride
self.padding = padding
self.padding_mode = padding_mode
self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size))
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, channels, _ = x.shape
if self.padding:
x = F.pad(x, (self.pad_left, self.pad_right), mode=self.padding_mode)
return F.conv1d(x, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels)
class UpSample1d(nn.Module):
def __init__(self, ratio: int = 2, kernel_size: int | None = None) -> None:
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.stride = ratio
self.pad = self.kernel_size // ratio - 1
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
self.register_buffer(
"filter", kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=self.kernel_size)
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, channels, _ = x.shape
x = F.pad(x, (self.pad, self.pad), mode="replicate")
x = self.ratio * F.conv_transpose1d(
x, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels
)
return x[..., self.pad_left : -self.pad_right]
class DownSample1d(nn.Module):
def __init__(self, ratio: int = 2, kernel_size: int | None = None) -> None:
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.lowpass = LowPassFilter1d(
cutoff=0.5 / ratio, half_width=0.6 / ratio, stride=ratio, kernel_size=self.kernel_size
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.lowpass(x)
class Activation1d(nn.Module):
def __init__(
self,
activation: nn.Module,
up_ratio: int = 2,
down_ratio: int = 2,
up_kernel_size: int = 12,
down_kernel_size: int = 12,
) -> None:
super().__init__()
self.up_ratio = up_ratio
self.down_ratio = down_ratio
self.act = activation
self.upsample = UpSample1d(up_ratio, up_kernel_size)
self.downsample = DownSample1d(down_ratio, down_kernel_size)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.downsample(self.act(self.upsample(x)))
class Snake(nn.Module):
def __init__(
self, in_features: int, alpha: float = 1.0, alpha_trainable: bool = True, alpha_logscale: bool = False
) -> None:
super().__init__()
self.in_features = in_features
self.alpha_logscale = alpha_logscale
initial = torch.zeros(in_features) * alpha if alpha_logscale else torch.ones(in_features) * alpha
self.alpha = nn.Parameter(initial, requires_grad=alpha_trainable)
self.no_div_by_zero = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
return x + (1.0 / (alpha + self.no_div_by_zero)) * torch.pow(
torch.sin(x * alpha), 2)
class SnakeBeta(nn.Module):
def __init__(
self, in_features: int, alpha: float = 1.0, alpha_trainable: bool = True, alpha_logscale: bool = False
) -> None:
super().__init__()
self.in_features = in_features
self.alpha_logscale = alpha_logscale
initial = torch.zeros(in_features) * alpha if alpha_logscale else torch.ones(in_features) * alpha
self.alpha = nn.Parameter(initial.clone(), requires_grad=alpha_trainable)
self.beta = nn.Parameter(initial.clone(), requires_grad=alpha_trainable)
self.no_div_by_zero = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
beta = self.beta.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
beta = torch.exp(beta)
return x + (1.0 / (beta + self.no_div_by_zero)) * torch.pow(
torch.sin(x * alpha), 2)
def _activation(name: str, channels: int, logscale: bool) -> Activation1d:
if name == "snake":
activation = Snake(channels, alpha_logscale=logscale)
elif name == "snakebeta":
activation = SnakeBeta(channels, alpha_logscale=logscale)
else:
raise ValueError(f"Unsupported BigVGAN activation: {name}")
return Activation1d(activation)
class AMPBlock1(nn.Module):
def __init__(
self,
config: AttrDict,
channels: int,
kernel_size: int = 3,
dilation: tuple[int, ...] = (1, 3, 5),
activation: str = "snake",
) -> None:
super().__init__()
self.convs1 = nn.ModuleList(
[
weight_norm(
Conv1d(
channels, channels, kernel_size, stride=1, dilation=rate, padding=get_padding(kernel_size, rate)
)
)
for rate in dilation
]
)
self.convs1.apply(init_weights)
self.convs2 = nn.ModuleList(
[
weight_norm(
Conv1d(channels, channels, kernel_size, stride=1, dilation=1, padding=get_padding(kernel_size, 1))
)
for _ in dilation
]
)
self.convs2.apply(init_weights)
self.activations = nn.ModuleList(
[
_activation(activation, channels, config.snake_logscale)
for _ in range(len(self.convs1) + len(self.convs2))
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
acts1, acts2 = self.activations[::2], self.activations[1::2]
for conv1, conv2, act1, act2 in zip(self.convs1, self.convs2, acts1, acts2, strict=True):
residual = conv2(act2(conv1(act1(x))))
x = residual + x
return x
def remove_weight_norm(self) -> None:
for layer in self.convs1:
remove_parametrizations(layer, "weight")
for layer in self.convs2:
remove_parametrizations(layer, "weight")
class AMPBlock2(nn.Module):
def __init__(
self,
config: AttrDict,
channels: int,
kernel_size: int = 3,
dilation: tuple[int, ...] = (1, 3, 5),
activation: str = "snake",
) -> None:
super().__init__()
self.convs = nn.ModuleList(
[
weight_norm(
Conv1d(
channels, channels, kernel_size, stride=1, dilation=rate, padding=get_padding(kernel_size, rate)
)
)
for rate in dilation
]
)
self.convs.apply(init_weights)
self.activations = nn.ModuleList([_activation(activation, channels, config.snake_logscale) for _ in self.convs])
def forward(self, x: torch.Tensor) -> torch.Tensor:
for conv, activation in zip(self.convs, self.activations, strict=True):
x = conv(activation(x)) + x
return x
def remove_weight_norm(self) -> None:
for layer in self.convs:
remove_parametrizations(layer, "weight")
class BigVGANV2(nn.Module):
"""BigVGAN-v2 generator compatible with NVIDIA checkpoint keys."""
def __init__(self, config: dict[str, Any]) -> None:
super().__init__()
config = dict(config)
config.pop("_class_name", None)
config.pop("architectures", None)
weight_norm_removed = bool(config.pop("weight_norm_removed", False))
self.config = AttrDict(config)
if self.config.get("use_cuda_kernel", False):
raise ValueError("FastVideo BigVGANV2 supports only the portable PyTorch path")
self.config["use_cuda_kernel"] = False
self.num_kernels = len(self.config.resblock_kernel_sizes)
self.num_upsamples = len(self.config.upsample_rates)
self.conv_pre = weight_norm(Conv1d(self.config.num_mels, self.config.upsample_initial_channel, 7, 1, padding=3))
if self.config.resblock == "1":
block_class = AMPBlock1
elif self.config.resblock == "2":
block_class = AMPBlock2
else:
raise ValueError(f"Unsupported BigVGAN resblock: {self.config.resblock}")
self.ups = nn.ModuleList()
for index, (rate, kernel) in enumerate(
zip(self.config.upsample_rates, self.config.upsample_kernel_sizes, strict=True)
):
self.ups.append(
nn.ModuleList(
[
weight_norm(
ConvTranspose1d(
self.config.upsample_initial_channel // (2**index),
self.config.upsample_initial_channel // (2 ** (index + 1)),
kernel,
rate,
padding=(kernel - rate) // 2,
)
)
]
)
)
self.resblocks = nn.ModuleList()
for index in range(len(self.ups)):
channels = self.config.upsample_initial_channel // (2 ** (index + 1))
for kernel, dilation in zip(
self.config.resblock_kernel_sizes, self.config.resblock_dilation_sizes, strict=True
):
self.resblocks.append(
block_class(self.config, channels, kernel, tuple(dilation), activation=self.config.activation)
)
channels = self.config.upsample_initial_channel // (2 ** len(self.ups))
self.activation_post = _activation(self.config.activation, channels, self.config.snake_logscale)
self.use_bias_at_final = self.config.get("use_bias_at_final", True)
self.conv_post = weight_norm(Conv1d(channels, 1, 7, 1, padding=3, bias=self.use_bias_at_final))
for upsampler in self.ups:
upsampler.apply(init_weights)
self.conv_post.apply(init_weights)
self.use_tanh_at_final = self.config.get("use_tanh_at_final", True)
if weight_norm_removed:
self.remove_weight_norm()
def forward(self, mel: torch.Tensor) -> torch.Tensor:
hidden = self.conv_pre(mel)
for index in range(self.num_upsamples):
for upsampler in self.ups[index]:
hidden = upsampler(hidden)
accumulated = None
for kernel in range(self.num_kernels):
block_output = self.resblocks[
index * self.num_kernels + kernel](hidden)
if accumulated is None:
accumulated = block_output
else:
# Preserve BigVGAN's published sequential accumulation
# order. Tree reduction drifts through later nonlinear
# upsampling stages with the full checkpoint.
accumulated += block_output
assert accumulated is not None
hidden = accumulated / self.num_kernels
hidden = self.conv_post(self.activation_post(hidden))
if self.use_tanh_at_final:
return torch.tanh(hidden)
return torch.clamp(hidden, min=-1.0, max=1.0)
def remove_weight_norm(self) -> None:
try:
for upsamplers in self.ups:
for upsampler in upsamplers:
remove_parametrizations(upsampler, "weight")
for block in self.resblocks:
block.remove_weight_norm()
remove_parametrizations(self.conv_pre, "weight")
remove_parametrizations(self.conv_post, "weight")
except ValueError:
# Idempotent for pipeline setup and converted checkpoints.
return
EntryClass = BigVGANV2
+827
View File
@@ -0,0 +1,827 @@
# SPDX-License-Identifier: MIT
#
# MIT License
#
# Copyright (c) 2024 Sony Research Inc.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
"""Native 1D audio VAE used by MMAudio."""
from __future__ import annotations
import logging
import math
from collections.abc import Iterable
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from fastvideo.models.loader.weight_utils import default_weight_loader
logger = logging.getLogger(__name__)
DATA_MEAN_80D = [
-1.6058,
-1.3676,
-1.2520,
-1.2453,
-1.2078,
-1.2224,
-1.2419,
-1.2439,
-1.2922,
-1.2927,
-1.3170,
-1.3543,
-1.3401,
-1.3836,
-1.3907,
-1.3912,
-1.4313,
-1.4152,
-1.4527,
-1.4728,
-1.4568,
-1.5101,
-1.5051,
-1.5172,
-1.5623,
-1.5373,
-1.5746,
-1.5687,
-1.6032,
-1.6131,
-1.6081,
-1.6331,
-1.6489,
-1.6489,
-1.6700,
-1.6738,
-1.6953,
-1.6969,
-1.7048,
-1.7280,
-1.7361,
-1.7495,
-1.7658,
-1.7814,
-1.7889,
-1.8064,
-1.8221,
-1.8377,
-1.8417,
-1.8643,
-1.8857,
-1.8929,
-1.9173,
-1.9379,
-1.9531,
-1.9673,
-1.9824,
-2.0042,
-2.0215,
-2.0436,
-2.0766,
-2.1064,
-2.1418,
-2.1855,
-2.2319,
-2.2767,
-2.3161,
-2.3572,
-2.3954,
-2.4282,
-2.4659,
-2.5072,
-2.5552,
-2.6074,
-2.6584,
-2.7107,
-2.7634,
-2.8266,
-2.8981,
-2.9673,
]
DATA_STD_80D = [
1.0291,
1.0411,
1.0043,
0.9820,
0.9677,
0.9543,
0.9450,
0.9392,
0.9343,
0.9297,
0.9276,
0.9263,
0.9242,
0.9254,
0.9232,
0.9281,
0.9263,
0.9315,
0.9274,
0.9247,
0.9277,
0.9199,
0.9188,
0.9194,
0.9160,
0.9161,
0.9146,
0.9161,
0.9100,
0.9095,
0.9145,
0.9076,
0.9066,
0.9095,
0.9032,
0.9043,
0.9038,
0.9011,
0.9019,
0.9010,
0.8984,
0.8983,
0.8986,
0.8961,
0.8962,
0.8978,
0.8962,
0.8973,
0.8993,
0.8976,
0.8995,
0.9016,
0.8982,
0.8972,
0.8974,
0.8949,
0.8940,
0.8947,
0.8936,
0.8939,
0.8951,
0.8956,
0.9017,
0.9167,
0.9436,
0.9690,
1.0003,
1.0225,
1.0381,
1.0491,
1.0545,
1.0604,
1.0761,
1.0929,
1.1089,
1.1196,
1.1176,
1.1156,
1.1117,
1.1070,
]
DATA_MEAN_128D = [
-3.3462,
-2.6723,
-2.4893,
-2.3143,
-2.2664,
-2.3317,
-2.1802,
-2.4006,
-2.2357,
-2.4597,
-2.3717,
-2.4690,
-2.5142,
-2.4919,
-2.6610,
-2.5047,
-2.7483,
-2.5926,
-2.7462,
-2.7033,
-2.7386,
-2.8112,
-2.7502,
-2.9594,
-2.7473,
-3.0035,
-2.8891,
-2.9922,
-2.9856,
-3.0157,
-3.1191,
-2.9893,
-3.1718,
-3.0745,
-3.1879,
-3.2310,
-3.1424,
-3.2296,
-3.2791,
-3.2782,
-3.2756,
-3.3134,
-3.3509,
-3.3750,
-3.3951,
-3.3698,
-3.4505,
-3.4509,
-3.5089,
-3.4647,
-3.5536,
-3.5788,
-3.5867,
-3.6036,
-3.6400,
-3.6747,
-3.7072,
-3.7279,
-3.7283,
-3.7795,
-3.8259,
-3.8447,
-3.8663,
-3.9182,
-3.9605,
-3.9861,
-4.0105,
-4.0373,
-4.0762,
-4.1121,
-4.1488,
-4.1874,
-4.2461,
-4.3170,
-4.3639,
-4.4452,
-4.5282,
-4.6297,
-4.7019,
-4.7960,
-4.8700,
-4.9507,
-5.0303,
-5.0866,
-5.1634,
-5.2342,
-5.3242,
-5.4053,
-5.4927,
-5.5712,
-5.6464,
-5.7052,
-5.7619,
-5.8410,
-5.9188,
-6.0103,
-6.0955,
-6.1673,
-6.2362,
-6.3120,
-6.3926,
-6.4797,
-6.5565,
-6.6511,
-6.8130,
-6.9961,
-7.1275,
-7.2457,
-7.3576,
-7.4663,
-7.6136,
-7.7469,
-7.8815,
-8.0132,
-8.1515,
-8.3071,
-8.4722,
-8.7418,
-9.3975,
-9.6628,
-9.7671,
-9.8863,
-9.9992,
-10.0860,
-10.1709,
-10.5418,
-11.2795,
-11.3861,
]
DATA_STD_128D = [
2.3804,
2.4368,
2.3772,
2.3145,
2.2803,
2.2510,
2.2316,
2.2083,
2.1996,
2.1835,
2.1769,
2.1659,
2.1631,
2.1618,
2.1540,
2.1606,
2.1571,
2.1567,
2.1612,
2.1579,
2.1679,
2.1683,
2.1634,
2.1557,
2.1668,
2.1518,
2.1415,
2.1449,
2.1406,
2.1350,
2.1313,
2.1415,
2.1281,
2.1352,
2.1219,
2.1182,
2.1327,
2.1195,
2.1137,
2.1080,
2.1179,
2.1036,
2.1087,
2.1036,
2.1015,
2.1068,
2.0975,
2.0991,
2.0902,
2.1015,
2.0857,
2.0920,
2.0893,
2.0897,
2.0910,
2.0881,
2.0925,
2.0873,
2.0960,
2.0900,
2.0957,
2.0958,
2.0978,
2.0936,
2.0886,
2.0905,
2.0845,
2.0855,
2.0796,
2.0840,
2.0813,
2.0817,
2.0838,
2.0840,
2.0917,
2.1061,
2.1431,
2.1976,
2.2482,
2.3055,
2.3700,
2.4088,
2.4372,
2.4609,
2.4731,
2.4847,
2.5072,
2.5451,
2.5772,
2.6147,
2.6529,
2.6596,
2.6645,
2.6726,
2.6803,
2.6812,
2.6899,
2.6916,
2.6931,
2.6998,
2.7062,
2.7262,
2.7222,
2.7158,
2.7041,
2.7485,
2.7491,
2.7451,
2.7485,
2.7233,
2.7297,
2.7233,
2.7145,
2.6958,
2.6788,
2.6439,
2.6007,
2.4786,
2.2469,
2.1877,
2.1392,
2.0717,
2.0107,
1.9676,
1.9140,
1.7102,
0.9101,
0.7164,
]
_NORMALIZATION_EPSILON = 1e-4
_SILU_DIVISOR = 0.596
_RESIDUAL_BLEND = 0.3
_RESIDUAL_DIVISOR = math.hypot(1.0 - _RESIDUAL_BLEND, _RESIDUAL_BLEND)
def _rms_normalize(x: torch.Tensor, dim: int | tuple[int, ...]) -> torch.Tensor:
dims = (dim,) if isinstance(dim, int) else dim
element_count = math.prod(x.shape[axis] for axis in dims)
l2_norm = torch.linalg.vector_norm(x, dim=dims, keepdim=True, dtype=torch.float32)
rms = torch.add(_NORMALIZATION_EPSILON, l2_norm, alpha=element_count**-0.5)
return x / rms.to(x.dtype)
def _conv1d(in_channels: int, out_channels: int, kernel_size: int) -> nn.Conv1d:
return nn.Conv1d(in_channels, out_channels, kernel_size, padding=kernel_size // 2, bias=False)
def _conv1d_with_gain(conv: nn.Conv1d, x: torch.Tensor, gain: torch.Tensor | float) -> torch.Tensor:
return F.conv1d(x, conv.weight * gain, stride=conv.stride, padding=conv.padding, dilation=conv.dilation,
groups=conv.groups)
class DiagonalGaussianDistribution:
def __init__(self, parameters: torch.Tensor, deterministic: bool = False) -> None:
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if deterministic:
self.var = self.std = torch.zeros_like(self.mean)
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
noise = torch.empty_like(self.mean).normal_(generator=generator)
return self.mean + self.std * noise
def mode(self) -> torch.Tensor:
return self.mean
class ResnetBlock1D(nn.Module):
def __init__(
self,
*,
in_dim: int,
out_dim: int | None = None,
conv_shortcut: bool = False,
kernel_size: int = 3,
use_norm: bool = True,
) -> None:
super().__init__()
self.in_dim = in_dim
self.out_dim = in_dim if out_dim is None else out_dim
self.use_conv_shortcut = conv_shortcut
self.use_norm = use_norm
self.conv1 = _conv1d(in_dim, self.out_dim, kernel_size)
self.conv2 = _conv1d(self.out_dim, self.out_dim, kernel_size)
if self.in_dim != self.out_dim:
if conv_shortcut:
self.conv_shortcut = _conv1d(in_dim, self.out_dim, kernel_size)
else:
self.nin_shortcut = _conv1d(in_dim, self.out_dim, 1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.use_norm:
x = _rms_normalize(x, dim=1)
hidden = self.conv1(F.silu(x) / _SILU_DIVISOR)
hidden = self.conv2(F.silu(hidden) / _SILU_DIVISOR)
if self.in_dim != self.out_dim:
shortcut = self.conv_shortcut if self.use_conv_shortcut else self.nin_shortcut
x = shortcut(x)
return torch.lerp(x, hidden, _RESIDUAL_BLEND) / _RESIDUAL_DIVISOR
class AttnBlock1D(nn.Module):
def __init__(self, in_channels: int, num_heads: int = 1) -> None:
super().__init__()
self.in_channels = in_channels
self.num_heads = num_heads
self.qkv = _conv1d(in_channels, in_channels * 3, 1)
self.proj_out = _conv1d(in_channels, in_channels, 1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
qkv = self.qkv(x).reshape(x.shape[0], self.num_heads, -1, 3, x.shape[-1])
query, key, value = _rms_normalize(qkv, dim=2).unbind(3)
query = rearrange(query, "b h c l -> b h l c")
key = rearrange(key, "b h c l -> b h l c")
value = rearrange(value, "b h c l -> b h l c")
hidden = F.scaled_dot_product_attention(query, key, value)
hidden = rearrange(hidden, "b h l c -> b (h c) l")
return torch.lerp(x, self.proj_out(hidden), _RESIDUAL_BLEND) / _RESIDUAL_DIVISOR
class Upsample1D(nn.Module):
def __init__(self, in_channels: int, with_conv: bool) -> None:
super().__init__()
self.with_conv = with_conv
if with_conv:
self.conv = _conv1d(in_channels, in_channels, 3)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = F.interpolate(x, scale_factor=2.0, mode="nearest-exact")
return self.conv(x) if self.with_conv else x
class Downsample1D(nn.Module):
def __init__(self, in_channels: int, with_conv: bool) -> None:
super().__init__()
self.with_conv = with_conv
if with_conv:
self.conv1 = _conv1d(in_channels, in_channels, 1)
self.conv2 = _conv1d(in_channels, in_channels, 1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.with_conv:
x = self.conv1(x)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
return self.conv2(x) if self.with_conv else x
class Encoder1D(nn.Module):
def __init__(
self,
*,
dim: int,
ch_mult: tuple[int, ...],
num_res_blocks: int,
attn_layers: list[int],
down_layers: list[int],
in_dim: int,
embed_dim: int,
resamp_with_conv: bool = True,
double_z: bool = True,
kernel_size: int = 3,
clip_act: float = 256.0,
) -> None:
super().__init__()
self.dim = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = down_layers
self.attn_layers = attn_layers
self.conv_in = _conv1d(in_dim, dim, kernel_size)
in_ch_mult = (1,) + ch_mult
self.in_ch_mult = in_ch_mult
self.down = nn.ModuleList()
for level in range(self.num_layers):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = dim * in_ch_mult[level]
block_out = dim * ch_mult[level]
for _ in range(num_res_blocks):
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, kernel_size=kernel_size, use_norm=True))
block_in = block_out
if level in attn_layers:
attn.append(AttnBlock1D(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if level in down_layers:
down.downsample = Downsample1D(block_in, resamp_with_conv)
self.down.append(down)
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, kernel_size=kernel_size, use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, kernel_size=kernel_size, use_norm=True)
output_dim = 2 * embed_dim if double_z else embed_dim
self.conv_out = _conv1d(block_in, output_dim, kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, x: torch.Tensor) -> torch.Tensor:
states = [self.conv_in(x)]
for level in range(self.num_layers):
for block_index in range(self.num_res_blocks):
hidden = self.down[level].block[block_index](states[-1])
if len(self.down[level].attn) > 0:
hidden = self.down[level].attn[block_index](hidden)
states.append(hidden.clamp(-self.clip_act, self.clip_act))
if level in self.down_layers:
states.append(self.down[level].downsample(states[-1]))
hidden = self.mid.block_1(states[-1])
hidden = self.mid.attn_1(hidden)
hidden = self.mid.block_2(hidden).clamp(-self.clip_act, self.clip_act)
return _conv1d_with_gain(self.conv_out, F.silu(hidden) / _SILU_DIVISOR, self.learnable_gain + 1)
class Decoder1D(nn.Module):
def __init__(
self,
*,
dim: int,
out_dim: int,
ch_mult: tuple[int, ...],
num_res_blocks: int,
attn_layers: list[int],
down_layers: list[int],
in_dim: int,
embed_dim: int,
kernel_size: int = 3,
resamp_with_conv: bool = True,
clip_act: float = 256.0,
) -> None:
super().__init__()
self.ch = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = [level + 1 for level in down_layers]
block_in = dim * ch_mult[-1]
self.conv_in = _conv1d(embed_dim, block_in, kernel_size)
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
self.up = nn.ModuleList()
for level in reversed(range(self.num_layers)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = dim * ch_mult[level]
for _ in range(num_res_blocks + 1):
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, use_norm=True))
block_in = block_out
if level in attn_layers:
attn.append(AttnBlock1D(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if level in self.down_layers:
up.upsample = Upsample1D(block_in, resamp_with_conv)
self.up.insert(0, up)
self.conv_out = _conv1d(block_in, out_dim, kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, z: torch.Tensor) -> torch.Tensor:
hidden = self.conv_in(z)
hidden = self.mid.block_1(hidden)
hidden = self.mid.attn_1(hidden)
hidden = self.mid.block_2(hidden).clamp(-self.clip_act, self.clip_act)
for level in reversed(range(self.num_layers)):
for block_index in range(self.num_res_blocks + 1):
hidden = self.up[level].block[block_index](hidden)
if len(self.up[level].attn) > 0:
hidden = self.up[level].attn[block_index](hidden)
hidden = hidden.clamp(-self.clip_act, self.clip_act)
if level in self.down_layers:
hidden = self.up[level].upsample(hidden)
return _conv1d_with_gain(self.conv_out, F.silu(hidden) / _SILU_DIVISOR, self.learnable_gain + 1)
class MMAudioVAE(nn.Module):
"""MMAudio mel-spectrogram VAE for 16 kHz or 44.1 kHz audio."""
def __init__(self, mode: str | dict[str, Any] = "44k", need_encoder: bool = False) -> None:
super().__init__()
if isinstance(mode, dict):
config = mode
mode = config.get("mode", "44k")
need_encoder = config.get("need_encoder", need_encoder)
if mode == "16k":
data_dim, embed_dim, hidden_dim = 80, 20, 384
data_mean, data_std = DATA_MEAN_80D, DATA_STD_80D
elif mode == "44k":
data_dim, embed_dim, hidden_dim = 128, 40, 512
data_mean, data_std = DATA_MEAN_128D, DATA_STD_128D
else:
raise ValueError(f"Unknown MMAudio VAE mode: {mode}")
self.mode = mode
self.embed_dim = embed_dim
self._weights_normalized = False
self.register_buffer("data_mean", torch.tensor(data_mean, dtype=torch.float32).view(1, -1, 1))
self.register_buffer("data_std", torch.tensor(data_std, dtype=torch.float32).view(1, -1, 1))
if need_encoder:
self.encoder = Encoder1D(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
embed_dim=embed_dim,
)
self.decoder = Decoder1D(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
out_dim=data_dim,
embed_dim=embed_dim,
)
def encode(self, mel: torch.Tensor, normalize_input: bool = True) -> DiagonalGaussianDistribution:
self._require_normalized_weights()
if not hasattr(self, "encoder"):
raise RuntimeError("This MMAudio VAE was loaded decoder-only")
if normalize_input:
mel = self.normalize(mel)
return DiagonalGaussianDistribution(self.encoder(mel))
def decode(self, latent: torch.Tensor, unnormalize_output: bool = True) -> torch.Tensor:
self._require_normalized_weights()
mel = self.decoder(latent)
return self.unnormalize(mel) if unnormalize_output else mel
def forward(self, latent: torch.Tensor) -> torch.Tensor:
return self.decode(latent)
def normalize(self, mel: torch.Tensor) -> torch.Tensor:
return (mel - self.data_mean) / self.data_std
def unnormalize(self, mel: torch.Tensor) -> torch.Tensor:
return mel * self.data_std + self.data_mean
def _require_normalized_weights(self) -> None:
if not self._weights_normalized:
raise RuntimeError("call remove_weight_norm() before inference")
@torch.no_grad()
def remove_weight_norm(self):
for name, module in self.named_modules():
if isinstance(module, nn.Conv1d):
weight = _rms_normalize(module.weight.to(torch.float32), dim=(1, 2))
weight = weight / math.sqrt(weight[0].numel())
module.weight.copy_(weight.to(module.weight.dtype))
logger.debug("Removed weight norm from %s", name)
self._weights_normalized = True
return self
def load_weights(
self,
weights: Iterable[tuple[str, torch.Tensor]],
) -> set[str]:
params = dict(self.named_parameters())
loaded: set[str] = set()
for name, tensor in weights:
if name not in params:
continue
parameter = params[name]
loader = getattr(parameter, "weight_loader", default_weight_loader)
loader(parameter, tensor)
loaded.add(name)
return loaded
EntryClass = MMAudioVAE
+588
View File
@@ -0,0 +1,588 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio multimodal flow-prediction transformer.
The model operates on one-dimensional audio latent sequences and jointly
attends to semantic video, synchronization-video, and text conditions. The
initial implementation intentionally uses torch SDPA so its single-GPU numeric
contract remains explicit; sequence/tensor parallel support is a later,
separately verified optimization.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
import torch.nn.functional as F
from einops import rearrange
from einops.layers.torch import Rearrange
from torch import nn
from fastvideo.configs.models.dits.mmaudio import MMAudioTransformerConfig
from fastvideo.models.dits.base import BaseDiT
def compute_rope_rotations(
length: int, dim: int, theta: int, *, freq_scaling: float = 1.0, device: torch.device | str = "cpu"
) -> torch.Tensor:
if dim % 2 != 0:
raise ValueError(f"RoPE dimension must be even, got {dim}.")
with torch.amp.autocast(device_type="cuda", enabled=False):
pos = torch.arange(length, dtype=torch.float32, device=device)
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))
freqs *= freq_scaling
rotations = torch.einsum("..., f -> ... f", pos, freqs)
rotations = torch.stack(
[
torch.cos(rotations),
-torch.sin(rotations),
torch.sin(rotations),
torch.cos(rotations),
],
dim=-1,
)
return rearrange(rotations, "n d (i j) -> 1 n d i j", i=2, j=2)
def apply_rope(x: torch.Tensor, rotations: torch.Tensor) -> torch.Tensor:
with torch.amp.autocast(device_type="cuda", enabled=False):
source = x.float().view(*x.shape[:-1], -1, 1, 2)
output = rotations[..., 0] * source[..., 0] + rotations[..., 1] * source[..., 1]
return output.reshape(*x.shape).to(dtype=x.dtype)
class ChannelLastConv1d(nn.Conv1d):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return super().forward(x.permute(0, 2, 1)).permute(0, 2, 1)
class MMAudioMLP(nn.Module):
def __init__(self, dim: int, hidden_dim: int, multiple_of: int = 256) -> None:
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class MMAudioConvMLP(nn.Module):
def __init__(
self, dim: int, hidden_dim: int, multiple_of: int = 256, kernel_size: int = 3, padding: int = 1
) -> None:
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.w1 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w2 = ChannelLastConv1d(hidden_dim, dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w3 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TimestepEmbedder(nn.Module):
def __init__(self, dim: int, frequency_embedding_size: int, max_period: int) -> None:
super().__init__()
if dim % 2 != 0:
raise ValueError(f"Timestep embedding dim must be even, got {dim}.")
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, dim),
nn.SiLU(),
nn.Linear(dim, dim),
)
self.dim = dim
self.max_period = max_period
with torch.autocast("cuda", enabled=False):
freqs = 1.0 / (
10000 ** (torch.arange(0, frequency_embedding_size, 2, dtype=torch.float32) / frequency_embedding_size)
)
self.register_buffer("freqs", (10000 / max_period) * freqs, persistent=False)
def timestep_embedding(self, timestep: torch.Tensor) -> torch.Tensor:
args = timestep[:, None].float() * self.freqs[None]
return torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
def forward(self, timestep: torch.Tensor) -> torch.Tensor:
embedding = self.timestep_embedding(timestep).to(self.mlp[0].weight.dtype)
return self.mlp(embedding)
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
return x * (1 + scale) + shift
def mmaudio_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
output = F.scaled_dot_product_attention(q.contiguous(), k.contiguous(), v.contiguous())
return rearrange(output, "b h n d -> b n (h d)").contiguous()
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int) -> None:
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.qkv = nn.Linear(dim, dim * 3, bias=True)
self.q_norm = nn.RMSNorm(dim // num_heads)
self.k_norm = nn.RMSNorm(dim // num_heads)
self.split_into_heads = Rearrange(
"b n (h d j) -> b h n d j",
h=num_heads,
d=dim // num_heads,
j=3,
)
def pre_attention(
self, x: torch.Tensor, rotations: torch.Tensor | None
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q, k, v = self.split_into_heads(self.qkv(x)).chunk(3, dim=-1)
q = self.q_norm(q.squeeze(-1))
k = self.k_norm(k.squeeze(-1))
v = v.squeeze(-1)
if rotations is not None:
q = apply_rope(q, rotations)
k = apply_rope(k, rotations)
return q, k, v
def forward(self, x: torch.Tensor) -> torch.Tensor:
return mmaudio_attention(*self.pre_attention(x, None))
class MMDitSingleBlock(nn.Module):
def __init__(
self,
dim: int,
num_heads: int,
mlp_ratio: float = 4.0,
pre_only: bool = False,
kernel_size: int = 7,
padding: int = 3,
) -> None:
super().__init__()
self.norm1 = nn.LayerNorm(dim, elementwise_affine=False)
self.attn = SelfAttention(dim, num_heads)
self.pre_only = pre_only
if pre_only:
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(dim, 2 * dim, bias=True))
else:
self.linear1 = (
nn.Linear(dim, dim)
if kernel_size == 1
else ChannelLastConv1d(dim, dim, kernel_size=kernel_size, padding=padding)
)
self.norm2 = nn.LayerNorm(dim, elementwise_affine=False)
self.ffn = (
MMAudioMLP(dim, int(dim * mlp_ratio))
if kernel_size == 1
else MMAudioConvMLP(
dim,
int(dim * mlp_ratio),
kernel_size=kernel_size,
padding=padding,
)
)
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True))
def pre_attention(self, x: torch.Tensor, condition: torch.Tensor, rotations: torch.Tensor | None):
modulation = self.adaLN_modulation(condition)
if self.pre_only:
shift_msa, scale_msa = modulation.chunk(2, dim=-1)
gate_msa = shift_mlp = scale_mlp = gate_mlp = None
else:
(shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp) = modulation.chunk(6, dim=-1)
normalized = modulate(self.norm1(x), shift_msa, scale_msa)
qkv = self.attn.pre_attention(normalized, rotations)
return qkv, (gate_msa, shift_mlp, scale_mlp, gate_mlp)
def post_attention(self, x: torch.Tensor, attention_output: torch.Tensor, condition):
if self.pre_only:
return x
gate_msa, shift_mlp, scale_mlp, gate_mlp = condition
x = x + self.linear1(attention_output) * gate_msa
residual = modulate(self.norm2(x), shift_mlp, scale_mlp)
return x + self.ffn(residual) * gate_mlp
def forward(self, x: torch.Tensor, condition: torch.Tensor, rotations: torch.Tensor | None) -> torch.Tensor:
qkv, modulation = self.pre_attention(x, condition, rotations)
return self.post_attention(x, mmaudio_attention(*qkv), modulation)
class JointBlock(nn.Module):
def __init__(self, dim: int, num_heads: int, mlp_ratio: float = 4.0, pre_only: bool = False) -> None:
super().__init__()
self.pre_only = pre_only
self.latent_block = MMDitSingleBlock(dim, num_heads, mlp_ratio, pre_only=False, kernel_size=3, padding=1)
self.clip_block = MMDitSingleBlock(dim, num_heads, mlp_ratio, pre_only=pre_only, kernel_size=3, padding=1)
self.text_block = MMDitSingleBlock(dim, num_heads, mlp_ratio, pre_only=pre_only, kernel_size=1)
def forward(
self,
latent: torch.Tensor,
clip_features: torch.Tensor,
text_features: torch.Tensor,
global_condition: torch.Tensor,
extended_condition: torch.Tensor,
latent_rotations: torch.Tensor,
clip_rotations: torch.Tensor,
):
latent_qkv, latent_mod = self.latent_block.pre_attention(latent, extended_condition, latent_rotations)
clip_qkv, clip_mod = self.clip_block.pre_attention(clip_features, global_condition, clip_rotations)
text_qkv, text_mod = self.text_block.pre_attention(text_features, global_condition, None)
latent_len = latent.shape[1]
clip_len = clip_features.shape[1]
joint_qkv = [torch.cat([latent_qkv[i], clip_qkv[i], text_qkv[i]], dim=2) for i in range(3)]
attention_output = mmaudio_attention(*joint_qkv)
latent_output = attention_output[:, :latent_len]
clip_output = attention_output[:, latent_len : latent_len + clip_len]
text_output = attention_output[:, latent_len + clip_len :]
latent = self.latent_block.post_attention(latent, latent_output, latent_mod)
if not self.pre_only:
clip_features = self.clip_block.post_attention(clip_features, clip_output, clip_mod)
text_features = self.text_block.post_attention(text_features, text_output, text_mod)
return latent, clip_features, text_features
class FinalBlock(nn.Module):
def __init__(self, dim: int, out_dim: int) -> None:
super().__init__()
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(dim, 2 * dim, bias=True))
self.norm = nn.LayerNorm(dim, elementwise_affine=False)
self.conv = ChannelLastConv1d(dim, out_dim, kernel_size=7, padding=3)
def forward(self, latent: torch.Tensor, condition: torch.Tensor) -> torch.Tensor:
shift, scale = self.adaLN_modulation(condition).chunk(2, dim=-1)
return self.conv(modulate(self.norm(latent), shift, scale))
@dataclass
class PreprocessedConditions:
clip_f: torch.Tensor
sync_f: torch.Tensor
text_f: torch.Tensor
clip_f_c: torch.Tensor
text_f_c: torch.Tensor
_DEFAULT_CONFIG = MMAudioTransformerConfig()
class MMAudioTransformer(BaseDiT):
_fsdp_shard_conditions = _DEFAULT_CONFIG.arch_config._fsdp_shard_conditions
_compile_conditions = _DEFAULT_CONFIG.arch_config._compile_conditions
_supported_attention_backends = _DEFAULT_CONFIG.arch_config._supported_attention_backends
# The shared FSDP loader converts this mapping dictionary into its callable
# form. Keeping the class attribute as a dict matches every other native
# FastVideo DiT and also supports identity-mapped converted checkpoints.
param_names_mapping = _DEFAULT_CONFIG.arch_config.param_names_mapping
reverse_param_names_mapping = _DEFAULT_CONFIG.arch_config.reverse_param_names_mapping
def __init__(self, config: MMAudioTransformerConfig, hf_config: dict[str, Any], **kwargs) -> None:
del kwargs
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
self.v2 = arch.v2
self.latent_dim = arch.latent_dim
self._latent_seq_len = arch.latent_seq_len
self._clip_seq_len = arch.clip_seq_len
self._sync_seq_len = arch.sync_seq_len
self._text_seq_len = arch.text_seq_len
self.hidden_dim = arch.hidden_dim
self.num_heads = arch.num_heads
self.hidden_size = arch.hidden_size
self.num_attention_heads = arch.num_attention_heads
self.num_channels_latents = arch.num_channels_latents
activation = nn.SiLU if arch.v2 else nn.SELU
self.audio_input_proj = nn.Sequential(
ChannelLastConv1d(arch.latent_dim, arch.hidden_dim, kernel_size=7, padding=3),
activation(),
MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=7, padding=3),
)
clip_layers: list[nn.Module] = [nn.Linear(arch.clip_dim, arch.hidden_dim)]
if arch.v2:
clip_layers.append(nn.SiLU())
clip_layers.append(MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=3, padding=1))
self.clip_input_proj = nn.Sequential(*clip_layers)
self.sync_input_proj = nn.Sequential(
ChannelLastConv1d(arch.sync_dim, arch.hidden_dim, kernel_size=7, padding=3),
activation(),
MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=3, padding=1),
)
text_layers: list[nn.Module] = [nn.Linear(arch.text_dim, arch.hidden_dim)]
if arch.v2:
text_layers.append(nn.SiLU())
text_layers.append(MMAudioMLP(arch.hidden_dim, arch.hidden_dim * 4))
self.text_input_proj = nn.Sequential(*text_layers)
self.clip_cond_proj = nn.Linear(arch.hidden_dim, arch.hidden_dim)
self.text_cond_proj = nn.Linear(arch.hidden_dim, arch.hidden_dim)
self.global_cond_mlp = MMAudioMLP(arch.hidden_dim, arch.hidden_dim * 4)
self.sync_pos_emb = nn.Parameter(torch.zeros((1, 1, 8, arch.sync_dim)))
self.final_layer = FinalBlock(arch.hidden_dim, arch.latent_dim)
self.t_embed = TimestepEmbedder(
arch.hidden_dim,
frequency_embedding_size=(arch.hidden_dim if arch.v2 else 256),
max_period=(1 if arch.v2 else 10000),
)
self.joint_blocks = nn.ModuleList(
[
JointBlock(
arch.hidden_dim,
arch.num_heads,
mlp_ratio=arch.mlp_ratio,
pre_only=(index == arch.depth - arch.fused_depth - 1),
)
for index in range(arch.depth - arch.fused_depth)
]
)
self.fused_blocks = nn.ModuleList(
[
MMDitSingleBlock(arch.hidden_dim, arch.num_heads, mlp_ratio=arch.mlp_ratio, kernel_size=3, padding=1)
for _ in range(arch.fused_depth)
]
)
self.latent_mean = nn.Parameter(torch.full((1, 1, arch.latent_dim), float("nan")), requires_grad=False)
self.latent_std = nn.Parameter(torch.full((1, 1, arch.latent_dim), float("nan")), requires_grad=False)
self.empty_string_feat = nn.Parameter(torch.zeros((arch.text_seq_len, arch.text_dim)), requires_grad=False)
self.empty_clip_feat = nn.Parameter(torch.zeros(1, arch.clip_dim), requires_grad=True)
self.empty_sync_feat = nn.Parameter(torch.zeros(1, arch.sync_dim), requires_grad=True)
self.initialize_weights()
self.initialize_rotations()
self.__post_init__()
@property
def device(self) -> torch.device:
return self.latent_mean.device
@property
def latent_seq_len(self) -> int:
return self._latent_seq_len
@property
def clip_seq_len(self) -> int:
return self._clip_seq_len
@property
def sync_seq_len(self) -> int:
return self._sync_seq_len
def initialize_rotations(self) -> None:
head_dim = self.hidden_dim // self.num_heads
latent_rotations = compute_rope_rotations(self._latent_seq_len, head_dim, 10000, device=self.device)
clip_rotations = compute_rope_rotations(
self._clip_seq_len,
head_dim,
10000,
freq_scaling=self._latent_seq_len / self._clip_seq_len,
device=self.device,
)
self.register_buffer("latent_rot", latent_rotations, persistent=False)
self.register_buffer("clip_rot", clip_rotations, persistent=False)
def materialize_non_persistent_buffers(
self,
device: torch.device,
dtype: torch.dtype | None = None,
) -> None:
"""Rebuild derived buffers after meta-device production loading."""
if self.t_embed.freqs.is_meta:
frequency_dim = self.t_embed.mlp[0].in_features
freqs = 1.0 / (
10000
** (
torch.arange(0, frequency_dim, 2, dtype=torch.float32, device=device)
/ frequency_dim
)
)
freqs = (10000 / self.t_embed.max_period) * freqs
self.t_embed._buffers["freqs"] = freqs.to(dtype=dtype or torch.float32)
if self.latent_rot.is_meta or self.clip_rot.is_meta:
head_dim = self.hidden_dim // self.num_heads
self._buffers["latent_rot"] = compute_rope_rotations(
self._latent_seq_len, head_dim, 10000, device=device
)
self._buffers["clip_rot"] = compute_rope_rotations(
self._clip_seq_len,
head_dim,
10000,
freq_scaling=self._latent_seq_len / self._clip_seq_len,
device=device,
)
def update_seq_lengths(self, latent_seq_len: int, clip_seq_len: int, sync_seq_len: int) -> None:
self._latent_seq_len = latent_seq_len
self._clip_seq_len = clip_seq_len
self._sync_seq_len = sync_seq_len
self.initialize_rotations()
def initialize_weights(self) -> None:
def basic_init(module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
self.apply(basic_init)
nn.init.normal_(self.t_embed.mlp[0].weight, std=0.02)
nn.init.normal_(self.t_embed.mlp[2].weight, std=0.02)
for block in self.joint_blocks:
for stream in (block.latent_block, block.clip_block, block.text_block):
nn.init.constant_(stream.adaLN_modulation[-1].weight, 0)
nn.init.constant_(stream.adaLN_modulation[-1].bias, 0)
for block in self.fused_blocks:
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
nn.init.constant_(self.final_layer.conv.weight, 0)
nn.init.constant_(self.final_layer.conv.bias, 0)
nn.init.constant_(self.sync_pos_emb, 0)
nn.init.constant_(self.empty_clip_feat, 0)
nn.init.constant_(self.empty_sync_feat, 0)
def normalize(self, latent: torch.Tensor) -> torch.Tensor:
return latent.sub_(self.latent_mean).div_(self.latent_std)
def unnormalize(self, latent: torch.Tensor) -> torch.Tensor:
return latent.mul_(self.latent_std).add_(self.latent_mean)
def preprocess_conditions(
self, clip_features: torch.Tensor, sync_features: torch.Tensor, text_features: torch.Tensor
) -> PreprocessedConditions:
if clip_features.shape[1] != self._clip_seq_len:
raise ValueError(f"Expected {self._clip_seq_len} CLIP tokens, got {clip_features.shape}.")
if sync_features.shape[1] != self._sync_seq_len:
raise ValueError(f"Expected {self._sync_seq_len} sync tokens, got {sync_features.shape}.")
if text_features.shape[1] != self._text_seq_len:
raise ValueError(f"Expected {self._text_seq_len} text tokens, got {text_features.shape}.")
batch_size = clip_features.shape[0]
sync_features = sync_features.view(batch_size, self._sync_seq_len // 8, 8, -1) + self.sync_pos_emb
sync_features = sync_features.flatten(1, 2)
clip_features = self.clip_input_proj(clip_features)
sync_features = self.sync_input_proj(sync_features)
text_features = self.text_input_proj(text_features)
sync_features = F.interpolate(
sync_features.transpose(1, 2),
size=self._latent_seq_len,
mode="nearest-exact",
).transpose(1, 2)
return PreprocessedConditions(
clip_f=clip_features,
sync_f=sync_features,
text_f=text_features,
clip_f_c=self.clip_cond_proj(clip_features.mean(dim=1)),
text_f_c=self.text_cond_proj(text_features.mean(dim=1)),
)
def predict_flow(
self, latent: torch.Tensor, timestep: torch.Tensor, conditions: PreprocessedConditions
) -> torch.Tensor:
if latent.shape[1] != self._latent_seq_len:
raise ValueError(f"Expected latent length {self._latent_seq_len}, got {latent.shape}.")
clip_features = conditions.clip_f
text_features = conditions.text_f
latent = self.audio_input_proj(latent)
global_condition = self.global_cond_mlp(conditions.clip_f_c + conditions.text_f_c)
global_condition = self.t_embed(timestep).unsqueeze(1) + global_condition.unsqueeze(1)
extended_condition = global_condition + conditions.sync_f
for block in self.joint_blocks:
latent, clip_features, text_features = block(
latent,
clip_features,
text_features,
global_condition,
extended_condition,
self.latent_rot,
self.clip_rot,
)
for block in self.fused_blocks:
latent = block(latent, extended_condition, self.latent_rot)
# The released checkpoint was trained with global rather than sync-
# extended conditioning at the final layer; preserve that contract.
return self.final_layer(latent, global_condition)
def get_empty_string_sequence(self, batch_size: int) -> torch.Tensor:
return self.empty_string_feat.unsqueeze(0).expand(batch_size, -1, -1)
def get_empty_clip_sequence(self, batch_size: int) -> torch.Tensor:
return self.empty_clip_feat.unsqueeze(0).expand(batch_size, self._clip_seq_len, -1)
def get_empty_sync_sequence(self, batch_size: int) -> torch.Tensor:
return self.empty_sync_feat.unsqueeze(0).expand(batch_size, self._sync_seq_len, -1)
def get_empty_conditions(
self,
batch_size: int,
*,
negative_text_features: torch.Tensor | None = None,
) -> PreprocessedConditions:
empty_text = negative_text_features if negative_text_features is not None else self.get_empty_string_sequence(1)
conditions = self.preprocess_conditions(
self.get_empty_clip_sequence(1),
self.get_empty_sync_sequence(1),
empty_text,
)
conditions.clip_f = conditions.clip_f.expand(batch_size, -1, -1)
conditions.sync_f = conditions.sync_f.expand(batch_size, -1, -1)
conditions.clip_f_c = conditions.clip_f_c.expand(batch_size, -1)
if negative_text_features is None:
conditions.text_f = conditions.text_f.expand(batch_size, -1, -1)
conditions.text_f_c = conditions.text_f_c.expand(batch_size, -1)
return conditions
def guided_flow(
self,
timestep: torch.Tensor,
latent: torch.Tensor,
conditions: PreprocessedConditions,
empty_conditions: PreprocessedConditions,
guidance_scale: float,
) -> torch.Tensor:
timestep = timestep * torch.ones(len(latent), device=latent.device, dtype=latent.dtype)
if guidance_scale < 1.0:
return self.predict_flow(latent, timestep, conditions)
return guidance_scale * self.predict_flow(latent, timestep, conditions) + (
1 - guidance_scale
) * self.predict_flow(latent, timestep, empty_conditions)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: PreprocessedConditions
| tuple[torch.Tensor, torch.Tensor, torch.Tensor]
| list[torch.Tensor]
| dict[str, torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance=None,
**kwargs,
) -> torch.Tensor:
del encoder_hidden_states_image, guidance, kwargs
if isinstance(encoder_hidden_states, PreprocessedConditions):
conditions = encoder_hidden_states
elif isinstance(encoder_hidden_states, dict):
conditions = self.preprocess_conditions(
encoder_hidden_states["clip_features"],
encoder_hidden_states["sync_features"],
encoder_hidden_states["text_features"],
)
elif isinstance(encoder_hidden_states, (tuple, list)) and len(encoder_hidden_states) == 3:
conditions = self.preprocess_conditions(*encoder_hidden_states)
else:
raise TypeError(
"MMAudio encoder_hidden_states must be PreprocessedConditions, "
"a (clip, sync, text) triple, or a named condition dict."
)
return self.predict_flow(hidden_states, timestep, conditions)
EntryClass = MMAudioTransformer
+172
View File
@@ -0,0 +1,172 @@
# SPDX-License-Identifier: Apache-2.0
"""Native split DFN5B CLIP encoders used by MMAudio.
MMAudio uses one OpenCLIP checkpoint in two different ways: normalized
projected image embeddings and normalized token-wise text hidden states. The
classes here reuse FastVideo's native CLIP transformer while exposing those
two contracts as independently loadable pipeline components.
"""
from __future__ import annotations
from collections.abc import Iterable
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.attention.backends.sdpa import SDPAMetadata
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextConfig,
MMAudioDFNCLIPVisionConfig,
)
from fastvideo.models.encoders.clip import CLIPTextModel, CLIPVisionModel
from fastvideo.forward_context import set_forward_context
from fastvideo.models.loader.weight_utils import default_weight_loader
class MMAudioDFNCLIPTextEncoder(CLIPTextModel):
"""Return normalized CLIP hidden states for every text token."""
def __init__(self, config: MMAudioDFNCLIPTextConfig) -> None:
super().__init__(config)
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
if input_ids is not None:
sequence_length = input_ids.shape[-1]
device = input_ids.device
elif inputs_embeds is not None:
sequence_length = inputs_embeds.shape[-2]
device = inputs_embeds.device
else:
raise ValueError("MMAudio CLIP text encoding requires input_ids or inputs_embeds.")
model_dtype = next(self.parameters()).dtype
causal_mask = torch.empty(
(sequence_length, sequence_length),
device=device,
dtype=model_dtype,
)
causal_mask.fill_(float("-inf"))
causal_mask.triu_(1)
metadata = SDPAMetadata(
current_timestep=0,
attn_mask=causal_mask[None, None],
is_causal=True,
)
with set_forward_context(current_timestep=0, attn_metadata=metadata):
output = super().forward(
input_ids=input_ids,
position_ids=position_ids,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
output_hidden_states=output_hidden_states,
**kwargs,
)
assert output.last_hidden_state is not None
normalized = F.normalize(output.last_hidden_state, dim=-1)
return BaseEncoderOutput(
last_hidden_state=normalized,
pooler_output=output.pooler_output,
hidden_states=output.hidden_states,
attentions=output.attentions,
attention_mask=output.attention_mask,
)
def load_weights(
self,
weights: Iterable[tuple[str, torch.Tensor]],
) -> set[str]:
"""Load either split OpenCLIP Q/K/V or converted fused QKV weights."""
params = dict(self.named_parameters())
loaded: set[str] = set()
for name, tensor in weights:
if name in params:
parameter = params[name]
loader = getattr(parameter, "weight_loader", default_weight_loader)
loader(parameter, tensor)
loaded.add(name)
continue
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
target_name = name.replace(weight_name, param_name)
if target_name not in params:
continue
parameter = params[target_name]
parameter.weight_loader(parameter, tensor, shard_id)
loaded.add(target_name)
break
return loaded
class MMAudioDFNCLIPVisionEncoder(CLIPVisionModel):
"""Return normalized projected CLS embeddings for individual frames."""
def __init__(self, config: MMAudioDFNCLIPVisionConfig) -> None:
super().__init__(config)
self.visual_projection = nn.Linear(config.hidden_size, config.projection_dim, bias=False)
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: list[int] | None = None,
**kwargs,
) -> BaseEncoderOutput:
del kwargs
if feature_sample_layers is not None:
raise ValueError("MMAudio DFN5B vision encoding requires the final CLS token")
tokens = self.vision_model(pixel_values, feature_sample_layers=None)
image_features = F.normalize(self.visual_projection(tokens[:, 0]), dim=-1)
return BaseEncoderOutput(last_hidden_state=image_features, pooler_output=image_features)
def load_weights(
self,
weights: Iterable[tuple[str, torch.Tensor]],
) -> set[str]:
params = dict(self.named_parameters())
loaded: set[str] = set()
layer_count = len(self.vision_model.encoder.layers)
for name, tensor in weights:
if name.startswith("vision_model.encoder.layers"):
layer_index = int(name.split(".")[3])
if layer_index >= layer_count:
continue
# Converted checkpoints already contain fused qkv_proj tensors.
# Check exact names before matching the ``v_proj`` substring that
# also occurs at the end of ``qkv_proj``.
if name in params:
parameter = params[name]
loader = getattr(parameter, "weight_loader", default_weight_loader)
loader(parameter, tensor)
loaded.add(name)
continue
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
target_name = name.replace(weight_name, param_name)
if target_name not in params:
continue
parameter = params[target_name]
parameter.weight_loader(parameter, tensor, shard_id)
loaded.add(target_name)
break
else:
if name not in params:
continue
parameter = params[name]
loader = getattr(parameter, "weight_loader", default_weight_loader)
loader(parameter, tensor)
loaded.add(name)
return loaded
EntryClass = [MMAudioDFNCLIPTextEncoder, MMAudioDFNCLIPVisionEncoder]
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""Synchformer visual conditioner used by MMAudio.
The MotionFormer implementation is shared with FastVideo's existing
audio/video synchronization evaluator. This production adapter deliberately
owns only the visual feature extractor used by MMAudio, so its state-dict and
forward numerics match the official ``Synchformer`` module without carrying
the evaluator's unused audio and classification heads.
"""
from __future__ import annotations
from collections.abc import Iterable
import torch
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerConfig,
)
from fastvideo.models.encoders.base import ImageEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.third_party.synchformer.motionformer import MotionFormer
class MMAudioSynchformerVisualEncoder(ImageEncoder):
"""Extract temporal synchronization tokens from 25 FPS video frames."""
def __init__(self, config: MMAudioSynchformerConfig) -> None:
super().__init__(config)
self.vfeat_extractor = MotionFormer(
extract_features=True,
factorize_space_time=True,
agg_space_module="TransformerEncoderLayer",
agg_time_module="torch.nn.Identity",
add_global_repr=False,
)
def forward_segmented(self, segments: torch.Tensor) -> torch.Tensor:
"""Encode ``[B, S, 16, 3, 224, 224]`` frame segments."""
if segments.ndim != 6:
raise ValueError(f"Synchformer segments must have shape [B, S, T, C, H, W], got {tuple(segments.shape)}")
_, _, frames, channels, height, width = segments.shape
expected = self.config.arch_config
if (
frames != expected.segment_size
or channels != expected.num_channels
or height != expected.image_size
or width != expected.image_size
):
raise ValueError(
"Synchformer requires segments shaped "
f"[B, S, {expected.segment_size}, {expected.num_channels}, "
f"{expected.image_size}, {expected.image_size}], got "
f"{tuple(segments.shape)}"
)
visual = segments.permute(0, 1, 3, 2, 4, 5)
return self.vfeat_extractor(visual)
def forward(
self,
pixel_values: torch.Tensor,
**kwargs,
) -> BaseEncoderOutput:
"""Encode contiguous ``[B, T, 3, 224, 224]`` 25 FPS frames."""
del kwargs
if pixel_values.ndim != 5:
raise ValueError(f"Synchformer video must have shape [B, T, C, H, W], got {tuple(pixel_values.shape)}")
config = self.config.arch_config
frame_count = pixel_values.shape[1]
if frame_count < config.segment_size:
raise ValueError(f"Synchformer needs at least {config.segment_size} frames, got {frame_count}")
segments = (
pixel_values.unfold(
dimension=1,
size=config.segment_size,
step=config.segment_stride,
)
.permute(0, 1, 5, 2, 3, 4)
.contiguous()
)
features = self.forward_segmented(segments)
features = features.flatten(1, 2)
return BaseEncoderOutput(last_hidden_state=features)
def load_weights(
self,
weights: Iterable[tuple[str, torch.Tensor]],
) -> set[str]:
params = dict(self.named_parameters())
loaded: set[str] = set()
for name, tensor in weights:
if not name.startswith("vfeat_extractor.") or name not in params:
continue
parameter = params[name]
weight_loader = getattr(parameter, "weight_loader", default_weight_loader)
weight_loader(parameter, tensor)
loaded.add(name)
return loaded
EntryClass = MMAudioSynchformerVisualEncoder
+65 -11
View File
@@ -102,6 +102,8 @@ class ComponentLoader(ABC):
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
"image_encoder_2": (ImageEncoderLoader, "transformers"),
"image_encoder_3": (ImageEncoderLoader, "transformers"),
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
"processor": (ProcessorLoader, "transformers"),
"upsampler": (UpsamplerLoader, "diffusers"),
@@ -340,13 +342,15 @@ class TextEncoderLoader(ComponentLoader):
fastvideo_args: FastVideoArgs,
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
cpu_offload: bool | None = None,
):
use_cpu_offload = (fastvideo_args.text_encoder_cpu_offload
and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
from fastvideo.platforms import current_platform
if fastvideo_args.text_encoder_cpu_offload:
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
@@ -387,7 +391,7 @@ class TextEncoderLoader(ComponentLoader):
self._get_all_weights(
model,
model_path,
to_cpu=fastvideo_args.text_encoder_cpu_offload,
to_cpu=cpu_offload,
)) # type: ignore
self.counter_after_loading_weights = time.perf_counter()
@@ -463,7 +467,36 @@ class ImageEncoderLoader(TextEncoderLoader):
model_config.pop("model_type", None)
logger.info("HF Model config: %s", model_config)
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
base = os.path.basename(os.path.normpath(model_path))
index = 0
if base.startswith("image_encoder_"):
try:
index = int(base.split("_")[-1]) - 1
except ValueError:
index = 0
encoder_configs = getattr(
fastvideo_args.pipeline_config, "image_encoder_configs", None)
encoder_precisions = getattr(
fastvideo_args.pipeline_config, "image_encoder_precisions", None)
if encoder_configs is None:
if index != 0:
raise IndexError(
f"image encoder index {index} requires pipeline_config."
"image_encoder_configs")
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_precision = fastvideo_args.pipeline_config.image_encoder_precision
else:
if index < 0 or index >= len(encoder_configs):
raise IndexError(
f"image encoder index {index} out of range for "
f"image_encoder_configs (len={len(encoder_configs)})")
encoder_config = encoder_configs[index]
if encoder_precisions is None or index >= len(
encoder_precisions):
raise IndexError(
f"image encoder index {index} has no matching precision")
encoder_precision = encoder_precisions[index]
encoder_config.update_model_arch(model_config)
record_resolved_attention_backend(encoder_config)
@@ -479,7 +512,8 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config,
target_device,
fastvideo_args,
fastvideo_args.pipeline_config.image_encoder_precision,
encoder_precision,
cpu_offload=fastvideo_args.image_encoder_cpu_offload,
)
@@ -859,14 +893,25 @@ class AudioDecoderLoader(ComponentLoader):
return audio_vae.eval()
precision = getattr(fastvideo_args.pipeline_config, "audio_decoder_precision", "bf16")
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
audio_decoder = model_cls(config).to(target_device)
# MMAudio normalizes its magnitude-preserving convolution weights in
# fp32 and only then casts the whole feature utility module to bf16.
# Constructing/loading directly in bf16 quantizes the unnormalized
# checkpoint first and changes the decoded mel trajectory.
construction_precision = "fp32" if class_name == "MMAudioVAE" else precision
construction_device = torch.device("cpu") if class_name == "MMAudioVAE" else target_device
with set_default_torch_dtype(PRECISION_TO_TYPE[construction_precision]):
audio_decoder = model_cls(config).to(construction_device)
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
if class_name == "MMAudioVAE":
audio_decoder.load_state_dict(loaded, strict=True)
audio_decoder.remove_weight_norm()
return audio_decoder.to(device=target_device, dtype=PRECISION_TO_TYPE[precision]).eval()
decoder_state = {}
for name, tensor in loaded.items():
if name.startswith("decoder."):
@@ -880,7 +925,7 @@ class AudioDecoderLoader(ComponentLoader):
class VocoderLoader(ComponentLoader):
"""Loader for LTX-2 vocoder."""
"""Loader for native vocoders."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
@@ -890,14 +935,23 @@ class VocoderLoader(ComponentLoader):
target_device = get_local_torch_device()
precision = getattr(fastvideo_args.pipeline_config, "vocoder_precision", "bf16")
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
vocoder = model_cls(config).to(target_device)
# Canonical BigVGAN likewise removes parametrized weight norm in fp32
# before the official MMAudio feature module is cast to bf16.
construction_precision = "fp32" if class_name == "BigVGANV2" else precision
construction_device = torch.device("cpu") if class_name == "BigVGANV2" else target_device
with set_default_torch_dtype(PRECISION_TO_TYPE[construction_precision]):
vocoder = model_cls(config).to(construction_device)
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
if class_name == "BigVGANV2":
vocoder.load_state_dict(loaded, strict=True)
vocoder.remove_weight_norm()
return vocoder.to(device=target_device, dtype=PRECISION_TO_TYPE[precision]).eval()
target_module = getattr(vocoder, "model", vocoder)
target_module.load_state_dict(loaded, strict=False)
return vocoder.eval()
+10
View File
@@ -23,6 +23,7 @@ logger = init_logger(__name__)
# huggingface class name: (component_name, fastvideo module name, fastvideo class name)
_TEXT_TO_VIDEO_DIT_MODELS = {
"MMAudioTransformer": ("dits", "mmaudio", "MMAudioTransformer"),
"HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"HunyuanGameCraftTransformer3DModel": ("dits", "hunyuangamecraft", "HunyuanGameCraftTransformer3DModel"),
"HunyuanVideo15Transformer3DModel": ("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
@@ -75,6 +76,7 @@ _TEXT_TO_IMAGE_DIT_MODELS = {
}
_TEXT_ENCODER_MODELS = {
"MMAudioDFNCLIPTextEncoder": ("encoders", "mmaudio_clip", "MMAudioDFNCLIPTextEncoder"),
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"CLIPTextModelWithProjection": ("encoders", "clip", "CLIPTextModelWithProjection"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
@@ -94,6 +96,12 @@ _TEXT_ENCODER_MODELS = {
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
"MMAudioDFNCLIPVisionEncoder": ("encoders", "mmaudio_clip", "MMAudioDFNCLIPVisionEncoder"),
"MMAudioSynchformerVisualEncoder": (
"encoders",
"mmaudio_synchformer",
"MMAudioSynchformerVisualEncoder",
),
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
"CLIPVisionModel": ("encoders", "clip", "CLIPVisionModel"),
@@ -118,6 +126,8 @@ _VAE_MODELS = {
}
_AUDIO_MODELS = {
"MMAudioVAE": ("audio", "mmaudio_vae", "MMAudioVAE"),
"BigVGANV2": ("audio", "bigvgan", "BigVGANV2"),
"LTX2AudioEncoder": ("audio", "ltx2_audio_vae", "LTX2AudioEncoder"),
"LTX2AudioDecoder": ("audio", "ltx2_audio_vae", "LTX2AudioDecoder"),
"LTX2Vocoder": ("audio", "ltx2_audio_vae", "LTX2Vocoder"),
+222
View File
@@ -0,0 +1,222 @@
# MMAudio Video-to-Audio Inference
This directory contains FastVideo's native `MMAudioPipeline` implementation for
video-to-audio (V2A) and text-to-audio (T2A) generation. The first supported
checkpoint is `large_44k_v2`, which produces mono audio at 44.1 kHz.
The production pipeline is implemented with FastVideo components and does not
import the upstream `mmaudio` Python package. It loads the transformer, DFN5B
CLIP text/vision encoders, Synchformer visual encoder, audio VAE, BigVGAN-v2,
tokenizer, and scheduler from a standard FastVideo/Diffusers-style checkpoint
directory.
## Status
- `large_44k_v2` V2A inference: supported
- T2A pipeline routing: supported; the provided example currently targets V2A
- Output: mono WAV, 44.1 kHz
- Default inference duration: 8 seconds
- Variable-duration inference: supported
- Single-GPU inference with FastVideo's default offloading: supported
- Source-video/audio muxing: not yet part of the pipeline output
- Training and small/medium/16 kHz variants: not included in this port
The native pipeline has passed exact official-vs-FastVideo real-weight parity
for a 25-step two-second waveform and a real ten-second V2A inference smoke
test.
## Requirements
Follow FastVideo's main NVIDIA installation guide. The relevant baseline is:
- Linux or Windows WSL
- Python 3.10-3.12
- CUDA 12.6 or CUDA 13.0
- PyTorch 2.12.0
- One NVIDIA GPU
The port was validated with Python 3.12, PyTorch 2.12.0+cu126, and an RTX 6000
Ada 48 GB. This is a validated configuration, not a minimum VRAM claim. Default
layerwise/component offloading is enabled. FlashAttention is optional; the
pipeline falls back to Torch SDPA when FlashAttention is unavailable.
Install FastVideo from this source tree with `uv`:
```bash
cd FastVideo
uv venv --python 3.12 --seed
source .venv/bin/activate
UV_TORCH_BACKEND=cu126 uv pip install -e .
```
Use `UV_TORCH_BACKEND=cu130` instead on CUDA 13. Conda is not required.
## Checkpoint Layout
Model weights are intentionally excluded from the FastVideo Git repository.
For local development, the example expects a converted checkpoint at:
```text
converted_weights/mmaudio/large_44k_v2/
├── model_index.json
├── transformer/
├── text_encoder/
├── tokenizer/
├── image_encoder/
├── image_encoder_2/
├── audio_vae/
├── vocoder/
└── scheduler/
```
`official_weights/`, `converted_weights/`, and inference outputs are ignored by
Git. Cloning a code branch therefore does not clone the model.
### Pre-converted checkpoint
FastVideo accepts either a local directory or a Hugging Face model ID. When a
complete converted checkpoint is published, select it with:
```bash
export MMAUDIO_MODEL_PATH=ORG/MMAudio-large-44k-v2-Diffusers
```
FastVideo will then download the complete snapshot on first use and reuse the
Hugging Face cache on later runs. At the time of this port, the registered
`FastVideo/MMAudio-large-44k-v2-Diffusers` name is reserved but is not yet a
public checkpoint, so use the local conversion below.
### Convert the official weights locally
Only checkpoint conversion requires `open_clip_torch`; native FastVideo
inference does not depend on the upstream MMAudio package.
```bash
uv pip install open_clip_torch
mkdir -p official_weights/mmaudio/raw/weights
mkdir -p official_weights/mmaudio/raw/ext_weights
mkdir -p official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384
mkdir -p official_weights/mmaudio/bigvgan_v2_44khz_128band_512x
```
Download the three MMAudio assets:
```bash
curl -L --continue-at - \
https://huggingface.co/hkchengrex/MMAudio/resolve/main/weights/mmaudio_large_44k_v2.pth \
-o official_weights/mmaudio/raw/weights/mmaudio_large_44k_v2.pth
curl -L --continue-at - \
https://github.com/hkchengrex/MMAudio/releases/download/v0.1/v1-44.pth \
-o official_weights/mmaudio/raw/ext_weights/v1-44.pth
curl -L --continue-at - \
https://github.com/hkchengrex/MMAudio/releases/download/v0.1/synchformer_state_dict.pth \
-o official_weights/mmaudio/raw/ext_weights/synchformer_state_dict.pth
```
Download only the DFN5B and BigVGAN files used by the converter:
```bash
python - <<'PY'
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="apple/DFN5B-CLIP-ViT-H-14-384",
local_dir="official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384",
allow_patterns=["open_clip_config.json", "open_clip_pytorch_model.bin"],
)
snapshot_download(
repo_id="nvidia/bigvgan_v2_44khz_128band_512x",
local_dir="official_weights/mmaudio/bigvgan_v2_44khz_128band_512x",
allow_patterns=["config.json", "bigvgan_generator.pt"],
)
PY
```
Convert the source assets into the component tree consumed by FastVideo:
```bash
python scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py \
--transformer-checkpoint official_weights/mmaudio/raw/weights/mmaudio_large_44k_v2.pth \
--audio-vae-checkpoint official_weights/mmaudio/raw/ext_weights/v1-44.pth \
--synchformer-checkpoint official_weights/mmaudio/raw/ext_weights/synchformer_state_dict.pth \
--dfn5b-dir official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384 \
--bigvgan-dir official_weights/mmaudio/bigvgan_v2_44khz_128band_512x \
--output converted_weights/mmaudio/large_44k_v2
```
The converted checkpoint is approximately 9 GB.
## Run V2A Inference
The runnable example is
[`examples/inference/basic/basic_mmaudio.py`](../../../../examples/inference/basic/basic_mmaudio.py).
Run all commands from the FastVideo repository root.
```bash
export MMAUDIO_MODEL_PATH=converted_weights/mmaudio/large_44k_v2
python examples/inference/basic/basic_mmaudio.py \
--video-path /path/to/input.mp4 \
--duration-seconds 8 \
--prompt "A skateboarder rolls over concrete and lands on a metal rail." \
--negative-prompt "music, speech" \
--output-path outputs_audio/mmaudio.wav
```
To select a GPU explicitly:
```bash
CUDA_VISIBLE_DEVICES=0 \
MMAUDIO_MODEL_PATH=converted_weights/mmaudio/large_44k_v2 \
python examples/inference/basic/basic_mmaudio.py \
--video-path /path/to/input.mp4 \
--duration-seconds 10 \
--output-path outputs_audio/mmaudio_10s.wav
```
The text prompt is optional, but a short description of audible events usually
provides better control. The negative prompt can suppress unwanted categories
such as music or speech.
Eight seconds is the published training/default duration, not a hard inference
limit. Shorter and longer clips use dynamic sequence lengths. As in the
official MMAudio demo, quality may decrease when the requested duration is far
from eight seconds. If the source video is shorter than the requested duration,
the pipeline uses the available decoded duration. Synchformer requires at least
16 frames at 25 FPS, so V2A input must cover at least 0.64 seconds.
## Listen with the Source Video
The pipeline currently writes a WAV file. To create a preview MP4 that replaces
the source audio with the generated waveform:
```bash
ffmpeg -y \
-i /path/to/input.mp4 \
-i outputs_audio/mmaudio.wav \
-map 0:v:0 -map 1:a:0 \
-c:v copy -c:a aac -shortest \
outputs_audio/mmaudio_preview.mp4
```
## Troubleshooting
- **Model download fails:** verify that `MMAUDIO_MODEL_PATH` is an existing
local converted directory or a public Hugging Face repository containing
`model_index.json` and all eight components.
- **`FlashAttention-2 ... not found`:** this is informational. Torch SDPA is a
supported fallback and was used for exact parity validation.
- **The original video already has audio:** V2A conditioning reads video frames
only. The source audio is not passed into MMAudio.
- **No MP4 is returned:** the native result is audio-only by design; use the
`ffmpeg` command above for a preview mux.
- **Do not commit checkpoints:** never force-add `official_weights/` or
`converted_weights/` to the FastVideo Git repository.
## License and Attribution
The upstream MMAudio checkpoint is distributed as CC-BY-NC 4.0. Preserve its
license and attribution when redistributing raw or converted checkpoints.
@@ -0,0 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.pipelines.basic.mmaudio.mmaudio_pipeline import MMAudioPipeline
__all__ = ["MMAudioPipeline"]
@@ -0,0 +1,63 @@
# SPDX-License-Identifier: Apache-2.0
"""Native MMAudio video/text-to-audio pipeline."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.basic.mmaudio.stages import (
MMAudioDecodingStage,
MMAudioDenoisingStage,
MMAudioInputValidationStage,
MMAudioLatentPreparationStage,
MMAudioTextConditioningStage,
MMAudioVideoConditioningStage,
)
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
class MMAudioPipeline(ComposedPipelineBase):
"""MMAudio large-44k-v2 V2A/T2A inference pipeline."""
_required_config_modules = [
"transformer",
"scheduler",
"text_encoder",
"tokenizer",
"image_encoder",
"image_encoder_2",
"audio_vae",
"vocoder",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
transformer = self.get_module("transformer")
self.add_stage("input_validation_stage", MMAudioInputValidationStage())
self.add_stage(
"video_conditioning_stage",
MMAudioVideoConditioningStage(
image_encoder=self.get_module("image_encoder"),
sync_encoder=self.get_module("image_encoder_2"),
transformer=transformer,
),
)
self.add_stage(
"text_conditioning_stage",
MMAudioTextConditioningStage(
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
transformer=transformer,
),
)
self.add_stage("latent_preparation_stage", MMAudioLatentPreparationStage(transformer))
self.add_stage(
"denoising_stage",
MMAudioDenoisingStage(transformer, self.get_module("scheduler")),
)
self.add_stage(
"decoding_stage",
MMAudioDecodingStage(
audio_vae=self.get_module("audio_vae"),
vocoder=self.get_module("vocoder"),
),
)
EntryClass = MMAudioPipeline
@@ -0,0 +1,40 @@
# SPDX-License-Identifier: Apache-2.0
"""Published MMAudio large-44k-v2 inference defaults."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="MMAudio forward-time Euler flow with multimodal CFG.",
allowed_overrides=frozenset({"num_inference_steps", "guidance_scale"}),
)
_DEFAULTS = {
"seed": 42,
"guidance_scale": 4.5,
"num_inference_steps": 25,
"negative_prompt": "",
"audio_start_in_s": 0.0,
"audio_end_in_s": 8.0,
# The shared generator still validates these fields, but MMAudio returns
# audio metadata rather than materializing the placeholder pixels.
"height": 8,
"width": 8,
"num_frames": 1,
"fps": 25,
"return_frames": False,
}
MMAUDIO_LARGE_44K_V2 = InferencePreset(
name="mmaudio_large_44k_v2",
version=1,
model_family="mmaudio",
description=("MMAudio large-44k-v2 video-to-audio generation with DFN5B CLIP, "
"Synchformer, a 44.1 kHz audio VAE, and BigVGAN-v2."),
workload_type="v2a",
stage_schemas=(_DENOISE_STAGE, ),
defaults=dict(_DEFAULTS),
)
ALL_PRESETS = (MMAUDIO_LARGE_44K_V2, )
+455
View File
@@ -0,0 +1,455 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference stages for MMAudio V2A/T2A.
The video sampler intentionally mirrors ``mmaudio.data.av_utils.read_frames``:
timestamps are sampled independently at 8 FPS and 25 FPS, and a decoded frame
is repeated when the source FPS is lower than a requested sampling rate.
"""
from __future__ import annotations
import math
from pathlib import Path
import numpy as np
import torch
from torchvision.transforms import v2
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
_CLIP_MEAN = (0.48145466, 0.4578275, 0.40821073)
_CLIP_STD = (0.26862954, 0.26130258, 0.27577711)
def _read_frames_at_fps(
video_path: str | Path,
frame_rates: tuple[float, ...],
*,
start_s: float,
end_s: float,
) -> list[np.ndarray]:
"""Decode RGB frames with MMAudio's timestamp/duplication contract."""
import av
outputs: list[list[np.ndarray]] = [[] for _ in frame_rates]
next_times = [0.0 for _ in frame_rates]
deltas = [1.0 / fps for fps in frame_rates]
with av.open(str(video_path)) as container:
stream = container.streams.video[0]
stream.thread_type = "AUTO"
for packet in container.demux(stream):
for frame in packet.decode():
frame_time = frame.time
if frame_time is None or frame_time < start_s:
continue
if frame_time > end_s:
break
frame_array = None
for index in range(len(frame_rates)):
while frame_time >= next_times[index]:
if frame_array is None:
frame_array = frame.to_ndarray(format="rgb24")
outputs[index].append(frame_array)
next_times[index] += deltas[index]
if any(not frames for frames in outputs):
raise ValueError(f"Could not decode enough video frames from {video_path} in "
f"[{start_s}, {end_s}] seconds.")
return [np.stack(frames) for frames in outputs]
def preprocess_mmaudio_video(
video_path: str | Path,
*,
duration_s: float,
clip_fps: int = 8,
sync_fps: int = 25,
clip_size: int = 384,
sync_size: int = 224,
) -> tuple[torch.Tensor, torch.Tensor, float]:
"""Return official-format CLIP frames, sync frames, and effective duration.
CLIP output is float32 ``[T,3,384,384]`` in ``[0,1]``. Synchformer
output is float32 ``[T,3,224,224]`` in ``[-1,1]``.
"""
clip_array, sync_array = _read_frames_at_fps(
video_path,
(float(clip_fps), float(sync_fps)),
start_s=0.0,
end_s=duration_s,
)
clip_frames = torch.from_numpy(clip_array).permute(0, 3, 1, 2)
sync_frames = torch.from_numpy(sync_array).permute(0, 3, 1, 2)
clip_transform = v2.Compose([
v2.Resize((clip_size, clip_size), interpolation=v2.InterpolationMode.BICUBIC),
v2.ToImage(),
v2.ToDtype(torch.float32, scale=True),
])
sync_transform = v2.Compose([
v2.Resize(sync_size, interpolation=v2.InterpolationMode.BICUBIC),
v2.CenterCrop(sync_size),
v2.ToImage(),
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
clip_frames = clip_transform(clip_frames)
sync_frames = sync_transform(sync_frames)
effective_duration = min(
duration_s,
clip_frames.shape[0] / clip_fps,
sync_frames.shape[0] / sync_fps,
)
clip_frames = clip_frames[:int(clip_fps * effective_duration)]
sync_frames = sync_frames[:int(sync_fps * effective_duration)]
return clip_frames, sync_frames, effective_duration
def mmaudio_sequence_lengths(duration_s: float, pc) -> tuple[int, int, int]:
latent_length = math.ceil(duration_s * pc.sampling_rate / pc.spectrogram_frame_rate / pc.latent_downsample_rate)
clip_length = int(duration_s * pc.clip_frame_rate)
sync_frame_count = duration_s * pc.sync_frame_rate
sync_segments = ((sync_frame_count - pc.sync_segment_size) // pc.sync_segment_stride + 1)
sync_length = int(sync_segments * pc.sync_segment_size / pc.sync_downsample_rate)
return latent_length, clip_length, sync_length
class MMAudioInputValidationStage(PipelineStage):
"""Validate audio-generation inputs without invoking video pipeline logic."""
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.num_inference_steps <= 0:
raise ValueError("MMAudio num_inference_steps must be positive.")
if batch.num_videos_per_prompt != 1:
raise ValueError("MMAudio currently supports one output per request.")
if isinstance(batch.prompt, list) and len(batch.prompt) != 1:
raise ValueError("MMAudio currently supports one prompt per request.")
if isinstance(batch.negative_prompt, list) and len(batch.negative_prompt) != 1:
raise ValueError("MMAudio currently supports one negative prompt per request.")
direct_video = (batch.extra.get("mmaudio_clip_frames") is not None
and batch.extra.get("mmaudio_sync_frames") is not None)
if (fastvideo_args.workload_type is WorkloadType.V2A and batch.video_path is None and not direct_video):
raise ValueError("MMAudio V2A requires `video_path` or preprocessed MMAudio frame tensors.")
pc = fastvideo_args.pipeline_config
duration_s = batch.audio_end_in_s
duration_s = pc.duration_s if duration_s is None else float(duration_s)
start_s = 0.0 if batch.audio_start_in_s is None else float(batch.audio_start_in_s)
if start_s != 0.0:
raise ValueError("MMAudio currently generates from time zero; audio_start_in_s must be 0.")
max_duration_s = pc.max_audio_duration_s
if max_duration_s is not None and duration_s > max_duration_s:
raise ValueError(f"MMAudio duration {duration_s}s exceeds this checkpoint's "
f"{max_duration_s}s maximum.")
minimum_duration = pc.sync_segment_size / pc.sync_frame_rate
if duration_s < minimum_duration:
raise ValueError(f"MMAudio needs at least {minimum_duration:.2f}s "
f"({pc.sync_segment_size} sync frames), got {duration_s}s.")
batch.extra["mmaudio_duration_s"] = duration_s
batch.seed = 0 if batch.seed is None else int(batch.seed)
return batch
class MMAudioVideoConditioningStage(PipelineStage):
"""Decode video and run DFN5B/Synchformer with official preprocessing."""
def __init__(self, image_encoder, sync_encoder, transformer) -> None:
super().__init__()
self.image_encoder = image_encoder
self.sync_encoder = sync_encoder
self.transformer = transformer
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
pc = fastvideo_args.pipeline_config
duration_s = float(batch.extra["mmaudio_duration_s"])
clip_frames = batch.extra.get("mmaudio_clip_frames")
sync_frames = batch.extra.get("mmaudio_sync_frames")
use_video = batch.video_path is not None or (clip_frames is not None and sync_frames is not None)
if batch.video_path is not None:
clip_frames, sync_frames, duration_s = preprocess_mmaudio_video(
batch.video_path,
duration_s=duration_s,
clip_fps=pc.clip_frame_rate,
sync_fps=pc.sync_frame_rate,
clip_size=pc.clip_image_size,
sync_size=pc.sync_image_size,
)
elif use_video:
if clip_frames.ndim == 5:
if clip_frames.shape[0] != 1:
raise ValueError("MMAudio preprocessed CLIP frames must have batch size 1.")
clip_frames = clip_frames[0]
if sync_frames.ndim == 5:
if sync_frames.shape[0] != 1:
raise ValueError("MMAudio preprocessed sync frames must have batch size 1.")
sync_frames = sync_frames[0]
duration_s = min(
duration_s,
clip_frames.shape[0] / pc.clip_frame_rate,
sync_frames.shape[0] / pc.sync_frame_rate,
)
clip_frames = clip_frames[:int(duration_s * pc.clip_frame_rate)]
sync_frames = sync_frames[:int(duration_s * pc.sync_frame_rate)]
latent_length, clip_length, sync_length = mmaudio_sequence_lengths(duration_s, pc)
if clip_length <= 0 or sync_length <= 0:
raise ValueError(f"MMAudio duration {duration_s}s produces an empty condition sequence.")
self.transformer.update_seq_lengths(latent_length, clip_length, sync_length)
batch.extra["mmaudio_duration_s"] = duration_s
batch.extra["mmaudio_sequence_lengths"] = (latent_length, clip_length, sync_length)
if not use_video:
batch.extra["mmaudio_clip_features"] = self.transformer.get_empty_clip_sequence(1)
batch.extra["mmaudio_sync_features"] = self.transformer.get_empty_sync_sequence(1)
return batch
assert clip_frames is not None and sync_frames is not None
if clip_frames.shape != (clip_length, 3, pc.clip_image_size, pc.clip_image_size):
raise ValueError(
"MMAudio CLIP frames must have shape "
f"[{clip_length},3,{pc.clip_image_size},{pc.clip_image_size}], got {tuple(clip_frames.shape)}.")
expected_sync_frames = int(duration_s * pc.sync_frame_rate)
if sync_frames.shape != (expected_sync_frames, 3, pc.sync_image_size, pc.sync_image_size):
raise ValueError("MMAudio sync frames must have shape "
f"[{expected_sync_frames},3,{pc.sync_image_size},{pc.sync_image_size}], "
f"got {tuple(sync_frames.shape)}.")
device = get_local_torch_device()
model_dtype = next(self.transformer.parameters()).dtype
self.image_encoder = self.image_encoder.to(device)
clip_video = clip_frames.to(device=device, dtype=model_dtype, non_blocking=True)
mean = torch.tensor(_CLIP_MEAN, device=device, dtype=model_dtype).view(1, 3, 1, 1)
std = torch.tensor(_CLIP_STD, device=device, dtype=model_dtype).view(1, 3, 1, 1)
clip_video = (clip_video - mean) / std
clip_outputs: list[torch.Tensor] = []
chunk_size = pc.clip_batch_size_multiplier
with set_forward_context(current_timestep=0, attn_metadata=None):
for start in range(0, clip_length, chunk_size):
encoded = self.image_encoder(clip_video[start:start + chunk_size]).last_hidden_state
clip_outputs.append(encoded)
clip_features = torch.cat(clip_outputs, dim=0).unsqueeze(0)
if fastvideo_args.image_encoder_cpu_offload:
self.image_encoder = self.image_encoder.to("cpu")
self.sync_encoder = self.sync_encoder.to(device)
sync_video = sync_frames.to(device=device, dtype=model_dtype, non_blocking=True).unsqueeze(0)
with set_forward_context(current_timestep=0, attn_metadata=None):
sync_features = self.sync_encoder(sync_video).last_hidden_state
if fastvideo_args.image_encoder_cpu_offload:
self.sync_encoder = self.sync_encoder.to("cpu")
if sync_features.shape[1] != sync_length:
raise RuntimeError(f"Synchformer produced {sync_features.shape[1]} tokens; expected {sync_length}.")
batch.extra["mmaudio_clip_features"] = clip_features
batch.extra["mmaudio_sync_features"] = sync_features
return batch
class MMAudioTextConditioningStage(PipelineStage):
"""Encode positive/negative OpenCLIP token sequences and project conditions."""
def __init__(self, text_encoder, tokenizer, transformer) -> None:
super().__init__()
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.transformer = transformer
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
device = get_local_torch_device()
model_dtype = next(self.transformer.parameters()).dtype
self.text_encoder = self.text_encoder.to(device)
def encode(text: str | list[str] | None) -> torch.Tensor | None:
if text is None:
return None
texts = [text] if isinstance(text, str) else text
tokens = self.tokenizer(
texts,
padding="max_length",
truncation=True,
max_length=77,
return_tensors="pt",
)
# Hugging Face CLIPTokenizer pads with EOS by default, while
# OpenCLIP/MMAudio pads with token id zero.
input_ids = tokens.input_ids.masked_fill(tokens.attention_mask == 0, 0).to(device)
with set_forward_context(current_timestep=0, attn_metadata=None):
return self.text_encoder(input_ids).last_hidden_state.to(model_dtype)
text_features = encode(batch.prompt)
if text_features is None:
text_features = self.transformer.get_empty_string_sequence(1)
negative_text_features = encode(batch.negative_prompt)
if fastvideo_args.text_encoder_cpu_offload:
self.text_encoder = self.text_encoder.to("cpu")
conditions = self.transformer.preprocess_conditions(
batch.extra["mmaudio_clip_features"],
batch.extra["mmaudio_sync_features"],
text_features,
)
empty_conditions = self.transformer.get_empty_conditions(
1,
negative_text_features=negative_text_features,
)
batch.extra["mmaudio_conditions"] = conditions
batch.extra["mmaudio_empty_conditions"] = empty_conditions
return batch
class MMAudioLatentPreparationStage(PipelineStage):
"""Sample the Gaussian flow prior using a device-local seeded generator."""
def __init__(self, transformer) -> None:
super().__init__()
self.transformer = transformer
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is not None:
return batch
device = get_local_torch_device()
dtype = next(self.transformer.parameters()).dtype
generator = torch.Generator(device=device).manual_seed(int(batch.seed))
batch.generator = generator
batch.latents = torch.randn(
(1, self.transformer.latent_seq_len, self.transformer.latent_dim),
device=device,
dtype=dtype,
generator=generator,
)
return batch
class MMAudioDenoisingStage(PipelineStage):
"""Run MMAudio's forward-time Euler flow with FastVideo's shared scheduler."""
def __init__(self, transformer, scheduler) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None
# Official MMAudio builds ``torch.linspace`` on CPU. A CPU scalar in
# the bf16 CUDA update follows scalar-promotion rules, whereas moving
# the same float32 scalar to CUDA changes the rounded trajectory.
self.scheduler.set_timesteps(batch.num_inference_steps, device="cpu")
conditions = batch.extra["mmaudio_conditions"]
empty_conditions = batch.extra["mmaudio_empty_conditions"]
latents = batch.latents
for index, timestep in enumerate(self.scheduler.timesteps):
flow = self.transformer.guided_flow(
timestep / self.scheduler.config.num_train_timesteps,
latents,
conditions,
empty_conditions,
float(batch.guidance_scale),
)
# The shared FlowMatch scheduler supplies the exact inverted
# forward-time sigma schedule. Its generic ``step`` deliberately
# upcasts samples to fp32, however, while MMAudio's published
# Euler loop performs ``x += dt * flow`` in the model's bf16
# dtype. Preserve that operation order/precision here; changing
# it shifts the full 25-step trajectory.
delta = self.scheduler.sigmas[index + 1] - self.scheduler.sigmas[index]
latents = latents + delta * flow
batch.step_index = index
batch.timestep = timestep
batch.latents = self.transformer.unnormalize(latents)
return batch
class MMAudioDecodingStage(PipelineStage):
"""Decode MMAudio latents to a mono 44.1 kHz waveform."""
def __init__(self, audio_vae, vocoder) -> None:
super().__init__()
self.audio_vae = audio_vae
self.vocoder = vocoder
def verify_input(self, batch, fastvideo_args):
return VerificationResult()
def verify_output(self, batch, fastvideo_args):
return VerificationResult()
@torch.inference_mode()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None
if fastvideo_args.output_type == "latent":
batch.output = batch.latents.detach().cpu()
return batch
device = get_local_torch_device()
self.audio_vae = self.audio_vae.to(device)
self.vocoder = self.vocoder.to(device)
decoder_dtype = next(self.audio_vae.parameters()).dtype
mel = self.audio_vae.decode(batch.latents.transpose(1, 2).to(decoder_dtype))
audio = self.vocoder(mel.to(next(self.vocoder.parameters()).dtype))
pc = fastvideo_args.pipeline_config
expected_samples = batch.extra["mmaudio_sequence_lengths"][0]
expected_samples *= pc.spectrogram_frame_rate * pc.latent_downsample_rate
if audio.shape[-1] != expected_samples:
raise RuntimeError(f"MMAudio vocoder produced {audio.shape[-1]} samples; expected {expected_samples}.")
decoded_audio = audio.detach().float().cpu()
batch.extra["audio"] = decoded_audio[0].T.contiguous().numpy()
batch.extra["audio_sample_rate"] = int(pc.sampling_rate)
batch.extra["audio_only"] = True
batch.extra["decoded_audio"] = decoded_audio
batch.data_type = "audio"
# VideoGenerator currently transports a tensor for every workload;
# audio-only paths never materialize or save these pixels.
batch.output = torch.zeros((1, 3, 1, 8, 8), dtype=torch.uint8)
return batch
+24 -4
View File
@@ -41,6 +41,7 @@ from fastvideo.configs.pipelines.flux_2 import (
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionI2V_A14B_Config,
TurboDiffusionT2V_14B_Config,
@@ -239,6 +240,22 @@ def _get_config_info(
def _register_configs() -> None:
# MMAudio large-44k-v2 (video/text-to-audio). The checkpoint is converted
# into standard per-component FastVideo/Diffusers-style directories by
# scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=MMAudioV2AConfig,
workload_types=(WorkloadType.V2A, WorkloadType.T2A),
hf_model_paths=["FastVideo/MMAudio-large-44k-v2-Diffusers"],
model_detectors=[
lambda path: "mmaudio" in path.lower() or "mmaudiopipeline" in path.lower(),
],
model_family="mmaudio",
default_preset="mmaudio_large_44k_v2",
pipeline_cls_name="MMAudioPipeline",
)
# LTX-2 (distilled) — registered FIRST so its detector wins over
# the base detector when both fire. The detector loop in
# ``get_model_name_for_path`` ORs the path-based check with a
@@ -310,12 +327,12 @@ def _register_configs() -> None:
# ship `model.safetensors` as a single monolithic checkpoint with
# no per-component subfolders our standard loader can consume. See
# `scripts/checkpoint_conversion/stable_audio_to_diffusers.py`.
# NOTE: WorkloadType has no T2A variant yet; use T2V as the
# compatibility placeholder until the enum is extended.
register_configs(
sampling_param_cls=None,
pipeline_config_cls=StableAudioT2AConfig,
workload_types=(WorkloadType.T2V, ),
# Keep T2V as a backward-compatible alias for existing callers while
# exposing the semantically correct audio workload to new clients.
workload_types=(WorkloadType.T2A, WorkloadType.T2V),
hf_model_paths=[
"FastVideo/stable-audio-open-1.0-Diffusers",
],
@@ -334,7 +351,7 @@ def _register_configs() -> None:
register_configs(
sampling_param_cls=None,
pipeline_config_cls=StableAudioOpenSmallConfig,
workload_types=(WorkloadType.T2V, ),
workload_types=(WorkloadType.T2A, WorkloadType.T2V),
hf_model_paths=[
"FastVideo/stable-audio-open-small-Diffusers",
],
@@ -1302,6 +1319,8 @@ def _register_presets() -> None:
ALL_PRESETS as MATRIXGAME3_PRESETS, )
from fastvideo.pipelines.basic.minimax_h3.presets import (
ALL_PRESETS as MINIMAX_H3_PRESETS, )
from fastvideo.pipelines.basic.mmaudio.presets import (
ALL_PRESETS as MMAUDIO_PRESETS, )
from fastvideo.pipelines.basic.sd35.presets import (
ALL_PRESETS as SD35_PRESETS, )
from fastvideo.pipelines.basic.stable_audio.presets import (
@@ -1333,6 +1352,7 @@ def _register_presets() -> None:
MATRIXGAME2_PRESETS,
MATRIXGAME3_PRESETS,
MINIMAX_H3_PRESETS,
MMAUDIO_PRESETS,
SD35_PRESETS,
STABLE_AUDIO_PRESETS,
TURBODIFFUSION_PRESETS,
@@ -85,6 +85,50 @@ def test_no_pad_inference_matches_original(no_pad_impls, dtype):
torch.testing.assert_close(out_test, out_ref, atol=0, rtol=0)
def test_backend_translates_explicit_causal_mask(no_pad_impls):
"""MMAudio's additive mask must route through native causal attention."""
from fastvideo.attention.backends.flash_attn import (
FlashAttentionImpl,
flash_attn_func_compilable,
)
from fastvideo.attention.backends.sdpa import SDPAMetadata
torch.manual_seed(0)
device = torch.device("cuda")
batch, seqlen, heads, head_dim = 1, 32, 2, 64
query = torch.randn(batch, seqlen, heads, head_dim, device=device, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
causal_mask = torch.full(
(1, 1, seqlen, seqlen),
float("-inf"),
device=device,
dtype=query.dtype,
).triu_(1)
impl = FlashAttentionImpl(
num_heads=heads,
head_size=head_dim,
causal=False,
softmax_scale=head_dim**-0.5,
)
with torch.inference_mode():
expected = flash_attn_func_compilable(
query,
key,
value,
softmax_scale=head_dim**-0.5,
causal=True,
)
actual = impl.forward(
query,
key,
value,
SDPAMetadata(current_timestep=0, attn_mask=causal_mask, is_causal=True),
)
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_no_pad_training_backward_through_registered_autograd(no_pad_impls, dtype):
"""FA2: grads flow through the registered op and match the original."""
@@ -0,0 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.models.dits.mmaudio import TimestepEmbedder
def test_timestep_embedder_accepts_integer_timesteps() -> None:
embedder = TimestepEmbedder(8, 8, 10_000)
actual = embedder(torch.tensor([1], dtype=torch.long))
expected = embedder(torch.tensor([1.0]))
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
assert actual.dtype == embedder.mlp[0].weight.dtype
+2 -2
View File
@@ -4,8 +4,8 @@ from transformers.modeling_outputs import BaseModelOutputWithPooling
# importing modified version of AST
from fastvideo.third_party.eval.synchformer.hf_src.modeling_ast import (ASTConfig, ASTForAudioClassification)
from fastvideo.third_party.eval.synchformer.motionformer import (AveragePooling, BaseEncoderLayer,
TemporalTransformerEncoderLayer)
from fastvideo.third_party.synchformer.motionformer import (AveragePooling, BaseEncoderLayer,
TemporalTransformerEncoderLayer)
class AST(torch.nn.Module):
+2 -387
View File
@@ -1,388 +1,3 @@
import logging
from pathlib import Path
"""Backward-compatible import for the shared Synchformer backbone."""
import einops
import torch
from omegaconf import OmegaConf
from timm.layers import trunc_normal_
from torch import nn
from fastvideo.third_party.eval.synchformer.utils import check_if_file_exists_else_download
from fastvideo.third_party.eval.synchformer.video_model_builder import VisionTransformer
FILE2URL = {
# cfg
'motionformer_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/motionformer_224_16x4.yaml',
'joint_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/joint_224_16x4.yaml',
'divided_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/divided_224_16x4.yaml',
# ckpt
'ssv2_motionformer_224_16x4.pyth': 'https://dl.fbaipublicfiles.com/motionformer/ssv2_motionformer_224_16x4.pyth',
'ssv2_joint_224_16x4.pyth': 'https://dl.fbaipublicfiles.com/motionformer/ssv2_joint_224_16x4.pyth',
'ssv2_divided_224_16x4.pyth': 'https://dl.fbaipublicfiles.com/motionformer/ssv2_divided_224_16x4.pyth',
}
class MotionFormer(VisionTransformer):
''' This class serves three puposes:
1. Renames the class to MotionFormer.
2. Downloads the cfg from the original repo and patches it if needed.
3. Takes care of feature extraction by redefining .forward()
- if `extract_features=True` and `factorize_space_time=False`,
the output is of shape (B, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
- if `extract_features=True` and `factorize_space_time=True`, the output is of shape (B*S, D)
and spatial and temporal transformer encoder layers are used.
- if `extract_features=True` and `factorize_space_time=True` as well as `add_global_repr=True`
the output is of shape (B, D) and spatial and temporal transformer encoder layers
are used as well as the global representation is extracted from segments (extra pos emb
is added).
'''
def __init__(
self,
extract_features: bool = False,
ckpt_path: str = None,
factorize_space_time: bool = None,
agg_space_module: str = None,
agg_time_module: str = None,
add_global_repr: bool = True,
agg_segments_module: str = None,
max_segments: int = None,
):
self.extract_features = extract_features
self.ckpt_path = ckpt_path
self.factorize_space_time = factorize_space_time
if self.ckpt_path is not None:
check_if_file_exists_else_download(self.ckpt_path, FILE2URL)
ckpt = torch.load(self.ckpt_path, map_location='cpu')
mformer_ckpt2cfg = {
'ssv2_motionformer_224_16x4.pyth': 'motionformer_224_16x4.yaml',
'ssv2_joint_224_16x4.pyth': 'joint_224_16x4.yaml',
'ssv2_divided_224_16x4.pyth': 'divided_224_16x4.yaml',
}
# init from motionformer ckpt or from our Stage I ckpt
# depending on whether the feat extractor was pre-trained on AVCLIPMoCo or not, we need to
# load the state dict differently
was_pt_on_avclip = self.ckpt_path.endswith('.pt') # checks if it is a stage I ckpt (FIXME: a bit generic)
if self.ckpt_path.endswith(tuple(mformer_ckpt2cfg.keys())):
cfg_fname = mformer_ckpt2cfg[Path(self.ckpt_path).name]
elif was_pt_on_avclip:
# TODO: this is a hack, we should be able to get the cfg from the ckpt (earlier ckpt didn't have it)
s1_cfg = ckpt.get('args', None) # Stage I cfg
if s1_cfg is not None:
s1_vfeat_extractor_ckpt_path = s1_cfg.model.params.vfeat_extractor.params.ckpt_path
# if the stage I ckpt was initialized from a motionformer ckpt or train from scratch
if s1_vfeat_extractor_ckpt_path is not None:
cfg_fname = mformer_ckpt2cfg[Path(s1_vfeat_extractor_ckpt_path).name]
else:
cfg_fname = 'divided_224_16x4.yaml'
else:
cfg_fname = 'divided_224_16x4.yaml'
else:
raise ValueError(f'ckpt_path {self.ckpt_path} is not supported.')
else:
was_pt_on_avclip = False
cfg_fname = 'divided_224_16x4.yaml'
# logging.info(f'No ckpt_path provided, using {cfg_fname} config.')
if cfg_fname in ['motionformer_224_16x4.yaml', 'divided_224_16x4.yaml']:
pos_emb_type = 'separate'
elif cfg_fname == 'joint_224_16x4.yaml':
pos_emb_type = 'joint'
self.mformer_cfg_path = Path(__file__).absolute().parent / cfg_fname
check_if_file_exists_else_download(self.mformer_cfg_path, FILE2URL)
mformer_cfg = OmegaConf.load(self.mformer_cfg_path)
logging.info(f'Loading MotionFormer config from {self.mformer_cfg_path.absolute()}')
# patch the cfg (from the default cfg defined in the repo `Motionformer/slowfast/config/defaults.py`)
mformer_cfg.VIT.ATTN_DROPOUT = 0.0
mformer_cfg.VIT.POS_EMBED = pos_emb_type
mformer_cfg.VIT.USE_ORIGINAL_TRAJ_ATTN_CODE = True
mformer_cfg.VIT.APPROX_ATTN_TYPE = 'none' # guessing
mformer_cfg.VIT.APPROX_ATTN_DIM = 64 # from ckpt['cfg']
# finally init VisionTransformer with the cfg
super().__init__(mformer_cfg)
# load the ckpt now if ckpt is provided and not from AVCLIPMoCo-pretrained ckpt
if (self.ckpt_path is not None) and (not was_pt_on_avclip):
_ckpt_load_status = self.load_state_dict(ckpt['model_state'], strict=False)
if len(_ckpt_load_status.missing_keys) > 0 or len(_ckpt_load_status.unexpected_keys) > 0:
logging.warning(f'Loading exact vfeat_extractor ckpt from {self.ckpt_path} failed.' \
f'Missing keys: {_ckpt_load_status.missing_keys}, ' \
f'Unexpected keys: {_ckpt_load_status.unexpected_keys}')
else:
logging.info(f'Loading vfeat_extractor ckpt from {self.ckpt_path} succeeded.')
if self.extract_features:
assert isinstance(self.norm, nn.LayerNorm), 'early x[:, 1:, :] may not be safe for per-tr weights'
# pre-logits are Sequential(nn.Linear(emb, emd), act) and `act` is tanh but see the logger
self.pre_logits = nn.Identity()
# we don't need the classification head (saving memory)
self.head = nn.Identity()
self.head_drop = nn.Identity()
# avoiding code duplication (used only if agg_*_module is TransformerEncoderLayer)
transf_enc_layer_kwargs = dict(
d_model=self.embed_dim,
nhead=self.num_heads,
activation=nn.GELU(),
batch_first=True,
dim_feedforward=self.mlp_ratio * self.embed_dim,
dropout=self.drop_rate,
layer_norm_eps=1e-6,
norm_first=True,
)
# define adapters if needed
if self.factorize_space_time:
if agg_space_module == 'TransformerEncoderLayer':
self.spatial_attn_agg = SpatialTransformerEncoderLayer(**transf_enc_layer_kwargs)
elif agg_space_module == 'AveragePooling':
self.spatial_attn_agg = AveragePooling(avg_pattern='BS D t h w -> BS D t',
then_permute_pattern='BS D t -> BS t D')
if agg_time_module == 'TransformerEncoderLayer':
self.temp_attn_agg = TemporalTransformerEncoderLayer(**transf_enc_layer_kwargs)
elif agg_time_module == 'AveragePooling':
self.temp_attn_agg = AveragePooling(avg_pattern='BS t D -> BS D')
elif 'Identity' in agg_time_module:
self.temp_attn_agg = nn.Identity()
# define a global aggregation layer (aggregarate over segments)
self.add_global_repr = add_global_repr
if add_global_repr:
if agg_segments_module == 'TransformerEncoderLayer':
# we can reuse the same layer as for temporal factorization (B, dim_to_agg, D) -> (B, D)
# we need to add pos emb (PE) because previously we added the same PE for each segment
pos_max_len = max_segments if max_segments is not None else 16 # 16 = 10sec//0.64sec + 1
self.global_attn_agg = TemporalTransformerEncoderLayer(add_pos_emb=True,
pos_emb_drop=mformer_cfg.VIT.POS_DROPOUT,
pos_max_len=pos_max_len,
**transf_enc_layer_kwargs)
elif agg_segments_module == 'AveragePooling':
self.global_attn_agg = AveragePooling(avg_pattern='B S D -> B D')
if was_pt_on_avclip:
# we need to filter out the state_dict of the AVCLIP model (has both A and V extractors)
# and keep only the state_dict of the feat extractor
ckpt_weights = dict()
for k, v in ckpt['state_dict'].items():
if k.startswith(('module.v_encoder.', 'v_encoder.')):
k = k.replace('module.', '').replace('v_encoder.', '')
ckpt_weights[k] = v
_load_status = self.load_state_dict(ckpt_weights, strict=False)
if len(_load_status.missing_keys) > 0 or len(_load_status.unexpected_keys) > 0:
logging.warning(f'Loading exact vfeat_extractor ckpt from {self.ckpt_path} failed. \n' \
f'Missing keys ({len(_load_status.missing_keys)}): ' \
f'{_load_status.missing_keys}, \n' \
f'Unexpected keys ({len(_load_status.unexpected_keys)}): ' \
f'{_load_status.unexpected_keys} \n' \
f'temp_attn_agg are expected to be missing if ckpt was pt contrastively.')
else:
logging.info(f'Loading vfeat_extractor ckpt from {self.ckpt_path} succeeded.')
# patch_embed is not used in MotionFormer, only patch_embed_3d, because cfg.VIT.PATCH_SIZE_TEMP > 1
# but it used to calculate the number of patches, so we need to set keep it
self.patch_embed.requires_grad_(False)
def forward(self, x):
'''
x is of shape (B, S, C, T, H, W) where S is the number of segments.
'''
# Batch, Segments, Channels, T=frames, Height, Width
B, S, C, T, H, W = x.shape
# Motionformer expects a tensor of shape (1, B, C, T, H, W).
# The first dimension (1) is a dummy dimension to make the input tensor and won't be used:
# see `video_model_builder.video_input`.
# x = x.unsqueeze(0) # (1, B, S, C, T, H, W)
orig_shape = (B, S, C, T, H, W)
x = x.view(B * S, C, T, H, W) # flatten batch and segments
x = self.forward_segments(x, orig_shape=orig_shape)
# unpack the segments (using rest dimensions to support different shapes e.g. (BS, D) or (BS, t, D))
x = x.view(B, S, *x.shape[1:])
# x is now of shape (B*S, D) or (B*S, t, D) if `self.temp_attn_agg` is `Identity`
return x # x is (B, S, ...)
def forward_segments(self, x, orig_shape: tuple) -> torch.Tensor:
'''x is of shape (1, BS, C, T, H, W) where S is the number of segments.'''
x, x_mask = self.forward_features(x)
assert self.extract_features
# (BS, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
x = x[:, 1:, :] # without the CLS token for efficiency (should be safe for LayerNorm and FC)
x = self.norm(x)
x = self.pre_logits(x)
if self.factorize_space_time:
x = self.restore_spatio_temp_dims(x, orig_shape) # (B*S, D, t, h, w) <- (B*S, t*h*w, D)
x = self.spatial_attn_agg(x, x_mask) # (B*S, t, D)
x = self.temp_attn_agg(x) # (B*S, D) or (BS, t, D) if `self.temp_attn_agg` is `Identity`
return x
def restore_spatio_temp_dims(self, feats: torch.Tensor, orig_shape: tuple) -> torch.Tensor:
'''
feats are of shape (B*S, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
Our goal is to make them of shape (B*S, t, h, w, D) where h, w are the spatial dimensions.
From `self.patch_embed_3d`, it follows that we could reshape feats with:
`feats.transpose(1, 2).view(B*S, D, t, h, w)`
'''
B, S, C, T, H, W = orig_shape
D = self.embed_dim
# num patches in each dimension
t = T // self.patch_embed_3d.z_block_size
h = self.patch_embed_3d.height
w = self.patch_embed_3d.width
feats = feats.permute(0, 2, 1) # (B*S, D, T)
feats = feats.view(B * S, D, t, h, w) # (B*S, D, t, h, w)
return feats
class BaseEncoderLayer(nn.TransformerEncoderLayer):
'''
This is a wrapper around nn.TransformerEncoderLayer that adds a CLS token
to the sequence and outputs the CLS token's representation.
This base class parents both SpatialEncoderLayer and TemporalEncoderLayer for the RGB stream
and the FrequencyEncoderLayer and TemporalEncoderLayer for the audio stream stream.
We also, optionally, add a positional embedding to the input sequence which
allows to reuse it for global aggregation (of segments) for both streams.
'''
def __init__(self,
add_pos_emb: bool = False,
pos_emb_drop: float = None,
pos_max_len: int = None,
*args_transformer_enc,
**kwargs_transformer_enc):
super().__init__(*args_transformer_enc, **kwargs_transformer_enc)
self.cls_token = nn.Parameter(torch.zeros(1, 1, self.self_attn.embed_dim))
trunc_normal_(self.cls_token, std=.02)
# add positional embedding
self.add_pos_emb = add_pos_emb
if add_pos_emb:
self.pos_max_len = 1 + pos_max_len # +1 (for CLS)
self.pos_emb = nn.Parameter(torch.zeros(1, self.pos_max_len, self.self_attn.embed_dim))
self.pos_drop = nn.Dropout(pos_emb_drop)
trunc_normal_(self.pos_emb, std=.02)
self.apply(self._init_weights)
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None):
''' x is of shape (B, N, D); if provided x_mask is of shape (B, N)'''
batch_dim = x.shape[0]
# add CLS token
cls_tokens = self.cls_token.expand(batch_dim, -1, -1) # expanding to match batch dimension
x = torch.cat((cls_tokens, x), dim=-2) # (batch_dim, 1+seq_len, D)
if x_mask is not None:
cls_mask = torch.ones((batch_dim, 1), dtype=torch.bool, device=x_mask.device) # 1=keep; 0=mask
x_mask_w_cls = torch.cat((cls_mask, x_mask), dim=-1) # (batch_dim, 1+seq_len)
B, N = x_mask_w_cls.shape
# torch expects (N, N) or (B*num_heads, N, N) mask (sadness ahead); torch masks
x_mask_w_cls = x_mask_w_cls.reshape(B, 1, 1, N)\
.expand(-1, self.self_attn.num_heads, N, -1)\
.reshape(B * self.self_attn.num_heads, N, N)
assert x_mask_w_cls.dtype == x_mask_w_cls.bool().dtype, 'x_mask_w_cls.dtype != bool'
x_mask_w_cls = ~x_mask_w_cls # invert mask (1=mask)
else:
x_mask_w_cls = None
# add positional embedding
if self.add_pos_emb:
seq_len = x.shape[1] # (don't even think about moving it before the CLS token concatenation)
assert seq_len <= self.pos_max_len, f'Seq len ({seq_len}) > pos_max_len ({self.pos_max_len})'
x = x + self.pos_emb[:, :seq_len, :]
x = self.pos_drop(x)
# apply encoder layer (calls nn.TransformerEncoderLayer.forward);
x = super().forward(src=x, src_mask=x_mask_w_cls) # (batch_dim, 1+seq_len, D)
# CLS token is expected to hold spatial information for each frame
x = x[:, 0, :] # (batch_dim, D)
return x
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
@torch.jit.ignore
def no_weight_decay(self):
return {'cls_token', 'pos_emb'}
class SpatialTransformerEncoderLayer(BaseEncoderLayer):
''' Aggregates spatial dimensions by applying attention individually to each frame. '''
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None) -> torch.Tensor:
''' x is of shape (B*S, D, t, h, w) where S is the number of segments.
if specified x_mask (B*S, t, h, w), 0=masked, 1=kept
Returns a tensor of shape (B*S, t, D) pooling spatial information for each frame. '''
BS, D, t, h, w = x.shape
# time as a batch dimension and flatten spatial dimensions as sequence
x = einops.rearrange(x, 'BS D t h w -> (BS t) (h w) D')
# similar to mask
if x_mask is not None:
x_mask = einops.rearrange(x_mask, 'BS t h w -> (BS t) (h w)')
# apply encoder layer (BaseEncoderLayer.forward) - it will add CLS token and output its representation
x = super().forward(x=x, x_mask=x_mask) # (B*S*t, D)
# reshape back to (B*S, t, D)
x = einops.rearrange(x, '(BS t) D -> BS t D', BS=BS, t=t)
# (B*S, t, D)
return x
class TemporalTransformerEncoderLayer(BaseEncoderLayer):
''' Aggregates temporal dimension with attention. Also used with pos emb as global aggregation
in both streams. '''
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x):
''' x is of shape (B*S, t, D) where S is the number of segments.
Returns a tensor of shape (B*S, D) pooling temporal information. '''
BS, t, D = x.shape
# apply encoder layer (BaseEncoderLayer.forward) - it will add CLS token and output its representation
x = super().forward(x) # (B*S, D)
return x # (B*S, D)
class AveragePooling(nn.Module):
def __init__(self, avg_pattern: str, then_permute_pattern: str = None) -> None:
''' patterns are e.g. "bs t d -> bs d" '''
super().__init__()
# TODO: need to register them as buffers (but fails because these are strings)
self.reduce_fn = 'mean'
self.avg_pattern = avg_pattern
self.then_permute_pattern = then_permute_pattern
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None) -> torch.Tensor:
x = einops.reduce(x, self.avg_pattern, self.reduce_fn)
if self.then_permute_pattern is not None:
x = einops.rearrange(x, self.then_permute_pattern)
return x
from fastvideo.third_party.synchformer.motionformer import * # noqa: F403
+1 -1
View File
@@ -5,7 +5,7 @@ import torch
from torch import nn
from fastvideo.third_party.eval.synchformer.ast import AST
from fastvideo.third_party.eval.synchformer.motionformer import MotionFormer
from fastvideo.third_party.synchformer.motionformer import MotionFormer
from fastvideo.third_party.eval.synchformer.transformer import GlobalTransformer
+2 -68
View File
@@ -1,69 +1,3 @@
from hashlib import md5
from pathlib import Path
"""Backward-compatible import for shared Synchformer utilities."""
import requests
from tqdm import tqdm
PARENT_LINK = 'https://a3s.fi/swift/v1/AUTH_a235c0f452d648828f745589cde1219a'
FNAME2LINK = {
# S3: Synchability: AudioSet (run 2)
'24-01-22T20-34-52.pt': f'{PARENT_LINK}/sync/sync_models/24-01-22T20-34-52/24-01-22T20-34-52.pt',
'cfg-24-01-22T20-34-52.yaml': f'{PARENT_LINK}/sync/sync_models/24-01-22T20-34-52/cfg-24-01-22T20-34-52.yaml',
# S2: Synchformer: AudioSet (run 2)
'24-01-04T16-39-21.pt': f'{PARENT_LINK}/sync/sync_models/24-01-04T16-39-21/24-01-04T16-39-21.pt',
'cfg-24-01-04T16-39-21.yaml': f'{PARENT_LINK}/sync/sync_models/24-01-04T16-39-21/cfg-24-01-04T16-39-21.yaml',
# S2: Synchformer: AudioSet (run 1)
'23-08-28T11-23-23.pt': f'{PARENT_LINK}/sync/sync_models/23-08-28T11-23-23/23-08-28T11-23-23.pt',
'cfg-23-08-28T11-23-23.yaml': f'{PARENT_LINK}/sync/sync_models/23-08-28T11-23-23/cfg-23-08-28T11-23-23.yaml',
# S2: Synchformer: LRS3 (run 2)
'23-12-23T18-33-57.pt': f'{PARENT_LINK}/sync/sync_models/23-12-23T18-33-57/23-12-23T18-33-57.pt',
'cfg-23-12-23T18-33-57.yaml': f'{PARENT_LINK}/sync/sync_models/23-12-23T18-33-57/cfg-23-12-23T18-33-57.yaml',
# S2: Synchformer: VGS (run 2)
'24-01-02T10-00-53.pt': f'{PARENT_LINK}/sync/sync_models/24-01-02T10-00-53/24-01-02T10-00-53.pt',
'cfg-24-01-02T10-00-53.yaml': f'{PARENT_LINK}/sync/sync_models/24-01-02T10-00-53/cfg-24-01-02T10-00-53.yaml',
# SparseSync: ft VGGSound-Full
'22-09-21T21-00-52.pt': f'{PARENT_LINK}/sync/sync_models/22-09-21T21-00-52/22-09-21T21-00-52.pt',
'cfg-22-09-21T21-00-52.yaml': f'{PARENT_LINK}/sync/sync_models/22-09-21T21-00-52/cfg-22-09-21T21-00-52.yaml',
# SparseSync: ft VGGSound-Sparse
'22-07-28T15-49-45.pt': f'{PARENT_LINK}/sync/sync_models/22-07-28T15-49-45/22-07-28T15-49-45.pt',
'cfg-22-07-28T15-49-45.yaml': f'{PARENT_LINK}/sync/sync_models/22-07-28T15-49-45/cfg-22-07-28T15-49-45.yaml',
# SparseSync: only pt on LRS3
'22-07-13T22-25-49.pt': f'{PARENT_LINK}/sync/sync_models/22-07-13T22-25-49/22-07-13T22-25-49.pt',
'cfg-22-07-13T22-25-49.yaml': f'{PARENT_LINK}/sync/sync_models/22-07-13T22-25-49/cfg-22-07-13T22-25-49.yaml',
# SparseSync: feature extractors
'ResNetAudio-22-08-04T09-51-04.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-08-04T09-51-04.pt', # 2s
'ResNetAudio-22-08-03T23-14-49.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-08-03T23-14-49.pt', # 3s
'ResNetAudio-22-08-03T23-14-28.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-08-03T23-14-28.pt', # 4s
'ResNetAudio-22-06-24T08-10-33.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T08-10-33.pt', # 5s
'ResNetAudio-22-06-24T17-31-07.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T17-31-07.pt', # 6s
'ResNetAudio-22-06-24T23-57-11.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T23-57-11.pt', # 7s
'ResNetAudio-22-06-25T04-35-42.pt': f'{PARENT_LINK}/sync/ResNetAudio-22-06-25T04-35-42.pt', # 8s
}
def check_if_file_exists_else_download(path, fname2link=FNAME2LINK, chunk_size=1024):
'''Checks if file exists, if not downloads it from the link to the path'''
path = Path(path)
if not path.exists():
path.parent.mkdir(exist_ok=True, parents=True)
link = fname2link.get(path.name, None)
if link is None:
raise ValueError(f'Cant find the checkpoint file: {path}.',
f'Please download it manually and ensure the path exists.')
with requests.get(fname2link[path.name], stream=True) as r:
total_size = int(r.headers.get('content-length', 0))
with tqdm(total=total_size, unit='B', unit_scale=True) as pbar:
with open(path, 'wb') as f:
for data in r.iter_content(chunk_size=chunk_size):
if data:
f.write(data)
pbar.update(chunk_size)
def get_md5sum(path):
hash_md5 = md5()
with open(path, 'rb') as f:
for chunk in iter(lambda: f.read(4096 * 8), b''):
hash_md5.update(chunk)
md5sum = hash_md5.hexdigest()
return md5sum
from fastvideo.third_party.synchformer.utils import * # noqa: F403
+2 -271
View File
@@ -1,272 +1,3 @@
#!/usr/bin/env python3
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
# Copyright 2020 Ross Wightman
# Modified Model definition
"""Backward-compatible import for shared Synchformer model builders."""
from collections import OrderedDict
from functools import partial
import torch
import torch.nn as nn
from timm.layers import trunc_normal_
from fastvideo.third_party.eval.synchformer import vit_helper
class VisionTransformer(nn.Module):
""" Vision Transformer with support for patch or hybrid CNN input stage """
def __init__(self, cfg):
super().__init__()
self.img_size = cfg.DATA.TRAIN_CROP_SIZE
self.patch_size = cfg.VIT.PATCH_SIZE
self.in_chans = cfg.VIT.CHANNELS
if cfg.TRAIN.DATASET == "Epickitchens":
self.num_classes = [97, 300]
else:
self.num_classes = cfg.MODEL.NUM_CLASSES
self.embed_dim = cfg.VIT.EMBED_DIM
self.depth = cfg.VIT.DEPTH
self.num_heads = cfg.VIT.NUM_HEADS
self.mlp_ratio = cfg.VIT.MLP_RATIO
self.qkv_bias = cfg.VIT.QKV_BIAS
self.drop_rate = cfg.VIT.DROP
self.drop_path_rate = cfg.VIT.DROP_PATH
self.head_dropout = cfg.VIT.HEAD_DROPOUT
self.video_input = cfg.VIT.VIDEO_INPUT
self.temporal_resolution = cfg.VIT.TEMPORAL_RESOLUTION
self.use_mlp = cfg.VIT.USE_MLP
self.num_features = self.embed_dim
norm_layer = partial(nn.LayerNorm, eps=1e-6)
self.attn_drop_rate = cfg.VIT.ATTN_DROPOUT
self.head_act = cfg.VIT.HEAD_ACT
self.cfg = cfg
# Patch Embedding
self.patch_embed = vit_helper.PatchEmbed(img_size=224,
patch_size=self.patch_size,
in_chans=self.in_chans,
embed_dim=self.embed_dim)
# 3D Patch Embedding
self.patch_embed_3d = vit_helper.PatchEmbed3D(img_size=self.img_size,
temporal_resolution=self.temporal_resolution,
patch_size=self.patch_size,
in_chans=self.in_chans,
embed_dim=self.embed_dim,
z_block_size=self.cfg.VIT.PATCH_SIZE_TEMP)
self.patch_embed_3d.proj.weight.data = torch.zeros_like(self.patch_embed_3d.proj.weight.data)
# Number of patches
if self.video_input:
num_patches = self.patch_embed.num_patches * self.temporal_resolution
else:
num_patches = self.patch_embed.num_patches
self.num_patches = num_patches
# CLS token
self.cls_token = nn.Parameter(torch.zeros(1, 1, self.embed_dim))
trunc_normal_(self.cls_token, std=.02)
# Positional embedding
self.pos_embed = nn.Parameter(torch.zeros(1, self.patch_embed.num_patches + 1, self.embed_dim))
self.pos_drop = nn.Dropout(p=cfg.VIT.POS_DROPOUT)
trunc_normal_(self.pos_embed, std=.02)
if self.cfg.VIT.POS_EMBED == "joint":
self.st_embed = nn.Parameter(torch.zeros(1, num_patches + 1, self.embed_dim))
trunc_normal_(self.st_embed, std=.02)
elif self.cfg.VIT.POS_EMBED == "separate":
self.temp_embed = nn.Parameter(torch.zeros(1, self.temporal_resolution, self.embed_dim))
# Layer Blocks
dpr = [x.item() for x in torch.linspace(0, self.drop_path_rate, self.depth)]
if self.cfg.VIT.ATTN_LAYER == "divided":
self.blocks = nn.ModuleList([
vit_helper.DividedSpaceTimeBlock(
attn_type=cfg.VIT.ATTN_LAYER,
dim=self.embed_dim,
num_heads=self.num_heads,
mlp_ratio=self.mlp_ratio,
qkv_bias=self.qkv_bias,
drop=self.drop_rate,
attn_drop=self.attn_drop_rate,
drop_path=dpr[i],
norm_layer=norm_layer,
) for i in range(self.depth)
])
else:
self.blocks = nn.ModuleList([
vit_helper.Block(attn_type=cfg.VIT.ATTN_LAYER,
dim=self.embed_dim,
num_heads=self.num_heads,
mlp_ratio=self.mlp_ratio,
qkv_bias=self.qkv_bias,
drop=self.drop_rate,
attn_drop=self.attn_drop_rate,
drop_path=dpr[i],
norm_layer=norm_layer,
use_original_code=self.cfg.VIT.USE_ORIGINAL_TRAJ_ATTN_CODE) for i in range(self.depth)
])
self.norm = norm_layer(self.embed_dim)
# MLP head
if self.use_mlp:
hidden_dim = self.embed_dim
if self.head_act == 'tanh':
# logging.info("Using TanH activation in MLP")
act = nn.Tanh()
elif self.head_act == 'gelu':
# logging.info("Using GELU activation in MLP")
act = nn.GELU()
else:
# logging.info("Using ReLU activation in MLP")
act = nn.ReLU()
self.pre_logits = nn.Sequential(OrderedDict([
('fc', nn.Linear(self.embed_dim, hidden_dim)),
('act', act),
]))
else:
self.pre_logits = nn.Identity()
# Classifier Head
self.head_drop = nn.Dropout(p=self.head_dropout)
if isinstance(self.num_classes, (list, )) and len(self.num_classes) > 1:
for a, i in enumerate(range(len(self.num_classes))):
setattr(self, "head%d" % a, nn.Linear(self.embed_dim, self.num_classes[i]))
else:
self.head = nn.Linear(self.embed_dim, self.num_classes) if self.num_classes > 0 else nn.Identity()
# Initialize weights
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
@torch.jit.ignore
def no_weight_decay(self):
if self.cfg.VIT.POS_EMBED == "joint":
return {'pos_embed', 'cls_token', 'st_embed'}
else:
return {'pos_embed', 'cls_token', 'temp_embed'}
def get_classifier(self):
return self.head
def reset_classifier(self, num_classes, global_pool=''):
self.num_classes = num_classes
self.head = (nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity())
def forward_features(self, x):
# if self.video_input:
# x = x[0]
B = x.shape[0]
# Tokenize input
# if self.cfg.VIT.PATCH_SIZE_TEMP > 1:
# for simplicity of mapping between content dimensions (input x) and token dims (after patching)
# we use the same trick as for AST (see modeling_ast.ASTModel.forward for the details):
# apply patching on input
x = self.patch_embed_3d(x)
tok_mask = None
# else:
# tok_mask = None
# # 2D tokenization
# if self.video_input:
# x = x.permute(0, 2, 1, 3, 4)
# (B, T, C, H, W) = x.shape
# x = x.reshape(B * T, C, H, W)
# x = self.patch_embed(x)
# if self.video_input:
# (B2, T2, D2) = x.shape
# x = x.reshape(B, T * T2, D2)
# Append CLS token
cls_tokens = self.cls_token.expand(B, -1, -1)
x = torch.cat((cls_tokens, x), dim=1)
# if tok_mask is not None:
# # prepend 1(=keep) to the mask to account for the CLS token as well
# tok_mask = torch.cat((torch.ones_like(tok_mask[:, [0]]), tok_mask), dim=1)
# Interpolate positinoal embeddings
# if self.cfg.DATA.TRAIN_CROP_SIZE != 224:
# pos_embed = self.pos_embed
# N = pos_embed.shape[1] - 1
# npatch = int((x.size(1) - 1) / self.temporal_resolution)
# class_emb = pos_embed[:, 0]
# pos_embed = pos_embed[:, 1:]
# dim = x.shape[-1]
# pos_embed = torch.nn.functional.interpolate(
# pos_embed.reshape(1, int(math.sqrt(N)), int(math.sqrt(N)), dim).permute(0, 3, 1, 2),
# scale_factor=math.sqrt(npatch / N),
# mode='bicubic',
# )
# pos_embed = pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)
# new_pos_embed = torch.cat((class_emb.unsqueeze(0), pos_embed), dim=1)
# else:
new_pos_embed = self.pos_embed
npatch = self.patch_embed.num_patches
# Add positional embeddings to input
if self.video_input:
if self.cfg.VIT.POS_EMBED == "separate":
cls_embed = self.pos_embed[:, 0, :].unsqueeze(1)
tile_pos_embed = new_pos_embed[:, 1:, :].repeat(1, self.temporal_resolution, 1)
tile_temporal_embed = self.temp_embed.repeat_interleave(npatch, 1)
total_pos_embed = tile_pos_embed + tile_temporal_embed
total_pos_embed = torch.cat([cls_embed, total_pos_embed], dim=1)
x = x + total_pos_embed
elif self.cfg.VIT.POS_EMBED == "joint":
x = x + self.st_embed
else:
# image input
x = x + new_pos_embed
# Apply positional dropout
x = self.pos_drop(x)
# Encoding using transformer layers
for i, blk in enumerate(self.blocks):
x = blk(x,
seq_len=npatch,
num_frames=self.temporal_resolution,
approx=self.cfg.VIT.APPROX_ATTN_TYPE,
num_landmarks=self.cfg.VIT.APPROX_ATTN_DIM,
tok_mask=tok_mask)
### v-iashin: I moved it to the forward pass
# x = self.norm(x)[:, 0]
# x = self.pre_logits(x)
###
return x, tok_mask
# def forward(self, x):
# x = self.forward_features(x)
# ### v-iashin: here. This should leave the same forward output as before
# x = self.norm(x)[:, 0]
# x = self.pre_logits(x)
# ###
# x = self.head_drop(x)
# if isinstance(self.num_classes, (list, )) and len(self.num_classes) > 1:
# output = []
# for head in range(len(self.num_classes)):
# x_out = getattr(self, "head%d" % head)(x)
# if not self.training:
# x_out = torch.nn.functional.softmax(x_out, dim=-1)
# output.append(x_out)
# return output
# else:
# x = self.head(x)
# if not self.training:
# x = torch.nn.functional.softmax(x, dim=-1)
# return x
from fastvideo.third_party.synchformer.video_model_builder import * # noqa: F403
+2 -361
View File
@@ -1,362 +1,3 @@
#!/usr/bin/env python3
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
# Copyright 2020 Ross Wightman
# Modified Model definition
"""Video models."""
"""Backward-compatible import for shared Synchformer ViT helpers."""
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from timm.layers import to_2tuple
from torch import einsum
from torch.nn import functional as F
default_cfgs = {
'vit_1k':
'https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_base_p16_224-80ecf9dd.pth',
'vit_1k_large':
'https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_large_p16_224-4ee7a4dc.pth',
}
def qkv_attn(q, k, v, tok_mask: torch.Tensor = None):
sim = einsum('b i d, b j d -> b i j', q, k)
# apply masking if provided, tok_mask is (B*S*H, N): 1s - keep; sim is (B*S*H, H, N, N)
if tok_mask is not None:
BSH, N = tok_mask.shape
sim = sim.masked_fill(tok_mask.view(BSH, 1, N) == 0, float('-inf')) # 1 - broadcasts across N
attn = sim.softmax(dim=-1)
out = einsum('b i j, b j d -> b i d', attn, v)
return out
class DividedAttention(nn.Module):
def __init__(self, dim, num_heads=8, qkv_bias=False, attn_drop=0., proj_drop=0.):
super().__init__()
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim**-0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.proj = nn.Linear(dim, dim)
# init to zeros
self.qkv.weight.data.fill_(0)
self.qkv.bias.data.fill_(0)
self.proj.weight.data.fill_(1)
self.proj.bias.data.fill_(0)
self.attn_drop = nn.Dropout(attn_drop)
self.proj_drop = nn.Dropout(proj_drop)
def forward(self, x, einops_from, einops_to, tok_mask: torch.Tensor = None, **einops_dims):
# num of heads variable
h = self.num_heads
# project x to q, k, v vaalues
q, k, v = self.qkv(x).chunk(3, dim=-1)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
if tok_mask is not None:
# replicate token mask across heads (b, n) -> (b, h, n) -> (b*h, n) -- same as qkv but w/o d
assert len(tok_mask.shape) == 2
tok_mask = tok_mask.unsqueeze(1).expand(-1, h, -1).reshape(-1, tok_mask.shape[1])
# Scale q
q *= self.scale
# Take out cls_q, cls_k, cls_v
(cls_q, q_), (cls_k, k_), (cls_v, v_) = map(lambda t: (t[:, 0:1], t[:, 1:]), (q, k, v))
# the same for masking
if tok_mask is not None:
cls_mask, mask_ = tok_mask[:, 0:1], tok_mask[:, 1:]
else:
cls_mask, mask_ = None, None
# let CLS token attend to key / values of all patches across time and space
cls_out = qkv_attn(cls_q, k, v, tok_mask=tok_mask)
# rearrange across time or space
q_, k_, v_ = map(lambda t: rearrange(t, f'{einops_from} -> {einops_to}', **einops_dims), (q_, k_, v_))
# expand CLS token keys and values across time or space and concat
r = q_.shape[0] // cls_k.shape[0]
cls_k, cls_v = map(lambda t: repeat(t, 'b () d -> (b r) () d', r=r), (cls_k, cls_v))
k_ = torch.cat((cls_k, k_), dim=1)
v_ = torch.cat((cls_v, v_), dim=1)
# the same for masking (if provided)
if tok_mask is not None:
# since mask does not have the latent dim (d), we need to remove it from einops dims
mask_ = rearrange(mask_, f'{einops_from} -> {einops_to}'.replace(' d', ''), **einops_dims)
cls_mask = repeat(cls_mask, 'b () -> (b r) ()', r=r) # expand cls_mask across time or space
mask_ = torch.cat((cls_mask, mask_), dim=1)
# attention
out = qkv_attn(q_, k_, v_, tok_mask=mask_)
# merge back time or space
out = rearrange(out, f'{einops_to} -> {einops_from}', **einops_dims)
# concat back the cls token
out = torch.cat((cls_out, out), dim=1)
# merge back the heads
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
## to out
x = self.proj(out)
x = self.proj_drop(x)
return x
class DividedSpaceTimeBlock(nn.Module):
def __init__(self,
dim=768,
num_heads=12,
attn_type='divided',
mlp_ratio=4.,
qkv_bias=False,
drop=0.,
attn_drop=0.,
drop_path=0.,
act_layer=nn.GELU,
norm_layer=nn.LayerNorm):
super().__init__()
self.einops_from_space = 'b (f n) d'
self.einops_to_space = '(b f) n d'
self.einops_from_time = 'b (f n) d'
self.einops_to_time = '(b n) f d'
self.norm1 = norm_layer(dim)
self.attn = DividedAttention(dim, num_heads=num_heads, qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop)
self.timeattn = DividedAttention(dim,
num_heads=num_heads,
qkv_bias=qkv_bias,
attn_drop=attn_drop,
proj_drop=drop)
# self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
self.drop_path = nn.Identity()
self.norm2 = norm_layer(dim)
mlp_hidden_dim = int(dim * mlp_ratio)
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
self.norm3 = norm_layer(dim)
def forward(self, x, seq_len=196, num_frames=8, approx='none', num_landmarks=128, tok_mask: torch.Tensor = None):
time_output = self.timeattn(self.norm3(x),
self.einops_from_time,
self.einops_to_time,
n=seq_len,
tok_mask=tok_mask)
time_residual = x + time_output
space_output = self.attn(self.norm1(time_residual),
self.einops_from_space,
self.einops_to_space,
f=num_frames,
tok_mask=tok_mask)
space_residual = time_residual + self.drop_path(space_output)
x = space_residual
x = x + self.drop_path(self.mlp(self.norm2(x)))
return x
class Mlp(nn.Module):
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
class PatchEmbed(nn.Module):
""" Image to Patch Embedding
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
img_size = img_size if type(img_size) is tuple else to_2tuple(img_size)
patch_size = img_size if type(patch_size) is tuple else to_2tuple(patch_size)
num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = num_patches
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x):
B, C, H, W = x.shape
x = self.proj(x).flatten(2).transpose(1, 2)
return x
class PatchEmbed3D(nn.Module):
""" Image to Patch Embedding """
def __init__(self,
img_size=224,
temporal_resolution=4,
in_chans=3,
patch_size=16,
z_block_size=2,
embed_dim=768,
flatten=True):
super().__init__()
self.height = (img_size // patch_size)
self.width = (img_size // patch_size)
### v-iashin: these two are incorrect
# self.frames = (temporal_resolution // z_block_size)
# self.num_patches = self.height * self.width * self.frames
self.z_block_size = z_block_size
###
self.proj = nn.Conv3d(in_chans,
embed_dim,
kernel_size=(z_block_size, patch_size, patch_size),
stride=(z_block_size, patch_size, patch_size))
self.flatten = flatten
def forward(self, x):
B, C, T, H, W = x.shape
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2)
return x
class HeadMLP(nn.Module):
def __init__(self, n_input, n_classes, n_hidden=512, p=0.1):
super(HeadMLP, self).__init__()
self.n_input = n_input
self.n_classes = n_classes
self.n_hidden = n_hidden
if n_hidden is None:
# use linear classifier
self.block_forward = nn.Sequential(nn.Dropout(p=p), nn.Linear(n_input, n_classes, bias=True))
else:
# use simple MLP classifier
self.block_forward = nn.Sequential(nn.Dropout(p=p), nn.Linear(n_input, n_hidden, bias=True),
nn.BatchNorm1d(n_hidden), nn.ReLU(inplace=True), nn.Dropout(p=p),
nn.Linear(n_hidden, n_classes, bias=True))
print(f"Dropout-NLP: {p}")
def forward(self, x):
return self.block_forward(x)
def _conv_filter(state_dict, patch_size=16):
""" convert patch embedding weight from manual patchify + linear proj to conv"""
out_dict = {}
for k, v in state_dict.items():
if 'patch_embed.proj.weight' in k:
v = v.reshape((v.shape[0], 3, patch_size, patch_size))
out_dict[k] = v
return out_dict
def adapt_input_conv(in_chans, conv_weight, agg='sum'):
conv_type = conv_weight.dtype
conv_weight = conv_weight.float()
O, I, J, K = conv_weight.shape
if in_chans == 1:
if I > 3:
assert conv_weight.shape[1] % 3 == 0
# For models with space2depth stems
conv_weight = conv_weight.reshape(O, I // 3, 3, J, K)
conv_weight = conv_weight.sum(dim=2, keepdim=False)
else:
if agg == 'sum':
print("Summing conv1 weights")
conv_weight = conv_weight.sum(dim=1, keepdim=True)
else:
print("Averaging conv1 weights")
conv_weight = conv_weight.mean(dim=1, keepdim=True)
elif in_chans != 3:
if I != 3:
raise NotImplementedError('Weight format not supported by conversion.')
else:
if agg == 'sum':
print("Summing conv1 weights")
repeat = int(math.ceil(in_chans / 3))
conv_weight = conv_weight.repeat(1, repeat, 1, 1)[:, :in_chans, :, :]
conv_weight *= (3 / float(in_chans))
else:
print("Averaging conv1 weights")
conv_weight = conv_weight.mean(dim=1, keepdim=True)
conv_weight = conv_weight.repeat(1, in_chans, 1, 1)
conv_weight = conv_weight.to(conv_type)
return conv_weight
def load_pretrained(model, cfg=None, num_classes=1000, in_chans=3, filter_fn=None, strict=True, progress=False):
# Load state dict
assert (f"{cfg.VIT.PRETRAINED_WEIGHTS} not in [vit_1k, vit_1k_large]")
state_dict = torch.hub.load_state_dict_from_url(url=default_cfgs[cfg.VIT.PRETRAINED_WEIGHTS])
if filter_fn is not None:
state_dict = filter_fn(state_dict)
input_convs = 'patch_embed.proj'
if input_convs is not None and in_chans != 3:
if isinstance(input_convs, str):
input_convs = (input_convs, )
for input_conv_name in input_convs:
weight_name = input_conv_name + '.weight'
try:
state_dict[weight_name] = adapt_input_conv(in_chans, state_dict[weight_name], agg='avg')
print(f'Converted input conv {input_conv_name} pretrained weights from 3 to {in_chans} channel(s)')
except NotImplementedError as e:
del state_dict[weight_name]
strict = False
print(f'Unable to convert pretrained {input_conv_name} weights, using random init for this layer.')
classifier_name = 'head'
label_offset = cfg.get('label_offset', 0)
pretrain_classes = 1000
if num_classes != pretrain_classes:
# completely discard fully connected if model num_classes doesn't match pretrained weights
del state_dict[classifier_name + '.weight']
del state_dict[classifier_name + '.bias']
strict = False
elif label_offset > 0:
# special case for pretrained weights with an extra background class in pretrained weights
classifier_weight = state_dict[classifier_name + '.weight']
state_dict[classifier_name + '.weight'] = classifier_weight[label_offset:]
classifier_bias = state_dict[classifier_name + '.bias']
state_dict[classifier_name + '.bias'] = classifier_bias[label_offset:]
loaded_state = state_dict
self_state = model.state_dict()
all_names = set(self_state.keys())
saved_names = set([])
for name, param in loaded_state.items():
param = param
if 'module.' in name:
name = name.replace('module.', '')
if name in self_state.keys() and param.shape == self_state[name].shape:
saved_names.add(name)
self_state[name].copy_(param)
else:
print(f"didnt load: {name} of shape: {param.shape}")
print("Missing Keys:")
print(all_names - saved_names)
from fastvideo.third_party.synchformer.vit_helper import * # noqa: F403
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 Vladimir Iashin
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+5
View File
@@ -0,0 +1,5 @@
"""Shared vendored Synchformer visual backbone."""
from fastvideo.third_party.synchformer.motionformer import MotionFormer
__all__ = ["MotionFormer"]
@@ -1,3 +1,4 @@
# Shared MotionFormer configuration used by Synchformer.
TRAIN:
ENABLE: True
DATASET: Ssv2
+400
View File
@@ -0,0 +1,400 @@
import logging
from pathlib import Path
import einops
import torch
from omegaconf import OmegaConf
from timm.layers import trunc_normal_
from torch import nn
from fastvideo.third_party.synchformer.utils import check_if_file_exists_else_download
from fastvideo.third_party.synchformer.video_model_builder import VisionTransformer
FILE2URL = {
# cfg
'motionformer_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/motionformer_224_16x4.yaml',
'joint_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/joint_224_16x4.yaml',
'divided_224_16x4.yaml':
'https://raw.githubusercontent.com/facebookresearch/Motionformer/bf43d50/configs/SSV2/divided_224_16x4.yaml',
# ckpt
'ssv2_motionformer_224_16x4.pyth':
'https://dl.fbaipublicfiles.com/motionformer/ssv2_motionformer_224_16x4.pyth',
'ssv2_joint_224_16x4.pyth':
'https://dl.fbaipublicfiles.com/motionformer/ssv2_joint_224_16x4.pyth',
'ssv2_divided_224_16x4.pyth':
'https://dl.fbaipublicfiles.com/motionformer/ssv2_divided_224_16x4.pyth',
}
class MotionFormer(VisionTransformer):
''' This class serves three puposes:
1. Renames the class to MotionFormer.
2. Downloads the cfg from the original repo and patches it if needed.
3. Takes care of feature extraction by redefining .forward()
- if `extract_features=True` and `factorize_space_time=False`,
the output is of shape (B, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
- if `extract_features=True` and `factorize_space_time=True`, the output is of shape (B*S, D)
and spatial and temporal transformer encoder layers are used.
- if `extract_features=True` and `factorize_space_time=True` as well as `add_global_repr=True`
the output is of shape (B, D) and spatial and temporal transformer encoder layers
are used as well as the global representation is extracted from segments (extra pos emb
is added).
'''
def __init__(
self,
extract_features: bool = False,
ckpt_path: str = None,
factorize_space_time: bool = None,
agg_space_module: str = None,
agg_time_module: str = None,
add_global_repr: bool = True,
agg_segments_module: str = None,
max_segments: int = None,
):
self.extract_features = extract_features
self.ckpt_path = ckpt_path
self.factorize_space_time = factorize_space_time
if self.ckpt_path is not None:
check_if_file_exists_else_download(self.ckpt_path, FILE2URL)
ckpt = torch.load(self.ckpt_path, map_location='cpu')
mformer_ckpt2cfg = {
'ssv2_motionformer_224_16x4.pyth': 'motionformer_224_16x4.yaml',
'ssv2_joint_224_16x4.pyth': 'joint_224_16x4.yaml',
'ssv2_divided_224_16x4.pyth': 'divided_224_16x4.yaml',
}
# init from motionformer ckpt or from our Stage I ckpt
# depending on whether the feat extractor was pre-trained on AVCLIPMoCo or not, we need to
# load the state dict differently
was_pt_on_avclip = self.ckpt_path.endswith(
'.pt') # checks if it is a stage I ckpt (FIXME: a bit generic)
if self.ckpt_path.endswith(tuple(mformer_ckpt2cfg.keys())):
cfg_fname = mformer_ckpt2cfg[Path(self.ckpt_path).name]
elif was_pt_on_avclip:
# TODO: this is a hack, we should be able to get the cfg from the ckpt (earlier ckpt didn't have it)
s1_cfg = ckpt.get('args', None) # Stage I cfg
if s1_cfg is not None:
s1_vfeat_extractor_ckpt_path = s1_cfg.model.params.vfeat_extractor.params.ckpt_path
# if the stage I ckpt was initialized from a motionformer ckpt or train from scratch
if s1_vfeat_extractor_ckpt_path is not None:
cfg_fname = mformer_ckpt2cfg[Path(s1_vfeat_extractor_ckpt_path).name]
else:
cfg_fname = 'divided_224_16x4.yaml'
else:
cfg_fname = 'divided_224_16x4.yaml'
else:
raise ValueError(f'ckpt_path {self.ckpt_path} is not supported.')
else:
was_pt_on_avclip = False
cfg_fname = 'divided_224_16x4.yaml'
# logging.info(f'No ckpt_path provided, using {cfg_fname} config.')
if cfg_fname in ['motionformer_224_16x4.yaml', 'divided_224_16x4.yaml']:
pos_emb_type = 'separate'
elif cfg_fname == 'joint_224_16x4.yaml':
pos_emb_type = 'joint'
self.mformer_cfg_path = Path(__file__).absolute().parent / cfg_fname
check_if_file_exists_else_download(self.mformer_cfg_path, FILE2URL)
mformer_cfg = OmegaConf.load(self.mformer_cfg_path)
logging.info(f'Loading MotionFormer config from {self.mformer_cfg_path.absolute()}')
# patch the cfg (from the default cfg defined in the repo `Motionformer/slowfast/config/defaults.py`)
mformer_cfg.VIT.ATTN_DROPOUT = 0.0
mformer_cfg.VIT.POS_EMBED = pos_emb_type
mformer_cfg.VIT.USE_ORIGINAL_TRAJ_ATTN_CODE = True
mformer_cfg.VIT.APPROX_ATTN_TYPE = 'none' # guessing
mformer_cfg.VIT.APPROX_ATTN_DIM = 64 # from ckpt['cfg']
# finally init VisionTransformer with the cfg
super().__init__(mformer_cfg)
# load the ckpt now if ckpt is provided and not from AVCLIPMoCo-pretrained ckpt
if (self.ckpt_path is not None) and (not was_pt_on_avclip):
_ckpt_load_status = self.load_state_dict(ckpt['model_state'], strict=False)
if len(_ckpt_load_status.missing_keys) > 0 or len(
_ckpt_load_status.unexpected_keys) > 0:
logging.warning(f'Loading exact vfeat_extractor ckpt from {self.ckpt_path} failed.' \
f'Missing keys: {_ckpt_load_status.missing_keys}, ' \
f'Unexpected keys: {_ckpt_load_status.unexpected_keys}')
else:
logging.info(f'Loading vfeat_extractor ckpt from {self.ckpt_path} succeeded.')
if self.extract_features:
assert isinstance(self.norm,
nn.LayerNorm), 'early x[:, 1:, :] may not be safe for per-tr weights'
# pre-logits are Sequential(nn.Linear(emb, emd), act) and `act` is tanh but see the logger
self.pre_logits = nn.Identity()
# we don't need the classification head (saving memory)
self.head = nn.Identity()
self.head_drop = nn.Identity()
# avoiding code duplication (used only if agg_*_module is TransformerEncoderLayer)
transf_enc_layer_kwargs = dict(
d_model=self.embed_dim,
nhead=self.num_heads,
activation=nn.GELU(),
batch_first=True,
dim_feedforward=self.mlp_ratio * self.embed_dim,
dropout=self.drop_rate,
layer_norm_eps=1e-6,
norm_first=True,
)
# define adapters if needed
if self.factorize_space_time:
if agg_space_module == 'TransformerEncoderLayer':
self.spatial_attn_agg = SpatialTransformerEncoderLayer(
**transf_enc_layer_kwargs)
elif agg_space_module == 'AveragePooling':
self.spatial_attn_agg = AveragePooling(avg_pattern='BS D t h w -> BS D t',
then_permute_pattern='BS D t -> BS t D')
if agg_time_module == 'TransformerEncoderLayer':
self.temp_attn_agg = TemporalTransformerEncoderLayer(**transf_enc_layer_kwargs)
elif agg_time_module == 'AveragePooling':
self.temp_attn_agg = AveragePooling(avg_pattern='BS t D -> BS D')
elif 'Identity' in agg_time_module:
self.temp_attn_agg = nn.Identity()
# define a global aggregation layer (aggregarate over segments)
self.add_global_repr = add_global_repr
if add_global_repr:
if agg_segments_module == 'TransformerEncoderLayer':
# we can reuse the same layer as for temporal factorization (B, dim_to_agg, D) -> (B, D)
# we need to add pos emb (PE) because previously we added the same PE for each segment
pos_max_len = max_segments if max_segments is not None else 16 # 16 = 10sec//0.64sec + 1
self.global_attn_agg = TemporalTransformerEncoderLayer(
add_pos_emb=True,
pos_emb_drop=mformer_cfg.VIT.POS_DROPOUT,
pos_max_len=pos_max_len,
**transf_enc_layer_kwargs)
elif agg_segments_module == 'AveragePooling':
self.global_attn_agg = AveragePooling(avg_pattern='B S D -> B D')
if was_pt_on_avclip:
# we need to filter out the state_dict of the AVCLIP model (has both A and V extractors)
# and keep only the state_dict of the feat extractor
ckpt_weights = dict()
for k, v in ckpt['state_dict'].items():
if k.startswith(('module.v_encoder.', 'v_encoder.')):
k = k.replace('module.', '').replace('v_encoder.', '')
ckpt_weights[k] = v
_load_status = self.load_state_dict(ckpt_weights, strict=False)
if len(_load_status.missing_keys) > 0 or len(_load_status.unexpected_keys) > 0:
logging.warning(f'Loading exact vfeat_extractor ckpt from {self.ckpt_path} failed. \n' \
f'Missing keys ({len(_load_status.missing_keys)}): ' \
f'{_load_status.missing_keys}, \n' \
f'Unexpected keys ({len(_load_status.unexpected_keys)}): ' \
f'{_load_status.unexpected_keys} \n' \
f'temp_attn_agg are expected to be missing if ckpt was pt contrastively.')
else:
logging.info(f'Loading vfeat_extractor ckpt from {self.ckpt_path} succeeded.')
# patch_embed is not used in MotionFormer, only patch_embed_3d, because cfg.VIT.PATCH_SIZE_TEMP > 1
# but it used to calculate the number of patches, so we need to set keep it
self.patch_embed.requires_grad_(False)
def forward(self, x):
'''
x is of shape (B, S, C, T, H, W) where S is the number of segments.
'''
# Batch, Segments, Channels, T=frames, Height, Width
B, S, C, T, H, W = x.shape
# Motionformer expects a tensor of shape (1, B, C, T, H, W).
# The first dimension (1) is a dummy dimension to make the input tensor and won't be used:
# see `video_model_builder.video_input`.
# x = x.unsqueeze(0) # (1, B, S, C, T, H, W)
orig_shape = (B, S, C, T, H, W)
x = x.view(B * S, C, T, H, W) # flatten batch and segments
x = self.forward_segments(x, orig_shape=orig_shape)
# unpack the segments (using rest dimensions to support different shapes e.g. (BS, D) or (BS, t, D))
x = x.view(B, S, *x.shape[1:])
# x is now of shape (B*S, D) or (B*S, t, D) if `self.temp_attn_agg` is `Identity`
return x # x is (B, S, ...)
def forward_segments(self, x, orig_shape: tuple) -> torch.Tensor:
'''x is of shape (1, BS, C, T, H, W) where S is the number of segments.'''
x, x_mask = self.forward_features(x)
assert self.extract_features
# (BS, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
x = x[:,
1:, :] # without the CLS token for efficiency (should be safe for LayerNorm and FC)
x = self.norm(x)
x = self.pre_logits(x)
if self.factorize_space_time:
x = self.restore_spatio_temp_dims(x, orig_shape) # (B*S, D, t, h, w) <- (B*S, t*h*w, D)
x = self.spatial_attn_agg(x, x_mask) # (B*S, t, D)
x = self.temp_attn_agg(
x) # (B*S, D) or (BS, t, D) if `self.temp_attn_agg` is `Identity`
return x
def restore_spatio_temp_dims(self, feats: torch.Tensor, orig_shape: tuple) -> torch.Tensor:
'''
feats are of shape (B*S, T, D) where T = 1 + (224 // 16) * (224 // 16) * 8
Our goal is to make them of shape (B*S, t, h, w, D) where h, w are the spatial dimensions.
From `self.patch_embed_3d`, it follows that we could reshape feats with:
`feats.transpose(1, 2).view(B*S, D, t, h, w)`
'''
B, S, C, T, H, W = orig_shape
D = self.embed_dim
# num patches in each dimension
t = T // self.patch_embed_3d.z_block_size
h = self.patch_embed_3d.height
w = self.patch_embed_3d.width
feats = feats.permute(0, 2, 1) # (B*S, D, T)
feats = feats.view(B * S, D, t, h, w) # (B*S, D, t, h, w)
return feats
class BaseEncoderLayer(nn.TransformerEncoderLayer):
'''
This is a wrapper around nn.TransformerEncoderLayer that adds a CLS token
to the sequence and outputs the CLS token's representation.
This base class parents both SpatialEncoderLayer and TemporalEncoderLayer for the RGB stream
and the FrequencyEncoderLayer and TemporalEncoderLayer for the audio stream stream.
We also, optionally, add a positional embedding to the input sequence which
allows to reuse it for global aggregation (of segments) for both streams.
'''
def __init__(self,
add_pos_emb: bool = False,
pos_emb_drop: float = None,
pos_max_len: int = None,
*args_transformer_enc,
**kwargs_transformer_enc):
super().__init__(*args_transformer_enc, **kwargs_transformer_enc)
self.cls_token = nn.Parameter(torch.zeros(1, 1, self.self_attn.embed_dim))
trunc_normal_(self.cls_token, std=.02)
# add positional embedding
self.add_pos_emb = add_pos_emb
if add_pos_emb:
self.pos_max_len = 1 + pos_max_len # +1 (for CLS)
self.pos_emb = nn.Parameter(torch.zeros(1, self.pos_max_len, self.self_attn.embed_dim))
self.pos_drop = nn.Dropout(pos_emb_drop)
trunc_normal_(self.pos_emb, std=.02)
self.apply(self._init_weights)
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None):
''' x is of shape (B, N, D); if provided x_mask is of shape (B, N)'''
batch_dim = x.shape[0]
# add CLS token
cls_tokens = self.cls_token.expand(batch_dim, -1, -1) # expanding to match batch dimension
x = torch.cat((cls_tokens, x), dim=-2) # (batch_dim, 1+seq_len, D)
if x_mask is not None:
cls_mask = torch.ones((batch_dim, 1), dtype=torch.bool,
device=x_mask.device) # 1=keep; 0=mask
x_mask_w_cls = torch.cat((cls_mask, x_mask), dim=-1) # (batch_dim, 1+seq_len)
B, N = x_mask_w_cls.shape
# torch expects (N, N) or (B*num_heads, N, N) mask (sadness ahead); torch masks
x_mask_w_cls = x_mask_w_cls.reshape(B, 1, 1, N)\
.expand(-1, self.self_attn.num_heads, N, -1)\
.reshape(B * self.self_attn.num_heads, N, N)
assert x_mask_w_cls.dtype == x_mask_w_cls.bool().dtype, 'x_mask_w_cls.dtype != bool'
x_mask_w_cls = ~x_mask_w_cls # invert mask (1=mask)
else:
x_mask_w_cls = None
# add positional embedding
if self.add_pos_emb:
seq_len = x.shape[
1] # (don't even think about moving it before the CLS token concatenation)
assert seq_len <= self.pos_max_len, f'Seq len ({seq_len}) > pos_max_len ({self.pos_max_len})'
x = x + self.pos_emb[:, :seq_len, :]
x = self.pos_drop(x)
# apply encoder layer (calls nn.TransformerEncoderLayer.forward);
x = super().forward(src=x, src_mask=x_mask_w_cls) # (batch_dim, 1+seq_len, D)
# CLS token is expected to hold spatial information for each frame
x = x[:, 0, :] # (batch_dim, D)
return x
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
@torch.jit.ignore
def no_weight_decay(self):
return {'cls_token', 'pos_emb'}
class SpatialTransformerEncoderLayer(BaseEncoderLayer):
''' Aggregates spatial dimensions by applying attention individually to each frame. '''
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None) -> torch.Tensor:
''' x is of shape (B*S, D, t, h, w) where S is the number of segments.
if specified x_mask (B*S, t, h, w), 0=masked, 1=kept
Returns a tensor of shape (B*S, t, D) pooling spatial information for each frame. '''
BS, D, t, h, w = x.shape
# time as a batch dimension and flatten spatial dimensions as sequence
x = einops.rearrange(x, 'BS D t h w -> (BS t) (h w) D')
# similar to mask
if x_mask is not None:
x_mask = einops.rearrange(x_mask, 'BS t h w -> (BS t) (h w)')
# apply encoder layer (BaseEncoderLayer.forward) - it will add CLS token and output its representation
x = super().forward(x=x, x_mask=x_mask) # (B*S*t, D)
# reshape back to (B*S, t, D)
x = einops.rearrange(x, '(BS t) D -> BS t D', BS=BS, t=t)
# (B*S, t, D)
return x
class TemporalTransformerEncoderLayer(BaseEncoderLayer):
''' Aggregates temporal dimension with attention. Also used with pos emb as global aggregation
in both streams. '''
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x):
''' x is of shape (B*S, t, D) where S is the number of segments.
Returns a tensor of shape (B*S, D) pooling temporal information. '''
BS, t, D = x.shape
# apply encoder layer (BaseEncoderLayer.forward) - it will add CLS token and output its representation
x = super().forward(x) # (B*S, D)
return x # (B*S, D)
class AveragePooling(nn.Module):
def __init__(self, avg_pattern: str, then_permute_pattern: str = None) -> None:
''' patterns are e.g. "bs t d -> bs d" '''
super().__init__()
# TODO: need to register them as buffers (but fails because these are strings)
self.reduce_fn = 'mean'
self.avg_pattern = avg_pattern
self.then_permute_pattern = then_permute_pattern
def forward(self, x: torch.Tensor, x_mask: torch.Tensor = None) -> torch.Tensor:
x = einops.reduce(x, self.avg_pattern, self.reduce_fn)
if self.then_permute_pattern is not None:
x = einops.rearrange(x, self.then_permute_pattern)
return x
+94
View File
@@ -0,0 +1,94 @@
"""Checkpoint utilities shared by Synchformer consumers."""
from hashlib import md5
from pathlib import Path
import requests
from tqdm import tqdm
PARENT_LINK = 'https://a3s.fi/swift/v1/AUTH_a235c0f452d648828f745589cde1219a'
FNAME2LINK = {
# S3: Synchability: AudioSet (run 2)
'24-01-22T20-34-52.pt':
f'{PARENT_LINK}/sync/sync_models/24-01-22T20-34-52/24-01-22T20-34-52.pt',
'cfg-24-01-22T20-34-52.yaml':
f'{PARENT_LINK}/sync/sync_models/24-01-22T20-34-52/cfg-24-01-22T20-34-52.yaml',
# S2: Synchformer: AudioSet (run 2)
'24-01-04T16-39-21.pt':
f'{PARENT_LINK}/sync/sync_models/24-01-04T16-39-21/24-01-04T16-39-21.pt',
'cfg-24-01-04T16-39-21.yaml':
f'{PARENT_LINK}/sync/sync_models/24-01-04T16-39-21/cfg-24-01-04T16-39-21.yaml',
# S2: Synchformer: AudioSet (run 1)
'23-08-28T11-23-23.pt':
f'{PARENT_LINK}/sync/sync_models/23-08-28T11-23-23/23-08-28T11-23-23.pt',
'cfg-23-08-28T11-23-23.yaml':
f'{PARENT_LINK}/sync/sync_models/23-08-28T11-23-23/cfg-23-08-28T11-23-23.yaml',
# S2: Synchformer: LRS3 (run 2)
'23-12-23T18-33-57.pt':
f'{PARENT_LINK}/sync/sync_models/23-12-23T18-33-57/23-12-23T18-33-57.pt',
'cfg-23-12-23T18-33-57.yaml':
f'{PARENT_LINK}/sync/sync_models/23-12-23T18-33-57/cfg-23-12-23T18-33-57.yaml',
# S2: Synchformer: VGS (run 2)
'24-01-02T10-00-53.pt':
f'{PARENT_LINK}/sync/sync_models/24-01-02T10-00-53/24-01-02T10-00-53.pt',
'cfg-24-01-02T10-00-53.yaml':
f'{PARENT_LINK}/sync/sync_models/24-01-02T10-00-53/cfg-24-01-02T10-00-53.yaml',
# SparseSync: ft VGGSound-Full
'22-09-21T21-00-52.pt':
f'{PARENT_LINK}/sync/sync_models/22-09-21T21-00-52/22-09-21T21-00-52.pt',
'cfg-22-09-21T21-00-52.yaml':
f'{PARENT_LINK}/sync/sync_models/22-09-21T21-00-52/cfg-22-09-21T21-00-52.yaml',
# SparseSync: ft VGGSound-Sparse
'22-07-28T15-49-45.pt':
f'{PARENT_LINK}/sync/sync_models/22-07-28T15-49-45/22-07-28T15-49-45.pt',
'cfg-22-07-28T15-49-45.yaml':
f'{PARENT_LINK}/sync/sync_models/22-07-28T15-49-45/cfg-22-07-28T15-49-45.yaml',
# SparseSync: only pt on LRS3
'22-07-13T22-25-49.pt':
f'{PARENT_LINK}/sync/sync_models/22-07-13T22-25-49/22-07-13T22-25-49.pt',
'cfg-22-07-13T22-25-49.yaml':
f'{PARENT_LINK}/sync/sync_models/22-07-13T22-25-49/cfg-22-07-13T22-25-49.yaml',
# SparseSync: feature extractors
'ResNetAudio-22-08-04T09-51-04.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-08-04T09-51-04.pt', # 2s
'ResNetAudio-22-08-03T23-14-49.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-08-03T23-14-49.pt', # 3s
'ResNetAudio-22-08-03T23-14-28.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-08-03T23-14-28.pt', # 4s
'ResNetAudio-22-06-24T08-10-33.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T08-10-33.pt', # 5s
'ResNetAudio-22-06-24T17-31-07.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T17-31-07.pt', # 6s
'ResNetAudio-22-06-24T23-57-11.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-06-24T23-57-11.pt', # 7s
'ResNetAudio-22-06-25T04-35-42.pt':
f'{PARENT_LINK}/sync/ResNetAudio-22-06-25T04-35-42.pt', # 8s
}
def check_if_file_exists_else_download(path, fname2link=FNAME2LINK, chunk_size=1024):
'''Checks if file exists, if not downloads it from the link to the path'''
path = Path(path)
if not path.exists():
path.parent.mkdir(exist_ok=True, parents=True)
link = fname2link.get(path.name, None)
if link is None:
raise ValueError(f'Cant find the checkpoint file: {path}.',
f'Please download it manually and ensure the path exists.')
with requests.get(fname2link[path.name], stream=True) as r:
total_size = int(r.headers.get('content-length', 0))
with tqdm(total=total_size, unit='B', unit_scale=True) as pbar:
with open(path, 'wb') as f:
for data in r.iter_content(chunk_size=chunk_size):
if data:
f.write(data)
pbar.update(chunk_size)
def get_md5sum(path):
hash_md5 = md5()
with open(path, 'rb') as f:
for chunk in iter(lambda: f.read(4096 * 8), b''):
hash_md5.update(chunk)
md5sum = hash_md5.hexdigest()
return md5sum
+277
View File
@@ -0,0 +1,277 @@
#!/usr/bin/env python3
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
# Copyright 2020 Ross Wightman
# Modified Model definition
from collections import OrderedDict
from functools import partial
import torch
import torch.nn as nn
from timm.layers import trunc_normal_
from fastvideo.third_party.synchformer import vit_helper
class VisionTransformer(nn.Module):
""" Vision Transformer with support for patch or hybrid CNN input stage """
def __init__(self, cfg):
super().__init__()
self.img_size = cfg.DATA.TRAIN_CROP_SIZE
self.patch_size = cfg.VIT.PATCH_SIZE
self.in_chans = cfg.VIT.CHANNELS
if cfg.TRAIN.DATASET == "Epickitchens":
self.num_classes = [97, 300]
else:
self.num_classes = cfg.MODEL.NUM_CLASSES
self.embed_dim = cfg.VIT.EMBED_DIM
self.depth = cfg.VIT.DEPTH
self.num_heads = cfg.VIT.NUM_HEADS
self.mlp_ratio = cfg.VIT.MLP_RATIO
self.qkv_bias = cfg.VIT.QKV_BIAS
self.drop_rate = cfg.VIT.DROP
self.drop_path_rate = cfg.VIT.DROP_PATH
self.head_dropout = cfg.VIT.HEAD_DROPOUT
self.video_input = cfg.VIT.VIDEO_INPUT
self.temporal_resolution = cfg.VIT.TEMPORAL_RESOLUTION
self.use_mlp = cfg.VIT.USE_MLP
self.num_features = self.embed_dim
norm_layer = partial(nn.LayerNorm, eps=1e-6)
self.attn_drop_rate = cfg.VIT.ATTN_DROPOUT
self.head_act = cfg.VIT.HEAD_ACT
self.cfg = cfg
# Patch Embedding
self.patch_embed = vit_helper.PatchEmbed(img_size=224,
patch_size=self.patch_size,
in_chans=self.in_chans,
embed_dim=self.embed_dim)
# 3D Patch Embedding
self.patch_embed_3d = vit_helper.PatchEmbed3D(img_size=self.img_size,
temporal_resolution=self.temporal_resolution,
patch_size=self.patch_size,
in_chans=self.in_chans,
embed_dim=self.embed_dim,
z_block_size=self.cfg.VIT.PATCH_SIZE_TEMP)
self.patch_embed_3d.proj.weight.data = torch.zeros_like(
self.patch_embed_3d.proj.weight.data)
# Number of patches
if self.video_input:
num_patches = self.patch_embed.num_patches * self.temporal_resolution
else:
num_patches = self.patch_embed.num_patches
self.num_patches = num_patches
# CLS token
self.cls_token = nn.Parameter(torch.zeros(1, 1, self.embed_dim))
trunc_normal_(self.cls_token, std=.02)
# Positional embedding
self.pos_embed = nn.Parameter(
torch.zeros(1, self.patch_embed.num_patches + 1, self.embed_dim))
self.pos_drop = nn.Dropout(p=cfg.VIT.POS_DROPOUT)
trunc_normal_(self.pos_embed, std=.02)
if self.cfg.VIT.POS_EMBED == "joint":
self.st_embed = nn.Parameter(torch.zeros(1, num_patches + 1, self.embed_dim))
trunc_normal_(self.st_embed, std=.02)
elif self.cfg.VIT.POS_EMBED == "separate":
self.temp_embed = nn.Parameter(torch.zeros(1, self.temporal_resolution, self.embed_dim))
# Layer Blocks
dpr = [x.item() for x in torch.linspace(0, self.drop_path_rate, self.depth)]
if self.cfg.VIT.ATTN_LAYER == "divided":
self.blocks = nn.ModuleList([
vit_helper.DividedSpaceTimeBlock(
attn_type=cfg.VIT.ATTN_LAYER,
dim=self.embed_dim,
num_heads=self.num_heads,
mlp_ratio=self.mlp_ratio,
qkv_bias=self.qkv_bias,
drop=self.drop_rate,
attn_drop=self.attn_drop_rate,
drop_path=dpr[i],
norm_layer=norm_layer,
) for i in range(self.depth)
])
else:
self.blocks = nn.ModuleList([
vit_helper.Block(attn_type=cfg.VIT.ATTN_LAYER,
dim=self.embed_dim,
num_heads=self.num_heads,
mlp_ratio=self.mlp_ratio,
qkv_bias=self.qkv_bias,
drop=self.drop_rate,
attn_drop=self.attn_drop_rate,
drop_path=dpr[i],
norm_layer=norm_layer,
use_original_code=self.cfg.VIT.USE_ORIGINAL_TRAJ_ATTN_CODE)
for i in range(self.depth)
])
self.norm = norm_layer(self.embed_dim)
# MLP head
if self.use_mlp:
hidden_dim = self.embed_dim
if self.head_act == 'tanh':
# logging.info("Using TanH activation in MLP")
act = nn.Tanh()
elif self.head_act == 'gelu':
# logging.info("Using GELU activation in MLP")
act = nn.GELU()
else:
# logging.info("Using ReLU activation in MLP")
act = nn.ReLU()
self.pre_logits = nn.Sequential(
OrderedDict([
('fc', nn.Linear(self.embed_dim, hidden_dim)),
('act', act),
]))
else:
self.pre_logits = nn.Identity()
# Classifier Head
self.head_drop = nn.Dropout(p=self.head_dropout)
if isinstance(self.num_classes, (list, )) and len(self.num_classes) > 1:
for a, i in enumerate(range(len(self.num_classes))):
setattr(self, "head%d" % a, nn.Linear(self.embed_dim, self.num_classes[i]))
else:
self.head = nn.Linear(self.embed_dim,
self.num_classes) if self.num_classes > 0 else nn.Identity()
# Initialize weights
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
@torch.jit.ignore
def no_weight_decay(self):
if self.cfg.VIT.POS_EMBED == "joint":
return {'pos_embed', 'cls_token', 'st_embed'}
else:
return {'pos_embed', 'cls_token', 'temp_embed'}
def get_classifier(self):
return self.head
def reset_classifier(self, num_classes, global_pool=''):
self.num_classes = num_classes
self.head = (nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity())
def forward_features(self, x):
# if self.video_input:
# x = x[0]
B = x.shape[0]
# Tokenize input
# if self.cfg.VIT.PATCH_SIZE_TEMP > 1:
# for simplicity of mapping between content dimensions (input x) and token dims (after patching)
# we use the same trick as for AST (see modeling_ast.ASTModel.forward for the details):
# apply patching on input
x = self.patch_embed_3d(x)
tok_mask = None
# else:
# tok_mask = None
# # 2D tokenization
# if self.video_input:
# x = x.permute(0, 2, 1, 3, 4)
# (B, T, C, H, W) = x.shape
# x = x.reshape(B * T, C, H, W)
# x = self.patch_embed(x)
# if self.video_input:
# (B2, T2, D2) = x.shape
# x = x.reshape(B, T * T2, D2)
# Append CLS token
cls_tokens = self.cls_token.expand(B, -1, -1)
x = torch.cat((cls_tokens, x), dim=1)
# if tok_mask is not None:
# # prepend 1(=keep) to the mask to account for the CLS token as well
# tok_mask = torch.cat((torch.ones_like(tok_mask[:, [0]]), tok_mask), dim=1)
# Interpolate positinoal embeddings
# if self.cfg.DATA.TRAIN_CROP_SIZE != 224:
# pos_embed = self.pos_embed
# N = pos_embed.shape[1] - 1
# npatch = int((x.size(1) - 1) / self.temporal_resolution)
# class_emb = pos_embed[:, 0]
# pos_embed = pos_embed[:, 1:]
# dim = x.shape[-1]
# pos_embed = torch.nn.functional.interpolate(
# pos_embed.reshape(1, int(math.sqrt(N)), int(math.sqrt(N)), dim).permute(0, 3, 1, 2),
# scale_factor=math.sqrt(npatch / N),
# mode='bicubic',
# )
# pos_embed = pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)
# new_pos_embed = torch.cat((class_emb.unsqueeze(0), pos_embed), dim=1)
# else:
new_pos_embed = self.pos_embed
npatch = self.patch_embed.num_patches
# Add positional embeddings to input
if self.video_input:
if self.cfg.VIT.POS_EMBED == "separate":
cls_embed = self.pos_embed[:, 0, :].unsqueeze(1)
tile_pos_embed = new_pos_embed[:, 1:, :].repeat(1, self.temporal_resolution, 1)
tile_temporal_embed = self.temp_embed.repeat_interleave(npatch, 1)
total_pos_embed = tile_pos_embed + tile_temporal_embed
total_pos_embed = torch.cat([cls_embed, total_pos_embed], dim=1)
x = x + total_pos_embed
elif self.cfg.VIT.POS_EMBED == "joint":
x = x + self.st_embed
else:
# image input
x = x + new_pos_embed
# Apply positional dropout
x = self.pos_drop(x)
# Encoding using transformer layers
for i, blk in enumerate(self.blocks):
x = blk(x,
seq_len=npatch,
num_frames=self.temporal_resolution,
approx=self.cfg.VIT.APPROX_ATTN_TYPE,
num_landmarks=self.cfg.VIT.APPROX_ATTN_DIM,
tok_mask=tok_mask)
### v-iashin: I moved it to the forward pass
# x = self.norm(x)[:, 0]
# x = self.pre_logits(x)
###
return x, tok_mask
# def forward(self, x):
# x = self.forward_features(x)
# ### v-iashin: here. This should leave the same forward output as before
# x = self.norm(x)[:, 0]
# x = self.pre_logits(x)
# ###
# x = self.head_drop(x)
# if isinstance(self.num_classes, (list, )) and len(self.num_classes) > 1:
# output = []
# for head in range(len(self.num_classes)):
# x_out = getattr(self, "head%d" % head)(x)
# if not self.training:
# x_out = torch.nn.functional.softmax(x_out, dim=-1)
# output.append(x_out)
# return output
# else:
# x = self.head(x)
# if not self.training:
# x = torch.nn.functional.softmax(x, dim=-1)
# return x
+399
View File
@@ -0,0 +1,399 @@
#!/usr/bin/env python3
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
# Copyright 2020 Ross Wightman
# Modified Model definition
"""Video model building blocks shared by Synchformer consumers."""
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from timm.layers import to_2tuple
from torch import einsum
from torch.nn import functional as F
default_cfgs = {
'vit_1k':
'https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_base_p16_224-80ecf9dd.pth',
'vit_1k_large':
'https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-vitjx/jx_vit_large_p16_224-4ee7a4dc.pth',
}
def qkv_attn(q, k, v, tok_mask: torch.Tensor = None):
sim = einsum('b i d, b j d -> b i j', q, k)
# apply masking if provided, tok_mask is (B*S*H, N): 1s - keep; sim is (B*S*H, H, N, N)
if tok_mask is not None:
BSH, N = tok_mask.shape
sim = sim.masked_fill(tok_mask.view(BSH, 1, N) == 0,
float('-inf')) # 1 - broadcasts across N
attn = sim.softmax(dim=-1)
out = einsum('b i j, b j d -> b i d', attn, v)
return out
class DividedAttention(nn.Module):
def __init__(self, dim, num_heads=8, qkv_bias=False, attn_drop=0., proj_drop=0.):
super().__init__()
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim**-0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.proj = nn.Linear(dim, dim)
# init to zeros
self.qkv.weight.data.fill_(0)
self.qkv.bias.data.fill_(0)
self.proj.weight.data.fill_(1)
self.proj.bias.data.fill_(0)
self.attn_drop = nn.Dropout(attn_drop)
self.proj_drop = nn.Dropout(proj_drop)
def forward(self, x, einops_from, einops_to, tok_mask: torch.Tensor = None, **einops_dims):
# num of heads variable
h = self.num_heads
# project x to q, k, v vaalues
q, k, v = self.qkv(x).chunk(3, dim=-1)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
if tok_mask is not None:
# replicate token mask across heads (b, n) -> (b, h, n) -> (b*h, n) -- same as qkv but w/o d
assert len(tok_mask.shape) == 2
tok_mask = tok_mask.unsqueeze(1).expand(-1, h, -1).reshape(-1, tok_mask.shape[1])
# Scale q
q *= self.scale
# Take out cls_q, cls_k, cls_v
(cls_q, q_), (cls_k, k_), (cls_v, v_) = map(lambda t: (t[:, 0:1], t[:, 1:]), (q, k, v))
# the same for masking
if tok_mask is not None:
cls_mask, mask_ = tok_mask[:, 0:1], tok_mask[:, 1:]
else:
cls_mask, mask_ = None, None
# let CLS token attend to key / values of all patches across time and space
cls_out = qkv_attn(cls_q, k, v, tok_mask=tok_mask)
# rearrange across time or space
q_, k_, v_ = map(lambda t: rearrange(t, f'{einops_from} -> {einops_to}', **einops_dims),
(q_, k_, v_))
# expand CLS token keys and values across time or space and concat
r = q_.shape[0] // cls_k.shape[0]
cls_k, cls_v = map(lambda t: repeat(t, 'b () d -> (b r) () d', r=r), (cls_k, cls_v))
k_ = torch.cat((cls_k, k_), dim=1)
v_ = torch.cat((cls_v, v_), dim=1)
# the same for masking (if provided)
if tok_mask is not None:
# since mask does not have the latent dim (d), we need to remove it from einops dims
mask_ = rearrange(mask_, f'{einops_from} -> {einops_to}'.replace(' d', ''),
**einops_dims)
cls_mask = repeat(cls_mask, 'b () -> (b r) ()',
r=r) # expand cls_mask across time or space
mask_ = torch.cat((cls_mask, mask_), dim=1)
# attention
out = qkv_attn(q_, k_, v_, tok_mask=mask_)
# merge back time or space
out = rearrange(out, f'{einops_to} -> {einops_from}', **einops_dims)
# concat back the cls token
out = torch.cat((cls_out, out), dim=1)
# merge back the heads
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
## to out
x = self.proj(out)
x = self.proj_drop(x)
return x
class DividedSpaceTimeBlock(nn.Module):
def __init__(self,
dim=768,
num_heads=12,
attn_type='divided',
mlp_ratio=4.,
qkv_bias=False,
drop=0.,
attn_drop=0.,
drop_path=0.,
act_layer=nn.GELU,
norm_layer=nn.LayerNorm):
super().__init__()
self.einops_from_space = 'b (f n) d'
self.einops_to_space = '(b f) n d'
self.einops_from_time = 'b (f n) d'
self.einops_to_time = '(b n) f d'
self.norm1 = norm_layer(dim)
self.attn = DividedAttention(dim,
num_heads=num_heads,
qkv_bias=qkv_bias,
attn_drop=attn_drop,
proj_drop=drop)
self.timeattn = DividedAttention(dim,
num_heads=num_heads,
qkv_bias=qkv_bias,
attn_drop=attn_drop,
proj_drop=drop)
# self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
self.drop_path = nn.Identity()
self.norm2 = norm_layer(dim)
mlp_hidden_dim = int(dim * mlp_ratio)
self.mlp = Mlp(in_features=dim,
hidden_features=mlp_hidden_dim,
act_layer=act_layer,
drop=drop)
self.norm3 = norm_layer(dim)
def forward(self,
x,
seq_len=196,
num_frames=8,
approx='none',
num_landmarks=128,
tok_mask: torch.Tensor = None):
time_output = self.timeattn(self.norm3(x),
self.einops_from_time,
self.einops_to_time,
n=seq_len,
tok_mask=tok_mask)
time_residual = x + time_output
space_output = self.attn(self.norm1(time_residual),
self.einops_from_space,
self.einops_to_space,
f=num_frames,
tok_mask=tok_mask)
space_residual = time_residual + self.drop_path(space_output)
x = space_residual
x = x + self.drop_path(self.mlp(self.norm2(x)))
return x
class Mlp(nn.Module):
def __init__(self,
in_features,
hidden_features=None,
out_features=None,
act_layer=nn.GELU,
drop=0.):
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
class PatchEmbed(nn.Module):
""" Image to Patch Embedding
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
img_size = img_size if type(img_size) is tuple else to_2tuple(img_size)
patch_size = img_size if type(patch_size) is tuple else to_2tuple(patch_size)
num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = num_patches
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x):
B, C, H, W = x.shape
x = self.proj(x).flatten(2).transpose(1, 2)
return x
class PatchEmbed3D(nn.Module):
""" Image to Patch Embedding """
def __init__(self,
img_size=224,
temporal_resolution=4,
in_chans=3,
patch_size=16,
z_block_size=2,
embed_dim=768,
flatten=True):
super().__init__()
self.height = (img_size // patch_size)
self.width = (img_size // patch_size)
### v-iashin: these two are incorrect
# self.frames = (temporal_resolution // z_block_size)
# self.num_patches = self.height * self.width * self.frames
self.z_block_size = z_block_size
###
self.proj = nn.Conv3d(in_chans,
embed_dim,
kernel_size=(z_block_size, patch_size, patch_size),
stride=(z_block_size, patch_size, patch_size))
self.flatten = flatten
def forward(self, x):
B, C, T, H, W = x.shape
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2)
return x
class HeadMLP(nn.Module):
def __init__(self, n_input, n_classes, n_hidden=512, p=0.1):
super(HeadMLP, self).__init__()
self.n_input = n_input
self.n_classes = n_classes
self.n_hidden = n_hidden
if n_hidden is None:
# use linear classifier
self.block_forward = nn.Sequential(nn.Dropout(p=p),
nn.Linear(n_input, n_classes, bias=True))
else:
# use simple MLP classifier
self.block_forward = nn.Sequential(nn.Dropout(p=p),
nn.Linear(n_input, n_hidden, bias=True),
nn.BatchNorm1d(n_hidden), nn.ReLU(inplace=True),
nn.Dropout(p=p),
nn.Linear(n_hidden, n_classes, bias=True))
print(f"Dropout-NLP: {p}")
def forward(self, x):
return self.block_forward(x)
def _conv_filter(state_dict, patch_size=16):
""" convert patch embedding weight from manual patchify + linear proj to conv"""
out_dict = {}
for k, v in state_dict.items():
if 'patch_embed.proj.weight' in k:
v = v.reshape((v.shape[0], 3, patch_size, patch_size))
out_dict[k] = v
return out_dict
def adapt_input_conv(in_chans, conv_weight, agg='sum'):
conv_type = conv_weight.dtype
conv_weight = conv_weight.float()
O, I, J, K = conv_weight.shape
if in_chans == 1:
if I > 3:
assert conv_weight.shape[1] % 3 == 0
# For models with space2depth stems
conv_weight = conv_weight.reshape(O, I // 3, 3, J, K)
conv_weight = conv_weight.sum(dim=2, keepdim=False)
else:
if agg == 'sum':
print("Summing conv1 weights")
conv_weight = conv_weight.sum(dim=1, keepdim=True)
else:
print("Averaging conv1 weights")
conv_weight = conv_weight.mean(dim=1, keepdim=True)
elif in_chans != 3:
if I != 3:
raise NotImplementedError('Weight format not supported by conversion.')
else:
if agg == 'sum':
print("Summing conv1 weights")
repeat = int(math.ceil(in_chans / 3))
conv_weight = conv_weight.repeat(1, repeat, 1, 1)[:, :in_chans, :, :]
conv_weight *= (3 / float(in_chans))
else:
print("Averaging conv1 weights")
conv_weight = conv_weight.mean(dim=1, keepdim=True)
conv_weight = conv_weight.repeat(1, in_chans, 1, 1)
conv_weight = conv_weight.to(conv_type)
return conv_weight
def load_pretrained(model,
cfg=None,
num_classes=1000,
in_chans=3,
filter_fn=None,
strict=True,
progress=False):
# Load state dict
assert (f"{cfg.VIT.PRETRAINED_WEIGHTS} not in [vit_1k, vit_1k_large]")
state_dict = torch.hub.load_state_dict_from_url(url=default_cfgs[cfg.VIT.PRETRAINED_WEIGHTS])
if filter_fn is not None:
state_dict = filter_fn(state_dict)
input_convs = 'patch_embed.proj'
if input_convs is not None and in_chans != 3:
if isinstance(input_convs, str):
input_convs = (input_convs, )
for input_conv_name in input_convs:
weight_name = input_conv_name + '.weight'
try:
state_dict[weight_name] = adapt_input_conv(in_chans,
state_dict[weight_name],
agg='avg')
print(
f'Converted input conv {input_conv_name} pretrained weights from 3 to {in_chans} channel(s)'
)
except NotImplementedError as e:
del state_dict[weight_name]
strict = False
print(
f'Unable to convert pretrained {input_conv_name} weights, using random init for this layer.'
)
classifier_name = 'head'
label_offset = cfg.get('label_offset', 0)
pretrain_classes = 1000
if num_classes != pretrain_classes:
# completely discard fully connected if model num_classes doesn't match pretrained weights
del state_dict[classifier_name + '.weight']
del state_dict[classifier_name + '.bias']
strict = False
elif label_offset > 0:
# special case for pretrained weights with an extra background class in pretrained weights
classifier_weight = state_dict[classifier_name + '.weight']
state_dict[classifier_name + '.weight'] = classifier_weight[label_offset:]
classifier_bias = state_dict[classifier_name + '.bias']
state_dict[classifier_name + '.bias'] = classifier_bias[label_offset:]
loaded_state = state_dict
self_state = model.state_dict()
all_names = set(self_state.keys())
saved_names = set([])
for name, param in loaded_state.items():
param = param
if 'module.' in name:
name = name.replace('module.', '')
if name in self_state.keys() and param.shape == self_state[name].shape:
saved_names.add(name)
self_state[name].copy_(param)
else:
print(f"didnt load: {name} of shape: {param.shape}")
print("Missing Keys:")
print(all_names - saved_names)
@@ -0,0 +1,344 @@
# SPDX-License-Identifier: Apache-2.0
"""Convert MMAudio ``large_44k_v2`` assets into a FastVideo component tree.
The converter is deliberately offline: every large source asset must already
exist locally. It splits the shared DFN5B OpenCLIP checkpoint into native
FastVideo text and vision encoders, preserves the exact MMAudio transformer,
Synchformer, VAE, and BigVGAN weights, and emits the standard
``model_index.json`` layout consumed by ``ComposedPipelineBase``.
Example::
python scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py \
--transformer-checkpoint ../MMAudio/weights/mmaudio_large_44k_v2.pth \
--audio-vae-checkpoint ../MMAudio/ext_weights/v1-44.pth \
--synchformer-checkpoint ../MMAudio/ext_weights/synchformer_state_dict.pth \
--dfn5b-dir official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384 \
--bigvgan-dir official_weights/mmaudio/bigvgan_v2_44khz_128band_512x \
--output converted_weights/mmaudio/large_44k_v2
"""
from __future__ import annotations
import argparse
import gzip
import json
from pathlib import Path
from typing import Any
import torch
from safetensors.torch import save_file
TRANSFORMER_CONFIG = {
"_class_name": "MMAudioTransformer",
"latent_dim": 40,
"clip_dim": 1024,
"sync_dim": 768,
"text_dim": 1024,
"hidden_dim": 896,
"depth": 21,
"fused_depth": 14,
"num_heads": 14,
"mlp_ratio": 4.0,
"latent_seq_len": 345,
"clip_seq_len": 64,
"sync_seq_len": 192,
"text_seq_len": 77,
"v2": True,
}
TEXT_ENCODER_CONFIG = {
"architectures": ["MMAudioDFNCLIPTextEncoder"],
"vocab_size": 49408,
"hidden_size": 1024,
"intermediate_size": 4096,
"projection_dim": 1024,
"num_hidden_layers": 24,
"num_attention_heads": 16,
"max_position_embeddings": 77,
"text_len": 77,
"hidden_act": "quick_gelu",
"layer_norm_eps": 1e-5,
"pad_token_id": 0,
"bos_token_id": 49406,
"eos_token_id": 49407,
}
IMAGE_ENCODER_CONFIG = {
"architectures": ["MMAudioDFNCLIPVisionEncoder"],
"hidden_size": 1280,
"intermediate_size": 5120,
"projection_dim": 1024,
"num_hidden_layers": 32,
"num_attention_heads": 16,
"num_channels": 3,
"image_size": 378,
"patch_size": 14,
"hidden_act": "quick_gelu",
"layer_norm_eps": 1e-5,
}
SYNCHFORMER_CONFIG = {
"architectures": ["MMAudioSynchformerVisualEncoder"],
"image_size": 224,
"num_channels": 3,
"segment_size": 16,
"segment_stride": 8,
"hidden_size": 768,
"tokens_per_segment": 8,
}
MODEL_INDEX = {
"_class_name": "MMAudioPipeline",
"_diffusers_version": "0.36.0",
"_fastvideo_model_family": "mmaudio",
"_fastvideo_workload_types": ["V2A", "T2A"],
"transformer": [
"fastvideo.models.dits.mmaudio",
"MMAudioTransformer",
],
"text_encoder": [
"fastvideo.models.encoders.mmaudio_clip",
"MMAudioDFNCLIPTextEncoder",
],
"tokenizer": ["transformers", "CLIPTokenizer"],
"image_encoder": [
"fastvideo.models.encoders.mmaudio_clip",
"MMAudioDFNCLIPVisionEncoder",
],
"image_encoder_2": [
"fastvideo.models.encoders.mmaudio_synchformer",
"MMAudioSynchformerVisualEncoder",
],
"audio_vae": ["fastvideo.models.audio.mmaudio_vae", "MMAudioVAE"],
"vocoder": ["fastvideo.models.audio.bigvgan", "BigVGANV2"],
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
}
def _write_json(path: Path, value: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as handle:
json.dump(value, handle, indent=2, sort_keys=True)
handle.write("\n")
def _write_component(output: Path, name: str, state: dict[str, torch.Tensor], config: dict[str, Any]) -> None:
directory = output / name
directory.mkdir(parents=True, exist_ok=True)
contiguous = {key: tensor.detach().cpu().contiguous() for key, tensor in state.items()}
save_file(contiguous, directory / "diffusion_pytorch_model.safetensors", metadata={"format": "pt"})
_write_json(directory / "config.json", config)
def _load_torch_state(path: Path) -> dict[str, torch.Tensor]:
if not path.is_file():
raise FileNotFoundError(path)
value = torch.load(path, map_location="cpu", weights_only=True)
if not isinstance(value, dict):
raise TypeError(f"Expected a state dict in {path}, got {type(value)}")
for key in ("state_dict", "model", "generator"):
nested = value.get(key)
if isinstance(nested, dict) and nested:
value = nested
break
if not all(isinstance(tensor, torch.Tensor) for tensor in value.values()):
raise TypeError(f"Checkpoint {path} contains non-tensor state entries")
return value
def map_open_clip_text_state(
state: dict[str, torch.Tensor],
) -> dict[str, torch.Tensor]:
mapped: dict[str, torch.Tensor] = {
"text_model.embeddings.token_embedding.weight": state["token_embedding.weight"],
"text_model.embeddings.position_embedding.weight": state["positional_embedding"],
"text_model.final_layer_norm.weight": state["ln_final.weight"],
"text_model.final_layer_norm.bias": state["ln_final.bias"],
}
for name, tensor in state.items():
if not name.startswith("transformer.resblocks."):
continue
target = name.replace("transformer.resblocks.", "text_model.encoder.layers.")
target = target.replace(".ln_1.", ".layer_norm1.")
target = target.replace(".ln_2.", ".layer_norm2.")
target = target.replace(".attn.in_proj_", ".self_attn.qkv_proj.")
target = target.replace(".attn.out_proj.", ".self_attn.out_proj.")
target = target.replace(".mlp.c_fc.", ".mlp.fc1.")
target = target.replace(".mlp.c_proj.", ".mlp.fc2.")
mapped[target] = tensor
return mapped
def map_open_clip_vision_state(
state: dict[str, torch.Tensor],
) -> dict[str, torch.Tensor]:
mapped: dict[str, torch.Tensor] = {
"vision_model.embeddings.class_embedding": state["visual.class_embedding"],
"vision_model.embeddings.patch_embedding.weight": state["visual.conv1.weight"],
"vision_model.embeddings.position_embedding.weight": state["visual.positional_embedding"],
"vision_model.pre_layrnorm.weight": state["visual.ln_pre.weight"],
"vision_model.pre_layrnorm.bias": state["visual.ln_pre.bias"],
"vision_model.post_layernorm.weight": state["visual.ln_post.weight"],
"vision_model.post_layernorm.bias": state["visual.ln_post.bias"],
"visual_projection.weight": state["visual.proj"].t(),
}
for name, tensor in state.items():
if not name.startswith("visual.transformer.resblocks."):
continue
target = name.replace("visual.transformer.resblocks.", "vision_model.encoder.layers.")
target = target.replace(".ln_1.", ".layer_norm1.")
target = target.replace(".ln_2.", ".layer_norm2.")
target = target.replace(".attn.in_proj_", ".self_attn.qkv_proj.")
target = target.replace(".attn.out_proj.", ".self_attn.out_proj.")
target = target.replace(".mlp.c_fc.", ".mlp.fc1.")
target = target.replace(".mlp.c_proj.", ".mlp.fc2.")
mapped[target] = tensor
return mapped
def write_open_clip_tokenizer(output: Path) -> None:
"""Write the bundled OpenAI CLIP BPE as an AutoTokenizer component.
OpenCLIP pads its 77-token tensor with integer zero. ``CLIPTokenizer``
cannot use vocabulary ID zero as a special pad token without changing how
a literal exclamation mark is tokenized, so the MMAudio text stage zeros
positions selected by ``attention_mask`` after tokenization.
"""
from open_clip.tokenizer import bytes_to_unicode, default_bpe
with gzip.open(default_bpe()) as bpe_file:
merges_raw = bpe_file.read().decode("utf-8").split("\n")
merges = merges_raw[1 : 49152 - 256 - 2 + 1]
merge_pairs = [tuple(merge.split()) for merge in merges]
vocab = list(bytes_to_unicode().values())
vocab += [token + "</w>" for token in vocab]
vocab += ["".join(pair) for pair in merge_pairs]
vocab += ["<start_of_text>", "<end_of_text>"]
encoder = {token: index for index, token in enumerate(vocab)}
directory = output / "tokenizer"
directory.mkdir(parents=True, exist_ok=True)
_write_json(directory / "vocab.json", encoder)
with (directory / "merges.txt").open("w", encoding="utf-8") as handle:
handle.write("#version: 0.2\n")
for first, second in merge_pairs:
handle.write(f"{first} {second}\n")
_write_json(
directory / "tokenizer_config.json",
{
"tokenizer_class": "CLIPTokenizer",
"model_max_length": 77,
"bos_token": "<start_of_text>",
"eos_token": "<end_of_text>",
"unk_token": "<end_of_text>",
"pad_token": "<end_of_text>",
"do_lower_case": True,
},
)
_write_json(
directory / "special_tokens_map.json",
{
"bos_token": "<start_of_text>",
"eos_token": "<end_of_text>",
"unk_token": "<end_of_text>",
"pad_token": "<end_of_text>",
},
)
def _load_dfn5b_state(directory: Path) -> dict[str, torch.Tensor]:
if not directory.is_dir():
raise FileNotFoundError(directory)
from open_clip import create_model_from_pretrained
model = create_model_from_pretrained(f"local-dir:{directory}",
return_transform=False)
state = {key: tensor.detach().cpu() for key, tensor in model.state_dict().items()}
del model
return state
def convert(args: argparse.Namespace) -> None:
output = args.output.resolve()
output.mkdir(parents=True, exist_ok=True)
transformer_state = _load_torch_state(args.transformer_checkpoint)
# Official ``MMAudio.load_weights`` discards this derived buffer. Keeping
# it would make a standard strict FastVideo component load fail.
transformer_state.pop("t_embed.freqs", None)
transformer_state.pop("latent_rot", None)
transformer_state.pop("clip_rot", None)
_write_component(output, "transformer", transformer_state, TRANSFORMER_CONFIG)
vae_state = _load_torch_state(args.audio_vae_checkpoint)
decoder_state = {
key: tensor
for key, tensor in vae_state.items()
if key.startswith("decoder.") or key in {"data_mean", "data_std"}
}
if not decoder_state:
raise ValueError("Audio VAE checkpoint did not contain decoder weights")
_write_component(
output,
"audio_vae",
decoder_state,
{"_class_name": "MMAudioVAE", "mode": "44k", "need_encoder": False},
)
synchformer_state = _load_torch_state(args.synchformer_checkpoint)
synchformer_visual_state = {
name: tensor
for name, tensor in synchformer_state.items()
if name.startswith("vfeat_extractor.")
}
if not synchformer_visual_state:
raise ValueError(
"Synchformer checkpoint did not contain vfeat_extractor weights")
_write_component(output, "image_encoder_2", synchformer_visual_state,
SYNCHFORMER_CONFIG)
dfn_state = _load_dfn5b_state(args.dfn5b_dir)
_write_component(output, "text_encoder", map_open_clip_text_state(dfn_state), TEXT_ENCODER_CONFIG)
_write_component(output, "image_encoder", map_open_clip_vision_state(dfn_state), IMAGE_ENCODER_CONFIG)
write_open_clip_tokenizer(output)
bigvgan_config_path = args.bigvgan_dir / "config.json"
if not bigvgan_config_path.is_file():
raise FileNotFoundError(bigvgan_config_path)
with bigvgan_config_path.open(encoding="utf-8") as handle:
bigvgan_config = json.load(handle)
bigvgan_config["_class_name"] = "BigVGANV2"
bigvgan_config["weight_norm_removed"] = False
bigvgan_state = _load_torch_state(args.bigvgan_dir / "bigvgan_generator.pt")
_write_component(output, "vocoder", bigvgan_state, bigvgan_config)
_write_json(
output / "scheduler/scheduler_config.json",
{
"_class_name": "FlowMatchEulerDiscreteScheduler",
"num_train_timesteps": 1000,
"shift": 1.0,
"invert_sigmas": True,
"sigma_min": 0.0,
"use_reference_discrete_timesteps": True,
},
)
_write_json(output / "model_index.json", MODEL_INDEX)
print(f"Converted MMAudio large_44k_v2 components to {output}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
parser.add_argument("--transformer-checkpoint", type=Path, required=True)
parser.add_argument("--audio-vae-checkpoint", type=Path, required=True)
parser.add_argument("--synchformer-checkpoint", type=Path, required=True)
parser.add_argument("--dfn5b-dir", type=Path, required=True)
parser.add_argument("--bigvgan-dir", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
if __name__ == "__main__":
convert(parse_args())
+98
View File
@@ -0,0 +1,98 @@
# MMAudio Port Status
## Summary
- model_family: `mmaudio`
- workload_types: `V2A`, `T2A`
- official_ref: `../MMAudio` at `974010a026c731054592d8f777218bd9d85a6c24`
- first_variant: `large_44k_v2`
- phase: `native_pipeline_complete`
- status: `real_weight_parity_pass`
- last_updated: `2026-08-02`
## Native Components
| Component | FastVideo implementation | Reuse/port decision | Real-weight result |
|---|---|---|---|
| MMAudio transformer | `fastvideo/models/dits/mmaudio.py` | Native 1D multimodal DiT | exact |
| DFN5B text/vision | `fastvideo/models/encoders/mmaudio_clip.py` | Shared native CLIP core, MMAudio adapters | exact |
| Synchformer visual encoder | `fastvideo/models/encoders/mmaudio_synchformer.py` | Shared backbone under `fastvideo/third_party/synchformer` | exact, including 16-frame/stride-8 usage contract |
| 44.1 kHz VAE | `fastvideo/models/audio/mmaudio_vae.py` | Native audio component | exact state structure, FP32/BF16 random-weight decoder, and real-weight encode/decode parity |
| BigVGAN-v2 | `fastvideo/models/audio/bigvgan.py` | Shared native vocoder | exact |
| Euler flow schedule | shared `FlowMatchEulerDiscreteScheduler` | Reuse schedule; preserve official BF16 scalar update in MMAudio stage | exact |
## Pipeline Integration
- Pipeline: `fastvideo/pipelines/basic/mmaudio/MMAudioPipeline`
- Config: `fastvideo/configs/pipelines/mmaudio.py::MMAudioV2AConfig`
- Preset: `mmaudio_large_44k_v2`
- Registry: resolves both `WorkloadType.V2A` and `WorkloadType.T2A`
- Required production components: `transformer`, `scheduler`, `text_encoder`,
`tokenizer`, `image_encoder`, `image_encoder_2`, `audio_vae`, `vocoder`
- Output: mono `[B,1,samples]`, 44.1 kHz, exposed through FastVideo's
audio-only result contract and saved as WAV by `VideoGenerator`
- Duration: dynamic sequence lengths, with 8 seconds retained only as the
published training/default duration; longer and shorter inference is accepted
- Existing T2V/I2V/T2I pipelines are not routed through these stages.
The V2A preprocessing contract is identical to official MMAudio:
1. timestamp sampling at 8 FPS for DFN5B and 25 FPS for Synchformer;
2. DFN path: bicubic resize to `384x384`, float `[0,1]`, CLIP normalization;
3. sync path: bicubic short-side resize to 224, center crop, normalize to `[-1,1]`;
4. Synchformer: 16-frame windows, stride 8, `(segment,time)` token flattening.
## Converted Checkpoint
- Converter: `scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py`
- Local artifact: `converted_weights/mmaudio/large_44k_v2`
- Production strict load: pass for all eight components
- Total local artifact size: about 9.1 GB
- Converted weights and official assets remain ignored/untracked.
## Parity Evidence
| Scope | Result |
|---|---|
| Official/FastVideo video preprocessing | exact (`clip max_abs=0`, `sync max_abs=0`) |
| Condition features, random latent, projected conditions | exact |
| First flow prediction and 25-step final latent | exact |
| Final 2-second V2A waveform (89,088 samples) | exact (`atol=0`, `rtol=0`) |
| Real 10-second variable-duration V2A | pass (441,344 samples, 10.0078 s) |
| Default FastVideo offload path | real one-step smoke pass |
| Local suite | full rerun pending; VAE FP32/BF16 random-weight and real-weight cases pass |
| VAE component parity | state structure and FP32/BF16 random-weight decoder parity pass on GB200 (`atol=1e-6`, `rtol=1e-6`); real-weight encode/decode parity passes (`atol=1e-5`, `rtol=1e-5`) |
| Exact-head GB200 V2A examples | two 5-second, 25-step, seed-42 generations pass; each output is mono PCM16 at 44.1 kHz with 221,184 samples |
Commands:
```bash
pytest -q tests/local_tests/mmaudio
MMAUDIO_RUN_PIPELINE_PARITY=1 \
MMAUDIO_PARITY_VIDEO=/path/to/video-at-least-2s.mp4 \
pytest -q tests/local_tests/mmaudio/test_mmaudio_pipeline_parity.py::test_mmaudio_real_v2a_pipeline_waveform_parity -s
```
The opt-in real pipeline gate passed on an RTX 6000 Ada with the downloaded
official `large_44k_v2`, DFN5B, Synchformer, VAE, and BigVGAN assets.
## Important Numeric Decisions
- OpenCLIP text uses the explicit additive causal mask used by
`nn.MultiheadAttention`; SDPA's `is_causal` shortcut rounds differently in BF16.
- The Euler time/delta scalars stay on CPU, matching official MMAudio's
`torch.linspace` loop. Moving those float32 scalars to CUDA changes BF16 promotion.
- `t_embed.freqs` is materialized in BF16 after meta loading; dynamic RoPE buffers
are rebuilt in FP32 exactly as official `update_seq_lengths` does.
- MMAudio VAE and BigVGAN weight norm is removed on CPU in FP32 before casting to
BF16, matching the official feature utility construction order.
## Deferred Scope
- Publishing the converted checkpoint and immutable source revisions.
- Optional source-video mux/re-encode helper; the current V2A result is WAV/audio.
- 16 kHz and small/medium variants.
- Sequence/tensor-parallel optimization.
- Training integration. The official repository does not support training the
`_v2` variant; any training port should start from a v1 44.1 kHz checkpoint.
+149
View File
@@ -0,0 +1,149 @@
# MMAudio Local Tests
Local-only parity and smoke tests for the native MMAudio FastVideo port. These
tests compare FastVideo against the official reference implementation and are
not expected to run in CI unless explicitly promoted later.
Port progress, open questions, issues, and handoff notes live in
`tests/local_tests/mmaudio/PORT_STATUS.md`.
## Reference Assets
| Field | Value |
|---|---|
| Model family | `mmaudio` |
| Workload types | `V2A`, `T2A` (new native workload values; no T2V compatibility shim) |
| Official reference | `https://github.com/hkchengrex/MMAudio` |
| Local reference dir | `../MMAudio` relative to the FastVideo repository |
| Official commit/version | `974010a026c731054592d8f777218bd9d85a6c24` |
| HF weights | `hkchengrex/MMAudio` plus canonical DFN5B CLIP and BigVGAN component repos |
| HF revision | default; pin immutable revisions before publishing converted artifacts |
| Local weights dir | `official_weights/mmaudio` (downloaded locally and gitignored) |
| Source layout | mixed raw official checkpoints and external pretrained components |
| Needs conversion | yes |
The public reference assets have been downloaded locally. Never write token
values in this file; use `HF_TOKEN`, `HUGGINGFACE_HUB_TOKEN`, or `HF_API_KEY`
only in the shell when a future gated asset requires one.
## First Implementation Scope
- Inference parity target: `large_44k_v2`, 44.1 kHz, V2A and T2A.
- Training parity target: a v1 44.1 kHz checkpoint because the official
repository states that `_v2` training is unsupported.
- Inputs: video plus optional text for V2A, or text-only for T2A.
- Output: mono waveform and sample rate through FastVideo's audio-only result
contract. Source-video muxing is deferred.
- Duration: 8 seconds is the published training/default duration, but inference
uses dynamic sequence lengths and accepts shorter or longer clips. As in the
official demo, quality can fall when moving far away from 8 seconds.
- Deferred until the base parity gate passes: 16 kHz, small/medium variants,
sequence/tensor parallel optimization, and quality-reference publication.
## Shared Environment Setup
Run from the FastVideo repository root in the same environment used for
FastVideo. Do not create a separate upstream-only environment: both sides of a
parity test must share the same PyTorch/CUDA numeric stack.
The shared environment is `FastVideo/.venv`, created with uv-managed CPython
3.12.13. FastVideo is installed editable with its development dependencies;
MMAudio is installed editable with `--no-deps`, followed by the missing
reference dependencies. NumPy is pinned to 2.0.2, satisfying both MMAudio's
`numpy<2.1` constraint and FastVideo's installed OpenCV/SciPy constraints.
Recreate the environment with:
```bash
uv venv --python 3.12 --seed
source .venv/bin/activate
UV_TORCH_BACKEND=cu126 uv pip install -e ".[dev]"
uv pip install --no-deps -e ../MMAudio
uv pip install cython "gitpython>=3.1" "hydra-core>=1.3.2" \
"torchdiffeq>=0.2.5" "librosa>=0.8.1" nitrous-ema hydra_colorlog \
"tensordict>=0.6.1" colorlog "open_clip_torch>=2.29.0" "numpy==2.0.2"
```
## Official Environment Status
```text
dependency_changes: installed no-deps editable plus missing official dependencies in current env
official_env_status: imports_ok
private_dep_stubs: none planned
blocked_on: none
```
Run the FastVideo demo with an explicit duration:
```bash
MMAUDIO_MODEL_PATH=converted_weights/mmaudio/large_44k_v2 \
python examples/inference/basic/basic_mmaudio.py \
--video-path /path/to/video.mp4 \
--duration-seconds 10 \
--output-path outputs_audio/mmaudio_10s.wav
```
## Weight Setup
The approved reference set is present: `mmaudio_large_44k_v2.pth`,
`v1-44.pth`, Synchformer, DFN5B CLIP, and canonical 44.1 kHz BigVGAN-v2.
Converted artifacts will remain untracked under:
```text
official_weights/mmaudio/
converted_weights/mmaudio/
```
## Prototype And Conversion Artifacts
```text
official_key_dumps:
transformer: converted_weights/mmaudio/_mapping/transformer_official_keys.json
audio_vae: converted_weights/mmaudio/_mapping/audio_vae_official_keys.json
synchformer: converted_weights/mmaudio/_mapping/synchformer_official_keys.json
fastvideo_key_dumps:
transformer: converted_weights/mmaudio/_mapping/transformer_fastvideo_keys.json
audio_vae: converted_weights/mmaudio/_mapping/audio_vae_fastvideo_keys.json
synchformer: converted_weights/mmaudio/_mapping/synchformer_fastvideo_keys.json
conversion_script: scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py
conversion_source_layout: mixed
converted_weights_dir: converted_weights/mmaudio
strict_load_status: pass for all production components
```
## Expected Parity Tests
| Component | Official files / args | Test | Concerns | Status |
|---|---|---|---|---|
| Transformer | `mmaudio/model/networks.py`; `large_44k_v2()` | `tests/local_tests/mmaudio/test_mmaudio_transformer_parity.py` | RoPE scaling, `nearest-exact`, official final-layer `global_c` behavior | exact real-weight parity |
| DFN5B conditioner | `mmaudio/model/utils/features_utils.py`; `apple/DFN5B-CLIP-ViT-H-14-384` | `tests/local_tests/mmaudio/test_mmaudio_clip_parity.py` | patched tokenwise text output and normalized vision projection | exact real-weight parity |
| Synchformer | `mmaudio/ext/synchformer/`; 16-frame windows with stride 8 | `tests/local_tests/mmaudio/test_mmaudio_synchformer_parity.py` | shared backbone must stay outside eval-only namespace | exact real-weight model and usage parity |
| Audio VAE | `mmaudio/ext/autoencoder/`; `v1-44.pth` | `tests/local_tests/mmaudio/test_mmaudio_audio_vae_parity.py` | latent transpose and normalization statistics | exact real-weight parity |
| BigVGAN | canonical `nvidia/bigvgan_v2_44khz_128band_512x` instantiation used by `AutoEncoderModule` | `tests/local_tests/mmaudio/test_mmaudio_vocoder_parity.py` | exact resampling and weight-norm behavior | exact real-weight parity |
| Scheduler | `mmaudio/model/flow_matching.py::FlowMatching` | `tests/local_tests/mmaudio/test_mmaudio_scheduler_parity.py` | forward-time Euler convention | existing FastVideo scheduler parity passed |
| Pipeline | `mmaudio/eval_utils.py::generate` and `demo.py` | `tests/local_tests/mmaudio/test_mmaudio_pipeline_parity.py` | dual-FPS frame sampling, CFG order, Euler integration, waveform output | exact real-weight parity passed |
Planned commands:
```bash
pytest tests/local_tests/mmaudio/test_mmaudio_transformer_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_clip_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_synchformer_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_audio_vae_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_vocoder_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_scheduler_parity.py -v -s
pytest tests/local_tests/mmaudio/test_mmaudio_pipeline_parity.py -v -s
```
## Review Notes
- MMAudio must be a FastVideo-native component/pipeline port; production code
must not import `mmaudio.*`.
- Converted weights and large reference assets must remain ignored.
- Every stateful component, including reused DFN5B/BigVGAN implementations,
requires non-skip numerical parity before final handoff.
- Pipeline smoke and waveform parity are required before the new model is
presented as runnable.
- MMAudio checkpoints are documented upstream as CC-BY-NC 4.0; converted model
publishing must preserve the applicable license and attribution.
@@ -0,0 +1,105 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio 44.1 kHz audio VAE parity against the official reference.
Coverage scope: implementation_subcomponent. Production-loader coverage is
added after the converted component directory exists.
"""
from __future__ import annotations
import gc
import os
from pathlib import Path
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_WEIGHTS = Path(
os.environ.get(
"MMAUDIO_AUDIO_VAE_WEIGHTS",
REPO_ROOT.parent / "MMAudio/ext_weights/v1-44.pth",
)
)
def _official_model():
official_vae = pytest.importorskip("mmaudio.ext.autoencoder.vae")
return official_vae.VAE_44k()
def _fastvideo_model(*, need_encoder: bool):
from fastvideo.models.audio.mmaudio_vae import MMAudioVAE
return MMAudioVAE(mode="44k", need_encoder=need_encoder)
def test_mmaudio_44k_audio_vae_state_structure() -> None:
official = _official_model()
expected = {name: tensor.shape for name, tensor in official.state_dict().items()}
del official
gc.collect()
fastvideo = _fastvideo_model(need_encoder=True)
assert expected == {
name: tensor.shape for name, tensor in fastvideo.state_dict().items()
}
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16], ids=["fp32", "bf16"])
def test_mmaudio_44k_audio_vae_decoder_implementation_parity(dtype: torch.dtype) -> None:
if not torch.cuda.is_available():
pytest.skip("MMAudio audio VAE implementation parity requires CUDA")
official = _official_model()
del official.encoder
gc.collect()
fastvideo = _fastvideo_model(need_encoder=False)
fastvideo.load_state_dict(official.state_dict(), strict=True)
device = torch.device("cuda:0")
official.remove_weight_norm().to(device=device, dtype=dtype).eval()
fastvideo.remove_weight_norm().to(device=device, dtype=dtype).eval()
latent = torch.randn(
(1, 40, 4),
generator=torch.Generator(device=device).manual_seed(1234),
device=device,
dtype=dtype,
)
with torch.inference_mode():
expected = official.decode(latent)
actual = fastvideo.decode(latent)
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
def test_mmaudio_44k_audio_vae_numerical_parity() -> None:
if not torch.cuda.is_available():
pytest.skip("MMAudio audio VAE parity requires CUDA")
if not OFFICIAL_WEIGHTS.is_file():
pytest.skip(
"Official MMAudio 44.1 kHz VAE weights are absent. Set MMAUDIO_AUDIO_VAE_WEIGHTS or download v1-44.pth."
)
device = torch.device("cuda:0")
official = _official_model()
fastvideo = _fastvideo_model(need_encoder=True)
state = torch.load(OFFICIAL_WEIGHTS, map_location="cpu", weights_only=True)
official.load_state_dict(state, strict=True)
fastvideo.load_state_dict(state, strict=True)
official.remove_weight_norm().to(device).eval()
fastvideo.remove_weight_norm().to(device).eval()
generator = torch.Generator(device=device).manual_seed(1234)
mel = torch.randn((1, 128, 128), generator=generator, device=device)
latent = torch.randn((1, 40, 64), generator=generator, device=device)
with torch.inference_mode():
expected_posterior = official.encode(mel)
actual_posterior = fastvideo.encode(mel)
expected_mel = official.decode(latent)
actual_mel = fastvideo.decode(latent)
torch.testing.assert_close(actual_posterior.mean, expected_posterior.mean, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(actual_posterior.logvar, expected_posterior.logvar, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(actual_mel, expected_mel, atol=1e-5, rtol=1e-5)
@@ -0,0 +1,225 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio DFN5B OpenCLIP conditioner parity tests."""
from __future__ import annotations
import os
import importlib.util
from functools import cache
from pathlib import Path
import pytest
import torch
import torch.nn.functional as F
REPO_ROOT = Path(__file__).resolve().parents[3]
DFN5B_DIR = Path(
os.environ.get(
"MMAUDIO_DFN5B_DIR",
REPO_ROOT / "official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384",
)
)
@cache
def _conversion_module():
path = REPO_ROOT / "scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py"
spec = importlib.util.spec_from_file_location("mmaudio_converter", path)
assert spec is not None and spec.loader is not None
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
def test_converted_open_clip_tokenizer_parity(tmp_path: Path) -> None:
import open_clip
from transformers import CLIPTokenizer
_conversion_module().write_open_clip_tokenizer(tmp_path)
converted = CLIPTokenizer.from_pretrained(tmp_path / "tokenizer")
prompts = ["", "A dog runs past a red car!", "雨の日, cinematic"]
encoded = converted(
prompts,
padding="max_length",
truncation=True,
max_length=77,
return_tensors="pt",
)
actual = encoded.input_ids.masked_fill(encoded.attention_mask == 0, 0)
expected = open_clip.get_tokenizer("ViT-H-14-378-quickgelu")(prompts)
torch.testing.assert_close(actual, expected)
def _map_open_clip_text(state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
return _conversion_module().map_open_clip_text_state(state)
def _map_open_clip_vision(state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
return _conversion_module().map_open_clip_vision_state(state)
def test_mmaudio_dfn_clip_implementation_parity() -> None:
if not torch.cuda.is_available():
pytest.skip("DFN CLIP implementation parity requires CUDA")
from open_clip.model import CLIP, CLIPTextCfg, CLIPVisionCfg
from fastvideo.configs.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextArchConfig,
MMAudioDFNCLIPTextConfig,
MMAudioDFNCLIPVisionArchConfig,
MMAudioDFNCLIPVisionConfig,
)
from fastvideo.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextEncoder,
MMAudioDFNCLIPVisionEncoder,
)
from fastvideo.distributed import (
cleanup_dist_env_and_memory,
maybe_init_distributed_environment_and_model_parallel,
)
from fastvideo.forward_context import set_forward_context
from fastvideo.platforms import AttentionBackendEnum
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29591")
maybe_init_distributed_environment_and_model_parallel(1, 1)
embed_dim = 16
official = CLIP(
embed_dim=embed_dim,
vision_cfg=CLIPVisionCfg(layers=2, width=32, head_width=8, patch_size=8, image_size=16),
text_cfg=CLIPTextCfg(context_length=8, vocab_size=32, width=32, heads=4, layers=2, eos_id=31),
quick_gelu=True,
)
backends = (AttentionBackendEnum.TORCH_SDPA,)
text_arch = MMAudioDFNCLIPTextArchConfig(
vocab_size=32,
hidden_size=32,
intermediate_size=128,
projection_dim=embed_dim,
num_hidden_layers=2,
num_attention_heads=4,
max_position_embeddings=8,
text_len=8,
eos_token_id=31,
_supported_attention_backends=backends,
)
vision_arch = MMAudioDFNCLIPVisionArchConfig(
hidden_size=32,
intermediate_size=128,
projection_dim=embed_dim,
num_hidden_layers=2,
num_attention_heads=4,
image_size=16,
patch_size=8,
_supported_attention_backends=backends,
)
text_encoder = MMAudioDFNCLIPTextEncoder(MMAudioDFNCLIPTextConfig(arch_config=text_arch))
vision_encoder = MMAudioDFNCLIPVisionEncoder(MMAudioDFNCLIPVisionConfig(arch_config=vision_arch))
state = official.state_dict()
text_encoder.load_state_dict(_map_open_clip_text(state), strict=True)
vision_encoder.load_state_dict(_map_open_clip_vision(state), strict=True)
device = torch.device("cuda:0")
official.to(device).eval()
text_encoder.to(device).eval()
vision_encoder.to(device).eval()
tokens = torch.tensor([[1, 2, 3, 31, 0, 0, 0, 0]], device=device)
images = torch.randn((1, 3, 16, 16), generator=torch.Generator(device=device).manual_seed(1234), device=device)
with torch.inference_mode():
expected_text = official.token_embedding(tokens)
expected_text = expected_text + official.positional_embedding
expected_text = official.transformer(expected_text, attn_mask=official.attn_mask)
expected_text = F.normalize(official.ln_final(expected_text), dim=-1)
expected_image = official.encode_image(images, normalize=True)
with set_forward_context(current_timestep=0, attn_metadata=None):
actual_text = text_encoder(tokens).last_hidden_state
actual_image = vision_encoder(images).last_hidden_state
torch.testing.assert_close(actual_text, expected_text, atol=2e-5, rtol=2e-5)
torch.testing.assert_close(actual_image, expected_image, atol=2e-5, rtol=2e-5)
cleanup_dist_env_and_memory()
def test_mmaudio_dfn5b_real_weight_parity() -> None:
if not DFN5B_DIR.is_dir():
pytest.skip("DFN5B assets are absent. Set MMAUDIO_DFN5B_DIR to a local apple/DFN5B-CLIP-ViT-H-14-384 snapshot.")
if not torch.cuda.is_available():
pytest.skip("DFN5B parity requires CUDA")
import open_clip
from fastvideo.configs.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextConfig,
MMAudioDFNCLIPVisionConfig,
)
from fastvideo.distributed import (
cleanup_dist_env_and_memory,
maybe_init_distributed_environment_and_model_parallel,
)
from fastvideo.forward_context import set_forward_context
from fastvideo.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPTextEncoder,
MMAudioDFNCLIPVisionEncoder,
)
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29592")
maybe_init_distributed_environment_and_model_parallel(1, 1)
try:
official = open_clip.create_model_from_pretrained(
f"local-dir:{DFN5B_DIR}", return_transform=False)
state = official.state_dict()
text_encoder = MMAudioDFNCLIPTextEncoder(
MMAudioDFNCLIPTextConfig())
vision_encoder = MMAudioDFNCLIPVisionEncoder(
MMAudioDFNCLIPVisionConfig())
text_encoder.load_state_dict(_map_open_clip_text(state), strict=True)
vision_encoder.load_state_dict(_map_open_clip_vision(state),
strict=True)
device = torch.device("cuda:0")
official.to(device).eval()
text_encoder.to(device).eval()
vision_encoder.to(device).eval()
tokens = open_clip.get_tokenizer("ViT-H-14-378-quickgelu")(
["A dog runs past a red car!"]).to(device)
frames = torch.rand(
(1, 3, 384, 384),
generator=torch.Generator(device=device).manual_seed(9012),
device=device,
)
mean = torch.tensor([0.48145466, 0.4578275, 0.40821073],
device=device).view(1, 3, 1, 1)
std = torch.tensor([0.26862954, 0.26130258, 0.27577711],
device=device).view(1, 3, 1, 1)
frames = (frames - mean) / std
with torch.inference_mode():
expected_text = official.token_embedding(tokens)
expected_text = expected_text + official.positional_embedding
expected_text = official.transformer(
expected_text, attn_mask=official.attn_mask)
expected_text = F.normalize(official.ln_final(expected_text),
dim=-1)
expected_image = official.encode_image(frames, normalize=True)
with set_forward_context(current_timestep=0,
attn_metadata=None):
actual_text = text_encoder(tokens).last_hidden_state
actual_image = vision_encoder(frames).last_hidden_state
text_error = (actual_text.float() - expected_text.float()).abs()
image_error = (actual_image.float() - expected_image.float()).abs()
print("text_max_abs", text_error.max().item())
print("image_max_abs", image_error.max().item())
torch.testing.assert_close(actual_text,
expected_text,
atol=2e-5,
rtol=2e-5)
torch.testing.assert_close(actual_image,
expected_image,
atol=2e-5,
rtol=2e-5)
finally:
cleanup_dist_env_and_memory()
@@ -0,0 +1,62 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared component-loader extensions required by V2A pipelines."""
from __future__ import annotations
import json
from types import SimpleNamespace
def test_indexed_image_encoder_uses_matching_config_and_precision(tmp_path) -> None:
from fastvideo.configs.models.encoders.mmaudio_clip import (
MMAudioDFNCLIPVisionConfig,
)
from fastvideo.configs.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerConfig,
)
from fastvideo.models.loader.component_loader import ImageEncoderLoader
component = tmp_path / "image_encoder_2"
component.mkdir()
with (component / "config.json").open("w", encoding="utf-8") as handle:
json.dump(
{
"architectures": ["MMAudioSynchformerVisualEncoder"],
"segment_stride": 4,
},
handle,
)
class CaptureLoader(ImageEncoderLoader):
def load_model(
self,
model_path,
model_config,
target_device,
fastvideo_args,
dtype="fp16",
use_text_encoder_override=False,
cpu_offload=None,
):
del model_path, target_device, fastvideo_args, use_text_encoder_override
return model_config, dtype, cpu_offload
vision_config = MMAudioDFNCLIPVisionConfig()
sync_config = MMAudioSynchformerConfig()
pipeline_config = SimpleNamespace(
image_encoder_config=vision_config,
image_encoder_precision="fp32",
image_encoder_configs=(vision_config, sync_config),
image_encoder_precisions=("bf16", "fp16"),
)
args = SimpleNamespace(
pipeline_config=pipeline_config,
image_encoder_cpu_offload=True,
)
selected, precision, cpu_offload = CaptureLoader().load(str(component), args)
assert selected is sync_config
assert selected.arch_config.segment_stride == 4
assert vision_config.arch_config.image_size == 378
assert precision == "fp16"
assert cpu_offload is True
@@ -0,0 +1,245 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio pipeline contracts and opt-in real end-to-end parity."""
from __future__ import annotations
import gc
import math
import os
from pathlib import Path
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REPO = REPO_ROOT.parent / "MMAudio"
CONVERTED_MODEL = REPO_ROOT / "converted_weights/mmaudio/large_44k_v2"
DFN5B_DIR = REPO_ROOT / "official_weights/mmaudio/DFN5B-CLIP-ViT-H-14-384"
BIGVGAN_DIR = REPO_ROOT / "official_weights/mmaudio/bigvgan_v2_44khz_128band_512x"
def test_mmaudio_pipeline_config_registry_and_preset() -> None:
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.registry import get_model_info, get_preset_selection
assert WorkloadType.from_string("v2a") is WorkloadType.V2A
assert WorkloadType.from_string("t2a") is WorkloadType.T2A
config = MMAudioV2AConfig()
assert config.dit_config.prefix == "MMAudio"
assert len(config.image_encoder_configs or ()) == 2
assert config.text_encoder_precisions == ("bf16",)
assert config.image_encoder_precisions == ("bf16", "bf16")
assert config.audio_decoder_precision == "bf16"
assert config.vocoder_precision == "bf16"
assert config.duration_s == 8.0
assert config.max_audio_duration_s is None
assert get_preset_selection("FastVideo/MMAudio-large-44k-v2-Diffusers") == (
"mmaudio_large_44k_v2",
"mmaudio",
)
if CONVERTED_MODEL.is_dir():
info = get_model_info(str(CONVERTED_MODEL), workload_type=WorkloadType.V2A)
assert info.pipeline_cls.__name__ == "MMAudioPipeline"
assert info.pipeline_config_cls is MMAudioV2AConfig
def test_mmaudio_published_sequence_lengths() -> None:
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
from fastvideo.pipelines.basic.mmaudio.stages import mmaudio_sequence_lengths
assert mmaudio_sequence_lengths(8.0, MMAudioV2AConfig()) == (345, 64, 192)
assert mmaudio_sequence_lengths(10.0, MMAudioV2AConfig()) == (431, 80, 240)
assert mmaudio_sequence_lengths(2.0, MMAudioV2AConfig()) == (87, 16, 40)
def test_mmaudio_bf16_euler_stage_matches_official() -> None:
if not torch.cuda.is_available():
pytest.skip("MMAudio BF16 scheduler parity requires CUDA")
from mmaudio.model.flow_matching import FlowMatching
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
from fastvideo.pipelines.basic.mmaudio.stages import MMAudioDenoisingStage
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
class ToyTransformer:
@staticmethod
def guided_flow(timestep, latent, conditions, empty_conditions, guidance_scale):
del conditions, empty_conditions, guidance_scale
return torch.tanh(latent * 0.125 + timestep)
@staticmethod
def unnormalize(latent):
return latent
device = torch.device("cuda:0")
initial = torch.randn(
(1, 87, 40),
device=device,
dtype=torch.bfloat16,
generator=torch.Generator(device=device).manual_seed(42),
)
official = FlowMatching(min_sigma=0.0, inference_mode="euler", num_steps=25)
expected = official.to_data(lambda time, latent: torch.tanh(latent * 0.125 + time), initial.clone())
scheduler = FlowMatchEulerDiscreteScheduler(
shift=1.0,
invert_sigmas=True,
sigma_min=0.0,
use_reference_discrete_timesteps=True,
)
batch = ForwardBatch(data_type="audio", latents=initial.clone(), num_inference_steps=25, guidance_scale=4.5)
batch.extra["mmaudio_conditions"] = object()
batch.extra["mmaudio_empty_conditions"] = object()
actual = MMAudioDenoisingStage(ToyTransformer(), scheduler)(
batch,
FastVideoArgs(model_path="test", pipeline_config=MMAudioV2AConfig()),
).latents
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
@pytest.mark.skipif(
os.environ.get("MMAUDIO_RUN_PIPELINE_PARITY") != "1",
reason="Set MMAUDIO_RUN_PIPELINE_PARITY=1 and MMAUDIO_PARITY_VIDEO to run the 25-step real-weight gate.",
)
def test_mmaudio_real_v2a_pipeline_waveform_parity(monkeypatch: pytest.MonkeyPatch) -> None:
"""Compare official and FastVideo waveforms with the same 2s video/seed."""
video_value = os.environ.get("MMAUDIO_PARITY_VIDEO")
if not video_value:
pytest.skip("Set MMAUDIO_PARITY_VIDEO to a video containing at least two seconds.")
video_path = Path(video_value)
required = (
OFFICIAL_REPO / "weights/mmaudio_large_44k_v2.pth",
OFFICIAL_REPO / "ext_weights/v1-44.pth",
OFFICIAL_REPO / "ext_weights/synchformer_state_dict.pth",
CONVERTED_MODEL / "model_index.json",
DFN5B_DIR / "open_clip_pytorch_model.bin",
BIGVGAN_DIR / "bigvgan_generator.pt",
)
missing = [str(path) for path in required if not path.exists()]
if missing:
pytest.skip(f"MMAudio real pipeline assets are missing: {missing}")
if not torch.cuda.is_available():
pytest.skip("MMAudio real pipeline parity requires CUDA")
import open_clip
import mmaudio.ext.autoencoder.autoencoder as official_autoencoder
import mmaudio.model.utils.features_utils as official_features_module
from mmaudio.eval_utils import generate as official_generate
from mmaudio.eval_utils import load_video as official_load_video
from mmaudio.ext.bigvgan_v2.bigvgan import BigVGAN as OfficialBigVGAN
from mmaudio.model.flow_matching import FlowMatching
from mmaudio.model.networks import get_my_mmaudio
from mmaudio.model.utils.features_utils import FeaturesUtils
monkeypatch.setattr(
official_features_module,
"create_model_from_pretrained",
lambda *args, **kwargs: open_clip.create_model_from_pretrained(
f"local-dir:{DFN5B_DIR}", return_transform=False
),
)
class LocalBigVGAN:
@classmethod
def from_pretrained(cls, model_id, **kwargs):
del model_id
return OfficialBigVGAN.from_pretrained(str(BIGVGAN_DIR), **kwargs)
monkeypatch.setattr(official_autoencoder, "BigVGANv2", LocalBigVGAN)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
device = torch.device("cuda:0")
dtype = torch.bfloat16
with torch.inference_mode():
official_model = get_my_mmaudio("large_44k_v2").to(device, dtype).eval()
official_model.load_weights(
torch.load(required[0], map_location=device, weights_only=True)
)
official_features = FeaturesUtils(
tod_vae_ckpt=required[1],
synchformer_ckpt=required[2],
enable_conditions=True,
mode="44k",
need_vae_encoder=False,
).to(device, dtype).eval()
video = official_load_video(video_path, 2.0, load_all_frames=False)
duration = video.duration_sec
official_model.update_seq_lengths(
math.ceil(duration * 44100 / 512 / 2),
int(duration * 8),
int((((duration * 25) - 16) // 8 + 1) * 16 / 2),
)
expected = official_generate(
video.clip_frames[None],
video.sync_frames[None],
["A dog runs past a red car!"],
negative_text=[""],
feature_utils=official_features,
net=official_model,
fm=FlowMatching(min_sigma=0, inference_mode="euler", num_steps=25),
rng=torch.Generator(device=device).manual_seed(42),
cfg_strength=4.5,
).float().cpu()
del official_features, official_model
gc.collect()
torch.cuda.empty_cache()
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
from fastvideo.distributed import cleanup_dist_env_and_memory
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.pipelines.basic.mmaudio import MMAudioPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
args = FastVideoArgs(
model_path=str(CONVERTED_MODEL),
workload_type=WorkloadType.V2A,
pipeline_config=MMAudioV2AConfig(),
tp_size=1,
sp_size=1,
hsdp_shard_dim=1,
num_gpus=1,
dit_cpu_offload=False,
dit_layerwise_offload=False,
text_encoder_cpu_offload=False,
image_encoder_cpu_offload=False,
vae_cpu_offload=False,
pin_cpu_memory=False,
)
try:
pipeline = MMAudioPipeline(str(CONVERTED_MODEL), args)
pipeline.post_init()
output = pipeline.forward(
ForwardBatch(
data_type="video",
video_path=str(video_path),
prompt="A dog runs past a red car!",
negative_prompt="",
audio_start_in_s=0.0,
audio_end_in_s=2.0,
num_inference_steps=25,
guidance_scale=4.5,
seed=42,
num_videos_per_prompt=1,
height=8,
width=8,
num_frames=1,
save_video=False,
return_frames=False,
),
args,
)
actual = output.extra["decoded_audio"]
assert output.extra["audio_sample_rate"] == 44100
assert output.extra["audio_only"] is True
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
pipeline.close()
finally:
cleanup_dist_env_and_memory()
@@ -0,0 +1,37 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio flow-matching scheduler reuse parity."""
from __future__ import annotations
import torch
def test_mmaudio_euler_schedule_matches_fastvideo_flow_scheduler() -> None:
from mmaudio.model.flow_matching import FlowMatching
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
num_steps = 25
official = FlowMatching(min_sigma=0.0, inference_mode="euler", num_steps=num_steps)
fastvideo = FlowMatchEulerDiscreteScheduler(
shift=1.0,
invert_sigmas=True,
sigma_min=0.0,
use_reference_discrete_timesteps=True,
)
fastvideo.set_timesteps(num_steps, device="cpu")
initial = torch.randn((2, 7, 4), generator=torch.Generator().manual_seed(1234))
def flow(time: torch.Tensor, sample: torch.Tensor) -> torch.Tensor:
time = torch.as_tensor(time, dtype=sample.dtype, device=sample.device)
return torch.tanh(sample * 0.125 + time)
expected = official.to_data(flow, initial.clone())
actual = initial.clone()
for timestep in fastvideo.timesteps:
model_output = flow(timestep / fastvideo.config.num_train_timesteps, actual)
actual = fastvideo.step(model_output, timestep, actual).prev_sample
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
@@ -0,0 +1,146 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio Synchformer parity against the official reference.
Coverage scope: implementation_subcomponent. Production-loader coverage is
added after the converted component directory exists.
"""
from __future__ import annotations
import os
import gc
from pathlib import Path
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_WEIGHTS = Path(
os.environ.get(
"MMAUDIO_SYNCHFORMER_WEIGHTS",
REPO_ROOT.parent / "MMAudio/ext_weights/synchformer_state_dict.pth",
)
)
def _build_models(device: torch.device, dtype: torch.dtype):
from mmaudio.ext.synchformer import Synchformer
from fastvideo.configs.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerConfig,
)
from fastvideo.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerVisualEncoder,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
with torch.device(device), set_default_torch_dtype(dtype):
official = Synchformer()
fastvideo = MMAudioSynchformerVisualEncoder(MMAudioSynchformerConfig())
return official, fastvideo
def test_mmaudio_synchformer_state_structure() -> None:
# MotionFormer calls ``Tensor.item`` while building its stochastic-depth
# schedule, so it cannot be constructed on the meta device. Instantiate
# the two large models sequentially to keep peak host memory bounded.
from mmaudio.ext.synchformer import Synchformer
official = Synchformer()
official_state = official.state_dict()
official_shapes = {name: tensor.shape for name, tensor in official_state.items()}
del official_state, official
gc.collect()
from fastvideo.configs.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerConfig,
)
from fastvideo.models.encoders.mmaudio_synchformer import (
MMAudioSynchformerVisualEncoder,
)
fastvideo = MMAudioSynchformerVisualEncoder(MMAudioSynchformerConfig())
fastvideo_shapes = {name: tensor.shape for name, tensor in fastvideo.state_dict().items()}
assert official_shapes == fastvideo_shapes
def test_mmaudio_synchformer_numerical_parity() -> None:
if not torch.cuda.is_available():
pytest.skip("MMAudio Synchformer parity requires CUDA")
if not OFFICIAL_WEIGHTS.is_file():
pytest.skip(
"Official MMAudio Synchformer weights are absent. Set "
"MMAUDIO_SYNCHFORMER_WEIGHTS or download the released checkpoint."
)
device = torch.device("cuda:0")
dtype = torch.bfloat16
official, fastvideo = _build_models(device, dtype)
state = torch.load(OFFICIAL_WEIGHTS, map_location="cpu", weights_only=True)
official.load_state_dict(state, strict=True)
visual_state = {
name: tensor
for name, tensor in state.items()
if name.startswith("vfeat_extractor.")
}
fastvideo.load_state_dict(visual_state, strict=True)
official.eval()
fastvideo.eval()
generator = torch.Generator(device=device).manual_seed(1234)
segments = torch.randn((1, 1, 16, 3, 224, 224), generator=generator, device=device, dtype=dtype)
with torch.inference_mode(), torch.autocast("cuda", dtype=dtype):
expected = official(segments)
actual = fastvideo.forward_segmented(segments)
difference = (actual.float() - expected.float()).abs()
print("max_abs", difference.max().item())
print("mean_abs", difference.mean().item())
torch.testing.assert_close(actual, expected, atol=1e-3, rtol=1e-3)
def test_mmaudio_synchformer_feature_contract_parity() -> None:
"""Cover official windowing, batching, flatten order, and model forward."""
if not torch.cuda.is_available():
pytest.skip("MMAudio Synchformer feature parity requires CUDA")
if not OFFICIAL_WEIGHTS.is_file():
pytest.skip("Official MMAudio Synchformer weights are absent")
from mmaudio.model.utils.features_utils import FeaturesUtils
device = torch.device("cuda:0")
dtype = torch.bfloat16
official, fastvideo = _build_models(device, dtype)
state = torch.load(OFFICIAL_WEIGHTS, map_location="cpu", weights_only=True)
official.load_state_dict(state, strict=True)
fastvideo.load_state_dict(
{
name: tensor
for name, tensor in state.items()
if name.startswith("vfeat_extractor.")
},
strict=True,
)
official.eval()
fastvideo.eval()
feature_utils = FeaturesUtils(enable_conditions=False,
mode="44k",
need_vae_encoder=False)
feature_utils.synchformer = official
video = torch.randn(
(2, 24, 3, 224, 224),
generator=torch.Generator(device=device).manual_seed(5678),
device=device,
dtype=dtype,
).clamp_(-1, 1)
with torch.inference_mode(), torch.autocast("cuda", dtype=dtype):
expected = feature_utils.encode_video_with_sync(video, batch_size=2)
actual = fastvideo(video).last_hidden_state
assert expected.shape == actual.shape == (2, 16, 768)
torch.testing.assert_close(actual, expected, atol=1e-3, rtol=1e-3)
@@ -0,0 +1,126 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio transformer parity against the official reference.
Coverage scope: implementation_subcomponent. Production-loader coverage is
added after the converted Diffusers component directory exists.
"""
from __future__ import annotations
import os
from pathlib import Path
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_WEIGHTS = Path(
os.environ.get(
"MMAUDIO_TRANSFORMER_WEIGHTS",
REPO_ROOT.parent / "MMAudio/weights/mmaudio_large_44k_v2.pth",
)
)
def _require_cuda_and_weights() -> None:
if not torch.cuda.is_available():
pytest.skip("MMAudio full-transformer parity requires CUDA")
if not OFFICIAL_WEIGHTS.is_file():
pytest.skip(
"Official MMAudio transformer weights are absent. Set "
"MMAUDIO_TRANSFORMER_WEIGHTS or download large_44k_v2 assets."
)
def test_mmaudio_transformer_implementation_parity() -> None:
from mmaudio.model.networks import MMAudio
from fastvideo.configs.models.dits.mmaudio import (
MMAudioArchConfig,
MMAudioTransformerConfig,
)
from fastvideo.models.dits.mmaudio import MMAudioTransformer
kwargs = {
"latent_dim": 8,
"clip_dim": 16,
"sync_dim": 12,
"text_dim": 16,
"hidden_dim": 64,
"depth": 3,
"fused_depth": 2,
"num_heads": 4,
"mlp_ratio": 4.0,
"latent_seq_len": 5,
"clip_seq_len": 3,
"sync_seq_len": 8,
"text_seq_len": 4,
"v2": True,
}
official = MMAudio(**kwargs)
fastvideo = MMAudioTransformer(MMAudioTransformerConfig(arch_config=MMAudioArchConfig(**kwargs)), hf_config={})
assert {name: tensor.shape for name, tensor in official.state_dict().items()} == {
name: tensor.shape for name, tensor in fastvideo.state_dict().items()
}
fastvideo.load_state_dict(official.state_dict(), strict=True)
generator = torch.Generator().manual_seed(1234)
latent = torch.randn((1, 5, 8), generator=generator)
clip = torch.randn((1, 3, 16), generator=generator)
sync = torch.randn((1, 8, 12), generator=generator)
text = torch.randn((1, 4, 16), generator=generator)
timestep = torch.tensor([0.375])
official.eval()
fastvideo.eval()
with torch.inference_mode():
expected = official(latent, clip, sync, text, timestep)
actual = fastvideo(latent, (clip, sync, text), timestep)
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
def test_mmaudio_large_44k_v2_transformer_parity() -> None:
_require_cuda_and_weights()
from mmaudio.model.networks import large_44k_v2
from fastvideo.configs.models.dits.mmaudio import MMAudioTransformerConfig
from fastvideo.models.dits.mmaudio import MMAudioTransformer
from fastvideo.models.loader.utils import set_default_torch_dtype
device = torch.device("cuda:0")
dtype = torch.bfloat16
state = torch.load(OFFICIAL_WEIGHTS, map_location="cpu", weights_only=True)
with torch.device(device), set_default_torch_dtype(dtype):
official = large_44k_v2()
fastvideo = MMAudioTransformer(MMAudioTransformerConfig(), hf_config={})
official.load_weights(dict(state))
# The released checkpoint contains a stale derived buffer which the
# official ``load_weights`` explicitly discards. Validate it before
# applying the same canonicalization used by the converter.
checkpoint_freqs = state.pop("t_embed.freqs")
torch.testing.assert_close(checkpoint_freqs,
fastvideo.t_embed.freqs.cpu(),
atol=1e-6,
rtol=1e-6)
missing, unexpected = fastvideo.load_state_dict(state, strict=True)
assert missing == []
assert unexpected == []
del state
official.eval()
fastvideo.eval()
generator = torch.Generator(device=device).manual_seed(1234)
latent = torch.randn((1, 345, 40), generator=generator, device=device, dtype=dtype)
clip = torch.randn((1, 64, 1024), generator=generator, device=device, dtype=dtype)
sync = torch.randn((1, 192, 768), generator=generator, device=device, dtype=dtype)
text = torch.randn((1, 77, 1024), generator=generator, device=device, dtype=dtype)
timestep = torch.tensor([0.375], device=device, dtype=dtype)
with torch.inference_mode(), torch.autocast("cuda", dtype=dtype):
expected = official(latent, clip, sync, text, timestep)
actual = fastvideo(latent, (clip, sync, text), timestep)
difference = (actual.float() - expected.float()).abs()
print("max_abs", difference.max().item())
print("mean_abs", difference.mean().item())
torch.testing.assert_close(actual, expected, atol=1e-3, rtol=1e-3)
@@ -0,0 +1,96 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio BigVGAN-v2 parity against the canonical NVIDIA checkpoint."""
from __future__ import annotations
import json
import os
from pathlib import Path
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[3]
BIGVGAN_DIR = Path(
os.environ.get(
"MMAUDIO_BIGVGAN_DIR",
REPO_ROOT / "official_weights/mmaudio/bigvgan_v2_44khz_128band_512x",
)
)
def _require_assets() -> None:
required = (BIGVGAN_DIR / "config.json", BIGVGAN_DIR / "bigvgan_generator.pt")
if not all(path.is_file() for path in required):
pytest.skip(
"Canonical BigVGAN-v2 assets are absent. Set MMAUDIO_BIGVGAN_DIR "
"to a local nvidia/bigvgan_v2_44khz_128band_512x snapshot."
)
def test_bigvgan_v2_implementation_parity() -> None:
from mmaudio.ext.bigvgan_v2.bigvgan import BigVGAN
from mmaudio.ext.bigvgan_v2.env import AttrDict
from fastvideo.models.audio.bigvgan import BigVGANV2
config = {
"num_mels": 4,
"upsample_initial_channel": 16,
"resblock": "1",
"resblock_kernel_sizes": [3],
"resblock_dilation_sizes": [[1, 3, 5]],
"upsample_rates": [2],
"upsample_kernel_sizes": [4],
"activation": "snakebeta",
"snake_logscale": True,
"use_bias_at_final": True,
"use_tanh_at_final": True,
}
official = BigVGAN(AttrDict(config), use_cuda_kernel=False)
fastvideo = BigVGANV2(config)
assert {name: tensor.shape for name, tensor in official.state_dict().items()} == {
name: tensor.shape for name, tensor in fastvideo.state_dict().items()
}
fastvideo.load_state_dict(official.state_dict(), strict=True)
official.remove_weight_norm()
fastvideo.remove_weight_norm()
mel = torch.randn((1, 4, 8), generator=torch.Generator().manual_seed(1234))
with torch.inference_mode():
expected = official(mel)
actual = fastvideo(mel)
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
def test_mmaudio_bigvgan_v2_parity() -> None:
_require_assets()
if not torch.cuda.is_available():
pytest.skip("MMAudio BigVGAN parity requires CUDA")
from mmaudio.ext.bigvgan_v2.bigvgan import BigVGAN, load_hparams_from_json
from fastvideo.models.audio.bigvgan import BigVGANV2
config_path = BIGVGAN_DIR / "config.json"
with config_path.open(encoding="utf-8") as handle:
config = json.load(handle)
official = BigVGAN(load_hparams_from_json(config_path), use_cuda_kernel=False)
fastvideo = BigVGANV2(config)
state = torch.load(BIGVGAN_DIR / "bigvgan_generator.pt", map_location="cpu", weights_only=True)["generator"]
official.load_state_dict(state, strict=True)
fastvideo.load_state_dict(state, strict=True)
assert {name: tensor.shape for name, tensor in official.state_dict().items()} == {
name: tensor.shape for name, tensor in fastvideo.state_dict().items()
}
device = torch.device("cuda:0")
official.remove_weight_norm()
fastvideo.remove_weight_norm()
official.to(device).eval()
fastvideo.to(device).eval()
mel = torch.randn((1, 128, 8), generator=torch.Generator(device=device).manual_seed(1234), device=device)
with torch.inference_mode():
expected = official(mel)
actual = fastvideo(mel)
torch.testing.assert_close(actual, expected, atol=1e-5, rtol=1e-5)