[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_BYTES_PER_ELEM = 4
|
||||||
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
|
_MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000
|
||||||
_MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3
|
_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] = {}
|
_mqa_logits_budget_bytes: Dict[int, int] = {}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -1003,7 +1005,8 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
self, num_q: int, num_k: int, device_index: int
|
self, num_q: int, num_k: int, device_index: int
|
||||||
) -> Tuple[bool, 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)
|
Return: (need_chunk, logits_budget_bytes)
|
||||||
"""
|
"""
|
||||||
# Quick static check for normal batches
|
# 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_bytes = num_q * num_k * self._MQA_LOGITS_BYTES_PER_ELEM
|
||||||
logits_budget_bytes = self._get_mqa_logits_budget_bytes(device_index)
|
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
|
need_chunk = logits_bytes > logits_budget_bytes
|
||||||
return need_chunk, 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