Skip to content

vllm.v1.attention.ops.ultraquant.triton_store

Triton store kernel for the ultraquant KV cache format.

Per-token, per-head, per-group-of-32 the kernel computes:

s_raw = c * absmax(group)            # c = 0.156 by default
s     = 2^round(log2(s_raw))         # UE8M0 snap (power of 2)
code  = quantize_to_fp4(group / s)   # OCP E2M1, 4 bits/elt

and stores code as packed nibbles + s as a single UE8M0 byte per group. The output layout matches format.py (272 B/slot for D=256).

K is Hadamard-rotated in-register; V is not. There is no per-token L2 norm. The UE8M0 snap makes s a power of two so scaled F8F6F4 MFMA can consume it on the read side.

Functions:

  • ultraquant_store –

    Launch the ultraquant store kernel. Writes K/V into kv_cache

_get_hadamard(dim, device, dtype=torch.float32)

Sylvester Hadamard in fp32 (cached, normalised).

Source code in vllm/v1/attention/ops/ultraquant/triton_store.py
def _get_hadamard(
    dim: int, device: torch.device, dtype: torch.dtype = torch.float32
) -> torch.Tensor:
    """Sylvester Hadamard in fp32 (cached, normalised)."""
    key = (dim, device, dtype)
    cached = _PIT_CACHE.get(key)
    if cached is not None:
        return cached
    if dim <= 0 or (dim & (dim - 1)) != 0:
        raise ValueError(f"ultraquant store requires power-of-two dim, got {dim}")
    H = torch.tensor([[1.0]], dtype=torch.float64)
    while H.shape[0] < dim:
        H = torch.cat([torch.cat([H, H], dim=1), torch.cat([H, -H], dim=1)], dim=0)
    H = (H / (dim**0.5)).to(device=device, dtype=dtype).contiguous()
    _PIT_CACHE[key] = H
    return H

_kv_cache_flat(kv_cache)

1-D uint8 view of the raw KV allocation (handles padded layouts).

Source code in vllm/v1/attention/ops/ultraquant/triton_store.py
def _kv_cache_flat(kv_cache: torch.Tensor) -> torch.Tensor:
    """1-D uint8 view of the raw KV allocation (handles padded layouts)."""
    if kv_cache.is_contiguous():
        return kv_cache.reshape(-1)
    n = kv_cache.untyped_storage().nbytes() // kv_cache.element_size()
    return torch.empty(0, dtype=kv_cache.dtype, device=kv_cache.device).set_(
        kv_cache.untyped_storage(), 0, (n,)
    )

ultraquant_store(key, value, kv_cache, slot_mapping, *, PiT=None, constant_c=None)

Launch the ultraquant store kernel. Writes K/V into kv_cache in-place at the slots specified by slot_mapping.

Source code in vllm/v1/attention/ops/ultraquant/triton_store.py
def ultraquant_store(
    key: torch.Tensor,  # [N, H, D] bf16 or fp16
    value: torch.Tensor,  # [N, H, D] bf16 or fp16
    kv_cache: torch.Tensor,  # [num_blocks, block_size, Hk, slot_size_aligned] uint8
    slot_mapping: torch.Tensor,  # [N] int64
    *,
    PiT: torch.Tensor | None = None,  # unused; kept for call-site compatibility
    constant_c: float | None = None,
) -> None:
    """Launch the ultraquant store kernel. Writes K/V into `kv_cache`
    in-place at the slots specified by `slot_mapping`."""
    if key.dtype not in _SUPPORTED_DTYPES:
        raise ValueError(
            f"ultraquant_store: key.dtype must be one of "
            f"{_SUPPORTED_DTYPES}, got {key.dtype}"
        )
    if value.dtype != key.dtype:
        raise ValueError(
            f"ultraquant_store: key/value dtype mismatch ({key.dtype} vs {value.dtype})"
        )
    if slot_mapping.dtype != torch.int64:
        slot_mapping = slot_mapping.to(torch.int64)

    N, H, D = key.shape
    NH = N * H
    block_size = kv_cache.shape[1]
    num_kv_heads = kv_cache.shape[2]
    padded_slot = kv_cache.shape[3]

    expected_slot = slot_size(D)
    if padded_slot < expected_slot:
        raise ValueError(
            f"ultraquant_store: kv_cache slot {padded_slot} < expected "
            f"{expected_slot} for head_dim={D}"
        )
    if num_kv_heads != H:
        raise ValueError(
            f"ultraquant_store: kv_cache num_kv_heads {num_kv_heads} != key heads {H}"
        )

    midpoints = _get_midpoints_tensor(key.device)
    sorted_to_bits = _get_sorted_to_bits_tensor(key.device)

    k_flat = key.reshape(NH, D).contiguous()
    v_flat = value.reshape(NH, D).contiguous()

    gs = get_group_size()
    BLOCK_D = triton.next_power_of_2(D)
    N_GROUPS_C = n_groups(D, gs)
    K_SCALES_OFFSET = k_scales_offset(D, gs)
    V_CODES_OFFSET = v_codes_offset(D, gs)
    V_SCALES_OFFSET = v_scales_offset(D, gs)

    stride_block = kv_cache.stride(0)
    stride_pos = kv_cache.stride(1)
    stride_head = kv_cache.stride(2)

    c = constant_c if constant_c is not None else get_constant_c()

    _ensure_const_tables()
    grid = (NH,)
    _ultraquant_store_kernel[grid](
        k_flat,
        v_flat,
        _kv_cache_flat(kv_cache),
        slot_mapping,
        midpoints,
        sorted_to_bits,
        stride_cache_block=stride_block,
        stride_cache_pos=stride_pos,
        stride_cache_head=stride_head,
        HEAD_DIM=D,
        H=H,
        BLOCK_SIZE=block_size,
        BLOCK_D=BLOCK_D,
        LOG2_D=int(D).bit_length() - 1,
        GROUP_SIZE_C=gs,
        N_GROUPS_C=N_GROUPS_C,
        K_SCALES_OFFSET=K_SCALES_OFFSET,
        V_CODES_OFFSET=V_CODES_OFFSET,
        V_SCALES_OFFSET=V_SCALES_OFFSET,
        FP4_C=c,
        UE8M0_BIAS_C=UE8M0_BIAS,
        CONST_TABLES=1,
        num_warps=4,
        num_stages=1,
    )