Add get_parallel(): a structured accessor for parallel-topology state (#28567)

This commit is contained in:
Cheng Wan
2026-06-17 20:23:43 -07:00
committed by GitHub
parent d27d8b24de
commit 53318911ca
184 changed files with 1871 additions and 1733 deletions
@@ -5,11 +5,10 @@ from types import SimpleNamespace
import numpy as np
import torch
from sglang.srt.layers import dp_attention as _dp_attn
from sglang.srt.runtime_context import get_parallel
# Patch DP-attention globals before importing backends
# TODO: change the interface of both trtllm_mla and flashinfer backends to take tp_size as an argument instead of patching
_dp_attn.get_attention_tp_size = lambda: 1 # TP size = 1 for unit test
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
+3 -3
View File
@@ -5,11 +5,11 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.layers import dp_attention as _dp_attn
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_cuda_ci
# Patch DP-attention globals before importing backends
_dp_attn.get_attention_tp_size = lambda: 1 # TP size = 1 for unit test
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.dsa.dsa_indexer import (
@@ -3,8 +3,6 @@
# Adapted from https://github.com/vllm-project/vllm/blob/2c58742dff8613a3bd7496f2008ce927e18d38d1/tests/kernels/mamba/test_mamba_mixer2.py
from unittest.mock import patch
import pytest
import torch
@@ -15,6 +13,7 @@ from sglang.srt.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_device, get_device_count
from sglang.test.ci.ci_register import register_cuda_ci
@@ -109,10 +108,10 @@ def mixer2_gated_norm_tensor_parallel(
gate_states = torch.randn(batch_size, seq_len, hidden_size)
import sglang.srt.layers.attention.mamba.mixer2_rms_norm_gated as m2
import sglang.srt.model_loader.weight_utils as wu
# Convenience: Avoid calling initialize_dp_attention
with patch.object(wu, "get_attention_tp_rank", return_value=local_rank):
# Force attn-TP rank through the context (the weight loader reads it via
# get_parallel().attn_tp_rank); avoids calling initialize_dp_attention.
with get_parallel().override(attn_tp_rank=local_rank):
# create gated-norm with TP
mixer = m2.Mixer2RMSNormGated(
full_hidden_size=hidden_size,
@@ -120,10 +119,8 @@ def mixer2_gated_norm_tensor_parallel(
)
mixer.weight.weight_loader(mixer.weight, weight)
with (
patch.object(m2, "get_tensor_model_parallel_world_size", return_value=1),
patch.object(m2, "get_tensor_model_parallel_rank", return_value=0),
):
# m2 reads tp via get_parallel().tp_size/rank — force it through the context.
with get_parallel().override(tp_size=1, tp_rank=0):
# create gated-norm without TP to compute reference
mixer_single_gpu = m2.Mixer2RMSNormGated(
full_hidden_size=hidden_size,
@@ -14,6 +14,7 @@ import torch
import sglang.srt.batch_overlap.two_batch_overlap as tbo
from sglang.srt.batch_overlap.two_batch_overlap import TboForwardBatchPreparer
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -37,7 +38,7 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch:
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
with patch.object(tbo, "get_attention_tp_size", lambda: 1), patch.object(
with get_parallel().override(attn_tp_size=1), patch.object(
tbo, "get_global_server_args", lambda: fake_args
):
return TboForwardBatchPreparer.filter_batch(
@@ -14,6 +14,7 @@ from torch import nn
from sglang.srt.layers.moe import topk as topk_module
from sglang.srt.layers.moe.topk import TopKConfig
from sglang.srt.models.deepseek_v2 import DeepseekV2MoE
from sglang.srt.runtime_context import get_parallel
from sglang.test.test_utils import CustomTestCase
@@ -74,10 +75,7 @@ class TestDeepEPWaterfillEPLB(CustomTestCase):
patch.object(topk_module, "_is_cuda", True),
patch.object(topk_module, "_use_aiter", False),
patch.object(topk_module, "is_deepep_class_backend", return_value=True),
patch.object(
topk_module, "get_moe_expert_parallel_world_size", return_value=8
),
patch.object(topk_module, "get_moe_expert_parallel_rank", return_value=7),
get_parallel().override(moe_ep_size=8, moe_ep_rank=7),
patch.object(
topk_module,
"_biased_grouped_topk_postprocess",
@@ -5,6 +5,7 @@ from unittest.mock import patch
import torch
from sglang.srt.layers import flashinfer_comm_fusion as fusion
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-c", runner_config="4-gpu-h100")
@@ -149,11 +150,7 @@ class TestFlashInferCommFusion(unittest.TestCase):
patch.object(
fusion, "is_flashinfer_available", return_value=True
),
patch.object(
fusion,
"get_attn_tensor_model_parallel_world_size",
return_value=world_size,
),
get_parallel().override(attn_tp_size=world_size),
patch.object(
fusion, "ensure_workspace_initialized", return_value=True
),
@@ -904,12 +904,11 @@ class TestBuildDecodeRegistry(unittest.TestCase):
)
def test_num_token_non_padded_gathered_dp_branch(self):
import unittest.mock as mock
from sglang.srt.model_executor import forward_batch_info as fbi
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
build_decode_registry,
)
from sglang.srt.runtime_context import get_parallel
ntnp = torch.zeros(1, dtype=torch.int32)
src = SimpleNamespace(
@@ -926,9 +925,7 @@ class TestBuildDecodeRegistry(unittest.TestCase):
)
# Gathered (DP) path: post_fill overwrites the FB copy with the local
# count. Pin attn-TP (size=2, rank=0) so the result is deterministic.
with mock.patch.object(
fbi, "get_attention_tp_size", return_value=2
), mock.patch.object(fbi, "get_attention_tp_rank", return_value=0):
with get_parallel().override(attn_tp_size=2, attn_tp_rank=0):
reg = build_decode_registry(
device=torch.device("cpu"),
max_bs=4,
@@ -0,0 +1,146 @@
"""Unit tests for runtime_context: delegation, singletons, and override()."""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import unittest
from unittest.mock import patch
from sglang.srt.runtime_context import (
ParallelContext,
RuntimeContext,
get_context,
get_parallel,
)
from sglang.test.test_utils import CustomTestCase
_PS = "sglang.srt.distributed.parallel_state"
_DP = "sglang.srt.layers.dp_attention"
SIZE_RANK_DELEGATIONS = [
("world_size", f"{_PS}.get_world_size"),
("world_rank", f"{_PS}.get_world_rank"),
("tp_size", f"{_PS}.get_tensor_model_parallel_world_size"),
("tp_rank", f"{_PS}.get_tensor_model_parallel_rank"),
("pp_size", f"{_PS}.get_pipeline_model_parallel_world_size"),
("pp_rank", f"{_PS}.get_pipeline_model_parallel_rank"),
("moe_ep_size", f"{_PS}.get_moe_expert_parallel_world_size"),
("moe_ep_rank", f"{_PS}.get_moe_expert_parallel_rank"),
("moe_dp_size", f"{_PS}.get_moe_data_parallel_world_size"),
("moe_dp_rank", f"{_PS}.get_moe_data_parallel_rank"),
("moe_tp_size", f"{_PS}.get_moe_tensor_parallel_world_size"),
("moe_tp_rank", f"{_PS}.get_moe_tensor_parallel_rank"),
("attn_tp_size", f"{_PS}.get_attn_tensor_model_parallel_world_size"),
("attn_tp_rank", f"{_PS}.get_attn_tensor_model_parallel_rank"),
("attn_cp_size", f"{_PS}.get_attn_context_model_parallel_world_size"),
("attn_cp_rank", f"{_PS}.get_attn_context_model_parallel_rank"),
("attn_dp_size", f"{_DP}.get_attention_dp_size"),
("attn_dp_rank", f"{_DP}.get_attention_dp_rank"),
]
GROUP_DELEGATIONS = [
("world_group", f"{_PS}.get_world_group"),
("tp_group", f"{_PS}.get_tp_group"),
("pp_group", f"{_PS}.get_pp_group"),
("moe_ep_group", f"{_PS}.get_moe_ep_group"),
("moe_dp_group", f"{_PS}.get_moe_dp_group"),
("moe_tp_group", f"{_PS}.get_moe_tp_group"),
("attn_tp_group", f"{_PS}.get_attn_tp_group"),
("attn_cp_group", f"{_PS}.get_attn_cp_group"),
]
class TestRuntimeContextSingletons(CustomTestCase):
def test_singletons(self):
self.assertIs(get_parallel(), get_parallel())
self.assertIsInstance(get_parallel(), ParallelContext)
self.assertIsInstance(get_context(), RuntimeContext)
self.assertIs(get_context().parallel, get_parallel())
class _IsolatedOverrides(CustomTestCase):
"""Give each test a clean override map, restoring afterward only the overrides
installed outside it (e.g. by another test file sharing the process)."""
def setUp(self):
super().setUp()
p = get_parallel()
self._saved_overrides = dict(p._overrides)
p._overrides.clear()
def tearDown(self):
p = get_parallel()
p._overrides.clear()
p._overrides.update(self._saved_overrides)
super().tearDown()
class TestParallelDelegation(_IsolatedOverrides):
def test_size_rank_delegate_to_canonical_getters(self):
# Patch each getter to a distinct sentinel: a miswired attribute would read
# a different (unpatched) getter and fail.
for i, (attr, target) in enumerate(SIZE_RANK_DELEGATIONS):
sentinel = 1000 + i
with patch(target, return_value=sentinel):
self.assertEqual(
getattr(get_parallel(), attr),
sentinel,
msg=f"{attr} must delegate to {target}",
)
def test_groups_delegate_to_canonical_getters(self):
for attr, target in GROUP_DELEGATIONS:
sentinel = object()
with patch(target, return_value=sentinel):
self.assertIs(
getattr(get_parallel(), attr),
sentinel,
msg=f"{attr} must delegate to {target}",
)
def test_wrapper_holds_no_resolved_state(self):
# __slots__: no __dict__; the only instance state is the override hook.
self.assertFalse(hasattr(get_parallel(), "__dict__"))
# tp_group IS exposed: live delegation handles PD-multiplexing / the tp patch.
self.assertTrue(hasattr(ParallelContext, "tp_group"))
# local_attn_dp is intentionally not part of the wrapper surface.
self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size"))
class TestParallelOverride(_IsolatedOverrides):
def test_override_takes_precedence(self):
p = get_parallel()
with p.override(tp_size=99, tp_rank=3, attn_dp_size=8):
self.assertEqual(p.tp_size, 99)
self.assertEqual(p.tp_rank, 3)
self.assertEqual(p.attn_dp_size, 8)
# same singleton: a fresh get_parallel() sees the override too
self.assertEqual(get_parallel().tp_size, 99)
self.assertEqual(p._overrides, {})
def test_override_can_force_groups(self):
sentinel = object()
with get_parallel().override(tp_group=sentinel):
self.assertIs(get_parallel().tp_group, sentinel)
def test_override_nests_and_restores(self):
p = get_parallel()
with p.override(tp_size=2):
self.assertEqual(p.tp_size, 2)
with p.override(tp_size=4, pp_size=2):
self.assertEqual(p.tp_size, 4)
self.assertEqual(p.pp_size, 2)
self.assertEqual(p.tp_size, 2)
self.assertNotIn("pp_size", p._overrides)
def test_override_unknown_key_raises_and_does_not_mutate(self):
p = get_parallel()
with self.assertRaises(ValueError):
with p.override(tp_sizee=1): # typo
pass
self.assertEqual(p._overrides, {})
if __name__ == "__main__":
unittest.main()