diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 04bee16c6..962840768 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index a8a6dd0fe..7c95bc8b1 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -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