Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a0abcd837f | ||
|
|
cceacec997 | ||
|
|
4cabd0a05d | ||
|
|
be61633ac3 | ||
|
|
c4200f9ad7 | ||
|
|
78f719bf08 |
@@ -23,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_audio/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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", )
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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, )
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Shared vendored Synchformer visual backbone."""
|
||||
|
||||
from fastvideo.third_party.synchformer.motionformer import MotionFormer
|
||||
|
||||
__all__ = ["MotionFormer"]
|
||||
+1
@@ -1,3 +1,4 @@
|
||||
# Shared MotionFormer configuration used by Synchformer.
|
||||
TRAIN:
|
||||
ENABLE: True
|
||||
DATASET: Ssv2
|
||||
+400
@@ -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
@@ -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
|
||||
@@ -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
@@ -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())
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user