[diffusion] fix: fix --warmup-resolutions' conflict with CacheDiT (#16962)
This commit is contained in:
@@ -202,7 +202,6 @@ class Scheduler:
|
|||||||
prompt="",
|
prompt="",
|
||||||
is_warmup=True,
|
is_warmup=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.waiting_queue.append((None, req))
|
self.waiting_queue.append((None, req))
|
||||||
# if server is warmed-up, set this flag to avoid req-based warmup
|
# if server is warmed-up, set this flag to avoid req-based warmup
|
||||||
self.warmed_up = True
|
self.warmed_up = True
|
||||||
@@ -226,8 +225,7 @@ class Scheduler:
|
|||||||
warmup_req.set_as_warmup()
|
warmup_req.set_as_warmup()
|
||||||
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 = 0
|
||||||
logger.info("Processing warmup req... (1/1)")
|
|
||||||
self.warmed_up = True
|
self.warmed_up = True
|
||||||
return recv_reqs
|
return recv_reqs
|
||||||
|
|
||||||
|
|||||||
@@ -137,7 +137,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
setattr(module, "forward", compiled_forward)
|
setattr(module, "forward", compiled_forward)
|
||||||
return module
|
return module
|
||||||
|
|
||||||
def _maybe_enable_cache_dit(self, num_inference_steps: int) -> None:
|
def _maybe_enable_cache_dit(self, num_inference_steps: int, batch: Req) -> None:
|
||||||
"""Enable cache-dit on the transformers if configured (idempotent).
|
"""Enable cache-dit on the transformers if configured (idempotent).
|
||||||
|
|
||||||
This method should be called after the transformer is fully loaded
|
This method should be called after the transformer is fully loaded
|
||||||
@@ -158,7 +158,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
# check if cache-dit is enabled in config
|
# check if cache-dit is enabled in config
|
||||||
if not envs.SGLANG_CACHE_DIT_ENABLED:
|
if not envs.SGLANG_CACHE_DIT_ENABLED or batch.is_warmup:
|
||||||
return
|
return
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
@@ -496,13 +496,13 @@ class DenoisingStage(PipelineStage):
|
|||||||
server_args.model_paths["transformer"], server_args, "transformer"
|
server_args.model_paths["transformer"], server_args, "transformer"
|
||||||
)
|
)
|
||||||
# enable cache-dit before torch.compile (delayed mounting)
|
# enable cache-dit before torch.compile (delayed mounting)
|
||||||
self._maybe_enable_cache_dit(cache_dit_num_inference_steps)
|
self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
|
||||||
self.compile_module_with_torch_compile(self.transformer)
|
self.compile_module_with_torch_compile(self.transformer)
|
||||||
if pipeline:
|
if pipeline:
|
||||||
pipeline.add_module("transformer", self.transformer)
|
pipeline.add_module("transformer", self.transformer)
|
||||||
server_args.model_loaded["transformer"] = True
|
server_args.model_loaded["transformer"] = True
|
||||||
else:
|
else:
|
||||||
self._maybe_enable_cache_dit(cache_dit_num_inference_steps)
|
self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
|
||||||
|
|
||||||
# Prepare extra step kwargs for scheduler
|
# Prepare extra step kwargs for scheduler
|
||||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||||
|
|||||||
Reference in New Issue
Block a user