def __init__(
self,
config: MMAudioTransformerConfig,
hf_config: dict[str, Any],
latent_mean: torch.Tensor | None = None,
latent_std: torch.Tensor | None = None,
empty_string_feat: torch.Tensor | None = None,
**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)
]
)
def fixed_tensor(
value: torch.Tensor | None,
shape: tuple[int, ...],
*,
fill_value: float,
name: str,
) -> torch.Tensor:
if value is None:
return torch.full(
shape,
fill_value,
device=self.sync_pos_emb.device,
dtype=self.sync_pos_emb.dtype,
)
if value.numel() != math.prod(shape):
raise ValueError(
f"MMAudio {name} must have {math.prod(shape)} values, "
f"got {value.numel()}."
)
return value.to(
device=self.sync_pos_emb.device,
dtype=self.sync_pos_emb.dtype,
).reshape(shape).clone()
self.latent_mean = nn.Parameter(
fixed_tensor(
latent_mean,
(1, 1, arch.latent_dim),
fill_value=float("nan"),
name="latent_mean",
),
requires_grad=False,
)
self.latent_std = nn.Parameter(
fixed_tensor(
latent_std,
(1, 1, arch.latent_dim),
fill_value=float("nan"),
name="latent_std",
),
requires_grad=False,
)
self.empty_string_feat = nn.Parameter(
fixed_tensor(
empty_string_feat,
(arch.text_seq_len, arch.text_dim),
fill_value=0.0,
name="empty_string_feat",
),
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__()