[diffusion] chore: refactor warmup logic (#17027)
This commit is contained in:
@@ -219,13 +219,11 @@ class Scheduler:
|
||||
return recv_reqs
|
||||
|
||||
# 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]
|
||||
if isinstance(req, Req):
|
||||
warmup_req = deepcopy(req)
|
||||
warmup_req.is_warmup = True
|
||||
warmup_req.extra["cache_dit_num_inference_steps"] = req.num_inference_steps
|
||||
warmup_req.num_inference_steps = 1
|
||||
warmup_req.set_as_warmup()
|
||||
recv_reqs.insert(0, (identity, warmup_req))
|
||||
self._warmup_total = 1
|
||||
self._warmup_processed = 1
|
||||
|
||||
@@ -152,8 +152,7 @@ class Req:
|
||||
for name, value in kwargs.items():
|
||||
setattr(self, name, value)
|
||||
|
||||
if hasattr(self, "__post_init__"):
|
||||
self.__post_init__()
|
||||
self.validate()
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
"""
|
||||
@@ -231,7 +230,12 @@ class Req:
|
||||
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."""
|
||||
# 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:
|
||||
@@ -244,7 +248,7 @@ class Req:
|
||||
self.timings = RequestTimings(request_id=self.request_id)
|
||||
|
||||
if self.is_warmup:
|
||||
self.num_inference_steps = 1
|
||||
self.set_as_warmup()
|
||||
|
||||
def adjust_size(self, server_args: ServerArgs):
|
||||
pass
|
||||
|
||||
Reference in New Issue
Block a user