Skip to content

vllm_omni.model_executor.models.personaplex.personaplex_talker

PersonaPlex talker: the temporal transformer as a vLLM-native omni AR stage.

This is the stage-0 (LLM_AR) model the OmniGPUModelRunner drives. It composes three pieces, each verified in isolation against Moshi:

  • the Helium temporal transformer (:class:HeliumModel) on vLLM paged attention, consuming per-frame inputs_embeds and producing the per-frame hidden state;
  • the input embeddings (:class:PersonaPlexInputEmbeddings, embed_codes) that build those inputs_embeds from the delayed 17-row token stack;
  • the depformer (:class:PersonaPlexDepformer) that, conditioned on the temporal hidden state and the sampled text token, predicts the per-frame audio codes.

Per-frame protocol (OmniGPUModelRunner, gpu_model_runner.py):

  1. compute_logits produces the text logits; the engine samples the text token.
  2. preprocess (per-request, with that request's additional_information) carries Moshi's acoustic-delay cache and the precomputed user-audio code stream, builds the base inputs_embeds for the current frame, and exposes the previous frame's temporal hidden + text-step embedding via mtp_inputs.
  3. talker_mtp (batched, stateless) runs the depformer to predict the agent codes and finishes the next frame's inputs_embeds; the codes are stored under talker_mtp_output_key=("codes","audio") for the Mimi code2wav stage.

Phase 1 is turn-based: the user-audio rows (9..16) come from a precomputed Mimi encode of the input WAV (built in preprocess); live duplex is Phase 2.

PersonaPlexTalkerForConditionalGeneration

Bases: Module

vLLM-native PersonaPlex talker (temporal transformer + depformer).

Plain nn.Module (not SupportsPP): inheriting a vLLM Protocol puts Protocol in the MRO and breaks vLLM's runtime isinstance check for VllmModelForTextGeneration, which would misclassify the talker as a non-generate model. Qwen3-TTS's talker is plain nn.Module for the same reason; PersonaPlex does not need pipeline parallelism.

config instance-attribute

config = config

dep_q instance-attribute

dep_q = config.depformer_config.dep_q

depformer instance-attribute

depformer = PersonaPlexDepformer(
    config.depformer_config,
    temporal_hidden_size=hidden,
    text_card=config.text_vocab_size,
)

gpu_resident_buffer_keys instance-attribute

gpu_resident_buffer_keys: set[tuple[str, str]] = {
    ("hidden_states", "last")
}

has_postprocess instance-attribute

has_postprocess = True

has_preprocess instance-attribute

has_preprocess = True

have_multimodal_outputs instance-attribute

have_multimodal_outputs = True

input_embeddings instance-attribute

input_embeddings = PersonaPlexInputEmbeddings(config)

lm_head instance-attribute

lm_head = ParallelLMHead(
    config.text_vocab_size,
    hidden,
    quant_config=vllm_config.quant_config,
    prefix=maybe_prefix(prefix, "lm_head"),
)

logits_processor instance-attribute

logits_processor = LogitsProcessor(config.text_vocab_size)

make_empty_intermediate_tensors instance-attribute

make_empty_intermediate_tensors = (
    self.model.make_empty_intermediate_tensors
)

model instance-attribute

model = HeliumModel(
    vllm_config=vllm_config,
    prefix=maybe_prefix(prefix, "model"),
    config=self.temporal_config,
)

mtp_hidden_size instance-attribute

mtp_hidden_size = hidden

num_active_codebooks instance-attribute

num_active_codebooks = (
    config.depformer_config.num_active_codebooks
)

requires_full_prefix_cached_hidden_states instance-attribute

requires_full_prefix_cached_hidden_states = False

talker_mtp_output_key instance-attribute

talker_mtp_output_key = ('codes', 'audio')

temporal_config instance-attribute

temporal_config = config.temporal_config

vllm_config instance-attribute

vllm_config = vllm_config

compute_logits

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

embed_input_ids

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

Placeholder embedding for vLLM's VllmModel protocol.

The talker is driven by precomputed inputs_embeds (built in preprocess via embed_codes); input_ids are only in-vocab bookkeeping placeholders. Return a zero embedding of the temporal hidden size — the runner replaces it with the real per-frame inputs_embeds.

forward

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

load_weights

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

Route the Moshi checkpoint into the three components.

  • transformer.* / out_norm.alpha -> the Helium temporal backbone (same q/k-split + gate/up-split + alpha-squeeze map as HeliumForCausalLM).
  • text_linear.weight -> lm_head.
  • emb.* / text_emb.weight -> input embeddings.
  • depformer* / linears.* -> depformer.

make_omni_output

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

on_requests_finished

on_requests_finished(
    finished_req_ids: set[str] | list[str],
) -> None

post_sample_talker_mtp

post_sample_talker_mtp(
    *,
    input_ids: Tensor,
    hidden_states: Tensor,
    req_ids: list[str],
    req_infos: list[dict[str, Any]],
) -> Tensor

Generate depformer codes for a one-token resumable duplex segment.

The normal runner invokes talker_mtp at the start of the next decode step. PersonaPlex's unified duplex request instead appends one audio frame and stops after one sampled text token, so there is no next decode step. Run the same depformer dependency immediately from the current sampled text token and temporal hidden state.

postprocess

postprocess(
    hidden_states: Tensor, **_: Any
) -> dict[str, Any]

Capture this frame's last hidden for the next step's depformer (mtp).

preprocess

preprocess(
    input_ids: Tensor,
    input_embeds: Tensor | None,
    **info_dict: Any,
) -> tuple[Tensor, Tensor, dict[str, Any]]

Build the current frame's inputs_embeds and the mtp inputs.

Prefill emits the initial-token frame embedding; each decode step builds the delayed base frame (everything but agent cb0, which talker_mtp adds).

talker_mtp

talker_mtp(
    input_ids: Tensor,
    input_embeds: Tensor,
    last_talker_hidden: Tensor,
    text_step: Tensor,
    **kwargs: Any,
) -> tuple[Tensor, Tensor]

Run the depformer for the frame and finish its inputs_embeds.

Returns (inputs_embeds, audio_codes[B, dep_q]). audio_codes are stored under ("codes","audio"); inputs_embeds (base from preprocess + the fresh agent cb0 embedding, the delay-0 acoustic token) feeds the next temporal forward.