Add get_parallel(): a structured accessor for parallel-topology state (#28567)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user