From 193b9775723d2a18bbcdbd5e79bf8b92a1169edf Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 2 May 2026 09:30:06 +0800 Subject: [PATCH] [diffusion] chore: clean scheduler (#24229) --- .../runtime/managers/gpu_worker.py | 8 +- .../runtime/managers/scheduler.py | 108 ++++++++++-------- 2 files changed, 70 insertions(+), 46 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 389b65087..765bfcbfc 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -241,7 +241,7 @@ class GPUWorker: raise ValueError( "Grouped execute_forward does not support return_req=True" ) - # batched reqs is only possible with `num_outputs_per_prompt > 1` now + # grouped reqs currently come only from expanded num_outputs_per_prompt self._validate_group_forward_reqs(batch) return self._execute_forward_batch(batch) @@ -294,6 +294,7 @@ class GPUWorker: start_time = time.monotonic() + # capture memory baseline for each req in grouped forward on rank-0 request_metrics = [ item.metrics for item in log_reqs if item.metrics is not None ] @@ -311,6 +312,8 @@ class GPUWorker: ) result = forward_fn() + # disagg roles return raw Req so callers can keep and transfer intermediate tensors + # before converting it to OutputBatch if return_req and isinstance(result, Req): return result @@ -334,6 +337,8 @@ class GPUWorker: for metrics in output_metrics: metrics.total_duration_ms = duration_ms + # file-path-only responses avoid serializing generated tensors between + # scheduler_client and gpu_worker. if req.save_output and req.return_file_paths_only: save_output_paths(output_batch) output_batch.output = None @@ -350,6 +355,7 @@ class GPUWorker: if not req.is_warmup: PerformanceLogger.log_request_summary(metrics=output_batch.metrics) + # dump per-request perf report to the server-mode file path. if ( req.perf_dump_path is not None and not req.is_warmup diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 1b55f9119..ce60cde55 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -115,7 +115,6 @@ class Scheduler(SchedulerDisaggMixin): MergeLoraWeightsReq: self._handle_merge_lora, UnmergeLoraWeightsReq: self._handle_unmerge_lora, Req: self._handle_generation, - list: self._handle_generation, ListLorasReq: self._handle_list_loras, ShutdownReq: self._handle_shutdown, GetDisaggStatsReq: self._handle_get_disagg_stats, @@ -195,9 +194,64 @@ class Scheduler(SchedulerDisaggMixin): checksums = self.worker.get_weights_checksum(module_names=req.module_names) return OutputBatch(output=checksums) - def _handle_generation(self, reqs: List[Req] | list[list[Req]]): + @staticmethod + def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]: if len(reqs) == 1 and isinstance(reqs[0], list): - reqs = reqs[0] + return reqs[0] + return reqs + + @staticmethod + def _first_generation_req(req_or_group: Any) -> Req | None: + """Extract the first req""" + if isinstance(req_or_group, Req): + return req_or_group + if isinstance(req_or_group, list) and req_or_group: + first_req = req_or_group[0] + if isinstance(first_req, Req): + return first_req + return None + + @classmethod + def _is_warmup_item(cls, req_or_group: Any) -> bool: + req = cls._first_generation_req(req_or_group) + return req.is_warmup if req is not None else False + + def _dispatch_request(self, reqs: list[Any]) -> OutputBatch: + """dispatch req to its registered handler""" + req_or_group = reqs[0] + if isinstance(req_or_group, list): + return self._handle_generation(reqs) + + handler = self.request_handlers.get(type(req_or_group)) + if handler is None: + return OutputBatch(error=f"Unknown request type: {type(req_or_group)}") + return handler(reqs) + + def _log_warmup_result(self, output_batch: OutputBatch, is_warmup: bool) -> None: + if not is_warmup: + return + + if output_batch.error is None: + if self._warmup_total > 0: + logger.info( + f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds", + output_batch.metrics.total_duration_s, + ) + else: + logger.info( + f"Warmup req processed in {GREEN}%.2f{RESET} seconds", + output_batch.metrics.total_duration_s, + ) + else: + if self._warmup_total > 0: + logger.info( + f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed" + ) + else: + logger.info("Warmup req processing failed") + + def _handle_generation(self, reqs: list[Any]): + reqs = self._normalize_generation_reqs(reqs) warmup_reqs = [req for req in reqs if req.is_warmup] if warmup_reqs: self._warmup_processed += len(warmup_reqs) @@ -345,8 +399,8 @@ class Scheduler(SchedulerDisaggMixin): # handle server req-based warmup by inserting an identical req to the beginning of the waiting queue # only the very first req through server's lifetime will be warmed up identity, req_or_group = recv_reqs[0] - req = req_or_group[0] if isinstance(req_or_group, list) else req_or_group - if isinstance(req, Req): + req = self._first_generation_req(req_or_group) + if req is not None: warmup_req = req.copy_as_warmup(self.server_args.warmup_steps) recv_reqs.insert(0, (identity, warmup_req)) self._warmup_total = 1 @@ -467,21 +521,9 @@ class Scheduler(SchedulerDisaggMixin): reqs = [item[1] for item in items] try: - first_req = reqs[0] - if isinstance(first_req, list) and first_req: - is_warmup = first_req[0].is_warmup - else: - is_warmup = ( - first_req.is_warmup if isinstance(first_req, Req) else False - ) - - handler = self.request_handlers.get(type(first_req)) - if handler: - output_batch = handler(reqs) - else: - output_batch = OutputBatch( - error=f"Unknown request type: {type(first_req)}" - ) + req_or_group = reqs[0] + is_warmup = self._is_warmup_item(req_or_group) + output_batch = self._dispatch_request(reqs) except Exception as e: logger.error( f"Error executing request in scheduler event loop: {e}", @@ -491,31 +533,7 @@ class Scheduler(SchedulerDisaggMixin): # 3. return results try: - if isinstance(first_req, list) and first_req: - is_warmup = first_req[0].is_warmup - else: - is_warmup = ( - first_req.is_warmup if isinstance(first_req, Req) else False - ) - if is_warmup: - if output_batch.error is None: - if self._warmup_total > 0: - logger.info( - f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds", - output_batch.metrics.total_duration_s, - ) - else: - logger.info( - f"Warmup req processed in {GREEN}%.2f{RESET} seconds", - output_batch.metrics.total_duration_s, - ) - else: - if self._warmup_total > 0: - logger.info( - f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed" - ) - else: - logger.info("Warmup req processing failed") + self._log_warmup_result(output_batch, is_warmup) # TODO: Support sending back to multiple identities if batched self.return_result(output_batch, identities[0], is_warmup=is_warmup)