Skip to content

vllm_omni.model_executor.models.audex.tta

Audex text-to-audio (TTA) token space: RVQ phase masking and helpers.

TTA output is an interleaved 4-codebook RVQ stream over <audiocodec_N> tokens: generated position p (counted from <audiogen_start>) must come from codebook p % 4, whose codec-id range is [phase * 1024, (phase + 1) * 1024). The tokenizer defines 8192 <audiocodec_*> symbols; only codec ids 0..4095 (the 4 quantizers XCodec1 decodes) are generated — the upper half stays disallowed.

Unguided/unmasked sampling emits phase-invalid sequences that XCodec1 cannot decode, so :class:TTARVQPhaseMaskLogitsProcessor is required for validity, driven per request by SamplingParams.extra_args["tta_rvq"]:

{
    "phase_token_ids": [[...1024 ids...] x 4],
    "start_tid": <audiogen_start id>,
    "end_tid": <audiogen_end id>,
    "codec_cap": <max codec tokens or None>,
    "start_in_prompt": <bool>,
}

AUDEX_AUDIOCODEC_TOKEN_OFFSET module-attribute

AUDEX_AUDIOCODEC_TOKEN_OFFSET = 196613

AUDEX_AUDIOCODEC_VOCAB_SIZE module-attribute

AUDEX_AUDIOCODEC_VOCAB_SIZE = 8192

AUDEX_AUDIOGEN_END_TOKEN_ID module-attribute

AUDEX_AUDIOGEN_END_TOKEN_ID = 131074

AUDEX_AUDIOGEN_START_TOKEN_ID module-attribute

AUDEX_AUDIOGEN_START_TOKEN_ID = 131073

NEG_INF module-attribute

NEG_INF = float('-inf')

XCODEC1_CODEBOOK_SIZE module-attribute

XCODEC1_CODEBOOK_SIZE = 1024

XCODEC1_NUM_CODEBOOKS module-attribute

XCODEC1_NUM_CODEBOOKS = 4

logger module-attribute

logger = init_logger(__name__)

TTARVQPhaseMaskLogitsProcessor

Bases: LogitsProcessor

Mask logits to the RVQ-phase-valid token set for Audex TTA requests.

Per batch row

_req[idx] -> {"start_tid", "end_tid", "codec_cap", "start_in_prompt"} _output_tokens[idx] -> live reference to the request's output ids _start_pos[idx] -> cached index of once seen; -1 means the prompt already ended with it.

The disallow masks are vocab-sized boolean tensors built once on the first apply (using the real logits.shape[-1] so vocab padding is also forbidden).

device instance-attribute

device = device

apply

apply(logits: Tensor) -> Tensor

is_argmax_invariant

is_argmax_invariant() -> bool

update_state

update_state(batch_update: BatchUpdate | None) -> None

build_tta_phase_token_ids

build_tta_phase_token_ids(
    tokenizer: Any,
) -> tuple[list[list[int]], int, int]

Group tokenizer ids of <audiocodec_N> into the 4 RVQ phases.

Phase p collects the tokenizer ids for codec ids [p*1024, (p+1)*1024). Fails loudly on missing markers or an incomplete codec vocab.

Returns (phase_token_ids, audiogen_start_tid, audiogen_end_tid).

validate_rvq_phase

validate_rvq_phase(
    codec_ids: list[int],
    codebook_size: int = XCODEC1_CODEBOOK_SIZE,
    num_codebooks: int = XCODEC1_NUM_CODEBOOKS,
) -> dict[str, Any]

Check that codec_ids[i] falls in RVQ phase i % 4's codec range.