[diffusion] chore: clean scheduler (#24229)

This commit is contained in:
Mick
2026-05-02 09:30:06 +08:00
committed by GitHub
parent cb8fbd53fc
commit 193b977572
2 changed files with 70 additions and 46 deletions
@@ -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
@@ -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)