[diffusion] chore: refactor warmup logic (#17027)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user