[AMD] fix(triton-mla): cap max_kv_splits at 256 on gfx942 (Kimi-K2.6 hang) (#24762)

This commit is contained in:
Bingxu Chen
2026-06-03 00:13:18 -07:00
committed by GitHub
parent e5b8e3a66a
commit 8e77af1afc
5 changed files with 27 additions and 4 deletions
+2 -2
View File
@@ -898,7 +898,7 @@ jobs:
fail-fast: false fail-fast: false
matrix: matrix:
runner: [linux-mi325-8gpu-sglang] runner: [linux-mi325-8gpu-sglang]
part: [0, 1, 2] part: [0, 1, 2, 3]
runs-on: ${{matrix.runner}} runs-on: ${{matrix.runner}}
steps: steps:
- name: Checkout code - name: Checkout code
@@ -925,7 +925,7 @@ jobs:
- name: Run test - name: Run test
timeout-minutes: 120 timeout-minutes: 120
run: | run: |
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd --auto-partition-id ${{ matrix.part }} --auto-partition-size 3 --timeout-per-file 5400 ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }} bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd --auto-partition-id ${{ matrix.part }} --auto-partition-size 4 --timeout-per-file 5400 ${{ (github.event_name == 'schedule' || inputs.continue_on_error) && '--continue-on-error' || '' }}
stage-c-test-large-8-gpu-amd-mi35x-rocm720: stage-c-test-large-8-gpu-amd-mi35x-rocm720:
needs: [check-changes] needs: [check-changes]
+1 -1
View File
@@ -997,7 +997,7 @@ jobs:
- name: Run test - name: Run test
timeout-minutes: 120 timeout-minutes: 120
run: | run: |
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd --auto-partition-id ${{ matrix.part }} --auto-partition-size 3 --timeout-per-file 5400 ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }} bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd --auto-partition-id ${{ matrix.part }} --auto-partition-size 4 --timeout-per-file 5400 ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }}
stage-c-test-large-8-gpu-amd-mi35x: stage-c-test-large-8-gpu-amd-mi35x:
needs: [check-changes, call-gate, wait-for-stage-b-amd] needs: [check-changes, call-gate, wait-for-stage-b-amd]
@@ -22,10 +22,12 @@ from sglang.srt.utils import (
get_device_core_count, get_device_core_count,
get_int_env_var, get_int_env_var,
is_cuda, is_cuda,
is_gfx942_supported,
next_power_of_2, next_power_of_2,
) )
_is_cuda = is_cuda() _is_cuda = is_cuda()
_is_gfx942 = is_gfx942_supported()
if _is_cuda: if _is_cuda:
from sgl_kernel.utils import is_arch_support_pdl from sgl_kernel.utils import is_arch_support_pdl
@@ -163,6 +165,15 @@ class TritonAttnBackend(AttentionBackend):
self.device_core_count, self.device_core_count,
self.max_context_len, self.max_context_len,
) )
if _is_gfx942:
# gfx942 (MI300X / MI325X) has 304 CUs, so #20479's next_power_of_2(sm_count)
# rounds up to 512 — twice MI355X's natural cap of 256 — and the persistent
# cuda_graph_attn_logits fp32 buffer hits ~4 GiB on Kimi-K2.6 (v_head_dim=512),
# faulting in ROCm CUDA graph replay
# (https://github.com/sgl-project/sglang/actions/runs/25513282022/job/74877480809).
# Pin the cap at 256 so gfx942 matches the gfx950 (MI355X) behavior that we
# already validated end-to-end.
self.max_kv_splits = min(self.max_kv_splits, 256)
if _is_cuda: if _is_cuda:
self.use_pdl = is_arch_support_pdl() self.use_pdl = is_arch_support_pdl()
else: else:
+12
View File
@@ -3656,6 +3656,18 @@ def is_gfx95_supported():
return False return False
@lru_cache(maxsize=1)
def is_gfx942_supported():
"""
Returns whether the current platform is AMD CDNA3 (gfx942 — MI300X / MI325X).
"""
if torch.version.hip:
gcn_arch = torch.cuda.get_device_properties(0).gcnArchName
return any(gfx in gcn_arch for gfx in ["gfx942"])
else:
return False
def get_hip_version(): def get_hip_version():
if torch.version.hip: if torch.version.hip:
return tuple(map(int, torch.version.hip.split("-")[0].split("."))) return tuple(map(int, torch.version.hip.split("-")[0].split(".")))
+1 -1
View File
@@ -63,7 +63,7 @@ class TestKimiK2Instruct0905(CustomTestCase):
num_shots=8, num_shots=8,
data_path=None, data_path=None,
num_questions=1319, num_questions=1319,
parallel=1319, parallel=512,
max_new_tokens=512, max_new_tokens=512,
host="http://127.0.0.1", host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]), port=int(self.base_url.split(":")[-1]),