Skip to content

vllm.v1.attention.ops.ultraquant.triton_dequant

Full KV dequant for the UltraQuant cache format.

Used by continuation prefill: when a chunk brings many new query tokens on top of a long cached prefix, dequanting the prefix once and running a dense prefill kernel beats replaying the decode kernel per query token.

Reads FP4 codes plus UE8M0 group scales and writes K (Hadamard-rotated, as stored) and V (raw) into pre-allocated fp16/bf16 buffers.

Functions:

_get_fp4_decode_table(device, dtype)

16-entry FP4 E2M1 bit-pattern -> value table (cached per device).

Source code in vllm/v1/attention/ops/ultraquant/triton_dequant.py
def _get_fp4_decode_table(device: torch.device, dtype: torch.dtype) -> torch.Tensor:
    """16-entry FP4 E2M1 bit-pattern -> value table (cached per device)."""
    key = (device, dtype)
    t = _FP4_DECODE_CACHE.get(key)
    if t is None:
        t = torch.tensor(FP4_BITS_TO_VALUE, device=device, dtype=dtype).contiguous()
        _FP4_DECODE_CACHE[key] = t
    return t

ultraquant_full_dequant_kv(kv_cache, block_table, k_out, v_out, alloc_len)

Dequant alloc_len cached positions into k_out / v_out.

k_out / v_out are [B, Hk, alloc_len, D] in fp16 or bf16.

Source code in vllm/v1/attention/ops/ultraquant/triton_dequant.py
def ultraquant_full_dequant_kv(
    kv_cache: torch.Tensor,
    block_table: torch.Tensor,
    k_out: torch.Tensor,
    v_out: torch.Tensor,
    alloc_len: int,
) -> None:
    """Dequant ``alloc_len`` cached positions into ``k_out`` / ``v_out``.

    ``k_out`` / ``v_out`` are ``[B, Hk, alloc_len, D]`` in fp16 or bf16.
    """
    B = block_table.shape[0]
    Hk = kv_cache.shape[2]
    D = k_out.shape[3]
    block_size = kv_cache.shape[1]

    gs = get_group_size()
    _ultraquant_full_dequant_kv[(alloc_len, B * Hk)](
        _kv_cache_flat(kv_cache),
        block_table,
        _get_fp4_decode_table(kv_cache.device, k_out.dtype),
        k_out,
        v_out,
        k_out.stride(0),
        k_out.stride(1),
        k_out.stride(2),
        v_out.stride(0),
        v_out.stride(1),
        v_out.stride(2),
        kv_cache.stride(0),
        kv_cache.stride(1),
        kv_cache.stride(2),
        block_table.stride(0),
        HEAD_DIM=D,
        BLOCK_SIZE=block_size,
        NUM_KV_HEADS=Hk,
        K_SCALES_OFFSET=k_scales_offset(D, gs),
        V_CODES_OFFSET=v_codes_offset(D, gs),
        V_SCALES_OFFSET=v_scales_offset(D, gs),
        GROUP_SIZE_C=gs,
        N_GROUPS_C=n_groups(D, gs),
        BLOCK_D=triton.next_power_of_2(D),
        OUT_BF16=1 if k_out.dtype == torch.bfloat16 else 0,
        UE8M0_BIAS_C=UE8M0_BIAS,
        num_warps=4,
    )