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__()