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
+32 -6
View File
@@ -69,7 +69,9 @@ from sglang.srt.distributed.parallel_state import (
destroy_distributed_environment, destroy_distributed_environment,
destroy_model_parallel, destroy_model_parallel,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.entrypoints.engine import _set_envs_and_config from sglang.srt.entrypoints.engine import _set_envs_and_config
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
from sglang.srt.layers.moe import initialize_moe_config from sglang.srt.layers.moe import initialize_moe_config
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config
from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
@@ -300,16 +302,40 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size) moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
model_config = ModelConfig.from_server_args(server_args) model_config = ModelConfig.from_server_args(server_args)
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
compute_dp_attention_world_info(
server_args.enable_dp_attention,
tp_rank,
server_args.tp_size,
server_args.dp_size,
server_args.attn_cp_size,
)
)
ps = ParallelState(
tp_rank=tp_rank,
tp_size=server_args.tp_size,
pp_rank=0,
pp_size=1,
dp_rank=None,
dp_size=server_args.dp_size,
attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size,
attn_cp_rank=0,
attn_cp_size=server_args.attn_cp_size,
attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank,
moe_ep_size=server_args.ep_size,
moe_dp_rank=None,
moe_dp_size=server_args.moe_dp_size,
dcp_size=server_args.dcp_size,
gpu_id=gpu_id,
)
runner_kwargs = dict( runner_kwargs = dict(
model_config=model_config, model_config=model_config,
mem_fraction_static=server_args.mem_fraction_static, mem_fraction_static=server_args.mem_fraction_static,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, ps=ps,
tp_size=server_args.tp_size,
moe_ep_rank=moe_ep_rank,
moe_ep_size=server_args.ep_size,
pp_rank=0,
pp_size=1,
nccl_port=port_args.nccl_port, nccl_port=port_args.nccl_port,
server_args=server_args, server_args=server_args,
) )
+13 -10
View File
@@ -19,6 +19,7 @@ from sglang.srt.distributed import (
set_mscclpp_all_reduce, set_mscclpp_all_reduce,
set_torch_symm_mem_all_reduce, set_torch_symm_mem_all_reduce,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import initialize_dp_attention from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
@@ -52,20 +53,22 @@ def init_torch_distributed(
server_args: ServerArgs, server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
device: str, device: str,
gpu_id: int, ps: ParallelState,
tp_rank: int,
tp_size: int,
pp_rank: int,
pp_size: int,
dp_size: int,
attn_cp_size: int,
moe_ep_size: int,
moe_dp_size: int,
dcp_size: int,
dist_port: int, dist_port: int,
is_draft_worker: bool, is_draft_worker: bool,
local_omp_cpuid: Optional[List[int]], local_omp_cpuid: Optional[List[int]],
): ):
gpu_id = ps.gpu_id
tp_rank = ps.tp_rank
tp_size = ps.tp_size
pp_rank = ps.pp_rank
pp_size = ps.pp_size
dp_size = ps.attn_dp_size
attn_cp_size = ps.attn_cp_size
moe_ep_size = ps.moe_ep_size
moe_dp_size = ps.moe_dp_size
dcp_size = ps.dcp_size
tic = time.perf_counter() tic = time.perf_counter()
logger.info("Init torch distributed begin.") logger.info("Init torch distributed begin.")
@@ -20,4 +20,30 @@ class ParallelState:
moe_ep_size: int moe_ep_size: int
moe_dp_rank: Optional[int] moe_dp_rank: Optional[int]
moe_dp_size: int moe_dp_size: int
dcp_size: int
gpu_id: int gpu_id: int
@staticmethod
def trivial(**overrides: Optional[int]) -> "ParallelState":
kwargs: dict[str, Optional[int]] = dict(
tp_rank=0,
tp_size=1,
pp_rank=0,
pp_size=1,
dp_rank=0,
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=0,
moe_dp_size=1,
dcp_size=1,
gpu_id=0,
)
kwargs.update(overrides)
return ParallelState(**kwargs)
@@ -63,14 +63,8 @@ class MlxTpModelWorker(TpModelWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=self.server_args.mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
tp_rank=self.tp_rank, ps=self.ps,
tp_size=self.tp_size,
moe_ep_rank=self.moe_ep_rank,
moe_ep_size=self.ep_size,
pp_rank=self.pp_rank,
pp_size=self.pp_size,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
dp_rank=self.dp_rank,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
+3 -11
View File
@@ -387,6 +387,7 @@ class Scheduler(
moe_ep_size=server_args.ep_size, moe_ep_size=server_args.ep_size,
moe_dp_rank=moe_dp_rank, moe_dp_rank=moe_dp_rank,
moe_dp_size=server_args.moe_dp_size, moe_dp_size=server_args.moe_dp_size,
dcp_size=server_args.dcp_size,
gpu_id=gpu_id, gpu_id=gpu_id,
) )
@@ -747,12 +748,7 @@ class Scheduler(
worker_kwargs = dict( worker_kwargs = dict(
server_args=self.server_args, server_args=self.server_args,
gpu_id=self.ps.gpu_id, gpu_id=self.ps.gpu_id,
tp_rank=self.ps.tp_rank, ps=self.ps,
moe_ep_rank=self.ps.moe_ep_rank,
pp_rank=self.ps.pp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
moe_dp_rank=self.ps.moe_dp_rank,
dp_rank=self.ps.dp_rank,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
) )
@@ -776,13 +772,9 @@ class Scheduler(
draft_worker_kwargs = dict( draft_worker_kwargs = dict(
server_args=self.server_args, server_args=self.server_args,
gpu_id=self.ps.gpu_id, gpu_id=self.ps.gpu_id,
tp_rank=self.ps.tp_rank, ps=self.ps,
moe_ep_rank=self.ps.moe_ep_rank,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
target_worker=self.tp_worker, target_worker=self.tp_worker,
dp_rank=self.ps.dp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
moe_dp_rank=self.ps.moe_dp_rank,
) )
if self.server_args.speculative_draft_load_format is not None: if self.server_args.speculative_draft_load_format is not None:
+7 -31
View File
@@ -22,6 +22,7 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch import torch
from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
DestroyWeightsUpdateGroupReqInput, DestroyWeightsUpdateGroupReqInput,
GetWeightsByNameReqInput, GetWeightsByNameReqInput,
@@ -173,7 +174,7 @@ class BaseTpWorker(ABC):
monkey_patch_torch_reductions() monkey_patch_torch_reductions()
success, message = self.model_runner.weight_updater.update_weights_from_tensor( success, message = self.model_runner.weight_updater.update_weights_from_tensor(
named_tensors=MultiprocessingSerializer.deserialize( named_tensors=MultiprocessingSerializer.deserialize(
recv_req.serialized_named_tensors[self.tp_rank] recv_req.serialized_named_tensors[self.ps.tp_rank]
), ),
load_format=recv_req.load_format, load_format=recv_req.load_format,
) )
@@ -237,12 +238,7 @@ class TpModelWorker(BaseTpWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
moe_ep_rank: int,
pp_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
dp_rank: Optional[int],
nccl_port: int, nccl_port: int,
is_draft_worker: bool = False, is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None, req_to_token_pool: Optional[ReqToTokenPool] = None,
@@ -253,21 +249,13 @@ class TpModelWorker(BaseTpWorker):
): ):
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
self.tp_size = server_args.tp_size self.ps = ps
self.ep_size = server_args.ep_size
self.pp_size = server_args.pp_size
self.tp_rank = tp_rank
self.moe_ep_rank = moe_ep_rank
self.pp_rank = pp_rank
self.dp_rank = dp_rank
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.is_draft_worker = is_draft_worker self.is_draft_worker = is_draft_worker
self.is_multi_layer_eagle = is_multi_layer_eagle self.is_multi_layer_eagle = is_multi_layer_eagle
self.req_to_token_pool = req_to_token_pool self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.attn_cp_rank = attn_cp_rank
self.moe_dp_rank = moe_dp_rank
# Draft worker: target's resolved MemoryPoolConfig (forwarded to ModelRunner). # Draft worker: target's resolved MemoryPoolConfig (forwarded to ModelRunner).
self.memory_pool_config = memory_pool_config self.memory_pool_config = memory_pool_config
# Draft worker: target's effective context length; the draft runs at # Draft worker: target's effective context length; the draft runs at
@@ -317,7 +305,7 @@ class TpModelWorker(BaseTpWorker):
# Sync random seed across TP workers # Sync random seed across TP workers
self.random_seed = broadcast_pyobj( self.random_seed = broadcast_pyobj(
[server_args.random_seed], [server_args.random_seed],
self.tp_size * self.pp_rank + tp_rank, self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
self.world_group.cpu_group, self.world_group.cpu_group,
src=self.world_group.ranks[0], src=self.world_group.ranks[0],
)[0] )[0]
@@ -394,14 +382,8 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=self.server_args.mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
tp_rank=self.tp_rank, ps=self.ps,
tp_size=self.tp_size,
moe_ep_rank=self.moe_ep_rank,
moe_ep_size=self.ep_size,
pp_rank=self.pp_rank,
pp_size=self.pp_size,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
dp_rank=self.dp_rank,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
@@ -420,14 +402,8 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=self.server_args.mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
tp_rank=self.tp_rank, ps=self.ps,
tp_size=self.tp_size,
moe_ep_rank=self.moe_ep_rank,
moe_ep_size=self.ep_size,
pp_rank=self.pp_rank,
pp_size=self.pp_size,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
dp_rank=self.dp_rank,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
@@ -43,6 +43,7 @@ from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
prealloc_symmetric_memory_pool, prealloc_symmetric_memory_pool,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.elastic_ep.elastic_ep import ( from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager, ElasticEPStateManager,
@@ -242,17 +243,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
model_config: ModelConfig, model_config: ModelConfig,
mem_fraction_static: float, mem_fraction_static: float,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
tp_size: int,
moe_ep_rank: int,
moe_ep_size: int,
pp_rank: int,
pp_size: int,
nccl_port: int, nccl_port: int,
server_args: ServerArgs, server_args: ServerArgs,
dp_rank: Optional[int] = None,
attn_cp_rank: Optional[int] = None,
moe_dp_rank: Optional[int] = None,
is_draft_worker: bool = False, is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None, req_to_token_pool: Optional[ReqToTokenPool] = None,
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None, token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
@@ -267,22 +260,21 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.memory_pool_config = memory_pool_config self.memory_pool_config = memory_pool_config
self.device = server_args.device self.device = server_args.device
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.tp_rank = ps.tp_rank
self.tp_size = tp_size self.tp_size = ps.tp_size
self.dcp_size = server_args.dcp_size self.dcp_size = server_args.dcp_size
self.dcp_rank = self.tp_rank % self.dcp_size self.dcp_rank = ps.tp_rank % self.dcp_size
self.moe_ep_rank = moe_ep_rank self.ps = ps
self.moe_ep_size = moe_ep_size self.moe_ep_rank = ps.moe_ep_rank
self.dp_rank = dp_rank self.moe_ep_size = ps.moe_ep_size
self.attn_dp_size = ( self.dp_rank = ps.dp_rank
server_args.dp_size if server_args.enable_dp_attention else 1 self.attn_dp_size = ps.attn_dp_size
) self.pp_rank = ps.pp_rank
self.pp_rank = pp_rank self.pp_size = ps.pp_size
self.pp_size = pp_size self.attn_cp_rank = ps.attn_cp_rank
self.attn_cp_rank = attn_cp_rank self.attn_cp_size = ps.attn_cp_size
self.attn_cp_size = server_args.attn_cp_size self.moe_dp_rank = ps.moe_dp_rank
self.moe_dp_rank = moe_dp_rank self.moe_dp_size = ps.moe_dp_size
self.moe_dp_size = server_args.moe_dp_size
self.model_config = model_config self.model_config = model_config
self.dist_port = nccl_port self.dist_port = nccl_port
self.server_args = server_args self.server_args = server_args
@@ -356,7 +348,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.war_fastpath_read_done_event: Optional[torch.cuda.Event] = None self.war_fastpath_read_done_event: Optional[torch.cuda.Event] = None
# CPU offload # CPU offload
set_offloader(create_offloader_from_server_args(server_args, dp_rank=dp_rank)) set_offloader(
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
)
self._weight_checker = WeightChecker(model_runner=self) self._weight_checker = WeightChecker(model_runner=self)
@@ -813,16 +807,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
server_args=self.server_args, server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
device=self.device, device=self.device,
gpu_id=self.gpu_id, ps=self.ps,
tp_rank=self.tp_rank,
tp_size=self.tp_size,
pp_rank=self.pp_rank,
pp_size=self.pp_size,
dp_size=self.attn_dp_size,
attn_cp_size=self.attn_cp_size,
moe_ep_size=self.moe_ep_size,
moe_dp_size=self.moe_dp_size,
dcp_size=self.dcp_size,
dist_port=self.dist_port, dist_port=self.dist_port,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
local_omp_cpuid=self.local_omp_cpuid if self.device == "cpu" else None, local_omp_cpuid=self.local_omp_cpuid if self.device == "cpu" else None,
@@ -1,5 +1,6 @@
import logging import logging
import math import math
from dataclasses import replace
from typing import List, Optional from typing import List, Optional
import torch import torch
@@ -11,6 +12,7 @@ from sglang.kernels.ops.speculative.dflash import (
) )
from sglang.srt.configs.hybrid_arch import mambaish_config from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
@@ -108,21 +110,13 @@ class DFlashWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.ps = ps
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.attn_cp_rank = attn_cp_rank
self.moe_dp_rank = moe_dp_rank
self.nccl_port = nccl_port self.nccl_port = nccl_port
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
@@ -141,11 +135,7 @@ class DFlashWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker( bundle = build_draft_tp_worker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config, target_model_config=target_worker.model_runner.model_config,
algo_label="DFLASH", algo_label="DFLASH",
@@ -181,7 +171,7 @@ class DFlashWorkerV2(BaseSpecWorker):
mask_token=self._mask_token, mask_token=self._mask_token,
mask_token_id=self._mask_token_id_override, mask_token_id=self._mask_token_id_override,
) )
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s", "Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
bundle.resolved_attention_backend, bundle.resolved_attention_backend,
@@ -400,7 +390,7 @@ class DFlashWorkerV2(BaseSpecWorker):
break break
if fused_disable_reason is not None: if fused_disable_reason is not None:
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization disabled: %s", "DFLASH fused KV materialization disabled: %s",
fused_disable_reason, fused_disable_reason,
@@ -422,7 +412,7 @@ class DFlashWorkerV2(BaseSpecWorker):
max_position_hint=self.target_worker.model_runner.model_config.context_len max_position_hint=self.target_worker.model_runner.model_config.context_len
+ int(self.block_size), + int(self.block_size),
) )
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization enabled. " "DFLASH fused KV materialization enabled. "
"n_layers=%d, num_kv_heads=%d, head_dim=%d", "n_layers=%d, num_kv_heads=%d, head_dim=%d",
@@ -629,7 +619,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if resolved_id is None: if resolved_id is None:
resolved_id = tokenizer.convert_tokens_to_ids(mask_token) resolved_id = tokenizer.convert_tokens_to_ids(mask_token)
if added and self.tp_rank == 0: if added and self.ps.tp_rank == 0:
logger.info( logger.info(
"Added DFLASH mask token to tokenizer. token=%s, mask_token_id=%s, tokenizer_len=%s, model_vocab_size=%s", "Added DFLASH mask token to tokenizer. token=%s, mask_token_id=%s, tokenizer_len=%s, model_vocab_size=%s",
mask_token, mask_token,
@@ -1199,7 +1189,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if ( if (
not is_dflash_sampling_verify_available() not is_dflash_sampling_verify_available()
and not self._warned_sampling_fallback and not self._warned_sampling_fallback
and self.tp_rank == 0 and self.ps.tp_rank == 0
): ):
logger.warning( logger.warning(
"DFLASH non-greedy verification is unavailable on this build/device; " "DFLASH non-greedy verification is unavailable on this build/device; "
@@ -17,6 +17,7 @@ from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -57,11 +58,7 @@ def build_draft_tp_worker(
*, *,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_model_config: ModelConfig, target_model_config: ModelConfig,
algo_label: str, algo_label: str,
@@ -98,12 +95,7 @@ def build_draft_tp_worker(
draft_worker = TpModelWorker( draft_worker = TpModelWorker(
server_args=draft_server_args, server_args=draft_server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, ps=ps,
moe_ep_rank=moe_ep_rank,
pp_rank=0,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
dp_rank=dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
) )
@@ -1,9 +1,11 @@
import logging import logging
from contextlib import nullcontext from contextlib import nullcontext
from dataclasses import replace
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -63,21 +65,13 @@ class DSparkWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.ps = ps
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.attn_cp_rank = attn_cp_rank
self.moe_dp_rank = moe_dp_rank
self.nccl_port = nccl_port self.nccl_port = nccl_port
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
@@ -100,11 +94,7 @@ class DSparkWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker( bundle = build_draft_tp_worker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config, target_model_config=target_worker.model_runner.model_config,
algo_label="DSPARK", algo_label="DSPARK",
@@ -129,7 +119,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self.speculative_num_draft_tokens = self.verify_num_draft_tokens self.speculative_num_draft_tokens = self.verify_num_draft_tokens
self._mask_token_id = runtime_config.mask_token_id self._mask_token_id = runtime_config.mask_token_id
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
"Initialized DSpark draft runner. attention_backend=%s, model=%s, " "Initialized DSpark draft runner. attention_backend=%s, model=%s, "
"gamma=%s, verify_num_draft_tokens=%s, mask_token_id=%s, " "gamma=%s, verify_num_draft_tokens=%s, mask_token_id=%s, "
@@ -165,7 +155,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma, gamma=self.gamma,
model_runner=self.model_runner, model_runner=self.model_runner,
device=self.device, device=self.device,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
server_args=self.server_args, server_args=self.server_args,
verify_num_draft_tokens=self.verify_num_draft_tokens, verify_num_draft_tokens=self.verify_num_draft_tokens,
) )
@@ -256,7 +246,7 @@ class DSparkWorkerV2(BaseSpecWorker):
planner=self._verify_planner, planner=self._verify_planner,
gamma=self.gamma, gamma=self.gamma,
verify_num_draft_tokens=self.verify_num_draft_tokens, verify_num_draft_tokens=self.verify_num_draft_tokens,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
device=self.device, device=self.device,
simulate_acc_len=self._simulate_acc_len, simulate_acc_len=self._simulate_acc_len,
) )
@@ -332,7 +322,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma, gamma=self.gamma,
max_bs=max(self.server_args.cuda_graph_config.decode.bs), max_bs=max(self.server_args.cuda_graph_config.decode.bs),
device=self.device, device=self.device,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
confidence_fn=( confidence_fn=(
self._verify_planner.compute_confidence_tensor self._verify_planner.compute_confidence_tensor
if self._verify_planner.carries_confidence if self._verify_planner.carries_confidence
@@ -1,11 +1,13 @@
import contextlib import contextlib
import logging import logging
import time import time
from dataclasses import replace
from typing import List, Optional from typing import List, Optional
import torch import torch
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import ( from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import (
EAGLEDraftExtendNpuGraphRunner, EAGLEDraftExtendNpuGraphRunner,
@@ -127,24 +129,16 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: int,
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.ps = ps
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
self.attn_cp_rank = attn_cp_rank
self.moe_dp_rank = moe_dp_rank
# Args for easy access # Args for easy access
self.device = server_args.device self.device = server_args.device
@@ -173,12 +167,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, # spec workers don't support pipeline parallelism
pp_rank=0, # spec workers don't support pipeline parallelism ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
) )
@@ -1048,11 +1038,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -1061,7 +1047,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.tp_rank = tp_rank self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = server_args.device
self._target_worker = target_worker self._target_worker = target_worker
@@ -1079,11 +1065,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
self._draft_worker = EagleDraftWorker( self._draft_worker = EagleDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
tp_rank, ps,
dp_rank,
moe_ep_rank,
attn_cp_rank,
moe_dp_rank,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -1762,7 +1744,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
monkey_patch_torch_reductions() monkey_patch_torch_reductions()
named_tensors = MultiprocessingSerializer.deserialize( named_tensors = MultiprocessingSerializer.deserialize(
recv_req.serialized_named_tensors[self.tp_rank] recv_req.serialized_named_tensors[self.ps.tp_rank]
) )
success, message = ( success, message = (
self.draft_worker.draft_runner.weight_updater.update_weights_from_tensor( self.draft_worker.draft_runner.weight_updater.update_weights_from_tensor(
@@ -13,11 +13,13 @@
# ============================================================================== # ==============================================================================
import logging import logging
from typing import TYPE_CHECKING, List, Optional from dataclasses import replace
from typing import TYPE_CHECKING, List
import torch import torch
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import ( from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner, MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
@@ -92,20 +94,14 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: int,
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.ps = ps
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
self.draft_extend_attn_backend_list = [] self.draft_extend_attn_backend_list = []
@@ -136,12 +132,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, # spec workers don't support pipeline parallelism
pp_rank=0, # spec workers don't support pipeline parallelism ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
is_multi_layer_eagle=True, is_multi_layer_eagle=True,
@@ -665,11 +657,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -695,11 +683,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self._draft_worker = MultiLayerEagleDraftWorker( self._draft_worker = MultiLayerEagleDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
tp_rank, ps,
dp_rank,
moe_ep_rank,
attn_cp_rank,
moe_dp_rank,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -8,6 +8,7 @@ from sgl_kernel.speculative import reconstruct_indices_from_tree_mask
from sglang.kernels.ops.speculative.cache_locs import ( from sglang.kernels.ops.speculative.cache_locs import (
assign_extend_cache_locs_func as assign_extend_cache_locs_func, assign_extend_cache_locs_func as assign_extend_cache_locs_func,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -50,11 +51,7 @@ class NGRAMWorker(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -62,7 +59,7 @@ class NGRAMWorker(BaseSpecWorker):
self.enable_overlap = not server_args.disable_overlap_schedule self.enable_overlap = not server_args.disable_overlap_schedule
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
self.tp_rank = tp_rank self.tp_rank = ps.tp_rank
self.page_size = server_args.page_size self.page_size = server_args.page_size
self.draft_token_num: int = server_args.speculative_num_draft_tokens self.draft_token_num: int = server_args.speculative_num_draft_tokens
self.max_trie_depth: int = server_args.speculative_ngram_max_trie_depth self.max_trie_depth: int = server_args.speculative_ngram_max_trie_depth
@@ -1,8 +1,10 @@
import logging import logging
from dataclasses import replace
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.moe.utils import speculative_moe_backend_context from sglang.srt.layers.moe.utils import speculative_moe_backend_context
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -29,24 +31,16 @@ class StandaloneDraftWorker(EagleDraftWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: int,
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = tp_rank self.ps = ps
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
self.attn_cp_rank = attn_cp_rank
self.moe_dp_rank = moe_dp_rank
# Args for easy access # Args for easy access
self.device = server_args.device self.device = server_args.device
@@ -74,12 +68,8 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, # spec workers don't support pipeline parallelism
pp_rank=0, # spec workers don't support pipeline parallelism ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
) )
@@ -146,11 +136,7 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -177,11 +163,7 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self._draft_worker = StandaloneDraftWorker( self._draft_worker = StandaloneDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
tp_rank, ps,
dp_rank,
moe_ep_rank,
attn_cp_rank,
moe_dp_rank,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -7,6 +7,7 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
@@ -286,7 +287,9 @@ class TinyModelConfig:
num_key_value_heads=num_kv_heads, num_key_value_heads=num_kv_heads,
head_dim=head_dim, head_dim=head_dim,
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int: def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0 assert self.num_attention_heads % tp_size == 0
@@ -322,6 +325,7 @@ class MockModelRunner(ModelRunner):
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self.is_draft_worker = False self.is_draft_worker = False
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
# The runner lifecycle warms up kernels in capture() / first execute() # The runner lifecycle warms up kernels in capture() / first execute()
@@ -5,6 +5,7 @@ from typing import Any
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool
@@ -258,7 +259,9 @@ class TinyDSAModelConfig:
index_topk=index_topk, index_topk=index_topk,
num_hidden_layers=1, num_hidden_layers=1,
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
class DSAMockModelRunner(ModelRunner): class DSAMockModelRunner(ModelRunner):
@@ -308,6 +311,7 @@ class DSAMockModelRunner(ModelRunner):
self._kernel_warmed_up = True self._kernel_warmed_up = True
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args( self._server_args_override = get_context().override_server_args(
attention_backend=case.backend, attention_backend=case.backend,
chunked_prefill_size=-1, chunked_prefill_size=-1,
@@ -19,6 +19,7 @@ from typing import Any
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.dsv4.quant_k_cache import ( from sglang.srt.layers.attention.dsv4.quant_k_cache import (
@@ -282,7 +283,9 @@ class TinyDSV4ModelConfig:
num_hidden_layers=len(compression_ratios), num_hidden_layers=len(compression_ratios),
compress_ratios=list(compression_ratios), compress_ratios=list(compression_ratios),
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
class MockDSV4ModelRunner: class MockDSV4ModelRunner:
@@ -332,6 +335,7 @@ class MockDSV4ModelRunner:
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args( self._server_args_override = get_context().override_server_args(
attention_backend=case.backend, attention_backend=case.backend,
chunked_prefill_size=-1, chunked_prefill_size=-1,
@@ -4,6 +4,7 @@ from types import SimpleNamespace
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention import ( from sglang.srt.layers.attention import (
dual_chunk_flashattention_backend as _dual_chunk_backend, dual_chunk_flashattention_backend as _dual_chunk_backend,
) )
@@ -286,7 +287,9 @@ class TinyDualChunkModelConfig:
dual_chunk_attention_config or DUAL_CHUNK_CONFIG dual_chunk_attention_config or DUAL_CHUNK_CONFIG
), ),
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int: def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0 assert self.num_attention_heads % tp_size == 0
@@ -323,6 +326,7 @@ class DualChunkMockModelRunner(ModelRunner):
self._kernel_warmed_up = True self._kernel_warmed_up = True
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args( self._server_args_override = get_context().override_server_args(
attention_backend=case.backend, attention_backend=case.backend,
chunked_prefill_size=-1, chunked_prefill_size=-1,
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -182,7 +183,9 @@ class TinyGDNModelConfig:
self.attention_chunk_size = None self.attention_chunk_size = None
self.sliding_window_size = None self.sliding_window_size = None
self.hf_config = SimpleNamespace(architectures=["TinyGDNForCausalLM"]) self.hf_config = SimpleNamespace(architectures=["TinyGDNForCausalLM"])
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int: def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0 assert self.num_key_value_heads % tp_size == 0
@@ -210,6 +213,7 @@ class MockGDNModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateDType, Mamba2StateDType,
) )
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -188,7 +189,9 @@ class TinyKDAModelConfig:
self.attention_chunk_size = None self.attention_chunk_size = None
self.sliding_window_size = None self.sliding_window_size = None
self.hf_config = SimpleNamespace(architectures=["TinyKDAForCausalLM"]) self.hf_config = SimpleNamespace(architectures=["TinyKDAForCausalLM"])
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int: def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0 assert self.num_key_value_heads % tp_size == 0
@@ -216,6 +219,7 @@ class MockKDAModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.linear.lightning_backend import ( from sglang.srt.layers.attention.linear.lightning_backend import (
LightningAttentionBackend, LightningAttentionBackend,
@@ -198,7 +199,9 @@ class TinyLightningModelConfig:
num_hidden_layers=num_hidden_layers, num_hidden_layers=num_hidden_layers,
linear_backend=linear_backend, linear_backend=linear_backend,
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int: def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0 assert self.num_key_value_heads % tp_size == 0
@@ -224,6 +227,7 @@ class MockLightningModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
@@ -18,12 +18,14 @@ _parallel_override.__enter__()
# Provide a stub group with world_size=1 so use_symmetric_memory short-circuits. # Provide a stub group with world_size=1 so use_symmetric_memory short-circuits.
_linear_mod.get_tp_group = lambda: SimpleNamespace(world_size=1) _linear_mod.get_tp_group = lambda: SimpleNamespace(world_size=1)
from sglang.srt.configs.falcon_h1 import FalconH1Config # noqa: E402
from sglang.srt.configs.mamba_utils import ( # noqa: E402 from sglang.srt.configs.mamba_utils import ( # noqa: E402
Mamba2CacheParams, Mamba2CacheParams,
Mamba2StateDType, Mamba2StateDType,
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch # noqa: E402 from sglang.srt.configs.model_config import AttentionArch # noqa: E402
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ( # noqa: E402 from sglang.srt.layers.attention.attention_registry import ( # noqa: E402
ATTENTION_BACKENDS, ATTENTION_BACKENDS,
) )
@@ -278,14 +280,14 @@ class TinyMamba2ModelConfig:
self.is_local_attention_model = False self.is_local_attention_model = False
self.attention_chunk_size = None self.attention_chunk_size = None
self.sliding_window_size = None self.sliding_window_size = None
# Mamba2AttnBackend reads mamba2_config.mamba_chunk_size; expose it # Mamba2AttnBackend reads mamba2_config(model_config).mamba_chunk_size; expose it
# through a SimpleNamespace-as-hf_config so runner.mamba2_config returns self.hf_config = FalconH1Config(
# something non-None with the expected attribute.
self.hf_config = SimpleNamespace(
architectures=["TinyMamba2ForCausalLM"], architectures=["TinyMamba2ForCausalLM"],
mamba_chunk_size=case.mamba_chunk_size, mamba_chunk_size=case.mamba_chunk_size,
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int: def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0 assert self.num_key_value_heads % tp_size == 0
@@ -310,6 +312,7 @@ class MockMamba2ModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
@@ -7,6 +7,7 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool
@@ -191,7 +192,9 @@ class TinyMLAModelConfig:
qk_rope_head_dim=qk_rope_head_dim, qk_rope_head_dim=qk_rope_head_dim,
v_head_dim=kv_lora_rank, v_head_dim=kv_lora_rank,
) )
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int: def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0 assert self.num_attention_heads % tp_size == 0
@@ -233,6 +236,7 @@ class MockMLAModelRunner(ModelRunner):
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
speculative_num_draft_tokens = ( speculative_num_draft_tokens = (
max(case.input_lens) max(case.input_lens)
if case.forward_mode.is_target_verify() if case.forward_mode.is_target_verify()
+2 -6
View File
@@ -15,6 +15,7 @@ import torch
from sglang.benchmark.one_batch import TreeCacheNamespace from sglang.benchmark.one_batch import TreeCacheNamespace
from sglang.srt.configs.model_config import ModelConfig 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.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@@ -57,14 +58,9 @@ class TestForwardSplitPrefill(CustomTestCase):
model_config=cls.model_config, model_config=cls.model_config,
mem_fraction_static=cls.server_args.mem_fraction_static, mem_fraction_static=cls.server_args.mem_fraction_static,
gpu_id=0, gpu_id=0,
tp_rank=0, ps=ParallelState.trivial(tp_size=cls.tp_size),
tp_size=cls.tp_size,
pp_rank=0,
pp_size=1,
nccl_port=cls.port_args.nccl_port, nccl_port=cls.port_args.nccl_port,
server_args=cls.server_args, server_args=cls.server_args,
moe_ep_rank=0,
moe_ep_size=1,
) )
cls.tokenizer = get_tokenizer( 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 transformers import AutoModel, AutoProcessor, AutoTokenizer
from sglang.srt.configs.model_config import ModelConfig 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.entrypoints.openai.protocol import ChatCompletionRequest
from sglang.srt.managers.mm_utils import embed_mm_inputs, init_mm_embedding_cache from sglang.srt.managers.mm_utils import embed_mm_inputs, init_mm_embedding_cache
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
@@ -144,10 +145,7 @@ class VisionLLMLogitsBase(unittest.IsolatedAsyncioTestCase):
model_config=ModelConfig(self.model_path, model_override_args="{}"), model_config=ModelConfig(self.model_path, model_override_args="{}"),
mem_fraction_static=0.8, mem_fraction_static=0.8,
gpu_id=0, gpu_id=0,
tp_rank=0, ps=ParallelState.trivial(),
tp_size=1,
pp_rank=0,
pp_size=1,
nccl_port=12435, nccl_port=12435,
server_args=ServerArgs( server_args=ServerArgs(
model_path=self.model_path, model_path=self.model_path,
@@ -9,6 +9,7 @@ from sglang.srt.disaggregation.decode import (
HiCacheRestoreResult, HiCacheRestoreResult,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode 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.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -195,7 +196,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
scheduler.last_batch = None scheduler.last_batch = None
scheduler.cur_batch_for_debug = None scheduler.cur_batch_for_debug = None
scheduler.enable_overlap = False scheduler.enable_overlap = False
scheduler.ps = SimpleNamespace(pp_size=1) scheduler.ps = ParallelState.trivial()
scheduler.running_mbs = [] scheduler.running_mbs = []
scheduler.waiting_queue = [] scheduler.waiting_queue = []
scheduler.grammar_manager = SimpleNamespace(grammar_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: def _make_ps(**overrides) -> ParallelState:
defaults = dict( defaults = dict(
tp_rank=0,
tp_size=8, tp_size=8,
pp_rank=1, pp_rank=1,
pp_size=2, pp_size=2,
dp_rank=None, dp_rank=None,
dp_size=1,
attn_tp_rank=0,
attn_tp_size=2, attn_tp_size=2,
attn_cp_rank=0,
attn_cp_size=2, attn_cp_size=2,
attn_dp_rank=1, attn_dp_rank=1,
attn_dp_size=2, attn_dp_size=2,
moe_ep_rank=0,
moe_ep_size=1,
moe_dp_rank=None, moe_dp_rank=None,
moe_dp_size=1,
gpu_id=0,
) )
defaults.update(overrides) defaults.update(overrides)
return ParallelState(**defaults) return ParallelState.trivial(**defaults)
def _fake_group() -> SimpleNamespace: def _fake_group() -> SimpleNamespace:
@@ -17,26 +17,11 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
def _make_ps(**overrides) -> ParallelState: def _make_ps(**overrides) -> ParallelState:
"""Build a ParallelState with reasonable defaults for tests; override fields via kwargs.""" """Build a ParallelState with reasonable defaults for tests; override fields via kwargs."""
defaults = dict( defaults = dict(
tp_rank=0,
tp_size=1,
pp_rank=0,
pp_size=1,
dp_rank=None, 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_rank=None,
moe_dp_size=1,
gpu_id=0,
) )
defaults.update(overrides) defaults.update(overrides)
return ParallelState(**defaults) return ParallelState.trivial(**defaults)
class _FakeReq: class _FakeReq:
@@ -108,7 +93,7 @@ def _make_reporter(scheduler) -> SchedulerMetricsReporter:
enable_forward_pass_metrics=False, enable_forward_pass_metrics=False,
) )
if not hasattr(scheduler, "ps"): 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"): if not hasattr(scheduler, "kv_events_publisher"):
scheduler.kv_events_publisher = types.SimpleNamespace( scheduler.kv_events_publisher = types.SimpleNamespace(
init_kv_events=lambda *a, **kw: None, init_kv_events=lambda *a, **kw: None,