[diffusion] model: optimize torch.compile (#17472)
This commit is contained in:
@@ -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.compile_module_with_torch_compile(transformer)
|
self._maybe_enable_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,
|
||||||
|
|||||||
Reference in New Issue
Block a user