def __init__(
self,
*,
init_from: str | None = None,
training_config: TrainingConfig,
variant: str | None = None,
from_scratch: bool = False,
empty_string_features_path: str | None = None,
latent_statistics_path: str | None = None,
latent_statistics_chunk_size: int = 32,
trainable: bool = True,
allow_v2_training: bool = False,
lora: LoraConfig | dict[str, Any] | None = None,
attention_backend: AttentionBackendEnum | str | None = None,
transformer: torch.nn.Module | None = None,
) -> None:
super().__init__(
trainable=trainable,
lora=lora,
attention_backend=attention_backend,
)
if int(training_config.distributed.sp_size or 1) != 1:
raise ValueError("MMAudio training does not yet support sequence parallelism; set sp_size=1")
if int(training_config.distributed.tp_size or 1) != 1:
raise ValueError("MMAudio training does not yet support tensor parallelism; set tp_size=1")
use_ddp = is_ddp_strategy(training_config)
if use_ddp and (int(training_config.distributed.hsdp_replicate_dim) != 1
or int(training_config.distributed.hsdp_shard_dim) not in {-1, 1}):
raise ValueError("MMAudio DDP strategy does not use HSDP dimensions; "
"set hsdp_replicate_dim=1 and hsdp_shard_dim=1")
self.training_config = training_config
if variant is not None:
try:
pipeline_config_cls = MMAUDIO_PIPELINE_CONFIGS[variant]
except KeyError as exc:
supported = ", ".join(MMAUDIO_PIPELINE_CONFIGS)
raise ValueError(f"Unknown MMAudio variant {variant!r}; expected one of: {supported}") from exc
self.training_config.pipeline_config = pipeline_config_cls()
elif self.training_config.pipeline_config is None:
self.training_config.pipeline_config = MMAudioV2AConfig()
if transformer is None:
if from_scratch:
if init_from:
raise ValueError("MMAudio from_scratch=true cannot also set init_from; use "
"empty_string_features_path for the fixed CLIP embedding")
if variant is None:
raise ValueError("MMAudio from-scratch training requires variant")
if variant not in MMAUDIO_TRAINING_VARIANTS:
supported = ", ".join(MMAUDIO_TRAINING_VARIANTS)
raise ValueError(f"The official MMAudio recipe does not train {variant!r}; "
f"supported from-scratch variants: {supported}")
if not empty_string_features_path:
raise ValueError("MMAudio from-scratch training requires "
"empty_string_features_path")
latent_mean, latent_std = _load_or_compute_latent_stats(
self.training_config,
variant=variant,
cache_path=latent_statistics_path,
chunk_size=latent_statistics_chunk_size,
)
empty_string_feat = _load_empty_string_features(empty_string_features_path)
architecture = MMAUDIO_VARIANT_ARCHITECTURES[variant]
expected_text_shape = (77, 1024)
if tuple(empty_string_feat.shape) != expected_text_shape:
raise ValueError("MMAudio empty-string features must have shape "
f"{expected_text_shape}, got {tuple(empty_string_feat.shape)}")
model_config = get_mmaudio_transformer_config(variant)
hf_config = {
"_class_name": self._transformer_cls_name,
**architecture,
"clip_dim": 1024,
"clip_seq_len": 64,
"sync_dim": 768,
"sync_seq_len": 192,
"text_dim": 1024,
"text_seq_len": 77,
"mlp_ratio": 4.0,
}
default_dtype = PRECISION_TO_TYPE[self.training_config.dit_precision]
init_params = {
"config": model_config,
"hf_config": hf_config,
"latent_mean": latent_mean,
"latent_std": latent_std,
"empty_string_feat": empty_string_feat,
}
if use_ddp:
transformer = build_replicated_model_from_scratch(
MMAudioTransformer,
init_params,
device=get_local_torch_device(),
default_dtype=default_dtype,
seed=self.training_config.data.seed,
)
else:
transformer = build_fsdp_model_from_scratch(
model_cls=MMAudioTransformer,
init_params=init_params,
device=get_local_torch_device(),
hsdp_replicate_dim=self.training_config.distributed.hsdp_replicate_dim,
hsdp_shard_dim=self.training_config.distributed.hsdp_shard_dim,
default_dtype=default_dtype,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
seed=self.training_config.data.seed,
pin_cpu_memory=self.training_config.distributed.pin_cpu_memory,
)
self._init_from = f"scratch:{variant}"
else:
if not init_from:
raise ValueError("MMAudio pretrained training requires init_from, or set "
"from_scratch=true with a variant")
if use_ddp:
raise NotImplementedError("MMAudio DDP currently supports from_scratch=true only; "
"pretrained component loading remains on the FSDP path")
self._init_from = str(init_from)
transformer = load_module_from_path(
model_path=self._init_from,
module_type="transformer",
training_config=self.training_config,
override_transformer_cls_name=self._transformer_cls_name,
attention_backend=self.attention_backend,
)
else:
self._init_from = str(init_from or f"provided:{variant or 'unknown'}")
self.transformer = transformer
if bool(getattr(self.transformer, "v2", False)) and not allow_v2_training:
raise ValueError("The official MMAudio training recipe does not support `_v2` "
"checkpoints. Use small_44k, medium_44k, or large_44k.")
if not self._enable_lora_if_configured(self.transformer):
self.transformer = apply_trainable(self.transformer, trainable=self._trainable)
if self._trainable:
# These are checkpoint statistics/fixed text features in the
# official model, not optimization variables. ``apply_trainable``
# intentionally enables a whole module, so restore their contract.
for name in ("latent_mean", "latent_std", "empty_string_feat"):
parameter = getattr(self.transformer, name, None)
if isinstance(parameter, torch.nn.Parameter):
parameter.requires_grad_(False)
# The two learned null-video tokens are trained by MMAudio.
for name in ("empty_clip_feat", "empty_sync_feat"):
parameter = getattr(self.transformer, name, None)
if isinstance(parameter, torch.nn.Parameter):
parameter.requires_grad_(True)
if use_ddp:
# Official MMAudio wraps the fully initialized FP32 model after
# selecting trainable parameters. It disables buffer broadcasts;
# the fixed latent statistics and positional buffers are identical
# because every rank uses the same initialization seed.
self.transformer = wrap_module_ddp(
self.transformer,
device=get_local_torch_device(),
broadcast_buffers=False,
)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=1.0,
invert_sigmas=True,
sigma_min=0.0,
use_reference_discrete_timesteps=True,
)
self.dataloader: Any = None
self.start_step = 0