Introduce ModelRunner.ps ParallelState (#31161)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user