Skip to content

vllm_omni.model_executor.models.minimax_music3.talker

Stage-0 talker for MiniMax Music 3.

A Qwen3 backbone predicts one audio frame per decode step. Each frame is eight RVQ codebooks deep: the backbone's own lm_head emits c0 and the four-layer depth decoder walks c1..c7. The conditioning signal handed to the acoustic stage is the backbone hidden state concatenated with the seven depth hidden states, so 8 x 4096 = 32768 dimensions per frame.

Two things make this model unusual inside vLLM:

  • Every request decodes as two rows. Classifier-free guidance is not optional here. The conditioned row sees the real prompt and the unconditioned row sees the same prompt with the caption-and-lyrics span replaced by <|audio_cfg|>. Codes are drawn from the guided blend and then fed back to both rows, so the pair stays token-identical. Rows are paired by cfg_pair_id rather than by adjacency, because batch position is not stable across engine steps.
  • The decode input is not a token embedding. After prefill, each step's input is the composed embedding of the previous frame's eight codes. The sampled c0 token id still flows through vLLM's normal machinery so that stop handling and accounting work, but the embedding it would produce is replaced before the backbone runs.

logger module-attribute

logger = init_logger(__name__)

MiniMaxMusic3TalkerForConditionalGeneration

Bases: Module

Qwen3 backbone plus the audio embedding table and RVQ depth decoder.

audio_embeddings instance-attribute

audio_embeddings = nn.Embedding(
    AUDIO_VOCAB_SIZE * (NUM_CODEBOOKS - 1), hidden_size
)

config instance-attribute

config = config

frame_embedding_scale instance-attribute

frame_embedding_scale = float(NUM_CODEBOOKS) ** -0.5

has_postprocess class-attribute instance-attribute

has_postprocess: bool = True

have_multimodal_outputs class-attribute instance-attribute

have_multimodal_outputs: bool = True

lm_head instance-attribute

lm_head = self.model.embed_tokens

logits_processor instance-attribute

logits_processor = LogitsProcessor(int(config.vocab_size))

model instance-attribute

model = Qwen3Model(
    vllm_config=vllm_config,
    prefix=f"{prefix}.model" if prefix else "model",
)

omni_pooler_payload_include_hidden class-attribute instance-attribute

omni_pooler_payload_include_hidden: bool = False

postprocess_uses_hidden_states class-attribute instance-attribute

postprocess_uses_hidden_states: bool = False

postprocess_uses_multimodal_outputs class-attribute instance-attribute

postprocess_uses_multimodal_outputs: bool = False

postprocess_uses_req_infos class-attribute instance-attribute

postprocess_uses_req_infos: bool = True

prefer_model_sampler class-attribute instance-attribute

prefer_model_sampler: bool = True

requires_request_sample_eligibility class-attribute instance-attribute

requires_request_sample_eligibility: bool = True

rvq_decoder instance-attribute

rvq_decoder = RVQDepthDecoder(hidden_size=hidden_size)

skips_model_sampler_output_token_history class-attribute instance-attribute

skips_model_sampler_output_token_history: bool = True

supports_omni_query_start_loc class-attribute instance-attribute

supports_omni_query_start_loc: bool = True

vllm_config instance-attribute

vllm_config = vllm_config

compute_logits

compute_logits(
    hidden_states: Tensor, sampling_metadata: Any = None
) -> Tensor

embed_frames

embed_frames(codes: Tensor) -> Tensor

Compose one conditioning embedding per row from all eight codebooks.

Parameters:

Name Type Description Default
codes Tensor

[rows, 8] codebook indices, c0 zero-based.

required

Returns:

Type Description
Tensor

[rows, hidden] in the embedding table's dtype.

embed_input_ids

embed_input_ids(input_ids: Tensor, **_: Any) -> Tensor

forward

forward(
    input_ids: Tensor,
    positions: Tensor,
    intermediate_tensors: Any | None = None,
    inputs_embeds: Tensor | None = None,
    **kwargs: Any,
) -> Tensor

load_weights

load_weights(
    weights: Iterable[tuple[str, Tensor]],
) -> set[str]

Load the Qwen3 backbone, then the audio modules from their component.

The backbone comes from the stage's own checkpoint folder and loads through vLLM. The audio embedding table and the depth decoder live in a separate component of the repo and are loaded explicitly.

make_omni_output

make_omni_output(
    model_outputs: Tensor | OmniOutput, **kwargs: Any
) -> OmniOutput

on_requests_finished

on_requests_finished(
    finished_req_ids: Iterable[str],
) -> None

Release per-request state once the engine has retired the request.

The key must be normalized the same way it was when the state was created, or every finished song leaks its frame buffer.

postprocess

postprocess(
    hidden_states_slice: Tensor,
    multimodal_outputs: Any = None,
    **req_infos: Any,
) -> dict[str, Any]

Publish this request's next span of conditioning frames.

Mid-stream this is one window, released a frame late so that a window handed downstream is provably not the last. On the final step it is instead a single contiguous span covering every window not yet emitted, because one step produces exactly one payload and the tail cannot be split across two.

prepare_runner_inputs

prepare_runner_inputs(
    *,
    req_ids: list[str],
    num_computed_tokens: Any,
    num_scheduled_tokens: Any,
    input_ids: Tensor,
    positions: Tensor,
    **_: Any,
) -> tuple[Tensor, Tensor]

Bind this step's batch rows to requests, before the forward runs.

This is the only hook that sees the authoritative row-to-request mapping for the step about to execute. vLLM compacts and reorders the persistent batch whenever a request finishes or is preempted, so any state indexed by row number from a previous step is stale by the time it is read. Everything the model carries across steps is therefore keyed by request id and gathered into row order here.

It also resolves which rows are decoding. A row is decoding when its prompt is already computed and it is advancing by a single token; a row still working through a chunked prefill is not, and must keep its real prompt embeddings.

Python runs here on every step, including steps whose forward is replayed from a CUDA graph rather than executed, so this is also where the row-to-request binding that sample reads has to be rebuilt.

sample

sample(
    logits: Tensor, sampling_metadata: Any
) -> SamplerOutput | None

Draw one frame per conditioned request and feed it back to both rows.