[diffusion] feat: support torch compile for diffusers backend (#19673)

This commit is contained in:
DefTruth
2026-03-04 14:08:45 +08:00
committed by GitHub
parent c6850ac30c
commit e9b5706545
5 changed files with 74 additions and 9 deletions

View File

@@ -5,7 +5,6 @@ import asyncio
import os
import pickle
from collections import deque
from copy import deepcopy
from typing import Any, List
import zmq
@@ -233,7 +232,7 @@ class Scheduler:
height=height,
prompt="",
)
req.set_as_warmup()
req.set_as_warmup(self.server_args.warmup_steps)
self.waiting_queue.append((None, req))
# if server is warmed-up, set this flag to avoid req-based warmup
self.warmed_up = True
@@ -253,8 +252,7 @@ class Scheduler:
# only the very first req through server's lifetime will be warmed up
identity, req = recv_reqs[0]
if isinstance(req, Req):
warmup_req = deepcopy(req)
warmup_req.set_as_warmup()
warmup_req = req.copy_as_warmup(self.server_args.warmup_steps)
recv_reqs.insert(0, (identity, warmup_req))
self._warmup_total = 1
self._warmup_processed = 0

View File

@@ -378,6 +378,7 @@ class DiffusersPipeline(ComposedPipelineBase):
self.memory_usages: dict[str, float] = {}
self.post_init_called = False
self.executor = executor or SyncExecutor(server_args=server_args)
self._cache_dit_enabled = False
logger.info("Loading diffusers pipeline from %s", model_path)
self.diffusers_pipe = self._load_diffusers_pipeline(model_path, server_args)
@@ -462,6 +463,8 @@ class DiffusersPipeline(ComposedPipelineBase):
self._apply_attention_backend(pipe, server_args)
# Apply cache-dit acceleration if configured
pipe = self._apply_cache_dit(pipe, server_args)
# Apply torch.compile if enabled and supported
pipe = self._apply_torch_compile(pipe, server_args)
logger.info("Loaded diffusers pipeline: %s", pipe.__class__.__name__)
return pipe
@@ -562,6 +565,58 @@ class DiffusersPipeline(ComposedPipelineBase):
raise
logger.info("Enabled cache-dit for diffusers pipeline")
self._cache_dit_enabled = True
return pipe
def _apply_torch_compile(self, pipe: Any, server_args: ServerArgs) -> Any:
"""Apply torch.compile to the pipeline if configured and supported."""
if not server_args.enable_torch_compile:
return pipe
# check if the pipeline has 'transformer' or 'unet' components which are
# typically the most expensive parts to compile. 'transformer_2' for some
# video pipelines, e.g, Wan 2.2 series, also check for that.
compilable_components = ["transformer", "transformer_2", "unet"]
if not any(hasattr(pipe, comp) for comp in compilable_components):
logger.warning(
"Pipeline does not have 'transformer' or 'unet' components. "
"torch.compile may not provide significant benefits and could increase latency."
)
return pipe
if self._cache_dit_enabled:
try:
import cache_dit
if hasattr(cache_dit, "set_compile_configs"):
cache_dit.set_compile_configs()
except Exception as e:
logger.warning(
f"Failed to set torch_compile configs for cache-dit: {e}"
)
for comp in compilable_components:
if hasattr(pipe, comp):
try:
component = getattr(pipe, comp)
# TODO(DefTruth): Add support for 'compile_repeated_blocks' for 'transformer'
# modules which can significantly reduce compilation time for large models
# with repeated blocks.
if isinstance(component, torch.nn.Module) and hasattr(
component, "compile"
):
# Prefer in-place compilation if supported. According to PyTorch documentation:
# https://docs.pytorch.org/docs/stable/generated/torch.compile.html
component.compile()
else:
compiled_component = torch.compile(component)
setattr(pipe, comp, compiled_component)
logger.info(
f"Applied torch.compile to {comp} component of the pipeline"
)
except Exception as e:
logger.warning(f"Failed to apply torch.compile to {comp}: {e}")
return pipe
def _get_dtype(self, server_args: ServerArgs) -> torch.dtype:

View File

@@ -13,6 +13,7 @@ from __future__ import annotations
import os
import pprint
from copy import deepcopy
from dataclasses import MISSING, asdict, dataclass, field, fields
from typing import Any, Optional
@@ -251,12 +252,17 @@ class Req:
return None
return os.path.join(self.output_path, output_file_name)
def set_as_warmup(self):
def set_as_warmup(self, warmup_steps: int = 1):
self.is_warmup = True
self.save_output = False
self.suppress_logs = True
self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps
self.num_inference_steps = 1
self.num_inference_steps = warmup_steps
def copy_as_warmup(self, warmup_steps: int = 1) -> "Req":
req = deepcopy(self)
req.set_as_warmup(warmup_steps)
return req
def validate(self):
"""Initialize dependent fields after dataclass initialization."""
@@ -270,9 +276,6 @@ class Req:
self.metrics = RequestMetrics(request_id=self.request_id)
if self.is_warmup:
self.set_as_warmup()
def adjust_size(self, server_args: ServerArgs):
pass

View File

@@ -274,6 +274,7 @@ class ServerArgs:
# warmup
warmup: bool = False
warmup_resolutions: list[str] = None
warmup_steps: int = 1
disable_autocast: bool | None = None
@@ -755,6 +756,12 @@ class ServerArgs:
default=ServerArgs.warmup_resolutions,
help="Specify resolutions for server to warmup. e.g., `--warmup-resolutions 256x256, 720x720`",
)
parser.add_argument(
"--warmup-steps",
type=int,
default=ServerArgs.warmup_steps,
help="The number of warmup steps to perform for each resolution.",
)
parser.add_argument(
"--dit-cpu-offload",