diff --git a/python/sglang/kernels/jit/csrc/speculative/tree.cuh b/python/sglang/kernels/jit/csrc/speculative/tree.cuh index 70d3ccbd4..66a6bd6b5 100644 --- a/python/sglang/kernels/jit/csrc/speculative/tree.cuh +++ b/python/sglang/kernels/jit/csrc/speculative/tree.cuh @@ -48,7 +48,7 @@ inline void build_tree_kernel_efficient( TensorMatcher({batch_size, parent_width}) .with_strides({parent_list.size(1) == 0 ? -1 : parent_list.size(1), 1}) .with_dtype() - .with_device(device) + .with_device(device) .verify(parent_list); CHECK_HOST(depth == 1 || parent_width.unwrap() == topk * (depth - 1) + 1); TensorMatcher({batch_size, draft_token_num - 1}) @@ -138,7 +138,7 @@ inline void verify_tree_greedy( SymbolicDevice device; TensorMatcher({batch_size, draft_tokens}) .with_dtype() - .with_device(device) + .with_device(device) .verify(candidates) .verify(retrive_index) .verify(retrive_next_token) @@ -183,7 +183,7 @@ inline void reconstruct_indices_from_tree_mask( // Bytes, not element type -- same reasoning as build_tree_kernel_efficient: // the kernel casts straight to bool* and callers are free to spell a 1-byte // mask as bool or uint8. - TensorMatcher({-1}).with_device(device).verify(tree_mask); + TensorMatcher({-1}).with_device(device).verify(tree_mask); CHECK_HOST(tree_mask.dtype().bits == 8); CHECK_HOST(tree_mask.numel() >= batch_size * draft_token_num * draft_token_num); TensorMatcher({batch_size}).with_dtype().with_device(device).verify(verified_seq_len); diff --git a/test/registered/kernels/ops/attention/test_q8kv8_sparse_prefill_backend.py b/test/registered/kernels/ops/attention/test_q8kv8_sparse_prefill_backend.py index 9080914ba..57d061339 100644 --- a/test/registered/kernels/ops/attention/test_q8kv8_sparse_prefill_backend.py +++ b/test/registered/kernels/ops/attention/test_q8kv8_sparse_prefill_backend.py @@ -78,7 +78,7 @@ class _TokenToKVPool: extra_key_buffer if extra_key_buffer is not None else swa_key_buffer ) self.full_to_swa_index_mapping = full_to_swa_index_mapping - self.swa_page_size = page_size + self.swa_kv_pool = _Pool(page_size) def get_swa_key_buffer_radix(self, layer_id: int) -> torch.Tensor: _ = layer_id @@ -86,7 +86,7 @@ class _TokenToKVPool: def get_extra_key_page_size(self, layer_id: int) -> int: _ = layer_id - return self.swa_page_size + return self.swa_kv_pool.page_size def get_extra_key_buffer(self, layer_id: int) -> torch.Tensor: _ = layer_id diff --git a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py index 2e38e67ba..b9feaec83 100644 --- a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py +++ b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py @@ -20,6 +20,7 @@ import unittest from pathlib import Path from unittest.mock import MagicMock, patch +from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_mlx_ci register_mlx_ci(est_time=5, suite="stage-a-unit-test-mlx") @@ -267,6 +268,7 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase): torch.mps.profiler, "metal_capture", return_value=capture_ctx ), mock_patch("torch.distributed.barrier"), + get_parallel().override(tp_rank=0, pp_size=1, moe_ep_size=1), ): result = mgr._start_profile() self.assertTrue(result.success, result.message) diff --git a/test/registered/unit/managers/test_disagg_idle_step_counters.py b/test/registered/unit/managers/test_disagg_idle_step_counters.py index 781f3b0db..140d0f379 100644 --- a/test/registered/unit/managers/test_disagg_idle_step_counters.py +++ b/test/registered/unit/managers/test_disagg_idle_step_counters.py @@ -16,6 +16,7 @@ from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.runtime_context import get_parallel from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -120,12 +121,11 @@ class TestSchedulerIdleStepCounters(CustomTestCase): scheduler.disagg_decode_transfer_queue.queue = [ object() ] - parallel = SimpleNamespace( - pp_async_batch_depth=depth, - enable_dsa_prefill_context_parallel=False, - ) with ( - patch(f"{PP_MODULE}.get_parallel", return_value=parallel), + get_parallel().override( + pp_size=2, + pp_async_batch_depth=depth, + ), patch( f"{PP_MODULE}.get_disagg", return_value=SimpleNamespace( @@ -278,7 +278,10 @@ class TestSchedulerIdleStepCounters(CustomTestCase): scheduler.run_batch = run_batch scheduler.process_batch_result = process_batch_result - with self.assertRaises(StopIteration): + with ( + get_parallel().override(pp_rank=0, attn_tp_rank=0, attn_cp_rank=0), + self.assertRaises(StopIteration), + ): event_loop(scheduler) self.assertEqual(observed_idle_flags, [False, False, after_idle, False])