Skip to content

vllm_omni.model_executor.models.common.fused_sampling

Fused Gumbel selection after native nucleus probability reductions.

Keep ATen's top-k threshold, sort, softmax and cumulative sum, including boundary ties. Random numbers remain indexed by original token id. Only the final mask, Gumbel transform and argmax are fused; no vocabulary scatter is needed. This avoids changing which tokens survive a nucleus boundary through different reduction orders.

sample_top_k_top_p_gumbel

sample_top_k_top_p_gumbel(
    logits: Tensor,
    uniforms: Tensor,
    *,
    top_k: int,
    top_p: float,
) -> Tensor

Sample CUDA rows without scattering sorted logits back to the vocabulary.

Uniforms are caller-owned FP32 values with the same shape as logits; their generation and generator advancement are deliberately outside this kernel. Softmax/scan use the reference ATen operations. The Gumbel logarithms still use libdevice, so this is not a universal bitwise-equivalence claim.