[diffusion] chore: refactor warmup logic (#17027)

This commit is contained in:
Mick
2026-01-14 11:35:06 +08:00
committed by GitHub
parent 2122fea3c4
commit 9524040220
2 changed files with 10 additions and 8 deletions
@@ -219,13 +219,11 @@ class Scheduler:
return recv_reqs return recv_reqs
# handle server req-based warmup by inserting an identical req to the beginning of the waiting queue # 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 warmup # only the very first req through server's lifetime will be warmed up
identity, req = recv_reqs[0] identity, req = recv_reqs[0]
if isinstance(req, Req): if isinstance(req, Req):
warmup_req = deepcopy(req) warmup_req = deepcopy(req)
warmup_req.is_warmup = True warmup_req.set_as_warmup()
warmup_req.extra["cache_dit_num_inference_steps"] = req.num_inference_steps
warmup_req.num_inference_steps = 1
recv_reqs.insert(0, (identity, warmup_req)) recv_reqs.insert(0, (identity, warmup_req))
self._warmup_total = 1 self._warmup_total = 1
self._warmup_processed = 1 self._warmup_processed = 1
@@ -152,8 +152,7 @@ class Req:
for name, value in kwargs.items(): for name, value in kwargs.items():
setattr(self, name, value) setattr(self, name, value)
if hasattr(self, "__post_init__"): self.validate()
self.__post_init__()
def __getattr__(self, name: str) -> Any: def __getattr__(self, name: str) -> Any:
""" """
@@ -231,7 +230,12 @@ class Req:
else None else None
) )
def __post_init__(self): def set_as_warmup(self):
self.is_warmup = True
self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps
self.num_inference_steps = 1
def validate(self):
"""Initialize dependent fields after dataclass initialization.""" """Initialize dependent fields after dataclass initialization."""
# Set do_classifier_free_guidance based on guidance scale and negative prompt # Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0 and self.negative_prompt is not None: if self.guidance_scale > 1.0 and self.negative_prompt is not None:
@@ -244,7 +248,7 @@ class Req:
self.timings = RequestTimings(request_id=self.request_id) self.timings = RequestTimings(request_id=self.request_id)
if self.is_warmup: if self.is_warmup:
self.num_inference_steps = 1 self.set_as_warmup()
def adjust_size(self, server_args: ServerArgs): def adjust_size(self, server_args: ServerArgs):
pass pass