Check the topology identities where the layout is written, and build at the published widths (#40340)
This commit is contained in:
@@ -14,7 +14,7 @@ from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatch
|
||||
from sglang.srt.layers.moe.utils import initialize_moe_config
|
||||
from sglang.srt.runtime_context import get_context, publish
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
|
||||
class TestFlashinferDispatcher(CustomTestCase):
|
||||
@@ -44,9 +44,8 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
publish(server_args, role="scheduler")
|
||||
initialize_moe_config()
|
||||
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=world_size, expert_model_parallel_size=world_size
|
||||
)
|
||||
publish_build_topology(tp_size=world_size, ep_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
|
||||
@@ -18,7 +18,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
|
||||
def get_open_port() -> int:
|
||||
@@ -98,7 +98,8 @@ class TestCustomAllReduce(CustomTestCase):
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
|
||||
# Set global server args to avoid "Global server args is not set yet!" error
|
||||
@@ -161,7 +162,8 @@ class TestCustomAllReduce(CustomTestCase):
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
|
||||
# Set global server args to avoid "Global server args is not set yet!" error
|
||||
|
||||
@@ -23,7 +23,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
graph_capture,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
torch.manual_seed(42)
|
||||
random.seed(44) # keep the deterministic seed
|
||||
@@ -117,7 +117,8 @@ class TestQuickAllReduce(CustomTestCase):
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
|
||||
# A small all_reduce for warmup.
|
||||
@@ -186,7 +187,8 @@ class TestQuickAllReduce(CustomTestCase):
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
|
||||
for sz in self.TEST_SIZES:
|
||||
|
||||
@@ -24,6 +24,7 @@ import unittest
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
MODEL = "Qwen/Qwen2-0.5B"
|
||||
|
||||
@@ -43,7 +44,8 @@ def _init_model_parallel() -> None:
|
||||
local_rank=0,
|
||||
distributed_init_method="tcp://127.0.0.1:29634",
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=1)
|
||||
publish_build_topology(tp_size=1)
|
||||
initialize_model_parallel()
|
||||
monkey_patch_vllm_parallel_state()
|
||||
except AssertionError:
|
||||
pass
|
||||
|
||||
@@ -23,7 +23,11 @@ from sglang.srt.utils.rank_consensus_checker import (
|
||||
shutdown,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase, find_available_port
|
||||
from sglang.test.test_utils import (
|
||||
CustomTestCase,
|
||||
find_available_port,
|
||||
publish_build_topology,
|
||||
)
|
||||
|
||||
register_cpu_ci(est_time=193, suite="stage-a-test-cpu-intel")
|
||||
|
||||
@@ -80,11 +84,8 @@ def run_distributed_test(
|
||||
backend="gloo",
|
||||
)
|
||||
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=tp_size,
|
||||
pipeline_model_parallel_size=pp_size,
|
||||
backend="gloo",
|
||||
)
|
||||
publish_build_topology(tp_size=tp_size, pp_size=pp_size, world_rank=rank)
|
||||
initialize_model_parallel(backend="gloo")
|
||||
|
||||
fn()
|
||||
except Exception as e:
|
||||
|
||||
@@ -10,7 +10,11 @@ import torch
|
||||
from transformers import MistralConfig, PretrainedConfig
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTestCase
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
CustomTestCase,
|
||||
publish_build_topology,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-small")
|
||||
|
||||
@@ -77,7 +81,8 @@ class TestDraftEmbedScan(CustomTestCase):
|
||||
init_distributed_environment(
|
||||
world_size=1, rank=0, local_rank=0, distributed_init_method="env://"
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=1)
|
||||
publish_build_topology(tp_size=1)
|
||||
initialize_model_parallel()
|
||||
torch.set_default_dtype(torch.bfloat16)
|
||||
torch.cuda.set_device(0)
|
||||
|
||||
|
||||
@@ -38,6 +38,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cuda_ci(est_time=18, stage="base-b", runner_config="2-gpu-large")
|
||||
|
||||
@@ -238,10 +239,10 @@ def _worker_main(local_rank: int, world_size: int):
|
||||
init_distributed_environment(
|
||||
world_size=world_size, rank=local_rank, local_rank=local_rank
|
||||
)
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=world_size,
|
||||
expert_model_parallel_size=world_size,
|
||||
publish_build_topology(
|
||||
tp_size=world_size, ep_size=world_size, world_rank=local_rank
|
||||
)
|
||||
initialize_model_parallel()
|
||||
|
||||
from sglang.srt.eplb.lplb_solver import clear_global_lplb_solvers
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from sglang.srt.distributed import parallel_state as ps
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kernels.utils import multigpu_pytest_main
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
|
||||
|
||||
@@ -37,7 +38,8 @@ def group():
|
||||
local_rank=local_rank,
|
||||
distributed_init_method="env://",
|
||||
)
|
||||
ps.initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
ps.initialize_model_parallel()
|
||||
yield ps.get_tp_group()
|
||||
ps.destroy_model_parallel()
|
||||
ps.destroy_distributed_environment()
|
||||
|
||||
@@ -19,6 +19,7 @@ import pytest
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cuda_ci(est_time=12, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
@@ -52,12 +53,8 @@ def _runtime_scaffolding():
|
||||
if not torch.distributed.is_initialized():
|
||||
init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="gloo")
|
||||
if not model_parallel_is_initialized():
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
expert_model_parallel_size=1,
|
||||
pipeline_model_parallel_size=1,
|
||||
backend="gloo",
|
||||
)
|
||||
publish_build_topology(tp_size=1, ep_size=1, pp_size=1)
|
||||
initialize_model_parallel(backend="gloo")
|
||||
|
||||
|
||||
def _interleave_w13_rows(w13: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
@@ -17,6 +17,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
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, register_xpu_ci
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cuda_ci(est_time=30, stage="base-b", runner_config="2-gpu-large")
|
||||
register_xpu_ci(est_time=60, suite="nightly-xpu-2-gpu", nightly=True)
|
||||
@@ -105,7 +106,8 @@ def mixer2_gated_norm_tensor_parallel(
|
||||
local_rank=local_rank,
|
||||
backend=get_default_distributed_backend(device.type),
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=local_rank)
|
||||
initialize_model_parallel()
|
||||
|
||||
# create random weights an inputs
|
||||
weight = torch.rand((hidden_size,), dtype=dtype, device=device)
|
||||
|
||||
@@ -14,7 +14,7 @@ import torch
|
||||
from sglang.srt.layers import communicator as comm
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, ScatterMode
|
||||
from sglang.test.ci.ci_register import register_amd_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
register_amd_ci(est_time=240, suite="stage-c-test-large-8-gpu-amd")
|
||||
|
||||
@@ -64,7 +64,8 @@ def _run_residual_accuracy_check():
|
||||
distributed_init_method="env://",
|
||||
backend="nccl",
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
publish_build_topology(tp_size=world_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
|
||||
dtype = torch.bfloat16
|
||||
eps = 1e-6
|
||||
|
||||
@@ -44,6 +44,7 @@ import pytest
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
||||
|
||||
@@ -232,11 +233,8 @@ def test_parallel_group_construction_tp8_attn_cp2():
|
||||
mock_world_group.return_value = mock_world
|
||||
|
||||
# Call the actual function
|
||||
parallel_state.initialize_model_parallel(
|
||||
tensor_model_parallel_size=8,
|
||||
pipeline_model_parallel_size=1,
|
||||
attention_context_model_parallel_size=2,
|
||||
)
|
||||
publish_build_topology(tp_size=8, pp_size=1, attn_cp_size=2)
|
||||
parallel_state.initialize_model_parallel()
|
||||
|
||||
# Verify TP groups
|
||||
tp_groups = created_groups.get("tp", [])
|
||||
@@ -330,12 +328,8 @@ def test_parallel_group_construction_tp8_moe_ep4_cp2():
|
||||
mock_world_group.return_value = mock_world
|
||||
|
||||
# Call the actual function
|
||||
parallel_state.initialize_model_parallel(
|
||||
tensor_model_parallel_size=8,
|
||||
expert_model_parallel_size=4,
|
||||
pipeline_model_parallel_size=1,
|
||||
moe_data_model_parallel_size=2,
|
||||
)
|
||||
publish_build_topology(tp_size=8, ep_size=4, pp_size=1, moe_dp_size=2)
|
||||
parallel_state.initialize_model_parallel()
|
||||
|
||||
# Verify TP groups
|
||||
tp_groups = created_groups.get("tp", [])
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
from sglang.srt.layers.attention import vision
|
||||
from sglang.srt.runtime_context import get_context, get_parallel
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import publish_build_topology
|
||||
|
||||
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
|
||||
|
||||
@@ -52,7 +53,8 @@ def gloo_world():
|
||||
distributed_init_method=f"tcp://127.0.0.1:{port}",
|
||||
backend="gloo",
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=1, backend="gloo")
|
||||
publish_build_topology(tp_size=1)
|
||||
initialize_model_parallel(backend="gloo")
|
||||
yield
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
@@ -20,7 +20,7 @@ import torch
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
register_cuda_ci(est_time=28, stage="base-c", runner_config="4-gpu-b200")
|
||||
|
||||
@@ -54,10 +54,8 @@ def _run(rank: int, world: int, port: int):
|
||||
distributed_init_method=f"tcp://127.0.0.1:{port}",
|
||||
backend="nccl",
|
||||
)
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=world,
|
||||
attention_context_model_parallel_size=world,
|
||||
)
|
||||
publish_build_topology(tp_size=world, attn_cp_size=world, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
|
||||
from sglang.srt.mem_cache.dsa_cache_layer_split import (
|
||||
LayerSplitDSATokenToKVPool,
|
||||
|
||||
@@ -84,7 +84,7 @@ from sglang.srt.runtime_context import get_parallel, publish
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import ceil_div
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.test_utils import CustomTestCase, publish_build_topology
|
||||
|
||||
register_cpu_ci(est_time=30, suite="base-a-test-cpu")
|
||||
|
||||
@@ -1215,10 +1215,8 @@ def _dist_init(rank, world, port, attn_cp_size):
|
||||
ServerArgs(model_path="dummy", tp_size=world, attn_cp_size=attn_cp_size),
|
||||
role="scheduler",
|
||||
)
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=world,
|
||||
attention_context_model_parallel_size=attn_cp_size,
|
||||
)
|
||||
publish_build_topology(tp_size=world, attn_cp_size=attn_cp_size, world_rank=rank)
|
||||
initialize_model_parallel()
|
||||
|
||||
|
||||
def _gather_make_spec(shard_rank, max_prefix_groups=16, chunk_groups=4):
|
||||
|
||||
@@ -844,7 +844,9 @@ class TestShardConfig(unittest.TestCase):
|
||||
override.install()
|
||||
self.addCleanup(override.restore)
|
||||
with (
|
||||
get_parallel().override(tp_size=8, pp_size=1, moe_dp_size=2, moe_ep_size=4),
|
||||
get_parallel().override(
|
||||
tp_size=8, pp_size=1, moe_dp_size=2, moe_ep_size=4, moe_tp_size=1
|
||||
),
|
||||
mock.patch(
|
||||
"sglang.srt.layers.dp_attention.get_moe_cp_size",
|
||||
return_value=2,
|
||||
|
||||
@@ -83,8 +83,19 @@ class TestGlm5NextBfgFusion(unittest.TestCase):
|
||||
for attn_tp, rank in ((1, 0), (2, 0), (2, 1)):
|
||||
with (
|
||||
self.subTest(route=expected_route, attn_tp=attn_tp, rank=rank),
|
||||
# A width is a whole topology: the attention triple has
|
||||
# to factor `tp_size`, and this process has to sit where
|
||||
# the triple puts it.
|
||||
get_parallel().override(
|
||||
tp_size=4, tp_rank=3, attn_tp_size=attn_tp, attn_tp_rank=rank
|
||||
tp_size=4,
|
||||
tp_rank=rank,
|
||||
attn_tp_size=attn_tp,
|
||||
attn_tp_rank=rank,
|
||||
attn_dp_size=4 // attn_tp,
|
||||
attn_dp_rank=0,
|
||||
attn_cp_size=1,
|
||||
attn_cp_rank=0,
|
||||
moe_tp_size=4,
|
||||
),
|
||||
):
|
||||
quant = MockFp8Config(ignored)
|
||||
|
||||
@@ -61,7 +61,15 @@ class _FusionGateCase(CustomTestCase):
|
||||
def _reason(self, model_class, hf_config, quant_config=None, moe_ep_size=1):
|
||||
# The gates consult the live EP size; without a group installed the
|
||||
# canonical getter asserts, so every case states a topology.
|
||||
with get_parallel().override(moe_ep_size=moe_ep_size):
|
||||
with get_parallel().override(
|
||||
tp_size=moe_ep_size,
|
||||
attn_tp_size=moe_ep_size,
|
||||
attn_dp_size=1,
|
||||
attn_cp_size=1,
|
||||
moe_ep_size=moe_ep_size,
|
||||
moe_dp_size=1,
|
||||
moe_tp_size=1,
|
||||
):
|
||||
return model_class.shared_experts_fusion_disable_reason(
|
||||
hf_config, quant_config
|
||||
)
|
||||
|
||||
@@ -41,6 +41,7 @@ from sglang.srt.runtime_context import (
|
||||
RuntimeContext,
|
||||
SpawnRanks,
|
||||
_FlagGroupBase,
|
||||
_validate_parallel,
|
||||
assert_published,
|
||||
derive_parallel_widths,
|
||||
get_context,
|
||||
@@ -59,6 +60,56 @@ from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
_SRT = _pathlib.Path(next(iter(_sglang.__path__))).resolve() / "srt"
|
||||
_PACKAGE = _pathlib.Path(next(iter(_sglang.__path__))).resolve()
|
||||
|
||||
|
||||
def _sources():
|
||||
"""Every Python file this checkout ships, package and siblings alike.
|
||||
|
||||
The package alone is the wrong subject set for anything about entries or
|
||||
public names: `benchmark/`, `examples/` and the top-level `test/` call the
|
||||
same doors and are not covered by any suite that would notice them break.
|
||||
An installed package has no siblings, and then this is the package alone."""
|
||||
roots = [_PACKAGE]
|
||||
checkout = _PACKAGE.parents[1]
|
||||
roots += [
|
||||
checkout / name
|
||||
for name in ("benchmark", "examples", "scripts", "test")
|
||||
if (checkout / name).is_dir()
|
||||
]
|
||||
for root in roots:
|
||||
for path in root.rglob("*.py"):
|
||||
yield path
|
||||
|
||||
|
||||
def _scope_entries_that_say_nothing(paths):
|
||||
"""Draft-scope entries that do not state `owns_attention`, as `path:line`.
|
||||
|
||||
The scope either narrows the draft's attention and expert identity or
|
||||
leaves the target's in place, and only the worker knows which -- so the
|
||||
keyword has no default. Omitting it is a `TypeError`, but only on the path
|
||||
that runs, and those paths want a GPU and a draft model.
|
||||
"""
|
||||
import ast
|
||||
|
||||
missing = []
|
||||
for path in paths:
|
||||
try:
|
||||
tree = ast.parse(path.read_text(encoding="utf-8"))
|
||||
except (SyntaxError, UnicodeDecodeError):
|
||||
continue
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
func = node.func
|
||||
name = getattr(func, "attr", None) or getattr(func, "id", None)
|
||||
if name not in ("draft_tp_context", "patch_tensor_parallel_group"):
|
||||
continue
|
||||
if not any(kw.arg == "owns_attention" for kw in node.keywords):
|
||||
missing.append(f"{path}:{node.lineno}")
|
||||
return missing
|
||||
|
||||
|
||||
_PS = "sglang.srt.distributed.parallel_state"
|
||||
_DP = "sglang.srt.layers.dp_attention"
|
||||
|
||||
@@ -347,9 +398,8 @@ class TestStampedRanks(_IsolatedOverrides):
|
||||
"""`attn_dp_rank` comes from the stamp, and says so when there is none.
|
||||
|
||||
It is the one rank no group answers with: `initialize_dp_attention`
|
||||
computes it from this process's `tp_rank`, and an elastic scale-up
|
||||
replaces it with a rank in the expanded WORLD. Falling back to anything
|
||||
would be inventing a placement for this process.
|
||||
computes it from this process's `tp_rank`. Falling back to anything would
|
||||
be inventing a placement for this process.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
@@ -407,21 +457,71 @@ class TestStampedRanks(_IsolatedOverrides):
|
||||
)
|
||||
self.assertIs(mode, DpPaddingMode.SUM_LEN)
|
||||
|
||||
def test_a_scale_up_stamps_the_width_and_the_rank_together(self):
|
||||
"""The two describe one topology; a reader that saw only one moved
|
||||
would place this process in a group it is not in."""
|
||||
def test_the_gather_slot_follows_the_list_that_was_gathered(self):
|
||||
"""The DP sync gathers over the attention-DP replicas, or over the
|
||||
expanded WORLD once a scale-up has moved the gather there. The index
|
||||
into that list is a property of the gather, so it is read beside the
|
||||
flag that says which one happened rather than kept on the topology."""
|
||||
from sglang.srt.layers.dp_attention import dp_gather_slot
|
||||
|
||||
self.addCleanup(reset_context)
|
||||
dp_flags = get_flags().dp
|
||||
saved = (
|
||||
dp_flags.use_world_group_for_gather,
|
||||
dp_flags.joiner_skip_all_gather,
|
||||
)
|
||||
|
||||
def restore():
|
||||
(
|
||||
dp_flags.use_world_group_for_gather,
|
||||
dp_flags.joiner_skip_all_gather,
|
||||
) = saved
|
||||
|
||||
self.addCleanup(restore)
|
||||
publish(
|
||||
ServerArgs(
|
||||
model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True
|
||||
),
|
||||
role="test",
|
||||
ranks=SpawnRanks(world_rank=3),
|
||||
)
|
||||
parallel = get_parallel()
|
||||
dp_flags.use_world_group_for_gather = False
|
||||
self.assertEqual(dp_gather_slot(), parallel.attn_dp_rank)
|
||||
|
||||
# After a scale-up the gather spans the expanded WORLD, and the joining
|
||||
# cohort is numbered from its offset.
|
||||
dp_flags.use_world_group_for_gather = True
|
||||
dp_flags.joiner_skip_all_gather = False
|
||||
parallel.override_permanently(ep_join_rank_offset=8)
|
||||
self.assertEqual(dp_gather_slot(), 8 + parallel.tp_rank)
|
||||
# and the topology it was read off is untouched
|
||||
self.assertEqual(parallel.attn_dp_size, 8)
|
||||
self.assertEqual(parallel.tp_size, 8)
|
||||
|
||||
def test_a_scale_up_writes_no_width(self):
|
||||
"""The identities stay unconditional because nothing overrides them:
|
||||
the scale-up only points the gather at the expanded WORLD."""
|
||||
from sglang.srt.layers.dp_attention import update_dp_attention_post_scale
|
||||
|
||||
# It also flips a process-wide gather flag; put it back, or every
|
||||
# later test in this process runs as if a scale-up had happened.
|
||||
dp_flags = get_flags().dp
|
||||
saved_gather = dp_flags.use_world_group_for_gather
|
||||
self.addCleanup(setattr, dp_flags, "use_world_group_for_gather", saved_gather)
|
||||
|
||||
self.addCleanup(reset_context)
|
||||
publish(
|
||||
ServerArgs(
|
||||
model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True
|
||||
),
|
||||
role="test",
|
||||
ranks=SpawnRanks(world_rank=3),
|
||||
)
|
||||
parallel = get_parallel()
|
||||
before = (parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size)
|
||||
update_dp_attention_post_scale(new_dp_size=16, new_dp_rank=11)
|
||||
self.assertEqual(parallel.attn_dp_size, 16)
|
||||
self.assertEqual(parallel.attn_dp_rank, 11)
|
||||
self.assertTrue(dp_flags.use_world_group_for_gather)
|
||||
self.assertEqual(
|
||||
(parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size), before
|
||||
)
|
||||
|
||||
|
||||
class TestEveryDeclaredParallelNameIsStatable(_IsolatedOverrides):
|
||||
@@ -1898,14 +1998,26 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
|
||||
def test_a_topology_is_stated_by_naming_the_width(self):
|
||||
"""Overriding a leaf does not move the quotient -- the quotient is not
|
||||
recomputed on read. Naming it is how a test states one."""
|
||||
recomputed on read. Naming it is how a caller states one, and naming
|
||||
only some of them is refused: the caller owns the arithmetic, the
|
||||
context only checks it."""
|
||||
reset_context()
|
||||
self.addCleanup(reset_context)
|
||||
publish(ServerArgs(model_path="dummy", tp_size=8), role="test")
|
||||
self.assertEqual(get_parallel().attn_tp_size, 8)
|
||||
with get_parallel().override(tp_size=2):
|
||||
self.assertEqual(get_parallel().attn_tp_size, 8)
|
||||
with get_parallel().override(attn_tp_size=4):
|
||||
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
with get_parallel().override(tp_size=2):
|
||||
pass
|
||||
self.assertIn(
|
||||
"tp_size == attn_tp_size * attn_dp_size * attn_cp_size",
|
||||
str(caught.exception),
|
||||
)
|
||||
|
||||
with get_parallel().override(tp_size=2, attn_tp_size=2, moe_tp_size=2):
|
||||
self.assertEqual(get_parallel().tp_size, 2)
|
||||
self.assertEqual(get_parallel().attn_tp_size, 2)
|
||||
with get_parallel().override(attn_tp_size=4, tp_size=4, moe_tp_size=4):
|
||||
self.assertEqual(get_parallel().attn_tp_size, 4)
|
||||
|
||||
def test_an_unstated_topology_still_fails(self):
|
||||
@@ -2127,17 +2239,11 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
)
|
||||
self.assertEqual(published, recomputed)
|
||||
|
||||
def test_initialize_model_parallel_no_longer_touches_the_bag(self):
|
||||
"""`initialize_model_parallel` used to recompute and
|
||||
permanently override the six derived widths on `get_parallel()`
|
||||
after building its groups; that call is gone. Publish a placeholder
|
||||
config (tp_size defaults to 1), then build real groups at a
|
||||
different width -- the published leaf must now stay exactly what it
|
||||
was, because nothing corrects it. This is the behavior a caller
|
||||
relies on being told about, loudly, the first time it publishes and
|
||||
builds inconsistently -- see
|
||||
`test_recomputing_from_published_leaves_matches_the_publish_bag`
|
||||
for why every real caller must not do that.
|
||||
def test_initialize_model_parallel_builds_at_the_published_widths(self):
|
||||
"""The build takes every width from the context rather than from an
|
||||
argument, so "published one width, built another" is no longer a state
|
||||
a caller can reach -- there is nothing left to translate, and nothing
|
||||
to correct afterwards either.
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
@@ -2145,11 +2251,11 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
|
||||
reset_context()
|
||||
self.addCleanup(reset_context)
|
||||
publish(ServerArgs(model_path="dummy"), role="test")
|
||||
self.assertEqual(get_parallel().attn_tp_size, 1)
|
||||
self.assertEqual(get_parallel().moe_ep_size, 1)
|
||||
|
||||
world_size = 8
|
||||
publish(ServerArgs(model_path="dummy", tp_size=world_size), role="test")
|
||||
self.assertEqual(get_parallel().attn_tp_size, world_size)
|
||||
|
||||
built_at = []
|
||||
with (
|
||||
patch.object(parallel_state, "_WORLD", None),
|
||||
patch.object(parallel_state, "_TP", None),
|
||||
@@ -2168,25 +2274,20 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
patch.object(
|
||||
parallel_state,
|
||||
"init_model_parallel_group",
|
||||
return_value=Mock(device_group=Mock()),
|
||||
side_effect=lambda group_ranks, *a, **k: (
|
||||
built_at.append(group_ranks),
|
||||
Mock(device_group=Mock()),
|
||||
)[1],
|
||||
),
|
||||
patch.object(parallel_state, "get_world_group") as mock_world_group,
|
||||
):
|
||||
mock_world_group.return_value = Mock(device_group=Mock(), local_rank=0)
|
||||
parallel_state.initialize_model_parallel(
|
||||
tensor_model_parallel_size=world_size,
|
||||
expert_model_parallel_size=world_size,
|
||||
)
|
||||
parallel_state.initialize_model_parallel()
|
||||
self.addCleanup(parallel_state.destroy_model_parallel)
|
||||
|
||||
self.assertEqual(
|
||||
get_parallel().attn_tp_size,
|
||||
1,
|
||||
"initialize_model_parallel must not touch the published leaf -- "
|
||||
"a caller that needs it corrected must publish a config that "
|
||||
"already matches the width it is about to build",
|
||||
)
|
||||
self.assertEqual(get_parallel().moe_ep_size, 1)
|
||||
# The first group built is TP, one group spanning the published width.
|
||||
self.assertEqual(built_at[0], [list(range(world_size))])
|
||||
self.assertEqual(get_parallel().attn_tp_size, world_size)
|
||||
|
||||
|
||||
class TestTheDerivedHalfIsDeclared(CustomTestCase):
|
||||
@@ -2294,6 +2395,225 @@ class TestAnEntryThatBuildsARunnerHandsOverItsPlacement(CustomTestCase):
|
||||
)
|
||||
|
||||
|
||||
class TestTheAccessorsHaveNoCallersOutsideTheirPackage(CustomTestCase):
|
||||
"""`parallel_state`'s getters are the definition, not a second spelling.
|
||||
|
||||
Business code asks `get_parallel()`; a call that goes straight to the getter
|
||||
is a read the context cannot redirect, which is what a scope needs it to be
|
||||
able to do. The package that defines them is exempt -- a read there would
|
||||
go through the context back into itself -- and so is `multimodal_gen`, which
|
||||
has its own parallel state.
|
||||
"""
|
||||
|
||||
#: Not topology. `get_self_pp_group` builds the single-rank group a draft
|
||||
#: pipeline scope installs, so there is nothing for the context to answer
|
||||
#: with until the scope has installed it.
|
||||
ALLOWED = {
|
||||
"get_self_pp_group",
|
||||
"get_default_distributed_backend",
|
||||
"get_mooncake_transfer_engine",
|
||||
}
|
||||
|
||||
def _accessors(self):
|
||||
"""Derived from the source, not listed here: a guard whose subject set
|
||||
is written by hand stops watching whatever gets added next."""
|
||||
from sglang.srt.distributed import parallel_state as parallel_state_module
|
||||
|
||||
source = _pathlib.Path(parallel_state_module.__file__).read_text().splitlines()
|
||||
return {
|
||||
line[len("def ") : line.index("(")]
|
||||
for line in source
|
||||
if line.startswith("def get_") or line.startswith("def is_")
|
||||
}
|
||||
|
||||
def _callers(self, name):
|
||||
import re
|
||||
|
||||
from sglang.srt.distributed import parallel_state as parallel_state_module
|
||||
|
||||
root = _pathlib.Path(parallel_state_module.__file__).parents[2]
|
||||
pattern = re.compile(rf"(?<![.\w]){re.escape(name)}\(")
|
||||
hits = []
|
||||
for path in root.rglob("*.py"):
|
||||
rel = path.relative_to(root).as_posix()
|
||||
if rel.startswith(("srt/distributed/", "multimodal_gen/", "test/")):
|
||||
continue
|
||||
for number, line in enumerate(path.read_text().splitlines(), 1):
|
||||
if line.lstrip().startswith(("def ", "#")):
|
||||
continue
|
||||
if pattern.search(line):
|
||||
hits.append(f"{rel}:{number}")
|
||||
return hits
|
||||
|
||||
def test_no_business_code_calls_them(self):
|
||||
offenders = {}
|
||||
for name in sorted(self._accessors() - self.ALLOWED):
|
||||
callers = self._callers(name)
|
||||
if callers:
|
||||
offenders[name] = callers
|
||||
self.assertEqual(
|
||||
offenders,
|
||||
{},
|
||||
"read these through get_parallel() instead, or say here why the "
|
||||
"context cannot answer them",
|
||||
)
|
||||
|
||||
def test_the_guard_would_notice_a_caller(self):
|
||||
"""The subject set is derived, so this checks the search finds a real
|
||||
call rather than that the list happens to be empty: `get_self_pp_group`
|
||||
is exempt and does have one caller."""
|
||||
self.assertTrue(self._callers("get_self_pp_group"))
|
||||
|
||||
|
||||
class TestTheTopologyIdentities(CustomTestCase):
|
||||
"""One set of identities, checked wherever the layout is written.
|
||||
|
||||
Each one is injected in the direction that breaks it and in the direction
|
||||
that keeps it: a guard that only ever fires is as uninformative as one that
|
||||
never does. They hold unconditionally -- a caller that states one leaf owes
|
||||
the quotients that follow from it, because a namespace describing no real
|
||||
layout is what the guard exists to refuse.
|
||||
"""
|
||||
|
||||
def _publish_square(self):
|
||||
"""tp=4 over two attention-DP replicas of two: every identity holds."""
|
||||
reset_context()
|
||||
self.addCleanup(reset_context)
|
||||
publish(
|
||||
ServerArgs(
|
||||
model_path="dummy", tp_size=4, dp_size=2, enable_dp_attention=True
|
||||
),
|
||||
role="scheduler",
|
||||
ranks=SpawnRanks(world_rank=3, dp_rank=1),
|
||||
)
|
||||
|
||||
def test_a_published_topology_is_consistent(self):
|
||||
"""The quiet direction, and the reason publish can check everything:
|
||||
it is the one write that establishes the whole layout."""
|
||||
self._publish_square()
|
||||
parallel = get_parallel()
|
||||
self.assertEqual(parallel.tp_size, 4)
|
||||
self.assertEqual(parallel.attn_dp_size, 2)
|
||||
self.assertEqual(parallel.attn_tp_size, 2)
|
||||
self.assertEqual(parallel.tp_rank, 3)
|
||||
|
||||
def test_a_width_that_does_not_factor_is_refused(self):
|
||||
self._publish_square()
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
with get_parallel().override(
|
||||
tp_size=4, attn_tp_size=3, attn_dp_size=1, attn_cp_size=1
|
||||
):
|
||||
pass
|
||||
message = str(caught.exception)
|
||||
self.assertIn("set by override", message)
|
||||
self.assertIn("attn_tp_size * attn_dp_size * attn_cp_size", message)
|
||||
self.assertIn("4 != 3 * 1 * 1", message)
|
||||
|
||||
def test_a_rank_at_its_width_is_refused(self):
|
||||
self._publish_square()
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
with get_parallel().override(tp_rank=4, tp_size=4):
|
||||
pass
|
||||
self.assertIn("0 <= tp_rank < tp_size", str(caught.exception))
|
||||
|
||||
def test_a_moe_width_that_does_not_factor_is_refused(self):
|
||||
self._publish_square()
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
with get_parallel().override(
|
||||
tp_size=4, moe_ep_size=1, moe_dp_size=1, moe_tp_size=3
|
||||
):
|
||||
pass
|
||||
message = str(caught.exception)
|
||||
self.assertIn("moe_ep_size * moe_dp_size * moe_tp_size", message)
|
||||
self.assertIn("4 != 1 * 1 * 3", message)
|
||||
|
||||
def test_a_rank_the_attention_layout_cannot_produce_is_refused(self):
|
||||
"""`tp_rank` is not free of the attention ranks: the layout derives one
|
||||
from the other, so a set that does not satisfy it places this process
|
||||
in two different seats at once."""
|
||||
self._publish_square()
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
with get_parallel().override(
|
||||
tp_rank=0,
|
||||
attn_dp_rank=1,
|
||||
attn_cp_rank=0,
|
||||
attn_tp_rank=0,
|
||||
attn_cp_size=1,
|
||||
attn_tp_size=2,
|
||||
):
|
||||
pass
|
||||
self.assertIn("attn_tp_size + attn_tp_rank", str(caught.exception))
|
||||
|
||||
def test_the_published_ranks_satisfy_the_layout(self):
|
||||
"""The quiet direction for the same identity: publish derives the
|
||||
attention ranks from `tp_rank` through that very equation, so a
|
||||
published process always sits in one seat."""
|
||||
self._publish_square()
|
||||
parallel = get_parallel()
|
||||
self.assertEqual(
|
||||
parallel.tp_rank,
|
||||
(parallel.attn_dp_rank * parallel.attn_cp_size + parallel.attn_cp_rank)
|
||||
* parallel.attn_tp_size
|
||||
+ parallel.attn_tp_rank,
|
||||
)
|
||||
|
||||
def test_a_refused_write_leaves_nothing_behind(self):
|
||||
"""The scope never opened, so the value it tried to state must not be
|
||||
readable afterwards -- a half-applied override is the state this guard
|
||||
exists to prevent."""
|
||||
self._publish_square()
|
||||
with self.assertRaises(ValueError):
|
||||
with get_parallel().override(
|
||||
tp_size=4, attn_tp_size=3, attn_dp_size=1, attn_cp_size=1
|
||||
):
|
||||
pass
|
||||
self.assertEqual(get_parallel().attn_tp_size, 2)
|
||||
|
||||
def test_a_group_built_at_another_width_is_refused(self):
|
||||
"""The other end of the same identity: what the configuration says and
|
||||
what the coordinators were actually built at, checked where the
|
||||
disagreement is still attributable to the build."""
|
||||
from sglang.srt.distributed import parallel_state
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
|
||||
self._publish_square()
|
||||
wrong = GroupCoordinator.__new__(GroupCoordinator)
|
||||
wrong.world_size = 8
|
||||
wrong.rank_in_group = 0
|
||||
with patch.object(parallel_state, "_TP", wrong):
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
_validate_parallel(get_parallel(), "group build")
|
||||
message = str(caught.exception)
|
||||
self.assertIn("set by group build", message)
|
||||
self.assertIn("tp_group.world_size == tp_size", message)
|
||||
self.assertIn("built 8, configured 4", message)
|
||||
|
||||
def test_a_group_built_at_the_configured_width_is_quiet(self):
|
||||
from sglang.srt.distributed import parallel_state
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
|
||||
self._publish_square()
|
||||
right = GroupCoordinator.__new__(GroupCoordinator)
|
||||
right.world_size = 4
|
||||
right.rank_in_group = 3
|
||||
with patch.object(parallel_state, "_TP", right):
|
||||
_validate_parallel(get_parallel(), "group build")
|
||||
|
||||
def test_a_draft_scope_states_a_consistent_topology(self):
|
||||
"""The scope narrows four names at once, so the identity applies to it
|
||||
-- and holds, which is what step lets the guard stay on."""
|
||||
from sglang.srt.distributed import parallel_state
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
|
||||
self._publish_square()
|
||||
group = GroupCoordinator.__new__(GroupCoordinator)
|
||||
group.world_size = 2
|
||||
group.rank_in_group = 1
|
||||
with patch.object(parallel_state, "_TP", group):
|
||||
with parallel_state.patch_tensor_parallel_group(group, owns_attention=True):
|
||||
self.assertEqual(get_parallel().attn_tp_size, 2)
|
||||
|
||||
|
||||
class TestWhoAnswersDuringADraftScope(CustomTestCase):
|
||||
"""A draft worker runs in one process with the target, under a scope.
|
||||
|
||||
@@ -2397,29 +2717,36 @@ class TestWhoAnswersDuringADraftScope(CustomTestCase):
|
||||
worker states it, and a caller that forgets is the bug this catches --
|
||||
`owns_attention` has no default, but a missing one is a TypeError only
|
||||
on the path that runs, and these paths need a GPU and a draft model."""
|
||||
import ast
|
||||
self.assertEqual(
|
||||
_scope_entries_that_say_nothing(_sources()),
|
||||
[],
|
||||
"these enter the draft scope without saying",
|
||||
)
|
||||
|
||||
package = _pathlib.Path(next(iter(_sglang.__path__))).resolve()
|
||||
checkout = package.parents[1]
|
||||
roots = [package] + [
|
||||
checkout / name for name in ("test",) if (checkout / name).is_dir()
|
||||
]
|
||||
missing = []
|
||||
for path in (q for root in roots for q in root.rglob("*.py")):
|
||||
try:
|
||||
tree = ast.parse(path.read_text(encoding="utf-8"))
|
||||
except (SyntaxError, UnicodeDecodeError):
|
||||
continue
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
func = node.func
|
||||
name = getattr(func, "attr", None) or getattr(func, "id", None)
|
||||
if name not in ("draft_tp_context", "patch_tensor_parallel_group"):
|
||||
continue
|
||||
if not any(kw.arg == "owns_attention" for kw in node.keywords):
|
||||
missing.append(f"{path}:{node.lineno}")
|
||||
self.assertEqual(missing, [], "these enter the draft scope without saying")
|
||||
def test_the_census_would_notice_one(self):
|
||||
"""Two ways for it to report zero and still be wrong: the matcher does
|
||||
not recognise the call, or the walk never reaches the file. A draft
|
||||
scope entered from `benchmark/` breaks the same way as one in the
|
||||
package and no suite covers it, so the roots are part of the check."""
|
||||
trees = {
|
||||
part
|
||||
for path in _sources()
|
||||
for part in ("benchmark", "examples", "scripts", "test")
|
||||
if f"/{part}/" in path.as_posix()
|
||||
}
|
||||
self.assertEqual(
|
||||
trees,
|
||||
{"benchmark", "examples", "scripts", "test"},
|
||||
"the walk misses a tree that can enter the scope",
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
probe = _pathlib.Path(tmp) / "probe.py"
|
||||
probe.write_text(
|
||||
"with self.draft_tp_context(runner.tp_group):\n pass\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
self.assertEqual(_scope_entries_that_say_nothing([probe]), [f"{probe}:1"])
|
||||
|
||||
def test_a_full_width_swap_leaves_the_attention_layout_alone(self):
|
||||
"""The other caller. A draft built outside any scope carries the
|
||||
|
||||
Reference in New Issue
Block a user