[diffusion] feat: enable torch compile to eliminate GPU bubble (#13641)
Co-authored-by: jianyingzhu <53300651@qq.com> Co-authored-by: Jianying <53503712+jianyingzhu@users.noreply.github.com> Co-authored-by: root <root@2u2g-spr-0417.ipp4a1.colossus.nvidia.com> Co-authored-by: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com>
This commit is contained in:
co-authored by
jianyingzhu
Jianying
root
Xiaoyu Zhang
parent
45eeeb9a3c
commit
fca8e88f28
@@ -110,9 +110,6 @@ class ComposedPipelineBase(ABC):
|
|||||||
self.post_init_called = True
|
self.post_init_called = True
|
||||||
|
|
||||||
self.initialize_pipeline(self.server_args)
|
self.initialize_pipeline(self.server_args)
|
||||||
if self.server_args.enable_torch_compile:
|
|
||||||
self.modules["transformer"] = torch.compile(self.modules["transformer"])
|
|
||||||
logger.info("Torch Compile enabled for DiT")
|
|
||||||
|
|
||||||
logger.info("Creating pipeline stages...")
|
logger.info("Creating pipeline stages...")
|
||||||
self.create_pipeline_stages(self.server_args)
|
self.create_pipeline_stages(self.server_args)
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ Denoising stage for diffusion pipelines.
|
|||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
import math
|
import math
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
import weakref
|
import weakref
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
@@ -116,17 +117,9 @@ class DenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
# torch compile
|
# torch compile
|
||||||
if self.server_args.enable_torch_compile:
|
if self.server_args.enable_torch_compile:
|
||||||
full_graph = False
|
self.torch_compile_module(self.transformer)
|
||||||
self.transformer = torch.compile(
|
if transformer_2 is not None:
|
||||||
self.transformer, mode="max-autotune", fullgraph=full_graph
|
self.torch_compile_module(self.transformer_2)
|
||||||
)
|
|
||||||
self.transformer_2 = (
|
|
||||||
torch.compile(
|
|
||||||
self.transformer_2, mode="max-autotune", fullgraph=full_graph
|
|
||||||
)
|
|
||||||
if transformer_2 is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
@@ -150,11 +143,31 @@ class DenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
# misc
|
# misc
|
||||||
self.profiler = None
|
self.profiler = None
|
||||||
|
|
||||||
# cache-dit state (for delayed mounting and idempotent control)
|
# cache-dit state (for delayed mounting and idempotent control)
|
||||||
self._cache_dit_enabled = False
|
self._cache_dit_enabled = False
|
||||||
self._cached_num_steps = None
|
self._cached_num_steps = None
|
||||||
|
|
||||||
|
def torch_compile_module(self, module):
|
||||||
|
"""
|
||||||
|
Compile a module's forward with torch.compile, and enable inductor overlap tweak if available.
|
||||||
|
No-op if torch compile is disabled or the object has no forward.
|
||||||
|
"""
|
||||||
|
if not self.server_args.enable_torch_compile or module is None:
|
||||||
|
return module
|
||||||
|
if not hasattr(module, "forward"):
|
||||||
|
return module
|
||||||
|
try:
|
||||||
|
import torch._inductor.config as _inductor_cfg
|
||||||
|
|
||||||
|
_inductor_cfg.reorder_for_compute_comm_overlap = True
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs")
|
||||||
|
logger.info(f"Compiling transformer with mode: {mode}")
|
||||||
|
compiled_forward = torch.compile(getattr(module, "forward"), mode=mode)
|
||||||
|
setattr(module, "forward", compiled_forward)
|
||||||
|
return module
|
||||||
|
|
||||||
def _maybe_enable_cache_dit(self, num_inference_steps: int) -> None:
|
def _maybe_enable_cache_dit(self, num_inference_steps: int) -> None:
|
||||||
"""Enable cache-dit on the transformers if configured (idempotent).
|
"""Enable cache-dit on the transformers if configured (idempotent).
|
||||||
|
|
||||||
@@ -581,7 +594,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
image_kwargs = self.prepare_extra_func_kwargs(
|
image_kwargs = self.prepare_extra_func_kwargs(
|
||||||
self.transformer.forward,
|
getattr(self.transformer, "forward", self.transformer),
|
||||||
{
|
{
|
||||||
# TODO: make sure on-device
|
# TODO: make sure on-device
|
||||||
"encoder_hidden_states_image": image_embeds,
|
"encoder_hidden_states_image": image_embeds,
|
||||||
@@ -590,7 +603,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||||
self.transformer.forward,
|
getattr(self.transformer, "forward", self.transformer),
|
||||||
{
|
{
|
||||||
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||||
@@ -605,7 +618,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
if batch.do_classifier_free_guidance:
|
if batch.do_classifier_free_guidance:
|
||||||
neg_cond_kwargs = self.prepare_extra_func_kwargs(
|
neg_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||||
self.transformer.forward,
|
getattr(self.transformer, "forward", self.transformer),
|
||||||
{
|
{
|
||||||
"encoder_hidden_states_2": batch.clip_embedding_neg,
|
"encoder_hidden_states_2": batch.clip_embedding_neg,
|
||||||
"encoder_attention_mask": batch.negative_attention_mask,
|
"encoder_attention_mask": batch.negative_attention_mask,
|
||||||
|
|||||||
Reference in New Issue
Block a user