[diffusion] fix: mount Cache-DiT before torch.compile in native denoising (#25328)
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user