[diffusion] model: optimize torch.compile (#17472)

This commit is contained in:
triple-mu
2026-01-22 22:05:47 +08:00
committed by GitHub
parent 5d299c25c0
commit 3705f90629
@@ -15,6 +15,7 @@ from functools import lru_cache
from typing import Any from typing import Any
import torch import torch
import torch.nn as nn
from einops import rearrange from einops import rearrange
from tqdm.auto import tqdm from tqdm.auto import tqdm
@@ -92,9 +93,8 @@ class DenoisingStage(PipelineStage):
attn_head_size = hidden_size // num_attention_heads attn_head_size = hidden_size // num_attention_heads
# torch compile # torch compile
if self.server_args.enable_torch_compile: for transformer in filter(None, [self.transformer, self.transformer_2]):
for transformer in filter(None, [self.transformer, self.transformer_2]): self._maybe_enable_torch_compile(transformer)
self.compile_module_with_torch_compile(transformer)
self.scheduler = scheduler self.scheduler = scheduler
self.vae = vae self.vae = vae
@@ -116,15 +116,15 @@ class DenoisingStage(PipelineStage):
self._cached_num_steps = None self._cached_num_steps = None
self._is_warmed_up = False self._is_warmed_up = False
def compile_module_with_torch_compile(self, module): def _maybe_enable_torch_compile(self, module: object) -> None:
""" """
Compile a module's forward with torch.compile, and enable inductor overlap tweak if available. Compile a module with torch.compile, and enable inductor overlap tweak if available.
No-op if torch compile is disabled or the object has no forward. No-op if torch compile is disabled or the object is not a nn.Module.
""" """
if not self.server_args.enable_torch_compile or module is None: if not self.server_args.enable_torch_compile or not isinstance(
return module module, nn.Module
if not hasattr(module, "forward"): ):
return module return
try: try:
import torch._inductor.config as _inductor_cfg import torch._inductor.config as _inductor_cfg
@@ -133,9 +133,8 @@ class DenoisingStage(PipelineStage):
pass pass
mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs") mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs")
logger.info(f"Compiling transformer with mode: {mode}") logger.info(f"Compiling transformer with mode: {mode}")
compiled_forward = torch.compile(getattr(module, "forward"), mode=mode) # TODO(triple-mu): support customized fullgraph and dynamic in the future
setattr(module, "forward", compiled_forward) module.compile(mode=mode, fullgraph=False, dynamic=None)
return module
def _maybe_enable_cache_dit(self, num_inference_steps: int, batch: Req) -> 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).
@@ -497,7 +496,7 @@ class DenoisingStage(PipelineStage):
) )
# 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, batch) self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
self.compile_module_with_torch_compile(self.transformer) self._maybe_enable_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
@@ -801,8 +800,8 @@ class DenoisingStage(PipelineStage):
def _manage_device_placement( def _manage_device_placement(
self, self,
model_to_use: torch.nn.Module, model_to_use: nn.Module,
model_to_offload: torch.nn.Module | None, model_to_offload: nn.Module | None,
server_args: ServerArgs, server_args: ServerArgs,
): ):
""" """
@@ -1218,7 +1217,7 @@ class DenoisingStage(PipelineStage):
def _predict_noise_with_cfg( def _predict_noise_with_cfg(
self, self,
current_model: torch.nn.Module, current_model: nn.Module,
latent_model_input: torch.Tensor, latent_model_input: torch.Tensor,
timestep, timestep,
batch: Req, batch: Req,