[diffusion] fix: fix --warmup-resolutions' conflict with CacheDiT (#16962)

This commit is contained in:
HuangJi
2026-01-14 14:44:10 +08:00
committed by GitHub
parent c86ca12875
commit 030496eb06
2 changed files with 5 additions and 7 deletions
@@ -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(