Skip to content

vllm.model_executor.warmup.rocm_segmented_attn_autotune_warmup ¶

Startup-only, persistent launch tuning for segmented ROCm attention.

Functions:

_candidate_configs(default, batch, query_len, heads, dim, seq_len=None) ¶

Search launch tiles and split counts within a bounded workspace.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _candidate_configs(default, batch, query_len, heads, dim, seq_len=None):
    """Search launch tiles and split counts within a bounded workspace."""
    base = _normalized_config(default, batch, query_len, heads, dim)
    max_splits = base["splits"]
    long_d64_decode = (
        dim == 64
        and batch == 1
        and query_len == 1
        and seq_len is not None
        and seq_len >= 131072
    )
    if long_d64_decode:
        max_splits = 512
    elif dim == 64 and query_len <= 512 and seq_len is not None and seq_len >= 4096:
        max_splits = min(64, 2 * max_splits)
    candidates: list[dict] = []
    seen = set()

    def add(**updates):
        config = {**base, **updates}
        if (
            config["splits"] > max_splits
            or dim % config["bk"]
            or config["reduce_d"] > dim
        ):
            return
        identity = tuple(sorted(config.items()))
        if identity not in seen:
            seen.add(identity)
            candidates.append(config)

    if long_d64_decode:
        add()
        for splits in (128, 256, 512, 64):
            for bn, warps, stages in (
                (32, 1, 1),
                (32, 2, 1),
                (32, 2, 2),
                (64, 1, 1),
                (64, 2, 1),
                (64, 2, 2),
            ):
                add(
                    bm=16,
                    bn=bn,
                    bk=64,
                    splits=splits,
                    warps=warps,
                    stages=stages,
                )
        return candidates

    split_choices = [base["splits"]]
    if max_splits > base["splits"]:
        split_choices.append(max_splits)
    for divisor in (2, 4):
        split_choices.append(max(1, base["splits"] // divisor))
    split_choices.append(1)
    for splits in split_choices:
        add(splits=splits)

    tile_variants: tuple[tuple[int, int, int, int, int, int], ...] = (
        (16, 32, dim, 4, 1, 2),
        (16, 64, min(128, dim), 4, 1, 2),
        (32, 32, 64, 4, 2, 2),
        (32, 64, 64, 4, 1, 2),
        (64, 32, min(128, dim), 4, 1, 6),
        (128, 32, min(128, dim), 8, 1, 6),
    )
    if dim == 64:
        tile_variants += (
            (16, 16, 64, 4, 1, 2),
            (32, 16, 64, 4, 2, 2),
            (64, 32, 64, 4, 1, 2),
            (64, 64, 64, 4, 1, 2),
            (128, 32, 64, 4, 1, 2),
            (128, 64, 64, 4, 1, 2),
        )
    for bm, bn, bk, warps, stages, waves in tile_variants:
        add(
            bm=bm,
            bn=bn,
            bk=bk,
            warps=warps,
            stages=stages,
            waves_per_eu=waves,
        )

    half_splits = max(1, base["splits"] // 2)
    for bm, bn, bk, warps, stages, waves in tile_variants[:3]:
        add(
            bm=bm,
            bn=bn,
            bk=bk,
            splits=half_splits,
            warps=warps,
            stages=stages,
            waves_per_eu=waves,
        )

    if dim == 64:
        for splits in (max_splits, half_splits):
            for bm, bn, bk, warps, stages, waves in tile_variants[2:]:
                add(
                    bm=bm,
                    bn=bn,
                    bk=bk,
                    splits=splits,
                    warps=warps,
                    stages=stages,
                    waves_per_eu=waves,
                )

    add(waves_per_eu=6)
    add(stages=2 if base["stages"] == 1 else 1)
    add(reduce_d=64 if base["reduce_d"] != 64 else dim)
    return candidates

_load_records(path, identity, heads, kv_heads, dim, fp8) ¶

Load a canonical table and any crash-recovery TP shards.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _load_records(path, identity, heads, kv_heads, dim, fp8):
    """Load a canonical table and any crash-recovery TP shards."""
    records: dict[tuple, dict] = {}
    merged_shards = []
    sources = [path, *sorted(path.parent.glob(f"{path.name}.tp*.part"))]
    for source in sources:
        if not source.exists():
            continue
        try:
            saved = json.loads(source.read_text())
            if saved["identity"] != identity or not all(
                _valid_record(record, heads, kv_heads, dim, fp8)
                for record in saved["records"]
            ):
                raise ValueError("identity or record validation failed")
            records.update(
                (tuple(record["workload"]), record) for record in saved["records"]
            )
            if source != path:
                merged_shards.append(source)
        except (ValueError, KeyError, TypeError, AssertionError):
            logger.warning(
                "Ignoring invalid segmented attention tuning cache %s", source
            )
    return {"identity": identity, "records": list(records.values())}, merged_shards

_measure_configs(configs, run, device, eviction, rounds, graph_calls=0) ¶

Interleave candidates so clock and thermal drift affect them evenly.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _measure_configs(configs, run, device, eviction, rounds, graph_calls=0):
    """Interleave candidates so clock and thermal drift affect them evenly."""
    stream = torch.cuda.current_stream(device)
    samples: dict[tuple, list[float]] = {_config_key(config): [] for config in configs}
    graphs = {}
    if graph_calls:
        for config in configs:
            graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(graph):
                for _ in range(graph_calls):
                    run(config)
            graphs[_config_key(config)] = graph
        for graph in graphs.values():
            graph.replay()
        torch.accelerator.synchronize(device)
    for round_index in range(rounds):
        events = []
        order = _rotated(configs, round_index)
        if round_index % 2:
            order = list(reversed(order))
        for config in order:
            if not graph_calls:
                eviction.zero_()
            start = torch.cuda.Event(enable_timing=True)
            end = torch.cuda.Event(enable_timing=True)
            start.record(stream)
            if graph_calls:
                graphs[_config_key(config)].replay()
            else:
                run(config)
            end.record(stream)
            events.append((_config_key(config), start, end))
        events[-1][2].synchronize()
        for key, start, end in events:
            samples[key].append(start.elapsed_time(end) * 1000 / max(1, graph_calls))
    return samples

_select_tuned_winner(default, finalists, samples, min_promotion_speedup=_MIN_PROMOTION_SPEEDUP) ¶

Return a finalist only when it reliably beats the static incumbent.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _select_tuned_winner(
    default, finalists, samples, min_promotion_speedup=_MIN_PROMOTION_SPEEDUP
):
    """Return a finalist only when it reliably beats the static incumbent."""
    default_key = _config_key(default)
    default_samples = samples[default_key]
    comparisons = []
    for config in finalists:
        key = _config_key(config)
        config_samples = samples[key]
        paired_speedup = statistics.median(
            baseline / candidate
            for baseline, candidate in zip(default_samples, config_samples, strict=True)
        )
        comparisons.append(
            {
                "config": config,
                "us": statistics.median(config_samples),
                "paired_speedup_vs_default": paired_speedup,
            }
        )
    challenger = max(
        comparisons, key=lambda result: result["paired_speedup_vs_default"]
    )
    if challenger["paired_speedup_vs_default"] >= min_promotion_speedup:
        return challenger, comparisons
    return (
        next(result for result in comparisons if result["config"] == default),
        comparisons,
    )

_shard_workloads(workloads, world_size, max_tokens, heads, kv_heads, dim, fp8) ¶

Balance compile-affine (batch, query) groups across TP ranks.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _shard_workloads(workloads, world_size, max_tokens, heads, kv_heads, dim, fp8):
    """Balance compile-affine ``(batch, query)`` groups across TP ranks."""
    assignments: list[list[tuple[int, tuple]]] = [[] for _ in range(world_size)]
    loads = [0] * world_size
    groups: dict[tuple, dict] = {}
    for index, workload in enumerate(workloads):
        compile_group = groups.setdefault(workload[:2], {"weight": 0, "items": []})
        compile_group["weight"] += _workload_weight(
            workload, max_tokens, heads, kv_heads, dim, fp8
        )
        compile_group["items"].append((index, workload))
    for compile_group in sorted(
        groups.values(), key=lambda item: item["weight"], reverse=True
    ):
        rank = min(
            range(world_size), key=lambda candidate: (loads[candidate], candidate)
        )
        assignments[rank].extend(compile_group["items"])
        loads[rank] += compile_group["weight"]
    return [
        [workload for _, workload in sorted(assignment)] for assignment in assignments
    ]

_temporary_tuning_kernels(device) ¶

Release tuning-only HIP modules without invalidating existing graphs.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
@contextmanager
def _temporary_tuning_kernels(device):
    """Release tuning-only HIP modules without invalidating existing graphs."""
    device_index = torch.device(device).index
    if device_index is None:
        device_index = torch.accelerator.current_device_index()
    caches = []
    for kernel in (_segmented_attention_stage, _segmented_attention_reduce):
        cache = kernel.device_caches[device_index][0]
        caches.append((cache, set(cache)))
    try:
        yield
    finally:
        torch.accelerator.synchronize(device)
        for cache, existing_keys in caches:
            for key in cache.keys() - existing_keys:
                del cache[key]
        # CompiledKernel destruction unloads the HIP module. Keep pre-existing
        # kernels alive, since an earlier graph capture may reference them.
        gc.collect()
        torch.accelerator.empty_cache()

_tp_context() ¶

Return the initialized TP coordinator, or local-only tuning state.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _tp_context():
    """Return the initialized TP coordinator, or local-only tuning state."""
    from vllm.distributed.parallel_state import (
        get_tp_group,
        model_parallel_is_initialized,
    )

    if not model_parallel_is_initialized():
        return None, 0, 1
    group = get_tp_group()
    return group, group.rank_in_group, group.world_size

_workload_weight(workload, max_tokens, heads, kv_heads, dim, fp8) ¶

Estimate compile plus execution cost for TP workload balancing.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def _workload_weight(workload, max_tokens, heads, kv_heads, dim, fp8):
    """Estimate compile plus execution cost for TP workload balancing."""
    batch, query_len, seq_len = workload
    default = select_segmented_config(
        batch, query_len, seq_len, heads, kv_heads, dim, fp8
    )
    candidates = len(
        _candidate_configs(default, batch, query_len, heads, dim, seq_len=seq_len)
    )
    query_tokens = sum(_query_lengths(batch, query_len, max_tokens))
    attention_work = query_tokens * seq_len * heads * dim
    return candidates * (1 << 40) + attention_work

get_segmented_config(device, dtype, kv_dtype, heads, kv_heads, dim, page, scale, batch, query_len, seq_len, sliding_window=-1, causal=True, has_sinks=False) ¶

Return a warmed ceiling bucket, or None for the static fallback.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def get_segmented_config(
    device,
    dtype,
    kv_dtype,
    heads,
    kv_heads,
    dim,
    page,
    scale,
    batch,
    query_len,
    seq_len,
    sliding_window=-1,
    causal=True,
    has_sinks=False,
):
    """Return a warmed ceiling bucket, or None for the static fallback."""
    data = _TABLES.get(
        _key(
            device,
            dtype,
            kv_dtype,
            heads,
            kv_heads,
            dim,
            page,
            scale,
            sliding_window,
            causal,
            has_sinks,
        )
    )
    if data is None:
        return None
    records = data["records"]
    eligible = [
        record
        for record in records
        if record["workload"][0] >= batch
        and record["workload"][1] >= query_len
        and record["workload"][2] >= seq_len
    ]
    if not eligible:
        return None
    winner = min(
        eligible,
        key=lambda record: (
            record["workload"][1],
            record["workload"][0],
            record["workload"][2],
        ),
    )
    return dict(winner["best"])

rocm_segmented_attn_autotune_warmup(worker) ¶

Autotune RDNA segmented attention after KV allocation and before capture.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
def rocm_segmented_attn_autotune_warmup(worker: Worker) -> None:
    """Autotune RDNA segmented attention after KV allocation and before capture."""
    if (
        not current_platform.is_rocm()
        or not worker.vllm_config.kernel_config.enable_rocm_segmented_attn_autotune
    ):
        return

    from vllm.platforms.rocm import on_gfx1x
    from vllm.v1.attention.backends.rocm_segmented_attn import (
        RocmSegmentedAttentionImpl,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec

    if not on_gfx1x():
        return
    layers = []
    for layer in worker.vllm_config.compilation_config.static_forward_context.values():
        impl = getattr(layer, "impl", None)
        if not isinstance(impl, RocmSegmentedAttentionImpl):
            continue
        config = impl._segmented_attention_config
        if (
            config is not None
            and config.kernel_config.enable_rocm_segmented_attn_autotune
            and not impl._segmented_attention_warmed_up
            and impl.alibi_slopes is None
            and not impl.logits_soft_cap
        ):
            layers.append(layer)
    if not layers:
        return

    specs = worker.get_kv_cache_spec()
    layouts = tuple(
        (
            (spec.block_size, spec.page_size_bytes)
            if isinstance(spec, FullAttentionSpec)
            else (0, spec.max_memory_usage_bytes(worker.vllm_config))
        )
        for spec in specs.values()
    )
    cache_budget = sum(
        tensor.size
        for tensor in worker.model_runner.kv_cache_config.kv_cache_tensors
        if not tensor.host_resident
    )
    budget = _memory_budget(worker.device)
    for layer in layers:
        _warmup_segmented_attention(
            layer,
            worker.device,
            memory_budget_bytes=budget,
            cache_layouts=layouts,
            cache_budget_bytes=cache_budget,
        )

warmup_segmented_attention(device, dtype, heads, kv_heads, dim, page, scale, max_tokens, max_len, max_seqs, *, memory_budget_bytes=None, cache_layouts=(), cache_budget_bytes=None, kv_dtype=torch.bfloat16, sliding_window=-1, causal=True, max_query_len=MAX_QUERY_LEN, has_sinks=False, physical_max_len=None) ¶

Tune reachable segmented buckets before graph capture.

Source code in vllm/model_executor/warmup/rocm_segmented_attn_autotune_warmup.py
@torch.inference_mode()
def warmup_segmented_attention(
    device,
    dtype,
    heads,
    kv_heads,
    dim,
    page,
    scale,
    max_tokens,
    max_len,
    max_seqs,
    *,
    memory_budget_bytes=None,
    cache_layouts=(),
    cache_budget_bytes=None,
    kv_dtype=torch.bfloat16,
    sliding_window=-1,
    causal=True,
    max_query_len=MAX_QUERY_LEN,
    has_sinks=False,
    physical_max_len=None,
):
    """Tune reachable segmented buckets before graph capture."""
    if (
        dtype not in (torch.bfloat16, torch.float16)
        or kv_dtype
        not in (
            torch.bfloat16,
            torch.float16,
            torch.float8_e4m3fn,
            torch.float8_e4m3fnuz,
        )
        or dim not in (64, 128, 256)
        or kv_heads < 1
        or heads % kv_heads
        or not 1 <= heads // kv_heads <= 16
    ):
        return
    key = _key(
        device,
        dtype,
        kv_dtype,
        heads,
        kv_heads,
        dim,
        page,
        scale,
        sliding_window,
        causal,
        has_sinks,
    )
    identity = _identity(
        device,
        dtype,
        kv_dtype,
        heads,
        kv_heads,
        dim,
        page,
        scale,
        max_tokens,
        max_len,
        max_seqs,
        sliding_window,
        causal,
        max_query_len,
        has_sinks,
        physical_max_len,
    )
    if memory_budget_bytes is None:
        memory_budget_bytes = _memory_budget(device)
    workload_args = dict(
        memory_budget_bytes=memory_budget_bytes,
        dtype=dtype,
        kv_dtype=kv_dtype,
        heads=heads,
        kv_heads=kv_heads,
        dim=dim,
        page=page,
        cache_layouts=cache_layouts,
        cache_budget_bytes=cache_budget_bytes,
        max_query_len=max_query_len,
        physical_max_len=(
            physical_max_len if has_sinks and sliding_window >= 0 else None
        ),
    )
    workloads = list(_workloads(max_tokens, max_len, max_seqs, **workload_args))
    active = _TABLES.get(key, {})
    warmed = (
        {tuple(record["workload"]) for record in active.get("records", [])}
        if active.get("identity") == identity
        else set()
    )
    if all(workload in warmed for workload in workloads):
        return
    digest = hashlib.sha256(json.dumps(identity, sort_keys=True).encode()).hexdigest()
    path = Path(envs.VLLM_CACHE_ROOT) / "rocm_segmented_attention" / f"{digest}.json"
    path.parent.mkdir(parents=True, exist_ok=True)
    group, tp_rank, tp_size = _tp_context()
    if group is not None:
        digests = _gather_tp(group, digest)
        if any(peer != digest for peer in digests):
            logger.warning(
                "ROCm segmented attention TP ranks have different tuning "
                "identities; tuning this rank independently"
            )
            group, tp_rank, tp_size = None, 0, 1

    unpruned = len(
        list(_workloads(max_tokens, max_len, max_seqs, max_query_len=max_query_len))
    )
    start = time.monotonic()
    lock = FileLock(str(path) + ".lock") if tp_rank == 0 else None
    if lock is not None:
        lock.acquire()
    try:
        if tp_rank == 0:
            data, merged_shards = _load_records(
                path, identity, heads, kv_heads, dim, kv_dtype.itemsize == 1
            )
        else:
            data, merged_shards = None, []
        if group is not None:
            data = group.broadcast_object(data, src=0)
        assert data is not None
        records = {tuple(record["workload"]): record for record in data["records"]}
        missing = [workload for workload in workloads if workload not in records]
        loaded = len(workloads) - len(missing)
        assignments = _shard_workloads(
            missing,
            tp_size,
            max_tokens,
            heads,
            kv_heads,
            dim,
            kv_dtype.itemsize == 1,
        )
        local_workloads = assignments[tp_rank]
        logger.info(
            "ROCm segmented attention tuning plan: workloads=%d missing=%d "
            "pruned=%d tp_rank=%d/%d assigned=%d scratch_budget=%d MiB "
            "cache_budget=%s",
            len(workloads),
            len(missing),
            unpruned - len(workloads),
            tp_rank,
            tp_size,
            len(local_workloads),
            memory_budget_bytes // 2**20,
            cache_budget_bytes,
        )

        if not missing:
            if tp_rank == 0 and merged_shards:
                _save(path, data)
                for shard in merged_shards:
                    shard.unlink(missing_ok=True)
            _TABLES[key] = data
            logger.info(
                "ROCm segmented attention autotune ready: tp_rank=%d/%d "
                "local_tuned=0 tuned=0 loaded=%d failed=0 elapsed=%.2fs cache=%s",
                tp_rank,
                tp_size,
                loaded,
                time.monotonic() - start,
                path,
            )
            return

        shard_path = path.with_suffix(path.suffix + f".tp{tp_rank}.part")
        local_records = {}
        local_failed = 0
        fatal = None
        try:
            with _temporary_tuning_kernels(device):
                for workload in local_workloads:
                    try:
                        record = _tune_workload(
                            device,
                            dtype,
                            kv_dtype,
                            heads,
                            kv_heads,
                            dim,
                            page,
                            scale,
                            max_tokens,
                            workload,
                            sliding_window=sliding_window,
                            causal=causal,
                            has_sinks=has_sinks,
                            physical_seq_len=(
                                physical_max_len
                                if has_sinks and sliding_window >= 0
                                else None
                            ),
                        )
                    except (
                        torch.OutOfMemoryError,
                        RuntimeError,
                        AssertionError,
                    ) as error:
                        local_failed += 1
                        logger.warning(
                            "Skipping segmented attention tuning workload %s: %s",
                            workload,
                            error,
                        )
                    else:
                        local_records[workload] = record
                        _save(
                            shard_path,
                            {
                                "identity": identity,
                                "records": list(local_records.values()),
                            },
                        )
                    finally:
                        gc.collect()
                        torch.accelerator.empty_cache()
        except BaseException as error:
            fatal = f"{type(error).__name__}: {error}"

        gathered = _gather_tp(
            group,
            {
                "records": list(local_records.values()),
                "failed": local_failed,
                "fatal": fatal,
            },
        )
        if tp_rank == 0:
            fatal_errors = []
            for rank, contribution in enumerate(gathered):
                records.update(
                    (tuple(record["workload"]), record)
                    for record in contribution["records"]
                )
                if contribution["fatal"] is not None:
                    fatal_errors.append(f"TP rank {rank}: {contribution['fatal']}")
            data = {"identity": identity, "records": list(records.values())}
            _save(path, data)
            for shard in path.parent.glob(f"{path.name}.tp*.part"):
                shard.unlink(missing_ok=True)
            result = {
                "data": data,
                "tuned": sum(len(item["records"]) for item in gathered),
                "failed": sum(item["failed"] for item in gathered),
                "fatal": fatal_errors,
            }
        else:
            result = None
        if group is not None:
            result = group.broadcast_object(result, src=0)
        assert result is not None
        data = result["data"]
    finally:
        if lock is not None:
            lock.release()

    _TABLES[key] = data
    logger.info(
        "ROCm segmented attention autotune ready: tp_rank=%d/%d "
        "local_tuned=%d tuned=%d loaded=%d failed=%d elapsed=%.2fs cache=%s",
        tp_rank,
        tp_size,
        len(local_records),
        result["tuned"],
        loaded,
        result["failed"],
        time.monotonic() - start,
        path,
    )
    if result["fatal"]:
        raise RuntimeError(
            "ROCm segmented attention distributed tuning failed: "
            + "; ".join(result["fatal"])
        )