Skip to content

vllm_omni.model_executor.models.minimax_music3.dit

Condition encoder and flow-matching transformer for MiniMax Music 3.

Stage 1 turns the AR stage's conditioning frames into a DAC latent in two steps. The condition encoder collapses the eight 4096-wide layer slices of a frame into one 2048-wide vector and stretches the result from the AR frame grid onto the vocoder's latent grid. The transformer is then solved as a flow-matching velocity field: a fixed-step Euler integrator walks Gaussian noise to the latent the vocoder expects, under classifier-free guidance.

Module names mirror the checkpoint's condition_encoder/ and transformer/ component folders key for key, with one deliberate exception: the per-block to_q/to_k/to_v projections are fused into a single to_qkv so each block runs one GEMM instead of three. :func:remap_transformer_state performs that fusion, and is the only place the checkpoint's layout is reinterpreted.

Attention uses the project's diffusion attention layer when its selected backend supports the runtime dtype. This LLM_GENERATION stage does not set a diffusion backend and therefore gets the platform default. CUDA commonly resolves that default to FlashAttention, cuDNN, or FlashInfer, so the model routes its native float32 decode through SDPA while retaining the selected backend for supported lower-precision inputs. An explicitly selected backend remains authoritative.

logger module-attribute

logger = init_logger(__name__)

Attention

Bases: Module

Bidirectional self-attention with partial rotary position embeddings.

Q, K and V are kept in [B, S, H, D], which is the layout the diffusion attention backends consume and return.

backend instance-attribute

backend = _build_native_attention(
    num_heads=self.num_heads,
    head_size=head_dim,
    softmax_scale=self.softmax_scale,
    prefix=prefix,
)

backend_name instance-attribute

backend_name = f"TORCH_SDPA(float32)/{native_backend_name}(lower precision)"

head_dim instance-attribute

head_dim = head_dim

num_heads instance-attribute

num_heads = dim // head_dim

softmax_scale instance-attribute

softmax_scale = float(head_dim ** -0.5)

to_out instance-attribute

to_out = nn.Linear(dim, dim, bias=False)

to_qkv instance-attribute

to_qkv = nn.Linear(dim, dim * 3, bias=False)

forward

forward(
    x: Tensor, rope_cos: Tensor, rope_sin: Tensor
) -> Tensor

FourierFeatures

Bases: Module

Random-Fourier timestep features: [B, 1] -> [B, out_features].

weight instance-attribute

weight = nn.Parameter(
    torch.empty(out_features // 2, in_features)
)

forward

forward(value: Tensor) -> Tensor

MiniMaxMusic3ConditionEncoder

Bases: Module

Collapse an AR frame's eight layer slices and resample onto mel time.

layer_scale instance-attribute

layer_scale = nn.Parameter(torch.ones(1))

layer_weight_logits instance-attribute

layer_weight_logits = nn.Parameter(
    torch.zeros(NUM_CODEBOOKS)
)

proj instance-attribute

proj = nn.Conv1d(
    BACKBONE_HIDDEN_SIZE,
    _CONDITION_DIM,
    kernel_size=3,
    padding=1,
)

aligned_condition

aligned_condition(hidden: Tensor) -> Tensor

Project AR hidden states and resample them onto the latent grid.

aligned_mel_length

aligned_mel_length(frames: int) -> int

Return how many vocoder latent frames frames AR frames cover.

AR frames sit on a 24 kHz / 960-hop grid (25 fps) and the vocoder latent on a 44.1 kHz / 512-hop grid (about 86 fps), so one AR frame is roughly 3.445 latent frames. The result is truncated, never rounded.

condition

condition(hidden: Tensor) -> Tensor

Project [B, T, 32768] AR frames to [B, 2048, T].

Raises:

Type Description
ValueError

If the conditioning is not [B, T, 32768].

MiniMaxMusic3FlowMatchingDiT

Bases: Module

Condition encoder plus velocity field, with the Euler solver on top.

condition_encoder instance-attribute

condition_encoder = MiniMaxMusic3ConditionEncoder()

transformer instance-attribute

aligned_condition

aligned_condition(hidden: Tensor) -> Tensor

aligned_mel_length

aligned_mel_length(frames: int) -> int

condition

condition(hidden: Tensor) -> Tensor

forward

forward(
    align: Tensor,
    *,
    generator: Generator,
    initial_latent: Tensor | None = None,
    initial_condition: Tensor | None = None,
    num_steps: int = DEFAULT_DIT_STEPS,
    cfg_scale: float = DEFAULT_DIT_CFG_SCALE,
) -> Tensor

Solve an aligned condition into a vocoder latent [1, 128, T_mel].

The previous window's latent is re-imposed at every step rather than only at the start: the leading left frames are pinned to the exact noise-to-latent interpolation the solver would have produced had it generated them, so the two windows join without a seam.

Parameters:

Name Type Description Default
align Tensor

Aligned condition, [1, 2048, T_mel]. Overwritten in place over the prompt span when initial_condition is given, which the caller relies on when it saves the condition for the next window.

required
generator Generator

Seeded generator for the initial noise draw.

required
initial_latent Tensor | None

Previous window's tail latent, or None.

None
initial_condition Tensor | None

Previous window's tail condition, or None.

None
num_steps int

Euler steps.

DEFAULT_DIT_STEPS
cfg_scale float

Classifier-free guidance weight.

DEFAULT_DIT_CFG_SCALE

Returns:

Type Description
Tensor

The solved latent, [1, 128, T_mel].

Raises:

Type Description
ValueError

If num_steps is not positive.

MiniMaxMusic3Transformer1DModel

Bases: Module

The 36-block velocity field: ([B,128,T], [B], [B,2048,T]) -> [B,128,T].

attention_backend_name property

attention_backend_name: str

The attention path the blocks actually run, for startup logging.

postprocess_conv instance-attribute

postprocess_conv = nn.Conv1d(
    DIT_LATENT_CHANNELS, DIT_LATENT_CHANNELS, 1, bias=False
)

preprocess_conv instance-attribute

preprocess_conv = nn.Conv1d(
    _TRANSFORMER_IN_DIM, _TRANSFORMER_IN_DIM, 1, bias=False
)

proj_in instance-attribute

proj_in = nn.Linear(_TRANSFORMER_IN_DIM, _DIM, bias=False)

proj_out instance-attribute

proj_out = nn.Linear(_DIM, DIT_LATENT_CHANNELS, bias=False)

time_embed instance-attribute

time_embed = TimestepEmbedding(_FOURIER_DIM, _DIM)

time_proj instance-attribute

time_proj = FourierFeatures(1, _FOURIER_DIM)

transformer_blocks instance-attribute

transformer_blocks = nn.ModuleList(
    TransformerBlock(
        _DIM,
        head_dim=_HEAD_DIM,
        inner_dim=_FF_INNER_DIM,
        prefix=f"transformer_blocks.{index}.attn",
    )
    for index in range(_NUM_LAYERS)
)

forward

forward(x: Tensor, t: Tensor, condition: Tensor) -> Tensor

Predict the flow velocity at latent x and time t.

Parameters:

Name Type Description Default
x Tensor

Current latent, [B, 128, T_mel].

required
t Tensor

Flow time in [0, 1), [B].

required
condition Tensor

Aligned condition, [B, 2048, T_mel].

required

Returns:

Type Description
Tensor

The velocity, [B, 128, T_mel].

TimestepEmbedding

Bases: Module

The checkpoint's time_embed: Linear, SiLU, Linear.

act instance-attribute

act = nn.SiLU()

linear_1 instance-attribute

linear_1 = nn.Linear(in_dim, dim)

linear_2 instance-attribute

linear_2 = nn.Linear(dim, dim)

forward

forward(x: Tensor) -> Tensor

TransformerBlock

Bases: Module

Pre-norm attention followed by a pre-norm gated feed-forward.

attn instance-attribute

attn = Attention(dim, head_dim=head_dim, prefix=prefix)

ff_in instance-attribute

ff_in = nn.Linear(dim, inner_dim * 2)

ff_out instance-attribute

ff_out = nn.Linear(inner_dim, dim)

norm1 instance-attribute

norm1 = nn.LayerNorm(dim)

norm2 instance-attribute

norm2 = nn.LayerNorm(dim)

forward

forward(
    x: Tensor, rope_cos: Tensor, rope_sin: Tensor
) -> Tensor

remap_transformer_state

remap_transformer_state(
    state: dict[str, Tensor],
) -> dict[str, Tensor]

Rewrite transformer/ checkpoint keys onto this module tree.

Two rewrites, both confined to the attention block:

  • attn.to_q/to_k/to_v.weight are concatenated along the output axis into attn.to_qkv.weight. nn.Linear output feature i reads weight row i, and :meth:Attention.forward splits the projection with chunk(3, dim=-1) into q, k, v, so query rows come first and value rows last.
  • attn.to_out.0.weight loses its index: the checkpoint models to_out as a two-element list whose second element is a parameterless dropout, and inference only needs the projection.

Every other key is copied through untouched, including the fused ff_in whose halves are consumed as value, gate in that order.

Raises:

Type Description
ValueError

If a block is missing one of its three projections.