[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:
co-authored by
amdpilot-upstream-sync
parent
f8618714c5
commit
9978aaec8b
@@ -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"]))
|
||||
Reference in New Issue
Block a user