[diffusion] fix: mount Cache-DiT before torch.compile in native denoising (#25328)

This commit is contained in:
Chi McIsaac
2026-05-15 18:08:39 +08:00
committed by GitHub
parent eec5ba26cf
commit 20123e0b16
@@ -172,6 +172,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
super().__init__() super().__init__()
self.transformer = transformer self.transformer = transformer
self.transformer_2 = transformer_2 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 hidden_size = self.server_args.pipeline_config.dit_config.hidden_size
num_attention_heads = ( num_attention_heads = (
@@ -199,9 +203,6 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# misc # misc
self.profiler = None 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._is_warmed_up = False
self._extra_func_kwarg_names_cache: dict[int, tuple[bool, frozenset[str]]] = {} self._extra_func_kwarg_names_cache: dict[int, tuple[bool, frozenset[str]]] = {}
@@ -277,6 +278,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
module, nn.Module module, nn.Module
): ):
return 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} compile_kwargs: dict[str, Any] = {"fullgraph": False, "dynamic": None}
if current_platform.is_npu(): if current_platform.is_npu():
@@ -306,6 +314,15 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# TODO(triple-mu): support customized fullgraph and dynamic in the future # TODO(triple-mu): support customized fullgraph and dynamic in the future
module.compile(**compile_kwargs) 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 @staticmethod
def _needs_nvfp4_jit_prewarm(module: nn.Module) -> bool: def _needs_nvfp4_jit_prewarm(module: nn.Module) -> bool:
@@ -352,8 +369,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
) )
return return
# check if cache-dit is enabled in config # Keep cache-dit disabled for ordinary warmup, but allow torch.compile
if not envs.SGLANG_CACHE_DIT_ENABLED or batch.is_warmup: # 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 return
world_size = get_world_size() world_size = get_world_size()
@@ -593,20 +613,22 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
else: else:
cache_dit_num_inference_steps = num_inference_steps 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 # FIXME: reuse more code
loader = TransformerLoader() loader = TransformerLoader()
self.transformer = loader.load( self.transformer = loader.load(
server_args.model_paths["transformer"], server_args, "transformer" 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_cache_dit_and_torch_compile(
self._maybe_enable_torch_compile(self.transformer) cache_dit_num_inference_steps, batch
)
if not transformer_was_loaded:
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:
self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
if batch.rollout: if batch.rollout:
self._maybe_prepare_rollout(batch) self._maybe_prepare_rollout(batch)