Bringing the parallel runtime up becomes a phase, not a side effect (#40345)

This commit is contained in:
Cheng Wan
2026-09-21 12:29:50 -07:00
committed by GitHub
parent 1d3243d05f
commit bccf691b22
9 changed files with 271 additions and 109 deletions
@@ -45,7 +45,6 @@ def test_mooncake_te_condition(server_args: ServerArgs) -> bool:
"""
Test the condition logic for using MooncakeTransferEngine.
"""
from sglang.srt.model_executor.model_runner import ModelRunner
dummy_runner = SimpleNamespace(server_args=server_args, gpu_id=0)
init_called = False
@@ -69,7 +68,11 @@ def test_mooncake_te_condition(server_args: ServerArgs) -> bool:
return_value="127.0.0.1",
),
):
ModelRunner.init_shared_mooncake_transfer_engine(dummy_runner)
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
maybe_init_shared_mooncake_transfer_engine,
)
maybe_init_shared_mooncake_transfer_engine(gpu_id=dummy_runner.gpu_id)
return init_called
+13 -2
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 import bootstrap
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.forward_context import (
@@ -22,7 +23,7 @@ from sglang.srt.model_executor.forward_context import (
set_forward_context,
)
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import publish
from sglang.srt.runtime_context import SpawnRanks, publish
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -56,10 +57,20 @@ class TestForwardSplitPrefill(CustomTestCase):
cls.port_args = PortArgs.init_new(cls.server_args)
publish(cls.server_args, role="scheduler")
publish(
cls.server_args,
role="scheduler",
ranks=SpawnRanks(world_rank=0, gpu_id=0),
)
# Load model and tokenizer
cls.model_config = ModelConfig.from_server_args(cls.server_args)
bootstrap.init_parallel_runtime(
server_args=cls.server_args,
model_config=cls.model_config,
device=cls.device,
dist_port=cls.port_args.nccl_port,
)
cls.model_runner = ModelRunner(
model_config=cls.model_config,
mem_fraction_static=cls.server_args.mem_fraction_static,
+15 -3
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 import bootstrap
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 (
@@ -19,7 +20,7 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.parser.conversation import generate_chat_conv
from sglang.srt.runtime_context import publish
from sglang.srt.runtime_context import SpawnRanks, get_device, publish
from sglang.srt.server_args import ServerArgs
from sglang.test.test_utils import download_image_with_retry
@@ -145,9 +146,20 @@ class VisionLLMLogitsBase(unittest.IsolatedAsyncioTestCase):
model_path=self.model_path,
disable_cuda_graph=True,
)
publish(server_args, role="scheduler")
publish(
server_args,
role="scheduler",
ranks=SpawnRanks(world_rank=0, gpu_id=0),
)
model_config = ModelConfig(self.model_path, model_override_args="{}")
bootstrap.init_parallel_runtime(
server_args=server_args,
model_config=model_config,
device=get_device().device,
dist_port=12435,
)
self.model_runner = ModelRunner(
model_config=ModelConfig(self.model_path, model_override_args="{}"),
model_config=model_config,
mem_fraction_static=0.8,
gpu_id=0,
nccl_port=12435,