Extract Scheduler init methods and add skills to enforce the splitting requirements (#26271)

This commit is contained in:
fzyzcjy
2026-05-26 17:45:09 +08:00
committed by GitHub
parent 48f3264807
commit d9c82934c8
4 changed files with 235 additions and 161 deletions
@@ -0,0 +1,6 @@
# Must-Read Skills Before Modifying Components
Before modifying the following components, read the listed skill first.
- **Speculative decoding code** (anything under `python/sglang/srt/speculative/`, related attention backends, scheduler accumulators, IPC fields, observability metrics, or CLI flags) → [`speculative-naming`](../skills/speculative-naming/SKILL.md)
- **`Scheduler` / `TokenizerManager` / `ModelRunner` `__init__`** (`python/sglang/srt/managers/scheduler.py`, `python/sglang/srt/managers/tokenizer_manager.py`, `python/sglang/srt/model_executor/model_runner.py`) → [`large-class-init-style`](../skills/large-class-init-style/SKILL.md)
-3
View File
@@ -1,3 +0,0 @@
# Speculative Decoding — Naming Conventions
When adding, renaming, or reviewing identifiers in speculative decoding code (anything under `python/sglang/srt/speculative/`, related attention backends, scheduler accumulators, IPC fields, observability metrics, or CLI flags), read the [`speculative-naming`](../skills/speculative-naming/SKILL.md) skill first.
@@ -0,0 +1,32 @@
---
name: large-class-init-style
description: '`__init__` style for SGLang `Scheduler`, `TokenizerManager`, and `ModelRunner`. Use when modifying the `__init__` of any of these three classes, or reviewing changes that add new construction logic to them.'
---
# `__init__` Style for Scheduler / TokenizerManager / ModelRunner
Apply when modifying the `__init__` of:
- `Scheduler``python/sglang/srt/managers/scheduler.py`
- `TokenizerManager``python/sglang/srt/managers/tokenizer_manager.py`
- `ModelRunner``python/sglang/srt/model_executor/model_runner.py`
## Why
- Downstream forks override one piece (tokenizer, KV cache, IPC, …).
- Inline logic forces them to copy the whole `__init__`, which rots against upstream.
- Splitting into `init_*` helpers lets them override exactly what they need.
- Reference shape: `TokenizerManager.__init__` in `python/sglang/srt/managers/tokenizer_manager.py`.
## Rules
- **`__init__` is an orchestrator.** Sequence of `self.init_*(...)` calls + minimal glue. No non-trivial construction inlined.
- **One helper per overridable unit.** Each `init_*` = one concern a subclass might swap. Don't lump.
- **Naming:** `init_<thing>` (snake_case, names the component). Conditional construction → `maybe_init_<thing>`, gate inside the helper.
- **No silent state coupling.** A helper only reads `self.*` set by earlier helpers. Ordering lives in `__init__`. Shared intermediates → pass as args, not via `self.*`.
- **New logic = new helper.** Default to adding `init_<thing>`, not another inline block. One-line `self.foo = server_args.foo` is fine; structured logic is not.
- **Preserve override points.** Prefer additive changes to existing `init_*` signatures. Breaking changes → call out in PR.
## Scope
Only the three classes listed above. Not other manager-style classes, not small dataclass/utility constructors.
+197 -158
View File
@@ -510,11 +510,7 @@ class Scheduler(
self.init_watch_dog_memory_saver_input_blocker()
# Init profiler
self.profiler_manager = SchedulerProfilerManager(
ps=self.ps,
dp_tp_cpu_group=self.dp_tp_cpu_group,
get_forward_ct=lambda: self.forward_ct,
)
self.init_profiler()
# Init prefill-decodedisaggregation
self.init_disaggregation()
@@ -528,176 +524,35 @@ class Scheduler(
# Init prefill kv split size when deterministic inference is enabled with various attention backends
self.init_deterministic_inference_config()
self.weight_updater = SchedulerWeightUpdaterManager(
tp_worker=self.tp_worker,
draft_worker=self.draft_worker,
tp_cpu_group=self.tp_cpu_group,
memory_saver_adapter=self.memory_saver_adapter,
flush_cache=self.flush_cache,
is_fully_idle=self.is_fully_idle,
)
self.init_weight_updater()
# Init request dispatcher
self.init_request_dispatcher()
# Init LoRA drainer for fair scheduling
if self.server_args.lora_drain_wait_threshold > 0.0:
self.lora_drainer = LoRADrainer(
server_args.max_loras_per_batch,
server_args.lora_drain_wait_threshold,
)
else:
self.lora_drainer = None
self.init_lora_drainer()
# Init LoRA overlap loader
if self.enable_lora_overlap_loading:
self.lora_overlap_loader = LoRAOverlapLoader(
self.tp_worker.model_runner.lora_manager
)
self.init_lora_overlap_loader()
# Init the grammar backend for constrained generation
self.grammar_manager = GrammarManager(self)
self.init_grammar_manager()
self.request_receiver = SchedulerRequestReceiver(
recv_from_tokenizer=self.ipc_channels.recv_from_tokenizer,
recv_from_rpc=self.ipc_channels.recv_from_rpc,
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,
attn_tp_cpu_group=self.attn_tp_cpu_group,
attn_cp_group=self.attn_cp_group,
attn_cp_cpu_group=self.attn_cp_cpu_group,
world_group=self.world_group,
server_args=self.server_args,
model_config=self.model_config,
max_recv_per_poll=self.max_recv_per_poll,
stream_output=lambda *a, **kw: self.output_streamer.stream_output(*a, **kw),
get_last_forward_mode=lambda: (
self.last_batch.forward_mode if self.last_batch is not None else None
),
)
self.init_request_receiver()
self.dp_attn_adapter = SchedulerDPAttnAdapter(
tp_group=self.tp_group,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags,
ps=self.ps,
server_args=self.server_args,
model_config=self.model_config,
enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm,
get_require_mlp_sync=lambda: self.require_mlp_sync,
)
self.init_dp_attn_adapter()
self.pool_stats_observer = SchedulerPoolStatsObserver(
tree_cache=self.tree_cache,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
req_to_token_pool=self.req_to_token_pool,
session_controller=self.session_controller,
hisparse_coordinator=self.hisparse_coordinator,
is_hybrid_swa=self.is_hybrid_swa,
is_hybrid_ssm=self.is_hybrid_ssm,
enable_hisparse=self.enable_hisparse,
full_tokens_per_layer=self.full_tokens_per_layer,
swa_tokens_per_layer=self.swa_tokens_per_layer,
max_total_num_tokens=self.max_total_num_tokens,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
)
self.init_pool_stats_observer()
self.invariant_checker = SchedulerInvariantChecker(
is_hybrid_swa=self.is_hybrid_swa,
is_hybrid_ssm=self.is_hybrid_ssm,
disaggregation_mode=self.disaggregation_mode,
page_size=self.page_size,
full_tokens_per_layer=self.full_tokens_per_layer,
swa_tokens_per_layer=self.swa_tokens_per_layer,
max_total_num_tokens=self.max_total_num_tokens,
server_args=self.server_args,
tree_cache=self.tree_cache,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
req_to_token_pool=self.req_to_token_pool,
pool_stats_observer=self.pool_stats_observer,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
)
self.init_invariant_checker()
self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=self.server_args.kv_events_config,
ps=self.ps,
attn_tp_rank=self.ps.attn_tp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
attn_dp_rank=self.ps.attn_dp_rank,
dp_rank=self.ps.dp_rank,
tree_cache=self.tree_cache,
send_metrics_from_scheduler=self.ipc_channels.send_metrics_from_scheduler,
max_running_requests=self.max_running_requests,
max_total_num_tokens=self.max_total_num_tokens,
get_stats=lambda: self.metrics_reporter.stats,
)
self.init_kv_events_publisher()
self.load_inquirer = SchedulerLoadInquirer(
disaggregation_mode=self.disaggregation_mode,
ps=self.ps,
server_args=self.server_args,
max_total_num_tokens=self.max_total_num_tokens,
max_running_requests=self.max_running_requests,
pool_stats_observer=self.pool_stats_observer,
tp_worker=self.tp_worker,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
spec_algorithm=self.spec_algorithm,
get_running_batch=lambda: self.running_batch,
get_waiting_queue=lambda: self.waiting_queue,
get_stats=lambda: self.metrics_reporter.stats,
get_chunked_req=lambda: self.chunked_req,
get_disagg_prefill_bootstrap_queue=lambda: self.disagg_prefill_bootstrap_queue,
get_disagg_prefill_inflight_queue=lambda: self.disagg_prefill_inflight_queue,
get_disagg_decode_prealloc_queue=lambda: self.disagg_decode_prealloc_queue,
get_disagg_decode_transfer_queue=lambda: self.disagg_decode_transfer_queue,
get_spec_total_num_accept_tokens=lambda: self.metrics_reporter.spec_total_num_accept_tokens,
get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct,
)
self.init_load_inquirer()
self.output_streamer = SchedulerOutputStreamer(
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
tree_cache=self.tree_cache,
ps=self.ps,
server_args=self.server_args,
is_generation=self.is_generation,
spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode,
enable_hicache_storage=lambda: self.enable_hicache_storage,
load_inquirer_get_loads=lambda req: self.load_inquirer.get_loads(req),
)
self.init_output_streamer()
self.batch_result_processor = SchedulerBatchResultProcessor(
is_generation=self.is_generation,
disaggregation_mode=self.disaggregation_mode,
enable_overlap=self.enable_overlap,
enable_overlap_mlx=self.enable_overlap_mlx,
server_args=self.server_args,
model_config=self.model_config,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator,
req_to_token_pool=self.req_to_token_pool,
decode_offload_manager=self.decode_offload_manager,
metrics_collector=self.metrics_collector,
metrics_reporter=self.metrics_reporter,
draft_worker=self.draft_worker,
model_worker=self.model_worker,
logprob_result_processor=SchedulerLogprobResultProcessor(
server_args=self.server_args, model_config=self.model_config
),
output_streamer=self.output_streamer,
abort_request=self.abort_request,
)
self.init_batch_result_processor()
self.is_initializing = False
@@ -1647,6 +1502,190 @@ class Scheduler(
if self.external_corpus_manager is not None:
self.external_corpus_manager.check_pending_load()
def init_profiler(self) -> None:
self.profiler_manager = SchedulerProfilerManager(
ps=self.ps,
dp_tp_cpu_group=self.dp_tp_cpu_group,
get_forward_ct=lambda: self.forward_ct,
)
def init_weight_updater(self) -> None:
self.weight_updater = SchedulerWeightUpdaterManager(
tp_worker=self.tp_worker,
draft_worker=self.draft_worker,
tp_cpu_group=self.tp_cpu_group,
memory_saver_adapter=self.memory_saver_adapter,
flush_cache=self.flush_cache,
is_fully_idle=self.is_fully_idle,
)
def init_lora_drainer(self) -> None:
if self.server_args.lora_drain_wait_threshold > 0.0:
self.lora_drainer = LoRADrainer(
self.server_args.max_loras_per_batch,
self.server_args.lora_drain_wait_threshold,
)
else:
self.lora_drainer = None
def init_lora_overlap_loader(self) -> None:
if self.enable_lora_overlap_loading:
self.lora_overlap_loader = LoRAOverlapLoader(
self.tp_worker.model_runner.lora_manager
)
def init_grammar_manager(self) -> None:
self.grammar_manager = GrammarManager(self)
def init_request_receiver(self) -> None:
self.request_receiver = SchedulerRequestReceiver(
recv_from_tokenizer=self.ipc_channels.recv_from_tokenizer,
recv_from_rpc=self.ipc_channels.recv_from_rpc,
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,
attn_tp_cpu_group=self.attn_tp_cpu_group,
attn_cp_group=self.attn_cp_group,
attn_cp_cpu_group=self.attn_cp_cpu_group,
world_group=self.world_group,
server_args=self.server_args,
model_config=self.model_config,
max_recv_per_poll=self.max_recv_per_poll,
stream_output=lambda *a, **kw: self.output_streamer.stream_output(*a, **kw),
get_last_forward_mode=lambda: (
self.last_batch.forward_mode if self.last_batch is not None else None
),
)
def init_dp_attn_adapter(self) -> None:
self.dp_attn_adapter = SchedulerDPAttnAdapter(
tp_group=self.tp_group,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags,
ps=self.ps,
server_args=self.server_args,
model_config=self.model_config,
enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm,
get_require_mlp_sync=lambda: self.require_mlp_sync,
)
def init_pool_stats_observer(self) -> None:
self.pool_stats_observer = SchedulerPoolStatsObserver(
tree_cache=self.tree_cache,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
req_to_token_pool=self.req_to_token_pool,
session_controller=self.session_controller,
hisparse_coordinator=self.hisparse_coordinator,
is_hybrid_swa=self.is_hybrid_swa,
is_hybrid_ssm=self.is_hybrid_ssm,
enable_hisparse=self.enable_hisparse,
full_tokens_per_layer=self.full_tokens_per_layer,
swa_tokens_per_layer=self.swa_tokens_per_layer,
max_total_num_tokens=self.max_total_num_tokens,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
)
def init_invariant_checker(self) -> None:
self.invariant_checker = SchedulerInvariantChecker(
is_hybrid_swa=self.is_hybrid_swa,
is_hybrid_ssm=self.is_hybrid_ssm,
disaggregation_mode=self.disaggregation_mode,
page_size=self.page_size,
full_tokens_per_layer=self.full_tokens_per_layer,
swa_tokens_per_layer=self.swa_tokens_per_layer,
max_total_num_tokens=self.max_total_num_tokens,
server_args=self.server_args,
tree_cache=self.tree_cache,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
req_to_token_pool=self.req_to_token_pool,
pool_stats_observer=self.pool_stats_observer,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
)
def init_kv_events_publisher(self) -> None:
self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=self.server_args.kv_events_config,
ps=self.ps,
attn_tp_rank=self.ps.attn_tp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
attn_dp_rank=self.ps.attn_dp_rank,
dp_rank=self.ps.dp_rank,
tree_cache=self.tree_cache,
send_metrics_from_scheduler=self.ipc_channels.send_metrics_from_scheduler,
max_running_requests=self.max_running_requests,
max_total_num_tokens=self.max_total_num_tokens,
get_stats=lambda: self.metrics_reporter.stats,
)
def init_load_inquirer(self) -> None:
self.load_inquirer = SchedulerLoadInquirer(
disaggregation_mode=self.disaggregation_mode,
ps=self.ps,
server_args=self.server_args,
max_total_num_tokens=self.max_total_num_tokens,
max_running_requests=self.max_running_requests,
pool_stats_observer=self.pool_stats_observer,
tp_worker=self.tp_worker,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
spec_algorithm=self.spec_algorithm,
get_running_batch=lambda: self.running_batch,
get_waiting_queue=lambda: self.waiting_queue,
get_stats=lambda: self.metrics_reporter.stats,
get_chunked_req=lambda: self.chunked_req,
get_disagg_prefill_bootstrap_queue=lambda: self.disagg_prefill_bootstrap_queue,
get_disagg_prefill_inflight_queue=lambda: self.disagg_prefill_inflight_queue,
get_disagg_decode_prealloc_queue=lambda: self.disagg_decode_prealloc_queue,
get_disagg_decode_transfer_queue=lambda: self.disagg_decode_transfer_queue,
get_spec_total_num_accept_tokens=lambda: self.metrics_reporter.spec_total_num_accept_tokens,
get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct,
)
def init_output_streamer(self) -> None:
self.output_streamer = SchedulerOutputStreamer(
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
tree_cache=self.tree_cache,
ps=self.ps,
server_args=self.server_args,
is_generation=self.is_generation,
spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode,
enable_hicache_storage=lambda: self.enable_hicache_storage,
load_inquirer_get_loads=lambda req: self.load_inquirer.get_loads(req),
)
def init_batch_result_processor(self) -> None:
self.batch_result_processor = SchedulerBatchResultProcessor(
is_generation=self.is_generation,
disaggregation_mode=self.disaggregation_mode,
enable_overlap=self.enable_overlap,
enable_overlap_mlx=self.enable_overlap_mlx,
server_args=self.server_args,
model_config=self.model_config,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator,
req_to_token_pool=self.req_to_token_pool,
decode_offload_manager=self.decode_offload_manager,
metrics_collector=self.metrics_collector,
metrics_reporter=self.metrics_reporter,
draft_worker=self.draft_worker,
model_worker=self.model_worker,
logprob_result_processor=SchedulerLogprobResultProcessor(
server_args=self.server_args, model_config=self.model_config
),
output_streamer=self.output_streamer,
abort_request=self.abort_request,
)
def init_req_max_new_tokens(self, req):
input_len = len(req.origin_input_ids)
# Keep this bound consistent with PrefillAdder's admission budget: