diff --git a/.claude/rules/modify-component-must-read.md b/.claude/rules/modify-component-must-read.md new file mode 100644 index 000000000..948cc2227 --- /dev/null +++ b/.claude/rules/modify-component-must-read.md @@ -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) diff --git a/.claude/rules/speculative-naming.md b/.claude/rules/speculative-naming.md deleted file mode 100644 index 90070870d..000000000 --- a/.claude/rules/speculative-naming.md +++ /dev/null @@ -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. diff --git a/.claude/skills/large-class-init-style/SKILL.md b/.claude/skills/large-class-init-style/SKILL.md new file mode 100644 index 000000000..752c23bdc --- /dev/null +++ b/.claude/skills/large-class-init-style/SKILL.md @@ -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_` (snake_case, names the component). Conditional construction → `maybe_init_`, 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_`, 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. diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4c5561e41..3884ca164 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: