diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 18e760dff..27dd14586 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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 = ( diff --git a/test/registered/amd/test_moriep_small.py b/test/registered/amd/test_moriep_small.py index 76ccb42d6..1ffc27576 100644 --- a/test/registered/amd/test_moriep_small.py +++ b/test/registered/amd/test_moriep_small.py @@ -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,