Skip to content

vllm.v1.attention.ops.ultraquant.triton_unified_attention ¶

Unified Triton fallback for the UltraQuant 4-bit KV-cache format.

K and V use packed FP4 E2M1 values with UE8M0 group-of-32 scales. QK uses native scaled FP4×E4M3 MFMA on CDNA4. Production choices are fixed in code; there are no environment-variable tuning switches.

Functions:

_get_pit(dim, device, dtype) ¶

Sylvester Hadamard projection / sqrt(dim). PiT = PiT.T (symmetric).

Source code in vllm/v1/attention/ops/ultraquant/triton_unified_attention.py
def _get_pit(dim: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
    """Sylvester Hadamard projection / sqrt(dim). PiT = PiT.T (symmetric)."""
    key = (dim, device, dtype)
    H = _HADAMARD_CACHE.get(key)
    if H is None:
        assert (dim & (dim - 1)) == 0, f"dim={dim} must be power of 2"
        H = torch.tensor([[1.0]], dtype=torch.float32)
        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()
        _HADAMARD_CACHE[key] = H
    return H

_ultraquant_fp4_decode_arith(codes) ¶

Arithmetic FP4 E2M1 decode, bit-exact to format.FP4_BITS_TO_VALUE.

Source code in vllm/v1/attention/ops/ultraquant/triton_unified_attention.py
@triton.jit
def _ultraquant_fp4_decode_arith(codes):
    """Arithmetic FP4 E2M1 decode, bit-exact to format.FP4_BITS_TO_VALUE."""
    mag = codes & 7
    e = mag >> 1
    m = mag & 1
    exp_field = 126 + e
    mant_field = tl.where(e != 0, m, 0) << 22
    bits = (exp_field << 23) | mant_field
    magval = bits.to(tl.float32, bitcast=True)
    magval = tl.where(mag == 0, 0.0, magval)
    signf = tl.where((codes & 8) != 0, -1.0, 1.0)
    return magval * signf

_ultraquant_load_k_packed(KV_cache_ptr, data_bases, k_scales_addrs, d_half_offs, half_mask, tile_mask, BLOCK_D, HEAD_DIM, GROUP_SIZE_C, N_GROUPS_C, TILE_SIZE, UNMASKED) ¶

Load packed FP4 K codes and UE8M0 scales for scaled QK MFMA.

Source code in vllm/v1/attention/ops/ultraquant/triton_unified_attention.py
@triton.jit
def _ultraquant_load_k_packed(
    KV_cache_ptr,
    data_bases,
    k_scales_addrs,
    d_half_offs,
    half_mask,
    tile_mask,
    BLOCK_D: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    GROUP_SIZE_C: tl.constexpr,
    N_GROUPS_C: tl.constexpr,
    TILE_SIZE: tl.constexpr,
    UNMASKED: tl.constexpr,
):
    """Load packed FP4 K codes and UE8M0 scales for scaled QK MFMA."""
    addrs = data_bases[:, None] + d_half_offs[None, :]
    if UNMASKED:
        K_codes = tl.load(KV_cache_ptr + addrs, mask=half_mask[None, :], other=0)
    else:
        K_codes = tl.load(
            KV_cache_ptr + addrs,
            mask=tile_mask[:, None] & half_mask[None, :],
            other=0,
        )
    K_T_packed = tl.trans(K_codes)
    grp = tl.arange(0, N_GROUPS_C)
    scale_addrs = k_scales_addrs[:, None] + grp[None, :]
    if UNMASKED:
        K_scales = tl.load(KV_cache_ptr + scale_addrs)
    else:
        K_scales = tl.load(KV_cache_ptr + scale_addrs, mask=tile_mask[:, None], other=0)
    _ = HEAD_DIM
    _ = GROUP_SIZE_C
    _ = BLOCK_D
    return K_T_packed, K_scales

_ultraquant_load_v_tile(KV_cache_ptr, val_bases, v_scales_addrs, d_offs, d_mask, tile_mask, OUT_DTYPE, HEAD_DIM, BLOCK_D, GROUP_SIZE_C, N_GROUPS_C, TILE_SIZE, UNMASKED, UE8M0_BIAS_C) ¶

Load and dequant a V tile with arithmetic FP4 decode.

Source code in vllm/v1/attention/ops/ultraquant/triton_unified_attention.py
@triton.jit
def _ultraquant_load_v_tile(
    KV_cache_ptr,
    val_bases,
    v_scales_addrs,
    d_offs,
    d_mask,
    tile_mask,
    OUT_DTYPE: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    BLOCK_D: tl.constexpr,
    GROUP_SIZE_C: tl.constexpr,
    N_GROUPS_C: tl.constexpr,
    TILE_SIZE: tl.constexpr,
    UNMASKED: tl.constexpr,
    UE8M0_BIAS_C: tl.constexpr,
):
    """Load and dequant a V tile with arithmetic FP4 decode."""
    half_idx = d_offs // 2
    nibble_shift = (d_offs % 2) * 4
    addrs = val_bases[:, None] + half_idx[None, :]
    if UNMASKED:
        byte_raw = tl.load(KV_cache_ptr + addrs, mask=d_mask[None, :], other=0).to(
            tl.int32
        )
    else:
        byte_raw = tl.load(
            KV_cache_ptr + addrs,
            mask=tile_mask[:, None] & d_mask[None, :],
            other=0,
        ).to(tl.int32)
    codes = (byte_raw >> nibble_shift[None, :]) & 0xF
    grp = tl.arange(0, N_GROUPS_C)
    scale_addrs = v_scales_addrs[:, None] + grp[None, :]
    if UNMASKED:
        scale_bytes = tl.load(KV_cache_ptr + scale_addrs).to(tl.int32)
    else:
        scale_bytes = tl.load(
            KV_cache_ptr + scale_addrs, mask=tile_mask[:, None], other=0
        ).to(tl.int32)
    scale_exp = scale_bytes - UE8M0_BIAS_C
    fp4_vals = _ultraquant_fp4_decode_arith(codes)
    scales = tl.where(scale_bytes == 0, 0.0, tl.exp2(tl.cast(scale_exp, tl.float32)))
    V_g = tl.reshape(fp4_vals, [TILE_SIZE, N_GROUPS_C, GROUP_SIZE_C])
    return tl.reshape(V_g * scales[:, :, None], [TILE_SIZE, BLOCK_D]).to(OUT_DTYPE)

ultraquant_unified_attention(query, kv_cache, block_table, seq_lens, query_start_loc, scale, PiT=None, output=None, tile_size=None, max_query_len=None, max_seq_len=None, num_kv_splits=None, force_2d=False, sinks=None, sliding_window=None) ¶

Launch unified UltraQuant attention with scaled F8F6F4 MFMA.

Source code in vllm/v1/attention/ops/ultraquant/triton_unified_attention.py
def ultraquant_unified_attention(
    query: torch.Tensor,  # [num_tokens, Hq, D] fp16/bf16 (raw)
    kv_cache: torch.Tensor,  # [num_blocks, block_size, Hk, padded_slot] uint8
    block_table: torch.Tensor,  # [num_seqs, max_num_blocks] int32
    seq_lens: torch.Tensor,  # [num_seqs] int32
    query_start_loc: torch.Tensor,  # [num_seqs+1] int32
    scale: float,
    PiT: torch.Tensor | None = None,
    output: torch.Tensor | None = None,
    tile_size: int | None = None,
    max_query_len: int | None = None,
    max_seq_len: int | None = None,
    num_kv_splits: int | None = None,
    force_2d: bool = False,
    sinks: torch.Tensor | None = None,
    sliding_window: int | None = None,
) -> torch.Tensor:
    """Launch unified UltraQuant attention with scaled F8F6F4 MFMA."""
    assert query.dim() == 3, f"query must be [N, Hq, D], got {query.shape}"
    num_tokens, Hq, D = query.shape
    Hk = kv_cache.shape[2]
    block_size = kv_cache.shape[1]
    padded_slot = kv_cache.shape[3]
    kv_group_size = Hq // Hk
    num_seqs = int(query_start_loc.shape[0] - 1)
    device = query.device

    gs = get_group_size()
    if padded_slot < slot_size(D, gs):
        raise ValueError(
            f"ultraquant_unified: cache slot {padded_slot} < expected "
            f"{slot_size(D, gs)} for D={D} group_size={gs}"
        )
    if Hq % Hk != 0:
        raise ValueError(f"Hq={Hq} must be a multiple of Hk={Hk}")

    if PiT is None:
        PiT = _get_pit(D, device, torch.float32)
    elif PiT.dtype != torch.float32 or not PiT.is_contiguous():
        PiT = PiT.to(torch.float32).contiguous()

    q_rot = (query.float() @ PiT).contiguous()
    q_for_kernel = q_rot.to(torch.float8_e4m3fn).contiguous()

    if sinks is not None:
        sinks_f32 = sinks if sinks.dtype == torch.float32 else sinks.to(torch.float32)
        if not sinks_f32.is_contiguous():
            sinks_f32 = sinks_f32.contiguous()
        assert sinks_f32.numel() == Hq, (
            f"sinks must have shape [Hq={Hq}], got numel={sinks_f32.numel()}"
        )
        use_sinks = True
    else:
        sinks_f32 = q_for_kernel  # harmless dummy; never dereferenced when USE_SINKS=0
        use_sinks = False

    if output is None:
        output = torch.empty_like(query)

    # BLOCK_M heuristic shared with the compressed-KV fallback kernels.
    if max_query_len is not None:
        is_prefill_like = max_query_len > 1
    else:
        is_prefill_like = num_tokens > num_seqs

    if is_prefill_like:
        BLOCK_M = max(128, triton.next_power_of_2(kv_group_size))
    else:
        BLOCK_M = 16 if kv_group_size <= 16 else triton.next_power_of_2(kv_group_size)
    BLOCK_Q = BLOCK_M // kv_group_size

    total_num_q_blocks = num_tokens // BLOCK_Q + num_seqs

    if tile_size is None:
        tile_size = 32 if is_prefill_like else 16

    num_stages_2d = 1 if _is_hip else 2
    num_stages_3d = 3 if _is_hip else 2

    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)

    scale_for_kernel = float(scale)

    kv_flat = _kv_cache_flat(kv_cache)

    # Dispatch: 2D for prefill / chunked; 3D for pure decode with long KV.
    if max_seq_len is None:
        max_seq_len_hint = int(block_table.shape[1]) * int(block_size)
    else:
        max_seq_len_hint = int(max_seq_len)
    use_3d = (not force_2d) and (not is_prefill_like) and max_seq_len_hint >= 1024

    if not use_3d:
        kernel_ultraquant_unified_attention_2d[(total_num_q_blocks, Hk)](
            output_ptr=output,
            query_ptr=q_for_kernel,
            KV_cache_ptr=kv_flat,
            block_tables_ptr=block_table,
            seq_lens_ptr=seq_lens,
            query_start_len_ptr=query_start_loc,
            sinks_ptr=sinks_f32,
            scale=scale_for_kernel,
            num_query_heads=Hq,
            num_queries_per_kv=kv_group_size,
            block_table_stride=block_table.stride(0),
            query_stride_0=q_for_kernel.stride(0),
            query_stride_1=q_for_kernel.stride(1),
            output_stride_0=output.stride(0),
            output_stride_1=output.stride(1),
            stride_cache_block=kv_cache.stride(0),
            stride_cache_pos=kv_cache.stride(1),
            stride_cache_head=kv_cache.stride(2),
            BLOCK_SIZE=block_size,
            TILE_SIZE=tile_size,
            HEAD_SIZE=D,
            HEAD_SIZE_PADDED=BLOCK_D,
            BLOCK_Q=BLOCK_Q,
            BLOCK_M=BLOCK_M,
            num_seqs=num_seqs,
            K_SCALES_OFFSET=K_SCALES_OFFSET,
            V_CODES_OFFSET=V_CODES_OFFSET,
            V_SCALES_OFFSET=V_SCALES_OFFSET,
            GROUP_SIZE_C=gs,
            N_GROUPS_C=N_GROUPS_C,
            UE8M0_BIAS_C=UE8M0_BIAS,
            USE_SINKS=1 if use_sinks else 0,
            SLIDING_WINDOW=int(sliding_window)
            if sliding_window and sliding_window > 0
            else 0,
            num_warps=4,
            num_stages=num_stages_2d,
        )
        return output

    # 3D split-KV path
    if num_kv_splits is None:
        num_kv_splits = 16
    if num_kv_splits < 1:
        num_kv_splits = 1
    if num_kv_splits & (num_kv_splits - 1) != 0:
        num_kv_splits = 1 << (num_kv_splits.bit_length() - 1)
    max_possible_splits = max(1, (max_seq_len_hint + tile_size - 1) // tile_size)
    num_segments = max(1, min(num_kv_splits, max_possible_splits))

    segm_output = torch.empty(
        (num_tokens, Hq, num_segments, BLOCK_D),
        dtype=torch.float32,
        device=device,
    )
    segm_max = torch.empty(
        (num_tokens, Hq, num_segments),
        dtype=torch.float32,
        device=device,
    )
    segm_expsum = torch.empty(
        (num_tokens, Hq, num_segments),
        dtype=torch.float32,
        device=device,
    )

    kernel_ultraquant_unified_attention_3d[(total_num_q_blocks, Hk, num_segments)](
        segm_output_ptr=segm_output,
        segm_max_ptr=segm_max,
        segm_expsum_ptr=segm_expsum,
        query_ptr=q_for_kernel,
        KV_cache_ptr=kv_flat,
        block_tables_ptr=block_table,
        seq_lens_ptr=seq_lens,
        query_start_len_ptr=query_start_loc,
        sinks_ptr=sinks_f32,
        scale=scale_for_kernel,
        num_query_heads=Hq,
        num_queries_per_kv=kv_group_size,
        block_table_stride=block_table.stride(0),
        query_stride_0=q_for_kernel.stride(0),
        query_stride_1=q_for_kernel.stride(1),
        stride_cache_block=kv_cache.stride(0),
        stride_cache_pos=kv_cache.stride(1),
        stride_cache_head=kv_cache.stride(2),
        BLOCK_SIZE=block_size,
        TILE_SIZE=tile_size,
        HEAD_SIZE=D,
        HEAD_SIZE_PADDED=BLOCK_D,
        BLOCK_Q=BLOCK_Q,
        BLOCK_M=BLOCK_M,
        num_seqs=num_seqs,
        NUM_SEGMENTS_PER_SEQ=num_segments,
        K_SCALES_OFFSET=K_SCALES_OFFSET,
        V_CODES_OFFSET=V_CODES_OFFSET,
        V_SCALES_OFFSET=V_SCALES_OFFSET,
        GROUP_SIZE_C=gs,
        N_GROUPS_C=N_GROUPS_C,
        UE8M0_BIAS_C=UE8M0_BIAS,
        USE_SINKS=1 if use_sinks else 0,
        SLIDING_WINDOW=int(sliding_window)
        if sliding_window and sliding_window > 0
        else 0,
        num_warps=2,
        num_stages=num_stages_3d,
    )

    # Reduce split-KV partials with the shared vectorized reducer.
    reduce_segments[(num_tokens, Hq)](
        output_ptr=output,
        segm_output_ptr=segm_output,
        segm_max_ptr=segm_max,
        segm_expsum_ptr=segm_expsum,
        seq_lens_ptr=seq_lens,
        num_seqs=num_seqs,
        num_query_heads=Hq,
        out_scale_inv=1.0,
        output_stride_0=output.stride(0),
        output_stride_1=output.stride(1),
        block_table_stride=block_table.stride(0),
        TILE_SIZE=tile_size,
        HEAD_SIZE=D,
        HEAD_SIZE_PADDED=BLOCK_D,
        query_start_len_ptr=query_start_loc,
        BLOCK_Q=BLOCK_Q,
        NUM_SEGMENTS_PER_SEQ=num_segments,
        USE_FP8=False,
    )

    return output