Skip to content

vllm.v1.attention.ops.flydsl_ultraquant_decode

FlyDSL launcher for the optimized UltraQuant 4-bit D=256 decode kernel.

The production kernel uses scaled FP4×E4M3 QK MFMA, native V conversion, query hoisting, in-kernel Walsh-Hadamard rotation for GQA-6/8/16, and strided tile-group scheduling. These choices are fixed and are not environment-tunable.

Functions:

_SegmBufPool

Single-bucket buffer pool for segm_out/segm_max/segm_sum/output + the pooled Q-rotation intermediates. Identical to the TQ pool except that ultraquant needs an extra q_fp8 (float8_e4m3fn) slot for the Q precision haircut (q_rot_fp32 -> fp8_e4m3 -> bf16) so that no fresh allocation lands in the HIP graph memory pool post-capture.

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
class _SegmBufPool:
    """Single-bucket buffer pool for segm_out/segm_max/segm_sum/output + the
    pooled Q-rotation intermediates. Identical to the TQ pool except that
    ultraquant needs an extra ``q_fp8`` (float8_e4m3fn) slot for the Q precision
    haircut (``q_rot_fp32 -> fp8_e4m3 -> bf16``) so that no fresh allocation
    lands in the HIP graph memory pool post-capture.
    """

    __slots__ = ("_bufs", "_max_B")

    def __init__(self) -> None:
        self._bufs: dict[tuple, dict[str, torch.Tensor]] = {}
        self._max_B: int | None = None

    def get(
        self,
        B: int,
        Hk: int,
        Hq: int,
        num_partitions: int,
        QG: int,
        D: int,
        device: torch.device,
        q_dtype: torch.dtype,
    ) -> dict[str, torch.Tensor]:
        if self._max_B is None:
            self._max_B = _detect_max_capture_B()
        B_bucket = max(self._max_B, int(B))
        key = (
            int(Hk),
            int(Hq),
            int(num_partitions),
            int(QG),
            int(D),
            str(device),
            q_dtype,
        )
        bufs = self._bufs.get(key)
        if bufs is None or bufs["segm_out"].shape[0] < B_bucket:
            if bufs is not None and bufs["segm_out"].shape[0] < B_bucket:
                # First-time grow: warn so user knows cudagraphs may be
                # invalidated. This should only happen during warmup or in
                # always-eager configs (a batch larger than any captured
                # cudagraph size).
                logger.warning_once(
                    "FlyDSL UltraQuant buffer pool growing from %d "
                    "to %d (B=%d). If this happens AFTER cudagraph warmup the "
                    "captured graphs hold stale pointers and will GPU-fault.",
                    bufs["segm_out"].shape[0],
                    B_bucket,
                    B,
                )
            bufs = {
                "segm_out": torch.empty(
                    (B_bucket, Hk, num_partitions, QG, D),
                    dtype=torch.bfloat16,
                    device=device,
                ),
                "segm_max": torch.empty(
                    (B_bucket, Hk, num_partitions, QG),
                    dtype=torch.float32,
                    device=device,
                ),
                "segm_sum": torch.empty(
                    (B_bucket, Hk, num_partitions, QG),
                    dtype=torch.float32,
                    device=device,
                ),
                "output": torch.empty(
                    (B_bucket, Hq, D),
                    dtype=q_dtype,
                    device=device,
                ),
                "q_rot": torch.empty(
                    (B_bucket, Hq, D),
                    dtype=q_dtype,
                    device=device,
                ),
                "q_float": torch.empty(
                    (B_bucket, Hq, D),
                    dtype=torch.float32,
                    device=device,
                ),
                "q_rot_fp32": torch.empty(
                    (B_bucket, Hq, D),
                    dtype=torch.float32,
                    device=device,
                ),
                # ultraquant-specific: E4M3 haircut intermediate (pooled so the
                # cast does not allocate a fresh tensor post-capture).
                "q_fp8": torch.empty(
                    (B_bucket, Hq, D),
                    dtype=torch.float8_e4m3fn,
                    device=device,
                ),
            }
            self._bufs[key] = bufs
            self._max_B = B_bucket
            logger.info_once(
                "FlyDSL UltraQuant buffer pool allocated "
                "shape=(Hk=%d, Hq=%d, P=%d, QG=%d, D=%d, dtype=%s) "
                "B_bucket=%d. VRAM = %.1f MiB / shape.",
                Hk,
                Hq,
                num_partitions,
                QG,
                D,
                q_dtype,
                B_bucket,
                sum(t.numel() * t.element_size() for t in bufs.values()) / (1 << 20),
            )
        return {
            "segm_out": bufs["segm_out"][:B],
            "segm_max": bufs["segm_max"][:B],
            "segm_sum": bufs["segm_sum"][:B],
            "output": bufs["output"][:B],
            "q_rot": bufs["q_rot"][:B],
            "q_float": bufs["q_float"][:B],
            "q_rot_fp32": bufs["q_rot_fp32"][:B],
            "q_fp8": bufs["q_fp8"][:B],
        }

_fp4_value_lut(device)

Return the fixed 16-entry FP4 E2M1 value table as fp32 [16] on device.

Indexed by the stored 4-bit E2M1 code, so lut[nibble] is the dequant value directly (matches format.FP4_BITS_TO_VALUE).

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
def _fp4_value_lut(device: torch.device) -> torch.Tensor:
    """Return the fixed 16-entry FP4 E2M1 value table as fp32 [16] on device.

    Indexed by the stored 4-bit E2M1 code, so ``lut[nibble]`` is the dequant
    value directly (matches ``format.FP4_BITS_TO_VALUE``).
    """
    from vllm.v1.attention.ops.ultraquant.format import FP4_BITS_TO_VALUE

    key = (str(device), torch.float32)
    t = _FP4_LUT_CACHE.get(key)
    if t is None:
        t = torch.tensor(
            FP4_BITS_TO_VALUE, dtype=torch.float32, device=device
        ).contiguous()
        _FP4_LUT_CACHE[key] = t
    return t

_run_fused_q_rot(query, PiT_used, out, B, Hq, D, haircut)

Launch the fused prologue. out is the pooled [B_bucket, Hq, D] buf.

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
def _run_fused_q_rot(query, PiT_used, out, B, Hq, D, haircut):
    """Launch the fused prologue. ``out`` is the pooled [B_bucket, Hq, D] buf."""
    M = B * Hq
    BLOCK_M = 16
    BLOCK_K = 32
    BLOCK_D = 64
    grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(D, BLOCK_D))
    _fused_q_rot_haircut[grid](
        query,
        PiT_used,
        out,
        M,
        query.stride(0),
        query.stride(1),
        HQ=int(Hq),
        D=int(D),
        HAIRCUT=bool(haircut),
        BLOCK_M=BLOCK_M,
        BLOCK_K=BLOCK_K,
        BLOCK_D=BLOCK_D,
    )

flydsl_ultraquant_decode_attention(query, kv_cache, block_table, seq_lens, scale, PiT=None, max_seq_len=0, output_buf=None, buf_holder=None, max_num_kv_splits=32, sinks=None)

Launch the optimized FlyDSL UltraQuant D=256 decode.

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
def flydsl_ultraquant_decode_attention(
    query: torch.Tensor,
    kv_cache: torch.Tensor,
    block_table: torch.Tensor,
    seq_lens: torch.Tensor,
    scale: float,
    PiT: torch.Tensor | None = None,
    max_seq_len: int = 0,
    output_buf: torch.Tensor | None = None,
    buf_holder: Any = None,
    max_num_kv_splits: int = 32,
    sinks: torch.Tensor | None = None,
) -> torch.Tensor:
    """Launch the optimized FlyDSL UltraQuant D=256 decode."""
    if not is_ultraquant_flydsl_available():
        raise RuntimeError(
            "FlyDSL UltraQuant decode requires gfx950 and importable FlyDSL."
        )
    if sinks is not None:
        raise NotImplementedError("FlyDSL UltraQuant decode does not support sinks")

    B, Hq, D = query.shape
    Hk = kv_cache.shape[2]
    block_size = kv_cache.shape[1]
    padded_slot = int(kv_cache.shape[3])
    QG = Hq // Hk
    assert D == 256, f"UltraQuant FlyDSL requires head_dim=256, got {D}"
    assert block_size in (16, 32, 64, 128, 256), (
        f"UltraQuant FlyDSL supports kv_block_size 16/32/64/128/256, got {block_size}"
    )
    assert QG in _ULTRAQUANT_FLYDSL_GQA, (
        f"UltraQuant FlyDSL supports GQA factor "
        f"{', '.join(map(str, _ULTRAQUANT_FLYDSL_GQA))}, got {QG}"
    )

    device = query.device

    # ---- PiT (Hadamard) — cache the contiguous fp32 form on the layer -----
    PiT_f32: torch.Tensor
    if buf_holder is not None:
        PiT_f32 = getattr(buf_holder, "_ultraquant_PiT_f32", None)
        if PiT_f32 is None:
            assert PiT is not None, (
                "UltraQuant launcher requires PiT (the Hadamard rotation)"
            )
            PiT_f32 = PiT if PiT.dtype == torch.float32 else PiT.to(torch.float32)
            if not PiT_f32.is_contiguous():
                PiT_f32 = PiT_f32.contiguous()
            buf_holder._ultraquant_PiT_f32 = PiT_f32
    else:
        assert PiT is not None, (
            "UltraQuant launcher requires PiT (the Hadamard rotation)"
        )
        PiT_f32 = PiT if PiT.dtype == torch.float32 else PiT.to(torch.float32)
        if not PiT_f32.is_contiguous():
            PiT_f32 = PiT_f32.contiguous()

    centroids_c = _fp4_value_lut(device)

    # ---- Partition count (FA2 split-KV) -----------------------------------
    kv_compute_block = _ULTRAQUANT_MOD.KV_COMPUTE_BLOCK
    worst_case_max_seq_len = int(block_table.shape[1]) * int(block_size)
    if max_seq_len <= 0:
        max_seq_len = worst_case_max_seq_len
    sizing_max_seq_len = worst_case_max_seq_len

    # Split-KV parallelism cap. The old default (32) silently serialized long
    # contexts: once required_num_partitions exceeds the cap, TGPP grows and each
    # workgroup walks MORE tile-groups instead of the grid getting wider. The
    # decode kernel is latency-bound at ~1-2 waves/CU (grid = B*Hk*P workgroups
    # of a single wavefront), so that serialization dominated long-context decode.
    # Raising the cap to 256 keeps TGPP ~1 and measured (B=16, rocprof, per step
    # decode+reduce+store): 16k 82.7->40.0us, 32k 152.0->55.1us, 64k 339.4->95.7us,
    # 131k 671.3->179.0us (2.1x - 3.8x, growing with context). 512 regresses
    # because _reduce_partitions cost grows with partition count.
    # Cost: the segm pool scales with P (~8x vs P=32).
    # BATCH-ADAPTIVE split-KV cap (low-concurrency occupancy fix).
    # rocprof (B=4, 75k): the decode kernel is occupancy/latency-bound, not
    # compute-bound -- OccupancyPercent 5.8%, MfmaUtil 4.5%, VALUBusy 14.5%
    # (stalled ~85% on HBM), VALUUtilization 90% (no divergence). The grid is
    # B*Hk*num_partitions single-wavefront workgroups; at B=4/Hk=1/P=256 that's
    # only 1024 wg = 4/CU, too few to hide memory latency. Splitting MORE at low
    # batch adds resident wavefronts and cuts ITL; at high batch the grid is
    # already full so extra partitions only inflate _reduce_partitions + empty-
    # partition overhead. Measured end-to-end (graphs on, TP=8, uniform 75k,
    # median ITL vs kv8): cap512 vs cap256 -> C4 +9.0%->+4.9%, C8 +6.6%->+2.5%
    # (better), C16 +2.8%->+4.3% (worse). So use 512 at low batch, 256 at high.
    # This is only SAFE because the build-local-allocator fix in
    # ultraquant_decode_hd256.py lets the two num_partitions variants (256 & 512)
    # coexist across CUDA-graph batch buckets without the @flyc.jit
    # "global 'allocator' changed since first compile" drift crash.
    MAX_PARTITIONS = 512 if B <= 8 else 256
    required_num_partitions = (
        sizing_max_seq_len + kv_compute_block - 1
    ) // kv_compute_block
    parallelism_floor = min(MAX_PARTITIONS, max(1, max_num_kv_splits))
    num_partitions_actual = max(
        parallelism_floor,
        min(MAX_PARTITIONS, required_num_partitions),
    )
    # Cap partitions by how much work there actually IS to split. Nothing above
    # forces a partition to hold a useful number of tokens -- max_num_kv_splits
    # sets a parallelism FLOOR -- so a short context gets split into partitions
    # of a handful of tokens each, where every partition is almost entirely
    # fixed cost (and each one still writes a full QG x D partial for the
    # reduce to read back). Measured sweet spot is ~130-147 tokens/partition,
    # roughly flat in between; next_power_of_2 below rounds this target up, so
    # the effective range lands at ~96-192 tokens. At B=4 this is worth 27.9 ->
    # 24.4 us at seq=2048 and 27.0 -> 24.8 us at seq=16384, and is inert once
    # the context is long enough to want every partition (seq >= ~64k).
    _tok_per_part = 192
    if _tok_per_part > 0:
        _useful_parts = max(
            1, (int(sizing_max_seq_len) + _tok_per_part - 1) // _tok_per_part
        )
        num_partitions_actual = min(num_partitions_actual, _useful_parts)
    num_partitions = max(2, triton.next_power_of_2(num_partitions_actual))
    _tgpp_required = max(
        1,
        (required_num_partitions + num_partitions - 1) // num_partitions,
    )
    tile_groups_per_partition = int(triton.next_power_of_2(_tgpp_required))

    # ---- Pooled buffers ---------------------------------------------------
    pool_bufs = _SEGM_POOL.get(
        B,
        Hk,
        Hq,
        num_partitions,
        QG,
        D,
        device,
        query.dtype,
    )
    segm_out = pool_bufs["segm_out"]
    segm_max = pool_bufs["segm_max"]
    segm_sum = pool_bufs["segm_sum"]
    if output_buf is None:
        output = pool_bufs["output"]
    else:
        output = output_buf[:B] if output_buf.shape[0] != B else output_buf

    # Fold the scaled-MFMA head-dim permutation into PiT's columns so the
    # WHT / prologue produces native operand order. Always apply the E4M3
    # haircut: scaled QK consumes Q as E4M3.
    PiT_used = PiT_f32
    qperm_idx = _qperm_index(device, D=D)
    if buf_holder is not None:
        PiT_used = getattr(buf_holder, "_ultraquant_PiT_perm_f32", None)
        if PiT_used is None:
            PiT_used = PiT_f32.index_select(1, qperm_idx).contiguous()
            buf_holder._ultraquant_PiT_perm_f32 = PiT_used
    else:
        PiT_used = PiT_f32.index_select(1, qperm_idx).contiguous()
    _q_float = pool_bufs["q_float"]
    _q_rot_f32 = pool_bufs["q_rot_fp32"]
    _q_fp8 = pool_bufs["q_fp8"]
    _q_rot_out = pool_bufs["q_rot"]

    _inkernel_qrot = (
        _FUSE_QROT_INKERNEL
        and int(D) == 256
        and QG in _ULTRAQUANT_FLYDSL_GQA
        and query.dim() == 3
        and query.stride(2) == 1
        and query.stride(1) == int(D)
    )
    if _inkernel_qrot:
        q_for_kernel = query
    # The fused kernel indexes the head-dim contiguously; fall back if not.
    elif _FUSE_Q_ROT and query.stride(-1) == 1:
        # One kernel for the whole prologue instead of four ops.
        _run_fused_q_rot(query, PiT_used, _q_rot_out, B, Hq, D, haircut=True)
    else:
        _q_float.copy_(query)
        torch.mm(
            _q_float.view(B * Hq, D),
            PiT_used,
            out=_q_rot_f32.view(B * Hq, D),
        )
        _q_fp8.copy_(_q_rot_f32)  # fp32 -> e4m3
        _q_rot_out.copy_(_q_fp8)  # e4m3 -> bf16
    if not _inkernel_qrot:
        q_for_kernel = _q_rot_out

    # ---- FlyDSL kernel launch --------------------------------------------
    max_bps = int(block_table.shape[1])
    launch = _get_kernel(
        Hk,
        num_partitions,
        max_bps,
        scale,
        QG,
        block_size,
        padded_slot,
        fuse_qrot=_inkernel_qrot,
        num_seqs_hint=int(B),
        tile_groups_per_partition=int(tile_groups_per_partition),
        stride_q_seq=int(q_for_kernel.stride(0)),
        stride_q_head=int(q_for_kernel.stride(1)),
    )
    global _LOG_INVOKED_ONCE
    if not _LOG_INVOKED_ONCE:
        _LOG_INVOKED_ONCE = True
        logger.info(
            "FlyDSL UltraQuant launcher invoked: B=%d Hk=%d Hq=%d D=%d QG=%d "
            "num_partitions=%d (actual=%d, cap=%d) TGPP=%d max_bps=%d "
            "block_size=%d padded_slot=%d max_seq_len=%d "
            "(coverage=%d tokens, worst_case=%d tokens)",
            B,
            Hk,
            Hq,
            D,
            QG,
            num_partitions,
            num_partitions_actual,
            MAX_PARTITIONS,
            tile_groups_per_partition,
            max_bps,
            int(block_size),
            padded_slot,
            int(max_seq_len),
            num_partitions * tile_groups_per_partition * kv_compute_block,
            worst_case_max_seq_len,
        )
    launch(
        segm_out,
        segm_sum,
        segm_max,
        q_for_kernel,
        kv_cache,
        centroids_c,
        block_table,
        seq_lens,
        B,
        Hk,
        num_partitions,
        torch.cuda.current_stream(),
    )

    # ---- Reduce partitions -> [B, Hq, D] ---------------------------------
    # 64 splits the head-dim 4 ways, taking the grid from 32 to 128 workgroups.
    # Bit-identical by construction (the partition reduction is per-column).
    _red_block_d = min(int(D), 64)
    _reduce_partitions[(B, Hq, triton.cdiv(int(D), _red_block_d))](
        output_ptr=output,
        segm_out_ptr=segm_out,
        segm_max_ptr=segm_max,
        segm_sum_ptr=segm_sum,
        out_stride_n=output.stride(0),
        out_stride_h=output.stride(1),
        NUM_KV_HEADS=Hk,
        QG=QG,
        NUM_PARTS=num_partitions,
        HEAD_SIZE=D,
        BLOCK_D=_red_block_d,
    )
    return output

is_ultraquant_flydsl_available()

Return whether the gfx950 UltraQuant FlyDSL kernel can be loaded.

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
def is_ultraquant_flydsl_available() -> bool:
    """Return whether the gfx950 UltraQuant FlyDSL kernel can be loaded."""
    global _FLYDSL_AVAILABLE, _ULTRAQUANT_MOD
    global _FLYC, _FX, _TYPING_T, _CC, _IR
    if _FLYDSL_AVAILABLE is not None:
        return _FLYDSL_AVAILABLE
    try:
        from vllm.platforms.rocm import on_gfx950

        if not on_gfx950():
            _FLYDSL_AVAILABLE = False
            return False

        import flydsl.compiler as flyc
        import flydsl.expr as fx
        from flydsl._mlir import ir
        from flydsl.compiler.kernel_function import CompilationContext
        from flydsl.expr.typing import T

        from vllm.v1.attention.ops.flydsl_kernels import (
            ultraquant_decode_hd256 as ultraquant_mod,
        )

        _FLYC = flyc
        _FX = fx
        _TYPING_T = T
        _CC = CompilationContext
        _IR = ir
        _ULTRAQUANT_MOD = ultraquant_mod
        _FLYDSL_AVAILABLE = True
        logger.info_once("FlyDSL UltraQuant D=256 decode is available")
    except Exception as ex:  # noqa: BLE001
        _FLYDSL_AVAILABLE = False
        logger.warning_once(
            "FlyDSL UltraQuant decode is unavailable (%s); using Triton fallback.",
            ex,
        )
    return _FLYDSL_AVAILABLE

ultraquant_flydsl_decode_eligible(*, head_size, num_kv_groups, has_sinks, sliding_window, flydsl_loaded)

Return whether decode should use the gfx950 FlyDSL kernel.

Production FlyDSL covers D=256 and GQA in {6, 8, 16}. Sinks and sliding window fall back to unified Triton. flydsl_loaded is the process-level compiler/hardware probe so this policy can be unit-tested without a GPU.

Source code in vllm/v1/attention/ops/flydsl_ultraquant_decode.py
def ultraquant_flydsl_decode_eligible(
    *,
    head_size: int,
    num_kv_groups: int,
    has_sinks: bool,
    sliding_window: int | None,
    flydsl_loaded: bool,
) -> bool:
    """Return whether decode should use the gfx950 FlyDSL kernel.

    Production FlyDSL covers D=256 and GQA in {6, 8, 16}. Sinks and sliding
    window fall back to unified Triton. ``flydsl_loaded`` is the process-level
    compiler/hardware probe so this policy can be unit-tested without a GPU.
    """
    return (
        flydsl_loaded
        and head_size == 256
        and num_kv_groups in _ULTRAQUANT_FLYDSL_GQA
        and not has_sinks
        and not (sliding_window and sliding_window > 0)
    )