State the draft's whole topology in its scope, and read the rest from the context (#40339)

This commit is contained in:
Cheng Wan
2026-09-21 12:19:38 -07:00
committed by GitHub
parent ae7a516ba7
commit 0db1a93adb
22 changed files with 364 additions and 176 deletions
@@ -27,17 +27,15 @@ from abc import ABC, abstractmethod
from collections import deque
from itertools import count
from queue import Queue
from typing import TYPE_CHECKING, Any, Callable, Optional, Union
from typing import Any, Callable, Optional, Union
import msgspec
import zmq
from pydantic import BaseModel
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.network import NetworkAddress
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
logger = logging.getLogger(__name__)
@@ -63,17 +61,18 @@ def select_kv_publisher_dp_rank(
return dp_rank or 0
def is_kv_publisher_rank(kv_events_config: Optional[str], ps: "ParallelState") -> bool:
def is_kv_publisher_rank(kv_events_config: Optional[str]) -> bool:
"""Whether this scheduler owns a KV-event publisher slot: one per
independent KV cache (pp/attn-TP/attn-CP rank 0). Shared by
`SchedulerKvEventsPublisher` and `SchedulerLoadPublisher`, which must
gate identically or their /server_info-derived ports disagree.
"""
parallel = get_parallel()
return bool(
kv_events_config
and ps.pp_rank == 0
and ps.attn_tp_rank == 0
and ps.attn_cp_rank == 0
and parallel.pp_rank == 0
and parallel.attn_tp_rank == 0
and parallel.attn_cp_rank == 0
)
@@ -3045,17 +3045,37 @@ def patch_pipeline_parallel_group(pp_group: GroupCoordinator):
@contextmanager
def patch_tensor_parallel_group(tp_group: GroupCoordinator):
def patch_tensor_parallel_group(tp_group: GroupCoordinator, *, owns_attention: bool):
"""Run under a different tensor-parallel group until this scope ends.
This is for draft workers of speculative decoding, which run the draft model
at the target's attention-TP width rather than its global TP width.
The scope replaces both the module global that ``get_tp_group()`` reads and
the three members the runtime context answers with.
the members the runtime context answers with.
Which members depends on what the draft is, and only the worker knows: the
same call site hands over an attention-TP slice for one draft and the
target's whole TP group for another, so this cannot be read off the group.
``owns_attention`` says which. A draft that owns its attention topology
runs the whole model on the group being installed -- there is no
attention-DP replica inside it, so its attention identity is the group
itself, one replica, one context shard, and no expert dimension either.
Leaving those names on the target's answers is what lets a draft read
report a replica count the draft does not have.
A draft that does not own it was built outside any scope and keeps the
target's layout: the process is still one of several attention-DP replicas
and still gathers with them. Claiming one replica there is the same error
in the other direction, and the reader that acts on it is a collective -- a
DP gather takes its buffer size from the replica count and its communicator
from this group, so the two stop agreeing.
Args:
tp_group (GroupCoordinator): the tp group coordinator
owns_attention (bool): whether the draft's attention topology is this
group, decided by the worker where it builds its draft runner
"""
global _TP_STATE_PATCHED
@@ -3065,12 +3085,22 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator):
old_tp_group = get_tp_group()
global _TP
_TP = tp_group
narrowed = dict(
tp_size=tp_group.world_size,
tp_rank=tp_group.rank_in_group,
tp_group=tp_group,
)
if owns_attention:
narrowed.update(
attn_tp_size=tp_group.world_size,
attn_tp_rank=tp_group.rank_in_group,
attn_dp_size=1,
attn_dp_rank=0,
attn_cp_size=1,
attn_cp_rank=0,
)
try:
with get_parallel().override(
tp_size=tp_group.world_size,
tp_rank=tp_group.rank_in_group,
tp_group=tp_group,
):
with get_parallel().override(**narrowed):
yield
finally:
_TP_STATE_PATCHED = False
+3 -7
View File
@@ -801,7 +801,6 @@ class Scheduler(
) -> None:
self.metrics_collector_context = SchedulerMetricsCollector.init_new(
server_args=self.server_args,
ps=self.ps,
tp_rank=tp_rank,
pp_rank=pp_rank,
dp_rank=dp_rank,
@@ -843,7 +842,7 @@ class Scheduler(
try:
self.load_snapshot_writer = create_load_snapshot_writer(
port_args,
self.ps.dp_size,
get_parallel().dp_size,
dp_rank,
publish_interval=get_observability().load_snapshot_publish_interval,
)
@@ -1432,7 +1431,7 @@ class Scheduler(
)
else:
self.prefill_delayer = PrefillDelayer(
dp_size=self.ps.dp_size,
dp_size=get_parallel().dp_size,
attn_tp_size=get_parallel().attn_tp_size,
cpu_group=self.tp_cpu_group,
device_group=self.tp_group.device_group,
@@ -2340,7 +2339,6 @@ class Scheduler(
recv_skipper=self.recv_skipper,
input_blocker=self.input_blocker,
mm_receiver=self.mm_receiver,
ps=self.ps,
tp_group=self.tp_group,
tp_cpu_group=self.tp_cpu_group,
attn_tp_group=self.attn_tp_group,
@@ -2429,10 +2427,9 @@ class Scheduler(
def init_kv_events_publisher(self) -> None:
self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=get_observability().kv_events_config,
ps=self.ps,
attn_tp_rank=get_parallel().attn_tp_rank,
attn_cp_rank=get_parallel().attn_cp_rank,
attn_dp_rank=self.ps.attn_dp_rank,
attn_dp_rank=get_parallel().attn_dp_rank,
dp_rank=get_parallel().dp_rank,
tree_cache=self.tree_cache,
send_metrics_from_scheduler=self.ipc_channels.send_metrics_from_scheduler,
@@ -2448,7 +2445,6 @@ class Scheduler(
# instead of walking the queues itself.
self.load_publisher = SchedulerLoadPublisher(
kv_events_config=get_observability().kv_events_config,
ps=self.ps,
load_publish_endpoint=get_observability().load_publish_endpoint,
publish_interval=get_observability().load_snapshot_publish_interval,
)
@@ -22,7 +22,6 @@ from sglang.srt.managers.io_struct import hook_custom_types, sock_send
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
@@ -46,7 +45,6 @@ hook_custom_types(KvMetrics)
@dataclass(kw_only=True, slots=True)
class SchedulerKvEventsPublisher:
kv_events_config: Optional[str]
ps: ParallelState
attn_tp_rank: int
attn_cp_rank: int
attn_dp_rank: int
@@ -63,13 +61,14 @@ class SchedulerKvEventsPublisher:
self.init_kv_events(self.kv_events_config)
def init_kv_events(self, kv_events_config: Optional[str]):
self.enable_kv_cache_events = is_kv_publisher_rank(kv_events_config, self.ps)
self.enable_kv_cache_events = is_kv_publisher_rank(kv_events_config)
if self.enable_kv_cache_events:
parallel = get_parallel()
self.kv_event_publisher = EventPublisherFactory.create(
kv_events_config,
select_kv_publisher_dp_rank(
self.ps.attn_dp_size, self.ps.attn_dp_rank, get_parallel().dp_rank
parallel.attn_dp_size, parallel.attn_dp_rank, parallel.dp_rank
),
)
@@ -42,10 +42,10 @@ from sglang.srt.disaggregation.kv_events import (
resolve_load_pub_range,
select_kv_publisher_dp_rank,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.network import NetworkAddress, is_zmq_endpoint_ipv6
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.load_snapshot import LoadSnapshot
logger = logging.getLogger(__name__)
@@ -132,7 +132,6 @@ class SchedulerLoadPublisher:
self,
*,
kv_events_config: Optional[str],
ps: ParallelState,
load_publish_endpoint: Optional[str] = None,
publish_interval: int = LOAD_PUBLISH_INTERVAL,
) -> None:
@@ -147,7 +146,7 @@ class SchedulerLoadPublisher:
self._last_counts: Optional[tuple] = None
self._last_publish_ts = 0.0
self._publish_failed = False
if not is_kv_publisher_rank(kv_events_config, ps):
if not is_kv_publisher_rank(kv_events_config):
return
try:
cfg = KVEventsConfig.from_cli(kv_events_config)
@@ -165,7 +164,7 @@ class SchedulerLoadPublisher:
resolved, reason = resolve_load_pub_range(
kv_endpoint=cfg.endpoint,
replay_endpoint=cfg.replay_endpoint,
dp_size=ps.dp_size,
dp_size=get_parallel().dp_size,
load_publish_endpoint=load_publish_endpoint,
)
if resolved is None:
@@ -173,8 +172,9 @@ class SchedulerLoadPublisher:
logger.warning("load-publisher disabled: %s", reason)
return
host, base = resolved
parallel = get_parallel()
self._rank = select_kv_publisher_dp_rank(
ps.attn_dp_size, ps.attn_dp_rank, ps.dp_rank
parallel.attn_dp_size, parallel.attn_dp_rank, parallel.dp_rank
)
endpoint = NetworkAddress(host, base + self._rank).to_tcp()
try:
@@ -287,7 +287,7 @@ class SchedulerProfilerManager:
if get_parallel().tp_rank != 0:
return ""
if self.ps.dp_size > 1 and get_parallel().dp_rank != 0:
if get_parallel().dp_size > 1 and get_parallel().dp_rank != 0:
return ""
if get_parallel().pp_size > 1 and get_parallel().pp_rank != 0:
return ""
@@ -342,7 +342,7 @@ class SchedulerProfilerManager:
filename_parts = [self.profile_id, f"TP-{get_parallel().tp_rank}"]
# Only add other ranks if parallelism is enabled (size > 1)
if self.ps.dp_size > 1:
if get_parallel().dp_size > 1:
filename_parts.append(f"DP-{get_parallel().dp_rank}")
if get_parallel().pp_size > 1:
filename_parts.append(f"PP-{get_parallel().pp_rank}")
@@ -46,7 +46,6 @@ from sglang.srt.utils import (
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.rust_server.server import RustServer
from sglang.srt.server_args import ServerArgs
from sglang.test.scripted_runtime.scheduler_hook import ScriptedSchedulerHook
@@ -64,7 +63,6 @@ class SchedulerRequestReceiver:
recv_skipper: Any
input_blocker: Any
mm_receiver: Any
ps: ParallelState
tp_group: Any
tp_cpu_group: Any
attn_tp_group: Any
@@ -120,7 +118,7 @@ class SchedulerRequestReceiver:
def _pull_raw_reqs(self) -> Optional[List]:
if get_parallel().pp_rank == 0:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
if get_parallel().attn_tp_rank == 0 and get_parallel().attn_cp_rank == 0:
recv_reqs = []
# Rust ringbuffer backend: drain the in-process ring fed by the
@@ -152,16 +150,18 @@ class SchedulerRequestReceiver:
else:
recv_reqs = None
else:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
if get_parallel().attn_tp_rank == 0 and get_parallel().attn_cp_rank == 0:
dp_offset = (
self.ps.attn_dp_rank * self.ps.attn_cp_size * self.ps.attn_tp_size
get_parallel().attn_dp_rank
* get_parallel().attn_cp_size
* get_parallel().attn_tp_size
)
recv_reqs = point_to_point_pyobj(
[],
get_parallel().pp_rank * self.ps.tp_size + dp_offset,
get_parallel().pp_rank * get_parallel().tp_size + dp_offset,
self.world_group.cpu_group,
(get_parallel().pp_rank - 1) * self.ps.tp_size + dp_offset,
get_parallel().pp_rank * self.ps.tp_size + dp_offset,
(get_parallel().pp_rank - 1) * get_parallel().tp_size + dp_offset,
get_parallel().pp_rank * get_parallel().tp_size + dp_offset,
)
else:
recv_reqs = None
@@ -176,7 +176,7 @@ class SchedulerRequestReceiver:
"""
local_reqs = local_reqs or []
if get_parallel().enable_dp_attention:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
if get_parallel().attn_tp_rank == 0 and get_parallel().attn_cp_rank == 0:
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
work_reqs.extend(local_reqs)
else:
@@ -196,7 +196,7 @@ class SchedulerRequestReceiver:
)
if _local_ctrl:
control_reqs = attn_cp_tp_broadcast_pyobj(control_reqs)
elif self.ps.tp_size != 1:
elif get_parallel().tp_size != 1:
control_reqs = broadcast_pyobj(
control_reqs,
self.tp_group.rank,
@@ -207,7 +207,7 @@ class SchedulerRequestReceiver:
else:
if recv_reqs is not None:
recv_reqs = [*recv_reqs, *local_reqs]
if self.ps.tp_size != 1:
if get_parallel().tp_size != 1:
recv_reqs = broadcast_pyobj(
recv_reqs,
self.tp_group.rank,
@@ -269,11 +269,11 @@ class SchedulerRequestReceiver:
# 1. wait until every rank has opened the shared feature segments
parallel = get_parallel()
if parallel.enable_dp_attention:
if self.ps.attn_tp_size > 1:
if parallel.attn_tp_size > 1:
barrier(group=self.attn_tp_cpu_group)
if self.ps.attn_cp_size > 1:
if parallel.attn_cp_size > 1:
barrier(group=self.attn_cp_cpu_group)
elif self.ps.tp_size > 1:
elif parallel.tp_size > 1:
barrier(group=self.tp_cpu_group)
# 2. materialize independently so one bad VLM request does not stop the loop
@@ -293,11 +293,11 @@ class SchedulerRequestReceiver:
# 3. all ranks reject the same requests before entering model collectives
if parallel.enable_dp_attention:
if self.ps.attn_tp_size > 1:
if parallel.attn_tp_size > 1:
all_reduce(failed, op=ReduceOp.MAX, group=self.attn_tp_cpu_group)
if self.ps.attn_cp_size > 1:
if parallel.attn_cp_size > 1:
all_reduce(failed, op=ReduceOp.MAX, group=self.attn_cp_cpu_group)
elif self.ps.tp_size > 1:
elif parallel.tp_size > 1:
all_reduce(failed, op=ReduceOp.MAX, group=self.tp_cpu_group)
error = MMInputsProcessError(
@@ -815,7 +815,7 @@ class SchedulerPPMixin:
p2p_work = []
if get_parallel().attn_tp_rank == 0 and get_parallel().attn_cp_rank == 0:
dp_offset = (
self.ps.attn_dp_rank
get_parallel().attn_dp_rank
* get_parallel().attn_cp_size
* get_parallel().attn_tp_size
)
@@ -834,7 +834,7 @@ class SchedulerPPMixin:
def _pp_recv_pyobj_from_prev_stage(self: Scheduler):
if get_parallel().attn_tp_rank == 0 and get_parallel().attn_cp_rank == 0:
dp_offset = (
self.ps.attn_dp_rank
get_parallel().attn_dp_rank
* get_parallel().attn_cp_size
* get_parallel().attn_tp_size
)
@@ -1103,7 +1103,6 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
cls,
*,
server_args: ServerArgs,
ps: Any,
tp_rank: int,
pp_rank: int,
dp_rank: Optional[int],
@@ -1112,7 +1111,8 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
enable_hierarchical_cache: bool,
) -> SchedulerMetricsCollectorContext:
enable_metrics = get_observability().enable_metrics
is_stats_logging_rank = ps.attn_tp_rank == 0
parallel = get_parallel()
is_stats_logging_rank = parallel.attn_tp_rank == 0
current_scheduler_metrics_enabled = enable_metrics and (
is_stats_logging_rank
or get_observability().enable_metrics_for_all_schedulers
@@ -1120,8 +1120,8 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
enable_kv_cache_events = bool(
get_observability().kv_events_config
and get_parallel().pp_rank == 0
and ps.attn_tp_rank == 0
and ps.attn_cp_rank == 0
and parallel.attn_tp_rank == 0
and parallel.attn_cp_rank == 0
)
collector: Optional[SchedulerMetricsCollector] = None
if enable_metrics:
@@ -1136,7 +1136,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
"engine_type": engine_type,
"tp_rank": tp_rank,
"pp_rank": pp_rank,
"moe_ep_rank": ps.moe_ep_rank,
"moe_ep_rank": parallel.moe_ep_rank,
}
if enable_priority_scheduling:
labels["priority"] = ""
@@ -400,8 +400,16 @@ class DFlashWorkerV2(BaseSpecWorker):
self.draft_tp_context = (
draft_tp_context if get_parallel().enable_dp_attention else empty_context
)
if get_parallel().enable_dp_attention:
draft_init_ctx = draft_tp_context(get_parallel().attn_tp_group)
# One decision, used twice: whether the draft runs on an attention-TP
# slice of its own. It picks how the runner is built, and then what the
# scope may say about attention every time it is entered -- a draft
# built outside the scope keeps the target's replica count and still
# gathers with it.
self.draft_owns_attention = get_parallel().enable_dp_attention
if self.draft_owns_attention:
draft_init_ctx = draft_tp_context(
get_parallel().attn_tp_group, owns_attention=True
)
else:
draft_init_ctx = empty_context()
with draft_pp_context(), draft_init_ctx:
@@ -602,7 +610,10 @@ class DFlashWorkerV2(BaseSpecWorker):
def init_attention_backends(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
):
self._draft_worker.init_attention_backends()
self._need_mamba_verify_commit = mambaish_config(
@@ -615,7 +626,10 @@ class DFlashWorkerV2(BaseSpecWorker):
def init_cuda_graphs(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
):
capture_decode_cuda_graph = (
get_exec().graph.cuda_graph_config.decode.backend != Backend.DISABLED
@@ -1818,7 +1832,10 @@ class DFlashWorkerV2(BaseSpecWorker):
with (
torch.inference_mode(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
):
ctx_hidden = self.draft_model.project_target_hidden(target_hidden)
@@ -2513,7 +2530,10 @@ class DFlashWorkerV2(BaseSpecWorker):
with (
torch.inference_mode(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
):
draft_out = self.draft_model_runner.forward(forward_batch)
draft_logits_output = draft_out.logits_output
@@ -2572,7 +2592,10 @@ class DFlashWorkerV2(BaseSpecWorker):
self._draft_sampler.q_out[:bs],
)
elif self.selector is not None:
with self.draft_tp_context(self.draft_model_runner.tp_group):
with self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
):
draft_next = self._propose_selector_block(
draft_logits_output=draft_logits_output,
bs=bs,
@@ -2585,7 +2608,10 @@ class DFlashWorkerV2(BaseSpecWorker):
if draft_hidden is None:
raise RuntimeError("DFLASH draft model returned no hidden states.")
draft_hidden = draft_hidden.view(bs, int(self.block_size), -1)
with self.draft_tp_context(self.draft_model_runner.tp_group):
with self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
):
draft_next = self._greedy_sample_from_vocab_parallel_head(
hidden_states=draft_hidden[:, 1:, :].reshape(
-1, draft_hidden.shape[-1]
@@ -236,7 +236,7 @@ class DraftBlockProposer:
def _base_logits_context(self):
if self._dp_moe_sync:
return draft_tp_context(get_parallel().attn_tp_group)
return draft_tp_context(get_parallel().attn_tp_group, owns_attention=True)
return nullcontext()
def propose(
@@ -424,7 +424,7 @@ class DSparkWorkerV2(BaseSpecWorker):
def _draft_context(self):
if self._draft_dp_context_enabled:
return draft_tp_context(get_parallel().attn_tp_group)
return draft_tp_context(get_parallel().attn_tp_group, owns_attention=True)
return nullcontext()
def alloc_memory_pool(
@@ -265,11 +265,17 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self._rebuild_topk1_chain_buffers()
# Load draft model weights only.
if (
# One decision, used twice: whether the draft runs on an attention-TP
# slice of its own. It picks how the runner is built, and then what the
# scope may say about attention every time it is entered -- a draft
# built outside the scope keeps the target's replica count and still
# gathers with it.
self.draft_owns_attention = (
get_parallel().enable_dp_attention
and self.speculative_algorithm.is_eagle3()
):
ctx = draft_tp_context(get_parallel().attn_tp_group)
)
if self.draft_owns_attention:
ctx = draft_tp_context(get_parallel().attn_tp_group, owns_attention=True)
else:
ctx = empty_context()
with (
@@ -339,7 +345,10 @@ class EagleDraftWorker(EagleDraftWorkerBase):
def init_attention_backends(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_runner.tp_group),
self.draft_tp_context(
self.draft_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
):
@@ -349,7 +358,10 @@ class EagleDraftWorker(EagleDraftWorkerBase):
def init_cuda_graphs(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_runner.tp_group),
self.draft_tp_context(
self.draft_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
):
@@ -1358,7 +1370,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
if self.adaptive_controller is not None:
with (
self._draft_worker.draft_tp_context(
self._draft_worker.draft_runner.tp_group
self._draft_worker.draft_runner.tp_group,
owns_attention=self._draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -1418,7 +1431,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Draft prefill
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -1465,7 +1479,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
else:
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -1490,7 +1505,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
else:
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -1522,7 +1538,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -166,6 +166,10 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self.kv_context: Optional[FrozenKVMTPContext] = None
# Built above under the pipeline scope only, so this runner carries the
# target's attention topology: entering the tensor scope later swaps the
# communicator without giving the draft a replica of its own.
self.draft_owns_attention = False
self.draft_tp_context = (
draft_tp_context if get_parallel().enable_dp_attention else empty_context
)
@@ -206,7 +210,10 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
def init_attention_backends(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
):
@@ -217,7 +224,10 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
def init_cuda_graphs(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_model_runner.tp_group),
self.draft_tp_context(
self.draft_model_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
):
@@ -768,7 +778,8 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
# Draft prefill seed (no forward).
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -790,7 +801,8 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
batch.spec_info = self.draft_worker._idle_seed()
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -805,7 +817,8 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
on_publish(batch_output.new_seq_lens)
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
self.draft_worker.draft_runner.tp_group,
owns_attention=self.draft_worker.draft_owns_attention,
),
speculative_moe_backend_context(),
speculative_moe_a2a_backend_context(),
@@ -193,6 +193,11 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
"InklingForConditionalGenerationMTP",
"GigaChat35ForCausalLMNextN",
]
# The draft runner is built outside any tensor-parallel scope, so it
# carries the target's topology: entering the scope later swaps the
# communicator without making this process a draft with an attention
# replica of its own. It still gathers with the target's replicas.
self.draft_owns_attention = False
self.draft_tp_context = (
draft_tp_context if get_parallel().enable_dp_attention else empty_context
)
@@ -224,7 +229,10 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
def init_attention_backends(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_runner_list[0].tp_group),
self.draft_tp_context(
self.draft_runner_list[0].tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
):
super().init_attention_backends()
@@ -232,7 +240,10 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
def init_cuda_graphs(self):
with (
draft_pp_context(),
self.draft_tp_context(self.draft_runner_list[0].tp_group),
self.draft_tp_context(
self.draft_runner_list[0].tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
):
super().init_cuda_graphs()
+2 -2
View File
@@ -711,10 +711,10 @@ def draft_pp_context():
@contextmanager
def draft_tp_context(tp_group: GroupCoordinator):
def draft_tp_context(tp_group: GroupCoordinator, *, owns_attention: bool):
# Draft model doesn't use dp and has its own tp group.
# We disable mscclpp now because it doesn't support 2 comm groups.
with patch_tensor_parallel_group(tp_group):
with patch_tensor_parallel_group(tp_group, owns_attention=owns_attention):
yield
@@ -94,6 +94,11 @@ class StandaloneDraftWorker(EagleDraftWorker):
# Alias for better readability
self.draft_runner = self.draft_worker.model_runner
# The draft runner is built outside any tensor-parallel scope, so it
# carries the target's topology: entering the scope later swaps the
# communicator without making this process a draft with an attention
# replica of its own. It still gathers with the target's replicas.
self.draft_owns_attention = False
self.draft_tp_context = (
draft_tp_context if get_parallel().enable_dp_attention else empty_context
)
@@ -132,14 +137,20 @@ class StandaloneDraftWorker(EagleDraftWorker):
def init_attention_backends(self):
with (
self.draft_tp_context(self.draft_runner.tp_group),
self.draft_tp_context(
self.draft_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
):
super().init_attention_backends()
def init_cuda_graphs(self):
with (
self.draft_tp_context(self.draft_runner.tp_group),
self.draft_tp_context(
self.draft_runner.tp_group,
owns_attention=self.draft_owns_attention,
),
speculative_moe_backend_context(),
):
super().init_cuda_graphs()
+1 -1
View File
@@ -361,7 +361,7 @@ class _ProfilerTorch(_ProfilerConcreteBase):
filename_parts = [self.profile_id, f"TP-{get_parallel().tp_rank}"]
# Only add other ranks if parallelism is enabled (size > 1)
if self.ps.dp_size > 1:
if get_parallel().dp_size > 1:
filename_parts.append(f"DP-{get_parallel().dp_rank}")
if get_parallel().pp_size > 1:
filename_parts.append(f"PP-{get_parallel().pp_rank}")
@@ -23,13 +23,12 @@ from unittest.mock import MagicMock, patch
import msgspec.msgpack
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.scheduler_components.load_publisher import (
LoadStat,
SchedulerLoadPublisher,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
from sglang.test.test_utils import CustomTestCase, published_topology
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
@@ -97,19 +96,21 @@ class TestLoadPublisherGating(CustomTestCase):
connect-style one.
"""
def _build(
self, *, config=ZMQ_ENDPOINT, dp_size=1, explicit="auto", **ps_overrides
):
def _build(self, *, config=ZMQ_ENDPOINT, explicit="auto", ranks=None, **topology):
"""Construct a publisher with the socket bind stubbed out, returning
(publisher, captured _open_pub_socket mock). Opts in via explicit="auto"
by default (the feature is off without it). dp_size lives on the ps,
which the publisher reads (no separate param to disagree with it)."""
with patch(
"sglang.srt.managers.scheduler_components.load_publisher._open_pub_socket"
) as open_sock:
by default (the feature is off without it). The topology is published
rather than overridden, so the ranks the publisher reads are the ones a
layout of that shape actually produces; every read happens in the
constructor."""
with (
published_topology(ranks=ranks, **topology),
patch(
"sglang.srt.managers.scheduler_components.load_publisher._open_pub_socket"
) as open_sock,
):
pub = SchedulerLoadPublisher(
kv_events_config=config,
ps=ParallelState.trivial(dp_size=dp_size, **ps_overrides),
load_publish_endpoint=explicit,
)
return pub, open_sock
@@ -130,14 +131,17 @@ class TestLoadPublisherGating(CustomTestCase):
def test_disabled_off_pp_rank_zero(self):
# Every PP stage shares attn_tp_rank/attn_cp_rank 0, so without the
# pp_rank gate they all bind the same load port.
pub, open_sock = self._build(pp_rank=1, pp_size=2)
pub, open_sock = self._build(pp_size=2, ranks={"world_rank": 1})
self.assertFalse(pub.enable)
open_sock.assert_not_called()
def test_disabled_off_attn_tp_and_cp_rank_zero(self):
for override in ({"attn_tp_rank": 1}, {"attn_cp_rank": 1}):
with self.subTest(**override):
pub, open_sock = self._build(**override)
for layout in (
{"tp_size": 2},
{"tp_size": 2, "attn_cp_size": 2},
):
with self.subTest(**layout):
pub, open_sock = self._build(ranks={"world_rank": 1}, **layout)
self.assertFalse(pub.enable)
open_sock.assert_not_called()
@@ -145,11 +149,16 @@ class TestLoadPublisherGating(CustomTestCase):
# Pure DP: attn_dp_size == 1 and every worker has attn_dp_rank == 0, so
# the publisher must key off dp_rank or all replicas collide on one
# port. kv 5557 + dp_size 4 => base 5561; rank 2 binds 5563.
_, open_sock = self._build(attn_dp_size=1, attn_dp_rank=0, dp_rank=2, dp_size=4)
_, open_sock = self._build(dp_size=4, ranks={"world_rank": 0, "dp_rank": 2})
open_sock.assert_called_once_with("tcp://*:5563")
def test_dp_attention_keys_the_load_port_by_attn_dp_rank(self):
_, open_sock = self._build(attn_dp_size=4, attn_dp_rank=3, dp_rank=0, dp_size=4)
_, open_sock = self._build(
tp_size=4,
dp_size=4,
enable_dp_attention=True,
ranks={"world_rank": 3, "dp_rank": 0},
)
open_sock.assert_called_once_with("tcp://*:5564")
def test_load_port_is_packed_after_the_kv_range(self):
@@ -259,10 +268,8 @@ class TestLoadPublisherGating(CustomTestCase):
_, open_sock = self._build(
explicit="tcp://*:7000",
attn_dp_size=1,
attn_dp_rank=0,
dp_rank=2,
dp_size=4,
ranks={"world_rank": 0, "dp_rank": 2},
)
open_sock.assert_called_once_with("tcp://*:7002")
@@ -289,11 +296,11 @@ class TestLoadPublisherGating(CustomTestCase):
"sglang.srt.managers.scheduler_components.load_publisher._open_pub_socket",
side_effect=zmq.ZMQError,
) as open_sock:
pub = SchedulerLoadPublisher(
kv_events_config=ZMQ_ENDPOINT,
ps=ParallelState.trivial(),
load_publish_endpoint="auto",
)
with published_topology():
pub = SchedulerLoadPublisher(
kv_events_config=ZMQ_ENDPOINT,
load_publish_endpoint="auto",
)
open_sock.assert_called_once() # the bind was attempted and failed
self.assertFalse(pub.enable)
pub.publish_load_stat(MagicMock(), force=True) # still a no-op
@@ -468,11 +475,11 @@ class TestLoadStatIntegration(CustomTestCase):
with _socket.socket() as probe:
probe.bind(("", 0))
port = probe.getsockname()[1]
pub = SchedulerLoadPublisher(
kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}',
ps=ParallelState.trivial(),
load_publish_endpoint=f"tcp://*:{port}",
)
with published_topology():
pub = SchedulerLoadPublisher(
kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}',
load_publish_endpoint=f"tcp://*:{port}",
)
if pub.enable:
break
self.assertTrue(pub.enable, "load socket never bound a free port")
@@ -91,7 +91,7 @@ def _request(feature, rid: str = "vlm-request") -> TokenizedEmbeddingReqInput:
)
def _receiver(tp_size: int = 1) -> SchedulerRequestReceiver:
def _receiver() -> SchedulerRequestReceiver:
group = SimpleNamespace(rank=0, ranks=[0], cpu_group=object())
return SchedulerRequestReceiver(
recv_from_tokenizer=None,
@@ -99,14 +99,6 @@ def _receiver(tp_size: int = 1) -> SchedulerRequestReceiver:
recv_skipper=None,
input_blocker=None,
mm_receiver=None,
ps=SimpleNamespace(
pp_rank=0,
tp_size=tp_size,
attn_tp_rank=0,
attn_cp_rank=0,
attn_tp_size=1,
attn_cp_size=1,
),
tp_group=group,
tp_cpu_group=group,
attn_tp_group=group,
@@ -131,8 +123,8 @@ def _run_consensus_rank(rank: int, world_size: int, init_file: str) -> None:
)
try:
req = _request(_failed_pointer() if rank == 1 else _successful_pointer())
parallel = SimpleNamespace(enable_dp_attention=False)
receiver = _receiver(tp_size=world_size)
parallel = SimpleNamespace(enable_dp_attention=False, tp_size=world_size)
receiver = _receiver()
object.__setattr__(receiver, "tp_cpu_group", torch.distributed.group.WORLD)
with (
patch(
@@ -161,7 +153,7 @@ def _run_image_receiver(rank, init_file, pipe):
backend="gloo", init_method=Path(init_file).as_uri(), rank=rank, world_size=2
)
try:
receiver = _receiver(tp_size=2)
receiver = _receiver()
object.__setattr__(receiver, "tp_cpu_group", torch.distributed.group.WORLD)
torch.distributed.barrier()
torch.distributed.all_reduce(torch.zeros(1))
@@ -179,7 +171,7 @@ def _run_image_receiver(rank, init_file, pipe):
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.get_parallel",
return_value=SimpleNamespace(enable_dp_attention=False),
return_value=SimpleNamespace(enable_dp_attention=False, tp_size=2),
),
):
for base in [30, 90]:
@@ -382,7 +374,7 @@ class TestShmRequestFailureConsensus(unittest.TestCase):
def test_local_materialization_failure_becomes_request_error(self):
req = _request(_failed_pointer())
parallel = SimpleNamespace(enable_dp_attention=False)
parallel = SimpleNamespace(enable_dp_attention=False, tp_size=1)
with (
patch(
@@ -406,7 +398,7 @@ class TestShmRequestFailureConsensus(unittest.TestCase):
def test_peer_failure_rejects_the_local_request(self):
req = _request(torch.zeros(1))
parallel = SimpleNamespace(enable_dp_attention=False)
parallel = SimpleNamespace(enable_dp_attention=False, tp_size=2)
def inject_peer_failure(mask, **kwargs):
mask.fill_(1)
@@ -429,7 +421,7 @@ class TestShmRequestFailureConsensus(unittest.TestCase):
side_effect=inject_peer_failure,
) as all_reduce,
):
_receiver(tp_size=2)._finalize_shm_features([req])
_receiver()._finalize_shm_features([req])
all_reduce.assert_called_once()
self.assertIsInstance(req.mm_inputs, MMInputsProcessError)
@@ -438,7 +430,7 @@ class TestShmRequestFailureConsensus(unittest.TestCase):
failed_req = _request(torch.zeros(1), rid="failed")
healthy_req = _request(torch.zeros(1), rid="healthy")
batch = BatchTokenizedEmbeddingReqInput(batch=[failed_req, healthy_req])
parallel = SimpleNamespace(enable_dp_attention=False)
parallel = SimpleNamespace(enable_dp_attention=False, tp_size=1)
def materialize(req):
if req.rid == "failed":
@@ -14,7 +14,6 @@ from sglang.test.test_utils import (
maybe_stub_sgl_kernel()
from sglang.srt.distributed.parallel_state_wrapper import ParallelState # noqa: E402
from sglang.srt.managers.scheduler_components.request_receiver import ( # noqa: E402
SchedulerRequestReceiver,
)
@@ -23,29 +22,12 @@ from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin # noqa: E40
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
def _make_ps(**overrides) -> ParallelState:
defaults = dict(
tp_size=8,
pp_rank=1,
pp_size=2,
dp_rank=None,
attn_tp_size=2,
attn_cp_size=2,
attn_dp_rank=1,
attn_dp_size=2,
moe_dp_rank=None,
)
defaults.update(overrides)
return ParallelState.trivial(**defaults)
def _published_topology():
"""The topology `_make_ps` describes, published instead of stood in.
"""The topology these tests run in.
World rank 12 of a `tp=8, pp=2` world is `tp_rank=4` on the second stage,
which puts this process at `attn_dp_rank=1` with `attn_tp_rank=0`: the
context derives all of them from that one number and the widths, where the
record above had to be handed each.
context derives all of them from that one number and the widths.
"""
return published_topology(
role="scheduler",
@@ -62,7 +44,7 @@ def _fake_group() -> SimpleNamespace:
return SimpleNamespace(rank=0, ranks=[0], cpu_group=object())
def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
def _make_receiver() -> SchedulerRequestReceiver:
tp_group = _fake_group()
attn_tp_group = _fake_group()
attn_cp_group = _fake_group()
@@ -73,7 +55,6 @@ def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
recv_skipper=None,
input_blocker=None,
mm_receiver=None,
ps=ps,
tp_group=tp_group,
tp_cpu_group=tp_group,
attn_tp_group=attn_tp_group,
@@ -97,19 +78,17 @@ class TestRequestReceiverBroadcast(unittest.TestCase):
# Decode uses pure DP attention (attn_tp=attn_cp=1). The DP controller
# sends control requests to every local leader, so no per-tick Gloo
# broadcast should remain in SchedulerRequestReceiver.
ps = SimpleNamespace(
receiver = _make_receiver()
control_req = SimpleNamespace(kind="control")
parallel = SimpleNamespace(
enable_dp_attention=True,
enable_dp_attention_local_control_broadcast=True,
attn_tp_rank=0,
attn_cp_rank=0,
attn_tp_size=1,
attn_cp_size=1,
tp_size=32,
)
receiver = _make_receiver(ps)
control_req = SimpleNamespace(kind="control")
parallel = SimpleNamespace(
enable_dp_attention=True,
enable_dp_attention_local_control_broadcast=True,
)
with (
patch(
@@ -133,19 +112,17 @@ class TestRequestReceiverBroadcast(unittest.TestCase):
broadcast.assert_not_called()
def test_default_control_uses_full_tp_broadcast(self):
ps = SimpleNamespace(
receiver = _make_receiver()
control_req = SimpleNamespace(kind="control")
parallel = SimpleNamespace(
enable_dp_attention=True,
enable_dp_attention_local_control_broadcast=False,
attn_tp_rank=0,
attn_cp_rank=0,
attn_tp_size=1,
attn_cp_size=1,
tp_size=32,
)
receiver = _make_receiver(ps)
control_req = SimpleNamespace(kind="control")
parallel = SimpleNamespace(
enable_dp_attention=True,
enable_dp_attention_local_control_broadcast=False,
)
with (
patch(
@@ -183,7 +160,6 @@ class TestRequestReceiverBroadcast(unittest.TestCase):
class TestPPCPRankOffsets(unittest.TestCase):
def test_request_receiver_uses_cp_size_for_pp_recv_rank(self):
ps = _make_ps()
enter_scope(self, _published_topology())
calls = []
@@ -191,7 +167,7 @@ class TestPPCPRankOffsets(unittest.TestCase):
calls.append((rank, src, dst))
return ["req"]
receiver = _make_receiver(ps)
receiver = _make_receiver()
with patch(
"sglang.srt.managers.scheduler_components.request_receiver."
"point_to_point_pyobj",
@@ -202,10 +178,8 @@ class TestPPCPRankOffsets(unittest.TestCase):
self.assertEqual(calls, [(12, 4, 12)])
def test_pp_mixin_uses_cp_size_for_pyobj_send_and_recv_rank(self):
ps = _make_ps()
enter_scope(self, _published_topology())
scheduler = SchedulerPPMixin()
scheduler.ps = ps
scheduler.world_group = _fake_group()
scheduler.attn_tp_group = _fake_group()
scheduler.attn_tp_cpu_group = _fake_group()
@@ -2338,6 +2338,119 @@ class TestWhoAnswersDuringADraftScope(CustomTestCase):
self.assertEqual(get_parallel().pp_size, 2)
self.assertEqual(get_parallel().pp_rank, 1)
def _group(self, world_size, rank):
from sglang.srt.distributed.parallel_state import GroupCoordinator
group = GroupCoordinator.__new__(GroupCoordinator)
group.world_size = world_size
group.rank_in_group = rank
return group
def test_the_tensor_swap_states_the_draft_has_no_attention_replica(self):
"""The draft runs the whole model on the group being installed. Its
attention identity is therefore that group, with one replica -- while
the target this process also serves is attention-DP over four ranks."""
from sglang.srt.distributed import parallel_state
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy", tp_size=4, dp_size=2, enable_dp_attention=True
),
role="scheduler",
ranks=SpawnRanks(world_rank=0, dp_rank=0),
)
self.assertEqual(get_parallel().attn_dp_size, 2)
self.assertEqual(get_parallel().attn_tp_size, 2)
group = self._group(world_size=2, rank=1)
with patch.object(parallel_state, "_TP", group):
with parallel_state.patch_tensor_parallel_group(group, owns_attention=True):
parallel = get_parallel()
self.assertEqual(parallel.tp_size, 2)
self.assertEqual(parallel.attn_tp_size, 2)
self.assertEqual(parallel.attn_tp_rank, 1)
self.assertEqual(parallel.attn_dp_size, 1)
self.assertEqual(parallel.attn_dp_rank, 0)
self.assertEqual(parallel.attn_cp_size, 1)
self.assertEqual(parallel.attn_cp_rank, 0)
# `dp_size` is the deployment's replica count, not a property
# of the group being installed, so the scope leaves it alone --
# `require_mlp_tp_gather` asserts on it under dp attention.
self.assertEqual(parallel.dp_size, 2)
# The whole point of stating the rest: the identity the
# override path and the group build both check holds in here.
self.assertEqual(
parallel.tp_size,
parallel.attn_tp_size
* parallel.attn_dp_size
* parallel.attn_cp_size,
)
self.assertEqual(get_parallel().attn_dp_size, 2)
self.assertEqual(get_parallel().dp_size, 2)
def test_every_caller_says_whether_the_draft_owns_its_attention(self):
"""The scope cannot work it out from the group it is handed: the same
call site passes an attention-TP slice for one draft and the target's
whole TP group for another, and the two want opposite answers. So the
worker states it, and a caller that forgets is the bug this catches --
`owns_attention` has no default, but a missing one is a TypeError only
on the path that runs, and these paths need a GPU and a draft model."""
import ast
package = _pathlib.Path(next(iter(_sglang.__path__))).resolve()
checkout = package.parents[1]
roots = [package] + [
checkout / name for name in ("test",) if (checkout / name).is_dir()
]
missing = []
for path in (q for root in roots for q in root.rglob("*.py")):
try:
tree = ast.parse(path.read_text(encoding="utf-8"))
except (SyntaxError, UnicodeDecodeError):
continue
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
continue
func = node.func
name = getattr(func, "attr", None) or getattr(func, "id", None)
if name not in ("draft_tp_context", "patch_tensor_parallel_group"):
continue
if not any(kw.arg == "owns_attention" for kw in node.keywords):
missing.append(f"{path}:{node.lineno}")
self.assertEqual(missing, [], "these enter the draft scope without saying")
def test_a_full_width_swap_leaves_the_attention_layout_alone(self):
"""The other caller. A draft built outside any scope carries the
target's whole TP group, and the graph capture installs *that* -- so
the process is still one of two attention-DP replicas and still gathers
with the other one. Narrowing here would claim a replica count it does
not have, and the reader that acts on it is a collective: the DP gather
takes its buffer size from the replica count and its communicator from
this group, so the two stop agreeing and the all-gather is refused."""
from sglang.srt.distributed import parallel_state
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy", tp_size=4, dp_size=2, enable_dp_attention=True
),
role="scheduler",
ranks=SpawnRanks(world_rank=0, dp_rank=0),
)
whole_tp = self._group(world_size=4, rank=0)
with patch.object(parallel_state, "_TP", whole_tp):
with parallel_state.patch_tensor_parallel_group(
whole_tp, owns_attention=False
):
parallel = get_parallel()
self.assertEqual(parallel.tp_size, 4)
self.assertEqual(parallel.attn_dp_size, 2)
self.assertEqual(parallel.attn_tp_size, 2)
self.assertEqual(parallel.dp_size, 2)
def test_a_report_built_for_a_runner_follows_that_runner(self):
"""A weight check is an on-demand request served from the scheduler
loop, so it runs outside the scope that describes a draft runner. Its