Introduce ModelRunner.ps ParallelState (#31161)

This commit is contained in:
fzyzcjy
2026-07-14 16:01:14 +08:00
committed by GitHub
parent 1dc48c2c3b
commit 725920915f
28 changed files with 202 additions and 276 deletions
+2 -6
View File
@@ -15,6 +15,7 @@ import torch
from sglang.benchmark.one_batch import TreeCacheNamespace
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner
@@ -57,14 +58,9 @@ class TestForwardSplitPrefill(CustomTestCase):
model_config=cls.model_config,
mem_fraction_static=cls.server_args.mem_fraction_static,
gpu_id=0,
tp_rank=0,
tp_size=cls.tp_size,
pp_rank=0,
pp_size=1,
ps=ParallelState.trivial(tp_size=cls.tp_size),
nccl_port=cls.port_args.nccl_port,
server_args=cls.server_args,
moe_ep_rank=0,
moe_ep_size=1,
)
cls.tokenizer = get_tokenizer(
+2 -4
View File
@@ -9,6 +9,7 @@ import torch.nn.functional as F
from transformers import AutoModel, AutoProcessor, AutoTokenizer
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest
from sglang.srt.managers.mm_utils import embed_mm_inputs, init_mm_embedding_cache
from sglang.srt.managers.schedule_batch import (
@@ -144,10 +145,7 @@ class VisionLLMLogitsBase(unittest.IsolatedAsyncioTestCase):
model_config=ModelConfig(self.model_path, model_override_args="{}"),
mem_fraction_static=0.8,
gpu_id=0,
tp_rank=0,
tp_size=1,
pp_rank=0,
pp_size=1,
ps=ParallelState.trivial(),
nccl_port=12435,
server_args=ServerArgs(
model_path=self.model_path,
@@ -9,6 +9,7 @@ from sglang.srt.disaggregation.decode import (
HiCacheRestoreResult,
)
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler
from sglang.test.ci.ci_register import register_cpu_ci
@@ -195,7 +196,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
scheduler.last_batch = None
scheduler.cur_batch_for_debug = None
scheduler.enable_overlap = False
scheduler.ps = SimpleNamespace(pp_size=1)
scheduler.ps = ParallelState.trivial()
scheduler.running_mbs = []
scheduler.waiting_queue = []
scheduler.grammar_manager = SimpleNamespace(grammar_queue=[])
@@ -18,26 +18,18 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _make_ps(**overrides) -> ParallelState:
defaults = dict(
tp_rank=0,
tp_size=8,
pp_rank=1,
pp_size=2,
dp_rank=None,
dp_size=1,
attn_tp_rank=0,
attn_tp_size=2,
attn_cp_rank=0,
attn_cp_size=2,
attn_dp_rank=1,
attn_dp_size=2,
moe_ep_rank=0,
moe_ep_size=1,
moe_dp_rank=None,
moe_dp_size=1,
gpu_id=0,
)
defaults.update(overrides)
return ParallelState(**defaults)
return ParallelState.trivial(**defaults)
def _fake_group() -> SimpleNamespace:
@@ -17,26 +17,11 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
def _make_ps(**overrides) -> ParallelState:
"""Build a ParallelState with reasonable defaults for tests; override fields via kwargs."""
defaults = dict(
tp_rank=0,
tp_size=1,
pp_rank=0,
pp_size=1,
dp_rank=None,
dp_size=1,
attn_tp_rank=0,
attn_tp_size=1,
attn_cp_rank=0,
attn_cp_size=1,
attn_dp_rank=0,
attn_dp_size=1,
moe_ep_rank=0,
moe_ep_size=1,
moe_dp_rank=None,
moe_dp_size=1,
gpu_id=0,
)
defaults.update(overrides)
return ParallelState(**defaults)
return ParallelState.trivial(**defaults)
class _FakeReq:
@@ -108,7 +93,7 @@ def _make_reporter(scheduler) -> SchedulerMetricsReporter:
enable_forward_pass_metrics=False,
)
if not hasattr(scheduler, "ps"):
scheduler.ps = types.SimpleNamespace(attn_tp_rank=0, attn_cp_rank=0)
scheduler.ps = ParallelState.trivial()
if not hasattr(scheduler, "kv_events_publisher"):
scheduler.kv_events_publisher = types.SimpleNamespace(
init_kv_events=lambda *a, **kw: None,