Skip to content

vllm.models.glm5next.amd.sparse_indexer

Custom Sparse Attention Indexer layers.

Classes:

SparseAttnIndexerKpool

Bases: CustomOp

Sparse Attention Indexer Custom Op Layer. This layer is extracted as a separate custom op since it involves heavy custom kernels like mqa_logits, paged_mqa_logits and top_k_per_row, etc. Those kernels maybe requires specific memory layout or implementation for different hardware backends to achieve optimal performance.

For now, the default native path will use CUDA backend path. Other platform may requires add the corresponding Custom Op name sparse_attn_indexer to custom_ops in CompilationConfig to enable the platform specific path.

Source code in vllm/models/glm5next/amd/sparse_indexer.py
@CustomOp.register("sparse_attn_indexer_kpool")
class SparseAttnIndexerKpool(CustomOp):
    """Sparse Attention Indexer Custom Op Layer. This layer is extracted as a
    separate custom op since it involves heavy custom kernels like `mqa_logits`,
    `paged_mqa_logits` and `top_k_per_row`, etc. Those kernels maybe requires
    specific memory layout or implementation for different hardware backends to
    achieve optimal performance.

    For now, the default native path will use CUDA backend path. Other platform
    may requires add the corresponding Custom Op name `sparse_attn_indexer` to
    `custom_ops` in `CompilationConfig` to enable the platform specific path.
    """

    def __init__(
        self,
        k_cache,
        quant_block_size: int,
        scale_fmt: str,
        topk_tokens: int,
        head_dim: int,
        max_pool_len: int,
        max_total_seq_len: int,
        topk_indices_buffer: torch.Tensor,
        skip_k_cache_insert: bool = False,
        use_fp4_cache: bool = False,
        tail_cache=None,
    ):
        super().__init__()
        self.k_cache = k_cache
        self.tail_cache = tail_cache
        self.quant_block_size = quant_block_size
        self.scale_fmt = scale_fmt
        self.topk_tokens = topk_tokens
        self.head_dim = head_dim
        self.max_pool_len = max_pool_len
        self.max_total_seq_len = max_total_seq_len
        self.topk_indices_buffer = topk_indices_buffer
        self.skip_k_cache_insert = skip_k_cache_insert
        self.use_fp4_cache = use_fp4_cache
        cfg = get_current_vllm_config_or_none()
        self.topk_backend = (
            cfg.kernel_config.sparse_indexer_topk_backend if cfg is not None else "auto"
        )

    def forward_hip(
        self,
        hidden_states: torch.Tensor,
        q_quant: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
        k: torch.Tensor,
        weights: torch.Tensor,
        *,
        gate_score: torch.Tensor | None = None,
        compress_ape: torch.Tensor | None = None,
        index_kpool: int = 1,
        positions: torch.Tensor | None = None,
    ):
        assert not self.use_fp4_cache, "AMD platform doesn't support fp4 cache yet"
        assert isinstance(q_quant, torch.Tensor), (
            "AMD sparse_attn_indexer expects a single FP8 q_quant tensor"
        )
        if not rocm_aiter_ops.is_enabled():
            raise RuntimeError(
                "Sparse attention indexer ROCm path is only supported on AITER. "
                "Please enable aiter with VLLM_ROCM_USE_AITER=1"
            )
        if index_kpool <= 1:
            return torch.ops.vllm.rocm_aiter_sparse_attn_indexer(
                hidden_states,
                _encode_layer_name(self.k_cache.prefix),
                self.k_cache.kv_cache,
                q_quant,
                k,
                weights,
                self.quant_block_size,
                self.scale_fmt,
                self.topk_tokens,
                self.head_dim,
                self.max_pool_len,
                self.max_total_seq_len,
                self.topk_indices_buffer,
                skip_k_cache_insert=self.skip_k_cache_insert,
            )
        return sparse_attn_indexer_kpool(
            hidden_states,
            self.k_cache.prefix,
            self.k_cache.kv_cache,
            q_quant,
            None,
            k,
            weights,
            self.quant_block_size,
            self.scale_fmt,
            self.topk_tokens,
            self.head_dim,
            self.max_pool_len,
            self.max_total_seq_len,
            self.topk_indices_buffer,
            self.skip_k_cache_insert,
            self.use_fp4_cache,
            gate_score,
            compress_ape,
            index_kpool,
            positions,
            self.tail_cache.kv_cache if self.tail_cache is not None else None,
            self.tail_cache.prefix if self.tail_cache is not None else None,
            self.topk_backend,
        )

    def forward_native(
        self,
        hidden_states: torch.Tensor,
        q_quant: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
        k: torch.Tensor,
        weights: torch.Tensor,
        *,
        gate_score: torch.Tensor | None = None,
        compress_ape: torch.Tensor | None = None,
        index_kpool: int = 1,
        positions: torch.Tensor | None = None,
    ) -> torch.Tensor:
        return self.forward_hip(
            hidden_states,
            q_quant,
            k,
            weights,
            gate_score=gate_score,
            compress_ape=compress_ape,
            index_kpool=index_kpool,
            positions=positions,
        )

_kpool_compress_insert(k, gate_score, ape, kv_cache, slot_mapping, kpool, head_dim, round_scale)

Pool kpool consecutive tokens into one fp8 K and write at pool slots.

slot_mapping is pool-granular (compress_ratio == kpool on the spec): only the last token of each complete pool carries a valid (>=0) slot; intra-pool tokens are -1. Every position is treated as a pool-completion candidate and non-completions are masked off inside the kernel. Compacting the valid rows first costs two device syncs on the eager prefill path and buys nothing numerically. Assumes pool-aligned chunk starts.

Source code in vllm/models/glm5next/amd/sparse_indexer.py
def _kpool_compress_insert(
    k: torch.Tensor,
    gate_score: torch.Tensor,
    ape: torch.Tensor,
    kv_cache: torch.Tensor,
    slot_mapping: torch.Tensor,
    kpool: int,
    head_dim: int,
    round_scale: bool,
) -> None:
    """Pool ``kpool`` consecutive tokens into one fp8 K and write at pool slots.

    ``slot_mapping`` is pool-granular (compress_ratio == kpool on the spec):
    only the *last* token of each complete pool carries a valid (>=0) slot;
    intra-pool tokens are -1. Every position is treated as a pool-completion
    candidate and non-completions are masked off inside the kernel. Compacting
    the valid rows first costs two device syncs on the eager prefill path and
    buys nothing numerically. Assumes pool-aligned chunk starts.
    """
    n = slot_mapping.shape[0]
    # No pool can complete in a batch smaller than one pool; also keeps the
    # clamped gather indices below in bounds.
    if n < kpool:
        return
    pos = torch.arange(n, device=k.device)
    valid = slot_mapping >= 0
    # Drop pools whose start falls before the batch (leading padding); their
    # gate/k data is undefined anyway.
    write_mask = valid & (pos >= kpool - 1)
    offs = torch.arange(kpool, device=k.device)
    idx = (pos - (kpool - 1)).clamp_min(0)[:, None] + offs[None, :]
    kpool_ops.kpool_compress_and_write_cache(
        kv_cache,
        k[idx],  # [n, kpool, head_dim]
        gate_score[idx],
        ape,
        slot_mapping.to(torch.int64),
        pool_size=kpool,
        head_dim=head_dim,
        write_mask=write_mask,
        round_scale=round_scale,
        write_cache=True,
        return_compressed=False,
    )

_kpool_decode_topk_backend(configured, *, num_rows, max_valid_seq_len, select_k, index_kpool, full_cudagraph)

Select the topk backend for decodes.

Heuristic based on ctx lengths: - <16k pools: in-tree hip kernel - 16-256k pools: use aiter - >256k pools: in-tree because not measured on >1m ctx

Since under FULL cudagraphs we cannot access the context length, and the aiter kernel does not always outperform the in-tree kernel, we default to in-tree kernel under FULL cudagraphs. Users can opt in if they know their context length is long enough to benefit from aiter.

Source code in vllm/models/glm5next/amd/sparse_indexer.py
def _kpool_decode_topk_backend(
    configured: str,
    *,
    num_rows: int,
    max_valid_seq_len: int,
    select_k: int,
    index_kpool: int,
    full_cudagraph: bool,
) -> str:
    """Select the topk backend for decodes.

    Heuristic based on ctx lengths:
    - <16k pools: in-tree hip kernel
    - 16-256k pools: use aiter
    - >256k pools: in-tree because not measured on >1m ctx

    Since under FULL cudagraphs we cannot access the context length,
    and the aiter kernel does not always outperform the in-tree kernel,
    we default to in-tree kernel under FULL cudagraphs. Users can opt in
    if they know their context length is long enough to benefit from aiter.
    """
    if (
        configured != "auto"
        or full_cudagraph
        or not rocm_aiter_ops.is_indexer_top_k_enabled()
    ):
        return configured
    if rocm_aiter_ops.is_indexer_top_k_supported(
        is_prefill=False,
        compress_ratio=index_kpool,
        num_rows=num_rows,
        max_valid_seq_len=max_valid_seq_len,
    ):
        return "aiter"
    if (
        index_kpool > 1
        and select_k == 512
        and _GFX950_KPOOL_AITER_MIN_POOLS
        <= max_valid_seq_len
        <= _GFX950_KPOOL_AITER_MAX_POOLS
    ):
        return "aiter"
    return configured