Skip to content

vllm_omni.model_executor.models.audex.audex_xcodec

Audex TTA stage 1: XCodec1 RVQ codes → waveform.

Unlike the streaming causal speech decoder used for TTS, XCodec1 is a CNN codec that the official flow decodes over the full sequence, so this stage runs sync full-payload only: each request delivers its complete de-interleaved [frames, 4] codec-id payload once at finish and the decode happens in a single call.

The checkpoint (hf-audio/xcodec-hubert-general-balanced) is external to the Audex repo and loads through transformers remote code; the model skeleton is built from config at startup for vLLM's memory profiler and the weights arrive via the standard load_weights iteration over the snapshot's safetensors.

logger module-attribute

logger = init_logger(__name__)

AudexXCodec1

Bases: Module

Stage-1 model for Audex TTA (GenerationModelRunner).

codec instance-attribute

codec = AutoModel.from_config(
    config, trust_remote_code=True
)

enable_update_additional_information instance-attribute

enable_update_additional_information = True

has_postprocess instance-attribute

has_postprocess = False

has_preprocess instance-attribute

has_preprocess = False

have_multimodal_outputs instance-attribute

have_multimodal_outputs = True

input_modalities class-attribute instance-attribute

input_modalities = 'audio'

model_path instance-attribute

model_path = vllm_config.model_config.model

requires_raw_input_tokens instance-attribute

requires_raw_input_tokens = True

vllm_config instance-attribute

vllm_config = vllm_config

compute_logits

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

embed_input_ids

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

forward

forward(
    input_ids: Tensor | None = None,
    positions: Tensor | None = None,
    intermediate_tensors: Any = None,
    inputs_embeds: Tensor | None = None,
    runtime_additional_information: list[dict[str, Any]]
    | None = None,
    **kwargs: Any,
) -> OmniOutput

load_weights

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

make_omni_output

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