[AMD] fix tbo runtime error when initializing metadata for cuda graph (#22598)

This commit is contained in:
billishyahao
2026-04-19 12:42:48 -07:00
committed by GitHub
parent 32b7777f6c
commit b74a9dd854
2 changed files with 53 additions and 1 deletions
@@ -1193,8 +1193,26 @@ class AiterAttnBackend(AttentionBackend):
max_num_tokens: int,
kv_indices_buf: Optional[torch.Tensor] = None,
):
# PR #20978 pads max_bs beyond pool_size for higher cuda-graph
# coverage. Reallocate indptr buffers so they fit the padded max_bs.
# See: https://github.com/sgl-project/sglang/pull/20978
if max_bs + 1 > self.kv_indptr.shape[0]:
self.kv_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=self.device
)
self.qo_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=self.device
)
self.mask_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int64, device=self.device
)
if hasattr(self, "qo_indptr_"):
self.qo_indptr_ = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=self.device
)
self.cuda_graph_kv_last_page_len = torch.ones(
max_bs, dtype=torch.int, device=self.device
max_bs, dtype=torch.int32, device=self.device
)
if kv_indices_buf is None:
max_num_blocks_per_seq = (
+34
View File
@@ -4,7 +4,9 @@ from types import SimpleNamespace
import requests
from sglang.srt.server_args import ZMQ_TCP_PORT_DELTA
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.network import is_port_available
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
from sglang.test.test_utils import (
@@ -15,6 +17,30 @@ from sglang.test.test_utils import (
popen_launch_server,
)
def wait_all_ports_release(base_url, timeout_s=60):
"""Wait until all derived ports are fully released."""
import time
port = int(base_url.split(":")[-1])
# See https://github.com/sgl-project/sglang/blob/495ef8ec64b6b937e59cd530ad3150172061a008/python/sglang/srt/server_args.py#L6958-L6969
offsets = [
0, # no offset
ZMQ_TCP_PORT_DELTA, # dist_init_port
ZMQ_TCP_PORT_DELTA + 1, # detokenizer_port
ZMQ_TCP_PORT_DELTA + 2, # rpc_port
ZMQ_TCP_PORT_DELTA + 3, # metrics_port
ZMQ_TCP_PORT_DELTA + 4, # scheduler_input_port
]
for _ in range(timeout_s):
if all(is_port_available(port + off) for off in offsets):
return
time.sleep(1)
# Best-effort: log but don't raise so tearDown doesn't break the next class.
print(f"Warning: some ports still occupied after {timeout_s}s")
register_amd_ci(est_time=1200, suite="stage-c-test-large-8-gpu-amd")
common_args = [
@@ -84,6 +110,7 @@ class TestPureDP(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -128,6 +155,7 @@ class TestMTP(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -181,6 +209,7 @@ class TestNormal(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -230,6 +259,7 @@ class TestLowLatency(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -278,6 +308,7 @@ class TestTBOwithNormal(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -328,6 +359,7 @@ class TestTBOwithLowLatency(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -380,6 +412,7 @@ class TestMTPwithTBONormal(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,
@@ -440,6 +473,7 @@ class TestMTPwithTBOLowLatency(CustomTestCase):
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
wait_all_ports_release(cls.base_url)
def test_gsm8k(
self,