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-frameinputs_embedsand producing the per-frame hidden state; - the input embeddings (:class:
PersonaPlexInputEmbeddings,embed_codes) that build thoseinputs_embedsfrom 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):
compute_logitsproduces the text logits; the engine samples the text token.preprocess(per-request, with that request'sadditional_information) carries Moshi's acoustic-delay cache and the precomputed user-audio code stream, builds the baseinputs_embedsfor the current frame, and exposes the previous frame's temporal hidden + text-step embedding viamtp_inputs.talker_mtp(batched, stateless) runs the depformer to predict the agent codes and finishes the next frame'sinputs_embeds; the codes are stored undertalker_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.
depformer instance-attribute ¶
depformer = PersonaPlexDepformer(
config.depformer_config,
temporal_hidden_size=hidden,
text_card=config.text_vocab_size,
)
gpu_resident_buffer_keys instance-attribute ¶
lm_head instance-attribute ¶
lm_head = ParallelLMHead(
config.text_vocab_size,
hidden,
quant_config=vllm_config.quant_config,
prefix=maybe_prefix(prefix, "lm_head"),
)
make_empty_intermediate_tensors instance-attribute ¶
model instance-attribute ¶
model = HeliumModel(
vllm_config=vllm_config,
prefix=maybe_prefix(prefix, "model"),
config=self.temporal_config,
)
num_active_codebooks instance-attribute ¶
requires_full_prefix_cached_hidden_states instance-attribute ¶
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 ¶
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
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 ¶
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.