[ROCm][Bugfix] Cap the DSA MQA-logits budget at AITER's buffer_store limit (#36960)

Co-authored-by: amdpilot-upstream-sync <amdpilot-upstream-sync@users.noreply.github.com>
This commit is contained in:
Yikai Zhang
2026-09-01 15:40:12 -07:00
committed by GitHub
co-authored by amdpilot-upstream-sync
parent f8618714c5
commit 9978aaec8b
2 changed files with 66 additions and 1 deletions
@@ -203,6 +203,8 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
_MQA_LOGITS_BYTES_PER_ELEM = 4
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
_MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3
# aiter's fp8_mqa_logits only compiles below 2 GiB of logits (buffer_store).
_MQA_LOGITS_MAX_BYTES_ROCM = 2**31 - 1
_mqa_logits_budget_bytes: Dict[int, int] = {}
@staticmethod
@@ -1003,7 +1005,8 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
self, num_q: int, num_k: int, device_index: int
) -> Tuple[bool, int]:
"""
Detect whether we need to chunk the MQA logits computation to avoid OOM
Detect whether we need to chunk the MQA logits computation to avoid OOM,
and on ROCm to stay under aiter's 2 GiB logits limit
Return: (need_chunk, logits_budget_bytes)
"""
# Quick static check for normal batches
@@ -1012,6 +1015,10 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
logits_bytes = num_q * num_k * self._MQA_LOGITS_BYTES_PER_ELEM
logits_budget_bytes = self._get_mqa_logits_budget_bytes(device_index)
if _is_hip:
logits_budget_bytes = min(
logits_budget_bytes, self._MQA_LOGITS_MAX_BYTES_ROCM
)
need_chunk = logits_bytes > logits_budget_bytes
return need_chunk, logits_budget_bytes
@@ -0,0 +1,58 @@
"""Contract tests for the DSA indexer's MQA-logits chunk budget.
On ROCm the `[num_q x num_k]` fp32 logits tensor goes to aiter's
`fp8_mqa_logits`, which only compiles below 2 GiB, so the budget that decides
chunking is a correctness bound there and not only an out-of-memory guard.
The measured memory budget is stubbed: it is the only input the limit has to
beat, and stubbing it keeps these tests on CPU.
"""
from unittest import mock
import pytest
torch = pytest.importorskip("torch")
from sglang.srt.layers.attention.dsa import dsa_indexer # noqa: E402
from sglang.test.ci.ci_register import register_cpu_ci # noqa: E402
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
CEILING = dsa_indexer.Indexer._MQA_LOGITS_MAX_BYTES_ROCM
# More than any single logits tensor here needs, so it never decides a case.
HUGE_MEM_BUDGET = 64 * 2**30
def _decide(num_q, num_k, mem_budget=HUGE_MEM_BUDGET, is_hip=True):
# __new__ skips an __init__ that needs a model config and a device.
indexer = dsa_indexer.Indexer.__new__(dsa_indexer.Indexer)
with mock.patch.object(dsa_indexer, "_is_hip", is_hip), mock.patch.object(
dsa_indexer.Indexer,
"_get_mqa_logits_budget_bytes",
return_value=mem_budget,
):
return indexer._should_chunk_mqa_logits(num_q, num_k, 0)
def test_the_ceiling_is_the_largest_logits_aiter_still_takes():
# 16384 x 32768 x 4 bytes is exactly 2 GiB, and aiter compares `bytes <
# 2 GiB`, so that shape has to chunk and one KV token less must not.
assert _decide(16_384, 32_768) == (True, CEILING)
assert _decide(16_384, 32_767) == (False, CEILING)
def test_a_smaller_memory_budget_still_wins():
one_gib = 2**30
assert _decide(16_384, 32_767, mem_budget=one_gib) == (True, one_gib)
def test_off_rocm_the_budget_is_untouched():
# Elsewhere the logits go to DeepGEMM, which has no such limit.
assert _decide(16_384, 32_768, is_hip=False) == (False, HUGE_MEM_BUDGET)
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v"]))