Skip to content

vllm_omni.model_executor.models.audio8_tts.sampling

Audio8 TTS sampling primitives, ported 1:1 from the reference checkpoint.

Filter (top-k/top-p on raw logits) -> divide surviving logits by temperature -> Gumbel-max draw. This is the reverse of the usual filter order and is audibly different at temperature != 1, which is why this cannot reuse models/common/nucleus_ras_sampling.py. Per-row scalars are accepted so mixed batches run in one shot without a GPU->CPU sync.

SAMPLING_EPS module-attribute

SAMPLING_EPS = 1e-05

filter_top_k_top_p

filter_top_k_top_p(
    logits: Tensor,
    top_k: int | Tensor,
    top_p: float | Tensor,
) -> Tensor

Mask candidates outside the top-k / top-p set of the raw logits.

top_k <= 0 and top_p >= 1 disable each filter. The highest-scoring candidate is always kept, so no row can become all -inf.

gumbel_argmax_sample

gumbel_argmax_sample(
    scores: Tensor, generator: Generator | None = None
) -> Tensor

Draw one index per row via Gumbel-max over softmax(scores).

ras_sample_batch

ras_sample_batch(
    logits: Tensor,
    recent_ids: Tensor | None,
    *,
    temperature: float | Tensor,
    top_k: int | Tensor,
    top_p: float | Tensor,
    ras_temperature: float,
    ras_top_p: float,
    num_semantic_ids: int | None,
    generators: dict,
    num_reqs: int,
) -> Tensor

Batched RAS that honors each request's own RNG generator.

vLLM keys generators by batch slot, so a single shared generator would leak slot 0's RNG stream to every row and make a prompt's audio depend on its batch neighbours (undercutting per-request seed). The single-request case keeps the batched fast path; larger batches loop so each row draws from its own (possibly unseeded) generator.

ras_sample_semantic

ras_sample_semantic(
    logits: Tensor,
    recent_ids: Tensor | None,
    *,
    temperature: float | Tensor,
    top_k: int | Tensor,
    top_p: float | Tensor,
    ras_temperature: float,
    ras_top_p: float,
    do_sample: bool = True,
    num_semantic_ids: int | None = None,
    generator: Generator | None = None,
) -> Tensor

Repetition-Aware Sampling for the semantic (Slow AR) token.

Draws with the request's temperature / top_p; if the drawn id is already in recent_ids (pad with a value that cannot be sampled, e.g. -1), redraws from the flatter RAS distribution instead. The num_semantic_ids slot and beyond (EOS) are never substituted, so a repeat EOS still ends the utterance.

DualAR RAS resamples repeats rather than masking them — masking the dominant token inflates EOS probability and truncates the utterance.

sample_scores

sample_scores(
    logits: Tensor,
    *,
    temperature: float | Tensor,
    top_k: int | Tensor,
    top_p: float | Tensor,
    do_sample: bool = True,
    generator: Generator | None = None,
) -> Tensor

Filter -> temper -> draw, in the reference order.

Rows whose temperature is below :data:SAMPLING_EPS fall back to argmax.