diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index ba924e918..145b47a63 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -316,7 +316,9 @@ def compute_split_indices_for_cuda_graph_replay( class TboCudaGraphRunnerPlugin: def __init__(self): - self._tbo_children_num_token_non_padded = torch.zeros((2,), dtype=torch.int32) + self._tbo_children_num_token_non_padded = torch.zeros( + (2,), dtype=torch.int32, device=get_global_server_args().device + ) def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int): if not is_tbo_enabled(): diff --git a/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py b/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py new file mode 100644 index 000000000..0e01a46cd --- /dev/null +++ b/test/registered/unit/batch_overlap/test_tbo_cuda_graph_num_token_device.py @@ -0,0 +1,82 @@ +"""Regression: the TBO cuda-graph plugin must allocate its children +``num_token_non_padded`` buffer on the model device, not CPU. + +``ForwardBatch.num_token_non_padded`` is a scalar tensor on the model device +(see ``ForwardBatch.compute``, which does ``.to(device, ...)``). The eager TBO +split path already honors this -- ``compute_tbo_children_num_token_non_padded_raw`` +moves the tensor to ``get_global_server_args().device`` -- but +``TboCudaGraphRunnerPlugin`` preallocated its persistent buffer with a bare +``torch.zeros((2,), dtype=torch.int32)``, leaving it on CPU. + +During TBO cuda-graph replay that CPU buffer becomes the child's +``num_token_non_padded`` and is compared against a device ``arange`` in the +padded-region mask (``_mask_topk_ids_padded_region`` / +``_zero_topk_weights_padded_region``), raising a device-mismatch error +(an illegal memory access on HIP/aiter MoE). + +CPU-only: a ``meta`` device makes the buffer placement observable without a GPU. +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +import sglang.srt.batch_overlap.two_batch_overlap as tbo +from sglang.srt.batch_overlap.two_batch_overlap import ( + TboCudaGraphRunnerPlugin, + TboForwardBatchPreparer, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestTboCudaGraphNumTokenDevice(CustomTestCase): + def test_plugin_buffer_on_model_device(self): + # Use 'meta' so the configured device differs from the implicit CPU + # default; a bare torch.zeros() would leave the buffer on CPU and fail. + fake_args = SimpleNamespace(device="meta") + with patch.object(tbo, "get_global_server_args", lambda: fake_args): + plugin = TboCudaGraphRunnerPlugin() + + buf = plugin._tbo_children_num_token_non_padded + self.assertEqual(tuple(buf.shape), (2,)) + self.assertEqual(buf.dtype, torch.int32) + self.assertEqual(buf.device.type, "meta") + + def test_graph_and_eager_paths_agree_on_device(self): + # Both the preallocated cuda-graph buffer and the eager split tensor must + # land on the same (model) device, matching ForwardBatch's contract. + fake_args = SimpleNamespace(device="meta") + with patch.object(tbo, "get_global_server_args", lambda: fake_args): + eager = ( + TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw( + tbo_split_token_index=3, num_token_non_padded=8 + ) + ) + plugin = TboCudaGraphRunnerPlugin() + + self.assertEqual( + eager.device.type, + plugin._tbo_children_num_token_non_padded.device.type, + ) + + def test_eager_split_values(self): + # value_a = min(split, n); value_b = max(0, n - split). Computed on CPU + # so the values are materializable. + fake_args = SimpleNamespace(device="cpu") + with patch.object(tbo, "get_global_server_args", lambda: fake_args): + eager = ( + TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw( + tbo_split_token_index=3, num_token_non_padded=8 + ) + ) + self.assertEqual(eager.dtype, torch.int32) + self.assertEqual(eager.tolist(), [3, 5]) + + +if __name__ == "__main__": + unittest.main()