From 20123e0b165556cf2b8ec05fac866a3bd2a29ccd Mon Sep 17 00:00:00 2001 From: Chi McIsaac <153383231+qimcis@users.noreply.github.com> Date: Fri, 15 May 2026 06:08:39 -0400 Subject: [PATCH] [diffusion] fix: mount Cache-DiT before torch.compile in native denoising (#25328) --- .../pipelines_core/stages/denoising.py | 44 ++++++++++++++----- 1 file changed, 33 insertions(+), 11 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index fa99647c8..fa272c7a3 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -172,6 +172,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): super().__init__() self.transformer = transformer self.transformer_2 = transformer_2 + # cache-dit state (for delayed mounting and idempotent control) + self._cache_dit_enabled = False + self._cached_num_steps = None + self._torch_compiled_module_ids: set[int] = set() hidden_size = self.server_args.pipeline_config.dit_config.hidden_size num_attention_heads = ( @@ -199,9 +203,6 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): # misc self.profiler = None - # cache-dit state (for delayed mounting and idempotent control) - self._cache_dit_enabled = False - self._cached_num_steps = None self._is_warmed_up = False self._extra_func_kwarg_names_cache: dict[int, tuple[bool, frozenset[str]]] = {} @@ -277,6 +278,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): module, nn.Module ): return + if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled: + logger.debug("Deferring torch.compile until cache-dit is enabled") + return + module_id = id(module) + if module_id in self._torch_compiled_module_ids: + return + compile_kwargs: dict[str, Any] = {"fullgraph": False, "dynamic": None} if current_platform.is_npu(): @@ -306,6 +314,15 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): # TODO(triple-mu): support customized fullgraph and dynamic in the future module.compile(**compile_kwargs) + self._torch_compiled_module_ids.add(module_id) + + def _maybe_enable_cache_dit_and_torch_compile( + self, num_inference_steps: int | tuple[int, int], batch: Req + ) -> None: + """Apply request-dependent transformer acceleration in trace-safe order.""" + self._maybe_enable_cache_dit(num_inference_steps, batch) + for transformer in filter(None, [self.transformer, self.transformer_2]): + self._maybe_enable_torch_compile(transformer) @staticmethod def _needs_nvfp4_jit_prewarm(module: nn.Module) -> bool: @@ -352,8 +369,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): ) return - # check if cache-dit is enabled in config - if not envs.SGLANG_CACHE_DIT_ENABLED or batch.is_warmup: + # Keep cache-dit disabled for ordinary warmup, but allow torch.compile + # warmup to mount cache-dit before Dynamo traces the transformer. + if not envs.SGLANG_CACHE_DIT_ENABLED: + return + if batch.is_warmup and not self.server_args.enable_torch_compile: return world_size = get_world_size() @@ -593,20 +613,22 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): else: cache_dit_num_inference_steps = num_inference_steps - if not server_args.model_loaded["transformer"]: + transformer_was_loaded = server_args.model_loaded["transformer"] + if not transformer_was_loaded: # FIXME: reuse more code loader = TransformerLoader() self.transformer = loader.load( server_args.model_paths["transformer"], server_args, "transformer" ) - # enable cache-dit before torch.compile (delayed mounting) - self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch) - self._maybe_enable_torch_compile(self.transformer) + + self._maybe_enable_cache_dit_and_torch_compile( + cache_dit_num_inference_steps, batch + ) + + if not transformer_was_loaded: if pipeline: pipeline.add_module("transformer", self.transformer) server_args.model_loaded["transformer"] = True - else: - self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch) if batch.rollout: self._maybe_prepare_rollout(batch)