Retire the per-runner parallel record (#40343)
This commit is contained in:
@@ -75,7 +75,6 @@ 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.hardware_backend.mlx.runtime import use_mlx
|
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
|
||||||
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
|
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
|
||||||
@@ -329,32 +328,10 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
|||||||
cfg.attn_cp_size,
|
cfg.attn_cp_size,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
ps = ParallelState(
|
|
||||||
tp_rank=tp_rank,
|
|
||||||
tp_size=cfg.tp_size,
|
|
||||||
pp_rank=0,
|
|
||||||
pp_size=1,
|
|
||||||
dp_rank=None,
|
|
||||||
dp_size=cfg.dp_size,
|
|
||||||
attn_tp_rank=attn_tp_rank,
|
|
||||||
attn_tp_size=attn_tp_size,
|
|
||||||
attn_cp_rank=0,
|
|
||||||
attn_cp_size=cfg.attn_cp_size,
|
|
||||||
attn_dcp_rank=tp_rank % cfg.dcp_size,
|
|
||||||
attn_dcp_size=cfg.dcp_size,
|
|
||||||
attn_dp_rank=attn_dp_rank,
|
|
||||||
attn_dp_size=attn_dp_size,
|
|
||||||
moe_ep_rank=moe_ep_rank,
|
|
||||||
moe_ep_size=cfg.ep_size,
|
|
||||||
moe_dp_rank=None,
|
|
||||||
moe_dp_size=cfg.moe_dp_size,
|
|
||||||
gpu_id=gpu_id,
|
|
||||||
)
|
|
||||||
runner_kwargs = dict(
|
runner_kwargs = dict(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
mem_fraction_static=cfg.mem_fraction_static,
|
mem_fraction_static=cfg.mem_fraction_static,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
ps=ps,
|
|
||||||
nccl_port=port_args.nccl_port,
|
nccl_port=port_args.nccl_port,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
@@ -571,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
|||||||
model_runner=model_runner,
|
model_runner=model_runner,
|
||||||
dp_size=get_parallel().dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
attn_tp_size=get_parallel().attn_tp_size,
|
attn_tp_size=get_parallel().attn_tp_size,
|
||||||
attn_cp_size=model_runner.ps.attn_cp_size,
|
attn_cp_size=model_runner.attn_cp_size,
|
||||||
tp_group=model_runner.tp_group,
|
tp_group=model_runner.tp_group,
|
||||||
get_idle_batch=None,
|
get_idle_batch=None,
|
||||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||||
@@ -709,13 +686,12 @@ def correctness_test(
|
|||||||
gpu_id,
|
gpu_id,
|
||||||
tp_rank,
|
tp_rank,
|
||||||
):
|
):
|
||||||
# With the placement this process was spawned with, so a rank read here
|
|
||||||
# does not need a process group -- the same bundle the runner is handed.
|
|
||||||
publish(
|
publish(
|
||||||
server_args,
|
server_args,
|
||||||
role="scheduler",
|
role="scheduler",
|
||||||
ranks=SpawnRanks(
|
ranks=SpawnRanks(
|
||||||
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0)
|
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0),
|
||||||
|
gpu_id=gpu_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -926,7 +902,8 @@ def latency_test(
|
|||||||
server_args,
|
server_args,
|
||||||
role="scheduler",
|
role="scheduler",
|
||||||
ranks=SpawnRanks(
|
ranks=SpawnRanks(
|
||||||
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0)
|
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0),
|
||||||
|
gpu_id=gpu_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
initialize_moe_config()
|
initialize_moe_config()
|
||||||
|
|||||||
@@ -92,6 +92,9 @@ class Arg(msgspec.Struct, frozen=True):
|
|||||||
fallback: Any = None
|
fallback: Any = None
|
||||||
|
|
||||||
|
|
||||||
|
_NO_DEFAULT = object()
|
||||||
|
|
||||||
|
|
||||||
class Derived(msgspec.Struct, frozen=True):
|
class Derived(msgspec.Struct, frozen=True):
|
||||||
"""Metadata for a field the configuration implies, not one anyone types.
|
"""Metadata for a field the configuration implies, not one anyone types.
|
||||||
|
|
||||||
@@ -122,6 +125,12 @@ class Derived(msgspec.Struct, frozen=True):
|
|||||||
|
|
||||||
doc: str = ""
|
doc: str = ""
|
||||||
fn: str = ""
|
fn: str = ""
|
||||||
|
# For a declaration with no ``fn`` whose absence is itself an answer:
|
||||||
|
# ``gpu_id`` is ``None`` in a process that runs on no device, and a reader
|
||||||
|
# wants that rather than an error. A rank has no such value -- the wrong
|
||||||
|
# one is a hang in a collective -- so it carries no default and a read
|
||||||
|
# before the write says so.
|
||||||
|
default: Any = _NO_DEFAULT
|
||||||
|
|
||||||
|
|
||||||
class NS(msgspec.Struct, frozen=True):
|
class NS(msgspec.Struct, frozen=True):
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from typing import (
|
|||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
from sglang.srt.arg_groups.arg_utils import A
|
from sglang.srt.arg_groups.arg_utils import A, Derived
|
||||||
|
|
||||||
|
|
||||||
class Device(msgspec.Struct):
|
class Device(msgspec.Struct):
|
||||||
@@ -40,6 +40,17 @@ class Device(msgspec.Struct):
|
|||||||
int,
|
int,
|
||||||
"The delta between consecutive GPU IDs that are used. For example, setting it to 2 will use GPU 0,2,4,...",
|
"The delta between consecutive GPU IDs that are used. For example, setting it to 2 will use GPU 0,2,4,...",
|
||||||
] = 1
|
] = 1
|
||||||
|
gpu_id = Derived(
|
||||||
|
doc=(
|
||||||
|
"Which device this process runs on. Nobody types it and nothing "
|
||||||
|
"computes it from the configuration: the parent decides -- "
|
||||||
|
"reindexing narrows the visible devices before the spawn, and Ray "
|
||||||
|
"allocates from its own pool -- so the entry states it in the "
|
||||||
|
"bundle it hands `publish`. `None` is an answer rather than an "
|
||||||
|
"absence: most roles run on no device at all."
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
)
|
||||||
random_seed: A[Optional[int], "The random seed."] = None
|
random_seed: A[Optional[int], "The random seed."] = None
|
||||||
mlx_enable_sampling: A[
|
mlx_enable_sampling: A[
|
||||||
bool,
|
bool,
|
||||||
|
|||||||
@@ -111,6 +111,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
|
|||||||
scheduler_stage_method,
|
scheduler_stage_method,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_memory,
|
get_memory,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -439,7 +440,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
if get_disagg().disaggregation_enable_kv_checksum:
|
if get_disagg().disaggregation_enable_kv_checksum:
|
||||||
kv_args = self.kv_manager.kv_args
|
kv_args = self.kv_manager.kv_args
|
||||||
self.scheduler.kv_checksum_computer = KvChecksumComputer(
|
self.scheduler.kv_checksum_computer = KvChecksumComputer(
|
||||||
device=torch.device(f"cuda:{self.scheduler.ps.gpu_id}"),
|
device=torch.device(f"cuda:{get_device().gpu_id}"),
|
||||||
kv_data_ptrs=kv_args.kv_data_ptrs,
|
kv_data_ptrs=kv_args.kv_data_ptrs,
|
||||||
kv_item_lens=kv_args.kv_item_lens,
|
kv_item_lens=kv_args.kv_item_lens,
|
||||||
state_data_ptrs=kv_args.state_data_ptrs,
|
state_data_ptrs=kv_args.state_data_ptrs,
|
||||||
@@ -565,7 +566,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
kv_args.engine_rank = self.tp_rank % (attn_tp_size)
|
kv_args.engine_rank = self.tp_rank % (attn_tp_size)
|
||||||
|
|
||||||
kv_args.pp_rank = self.pp_rank
|
kv_args.pp_rank = self.pp_rank
|
||||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
kv_args.system_dp_rank = get_parallel().dp_rank
|
||||||
kv_args.kv_cache_dtype_str = (
|
kv_args.kv_cache_dtype_str = (
|
||||||
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
|
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
|
||||||
)
|
)
|
||||||
@@ -633,7 +634,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
kv_args.gpu_id = get_device().gpu_id
|
||||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||||
kv_manager = kv_manager_class(
|
kv_manager = kv_manager_class(
|
||||||
kv_args,
|
kv_args,
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from sglang.srt.managers.schedule_batch import Modality, Req
|
|||||||
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
||||||
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
|
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_mm,
|
get_mm,
|
||||||
@@ -1921,7 +1922,7 @@ class MMReceiverBase(ABC):
|
|||||||
self.scheduler_embedding_port,
|
self.scheduler_embedding_port,
|
||||||
)
|
)
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self.gpu_id = scheduler.ps.gpu_id if scheduler is not None else 0
|
self.gpu_id = get_device().gpu_id if scheduler is not None else 0
|
||||||
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
|
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
|
||||||
self.embedding_pool = None
|
self.embedding_pool = None
|
||||||
|
|
||||||
|
|||||||
@@ -87,6 +87,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
|
|||||||
scheduler_stage_method,
|
scheduler_stage_method,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
@@ -212,7 +213,7 @@ class PrefillBootstrapQueue:
|
|||||||
if get_disagg().disaggregation_enable_kv_checksum:
|
if get_disagg().disaggregation_enable_kv_checksum:
|
||||||
kv_args = self.kv_manager.kv_args
|
kv_args = self.kv_manager.kv_args
|
||||||
self.scheduler.kv_checksum_computer = KvChecksumComputer(
|
self.scheduler.kv_checksum_computer = KvChecksumComputer(
|
||||||
device=torch.device(f"cuda:{self.scheduler.ps.gpu_id}"),
|
device=torch.device(f"cuda:{get_device().gpu_id}"),
|
||||||
kv_data_ptrs=kv_args.kv_data_ptrs,
|
kv_data_ptrs=kv_args.kv_data_ptrs,
|
||||||
kv_item_lens=kv_args.kv_item_lens,
|
kv_item_lens=kv_args.kv_item_lens,
|
||||||
state_data_ptrs=kv_args.state_data_ptrs,
|
state_data_ptrs=kv_args.state_data_ptrs,
|
||||||
@@ -226,7 +227,7 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args = kv_args_class()
|
kv_args = kv_args_class()
|
||||||
kv_args.engine_rank = self.tp_rank
|
kv_args.engine_rank = self.tp_rank
|
||||||
kv_args.pp_rank = self.pp_rank
|
kv_args.pp_rank = self.pp_rank
|
||||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
kv_args.system_dp_rank = get_parallel().dp_rank
|
||||||
kv_args.rust_http_port = (
|
kv_args.rust_http_port = (
|
||||||
self.scheduler.rust_server.http_port
|
self.scheduler.rust_server.http_port
|
||||||
if self.scheduler.rust_server is not None
|
if self.scheduler.rust_server is not None
|
||||||
@@ -303,7 +304,7 @@ class PrefillBootstrapQueue:
|
|||||||
self.metadata_buffers.get_buf_infos()
|
self.metadata_buffers.get_buf_infos()
|
||||||
)
|
)
|
||||||
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
kv_args.gpu_id = get_device().gpu_id
|
||||||
|
|
||||||
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
|
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
|
||||||
setup_state_kv_args(
|
setup_state_kv_args(
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Callable, Optional, Tuple
|
|||||||
from sglang.srt.disaggregation.common.conn import CommonKVReceiver
|
from sglang.srt.disaggregation.common.conn import CommonKVReceiver
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.managers.io_struct import PdRoleSwitchReqInput, PdRoleSwitchReqOutput
|
from sglang.srt.managers.io_struct import PdRoleSwitchReqInput, PdRoleSwitchReqOutput
|
||||||
from sglang.srt.runtime_context import get_context, get_disagg
|
from sglang.srt.runtime_context import get_context, get_device, get_disagg
|
||||||
from sglang.srt.utils import get_available_gpu_memory
|
from sglang.srt.utils import get_available_gpu_memory
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -95,7 +95,7 @@ def handle_pd_role_switch(
|
|||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
available_graph_gb = get_available_gpu_memory(
|
available_graph_gb = get_available_gpu_memory(
|
||||||
scheduler.device, scheduler.ps.gpu_id
|
scheduler.device, get_device().gpu_id
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return _fail(
|
return _fail(
|
||||||
|
|||||||
@@ -22,12 +22,12 @@ from sglang.srt.distributed.gated_launch import maybe_wait_for_gated_launch
|
|||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
_tag_groups_for_flashinfer_allreduce_only,
|
_tag_groups_for_flashinfer_allreduce_only,
|
||||||
)
|
)
|
||||||
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.layers.layernorm_sp import initialize_layernorm_sp
|
from sglang.srt.layers.layernorm_sp import initialize_layernorm_sp
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -63,17 +63,17 @@ def init_torch_distributed(
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
device: str,
|
device: str,
|
||||||
ps: ParallelState,
|
|
||||||
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]],
|
||||||
):
|
):
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
logger.info("Init torch distributed begin.")
|
logger.info("Init torch distributed begin.")
|
||||||
|
parallel = get_parallel()
|
||||||
|
|
||||||
backend = _resolve_backend(device=device)
|
backend = _resolve_backend(device=device)
|
||||||
|
|
||||||
before_avail_memory = get_available_gpu_memory(device, ps.gpu_id)
|
before_avail_memory = get_available_gpu_memory(device, get_device().gpu_id)
|
||||||
if not get_parallel().enable_p2p_check:
|
if not get_parallel().enable_p2p_check:
|
||||||
monkey_patch_p2p_access_check()
|
monkey_patch_p2p_access_check()
|
||||||
|
|
||||||
@@ -83,8 +83,8 @@ def init_torch_distributed(
|
|||||||
if not is_draft_worker:
|
if not is_draft_worker:
|
||||||
if device == "cpu":
|
if device == "cpu":
|
||||||
_init_cpu_threads_env(
|
_init_cpu_threads_env(
|
||||||
tp_size=ps.tp_size,
|
tp_size=parallel.tp_size,
|
||||||
tp_rank=ps.tp_rank,
|
tp_rank=parallel.tp_rank,
|
||||||
local_omp_cpuid=local_omp_cpuid,
|
local_omp_cpuid=local_omp_cpuid,
|
||||||
dist_init_method=dist_init_method,
|
dist_init_method=dist_init_method,
|
||||||
)
|
)
|
||||||
@@ -96,16 +96,18 @@ def init_torch_distributed(
|
|||||||
dist_init_method=dist_init_method,
|
dist_init_method=dist_init_method,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
gpu_id=ps.gpu_id,
|
gpu_id=get_device().gpu_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
||||||
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
||||||
if get_exec().comm.pre_warm_nccl and (
|
if get_exec().comm.pre_warm_nccl and (
|
||||||
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
|
parallel.tp_size > 1 or parallel.pp_size > 1 or parallel.moe_ep_size > 1
|
||||||
):
|
):
|
||||||
_prewarm_nccl(
|
_prewarm_nccl(
|
||||||
tp_size=ps.tp_size, pp_size=ps.pp_size, moe_ep_size=ps.moe_ep_size
|
tp_size=parallel.tp_size,
|
||||||
|
pp_size=parallel.pp_size,
|
||||||
|
moe_ep_size=parallel.moe_ep_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
# CUDA graph capture enables the PyNCCL communicator for TP LM-head
|
# CUDA graph capture enables the PyNCCL communicator for TP LM-head
|
||||||
@@ -115,7 +117,7 @@ def init_torch_distributed(
|
|||||||
if (
|
if (
|
||||||
device == "cuda"
|
device == "cuda"
|
||||||
and get_parallel().enable_tp_lm_head_all_to_all
|
and get_parallel().enable_tp_lm_head_all_to_all
|
||||||
and ps.tp_size > 1
|
and parallel.tp_size > 1
|
||||||
):
|
):
|
||||||
_prewarm_tp_lm_head_all_to_all()
|
_prewarm_tp_lm_head_all_to_all()
|
||||||
|
|
||||||
@@ -127,13 +129,13 @@ def init_torch_distributed(
|
|||||||
# including them in this WORLD reduction would deadlock on absent peers.
|
# including them in this WORLD reduction would deadlock on absent peers.
|
||||||
pre_model_load_memory = get_available_gpu_memory(
|
pre_model_load_memory = get_available_gpu_memory(
|
||||||
device,
|
device,
|
||||||
ps.gpu_id,
|
get_device().gpu_id,
|
||||||
distributed=get_world_group().world_size > 1 and not is_draft_worker,
|
distributed=get_world_group().world_size > 1 and not is_draft_worker,
|
||||||
cpu_group=get_world_group().cpu_group,
|
cpu_group=get_world_group().cpu_group,
|
||||||
)
|
)
|
||||||
# Check memory for tensor parallelism
|
# Check memory for tensor parallelism
|
||||||
local_gpu_memory = get_available_gpu_memory(device, ps.gpu_id)
|
local_gpu_memory = get_available_gpu_memory(device, get_device().gpu_id)
|
||||||
if ps.tp_size > 1 and not is_draft_worker:
|
if parallel.tp_size > 1 and not is_draft_worker:
|
||||||
_check_tp_memory_balance(
|
_check_tp_memory_balance(
|
||||||
pre_model_load_memory=pre_model_load_memory,
|
pre_model_load_memory=pre_model_load_memory,
|
||||||
local_gpu_memory=local_gpu_memory,
|
local_gpu_memory=local_gpu_memory,
|
||||||
|
|||||||
@@ -1,51 +0,0 @@
|
|||||||
from dataclasses import dataclass
|
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
||||||
class ParallelState:
|
|
||||||
tp_rank: int
|
|
||||||
tp_size: int
|
|
||||||
pp_rank: int
|
|
||||||
pp_size: int
|
|
||||||
dp_rank: Optional[int]
|
|
||||||
dp_size: int
|
|
||||||
attn_tp_rank: int
|
|
||||||
attn_tp_size: int
|
|
||||||
attn_cp_rank: int
|
|
||||||
attn_cp_size: int
|
|
||||||
attn_dcp_rank: int
|
|
||||||
attn_dcp_size: int
|
|
||||||
attn_dp_rank: int
|
|
||||||
attn_dp_size: int
|
|
||||||
moe_ep_rank: int
|
|
||||||
moe_ep_size: int
|
|
||||||
moe_dp_rank: Optional[int]
|
|
||||||
moe_dp_size: 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_dcp_rank=0,
|
|
||||||
attn_dcp_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,
|
|
||||||
gpu_id=0,
|
|
||||||
)
|
|
||||||
kwargs.update(overrides)
|
|
||||||
return ParallelState(**kwargs)
|
|
||||||
@@ -31,7 +31,6 @@ class EPLBManager:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
ps: Any,
|
|
||||||
get_model: Callable[[], nn.Module],
|
get_model: Callable[[], nn.Module],
|
||||||
get_expert_location_updater: Callable[[], ExpertLocationUpdater],
|
get_expert_location_updater: Callable[[], ExpertLocationUpdater],
|
||||||
get_expert_backup_client: Callable[[], Any],
|
get_expert_backup_client: Callable[[], Any],
|
||||||
@@ -42,7 +41,6 @@ class EPLBManager:
|
|||||||
# constructed (model load, expert_backup_client, weight_updater), so
|
# constructed (model load, expert_backup_client, weight_updater), so
|
||||||
# they are read through getters at rebalance time, not captured here.
|
# they are read through getters at rebalance time, not captured here.
|
||||||
self._model_config = model_config
|
self._model_config = model_config
|
||||||
self._ps = ps
|
|
||||||
self._get_model = get_model
|
self._get_model = get_model
|
||||||
self._get_expert_location_updater = get_expert_location_updater
|
self._get_expert_location_updater = get_expert_location_updater
|
||||||
self._get_expert_backup_client = get_expert_backup_client
|
self._get_expert_backup_client = get_expert_backup_client
|
||||||
@@ -163,7 +161,7 @@ class EPLBManager:
|
|||||||
tp_rank=(
|
tp_rank=(
|
||||||
self._elastic_global_rank()
|
self._elastic_global_rank()
|
||||||
if is_post_scale_rebalance
|
if is_post_scale_rebalance
|
||||||
else self._ps.tp_rank
|
else get_parallel().tp_rank
|
||||||
),
|
),
|
||||||
use_flat_topology=is_post_scale_rebalance,
|
use_flat_topology=is_post_scale_rebalance,
|
||||||
expert_backup_client=self._get_expert_backup_client(),
|
expert_backup_client=self._get_expert_backup_client(),
|
||||||
@@ -223,7 +221,7 @@ class EPLBManager:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _elastic_global_rank(self) -> int:
|
def _elastic_global_rank(self) -> int:
|
||||||
return self._ps.tp_rank + get_parallel().ep_join_rank_offset
|
return get_parallel().tp_rank + get_parallel().ep_join_rank_offset
|
||||||
|
|
||||||
def _check_rebalance_needed(self, average_utilization_rate_over_window):
|
def _check_rebalance_needed(self, average_utilization_rate_over_window):
|
||||||
if average_utilization_rate_over_window is None:
|
if average_utilization_rate_over_window is None:
|
||||||
@@ -248,7 +246,10 @@ class EPLBManager:
|
|||||||
return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
|
return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
|
||||||
|
|
||||||
def _should_log_expert_location_metadata(self) -> bool:
|
def _should_log_expert_location_metadata(self) -> bool:
|
||||||
return self._ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
|
return (
|
||||||
|
get_parallel().tp_rank == 0
|
||||||
|
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
|
||||||
|
)
|
||||||
|
|
||||||
def _log_rebalance_layout_before_update(
|
def _log_rebalance_layout_before_update(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -172,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
aux_state_size = get_schedule().max_mamba_cache_size
|
aux_state_size = get_schedule().max_mamba_cache_size
|
||||||
if aux_state_size is None:
|
if aux_state_size is None:
|
||||||
return None
|
return None
|
||||||
return aux_state_size // self.ps.attn_dp_size
|
return aux_state_size // self.attn_dp_size
|
||||||
|
|
||||||
def _resolve_max_running_requests(self) -> int:
|
def _resolve_max_running_requests(self) -> int:
|
||||||
"""Concurrency cap handed to the scheduler.
|
"""Concurrency cap handed to the scheduler.
|
||||||
@@ -197,7 +197,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
requested_per_worker = None
|
requested_per_worker = None
|
||||||
resolved = min(capacity_cap, 4096)
|
resolved = min(capacity_cap, 4096)
|
||||||
else:
|
else:
|
||||||
requested_per_worker = requested // self.ps.attn_dp_size
|
requested_per_worker = requested // self.attn_dp_size
|
||||||
resolved = min(requested_per_worker, capacity_cap)
|
resolved = min(requested_per_worker, capacity_cap)
|
||||||
|
|
||||||
aux_state_size = self._explicit_aux_state_size_per_worker()
|
aux_state_size = self._explicit_aux_state_size_per_worker()
|
||||||
@@ -209,7 +209,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
resolved = min(resolved, aux_state_size // ratio)
|
resolved = min(resolved, aux_state_size // ratio)
|
||||||
if resolved <= 0:
|
if resolved <= 0:
|
||||||
global_aux_state_size = get_schedule().max_mamba_cache_size
|
global_aux_state_size = get_schedule().max_mamba_cache_size
|
||||||
min_global_aux_state_size = ratio * self.ps.attn_dp_size
|
min_global_aux_state_size = ratio * self.attn_dp_size
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"MLX auxiliary-state cache is too small to serve any "
|
f"MLX auxiliary-state cache is too small to serve any "
|
||||||
f"requests: max_mamba_cache_size={global_aux_state_size} "
|
f"requests: max_mamba_cache_size={global_aux_state_size} "
|
||||||
|
|||||||
@@ -108,7 +108,6 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
is_draft_worker=self.is_draft_worker,
|
is_draft_worker=self.is_draft_worker,
|
||||||
|
|||||||
@@ -428,7 +428,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
self.is_dllm_model = True
|
self.is_dllm_model = True
|
||||||
self.dllm_block_size = self.dllm_config.block_size
|
self.dllm_block_size = self.dllm_config.block_size
|
||||||
|
|
||||||
self.attn_cp_size = model_runner.ps.attn_cp_size
|
self.attn_cp_size = model_runner.attn_cp_size
|
||||||
|
|
||||||
def _is_swa_layer(self, layer: RadixAttention) -> bool:
|
def _is_swa_layer(self, layer: RadixAttention) -> bool:
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -203,7 +203,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.kv_index_translator = model_runner.kv_index_translator
|
self.kv_index_translator = model_runner.kv_index_translator
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
self.attn_cp_size = model_runner.ps.attn_cp_size
|
self.attn_cp_size = model_runner.attn_cp_size
|
||||||
self._verify_mask = None
|
self._verify_mask = None
|
||||||
# The worker fetches the tree-mask scratch from the target backend
|
# The worker fetches the tree-mask scratch from the target backend
|
||||||
# only; draft-side instances must not allocate it.
|
# only; draft-side instances must not allocate it.
|
||||||
@@ -333,10 +333,10 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.head_dim = model_runner.model_config.head_dim
|
self.head_dim = model_runner.model_config.head_dim
|
||||||
self.num_attention_heads = (
|
self.num_attention_heads = (
|
||||||
model_runner.model_config.hf_text_config.num_attention_heads
|
model_runner.model_config.hf_text_config.num_attention_heads
|
||||||
// model_runner.ps.tp_size
|
// model_runner.tp_size
|
||||||
)
|
)
|
||||||
self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
|
self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
|
||||||
model_runner.ps.tp_size
|
model_runner.tp_size
|
||||||
)
|
)
|
||||||
_softcapping = getattr(
|
_softcapping = getattr(
|
||||||
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
|
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
|
||||||
|
|||||||
@@ -162,9 +162,8 @@ class HPCOpsAttnBackend(AttentionBackend):
|
|||||||
self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn
|
self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn
|
||||||
if self.use_fp8:
|
if self.use_fp8:
|
||||||
heads = (
|
heads = (
|
||||||
model_runner.model_config.num_attention_heads
|
model_runner.model_config.num_attention_heads // model_runner.tp_size,
|
||||||
// model_runner.ps.tp_size,
|
model_runner.model_config.get_num_kv_heads(model_runner.tp_size),
|
||||||
model_runner.model_config.get_num_kv_heads(model_runner.ps.tp_size),
|
|
||||||
)
|
)
|
||||||
if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS:
|
if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -177,8 +176,8 @@ class HPCOpsAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
config = model_runner.model_config
|
config = model_runner.model_config
|
||||||
head_dim = config.head_dim
|
head_dim = config.head_dim
|
||||||
num_q_heads = config.num_attention_heads // model_runner.ps.tp_size
|
num_q_heads = config.num_attention_heads // model_runner.tp_size
|
||||||
num_kv_heads = config.get_num_kv_heads(model_runner.ps.tp_size)
|
num_kv_heads = config.get_num_kv_heads(model_runner.tp_size)
|
||||||
gqa_group_size = num_q_heads // num_kv_heads
|
gqa_group_size = num_q_heads // num_kv_heads
|
||||||
if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in (
|
if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in (
|
||||||
_SUPPORTED_GQA_GROUP_SIZES
|
_SUPPORTED_GQA_GROUP_SIZES
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ class WaveAttnBackend(AttentionBackend):
|
|||||||
import wave_lang.kernel.wave.cache as cache
|
import wave_lang.kernel.wave.cache as cache
|
||||||
|
|
||||||
base_cache_dir = cache.CACHE_BASE_DIR
|
base_cache_dir = cache.CACHE_BASE_DIR
|
||||||
new_dir = base_cache_dir / f"worker_{model_runner.ps.tp_rank}"
|
new_dir = base_cache_dir / f"worker_{model_runner.tp_rank}"
|
||||||
logger.info(f"Setting Wave cache dir: {new_dir}")
|
logger.info(f"Setting Wave cache dir: {new_dir}")
|
||||||
cache.CACHE_BASE_DIR = new_dir
|
cache.CACHE_BASE_DIR = new_dir
|
||||||
|
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
self.num_attention_heads = (
|
self.num_attention_heads = (
|
||||||
model_runner.model_config.hf_text_config.num_attention_heads
|
model_runner.model_config.hf_text_config.num_attention_heads
|
||||||
)
|
)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.tp_size
|
||||||
assert self.num_attention_heads % self.tp_size == 0
|
assert self.num_attention_heads % self.tp_size == 0
|
||||||
self.num_local_heads = self.num_attention_heads // self.tp_size
|
self.num_local_heads = self.num_attention_heads // self.tp_size
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
|
|||||||
@@ -510,7 +510,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
|||||||
"PP-parallel DeepGEMM warmup start "
|
"PP-parallel DeepGEMM warmup start "
|
||||||
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
|
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
|
||||||
get_parallel().pp_rank,
|
get_parallel().pp_rank,
|
||||||
model_runner.ps.tp_rank,
|
model_runner.tp_rank,
|
||||||
batch_sizes,
|
batch_sizes,
|
||||||
disagg_mode,
|
disagg_mode,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -47,6 +47,24 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
|
|
||||||
|
def deployment_attn_dp_size() -> int:
|
||||||
|
"""Attention-DP replicas in the deployment, which no draft scope narrows.
|
||||||
|
|
||||||
|
A draft runs on one attention-DP replica and its scope says so, but the
|
||||||
|
metadata a draft gathers is shaped by the replicas it gathers *with* --
|
||||||
|
the target's. Those come from the configuration, which the scope leaves
|
||||||
|
alone, so this answers the same number inside the scope and outside it.
|
||||||
|
"""
|
||||||
|
parallel = get_parallel()
|
||||||
|
attn_dp_size, _ = derive_attention_widths(
|
||||||
|
tp_size=parallel.tp_size,
|
||||||
|
attn_cp_size=parallel.attn_cp_size,
|
||||||
|
dp_size=parallel.dp_size,
|
||||||
|
enable_dp_attention=parallel.enable_dp_attention,
|
||||||
|
)
|
||||||
|
return attn_dp_size
|
||||||
|
|
||||||
|
|
||||||
def dp_gather_width() -> int:
|
def dp_gather_width() -> int:
|
||||||
"""How many replicas the DP sync gathers over.
|
"""How many replicas the DP sync gathers over.
|
||||||
|
|
||||||
|
|||||||
@@ -104,12 +104,10 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
abort_distributed_environment,
|
abort_distributed_environment,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
|
||||||
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
|
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
|
||||||
from sglang.srt.environ import envs, exportable_env_vars
|
from sglang.srt.environ import envs, exportable_env_vars
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
|
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
|
||||||
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
|
||||||
@@ -449,12 +447,8 @@ class Scheduler(
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
port_args: PortArgs,
|
port_args: PortArgs,
|
||||||
gpu_id: int,
|
|
||||||
tp_rank: int,
|
tp_rank: int,
|
||||||
moe_ep_rank: int,
|
|
||||||
pp_rank: int,
|
pp_rank: int,
|
||||||
attn_cp_rank: int,
|
|
||||||
moe_dp_rank: int,
|
|
||||||
dp_rank: Optional[int],
|
dp_rank: Optional[int],
|
||||||
):
|
):
|
||||||
# NOTE: KEEP THE FOLLOWING CODE STYLE for this function:
|
# NOTE: KEEP THE FOLLOWING CODE STYLE for this function:
|
||||||
@@ -516,38 +510,6 @@ class Scheduler(
|
|||||||
self.enable_dp_attention = get_parallel().enable_dp_attention
|
self.enable_dp_attention = get_parallel().enable_dp_attention
|
||||||
self.enable_unified_memory = get_memory().enable_unified_memory
|
self.enable_unified_memory = get_memory().enable_unified_memory
|
||||||
|
|
||||||
# Distributed rank info
|
|
||||||
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
|
|
||||||
compute_dp_attention_world_info(
|
|
||||||
get_parallel().enable_dp_attention,
|
|
||||||
tp_rank,
|
|
||||||
get_parallel().tp_size,
|
|
||||||
get_parallel().dp_size,
|
|
||||||
get_parallel().attn_cp_size,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.ps = ParallelState(
|
|
||||||
tp_rank=tp_rank,
|
|
||||||
tp_size=get_parallel().tp_size,
|
|
||||||
pp_rank=pp_rank,
|
|
||||||
pp_size=get_parallel().pp_size,
|
|
||||||
dp_rank=dp_rank,
|
|
||||||
dp_size=get_parallel().dp_size,
|
|
||||||
attn_tp_rank=attn_tp_rank,
|
|
||||||
attn_tp_size=attn_tp_size,
|
|
||||||
attn_cp_rank=attn_cp_rank,
|
|
||||||
attn_cp_size=get_parallel().attn_cp_size,
|
|
||||||
attn_dcp_rank=tp_rank % get_parallel().dcp_size,
|
|
||||||
attn_dcp_size=get_parallel().dcp_size,
|
|
||||||
attn_dp_rank=attn_dp_rank,
|
|
||||||
attn_dp_size=attn_dp_size,
|
|
||||||
moe_ep_rank=moe_ep_rank,
|
|
||||||
moe_ep_size=get_parallel().ep_size,
|
|
||||||
moe_dp_rank=moe_dp_rank,
|
|
||||||
moe_dp_size=get_parallel().moe_dp_size,
|
|
||||||
gpu_id=gpu_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Init model configs
|
# Init model configs
|
||||||
self.init_model_config()
|
self.init_model_config()
|
||||||
|
|
||||||
@@ -779,7 +741,7 @@ class Scheduler(
|
|||||||
if get_parallel().pp_size > 1:
|
if get_parallel().pp_size > 1:
|
||||||
logger.error("only zbal mix mode support pp_size > 1!")
|
logger.error("only zbal mix mode support pp_size > 1!")
|
||||||
init_zbal(
|
init_zbal(
|
||||||
get_parallel().tp_size, self.ps.gpu_id, get_parallel().tp_rank
|
get_parallel().tp_size, get_device().gpu_id, get_parallel().tp_rank
|
||||||
) # only switch allocator if is mix mode
|
) # only switch allocator if is mix mode
|
||||||
|
|
||||||
def init_model_config(self):
|
def init_model_config(self):
|
||||||
@@ -1009,8 +971,7 @@ class Scheduler(
|
|||||||
def init_tp_model_worker(self):
|
def init_tp_model_worker(self):
|
||||||
worker_kwargs = dict(
|
worker_kwargs = dict(
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
gpu_id=self.ps.gpu_id,
|
gpu_id=get_device().gpu_id,
|
||||||
ps=self.ps,
|
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1047,8 +1008,7 @@ class Scheduler(
|
|||||||
# — is resolved per runner, not on a config copy.
|
# — is resolved per runner, not on a config copy.
|
||||||
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=get_device().gpu_id,
|
||||||
ps=self.ps,
|
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
target_worker=self.tp_worker,
|
target_worker=self.tp_worker,
|
||||||
)
|
)
|
||||||
@@ -1246,7 +1206,7 @@ class Scheduler(
|
|||||||
|
|
||||||
# Print debug info
|
# Print debug info
|
||||||
self.startup_available_gpu_memory_gb = get_available_gpu_memory(
|
self.startup_available_gpu_memory_gb = get_available_gpu_memory(
|
||||||
self.device, self.ps.gpu_id, empty_cache=False
|
self.device, get_device().gpu_id, empty_cache=False
|
||||||
)
|
)
|
||||||
if get_parallel().tp_rank == 0:
|
if get_parallel().tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -1582,7 +1542,7 @@ class Scheduler(
|
|||||||
tp_rank=get_parallel().tp_rank,
|
tp_rank=get_parallel().tp_rank,
|
||||||
tp_size=get_parallel().tp_size,
|
tp_size=get_parallel().tp_size,
|
||||||
dp_size=get_parallel().dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
gpu_id=self.ps.gpu_id,
|
gpu_id=get_device().gpu_id,
|
||||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
pp_rank=get_parallel().pp_rank,
|
pp_rank=get_parallel().pp_rank,
|
||||||
@@ -1613,7 +1573,7 @@ class Scheduler(
|
|||||||
metadata_buffers=self.disagg_metadata_buffers,
|
metadata_buffers=self.disagg_metadata_buffers,
|
||||||
tp_rank=get_parallel().tp_rank,
|
tp_rank=get_parallel().tp_rank,
|
||||||
tp_size=get_parallel().tp_size,
|
tp_size=get_parallel().tp_size,
|
||||||
gpu_id=self.ps.gpu_id,
|
gpu_id=get_device().gpu_id,
|
||||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||||
gloo_group=self.attn_tp_cpu_group,
|
gloo_group=self.attn_tp_cpu_group,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
@@ -2248,7 +2208,6 @@ class Scheduler(
|
|||||||
|
|
||||||
def init_profiler(self) -> None:
|
def init_profiler(self) -> None:
|
||||||
self.profiler_manager = SchedulerProfilerManager(
|
self.profiler_manager = SchedulerProfilerManager(
|
||||||
ps=self.ps,
|
|
||||||
dp_tp_cpu_group=self.dp_tp_cpu_group,
|
dp_tp_cpu_group=self.dp_tp_cpu_group,
|
||||||
get_forward_ct=lambda: self.forward_ct,
|
get_forward_ct=lambda: self.forward_ct,
|
||||||
)
|
)
|
||||||
@@ -2369,7 +2328,6 @@ class Scheduler(
|
|||||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
offload_tags=self.weight_updater.offload_tags,
|
offload_tags=self.weight_updater.offload_tags,
|
||||||
ps=self.ps,
|
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
enable_overlap=self.enable_overlap,
|
enable_overlap=self.enable_overlap,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
@@ -2457,7 +2415,6 @@ class Scheduler(
|
|||||||
self._sched_idled = False
|
self._sched_idled = False
|
||||||
self.load_inquirer = SchedulerLoadInquirer(
|
self.load_inquirer = SchedulerLoadInquirer(
|
||||||
disaggregation_mode=self.disaggregation_mode,
|
disaggregation_mode=self.disaggregation_mode,
|
||||||
ps=self.ps,
|
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
@@ -2500,7 +2457,6 @@ class Scheduler(
|
|||||||
self.output_streamer = self.get_output_streamer_class()(
|
self.output_streamer = self.get_output_streamer_class()(
|
||||||
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
|
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
ps=self.ps,
|
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
is_generation=self.is_generation,
|
is_generation=self.is_generation,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
@@ -3120,7 +3076,7 @@ class Scheduler(
|
|||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
return
|
||||||
|
|
||||||
if self.ps.pp_rank == 0 and getattr(
|
if get_parallel().pp_rank == 0 and getattr(
|
||||||
self.tree_cache.cache_controller, "pp_prefetch_command_group", None
|
self.tree_cache.cache_controller, "pp_prefetch_command_group", None
|
||||||
):
|
):
|
||||||
recv_req.pp_prefetch_ticketed = bool(self._prefetch_kvcache(req))
|
recv_req.pp_prefetch_ticketed = bool(self._prefetch_kvcache(req))
|
||||||
@@ -6078,6 +6034,7 @@ def run_scheduler_process(
|
|||||||
ranks=SpawnRanks(
|
ranks=SpawnRanks(
|
||||||
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=pp_rank),
|
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=pp_rank),
|
||||||
dp_rank=dp_rank,
|
dp_rank=dp_rank,
|
||||||
|
gpu_id=gpu_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
configure_scheduler_process(
|
configure_scheduler_process(
|
||||||
@@ -6115,12 +6072,8 @@ def run_scheduler_process(
|
|||||||
scheduler = Scheduler(
|
scheduler = Scheduler(
|
||||||
server_args,
|
server_args,
|
||||||
port_args,
|
port_args,
|
||||||
gpu_id,
|
|
||||||
tp_rank,
|
tp_rank,
|
||||||
moe_ep_rank,
|
|
||||||
pp_rank,
|
pp_rank,
|
||||||
attn_cp_rank,
|
|
||||||
moe_dp_rank,
|
|
||||||
dp_rank,
|
dp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
|
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
|
||||||
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.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.cp.utils import get_cp_strategy
|
from sglang.srt.layers.cp.utils import get_cp_strategy
|
||||||
from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled
|
from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled
|
||||||
@@ -537,7 +536,6 @@ class SchedulerDPAttnAdapter:
|
|||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
||||||
tree_cache: BasePrefixCache
|
tree_cache: BasePrefixCache
|
||||||
offload_tags: set[str]
|
offload_tags: set[str]
|
||||||
ps: ParallelState
|
|
||||||
model_config: ModelConfig
|
model_config: ModelConfig
|
||||||
enable_overlap: bool
|
enable_overlap: bool
|
||||||
spec_algorithm: SpeculativeAlgorithm
|
spec_algorithm: SpeculativeAlgorithm
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ from sglang.srt.managers.load_snapshot import (
|
|||||||
from sglang.srt.runtime_context import get_lora, get_parallel
|
from sglang.srt.runtime_context import get_lora, get_parallel
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
|
||||||
from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
||||||
SchedulerPoolStatsObserver,
|
SchedulerPoolStatsObserver,
|
||||||
)
|
)
|
||||||
@@ -33,7 +32,6 @@ logger = logging.getLogger(__name__)
|
|||||||
@dataclass(kw_only=True, slots=True, frozen=True)
|
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||||
class SchedulerLoadInquirer:
|
class SchedulerLoadInquirer:
|
||||||
disaggregation_mode: DisaggregationMode
|
disaggregation_mode: DisaggregationMode
|
||||||
ps: ParallelState
|
|
||||||
server_args: ServerArgs
|
server_args: ServerArgs
|
||||||
max_total_num_tokens: int
|
max_total_num_tokens: int
|
||||||
max_running_requests: int
|
max_running_requests: int
|
||||||
|
|||||||
@@ -308,7 +308,7 @@ class SchedulerMetricsReporter:
|
|||||||
self.scheduler.enable_fpm = False
|
self.scheduler.enable_fpm = False
|
||||||
if (
|
if (
|
||||||
get_observability().enable_forward_pass_metrics
|
get_observability().enable_forward_pass_metrics
|
||||||
and self.scheduler.ps.attn_tp_rank == 0
|
and get_parallel().attn_tp_rank == 0
|
||||||
and get_parallel().pp_rank == get_parallel().pp_size - 1
|
and get_parallel().pp_rank == get_parallel().pp_size - 1
|
||||||
):
|
):
|
||||||
from sglang.srt.observability.forward_pass_metrics import (
|
from sglang.srt.observability.forward_pass_metrics import (
|
||||||
@@ -316,9 +316,7 @@ class SchedulerMetricsReporter:
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.scheduler._fpm_dp_rank = (
|
self.scheduler._fpm_dp_rank = (
|
||||||
self.scheduler.ps.dp_rank
|
get_parallel().dp_rank if get_parallel().dp_rank is not None else 0
|
||||||
if self.scheduler.ps.dp_rank is not None
|
|
||||||
else 0
|
|
||||||
)
|
)
|
||||||
self.scheduler._fpm_worker_id = (
|
self.scheduler._fpm_worker_id = (
|
||||||
get_observability().forward_pass_metrics_worker_id
|
get_observability().forward_pass_metrics_worker_id
|
||||||
@@ -483,9 +481,9 @@ class SchedulerMetricsReporter:
|
|||||||
num_layers = float(getattr(model_config, "num_attention_layers", 0))
|
num_layers = float(getattr(model_config, "num_attention_layers", 0))
|
||||||
head_dim = float(getattr(model_config, "head_dim", 0))
|
head_dim = float(getattr(model_config, "head_dim", 0))
|
||||||
num_attn_heads = float(
|
num_attn_heads = float(
|
||||||
model_config.get_num_attention_heads(self.scheduler.ps.tp_size)
|
model_config.get_num_attention_heads(get_parallel().tp_size)
|
||||||
)
|
)
|
||||||
num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size))
|
num_kv_heads = float(model_config.get_num_kv_heads(get_parallel().tp_size))
|
||||||
intermediate_size = getattr(hf_text_config, "intermediate_size", None)
|
intermediate_size = getattr(hf_text_config, "intermediate_size", None)
|
||||||
if intermediate_size is None:
|
if intermediate_size is None:
|
||||||
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0)
|
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0)
|
||||||
|
|||||||
@@ -19,7 +19,6 @@ from sglang.srt.beam_search.output import (
|
|||||||
pack_beam_search_output,
|
pack_beam_search_output,
|
||||||
)
|
)
|
||||||
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.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
BatchEmbeddingOutput,
|
BatchEmbeddingOutput,
|
||||||
@@ -53,7 +52,6 @@ class SchedulerOutputStreamer:
|
|||||||
|
|
||||||
send_to_detokenizer: zmq.Socket
|
send_to_detokenizer: zmq.Socket
|
||||||
tree_cache: BasePrefixCache
|
tree_cache: BasePrefixCache
|
||||||
ps: ParallelState
|
|
||||||
server_args: ServerArgs
|
server_args: ServerArgs
|
||||||
is_generation: bool
|
is_generation: bool
|
||||||
spec_algorithm: SpeculativeAlgorithm
|
spec_algorithm: SpeculativeAlgorithm
|
||||||
|
|||||||
@@ -51,14 +51,12 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
class SchedulerProfilerManager:
|
class SchedulerProfilerManager:
|
||||||
ps: Any
|
|
||||||
dp_tp_cpu_group: Any
|
dp_tp_cpu_group: Any
|
||||||
get_forward_ct: Callable[[], int]
|
get_forward_ct: Callable[[], int]
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
if envs.SGLANG_PROFILE_V2.get():
|
if envs.SGLANG_PROFILE_V2.get():
|
||||||
self._profile_manager = ProfileManager(
|
self._profile_manager = ProfileManager(
|
||||||
ps=self.ps,
|
|
||||||
cpu_group=self.dp_tp_cpu_group,
|
cpu_group=self.dp_tp_cpu_group,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -274,7 +272,7 @@ class SchedulerProfilerManager:
|
|||||||
self.profile_in_progress = True
|
self.profile_in_progress = True
|
||||||
|
|
||||||
if "CUDA_PROFILER" in activities:
|
if "CUDA_PROFILER" in activities:
|
||||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
if get_device().gpu_id == get_device().base_gpu_id:
|
||||||
torch.cuda.cudart().cudaProfilerStart()
|
torch.cuda.cudart().cudaProfilerStart()
|
||||||
self.profile_in_progress = True
|
self.profile_in_progress = True
|
||||||
|
|
||||||
@@ -387,7 +385,7 @@ class SchedulerProfilerManager:
|
|||||||
torch.cuda.memory._record_memory_history(enabled=None)
|
torch.cuda.memory._record_memory_history(enabled=None)
|
||||||
|
|
||||||
if "CUDA_PROFILER" in self.profiler_activities:
|
if "CUDA_PROFILER" in self.profiler_activities:
|
||||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
if get_device().gpu_id == get_device().base_gpu_id:
|
||||||
torch.cuda.cudart().cudaProfilerStop()
|
torch.cuda.cudart().cudaProfilerStop()
|
||||||
|
|
||||||
merge_message = self._merge_profile_traces()
|
merge_message = self._merge_profile_traces()
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.beam_search.logits_capture import capture_pre_sample_logits
|
from sglang.srt.beam_search.logits_capture import capture_pre_sample_logits
|
||||||
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.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
DestroyWeightsUpdateGroupReqInput,
|
DestroyWeightsUpdateGroupReqInput,
|
||||||
@@ -216,12 +215,12 @@ class BaseTpWorker(ABC):
|
|||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
def _deserialize_own_rank(self, serialized_named_tensors):
|
def _deserialize_own_rank(self, serialized_named_tensors):
|
||||||
"""Each rank deserializes only its own payload (index ps.tp_rank);
|
"""Each rank deserializes only its own payload (index tp_rank);
|
||||||
deserializing another rank's copy would break producer-side CUDA-IPC
|
deserializing another rank's copy would break producer-side CUDA-IPC
|
||||||
refcounting."""
|
refcounting."""
|
||||||
monkey_patch_torch_reductions()
|
monkey_patch_torch_reductions()
|
||||||
return MultiprocessingSerializer.deserialize(
|
return MultiprocessingSerializer.deserialize(
|
||||||
serialized_named_tensors[self.ps.tp_rank]
|
serialized_named_tensors[self.model_runner.tp_rank]
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
@@ -290,12 +289,12 @@ class BaseTpWorker(ABC):
|
|||||||
extra = [n for n in tensors if n not in exp]
|
extra = [n for n in tensors if n not in exp]
|
||||||
if mismatch or missing or extra:
|
if mismatch or missing or extra:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
|
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
|
||||||
f"{len(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
|
f"{len(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
|
||||||
f"{len(extra)} extra {extra[:5]}"
|
f"{len(extra)} extra {extra[:5]}"
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
|
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
|
||||||
)
|
)
|
||||||
result = self.model_runner.load_lora_adapter_from_tensors(
|
result = self.model_runner.load_lora_adapter_from_tensors(
|
||||||
recv_req.to_ref(),
|
recv_req.to_ref(),
|
||||||
@@ -322,7 +321,6 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
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,
|
||||||
@@ -335,7 +333,6 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
):
|
):
|
||||||
# Parse args
|
# Parse args
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.ps = ps
|
|
||||||
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
|
||||||
@@ -411,14 +408,15 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
tp_group = self.model_runner.tp_group
|
tp_group = self.model_runner.tp_group
|
||||||
self.random_seed = broadcast_pyobj(
|
self.random_seed = broadcast_pyobj(
|
||||||
[get_device().random_seed],
|
[get_device().random_seed],
|
||||||
tp_group.ranks[self.ps.tp_rank],
|
tp_group.ranks[self.model_runner.tp_rank],
|
||||||
tp_group.cpu_group,
|
tp_group.cpu_group,
|
||||||
src=tp_group.ranks[0],
|
src=tp_group.ranks[0],
|
||||||
)[0]
|
)[0]
|
||||||
else:
|
else:
|
||||||
self.random_seed = broadcast_pyobj(
|
self.random_seed = broadcast_pyobj(
|
||||||
[get_device().random_seed],
|
[get_device().random_seed],
|
||||||
self.ps.tp_size * get_parallel().pp_rank + self.ps.tp_rank,
|
self.model_runner.tp_size * get_parallel().pp_rank
|
||||||
|
+ self.model_runner.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]
|
||||||
@@ -521,7 +519,6 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
is_draft_worker=self.is_draft_worker,
|
is_draft_worker=self.is_draft_worker,
|
||||||
@@ -542,7 +539,6 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
is_draft_worker=self.is_draft_worker,
|
is_draft_worker=self.is_draft_worker,
|
||||||
|
|||||||
@@ -1102,13 +1102,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
model_runner.lora_manager.prepare_lora_batch(ret)
|
model_runner.lora_manager.prepare_lora_batch(ret)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
model_runner.ps.attn_dcp_size > 1
|
model_runner.attn_dcp_size > 1
|
||||||
and ret.out_cache_loc is not None
|
and ret.out_cache_loc is not None
|
||||||
and is_hip()
|
and is_hip()
|
||||||
):
|
):
|
||||||
ret.dcp_kv_mask = (
|
ret.dcp_kv_mask = (
|
||||||
ret.positions % model_runner.ps.attn_dcp_size
|
ret.positions % model_runner.attn_dcp_size == model_runner.attn_dcp_rank
|
||||||
== model_runner.ps.attn_dcp_rank
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return ret
|
return ret
|
||||||
|
|||||||
@@ -37,7 +37,6 @@ from sglang.srt.distributed import bootstrap
|
|||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
maybe_init_shared_mooncake_transfer_engine,
|
maybe_init_shared_mooncake_transfer_engine,
|
||||||
)
|
)
|
||||||
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,
|
||||||
@@ -322,7 +321,6 @@ class ModelRunner:
|
|||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
mem_fraction_static: float,
|
mem_fraction_static: float,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
is_draft_worker: bool = False,
|
is_draft_worker: bool = False,
|
||||||
@@ -339,7 +337,6 @@ class ModelRunner:
|
|||||||
# `server_args._draft_pool_config` mutation hack).
|
# `server_args._draft_pool_config` mutation hack).
|
||||||
self.memory_pool_config = memory_pool_config
|
self.memory_pool_config = memory_pool_config
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.ps = ps
|
|
||||||
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
|
||||||
@@ -416,12 +413,12 @@ class ModelRunner:
|
|||||||
# Set device early so that TransferEngine init (e.g. Ascend NPU)
|
# Set device early so that TransferEngine init (e.g. Ascend NPU)
|
||||||
# can access the device context.
|
# can access the device context.
|
||||||
try:
|
try:
|
||||||
torch.get_device_module(self.device).set_device(ps.gpu_id)
|
torch.get_device_module(self.device).set_device(get_device().gpu_id)
|
||||||
except Exception:
|
except Exception:
|
||||||
import os
|
import os
|
||||||
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Context: {self.device=} {ps.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {get_parallel().tp_rank=} {get_parallel().tp_size=}"
|
f"Context: {self.device=} {get_device().gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {get_parallel().tp_rank=} {get_parallel().tp_size=}"
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
@@ -741,7 +738,6 @@ class ModelRunner:
|
|||||||
self.eplb_manager = (
|
self.eplb_manager = (
|
||||||
EPLBManager(
|
EPLBManager(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
ps=self.ps,
|
|
||||||
get_model=lambda: self.model,
|
get_model=lambda: self.model,
|
||||||
get_expert_location_updater=lambda: self.expert_location_updater,
|
get_expert_location_updater=lambda: self.expert_location_updater,
|
||||||
get_expert_backup_client=lambda: self.expert_backup_client,
|
get_expert_backup_client=lambda: self.expert_backup_client,
|
||||||
@@ -805,8 +801,8 @@ class ModelRunner:
|
|||||||
def get_pp_proxy_dspark_hidden_size(self) -> int:
|
def get_pp_proxy_dspark_hidden_size(self) -> int:
|
||||||
return misc_utils.resolve_pp_proxy_dspark_hidden_size(
|
return misc_utils.resolve_pp_proxy_dspark_hidden_size(
|
||||||
model=self.model,
|
model=self.model,
|
||||||
pp_size=self.ps.pp_size,
|
pp_size=self.pp_size,
|
||||||
pp_rank=self.ps.pp_rank,
|
pp_rank=self.pp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_pp_proxy_topk_size(self) -> Optional[int]:
|
def get_pp_proxy_topk_size(self) -> Optional[int]:
|
||||||
@@ -1175,7 +1171,6 @@ class ModelRunner:
|
|||||||
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,
|
||||||
ps=self.ps,
|
|
||||||
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,
|
||||||
@@ -1197,6 +1192,10 @@ class ModelRunner:
|
|||||||
self.pp_size = parallel.pp_size
|
self.pp_size = parallel.pp_size
|
||||||
self.attn_cp_rank = parallel.attn_cp_rank
|
self.attn_cp_rank = parallel.attn_cp_rank
|
||||||
self.attn_cp_size = parallel.attn_cp_size
|
self.attn_cp_size = parallel.attn_cp_size
|
||||||
|
self.attn_dcp_rank = parallel.attn_dcp_rank
|
||||||
|
self.attn_dcp_size = parallel.attn_dcp_size
|
||||||
|
self.moe_ep_size = parallel.moe_ep_size
|
||||||
|
self.dp_rank = parallel.dp_rank
|
||||||
|
|
||||||
def init_shared_mooncake_transfer_engine(self):
|
def init_shared_mooncake_transfer_engine(self):
|
||||||
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
|
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
|
||||||
|
|||||||
@@ -220,7 +220,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._cell_size == 0
|
self._cell_size == 0
|
||||||
and mambaish is not None
|
and mambaish is not None
|
||||||
and bool(mambaish.full_attention_layer_ids)
|
and bool(mambaish.full_attention_layer_ids)
|
||||||
and kvc.ps.pp_size > 1
|
and kvc.pp_size > 1
|
||||||
)
|
)
|
||||||
self._zero_kv_max_tokens = (
|
self._zero_kv_max_tokens = (
|
||||||
torch.iinfo(torch.int64).max
|
torch.iinfo(torch.int64).max
|
||||||
@@ -877,7 +877,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
self._swa_cap = compute_swa_request_cap(
|
self._swa_cap = compute_swa_request_cap(
|
||||||
page_size=kvc.page_size,
|
page_size=kvc.page_size,
|
||||||
window=kvc.sliding_window_size,
|
window=kvc.sliding_window_size,
|
||||||
attn_dp_size=kvc.ps.attn_dp_size,
|
attn_dp_size=kvc.attn_dp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -1015,7 +1015,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self.compression_ratios = cfg.compress_ratios[
|
self.compression_ratios = cfg.compress_ratios[
|
||||||
kvc.layer_info.start_layer : kvc.layer_info.end_layer
|
kvc.layer_info.start_layer : kvc.layer_info.end_layer
|
||||||
]
|
]
|
||||||
if kvc.ps.pp_size > 1:
|
if kvc.pp_size > 1:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} "
|
f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} "
|
||||||
f"layers=[{kvc.layer_info.start_layer},{kvc.layer_info.end_layer}) "
|
f"layers=[{kvc.layer_info.start_layer},{kvc.layer_info.end_layer}) "
|
||||||
@@ -1032,9 +1032,9 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self.page_size = kvc.page_size
|
self.page_size = kvc.page_size
|
||||||
self.is_speculative = get_spec().speculative_algorithm is not None
|
self.is_speculative = get_spec().speculative_algorithm is not None
|
||||||
self.online_c128_mtp_max_draft_tokens = max_speculative_num_draft_tokens() or 0
|
self.online_c128_mtp_max_draft_tokens = max_speculative_num_draft_tokens() or 0
|
||||||
self.attn_dp_size = kvc.ps.attn_dp_size
|
self.attn_dp_size = kvc.attn_dp_size
|
||||||
self.requested_max_running_requests_per_worker = (
|
self.requested_max_running_requests_per_worker = (
|
||||||
get_schedule().max_running_requests // kvc.ps.attn_dp_size
|
get_schedule().max_running_requests // kvc.attn_dp_size
|
||||||
if get_schedule().max_running_requests is not None
|
if get_schedule().max_running_requests is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -547,9 +547,9 @@ class BaseRunner(ABC):
|
|||||||
if (
|
if (
|
||||||
capture_forward_mode == ForwardMode.EXTEND
|
capture_forward_mode == ForwardMode.EXTEND
|
||||||
and get_parallel().pp_rank != 0
|
and get_parallel().pp_rank != 0
|
||||||
and mr.ps.attn_cp_size > 1
|
and mr.attn_cp_size > 1
|
||||||
):
|
):
|
||||||
pp_hidden_tokens = num_tokens // mr.ps.attn_cp_size
|
pp_hidden_tokens = num_tokens // mr.attn_cp_size
|
||||||
pp_proxy_tensors = PPProxyTensors(
|
pp_proxy_tensors = PPProxyTensors(
|
||||||
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -296,7 +296,7 @@ class EagerRunner(BaseRunner):
|
|||||||
or cp_active
|
or cp_active
|
||||||
or forward_batch.forward_mode.is_target_verify()
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
):
|
):
|
||||||
if model_runner.ps.attn_dcp_size > 1 and hasattr(
|
if model_runner.attn_dcp_size > 1 and hasattr(
|
||||||
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
|
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
|
||||||
):
|
):
|
||||||
# prepare kv cache buffer for dcp to gather kv cache
|
# prepare kv cache buffer for dcp to gather kv cache
|
||||||
|
|||||||
@@ -147,10 +147,10 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
|
|||||||
str(mr.dtype),
|
str(mr.dtype),
|
||||||
str(get_model().quantization),
|
str(get_model().quantization),
|
||||||
str(get_exec().moe.moe_runner_backend),
|
str(get_exec().moe.moe_runner_backend),
|
||||||
str(mr.ps.tp_size),
|
str(mr.tp_size),
|
||||||
str(get_parallel().pp_size),
|
str(get_parallel().pp_size),
|
||||||
str(mr.ps.attn_dp_size),
|
str(mr.attn_dp_size),
|
||||||
str(mr.ps.moe_ep_size),
|
str(mr.moe_ep_size),
|
||||||
str(mr.model_config.hf_config.__class__.__name__),
|
str(mr.model_config.hf_config.__class__.__name__),
|
||||||
]
|
]
|
||||||
# A different skip policy must not reuse previously tuned tactics.
|
# A different skip policy must not reuse previously tuned tactics.
|
||||||
@@ -171,7 +171,7 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
|
|||||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||||
return (
|
return (
|
||||||
cache_dir
|
cache_dir
|
||||||
/ f"rank_tp{mr.ps.tp_rank}_pp{get_parallel().pp_rank}_dp{mr.ps.dp_rank or 0}.json"
|
/ f"rank_tp{mr.tp_rank}_pp{get_parallel().pp_rank}_dp{mr.dp_rank or 0}.json"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ from typing import Any, Dict, Optional
|
|||||||
import ray
|
import ray
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
from sglang.srt.runtime_context import SpawnRanks, publish, spawn_world_rank
|
from sglang.srt.runtime_context import SpawnRanks, get_device, publish, spawn_world_rank
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -93,6 +93,7 @@ class SchedulerActor:
|
|||||||
server_args, tp_rank=tp_rank, pp_rank=pp_rank
|
server_args, tp_rank=tp_rank, pp_rank=pp_rank
|
||||||
),
|
),
|
||||||
dp_rank=dp_rank,
|
dp_rank=dp_rank,
|
||||||
|
gpu_id=actual_gpu_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -124,12 +125,8 @@ class SchedulerActor:
|
|||||||
self.scheduler = Scheduler(
|
self.scheduler = Scheduler(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
port_args=port_args,
|
port_args=port_args,
|
||||||
gpu_id=actual_gpu_id,
|
|
||||||
tp_rank=tp_rank,
|
tp_rank=tp_rank,
|
||||||
moe_ep_rank=moe_ep_rank,
|
|
||||||
pp_rank=pp_rank,
|
pp_rank=pp_rank,
|
||||||
attn_cp_rank=attn_cp_rank,
|
|
||||||
moe_dp_rank=moe_dp_rank,
|
|
||||||
dp_rank=dp_rank,
|
dp_rank=dp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -146,7 +143,7 @@ class SchedulerActor:
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
# Need to set the GPU id for the event loop for nccl to work
|
# Need to set the GPU id for the event loop for nccl to work
|
||||||
torch.cuda.set_device(self.scheduler.ps.gpu_id)
|
torch.cuda.set_device(get_device().gpu_id)
|
||||||
self.scheduler.run_event_loop()
|
self.scheduler.run_event_loop()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Scheduler PP{self._pp_rank} TP{self._tp_rank} crashed: {e}")
|
logger.error(f"Scheduler PP{self._pp_rank} TP{self._tp_rank} crashed: {e}")
|
||||||
|
|||||||
@@ -1119,7 +1119,7 @@ def _install_derived_leaves(tops: dict, server_args: Any) -> None:
|
|||||||
"""
|
"""
|
||||||
import importlib
|
import importlib
|
||||||
|
|
||||||
from sglang.srt.arg_groups.arg_utils import Derived
|
from sglang.srt.arg_groups.arg_utils import _NO_DEFAULT, Derived
|
||||||
from sglang.srt.arg_groups.overrides import resolved_view
|
from sglang.srt.arg_groups.overrides import resolved_view
|
||||||
|
|
||||||
namespaces = getattr(type(server_args), "_NAMESPACES", None)
|
namespaces = getattr(type(server_args), "_NAMESPACES", None)
|
||||||
@@ -1131,7 +1131,18 @@ def _install_derived_leaves(tops: dict, server_args: Any) -> None:
|
|||||||
if path is None:
|
if path is None:
|
||||||
continue
|
continue
|
||||||
for name, decl in vars(source).items():
|
for name, decl in vars(source).items():
|
||||||
if not isinstance(decl, Derived) or not decl.fn:
|
if not isinstance(decl, Derived):
|
||||||
|
continue
|
||||||
|
if not decl.fn:
|
||||||
|
# Nothing computes it. Seed the ones whose absence is itself an
|
||||||
|
# answer, so a process that never states one still reads it;
|
||||||
|
# the rest stay unwritten and say so when read.
|
||||||
|
if decl.default is not _NO_DEFAULT:
|
||||||
|
bag = tops.get(path.split(".")[0])
|
||||||
|
for segment in path.split(".")[1:]:
|
||||||
|
bag = bag and getattr(bag, segment, None)
|
||||||
|
if bag is not None:
|
||||||
|
bag._set(name, decl.default)
|
||||||
continue
|
continue
|
||||||
module, _, attr = decl.fn.rpartition(".")
|
module, _, attr = decl.fn.rpartition(".")
|
||||||
bag = tops.get(path.split(".")[0])
|
bag = tops.get(path.split(".")[0])
|
||||||
@@ -1251,7 +1262,20 @@ class RuntimeContext:
|
|||||||
# Snapshot resolved config into the namespace bags (the single source of
|
# Snapshot resolved config into the namespace bags (the single source of
|
||||||
# truth for config reads). Placed by `namespace_of`; a mock/partial
|
# truth for config reads). Placed by `namespace_of`; a mock/partial
|
||||||
# config that declares no namespace yields an empty tree (no bags).
|
# config that declares no namespace yields an empty tree (no bags).
|
||||||
|
# A name the configuration does not carry survives the re-projection.
|
||||||
|
# `gpu_id` is stated by the spawn, not derived, so rebuilding the bags
|
||||||
|
# from the record must not unwrite it -- the record has no field for it
|
||||||
|
# to be rebuilt from.
|
||||||
|
stated = {}
|
||||||
|
if self._config_bags is not None:
|
||||||
|
device = self._config_bags.get("device")
|
||||||
|
fields = object.__getattribute__(device, "_fields") if device else {}
|
||||||
|
if "gpu_id" in fields:
|
||||||
|
stated["gpu_id"] = fields["gpu_id"]
|
||||||
self._config_bags = _build_config_bags(server_args)
|
self._config_bags = _build_config_bags(server_args)
|
||||||
|
device = self._config_bags.get("device")
|
||||||
|
if stated and device is not None:
|
||||||
|
device._set("gpu_id", stated["gpu_id"])
|
||||||
spec = self._config_bags.get("spec")
|
spec = self._config_bags.get("spec")
|
||||||
if spec is not None:
|
if spec is not None:
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
@@ -1818,8 +1842,13 @@ def publish(
|
|||||||
# never built. With DCP on it is a position, and the bundle below states it.
|
# never built. With DCP on it is a position, and the bundle below states it.
|
||||||
if not _CONTEXT.parallel.dcp_enabled:
|
if not _CONTEXT.parallel.dcp_enabled:
|
||||||
_CONTEXT.parallel.override_permanently(attn_dcp_rank=0)
|
_CONTEXT.parallel.override_permanently(attn_dcp_rank=0)
|
||||||
if ranks is not None and ranks.gpu_id is not None:
|
# Stated on the bag directly: `gpu_id` is declared but not configured, so
|
||||||
_CONTEXT.override("spawn", gpu_id=ranks.gpu_id)
|
# it is not a leaf `override` can route, and the spawn is the only thing
|
||||||
|
# that knows it. Written whatever it is, `None` included -- most roles run
|
||||||
|
# on no device, and that is an answer rather than a name nobody wrote.
|
||||||
|
_CONTEXT.config_bag("device")._set(
|
||||||
|
"gpu_id", ranks.gpu_id if ranks is not None else None
|
||||||
|
)
|
||||||
if ranks is not None:
|
if ranks is not None:
|
||||||
# The placement, worked out here rather than carried: the widths are on
|
# The placement, worked out here rather than carried: the widths are on
|
||||||
# the bag a moment ago, and `world_rank` fixes the rest. A read of any
|
# the bag a moment ago, and `world_rank` fixes the rest. A read of any
|
||||||
@@ -1882,7 +1911,7 @@ def _attention_ranks(parallel, tp_rank: int) -> dict:
|
|||||||
The widths are already on the bag -- `publish` computed them a moment ago --
|
The widths are already on the bag -- `publish` computed them a moment ago --
|
||||||
and the rank comes from the spawn, so the position is known here, before any
|
and the rank comes from the spawn, so the position is known here, before any
|
||||||
process group exists. That is the point: a rank read then works in a process
|
process group exists. That is the point: a rank read then works in a process
|
||||||
that never initialises distributed, which is what `ParallelState` provided
|
that never initialises distributed, which is what the per-runner record used to provide
|
||||||
by being a plain frozen record.
|
by being a plain frozen record.
|
||||||
|
|
||||||
These are stamped rather than written as bag leaves because they are
|
These are stamped rather than written as bag leaves because they are
|
||||||
|
|||||||
@@ -119,17 +119,17 @@ class RustServer:
|
|||||||
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||||||
|
|
||||||
# Preserve the DP startup log; ports use node-local offsets.
|
# Preserve the DP startup log; ports use node-local offsets.
|
||||||
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
|
dp_rank = get_parallel().attn_dp_rank if get_parallel().dp_size > 1 else None
|
||||||
if get_exec().moe.is_ep_scale_joiner:
|
if get_exec().moe.is_ep_scale_joiner:
|
||||||
# The joining TP group is entirely local to this node.
|
# The joining TP group is entirely local to this node.
|
||||||
tp_size_per_node = scheduler.ps.tp_size
|
tp_size_per_node = get_parallel().tp_size
|
||||||
else:
|
else:
|
||||||
nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1)
|
nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1)
|
||||||
tp_size_per_node = scheduler.ps.tp_size // nnodes_per_pp_rank
|
tp_size_per_node = get_parallel().tp_size // nnodes_per_pp_rank
|
||||||
dp_group_width = scheduler.ps.attn_tp_size * scheduler.ps.attn_cp_size
|
dp_group_width = get_parallel().attn_tp_size * get_parallel().attn_cp_size
|
||||||
# Count DP leaders within this node's TP range. The first leader must
|
# Count DP leaders within this node's TP range. The first leader must
|
||||||
# use the base port even when a DP group spans multiple nodes.
|
# use the base port even when a DP group spans multiple nodes.
|
||||||
local_dp_rank = (scheduler.ps.tp_rank % tp_size_per_node) // dp_group_width
|
local_dp_rank = (get_parallel().tp_rank % tp_size_per_node) // dp_group_width
|
||||||
listen_port = get_serving().port + local_dp_rank
|
listen_port = get_serving().port + local_dp_rank
|
||||||
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str()
|
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str()
|
||||||
|
|
||||||
@@ -179,7 +179,7 @@ class RustServer:
|
|||||||
# Under DP every rank runs its own server on its own port, so the rank is
|
# Under DP every rank runs its own server on its own port, so the rank is
|
||||||
# what tells two otherwise identical startup lines apart.
|
# what tells two otherwise identical startup lines apart.
|
||||||
dp_note = (
|
dp_note = (
|
||||||
"" if dp_rank is None else f" (DP rank {dp_rank}/{scheduler.ps.dp_size})"
|
"" if dp_rank is None else f" (DP rank {dp_rank}/{get_parallel().dp_size})"
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
|
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
from dataclasses import replace
|
|
||||||
from typing import List, Optional, Tuple
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -19,7 +18,6 @@ from sglang.kernels.ops.speculative.dspark.dspark_accept import (
|
|||||||
accept_sampling,
|
accept_sampling,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||||
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.logits_processor import should_apply_lm_head_quant_method
|
from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
@@ -364,7 +362,6 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -372,7 +369,6 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.ps = ps
|
|
||||||
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
|
||||||
@@ -416,7 +412,6 @@ 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,
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
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",
|
||||||
@@ -453,7 +448,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
validate_domino_runtime(
|
validate_domino_runtime(
|
||||||
device=torch.device(self.device),
|
device=torch.device(self.device),
|
||||||
tp_size=int(get_parallel().tp_group.world_size),
|
tp_size=int(get_parallel().tp_group.world_size),
|
||||||
tp_rank=int(self.ps.tp_rank),
|
tp_rank=int(self.model_runner.tp_rank),
|
||||||
target_vocab_size=int(self.model_runner.model_config.vocab_size),
|
target_vocab_size=int(self.model_runner.model_config.vocab_size),
|
||||||
draft_vocab_size=int(self.draft_model_runner.model_config.vocab_size),
|
draft_vocab_size=int(self.draft_model_runner.model_config.vocab_size),
|
||||||
hidden_size=int(self.draft_model.config.hidden_size),
|
hidden_size=int(self.draft_model.config.hidden_size),
|
||||||
@@ -500,7 +495,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
self._maybe_merge_trained_mask_embedding()
|
self._maybe_merge_trained_mask_embedding()
|
||||||
self._cache_full_embed_weight()
|
self._cache_full_embed_weight()
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.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,
|
||||||
@@ -650,7 +645,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
# shared graph capture/replay; keep the draft eager under dp
|
# shared graph capture/replay; keep the draft eager under dp
|
||||||
# attention.
|
# attention.
|
||||||
capture_decode_cuda_graph = False
|
capture_decode_cuda_graph = False
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Disable DFLASH draft cuda graph because dp attention "
|
"Disable DFLASH draft cuda graph because dp attention "
|
||||||
"is enabled (draft runs eager)."
|
"is enabled (draft runs eager)."
|
||||||
@@ -795,7 +790,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
def _maybe_build_draft_sampler(self):
|
def _maybe_build_draft_sampler(self):
|
||||||
def _eager(reason):
|
def _eager(reason):
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
|
logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -819,7 +814,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
):
|
):
|
||||||
return _eager("unsupported quantized lm_head")
|
return _eager("unsupported quantized lm_head")
|
||||||
self.draft_model.lm_head = lm_head
|
self.draft_model.lm_head = lm_head
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DFLASH selector decode folded into the draft cuda graph "
|
"DFLASH selector decode folded into the draft cuda graph "
|
||||||
"(sampling_enabled=%s).",
|
"(sampling_enabled=%s).",
|
||||||
@@ -843,7 +838,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
embed_proj = self.draft_model.embed_proj
|
embed_proj = self.draft_model.embed_proj
|
||||||
if prefix_gru is None or embed_proj is None:
|
if prefix_gru is None or embed_proj is None:
|
||||||
return _eager("Domino projector modules are unavailable")
|
return _eager("Domino projector modules are unavailable")
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DFLASH Domino rollout folded into the draft cuda graph (tp=%s).",
|
"DFLASH Domino rollout folded into the draft cuda graph (tp=%s).",
|
||||||
int(tp_group.world_size),
|
int(tp_group.world_size),
|
||||||
@@ -882,7 +877,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
return _eager("added vocab")
|
return _eager("added vocab")
|
||||||
num_org = int(shard.num_org_elements)
|
num_org = int(shard.num_org_elements)
|
||||||
org_vocab_start = int(shard.org_vocab_start_index)
|
org_vocab_start = int(shard.org_vocab_start_index)
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DFLASH draft greedy head folded into the draft cuda graph (tp=%d).",
|
"DFLASH draft greedy head folded into the draft cuda graph (tp=%d).",
|
||||||
tp_group.world_size,
|
tp_group.world_size,
|
||||||
@@ -908,7 +903,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
fused_disable_reason = "draft model does not support fused context KV"
|
fused_disable_reason = "draft model does not support fused context KV"
|
||||||
|
|
||||||
if fused_disable_reason is not None:
|
if fused_disable_reason is not None:
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.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,
|
||||||
@@ -951,7 +946,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
break
|
break
|
||||||
|
|
||||||
if fused_disable_reason is not None:
|
if fused_disable_reason is not None:
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.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,
|
||||||
@@ -973,7 +968,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.ps.tp_rank == 0:
|
if self.model_runner.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",
|
||||||
@@ -1266,7 +1261,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
embedding_tensor.to(embed_module.weight.dtype)
|
embedding_tensor.to(embed_module.weight.dtype)
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Merged trained mask embedding into target model "
|
"Merged trained mask embedding into target model "
|
||||||
"(mask_token_id=%s, source=%s)",
|
"(mask_token_id=%s, source=%s)",
|
||||||
@@ -1301,7 +1296,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
parts = [torch.empty_like(shard_t) for _ in range(tp_size)]
|
parts = [torch.empty_like(shard_t) for _ in range(tp_size)]
|
||||||
dist.all_gather(parts, shard_t, group=tp_group.device_group)
|
dist.all_gather(parts, shard_t, group=tp_group.device_group)
|
||||||
self._full_embed_gpu = torch.cat(parts, dim=0)[:vocab_size]
|
self._full_embed_gpu = torch.cat(parts, dim=0)[:vocab_size]
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DFLASH cached full embed on GPU for dp attention: shape=%s",
|
"DFLASH cached full embed on GPU for dp attention: shape=%s",
|
||||||
list(self._full_embed_gpu.shape),
|
list(self._full_embed_gpu.shape),
|
||||||
@@ -1364,7 +1359,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.ps.tp_rank == 0:
|
if added and self.model_runner.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,
|
||||||
@@ -2179,7 +2174,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
if self.selector is not None:
|
if self.selector is not None:
|
||||||
if self._selector_sampling_enabled:
|
if self._selector_sampling_enabled:
|
||||||
return
|
return
|
||||||
if not self._warned_sampling_fallback and self.ps.tp_rank == 0:
|
if not self._warned_sampling_fallback and self.model_runner.tp_rank == 0:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"DFLASH non-greedy verification is unavailable on this "
|
"DFLASH non-greedy verification is unavailable on this "
|
||||||
"build/device; falling back to greedy argmax verification. "
|
"build/device; falling back to greedy argmax verification. "
|
||||||
@@ -2192,7 +2187,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.ps.tp_rank == 0
|
and self.model_runner.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; "
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ 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__)
|
||||||
@@ -63,7 +62,6 @@ def build_draft_tp_worker(
|
|||||||
*,
|
*,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_model_config: ModelConfig,
|
target_model_config: ModelConfig,
|
||||||
algo_label: str,
|
algo_label: str,
|
||||||
@@ -88,7 +86,6 @@ def build_draft_tp_worker(
|
|||||||
draft_worker = draft_worker_cls(
|
draft_worker = draft_worker_cls(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
ps=ps,
|
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
is_draft_worker=True,
|
||||||
random_seed=random_seed,
|
random_seed=random_seed,
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from dataclasses import replace
|
|
||||||
from typing import Callable, Optional, Protocol, runtime_checkable
|
from typing import Callable, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -9,8 +8,6 @@ from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
|||||||
is_unified_kv_triton,
|
is_unified_kv_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||||
from sglang.srt.distributed import get_pp_group
|
|
||||||
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.logprob_processor import compute_spec_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
|
||||||
from sglang.srt.lora.layers import unwrap_lora_layer
|
from sglang.srt.lora.layers import unwrap_lora_layer
|
||||||
@@ -134,7 +131,6 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
draft_worker_cls: type[TpModelWorker] = TpModelWorker,
|
draft_worker_cls: type[TpModelWorker] = TpModelWorker,
|
||||||
@@ -143,14 +139,13 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.ps = ps
|
|
||||||
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
|
||||||
self.page_size = get_schedule().page_size
|
self.page_size = get_schedule().page_size
|
||||||
self.device = target_worker.device
|
self.device = target_worker.device
|
||||||
self._draft_worker = None
|
self._draft_worker = None
|
||||||
self._hosts_draft = get_pp_group().is_last_rank
|
self._hosts_draft = get_parallel().pp_group.is_last_rank
|
||||||
if not self._hosts_draft:
|
if not self._hosts_draft:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -166,7 +161,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
if (
|
if (
|
||||||
get_parallel().enable_dp_attention
|
get_parallel().enable_dp_attention
|
||||||
and self._draft_is_moe
|
and self._draft_is_moe
|
||||||
and ps.attn_tp_size > 1
|
and get_parallel().attn_tp_size > 1
|
||||||
):
|
):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DSpark + dp attention with a DeepSeek-V4 (MoE) draft requires "
|
"DSpark + dp attention with a DeepSeek-V4 (MoE) draft requires "
|
||||||
@@ -178,7 +173,6 @@ 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,
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
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",
|
||||||
@@ -232,7 +226,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
else parallel.tp_group
|
else parallel.tp_group
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.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, query_token_num=%s, "
|
"gamma=%s, verify_num_draft_tokens=%s, query_token_num=%s, "
|
||||||
@@ -255,7 +249,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if getattr(self.draft_model, "uses_own_vocab_modules", False):
|
if getattr(self.draft_model, "uses_own_vocab_modules", False):
|
||||||
if self.ps.tp_rank == 0:
|
if self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DSpark draft uses its checkpoint-local embedding and LM head."
|
"DSpark draft uses its checkpoint-local embedding and LM head."
|
||||||
)
|
)
|
||||||
@@ -279,7 +273,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.ps.tp_rank,
|
tp_rank=self.model_runner.tp_rank,
|
||||||
verify_num_draft_tokens=self.verify_num_draft_tokens,
|
verify_num_draft_tokens=self.verify_num_draft_tokens,
|
||||||
tp_sync=self._tp_sync,
|
tp_sync=self._tp_sync,
|
||||||
)
|
)
|
||||||
@@ -327,7 +321,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
and self._verify_planner.mode_value == "static"
|
and self._verify_planner.mode_value == "static"
|
||||||
and self._draft_is_moe
|
and self._draft_is_moe
|
||||||
and not get_parallel().enable_dp_attention
|
and not get_parallel().enable_dp_attention
|
||||||
and self.ps.pp_size == 1
|
and self.model_runner.pp_size == 1
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
(self._verify_planner.is_compact_mode or static_epilogue_supported)
|
(self._verify_planner.is_compact_mode or static_epilogue_supported)
|
||||||
@@ -391,7 +385,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.ps.tp_rank,
|
tp_rank=self.model_runner.tp_rank,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
simulate_acc_len=self._simulate_acc_len,
|
simulate_acc_len=self._simulate_acc_len,
|
||||||
)
|
)
|
||||||
@@ -451,7 +445,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
draft_model=self.draft_model,
|
draft_model=self.draft_model,
|
||||||
is_deepseek_v4_draft=self._draft_is_moe,
|
is_deepseek_v4_draft=self._draft_is_moe,
|
||||||
)
|
)
|
||||||
if self._target_hidden_projection_enabled and self.ps.tp_rank == 0:
|
if self._target_hidden_projection_enabled and self.model_runner.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"DSpark prefill target-hidden projection runs before "
|
"DSpark prefill target-hidden projection runs before "
|
||||||
"sequence-parallel gather."
|
"sequence-parallel gather."
|
||||||
@@ -508,7 +502,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
gamma=self.gamma,
|
gamma=self.gamma,
|
||||||
max_bs=max(get_exec().graph.cuda_graph_config.decode.bs),
|
max_bs=max(get_exec().graph.cuda_graph_config.decode.bs),
|
||||||
device=self.device,
|
device=self.device,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.model_runner.tp_rank,
|
||||||
tp_sync=self._tp_sync,
|
tp_sync=self._tp_sync,
|
||||||
available_memory_gb=available_memory_gb,
|
available_memory_gb=available_memory_gb,
|
||||||
confidence_fn=(
|
confidence_fn=(
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_co
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
|
deployment_attn_dp_size,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
@@ -114,8 +115,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
# Fields the parent's capture() reads:
|
# Fields the parent's capture() reads:
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.tp_size
|
||||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
self.attn_dp_size = deployment_attn_dp_size()
|
||||||
self.pp_size = get_parallel().pp_size
|
self.pp_size = get_parallel().pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
|
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import torch
|
|||||||
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
|
deployment_attn_dp_size,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
@@ -111,8 +112,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
# Fields the parent's capture() reads:
|
# Fields the parent's capture() reads:
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.tp_size
|
||||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
self.attn_dp_size = deployment_attn_dp_size()
|
||||||
self.pp_size = get_parallel().pp_size
|
self.pp_size = get_parallel().pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
|
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
|
||||||
|
|||||||
@@ -1,14 +1,12 @@
|
|||||||
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.topk1 import draft_topk1_postprocess
|
from sglang.kernels.ops.speculative.topk1 import draft_topk1_postprocess
|
||||||
from sglang.srt.configs.model_config import get_dsa_mtp_topk_width
|
from sglang.srt.configs.model_config import get_dsa_mtp_topk_width
|
||||||
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,
|
||||||
@@ -238,7 +236,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -247,7 +244,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
# 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.ps = ps
|
|
||||||
self.nccl_port = nccl_port
|
self.nccl_port = nccl_port
|
||||||
self.target_worker = target_worker
|
self.target_worker = target_worker
|
||||||
|
|
||||||
@@ -288,8 +284,6 @@ 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,
|
||||||
# spec workers don't support pipeline parallelism
|
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
is_draft_worker=True,
|
||||||
# The draft runs at absolute target positions.
|
# The draft runs at absolute target positions.
|
||||||
@@ -1293,7 +1287,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -1304,7 +1297,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
self.topk = get_spec().speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.ps = ps
|
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.device = get_device().device
|
self.device = get_device().device
|
||||||
self._target_worker = target_worker
|
self._target_worker = target_worker
|
||||||
@@ -1320,7 +1312,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
EagleDraftWorker(
|
EagleDraftWorker(
|
||||||
server_args,
|
server_args,
|
||||||
gpu_id,
|
gpu_id,
|
||||||
ps,
|
|
||||||
nccl_port,
|
nccl_port,
|
||||||
target_worker,
|
target_worker,
|
||||||
)
|
)
|
||||||
@@ -1873,7 +1864,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.ps.tp_rank]
|
recv_req.serialized_named_tensors[self.model_runner.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(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import torch
|
|||||||
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
|
deployment_attn_dp_size,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
@@ -98,8 +99,8 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync()
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
self.tp_size = self.model_runner.ps.tp_size
|
self.tp_size = self.model_runner.tp_size
|
||||||
self.attn_dp_size = self.model_runner.ps.attn_dp_size
|
self.attn_dp_size = deployment_attn_dp_size()
|
||||||
self.pp_size = get_parallel().pp_size
|
self.pp_size = get_parallel().pp_size
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.topk = get_spec().speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
|
|||||||
@@ -23,12 +23,10 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
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 (
|
from sglang.srt.layers.moe.utils import (
|
||||||
draft_model_build_scope,
|
draft_model_build_scope,
|
||||||
speculative_moe_a2a_backend_context,
|
speculative_moe_a2a_backend_context,
|
||||||
@@ -102,7 +100,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -112,7 +109,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
|||||||
self.topk = get_spec().speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.ps = ps
|
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.device = get_device().device
|
self.device = get_device().device
|
||||||
self.target_worker = target_worker
|
self.target_worker = target_worker
|
||||||
@@ -146,8 +142,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
|||||||
self,
|
self,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
# spec workers don't support pipeline parallelism
|
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
is_draft_worker=True,
|
||||||
# The draft runs at absolute target positions.
|
# The draft runs at absolute target positions.
|
||||||
@@ -702,7 +696,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -715,7 +708,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
self.topk = get_spec().speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.ps = ps
|
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.device = get_device().device
|
self.device = get_device().device
|
||||||
self._target_worker = target_worker
|
self._target_worker = target_worker
|
||||||
@@ -730,7 +722,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
self._draft_worker = FrozenKVMTPDraftWorker(
|
self._draft_worker = FrozenKVMTPDraftWorker(
|
||||||
server_args,
|
server_args,
|
||||||
gpu_id,
|
gpu_id,
|
||||||
ps,
|
|
||||||
nccl_port,
|
nccl_port,
|
||||||
target_worker,
|
target_worker,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -155,7 +155,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
# Fields the parent's capture() reads:
|
# Fields the parent's capture() reads:
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.tp_size
|
||||||
self.dp_size = get_parallel().dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = get_parallel().pp_size
|
self.pp_size = get_parallel().pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ from typing import TYPE_CHECKING, List
|
|||||||
|
|
||||||
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.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,
|
||||||
@@ -121,7 +120,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -130,7 +128,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
# 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.ps = ps
|
|
||||||
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 = []
|
||||||
@@ -171,8 +168,6 @@ 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,
|
||||||
# spec workers don't support pipeline parallelism
|
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
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,
|
||||||
@@ -1013,7 +1008,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -1035,7 +1029,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
self._draft_worker = MultiLayerEagleDraftWorker(
|
self._draft_worker = MultiLayerEagleDraftWorker(
|
||||||
server_args,
|
server_args,
|
||||||
gpu_id,
|
gpu_id,
|
||||||
ps,
|
|
||||||
nccl_port,
|
nccl_port,
|
||||||
target_worker,
|
target_worker,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import torch
|
|||||||
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.logprob_processor import compute_spec_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_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
|
||||||
@@ -89,7 +88,6 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -99,7 +97,7 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
self.enable_overlap = not get_schedule().disable_overlap_schedule
|
self.enable_overlap = not get_schedule().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 = ps.tp_rank
|
self.tp_rank = self.model_runner.tp_rank
|
||||||
self.page_size = get_schedule().page_size
|
self.page_size = get_schedule().page_size
|
||||||
self.draft_token_num: int = get_spec().speculative_num_draft_tokens
|
self.draft_token_num: int = get_spec().speculative_num_draft_tokens
|
||||||
self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
|
self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
|
||||||
|
|||||||
@@ -1,10 +1,8 @@
|
|||||||
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 (
|
from sglang.srt.layers.moe.utils import (
|
||||||
draft_model_build_scope,
|
draft_model_build_scope,
|
||||||
speculative_moe_backend_context,
|
speculative_moe_backend_context,
|
||||||
@@ -45,7 +43,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -54,7 +51,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
|
|||||||
# 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.ps = ps
|
|
||||||
self.nccl_port = nccl_port
|
self.nccl_port = nccl_port
|
||||||
self.target_worker = target_worker
|
self.target_worker = target_worker
|
||||||
|
|
||||||
@@ -84,8 +80,6 @@ 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,
|
||||||
# spec workers don't support pipeline parallelism
|
|
||||||
ps=replace(ps, pp_rank=0, pp_size=1),
|
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
is_draft_worker=True,
|
||||||
# The draft runs at absolute target positions.
|
# The draft runs at absolute target positions.
|
||||||
@@ -167,7 +161,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -190,7 +183,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
|
|||||||
self._draft_worker = StandaloneDraftWorker(
|
self._draft_worker = StandaloneDraftWorker(
|
||||||
server_args,
|
server_args,
|
||||||
gpu_id,
|
gpu_id,
|
||||||
ps,
|
|
||||||
nccl_port,
|
nccl_port,
|
||||||
target_worker,
|
target_worker,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -45,8 +45,8 @@ def init_uno_lora_manager(
|
|||||||
dtype=model_runner.dtype,
|
dtype=model_runner.dtype,
|
||||||
server_args=model_runner.server_args,
|
server_args=model_runner.server_args,
|
||||||
lora_backend="uno_cublas", # fast path
|
lora_backend="uno_cublas", # fast path
|
||||||
tp_size=model_runner.ps.tp_size,
|
tp_size=model_runner.tp_size,
|
||||||
tp_rank=model_runner.ps.tp_rank,
|
tp_rank=model_runner.tp_rank,
|
||||||
# Infer these from the one trained adapter.
|
# Infer these from the one trained adapter.
|
||||||
max_lora_rank=None,
|
max_lora_rank=None,
|
||||||
target_modules=None,
|
target_modules=None,
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ from sglang.srt.utils.common import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
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.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
|
||||||
@@ -62,7 +61,6 @@ class UnoWorkerV2(BaseSpecWorker):
|
|||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
ps: ParallelState,
|
|
||||||
nccl_port: int,
|
nccl_port: int,
|
||||||
target_worker: TpModelWorker,
|
target_worker: TpModelWorker,
|
||||||
):
|
):
|
||||||
@@ -70,7 +68,6 @@ class UnoWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.ps = ps
|
|
||||||
self.nccl_port = nccl_port
|
self.nccl_port = nccl_port
|
||||||
|
|
||||||
self._target_worker = target_worker
|
self._target_worker = target_worker
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ from typing import Callable, Dict, List, 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.io_struct import ProfileReqOutput
|
from sglang.srt.managers.io_struct import ProfileReqOutput
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
@@ -71,14 +70,13 @@ def export_cuda_graph_capture_trace(prof_context, *, runner_name: str, tp_rank:
|
|||||||
|
|
||||||
|
|
||||||
class ProfileManager:
|
class ProfileManager:
|
||||||
def __init__(self, ps: ParallelState, cpu_group):
|
def __init__(self, cpu_group):
|
||||||
self.stage_based_trigger = _StageBasedTrigger(
|
self.stage_based_trigger = _StageBasedTrigger(
|
||||||
on_start=self._do_start,
|
on_start=self._do_start,
|
||||||
on_stop=self._do_stop,
|
on_stop=self._do_stop,
|
||||||
)
|
)
|
||||||
self.ps = ps
|
|
||||||
self.cpu_group = cpu_group
|
self.cpu_group = cpu_group
|
||||||
self.first_rank_in_node = ps.gpu_id == get_device().base_gpu_id
|
self.first_rank_in_node = get_device().gpu_id == get_device().base_gpu_id
|
||||||
self.profiler_kwargs = None
|
self.profiler_kwargs = None
|
||||||
self.profiler = None
|
self.profiler = None
|
||||||
self.detailed_annotations = False
|
self.detailed_annotations = False
|
||||||
@@ -154,7 +152,6 @@ class ProfileManager:
|
|||||||
set_detailed_annotations_enabled(self.detailed_annotations)
|
set_detailed_annotations_enabled(self.detailed_annotations)
|
||||||
self.profiler = _ProfilerBase.create(
|
self.profiler = _ProfilerBase.create(
|
||||||
**self.profiler_kwargs,
|
**self.profiler_kwargs,
|
||||||
ps=self.ps,
|
|
||||||
cpu_group=self.cpu_group,
|
cpu_group=self.cpu_group,
|
||||||
first_rank_in_node=self.first_rank_in_node,
|
first_rank_in_node=self.first_rank_in_node,
|
||||||
output_suffix=f"-{stage}" if stage else "",
|
output_suffix=f"-{stage}" if stage else "",
|
||||||
@@ -303,7 +300,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
|
|||||||
output_prefix: str,
|
output_prefix: str,
|
||||||
output_suffix: str,
|
output_suffix: str,
|
||||||
profile_id: str,
|
profile_id: str,
|
||||||
ps: ParallelState,
|
|
||||||
cpu_group,
|
cpu_group,
|
||||||
first_rank_in_node: bool,
|
first_rank_in_node: bool,
|
||||||
):
|
):
|
||||||
@@ -311,7 +307,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
|
|||||||
self.output_prefix = output_prefix
|
self.output_prefix = output_prefix
|
||||||
self.output_suffix = output_suffix
|
self.output_suffix = output_suffix
|
||||||
self.profile_id = profile_id
|
self.profile_id = profile_id
|
||||||
self.ps = ps
|
|
||||||
self.cpu_group = cpu_group
|
self.cpu_group = cpu_group
|
||||||
self.first_rank_in_node = first_rank_in_node
|
self.first_rank_in_node = first_rank_in_node
|
||||||
|
|
||||||
|
|||||||
@@ -274,15 +274,14 @@ class WeightCacheDaemon:
|
|||||||
from sglang.srt.model_loader.loader import get_model_loader
|
from sglang.srt.model_loader.loader import get_model_loader
|
||||||
|
|
||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
# The launcher told this daemon where it sits, and it builds the same
|
|
||||||
# groups a scheduler does, so the same one number places it.
|
|
||||||
publish(
|
publish(
|
||||||
server_args,
|
server_args,
|
||||||
role="weight_cache_daemon",
|
role="weight_cache_daemon",
|
||||||
ranks=SpawnRanks(
|
ranks=SpawnRanks(
|
||||||
world_rank=spawn_world_rank(
|
world_rank=spawn_world_rank(
|
||||||
server_args, tp_rank=self.tp_rank, pp_rank=self.pp_rank
|
server_args, tp_rank=self.tp_rank, pp_rank=self.pp_rank
|
||||||
)
|
),
|
||||||
|
gpu_id=self.gpu_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -503,17 +503,17 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
# Build engine's config fingerprint
|
# Build engine's config fingerprint
|
||||||
from sglang.srt.layers.dp_attention import get_moe_cp_size
|
from sglang.srt.layers.dp_attention import get_moe_cp_size
|
||||||
|
|
||||||
ps = get_parallel()
|
parallel = get_parallel()
|
||||||
tp_size = ps.tp_size
|
tp_size = parallel.tp_size
|
||||||
tp_rank = ps.tp_rank
|
tp_rank = parallel.tp_rank
|
||||||
|
|
||||||
pp_size = ps.pp_size
|
pp_size = parallel.pp_size
|
||||||
pp_rank = ps.pp_rank
|
pp_rank = parallel.pp_rank
|
||||||
|
|
||||||
ep_size = ps.moe_ep_size
|
ep_size = parallel.moe_ep_size
|
||||||
moe_dp_size = get_moe_cp_size()
|
moe_dp_size = get_moe_cp_size()
|
||||||
moe_dp_rank = ps.moe_dp_rank
|
moe_dp_rank = parallel.moe_dp_rank
|
||||||
moe_ep_rank = ps.moe_ep_rank
|
moe_ep_rank = parallel.moe_ep_rank
|
||||||
|
|
||||||
dp_size = get_parallel().dp_size
|
dp_size = get_parallel().dp_size
|
||||||
|
|
||||||
@@ -535,10 +535,10 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
moe_dp_size=moe_dp_size,
|
moe_dp_size=moe_dp_size,
|
||||||
moe_dp_rank=moe_dp_rank,
|
moe_dp_rank=moe_dp_rank,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
enable_dp_attention=ps.enable_dp_attention,
|
enable_dp_attention=parallel.enable_dp_attention,
|
||||||
enable_dp_lm_head=ps.enable_dp_lm_head,
|
enable_dp_lm_head=parallel.enable_dp_lm_head,
|
||||||
attn_cp_size=ps.attn_cp_size,
|
attn_cp_size=parallel.attn_cp_size,
|
||||||
moe_dense_tp_size=ps.moe_dense_tp_size,
|
moe_dense_tp_size=parallel.moe_dense_tp_size,
|
||||||
moe_a2a_backend=get_exec().moe.moe_a2a_backend,
|
moe_a2a_backend=get_exec().moe.moe_a2a_backend,
|
||||||
quant_method=quant_method,
|
quant_method=quant_method,
|
||||||
quant_config_hash=hash_quant_config(quant_config),
|
quant_config_hash=hash_quant_config(quant_config),
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ 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
|
||||||
@@ -339,7 +338,6 @@ 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.max_running_requests = pool_batch_size
|
self.max_running_requests = pool_batch_size
|
||||||
# trtllm_mha __init__ scans model.modules() for ENCODER_ONLY layers;
|
# trtllm_mha __init__ scans model.modules() for ENCODER_ONLY layers;
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ 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
|
||||||
@@ -322,7 +321,6 @@ 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,
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ from torch import nn
|
|||||||
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
|
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
|
||||||
quant_to_nope_fp8_rope_bf16_pack_triton,
|
quant_to_nope_fp8_rope_bf16_pack_triton,
|
||||||
)
|
)
|
||||||
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.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -349,7 +348,6 @@ 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,
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ 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,
|
||||||
@@ -228,7 +227,6 @@ class MockGDNModelRunner(ModelRunner):
|
|||||||
self.decode_attention_backend_str = case.backend
|
self.decode_attention_backend_str = case.backend
|
||||||
self.draft_attention_backend = None
|
self.draft_attention_backend = None
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
|
||||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ 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,
|
||||||
@@ -231,7 +230,6 @@ class MockKDAModelRunner(ModelRunner):
|
|||||||
self.decode_attention_backend_str = case.backend
|
self.decode_attention_backend_str = case.backend
|
||||||
self.draft_attention_backend = None
|
self.draft_attention_backend = None
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
|
||||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ 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,
|
||||||
@@ -239,7 +238,6 @@ class MockLightningModelRunner(ModelRunner):
|
|||||||
self.decode_attention_backend_str = case.backend
|
self.decode_attention_backend_str = case.backend
|
||||||
self.draft_attention_backend = None
|
self.draft_attention_backend = None
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
|
||||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -26,7 +26,6 @@ from sglang.srt.configs.mamba_utils import ( # noqa: E402
|
|||||||
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,
|
||||||
)
|
)
|
||||||
@@ -325,7 +324,6 @@ class MockMamba2ModelRunner(ModelRunner):
|
|||||||
self.decode_attention_backend_str = case.backend
|
self.decode_attention_backend_str = case.backend
|
||||||
self.draft_attention_backend = None
|
self.draft_attention_backend = None
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
|
||||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ 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
|
||||||
@@ -247,7 +246,6 @@ 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()
|
|
||||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
speculative_num_draft_tokens = (
|
speculative_num_draft_tokens = (
|
||||||
max(case.input_lens)
|
max(case.input_lens)
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import TYPE_CHECKING, Dict, Iterator, List, Optional
|
from typing import TYPE_CHECKING, Dict, Iterator, List, Optional
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.test.scripted_runtime.context.api import ScriptedContext
|
from sglang.test.scripted_runtime.context.api import ScriptedContext
|
||||||
@@ -12,7 +14,7 @@ def _get_all_reqs(ctx: ScriptedContext) -> Iterator[Req]:
|
|||||||
if s.chunked_req is not None:
|
if s.chunked_req is not None:
|
||||||
yield s.chunked_req
|
yield s.chunked_req
|
||||||
yield from s.waiting_queue
|
yield from s.waiting_queue
|
||||||
if s.ps.pp_size > 1:
|
if get_parallel().pp_size > 1:
|
||||||
for mb in (*s.mbs, *s.last_mbs, *s.running_mbs):
|
for mb in (*s.mbs, *s.last_mbs, *s.running_mbs):
|
||||||
if mb is not None:
|
if mb is not None:
|
||||||
yield from mb.reqs
|
yield from mb.reqs
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import zmq
|
|||||||
from sglang.srt.arg_groups.overrides import resolving_view
|
from sglang.srt.arg_groups.overrides import resolving_view
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils.network import get_zmq_socket
|
from sglang.srt.utils.network import get_zmq_socket
|
||||||
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
|
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
|
||||||
from sglang.test.scripted_runtime.context import ScriptedContext
|
from sglang.test.scripted_runtime.context import ScriptedContext
|
||||||
@@ -125,9 +126,9 @@ class ScriptedSchedulerHook:
|
|||||||
) -> None:
|
) -> None:
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self._is_driver = (
|
self._is_driver = (
|
||||||
scheduler.ps.pp_rank == 0
|
get_parallel().pp_rank == 0
|
||||||
and scheduler.ps.tp_rank == 0
|
and get_parallel().tp_rank == 0
|
||||||
and scheduler.ps.attn_cp_rank == 0
|
and get_parallel().attn_cp_rank == 0
|
||||||
)
|
)
|
||||||
self._batch_log: List[ScriptedBatchRecord] = []
|
self._batch_log: List[ScriptedBatchRecord] = []
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ 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.forward_context import (
|
from sglang.srt.model_executor.forward_context import (
|
||||||
@@ -65,7 +64,6 @@ 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,
|
||||||
ps=ParallelState.trivial(tp_size=cls.tp_size),
|
|
||||||
nccl_port=cls.port_args.nccl_port,
|
nccl_port=cls.port_args.nccl_port,
|
||||||
server_args=cls.server_args,
|
server_args=cls.server_args,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ 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 (
|
||||||
@@ -151,7 +150,6 @@ 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,
|
||||||
ps=ParallelState.trivial(),
|
|
||||||
nccl_port=12435,
|
nccl_port=12435,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ from sglang.srt.disaggregation.decode import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver
|
from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver
|
||||||
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.srt.runtime_context import get_context, publish, reset_context
|
from sglang.srt.runtime_context import get_context, publish, reset_context
|
||||||
@@ -440,7 +439,6 @@ 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 = 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=[])
|
||||||
|
|||||||
@@ -131,7 +131,6 @@ class TestHandlePdRoleSwitch(unittest.TestCase):
|
|||||||
def test_rejected_when_decode_graph_headroom_is_insufficient(self):
|
def test_rejected_when_decode_graph_headroom_is_insufficient(self):
|
||||||
s = self._scheduler(DisaggregationMode.PREFILL)
|
s = self._scheduler(DisaggregationMode.PREFILL)
|
||||||
s.device = "cuda"
|
s.device = "cuda"
|
||||||
s.ps = SimpleNamespace(gpu_id=0)
|
|
||||||
s.tp_worker.get_decode_cuda_graph_bs.return_value = []
|
s.tp_worker.get_decode_cuda_graph_bs.return_value = []
|
||||||
with patch.object(role_switch, "get_available_gpu_memory", return_value=0.5):
|
with patch.object(role_switch, "get_available_gpu_memory", return_value=0.5):
|
||||||
out = Scheduler.handle_pd_role_switch(
|
out = Scheduler.handle_pd_role_switch(
|
||||||
@@ -151,7 +150,6 @@ class TestHandlePdRoleSwitch(unittest.TestCase):
|
|||||||
def test_decode_graph_headroom_allows_flip(self):
|
def test_decode_graph_headroom_allows_flip(self):
|
||||||
s = self._scheduler(DisaggregationMode.PREFILL)
|
s = self._scheduler(DisaggregationMode.PREFILL)
|
||||||
s.device = "cuda"
|
s.device = "cuda"
|
||||||
s.ps = SimpleNamespace(gpu_id=0)
|
|
||||||
s.tp_worker.get_decode_cuda_graph_bs.return_value = []
|
s.tp_worker.get_decode_cuda_graph_bs.return_value = []
|
||||||
with patch.object(role_switch, "get_available_gpu_memory", return_value=1.0):
|
with patch.object(role_switch, "get_available_gpu_memory", return_value=1.0):
|
||||||
out = Scheduler.handle_pd_role_switch(
|
out = Scheduler.handle_pd_role_switch(
|
||||||
|
|||||||
@@ -43,18 +43,18 @@ def test_dp_leaders_reuse_node_local_ports(
|
|||||||
for dp_rank, tp_rank in enumerate(ranks):
|
for dp_rank, tp_rank in enumerate(ranks):
|
||||||
scheduler = SimpleNamespace(
|
scheduler = SimpleNamespace(
|
||||||
server_args=SimpleNamespace(),
|
server_args=SimpleNamespace(),
|
||||||
ps=SimpleNamespace(
|
|
||||||
tp_rank=tp_rank,
|
|
||||||
tp_size=parallel.tp_size,
|
|
||||||
pp_size=parallel.pp_size,
|
|
||||||
attn_tp_size=parallel.attn_tp_size,
|
|
||||||
attn_cp_size=parallel.attn_cp_size,
|
|
||||||
attn_dp_rank=dp_rank,
|
|
||||||
dp_size=dp_size,
|
|
||||||
),
|
|
||||||
model_config=SimpleNamespace(is_multimodal=False),
|
model_config=SimpleNamespace(is_multimodal=False),
|
||||||
)
|
)
|
||||||
ports.append(rust_server.RustServer.launch(scheduler).http_port)
|
# Where this rank sits, stated whole: the attention rank follows
|
||||||
|
# from the TP rank and the attention-TP width, and the identities
|
||||||
|
# refuse the combination if it describes no real layout.
|
||||||
|
with parallel.override(
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
attn_dp_rank=dp_rank,
|
||||||
|
attn_tp_rank=tp_rank % parallel.attn_tp_size,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
):
|
||||||
|
ports.append(rust_server.RustServer.launch(scheduler).http_port)
|
||||||
|
|
||||||
calls = extension.return_value.Server.call_args_list
|
calls = extension.return_value.Server.call_args_list
|
||||||
assert [c.kwargs["port_offset"] for c in calls] == expected
|
assert [c.kwargs["port_offset"] for c in calls] == expected
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ _HAS_MLX = importlib.util.find_spec("mlx") is not None
|
|||||||
_SKIP_REASON = "requires mlx"
|
_SKIP_REASON = "requires mlx"
|
||||||
|
|
||||||
if _HAS_MLX:
|
if _HAS_MLX:
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
|
||||||
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
|
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
|
||||||
MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO as RATIO,
|
MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO as RATIO,
|
||||||
)
|
)
|
||||||
@@ -62,7 +61,7 @@ def _stub_for_initialize(
|
|||||||
stub = MlxModelRunnerStub.__new__(MlxModelRunnerStub)
|
stub = MlxModelRunnerStub.__new__(MlxModelRunnerStub)
|
||||||
stub._mlx_pool_size = pool_size
|
stub._mlx_pool_size = pool_size
|
||||||
stub.device = "cpu"
|
stub.device = "cpu"
|
||||||
stub.ps = ParallelState.trivial(dp_size=dp_size, attn_dp_size=attn_dp_size)
|
stub.attn_dp_size = attn_dp_size
|
||||||
stub.server_args = server_args
|
stub.server_args = server_args
|
||||||
stub.model_config = SimpleNamespace(
|
stub.model_config = SimpleNamespace(
|
||||||
is_hybrid_swa=False,
|
is_hybrid_swa=False,
|
||||||
|
|||||||
@@ -208,14 +208,7 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase):
|
|||||||
SchedulerProfilerManager,
|
SchedulerProfilerManager,
|
||||||
)
|
)
|
||||||
|
|
||||||
class FakePS:
|
mgr = SchedulerProfilerManager(dp_tp_cpu_group=None, get_forward_ct=lambda: 0)
|
||||||
tp_rank = dp_rank = pp_rank = moe_ep_rank = 0
|
|
||||||
dp_size = pp_size = moe_ep_size = 1
|
|
||||||
gpu_id = 0
|
|
||||||
|
|
||||||
mgr = SchedulerProfilerManager(
|
|
||||||
ps=FakePS(), dp_tp_cpu_group=None, get_forward_ct=lambda: 0
|
|
||||||
)
|
|
||||||
mgr._init_profile(output_dir, None, None, None, None, None, False, "test")
|
mgr._init_profile(output_dir, None, None, None, None, None, False, "test")
|
||||||
return mgr
|
return mgr
|
||||||
|
|
||||||
|
|||||||
@@ -228,7 +228,7 @@ class TestMambaPrefillTrackMetadata(unittest.TestCase):
|
|||||||
prefill_attention_backend_str="torch_native",
|
prefill_attention_backend_str="torch_native",
|
||||||
ngram_embedding_manager=SimpleNamespace(enabled=False),
|
ngram_embedding_manager=SimpleNamespace(enabled=False),
|
||||||
lora_manager=None,
|
lora_manager=None,
|
||||||
ps=SimpleNamespace(attn_dcp_size=1),
|
attn_dcp_size=1,
|
||||||
attn_backend=SimpleNamespace(
|
attn_backend=SimpleNamespace(
|
||||||
get_cpu_graph_seq_len_fill_value=lambda: 1,
|
get_cpu_graph_seq_len_fill_value=lambda: 1,
|
||||||
get_cuda_graph_seq_len_fill_value=lambda: 1,
|
get_cuda_graph_seq_len_fill_value=lambda: 1,
|
||||||
|
|||||||
@@ -20,7 +20,11 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
|||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import (
|
||||||
|
CustomTestCase,
|
||||||
|
enter_scope,
|
||||||
|
published_topology,
|
||||||
|
)
|
||||||
|
|
||||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
@@ -58,6 +62,12 @@ def load_mlx_scheduler_module():
|
|||||||
|
|
||||||
|
|
||||||
class TestSchedulerIdleStepCounters(CustomTestCase):
|
class TestSchedulerIdleStepCounters(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
# The loop asks the context where this process sits; nothing here
|
||||||
|
# builds a process group, so the placement arrives by publishing one.
|
||||||
|
enter_scope(self, published_topology(role="scheduler"))
|
||||||
|
|
||||||
@parameterized.expand(
|
@parameterized.expand(
|
||||||
[
|
[
|
||||||
(
|
(
|
||||||
@@ -406,7 +416,6 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
|
|||||||
scheduler.forward_ct = 0
|
scheduler.forward_ct = 0
|
||||||
scheduler.processed_tokens_counter = 0
|
scheduler.processed_tokens_counter = 0
|
||||||
scheduler.spec_algorithm = SpeculativeAlgorithm.NONE
|
scheduler.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
scheduler.ps = SimpleNamespace(pp_rank=0, attn_tp_rank=0, attn_cp_rank=0)
|
|
||||||
scheduler._poll_timeout_aborts = Mock(return_value=[])
|
scheduler._poll_timeout_aborts = Mock(return_value=[])
|
||||||
scheduler.scheduler_stage_metrics = None
|
scheduler.scheduler_stage_metrics = None
|
||||||
scheduler.metrics_reporter = SimpleNamespace(record_scheduler_active=Mock())
|
scheduler.metrics_reporter = SimpleNamespace(record_scheduler_active=Mock())
|
||||||
@@ -455,7 +464,7 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
|
|||||||
return scheduler
|
return scheduler
|
||||||
|
|
||||||
def prepare_pp_scheduler(self, scheduler):
|
def prepare_pp_scheduler(self, scheduler):
|
||||||
scheduler.ps.pp_size = 2
|
enter_scope(self, get_parallel().override(pp_size=2, pp_rank=0))
|
||||||
scheduler.pp_group = SimpleNamespace(is_last_rank=True)
|
scheduler.pp_group = SimpleNamespace(is_last_rank=True)
|
||||||
scheduler.forward_stream_ctx = nullcontext()
|
scheduler.forward_stream_ctx = nullcontext()
|
||||||
scheduler.forward_stream = Mock()
|
scheduler.forward_stream = Mock()
|
||||||
|
|||||||
@@ -148,7 +148,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(
|
||||||
stream_interval=1,
|
stream_interval=1,
|
||||||
enable_request_time_stats_logging=False,
|
enable_request_time_stats_logging=False,
|
||||||
@@ -181,7 +180,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(
|
||||||
stream_interval=1,
|
stream_interval=1,
|
||||||
enable_request_time_stats_logging=False,
|
enable_request_time_stats_logging=False,
|
||||||
@@ -218,7 +216,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(
|
||||||
stream_interval=1,
|
stream_interval=1,
|
||||||
enable_request_time_stats_logging=False,
|
enable_request_time_stats_logging=False,
|
||||||
@@ -258,7 +255,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(
|
||||||
stream_interval=1,
|
stream_interval=1,
|
||||||
enable_request_time_stats_logging=False,
|
enable_request_time_stats_logging=False,
|
||||||
@@ -286,7 +282,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(),
|
server_args=SimpleNamespace(),
|
||||||
is_generation=True,
|
is_generation=True,
|
||||||
spec_algorithm=SpeculativeAlgorithm.NONE,
|
spec_algorithm=SpeculativeAlgorithm.NONE,
|
||||||
@@ -312,7 +307,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
streamer = Streamer(
|
streamer = Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
|
|
||||||
server_args=SimpleNamespace(),
|
server_args=SimpleNamespace(),
|
||||||
is_generation=True,
|
is_generation=True,
|
||||||
spec_algorithm=SpeculativeAlgorithm.NONE,
|
spec_algorithm=SpeculativeAlgorithm.NONE,
|
||||||
@@ -334,7 +328,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
Streamer(
|
Streamer(
|
||||||
send_to_detokenizer=SimpleNamespace(),
|
send_to_detokenizer=SimpleNamespace(),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
ps=SimpleNamespace(),
|
|
||||||
server_args=SimpleNamespace(),
|
server_args=SimpleNamespace(),
|
||||||
is_generation=True,
|
is_generation=True,
|
||||||
spec_algorithm=SpeculativeAlgorithm.NONE,
|
spec_algorithm=SpeculativeAlgorithm.NONE,
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
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.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_memory,
|
get_memory,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -36,10 +35,17 @@ def mock_cpu_env(kv_size=2, tp_size=1, swa_eviction_interval=4):
|
|||||||
|
|
||||||
with (
|
with (
|
||||||
patch("torch._utils._element_size", return_value=kv_size),
|
patch("torch._utils._element_size", return_value=kv_size),
|
||||||
# A width is a whole topology: state the TP siblings the identities
|
# The whole attention triple, not just one leaf: a width that does not
|
||||||
# relate it to, not the attention share alone.
|
# factor describes no layout.
|
||||||
get_parallel().override(
|
get_parallel().override(
|
||||||
tp_size=tp_size, attn_tp_size=tp_size, moe_tp_size=tp_size
|
tp_size=tp_size,
|
||||||
|
attn_tp_size=tp_size,
|
||||||
|
attn_dp_size=1,
|
||||||
|
attn_cp_size=1,
|
||||||
|
moe_ep_size=1,
|
||||||
|
moe_ep_group=None,
|
||||||
|
moe_dp_size=1,
|
||||||
|
moe_tp_size=tp_size,
|
||||||
),
|
),
|
||||||
envs.SGLANG_SWA_EVICTION_INTERVAL.override(swa_eviction_interval),
|
envs.SGLANG_SWA_EVICTION_INTERVAL.override(swa_eviction_interval),
|
||||||
):
|
):
|
||||||
@@ -174,7 +180,8 @@ def _make_model_runner(
|
|||||||
mr.layer_info = SimpleNamespace(
|
mr.layer_info = SimpleNamespace(
|
||||||
start_layer=0, end_layer=num_layers, num_effective_layers=num_layers
|
start_layer=0, end_layer=num_layers, num_effective_layers=num_layers
|
||||||
)
|
)
|
||||||
mr.ps = ParallelState.trivial()
|
mr.attn_dp_size = 1
|
||||||
|
mr.pp_size = 1
|
||||||
mr.pp_group = SimpleNamespace(rank_in_group=0)
|
mr.pp_group = SimpleNamespace(rank_in_group=0)
|
||||||
mr.spec_aux_config = SimpleNamespace(
|
mr.spec_aux_config = SimpleNamespace(
|
||||||
eagle_draft_num_layers=None,
|
eagle_draft_num_layers=None,
|
||||||
@@ -1271,7 +1278,8 @@ class TestSWAPoolFloor(CustomTestCase):
|
|||||||
kv_cache_dtype_str="fp8_e4m3",
|
kv_cache_dtype_str="fp8_e4m3",
|
||||||
model_config=cfg,
|
model_config=cfg,
|
||||||
layer_info=SimpleNamespace(start_layer=0, end_layer=40),
|
layer_info=SimpleNamespace(start_layer=0, end_layer=40),
|
||||||
ps=SimpleNamespace(pp_size=1, attn_dp_size=1),
|
pp_size=1,
|
||||||
|
attn_dp_size=1,
|
||||||
sliding_window_size=128,
|
sliding_window_size=128,
|
||||||
page_size=256,
|
page_size=256,
|
||||||
spec_algorithm=spec,
|
spec_algorithm=spec,
|
||||||
|
|||||||
@@ -29,14 +29,6 @@ class TestRustServerExtension(CustomTestCase):
|
|||||||
self.server.start_mm_workers(sentinel.spec, 8)
|
self.server.start_mm_workers(sentinel.spec, 8)
|
||||||
|
|
||||||
scheduler = SimpleNamespace(
|
scheduler = SimpleNamespace(
|
||||||
ps=SimpleNamespace(
|
|
||||||
dp_size=2,
|
|
||||||
attn_dp_rank=1,
|
|
||||||
tp_size=2,
|
|
||||||
tp_rank=1,
|
|
||||||
attn_tp_size=1,
|
|
||||||
attn_cp_size=1,
|
|
||||||
),
|
|
||||||
model_config=SimpleNamespace(is_multimodal=True),
|
model_config=SimpleNamespace(is_multimodal=True),
|
||||||
)
|
)
|
||||||
with (
|
with (
|
||||||
@@ -50,7 +42,16 @@ class TestRustServerExtension(CustomTestCase):
|
|||||||
patch.object(
|
patch.object(
|
||||||
server_module,
|
server_module,
|
||||||
"get_parallel",
|
"get_parallel",
|
||||||
return_value=SimpleNamespace(nnodes=1, pp_size=1),
|
return_value=SimpleNamespace(
|
||||||
|
nnodes=1,
|
||||||
|
pp_size=1,
|
||||||
|
dp_size=2,
|
||||||
|
attn_dp_rank=1,
|
||||||
|
tp_size=2,
|
||||||
|
tp_rank=1,
|
||||||
|
attn_tp_size=1,
|
||||||
|
attn_cp_size=1,
|
||||||
|
),
|
||||||
),
|
),
|
||||||
patch.object(ModelServer, "_partition_cores", return_value=(None, None)),
|
patch.object(ModelServer, "_partition_cores", return_value=(None, None)),
|
||||||
patch.object(
|
patch.object(
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ import unittest
|
|||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
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.scheduler_components.metrics_reporter import (
|
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
||||||
PrefillStats,
|
PrefillStats,
|
||||||
SchedulerMetricsReporter,
|
SchedulerMetricsReporter,
|
||||||
@@ -18,16 +17,6 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
|||||||
from sglang.test.test_utils import CustomTestCase, enter_scope
|
from sglang.test.test_utils import CustomTestCase, enter_scope
|
||||||
|
|
||||||
|
|
||||||
def _make_ps(**overrides) -> ParallelState:
|
|
||||||
"""Build a ParallelState with reasonable defaults for tests; override fields via kwargs."""
|
|
||||||
defaults = dict(
|
|
||||||
dp_rank=None,
|
|
||||||
moe_dp_rank=None,
|
|
||||||
)
|
|
||||||
defaults.update(overrides)
|
|
||||||
return ParallelState.trivial(**defaults)
|
|
||||||
|
|
||||||
|
|
||||||
class _FakeReq:
|
class _FakeReq:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -75,11 +64,30 @@ class _DummyPublisherThread:
|
|||||||
|
|
||||||
|
|
||||||
def _publish_server_args(test, **fields):
|
def _publish_server_args(test, **fields):
|
||||||
"""Publish a config for the reporter under test and return the instance."""
|
"""Publish a config for the reporter under test and return the instance.
|
||||||
|
|
||||||
|
The collector asks the context where this process sits, so the ranks are
|
||||||
|
stated too: without them a rank read falls through to a process group that
|
||||||
|
a unit test has not built.
|
||||||
|
"""
|
||||||
fields.setdefault("decode_log_interval", 40)
|
fields.setdefault("decode_log_interval", 40)
|
||||||
override = get_context().override_server_args(**fields)
|
override = get_context().override_server_args(**fields)
|
||||||
server_args = override.install()
|
server_args = override.install()
|
||||||
test.addCleanup(override.restore)
|
test.addCleanup(override.restore)
|
||||||
|
enter_scope(
|
||||||
|
test,
|
||||||
|
get_parallel().override(
|
||||||
|
tp_rank=0,
|
||||||
|
attn_tp_rank=0,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
moe_ep_rank=0,
|
||||||
|
attn_dp_rank=0,
|
||||||
|
dp_rank=0,
|
||||||
|
moe_ep_size=1,
|
||||||
|
moe_dp_size=1,
|
||||||
|
moe_tp_size=1,
|
||||||
|
),
|
||||||
|
)
|
||||||
return server_args
|
return server_args
|
||||||
|
|
||||||
|
|
||||||
@@ -93,8 +101,6 @@ def _make_reporter(test, scheduler) -> SchedulerMetricsReporter:
|
|||||||
enable_mfu_metrics=False,
|
enable_mfu_metrics=False,
|
||||||
enable_forward_pass_metrics=False,
|
enable_forward_pass_metrics=False,
|
||||||
)
|
)
|
||||||
if not hasattr(scheduler, "ps"):
|
|
||||||
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,
|
||||||
@@ -291,9 +297,9 @@ class TestForwardPassMetrics(unittest.TestCase):
|
|||||||
forward_pass_metrics_ipc_name=None,
|
forward_pass_metrics_ipc_name=None,
|
||||||
kv_events_config=None,
|
kv_events_config=None,
|
||||||
)
|
)
|
||||||
scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=2, pp_rank=0, pp_size=1)
|
# The reporter asks the context whether this is the last stage, and
|
||||||
# The reporter asks the context whether this is the last stage.
|
# which replica it is reporting for.
|
||||||
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=1))
|
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=1, dp_rank=2))
|
||||||
scheduler.enable_kv_cache_events = False
|
scheduler.enable_kv_cache_events = False
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
@@ -330,7 +336,6 @@ class TestForwardPassMetrics(unittest.TestCase):
|
|||||||
forward_pass_metrics_ipc_name=None,
|
forward_pass_metrics_ipc_name=None,
|
||||||
kv_events_config=None,
|
kv_events_config=None,
|
||||||
)
|
)
|
||||||
scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=0, pp_rank=0, pp_size=2)
|
|
||||||
# The reporter asks the context whether this is the last stage.
|
# The reporter asks the context whether this is the last stage.
|
||||||
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=2))
|
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=2))
|
||||||
scheduler.enable_kv_cache_events = False
|
scheduler.enable_kv_cache_events = False
|
||||||
|
|||||||
@@ -175,7 +175,6 @@ class TestDraftPerRunnerConfig(CustomTestCase):
|
|||||||
scheduler.tp_worker = SimpleNamespace(
|
scheduler.tp_worker = SimpleNamespace(
|
||||||
model_runner=SimpleNamespace(model_config=SimpleNamespace(context_len=4096))
|
model_runner=SimpleNamespace(model_config=SimpleNamespace(context_len=4096))
|
||||||
)
|
)
|
||||||
scheduler.ps = SimpleNamespace(gpu_id=0)
|
|
||||||
scheduler.nccl_port = 0
|
scheduler.nccl_port = 0
|
||||||
scheduler.spec_algorithm = SimpleNamespace(
|
scheduler.spec_algorithm = SimpleNamespace(
|
||||||
is_none=lambda: False,
|
is_none=lambda: False,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states
|
from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states
|
||||||
from sglang.srt.models.dspark import DSparkDraftMixin
|
from sglang.srt.models.dspark import DSparkDraftMixin
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.speculative.dspark_components.dspark_kv_inject import (
|
from sglang.srt.speculative.dspark_components.dspark_kv_inject import (
|
||||||
TargetHiddenKvInjector,
|
TargetHiddenKvInjector,
|
||||||
)
|
)
|
||||||
@@ -48,16 +49,13 @@ class DSparkTargetHiddenProjectionTest(CustomTestCase):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
with (
|
with (
|
||||||
mock.patch(
|
get_parallel().override(pp_group=SimpleNamespace(is_last_rank=False)),
|
||||||
"sglang.srt.speculative.dspark_components.dspark_worker_v2.get_pp_group",
|
|
||||||
return_value=SimpleNamespace(is_last_rank=False),
|
|
||||||
),
|
|
||||||
mock.patch(
|
mock.patch(
|
||||||
"sglang.srt.speculative.dspark_components.dspark_worker_v2.get_schedule",
|
"sglang.srt.speculative.dspark_components.dspark_worker_v2.get_schedule",
|
||||||
return_value=SimpleNamespace(page_size=1),
|
return_value=SimpleNamespace(page_size=1),
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
worker = DSparkWorkerV2(None, 0, None, 0, target)
|
worker = DSparkWorkerV2(None, 0, 0, target)
|
||||||
worker.alloc_memory_pool()
|
worker.alloc_memory_pool()
|
||||||
worker.init_attention_backends()
|
worker.init_attention_backends()
|
||||||
worker.init_cuda_graphs()
|
worker.init_cuda_graphs()
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ from sglang.srt.runtime_context import (
|
|||||||
assert_published,
|
assert_published,
|
||||||
derive_parallel_widths,
|
derive_parallel_widths,
|
||||||
get_context,
|
get_context,
|
||||||
|
get_device,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_flags,
|
get_flags,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -366,6 +367,35 @@ class TestSpawnIdentities(_IsolatedOverrides):
|
|||||||
self.assertEqual(parallel.pp_rank, 1)
|
self.assertEqual(parallel.pp_rank, 1)
|
||||||
self.assertEqual(parallel.dp_rank, 2)
|
self.assertEqual(parallel.dp_rank, 2)
|
||||||
|
|
||||||
|
def test_the_spawn_states_the_device_and_the_record_stays_clean(self):
|
||||||
|
"""The parent picks the device, so it arrives with the rest of the
|
||||||
|
placement. It is stamped onto the bag: the record is the startup
|
||||||
|
input and stays as the caller handed it over."""
|
||||||
|
server_args = ServerArgs(model_path="dummy")
|
||||||
|
publish(
|
||||||
|
server_args,
|
||||||
|
role="test",
|
||||||
|
ranks=SpawnRanks(world_rank=0, gpu_id=3),
|
||||||
|
)
|
||||||
|
self.assertEqual(get_device().gpu_id, 3)
|
||||||
|
# Not on the record at all. An `Arg` is the operator's input and is
|
||||||
|
# collected into `ServerArgs`; nobody types this one, so it is
|
||||||
|
# declared rather than carried, and the startup input has no field
|
||||||
|
# for the spawn to have to leave alone.
|
||||||
|
self.assertNotIn(
|
||||||
|
"gpu_id", {f.name for f in msgspec.structs.fields(type(server_args))}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_a_process_on_no_device_is_told_nothing(self):
|
||||||
|
"""Most roles run on no device at all, so the bundle leaves it out and
|
||||||
|
the bag keeps the declared default rather than inventing a zero."""
|
||||||
|
publish(
|
||||||
|
ServerArgs(model_path="dummy"),
|
||||||
|
role="test",
|
||||||
|
ranks=SpawnRanks(world_rank=0),
|
||||||
|
)
|
||||||
|
self.assertIsNone(get_device().gpu_id)
|
||||||
|
|
||||||
def test_no_controller_is_an_answer_not_a_failure(self):
|
def test_no_controller_is_an_answer_not_a_failure(self):
|
||||||
"""`dp_rank=None` means "not under a data parallel controller", which
|
"""`dp_rank=None` means "not under a data parallel controller", which
|
||||||
is a fact about the deployment, unlike never having been told. The
|
is a fact about the deployment, unlike never having been told. The
|
||||||
@@ -394,7 +424,7 @@ class TestSpawnIdentities(_IsolatedOverrides):
|
|||||||
class TestAttentionRanksComeFromPublish(_IsolatedOverrides):
|
class TestAttentionRanksComeFromPublish(_IsolatedOverrides):
|
||||||
"""With a spawn bundle, a rank read works before any group exists.
|
"""With a spawn bundle, a rank read works before any group exists.
|
||||||
|
|
||||||
This is what `ParallelState` provided by being a plain frozen record, and
|
This is what the per-runner record provided by being a plain frozen object, and
|
||||||
what the topology init could not: it needs the groups. Deriving at publish
|
what the topology init could not: it needs the groups. Deriving at publish
|
||||||
is what lets a reader ask the context in a process that never initialises
|
is what lets a reader ask the context in a process that never initialises
|
||||||
distributed -- every unit test that builds a scheduler component, for one.
|
distributed -- every unit test that builds a scheduler component, for one.
|
||||||
@@ -3068,6 +3098,58 @@ class TestWhoAnswersDuringADraftScope(CustomTestCase):
|
|||||||
self.assertEqual((info.pp_rank, info.pp_size), (0, 1))
|
self.assertEqual((info.pp_rank, info.pp_size), (0, 1))
|
||||||
|
|
||||||
|
|
||||||
|
class TestTheRecordIsNeverWrittenTo(CustomTestCase):
|
||||||
|
"""`server_args` is the startup record; the bags are the truth afterwards.
|
||||||
|
|
||||||
|
Writing a field onto it after `resolve_once()` has sealed it puts a second
|
||||||
|
answer where there is supposed to be one, and it is invisible to anything
|
||||||
|
reading the bag. The sanctioned writer is `RuntimeContext.override`, which
|
||||||
|
writes the bag and says so in its own contract. `arg_groups/` is exempt: it
|
||||||
|
is the resolution pipeline, so building the record is its job.
|
||||||
|
"""
|
||||||
|
|
||||||
|
#: Assignments here are the record being built, not mutated behind a reader.
|
||||||
|
EXEMPT = ("srt/arg_groups/",)
|
||||||
|
|
||||||
|
def test_nothing_assigns_a_field_of_the_record(self):
|
||||||
|
import ast as _ast
|
||||||
|
|
||||||
|
from sglang.srt.arg_groups.arg_utils import namespace_of
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
fields = set(namespace_of(ServerArgs))
|
||||||
|
offenders = []
|
||||||
|
for path in _sources():
|
||||||
|
rel = path.as_posix()
|
||||||
|
if "sglang/srt/" not in rel and "sglang/benchmark/" not in rel:
|
||||||
|
continue
|
||||||
|
if any(part in rel for part in self.EXEMPT):
|
||||||
|
continue
|
||||||
|
for node in _ast.walk(_ast.parse(path.read_text(encoding="utf-8-sig"))):
|
||||||
|
targets = (
|
||||||
|
node.targets
|
||||||
|
if isinstance(node, _ast.Assign)
|
||||||
|
else [node.target]
|
||||||
|
if isinstance(node, (_ast.AugAssign, _ast.AnnAssign))
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
for target in targets:
|
||||||
|
if not isinstance(target, _ast.Attribute):
|
||||||
|
continue
|
||||||
|
base = target.value
|
||||||
|
name = getattr(base, "id", getattr(base, "attr", None))
|
||||||
|
if target.attr.startswith("_"):
|
||||||
|
continue
|
||||||
|
if name == "server_args" and target.attr in fields:
|
||||||
|
offenders.append(f"{rel}:{target.lineno} .{target.attr}")
|
||||||
|
self.assertEqual(
|
||||||
|
offenders,
|
||||||
|
[],
|
||||||
|
"write the bag through get_context().override(source, ...) instead "
|
||||||
|
"-- the record is not a channel:\n " + "\n ".join(offenders),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestNothingReadsThePlacementBeforeItIsFrozen(CustomTestCase):
|
class TestNothingReadsThePlacementBeforeItIsFrozen(CustomTestCase):
|
||||||
"""`ModelRunner.__init__` freezes its placement partway through.
|
"""`ModelRunner.__init__` freezes its placement partway through.
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ from unittest.mock import patch
|
|||||||
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.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
quant_weight_ue8m0,
|
quant_weight_ue8m0,
|
||||||
transform_scale_ue8m0,
|
transform_scale_ue8m0,
|
||||||
@@ -181,15 +180,6 @@ class _FakeModelRunner:
|
|||||||
attn_dp_size: int | None = None,
|
attn_dp_size: int | None = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.ps = ParallelState.trivial(
|
|
||||||
tp_rank=tp_rank,
|
|
||||||
tp_size=tp_size,
|
|
||||||
dp_rank=dp_rank,
|
|
||||||
dp_size=dp_size,
|
|
||||||
attn_dp_size=attn_dp_size if attn_dp_size is not None else dp_size,
|
|
||||||
pp_rank=pp_rank,
|
|
||||||
pp_size=pp_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user