Skip to content

vllm.v1.attention.ops.segmented_attention ¶

ROCm token-major segmented attention kernels and dispatch.

Functions:

_launch_segmented_attention(q, out, kc, vc, table, starts, lengths, k_scale, v_scale, sinks, scale, cfg, partial, lse, qcap, sliding_window, causal, *, compile_only) ¶

Launch or compile the exact stage and reduction specializations.

Source code in vllm/v1/attention/ops/segmented_attention.py
def _launch_segmented_attention(
    q,
    out,
    kc,
    vc,
    table,
    starts,
    lengths,
    k_scale,
    v_scale,
    sinks,
    scale,
    cfg,
    partial,
    lse,
    qcap,
    sliding_window,
    causal,
    *,
    compile_only,
):
    """Launch or compile the exact stage and reduction specializations."""
    batch, hq, dim = lengths.numel(), q.shape[1], q.shape[2]
    hk = kc.shape[2]
    fp8 = kc.element_size() == 1
    splits = cfg["splits"]
    stage_grid = (
        batch * hk,
        triton.cdiv(qcap * (hq // hk), cfg["bm"]),
        splits,
    )
    stage_args = (
        q,
        kc,
        vc,
        table,
        starts,
        lengths,
        k_scale,
        v_scale,
        sinks,
        partial,
        lse,
        out,
        q.stride(0),
        q.stride(1),
        out.stride(0),
        out.stride(1),
        *kc.stride(),
        *vc.stride(),
        *table.stride(),
        kc.shape[1],
        hq,
        hk,
        dim,
        qcap,
        MAX_QUERY_LEN,
        scale,
        fp8,
        cfg["bm"],
        cfg["bn"],
        cfg["bk"],
        splits,
        sliding_window,
        causal,
        sinks is not None,
        cfg.get("prefix_fast", False),
        cfg.get("pv_split", False),
        cfg.get("qk_pipeline"),
    )
    stage_options = dict(
        num_warps=cfg["warps"],
        num_stages=cfg["stages"],
        waves_per_eu=cfg.get("waves_per_eu", 2),
    )
    if compile_only:
        _segmented_attention_stage.warmup(*stage_args, grid=stage_grid, **stage_options)
    else:
        _segmented_attention_stage[stage_grid](*stage_args, **stage_options)
    if splits > 1:
        reduce_d = cfg.get("reduce_d", 64 if batch * qcap * hq < 64 else dim)
        reduce_grid = (
            qcap * (hq // hk),
            batch * hk,
            triton.cdiv(dim, reduce_d),
        )
        reduce_args = (
            partial,
            lse,
            out,
            starts,
            out.stride(0),
            out.stride(1),
            hq,
            hk,
            dim,
            qcap,
            MAX_QUERY_LEN,
            splits,
            reduce_d,
        )
        reduce_options = {"num_warps": cfg.get("reduce_warps", 4)}
        if compile_only:
            _segmented_attention_reduce.warmup(
                *reduce_args, grid=reduce_grid, **reduce_options
            )
        else:
            _segmented_attention_reduce[reduce_grid](*reduce_args, **reduce_options)

_long_extend_splits(base_splits, batch, query_len, seq_len, hq, hk, dim, bm) ¶

Add long-prefix parallelism while bounding the extra split workspace.

Source code in vllm/v1/attention/ops/segmented_attention.py
def _long_extend_splits(
    base_splits: int,
    batch: int,
    query_len: int,
    seq_len: int,
    hq: int,
    hk: int,
    dim: int,
    bm: int,
) -> int:
    """Add long-prefix parallelism while bounding the extra split workspace."""
    if seq_len < 32768:
        return base_splits
    groups = batch * hk * triton.cdiv(query_len * (hq // hk), bm)
    split_cap = 16 if seq_len < 131072 else 32
    scratch_per_split = batch * segmented_query_capacity(query_len) * hq * (dim + 1) * 4
    scratch_limit = MAX_LONG_EXTEND_WORKSPACE_BYTES * dim // 128
    scratch_splits = scratch_limit // scratch_per_split
    scratch_cap = 1 << (scratch_splits.bit_length() - 1) if scratch_splits else 1
    occupancy_splits = triton.next_power_of_2(triton.cdiv(512, groups))
    return max(base_splits, min(split_cap, scratch_cap, occupancy_splits))

can_use_segmented_attention(q, k, v, out, kc, vc, table, starts, lengths, k_scale, v_scale) ¶

Check tensor metadata; unsupported feature checks live in the caller.

Source code in vllm/v1/attention/ops/segmented_attention.py
def can_use_segmented_attention(
    q, k, v, out, kc, vc, table, starts, lengths, k_scale, v_scale
):
    """Check tensor metadata; unsupported feature checks live in the caller."""
    if (
        k is None
        or v is None
        or q.ndim != 3
        or kc.ndim != 4
        or vc.ndim != 4
        or q.dtype not in (torch.bfloat16, torch.float16)
        or out.dtype != q.dtype
        or q.shape != out.shape
        or q.shape[-1] not in (64, 128, 256)
        or kc.dtype != vc.dtype
        or k.dtype != q.dtype
        or v.dtype != q.dtype
        or kc.dtype not in (q.dtype, torch.float8_e4m3fn, torch.float8_e4m3fnuz)
        or table.ndim != 2
        or starts.ndim != 1
        or lengths.ndim != 1
        or table.shape[0] != lengths.numel()
        or starts.numel() != lengths.numel() + 1
        or table.dtype != torch.int32
        or starts.dtype != torch.int32
        or lengths.dtype != torch.int32
        or starts.stride(0) != 1
        or lengths.stride(0) != 1
        or not q.is_cuda
    ):
        return False
    hkv = kc.shape[2]
    dim = q.shape[-1]
    page = kc.shape[1]
    if hkv < 1 or q.shape[1] % hkv or not 1 <= q.shape[1] // hkv <= 16:
        return False
    if (
        k.shape != (q.shape[0], hkv, dim)
        or v.shape != k.shape
        or kc.shape != (kc.shape[0], page, hkv, dim)
        or vc.shape != kc.shape
        or any(
            t.device != q.device for t in (k, v, out, kc, vc, table, starts, lengths)
        )
        or any(t.stride(-1) != 1 for t in (q, k, v, out, kc, vc))
    ):
        return False
    if kc.element_size() == 1:
        return all(
            isinstance(s, torch.Tensor)
            and s.numel() == 1
            and s.dtype == torch.float32
            and s.device == q.device
            for s in (k_scale, v_scale)
        )
    return True

compile_segmented_attention(q, out, kc, vc, table, starts, lengths, max_query_len, k_scale, v_scale, scale, config, workspace, *, sliding_window=-1, causal=True, sinks=None) ¶

Compile one segmented attention configuration without launching it.

Source code in vllm/v1/attention/ops/segmented_attention.py
def compile_segmented_attention(
    q,
    out,
    kc,
    vc,
    table,
    starts,
    lengths,
    max_query_len,
    k_scale,
    v_scale,
    scale,
    config,
    workspace,
    *,
    sliding_window=-1,
    causal=True,
    sinks=None,
):
    """Compile one segmented attention configuration without launching it."""
    batch, hq, dim = lengths.numel(), q.shape[1], q.shape[2]
    qcap = segmented_query_capacity(min(max_query_len, MAX_QUERY_LEN))
    hk = kc.shape[2]
    shapes = segmented_workspace_shapes(batch, qcap, hq, hk, dim, config["splits"])
    if shapes is None:
        partial = lse = out
    else:
        partial, lse = workspace
    _launch_segmented_attention(
        q,
        out,
        kc,
        vc,
        table,
        starts,
        lengths,
        k_scale,
        v_scale,
        sinks,
        scale,
        config,
        partial,
        lse,
        qcap,
        sliding_window,
        causal,
        compile_only=True,
    )

get_segmented_attention_workspace(device, sizes) ¶

Keep scratch storage stable and separate for each runtime lane/ubatch.

Source code in vllm/v1/attention/ops/segmented_attention.py
def get_segmented_attention_workspace(device, sizes):
    """Keep scratch storage stable and separate for each runtime lane/ubatch."""

    def allocate():
        return tuple(
            torch.empty(size, device=device, dtype=torch.float32) for size in sizes
        )

    if is_workspace_manager_initialized():
        return current_workspace_manager().get_persistent_resource(
            ("rocm_segmented_attention", device, sizes), allocate
        )
    return allocate()

get_segmented_config(*args, **kwargs) ¶

Load the tuner on demand to avoid a kernel/tuner import cycle.

Source code in vllm/v1/attention/ops/segmented_attention.py
def get_segmented_config(*args, **kwargs):
    """Load the tuner on demand to avoid a kernel/tuner import cycle."""
    from vllm.model_executor.warmup.rocm_segmented_attn_autotune_warmup import (
        get_segmented_config as tuned_config,
    )

    return tuned_config(*args, **kwargs)

run_segmented_attention(q, out, kc, vc, table, starts, lengths, max_query_len, max_seq_len, k_scale, v_scale, scale, *, sliding_window=-1, causal=True, sinks=None, config=None, workspace=None) ¶

Write eligible unified-cache attention rows; leave other requests untouched.

Source code in vllm/v1/attention/ops/segmented_attention.py
def run_segmented_attention(
    q,
    out,
    kc,
    vc,
    table,
    starts,
    lengths,
    max_query_len,
    max_seq_len,
    k_scale,
    v_scale,
    scale,
    *,
    sliding_window=-1,
    causal=True,
    sinks=None,
    config=None,
    workspace=None,
):
    """Write eligible unified-cache attention rows; leave other requests untouched."""
    batch, hq, dim = lengths.numel(), q.shape[1], q.shape[2]
    if batch == 0 or max_query_len == 0 or q.numel() == 0:
        return out
    query_len = min(max_query_len, MAX_QUERY_LEN)
    qcap = segmented_query_capacity(query_len)
    fp8 = kc.element_size() == 1
    hk = kc.shape[2]
    attention_span = max_seq_len
    if sliding_window >= 0:
        attention_span = min(max_seq_len, sliding_window + max_query_len)
    cfg = config or select_segmented_config(
        batch, query_len, attention_span, hq, hk, dim, fp8
    )
    splits = cfg["splits"]
    shapes = segmented_workspace_shapes(batch, qcap, hq, hk, dim, splits)
    if shapes is None:
        partial = lse = out
    elif workspace is not None:
        partial, lse = (
            buffer.view(-1)[: shape.numel()].view(shape)
            for buffer, shape in zip(workspace, map(torch.Size, shapes))
        )
    elif is_workspace_manager_initialized():
        partial, lse = current_workspace_manager().get_simultaneous(
            (shapes[0], torch.float32), (shapes[1], torch.float32)
        )
    else:
        partial = torch.empty(shapes[0], device=q.device, dtype=torch.float32)
        lse = torch.empty(shapes[1], device=q.device, dtype=torch.float32)
    _launch_segmented_attention(
        q,
        out,
        kc,
        vc,
        table,
        starts,
        lengths,
        k_scale,
        v_scale,
        sinks,
        scale,
        cfg,
        partial,
        lse,
        qcap,
        sliding_window,
        causal,
        compile_only=False,
    )
    return out

segmented_attention(query, key, value, output, kv_cache_dtype, key_cache, value_cache, block_table, query_start_loc, seq_lens, max_seq_len, max_query_len, k_scale, v_scale, sm_scale, *, sliding_window=-1, output_scale=None, sinks=None, causal=True, softcap=0.0, workspace=None, use_tuned_config=True) ¶

Run segmented attention when eligible, otherwise use unified Triton.

Source code in vllm/v1/attention/ops/segmented_attention.py
def segmented_attention(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    output: torch.Tensor,
    kv_cache_dtype: str,
    key_cache: torch.Tensor,
    value_cache: torch.Tensor,
    block_table: torch.Tensor,
    query_start_loc: torch.Tensor,
    seq_lens: torch.Tensor,
    max_seq_len: int,
    max_query_len: int,
    k_scale: torch.Tensor,
    v_scale: torch.Tensor,
    sm_scale: float,
    *,
    sliding_window: int = -1,
    output_scale: torch.Tensor | None = None,
    sinks: torch.Tensor | None = None,
    causal: bool | torch.Tensor = True,
    softcap: float = 0.0,
    workspace=None,
    use_tuned_config: bool = True,
) -> None:
    """Run segmented attention when eligible, otherwise use unified Triton."""
    if kv_cache_dtype in ("fp8", "fp8_e4m3"):
        fp8_dtype = current_platform.fp8_dtype()
        if key_cache.dtype == torch.uint8:
            key_cache = key_cache.view(fp8_dtype)
            value_cache = value_cache.view(fp8_dtype)

    from vllm.platforms.rocm import on_gfx1x, on_gfx12x

    segmented_pattern = (
        isinstance(causal, bool)
        and (on_gfx12x() if key_cache.element_size() == 1 else on_gfx1x())
        and not softcap
        and output_scale is None
        and (
            sinks is None
            or (
                sinks.ndim == 1
                and sinks.shape[0] == query.shape[1]
                and sinks.device == query.device
                and sinks.dtype in (torch.float16, torch.bfloat16, torch.float32)
            )
        )
    )
    if (
        segmented_pattern
        and 0 < max_query_len <= MAX_QUERY_LEN
        and (
            can_use_segmented_attention(
                query,
                key,
                value,
                output,
                key_cache,
                value_cache,
                block_table,
                query_start_loc,
                seq_lens,
                k_scale,
                v_scale,
            )
        )
    ):
        attention_span = max_seq_len
        if sliding_window >= 0:
            attention_span = min(max_seq_len, sliding_window + max_query_len)
        config = None
        if use_tuned_config:
            config = get_segmented_config(
                query.device,
                query.dtype,
                key_cache.dtype,
                query.shape[1],
                key_cache.shape[2],
                query.shape[2],
                key_cache.shape[1],
                sm_scale,
                len(seq_lens),
                max_query_len,
                attention_span,
                sliding_window,
                causal,
                has_sinks=sinks is not None,
            )
        if config is None:
            config = select_segmented_config(
                len(seq_lens),
                max_query_len,
                attention_span,
                query.shape[1],
                key_cache.shape[2],
                query.shape[2],
                key_cache.element_size() == 1,
            )
        run_segmented_attention(
            query,
            output,
            key_cache,
            value_cache,
            block_table,
            query_start_loc,
            seq_lens,
            max_query_len,
            max_seq_len,
            k_scale,
            v_scale,
            sm_scale,
            sliding_window=sliding_window,
            causal=causal,
            sinks=sinks,
            config=config,
            workspace=workspace,
        )
        return

    unified_attention(
        q=query,
        k=key_cache,
        v=value_cache,
        out=output,
        cu_seqlens_q=query_start_loc,
        max_seqlen_q=max_query_len,
        seqused_k=seq_lens,
        max_seqlen_k=max_seq_len,
        softmax_scale=sm_scale,
        causal=causal,
        window_size=(-1, -1) if sliding_window < 0 else (sliding_window, 0),
        block_table=block_table,
        softcap=softcap,
        q_descale=None,
        k_descale=k_scale,
        v_descale=v_scale,
        output_scale=output_scale,
        sinks=sinks,
        kv_quant_mode=get_kv_quant_mode(kv_cache_dtype),
    )

segmented_attention_workspace_size(max_batch, hq, hk, dim, max_seq_len, *, max_tokens=None, fp8=False, autotune=False) cached ¶

Size the largest selected attention workspace before graph capture.

max_tokens bounds reachable (batch, query_len) pairs using one longest query and one token for each remaining sequence. Omitting it keeps the legacy reservation behavior for callers without scheduler limits. autotune reserves additional splits for tuned D64 candidates.

Source code in vllm/v1/attention/ops/segmented_attention.py
@lru_cache(maxsize=128)
def segmented_attention_workspace_size(
    max_batch,
    hq,
    hk,
    dim,
    max_seq_len,
    *,
    max_tokens=None,
    fp8=False,
    autotune=False,
):
    """Size the largest selected attention workspace before graph capture.

    ``max_tokens`` bounds reachable ``(batch, query_len)`` pairs using one
    longest query and one token for each remaining sequence. Omitting it keeps
    the legacy reservation behavior for callers without scheduler limits.
    ``autotune`` reserves additional splits for tuned D64 candidates.
    """
    largest = 0
    previous_capacity = 0
    for query_capacity in _query_capacity_buckets():
        query_lengths = {previous_capacity + 1, query_capacity}
        for query_len in query_lengths:
            if max_tokens is not None and query_len > max_tokens:
                continue
            batch_limit = max_batch
            if max_tokens is not None:
                batch_limit = min(batch_limit, max_tokens - query_len + 1)
            for batch in range(1, batch_limit + 1):
                cfg = select_segmented_config(
                    batch, query_len, max_seq_len, hq, hk, dim, fp8
                )
                splits = cfg["splits"]
                if dim == 64 and query_len <= 512 and max_seq_len >= 4096 and autotune:
                    split_limit = (
                        512
                        if batch == 1 and query_len == 1 and max_seq_len >= 131072
                        else MAX_SPLITS
                    )
                    splits = min(split_limit, 2 * splits)
                if splits > 1:
                    largest = max(
                        largest,
                        splits * batch * query_capacity * hq,
                    )
        previous_capacity = query_capacity
    return largest * dim, largest

segmented_query_capacity(max_query_len) ¶

Round a supported query length to its compiled power-of-two capacity.

Source code in vllm/v1/attention/ops/segmented_attention.py
def segmented_query_capacity(max_query_len: int) -> int:
    """Round a supported query length to its compiled power-of-two capacity."""
    if not 1 <= max_query_len <= MAX_QUERY_LEN:
        raise ValueError(f"Query length must be in [1, {MAX_QUERY_LEN}]")
    return 1 << (max_query_len - 1).bit_length()

select_segmented_config(batch, max_query_len, max_seq_len, hq, hk, dim, fp8) cached ¶

Select the token-major cache configuration for segmented attention.

Source code in vllm/v1/attention/ops/segmented_attention.py
@lru_cache(maxsize=512)
def select_segmented_config(batch, max_query_len, max_seq_len, hq, hk, dim, fp8):
    """Select the token-major cache configuration for segmented attention."""
    qcap = min(max_query_len, MAX_QUERY_LEN)
    if batch < 1 or qcap < 1:
        raise ValueError("Batch and query capacity must be positive")

    if qcap <= 2:
        if dim == 64 and qcap == 1 and batch == 1 and max_seq_len >= 131072:
            return dict(
                bm=16,
                bn=32,
                bk=64,
                splits=256,
                warps=1,
                stages=1,
            )
        if dim == 64 and max_seq_len <= 256:
            return dict(
                bm=16,
                bn=32 if fp8 else 64,
                bk=64,
                splits=1,
                warps=4,
                stages=1,
            )
        if dim == 64 and max_seq_len >= 32768:
            groups = batch * hk * triton.cdiv(qcap * (hq // hk), 16)
            return dict(
                bm=16,
                bn=32,
                bk=64,
                splits=min(64, triton.next_power_of_2(triton.cdiv(128, groups))),
                warps=4,
                stages=1,
            )
        if batch * hk == 1:
            cfg = dict(
                bm=16,
                bn=64,
                bk=min(128, dim),
                splits=32,
                warps=4,
                stages=1,
            )
        else:
            groups = batch * hk * triton.cdiv(qcap * (hq // hk), 16)
            cfg = dict(
                bm=16,
                bn=32,
                bk=dim,
                splits=min(
                    MAX_SPLITS,
                    triton.next_power_of_2(triton.cdiv(64, groups)),
                ),
                warps=4,
                stages=1,
            )
        while cfg["splits"] > 1 and max_seq_len < cfg["splits"] * cfg["bn"] * 2:
            cfg["splits"] //= 2
        return cfg

    if qcap <= 8:
        bm, bn = 16, 64
        bk, stages = 64, 1
    elif qcap <= 32:
        bm, bk = 32, 64
        bn, stages = (32, 2) if fp8 else (64, 1)
    else:
        bm, bn, bk = 32, 64, 64
        stages = 1 if fp8 else 2
    groups = batch * hk * triton.cdiv(qcap * (hq // hk), bm)
    target = 96
    if fp8 and qcap > 32:
        target = 192
    splits = min(32, triton.next_power_of_2(triton.cdiv(target, groups)))
    while splits > 1 and max_seq_len < splits * 128:
        splits //= 2
    cfg = dict(
        bm=bm,
        bn=bn,
        bk=bk,
        splits=splits,
        warps=4,
        stages=stages,
    )
    if dim == 64 and qcap >= 64 and qcap < max_seq_len <= qcap + 256:
        cfg.update(
            bm=64 if fp8 or qcap <= 256 else 128,
            bn=32,
            bk=64,
            splits=1,
            warps=4,
            stages=1,
        )
    elif dim == 64 and fp8 and max_seq_len <= 256:
        cfg.update(bn=32, splits=1)
    elif dim == 64 and qcap == 8 and max_seq_len >= 32768:
        groups = batch * hk * triton.cdiv(qcap * (hq // hk), 32)
        cfg.update(
            bm=32,
            bn=32 if fp8 else 64,
            splits=min(64, triton.next_power_of_2(triton.cdiv(256, groups))),
        )
    elif dim == 64 and fp8 and qcap == 8 and max_seq_len >= 4096:
        groups = batch * hk * triton.cdiv(qcap * (hq // hk), 32)
        cfg.update(
            bm=32,
            bn=64,
            splits=min(32, triton.next_power_of_2(triton.cdiv(128, groups))),
        )
    elif dim == 64 and fp8 and qcap >= 64 and max_seq_len >= 4096:
        groups = batch * hk * triton.cdiv(qcap * (hq // hk), 64)
        cfg.update(
            bm=64,
            bn=32,
            splits=min(8, triton.next_power_of_2(triton.cdiv(1024, groups))),
            stages=1,
        )
    elif dim == 128 and fp8 and qcap <= 256 and max_seq_len <= qcap:
        cfg.update(splits=1)
    elif (
        dim == 128
        and fp8
        and qcap >= 256
        and max_seq_len >= 8192
        and (
            hq // hk >= 4
            or (max_seq_len >= 32768 and (hq // hk >= 2 or qcap >= 512))
            or max_seq_len >= 131072
        )
    ):
        # Shorter prefixes need sufficient grouped-query reuse; at GQA1/Q256
        # the wider tile pays off only once the prefix reaches 131K tokens.
        cfg.update(
            bm=64,
            bn=128,
            bk=64,
            warps=8,
            stages=3,
            waves_per_eu=6,
            prefix_fast=True,
        )
        if qcap >= 512 or max_seq_len >= 32768:
            cfg["qk_pipeline"] = 3
        cfg["splits"] = _long_extend_splits(
            cfg["splits"], batch, qcap, max_seq_len, hq, hk, dim, cfg["bm"]
        )
    elif dim == 256 and fp8 and qcap >= 256 and max_seq_len >= 8192:
        # Wider tiles amortize each FP8 KV load over twice as many queries and
        # keys. Six waves is the best gfx1201 pressure/occupancy tradeoff.
        cfg.update(
            bm=64,
            bn=128,
            bk=64,
            warps=8,
            stages=1,
            waves_per_eu=6,
        )
        cfg["splits"] = _long_extend_splits(
            cfg["splits"], batch, qcap, max_seq_len, hq, hk, dim, cfg["bm"]
        )
    elif dim == 256 and not fp8 and qcap >= 1024 and max_seq_len >= 8192:
        # Two half-width PV accumulators avoid spilling this BF16 D256 tile.
        # Keep Q256 on the incumbent: the wider tile loses at that capacity.
        cfg.update(
            bm=64,
            bn=64,
            bk=64,
            warps=8,
            stages=3,
            waves_per_eu=6,
            pv_split=True,
            qk_pipeline=3,
        )
    elif dim == 128 and not fp8 and qcap > 128:
        if max_seq_len >= 8192:
            cfg.update(
                bm=128,
                bn=32,
                bk=128,
                warps=8,
                stages=1,
                waves_per_eu=6,
            )
            groups = batch * hk * triton.cdiv(qcap * (hq // hk), cfg["bm"])
            cfg["splits"] = min(16, triton.next_power_of_2(triton.cdiv(192, groups)))
        elif qcap >= 4096:
            cfg.update(
                bm=128,
                bn=32,
                bk=128,
                splits=1,
                warps=8,
                stages=1,
                waves_per_eu=6,
            )
        elif qcap > 256 or hq > 10:
            cfg.update(
                bm=64,
                bn=32,
                bk=128,
                splits=1,
                warps=4,
                stages=1,
                waves_per_eu=6,
            )
        else:
            cfg.update(bm=32, bn=64, bk=64, splits=1, warps=4, stages=2)
    return cfg