Skip to content

vllm.models.qwen4_exp.amd.ops.qsa_pre_indexer

Fused QSA pre-indexer kernel for the AMD Qwen4Exp path.

Started as a copy of the NVIDIA kernel (nvidia/ops/qsa_prepare.py) and is kept separate so either vendor can change its own copy without re-validating the other. The state source select is two masked loads rather than a pointer select, which ROCm Triton rejects.

Functions:

_norm_rope(x, pos_t, pos_h, pos_w, cos_sin_ptr, cos_sin_stride, norm_weight_ptr, eps, IS_MROPE, MROPE_H, MROPE_W)

Apply Gemma RMSNorm and selected-axis NeoX RoPE to register rows.

Source code in vllm/models/qwen4_exp/amd/ops/qsa_pre_indexer.py
@triton.jit
def _norm_rope(
    x,
    pos_t,
    pos_h,
    pos_w,
    cos_sin_ptr,
    cos_sin_stride,
    norm_weight_ptr,
    eps,
    IS_MROPE: tl.constexpr,
    MROPE_H: tl.constexpr,
    MROPE_W: tl.constexpr,
):
    """Apply Gemma RMSNorm and selected-axis NeoX RoPE to register rows."""
    TILE_T: tl.constexpr = x.shape[0]
    TILE_H: tl.constexpr = x.shape[1]
    D: tl.constexpr = x.shape[2]
    ROWS: tl.constexpr = TILE_T * TILE_H
    HALF: tl.constexpr = D // 2
    QUARTER: tl.constexpr = D // 4
    pairs = tl.arange(0, QUARTER)
    if IS_MROPE:
        # Qwen interleaves temporal, height, and width rotary pairs. Each axis
        # still indexes the same position-major cos/sin table.
        h_mask = ((pairs % 3) == 1) & (pairs <= 3 * MROPE_H)
        w_mask = ((pairs % 3) == 2) & (pairs <= 3 * MROPE_W)
        t_mask = ~(h_mask | w_mask)
        base = cos_sin_ptr + pairs[None, :]
        pos_rows = (pos_t, pos_h, pos_w)
        axis_masks = (t_mask, h_mask, w_mask)
        cos = tl.zeros((TILE_T, QUARTER), dtype=cos_sin_ptr.dtype.element_ty)
        sin = tl.zeros((TILE_T, QUARTER), dtype=cos_sin_ptr.dtype.element_ty)
        for axis in tl.static_range(3):
            cos += tl.load(
                base + pos_rows[axis][:, None] * cos_sin_stride,
                mask=axis_masks[axis][None, :],
                other=0,
            )
            sin += tl.load(
                base + pos_rows[axis][:, None] * cos_sin_stride + QUARTER,
                mask=axis_masks[axis][None, :],
                other=0,
            )
    else:
        cos = tl.load(cos_sin_ptr + pos_t[:, None] * cos_sin_stride + pairs[None, :])
        sin = tl.load(
            cos_sin_ptr + pos_t[:, None] * cos_sin_stride + QUARTER + pairs[None, :]
        )

    cos = tl.reshape(
        tl.broadcast_to(cos[:, None, :], (TILE_T, TILE_H, QUARTER)),
        (ROWS, QUARTER),
    )
    sin = tl.reshape(
        tl.broadcast_to(sin[:, None, :], (TILE_T, TILE_H, QUARTER)),
        (ROWS, QUARTER),
    )
    x = tl.reshape(x, (ROWS, D)).to(tl.float32)
    weight = tl.load(norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + 1.0
    rrms = tl.rsqrt(tl.sum(x * x, axis=1) / D + eps)
    y = (x * rrms[:, None] * weight[None, :]).to(cos.dtype)
    rotated, passthrough = tl.split(
        tl.permute(tl.reshape(y, (ROWS, 2, HALF)), (0, 2, 1))
    )
    r0, r1 = tl.split(tl.permute(tl.reshape(rotated, (ROWS, 2, QUARTER)), (0, 2, 1)))
    out0 = r0 * cos - r1 * sin
    out1 = r1 * cos + r0 * sin
    rotated = tl.reshape(tl.permute(tl.join(out0, out1), (0, 2, 1)), (ROWS, HALF))
    result = tl.reshape(tl.permute(tl.join(rotated, passthrough), (0, 2, 1)), (ROWS, D))
    return tl.reshape(result, (TILE_T, TILE_H, D))

qsa_pre_indexer(q, k, positions, cos_sin_cache, q_norm_weight, k_norm_weight, eps, q_out, state_cache, state_slots, state_block_table, query_start_loc, logical_positions, compressed_cache, compressed_slots, k_work_metadata, *, compress_ratio, mrope_section, rope_pos_offset)

Normalize Q, compress K, then update the circular raw state.

Source code in vllm/models/qwen4_exp/amd/ops/qsa_pre_indexer.py
def qsa_pre_indexer(
    q: torch.Tensor,
    k: torch.Tensor,
    positions: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    q_norm_weight: torch.Tensor,
    k_norm_weight: torch.Tensor,
    eps: float,
    q_out: torch.Tensor,
    state_cache: torch.Tensor,
    state_slots: torch.Tensor,
    state_block_table: torch.Tensor,
    query_start_loc: torch.Tensor,
    logical_positions: torch.Tensor,
    compressed_cache: torch.Tensor,
    compressed_slots: torch.Tensor,
    k_work_metadata: torch.Tensor,
    *,
    compress_ratio: int,
    mrope_section: tuple[int, int, int] | None,
    rope_pos_offset: int | None,
) -> None:
    """Normalize Q, compress K, then update the circular raw state."""
    num_tokens = q.shape[0]
    if num_tokens == 0:
        return
    num_q_heads, head_dim = q_out.shape[1:]
    assert cos_sin_cache.shape[-1] * 2 == head_dim
    assert q.shape == (num_tokens, num_q_heads * head_dim)
    assert k.shape == (num_tokens, head_dim)
    assert q.stride(-1) == 1
    assert k.stride(-1) == 1
    assert q_out.stride(-1) == 1
    assert cos_sin_cache.is_contiguous()
    assert state_cache.stride(-1) == 1
    assert compressed_cache.stride(-1) == 1
    assert k_work_metadata.ndim == 2 and k_work_metadata.shape[1] == 2
    is_2d_positions = positions.ndim == 2
    is_k_mrope = bool(mrope_section)
    cache_has_rope_pos = rope_pos_offset is not None
    assert rope_pos_offset is None or rope_pos_offset == head_dim
    if is_2d_positions:
        assert positions.shape == (3, num_tokens)
        assert is_k_mrope
        pos_stride_axis, pos_stride_token = positions.stride()
    else:
        assert positions.shape == (num_tokens,)
        pos_stride_axis, pos_stride_token = 0, positions.stride(0)
    section = mrope_section if mrope_section is not None else (0, 0, 0)
    assert len(section) == 3

    # Swept on MI355X (64-lane wavefronts) over TILE_T_Q 1-16, TILE_H_Q 1-4 and
    # num_warps 1/2/4, from 4 to 16384 tokens: no config beat these by more
    # than noise, and 2 or 4 warps were slower everywhere.
    if num_tokens <= 4096:
        TILE_T_Q, TILE_H_Q = 2, 2
    else:
        TILE_T_Q, TILE_H_Q = 2, 4
    num_k_work = k_work_metadata.shape[0]
    num_q_work = triton.cdiv(num_tokens, TILE_T_Q) * triton.cdiv(num_q_heads, TILE_H_Q)
    _qsa_pre_indexer_kernel[(num_k_work + num_q_work,)](
        q,
        q.stride(0),
        k,
        k.stride(0),
        positions,
        pos_stride_axis,
        pos_stride_token,
        cos_sin_cache,
        q_norm_weight,
        k_norm_weight,
        eps,
        q_out,
        q_out.stride(0),
        q_out.stride(1),
        state_cache,
        state_cache.stride(0),
        state_cache.stride(1),
        state_slots,
        state_block_table,
        state_block_table.stride(0),
        query_start_loc,
        logical_positions,
        compressed_slots,
        k_work_metadata,
        compressed_cache,
        compressed_cache.stride(0),
        compressed_cache.stride(1),
        num_tokens,
        state_cache.shape[0],
        compressed_cache.shape[0],
        num_k_work,
        HQ=num_q_heads,
        D=head_dim,
        TILE_T_Q=TILE_T_Q,
        TILE_H_Q=TILE_H_Q,
        COMPRESS_RATIO=compress_ratio,
        STATE_SIZE=state_cache.shape[1],
        COMP_PAGE_SIZE=compressed_cache.shape[1],
        IS_2D_POSITIONS=is_2d_positions,
        IS_K_MROPE=is_k_mrope,
        CACHE_HAS_ROPE_POS=cache_has_rope_pos,
        MROPE_H=section[1],
        MROPE_W=section[2],
        num_warps=1,
    )

supports_fused_pre_indexer(rotary_emb, head_dim, num_kv_heads, compress_ratio)

Report whether this indexer's shapes match the fused kernel's assumptions.

The kernel hard-codes the rotary layout and the single-KV-head group compression it was written for; everything it rejects has a working unfused path.

Source code in vllm/models/qwen4_exp/amd/ops/qsa_pre_indexer.py
def supports_fused_pre_indexer(
    rotary_emb: nn.Module,
    head_dim: int,
    num_kv_heads: int,
    compress_ratio: int,
) -> bool:
    """Report whether this indexer's shapes match the fused kernel's assumptions.

    The kernel hard-codes the rotary layout and the single-KV-head group
    compression it was written for; everything it rejects has a working unfused
    path.
    """
    rotary_dim = int(rotary_emb.rotary_dim)
    mrope_section = getattr(rotary_emb, "mrope_section", None)
    return (
        bool(getattr(rotary_emb, "is_neox_style", False))
        and (
            not mrope_section
            or (
                len(mrope_section) == 3
                and sum(mrope_section) == rotary_dim // 2
                and bool(getattr(rotary_emb, "mrope_interleaved", False))
            )
        )
        and head_dim == 128
        and rotary_dim == 64
        and num_kv_heads == 1
        and compress_ratio > 1
        and compress_ratio & (compress_ratio - 1) == 0
    )