[diffusion] feat: add --offload-during-compile to fit max-autotune on tight-memory GPUs (#29862)

This commit is contained in:
Mick
2026-07-02 19:39:19 +08:00
committed by GitHub
parent 26b15a0825
commit 119b76567d
7 changed files with 156 additions and 8 deletions
@@ -11,6 +11,7 @@ import os
import time import time
import weakref import weakref
from collections.abc import Callable from collections.abc import Callable
from contextlib import contextmanager
from dataclasses import dataclass, field, fields from dataclasses import dataclass, field, fields
from functools import lru_cache from functools import lru_cache
from typing import Any from typing import Any
@@ -72,6 +73,10 @@ from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_c
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse, ComponentUse,
) )
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
is_layerwise_offloaded_module,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage, PipelineStage,
@@ -201,8 +206,12 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
attn_head_size = hidden_size // num_attention_heads attn_head_size = hidden_size // num_attention_heads
# torch compile # torch compile
# list of offloaded dit modules if torch compile is enabled. cleared after compile and warmup
self._offloaded_dit_modules_for_compile: list[torch.nn.Module] = []
# layerwise-offload and then compile to avoid OOM
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._maybe_offload_during_compile(transformer)
self._maybe_torch_compile(transformer)
self.scheduler = scheduler self.scheduler = scheduler
self.vae = vae self.vae = vae
@@ -298,7 +307,49 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
return name return name
return default_name return default_name
def _maybe_enable_torch_compile(self, module: object) -> None: def _maybe_offload_during_compile(self, module: object) -> None:
"""
Layerwise-offload the DiT so the compile warmup autotunes each layer
with only itself resident; forward() restores it after the warmup.
"""
args = self.server_args
if (
not args.enable_torch_compile
or not args.offload_during_compile
or not args.warmup
# a subclass with its own forward would never run the restore
or type(self).forward is not DenoisingStage.forward
or args.use_fsdp_inference
or envs.SGLANG_CACHE_DIT_ENABLED
or not isinstance(module, LayerwiseOffloadableModuleMixin)
or is_layerwise_offloaded_module(module)
):
return
module.configure_layerwise_offload(args)
if is_layerwise_offloaded_module(module):
self._offloaded_dit_modules_for_compile.append(module)
def _move_resident_components_for_warmup(self) -> list[torch.nn.Module]:
"""Move resident non-DiT components off-device while the warmup
denoising (the compile/autotune peak) runs; forward() moves them back."""
pipeline = self.pipeline() if self.pipeline is not None else None
if pipeline is None:
return []
dit_ids = {id(self.transformer), id(self.transformer_2)}
moved = []
for module in pipeline.modules.values():
if (
isinstance(module, torch.nn.Module)
and id(module) not in dit_ids
and not is_layerwise_offloaded_module(module)
):
param = next(module.parameters(), None)
if param is not None and param.device.type != "cpu":
module.to("cpu")
moved.append(module)
return moved
def _maybe_torch_compile(self, module: object) -> None:
""" """
Compile a module 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 is not a nn.Module. No-op if torch compile is disabled or the object is not a nn.Module.
@@ -354,7 +405,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
"""Apply request-dependent transformer acceleration in trace-safe order.""" """Apply request-dependent transformer acceleration in trace-safe order."""
self._maybe_enable_cache_dit(num_inference_steps, batch) self._maybe_enable_cache_dit(num_inference_steps, batch)
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._maybe_torch_compile(transformer)
@staticmethod @staticmethod
def _needs_nvfp4_jit_prewarm(module: nn.Module) -> bool: def _needs_nvfp4_jit_prewarm(module: nn.Module) -> bool:
@@ -1363,11 +1414,48 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
return latents return latents
@torch.no_grad() @contextmanager
def _offload_for_torch_compile_warmup(self, batch: Req):
"""wrap a denoise so the torch.compile warmup has spare VRAM.
No-op unless the DiT was layerwise-offloaded for compile.
Warmup: move resident non-DiT components off-device for the autotune window, restore
after.
"""
if not self._offloaded_dit_modules_for_compile:
yield
return
# if the dits have been layerwise-offloaded in preparation stage
if batch.is_warmup:
# if warmup is enabled, for warmup request, offload components to ensure sufficient VRAM headroom for torch compile
moved = self._move_resident_components_for_warmup()
try:
yield
finally:
device = get_local_torch_device()
for module in moved:
module.to(device)
return
# prepare for first real request: restore the user-configured residency
for module in self._offloaded_dit_modules_for_compile:
module.disable_offload()
# clear the list for avoid overhead during real request
self._offloaded_dit_modules_for_compile.clear()
yield
def forward( def forward(
self, self,
batch: Req, batch: Req,
server_args: ServerArgs, server_args: ServerArgs,
) -> Req:
with self._offload_for_torch_compile_warmup(batch):
return self._denoise(batch, server_args)
@torch.no_grad()
def _denoise(
self,
batch: Req,
server_args: ServerArgs,
) -> Req: ) -> Req:
""" """
Run the denoising loop. Run the denoising loop.
@@ -299,7 +299,7 @@ class Hunyuan3DShapeDenoisingStage(DenoisingStage):
server_args.model_paths["transformer"], server_args, "transformer" server_args.model_paths["transformer"], server_args, "transformer"
) )
self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch) self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
self._maybe_enable_torch_compile(self.transformer) self._maybe_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
@@ -268,7 +268,7 @@ class Ideogram4DenoisingStage(DenoisingStage):
pipeline=pipeline, pipeline=pipeline,
) )
self.unconditional_transformer = unconditional_transformer self.unconditional_transformer = unconditional_transformer
self._maybe_enable_torch_compile(self.unconditional_transformer) self._maybe_torch_compile(self.unconditional_transformer)
def _component_name_for_stage_module(self, module, default_name: str) -> str: def _component_name_for_stage_module(self, module, default_name: str) -> str:
if module is self.unconditional_transformer: if module is self.unconditional_transformer:
@@ -302,7 +302,7 @@ class Ideogram4DenoisingStage(DenoisingStage):
for transformer in filter( for transformer in filter(
None, [self.transformer, self.unconditional_transformer] None, [self.transformer, self.unconditional_transformer]
): ):
self._maybe_enable_torch_compile(transformer) self._maybe_torch_compile(transformer)
def _manage_unconditional_transformer_use_site(self, batch: Req) -> None: def _manage_unconditional_transformer_use_site(self, batch: Req) -> None:
manager = self._component_residency_manager manager = self._component_residency_manager
@@ -187,7 +187,7 @@ class Ideogram4ProgressiveDenoisingStage(
self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA
# Ideogram4DenoisingStage extra transformer # Ideogram4DenoisingStage extra transformer
self.unconditional_transformer = unconditional_transformer self.unconditional_transformer = unconditional_transformer
self._maybe_enable_torch_compile(self.unconditional_transformer) self._maybe_torch_compile(self.unconditional_transformer)
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Latent scale factor # Latent scale factor
@@ -216,6 +216,7 @@ class ServerArgs(DisaggServerArgsMixin):
dit_layerwise_offload: bool | None = None dit_layerwise_offload: bool | None = None
layerwise_offload_components: list[str] | None = None layerwise_offload_components: list[str] | None = None
dit_offload_prefetch_size: float = 0.0 dit_offload_prefetch_size: float = 0.0
offload_during_compile: bool = True
text_encoder_cpu_offload: bool | None = None text_encoder_cpu_offload: bool | None = None
image_encoder_cpu_offload: bool | None = None image_encoder_cpu_offload: bool | None = None
vae_cpu_offload: bool | None = False vae_cpu_offload: bool | None = False
@@ -1297,6 +1298,12 @@ class ServerArgs(DisaggServerArgsMixin):
help="Use torch.compile to speed up DiT inference." help="Use torch.compile to speed up DiT inference."
+ "However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)", + "However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
) )
parser.add_argument(
"--offload-during-compile",
action=StoreBoolean,
default=ServerArgs.offload_during_compile,
help="Offload components during the torch.compile warmup (the DiT layerwise) so max-autotune fits on tighter-memory GPUs, then restore the configured residency for serving. Skipped when the DiT is already layerwise-offloaded, or under cache-dit / FSDP.",
)
parser.add_argument( parser.add_argument(
"--enable-layerwise-nvtx-marker", "--enable-layerwise-nvtx-marker",
@@ -26,6 +26,13 @@ logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAS0lEQVR42u3PMQ0AAAwDoEqv9ErYvQQckD4XAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAYHLAB8+AWnmfUycAAAAAElFTkSuQmCC" MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAS0lEQVR42u3PMQ0AAAwDoEqv9ErYvQQckD4XAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAYHLAB8+AWnmfUycAAAAAElFTkSuQmCC"
def _is_ci_log_env() -> bool:
return (
os.environ.get("GITHUB_ACTIONS", "").lower() == "true"
or os.environ.get("CI", "").lower() == "true"
)
def get_first_generation_req(req_or_group: Any) -> Req | None: def get_first_generation_req(req_or_group: Any) -> Req | None:
"""Extract the first req""" """Extract the first req"""
if isinstance(req_or_group, Req): if isinstance(req_or_group, Req):
@@ -199,12 +206,20 @@ class SchedulerWarmupMixin:
if not self._show_warmup_progress: if not self._show_warmup_progress:
return return
ci_log_env = _is_ci_log_env()
if self._warmup_progress_bar is None: if self._warmup_progress_bar is None:
self._warmup_progress_bar = tqdm( self._warmup_progress_bar = tqdm(
total=self._warmup_progress_total(req_or_group), total=self._warmup_progress_total(req_or_group),
desc="Warmup requests", desc="Warmup requests",
unit="req", unit="req",
disable=ci_log_env,
) )
if ci_log_env:
logger.info(
"Warmup requests: 0/%s %s",
self._warmup_progress_bar.total,
self._format_warmup_req(req_or_group),
)
self._warmup_progress_bar.set_postfix_str( self._warmup_progress_bar.set_postfix_str(
self._format_warmup_req(req_or_group), refresh=False self._format_warmup_req(req_or_group), refresh=False
) )
@@ -225,6 +240,13 @@ class SchedulerWarmupMixin:
refresh=False, refresh=False,
) )
self._warmup_progress_bar.update(1) self._warmup_progress_bar.update(1)
if _is_ci_log_env():
logger.info(
"Warmup requests: %s/%s %s",
self._warmup_progress_bar.n,
self._warmup_progress_bar.total,
self._format_warmup_req(req_or_group),
)
if self._warmup_progress_bar.n >= self._warmup_progress_bar.total: if self._warmup_progress_bar.n >= self._warmup_progress_bar.total:
self._warmup_progress_bar.close() self._warmup_progress_bar.close()
@@ -81,6 +81,16 @@ _SERVER_FATAL_LOG_PATTERNS = (
"Segmentation fault", "Segmentation fault",
"Aborted (core dumped)", "Aborted (core dumped)",
) )
_CASE_LOG_SEPARATOR = "=" * 88
def _print_case_log_separator(case_id: str, state: str) -> None:
print(
f"\n{_CASE_LOG_SEPARATOR}\n"
f"[server-test] {state}: {case_id}\n"
f"{_CASE_LOG_SEPARATOR}",
flush=True,
)
@pytest.fixture @pytest.fixture
@@ -88,6 +98,7 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
"""Start a diffusion server for a single case and tear it down afterwards.""" """Start a diffusion server for a single case and tear it down afterwards."""
_fixture_start_time = time.perf_counter() _fixture_start_time = time.perf_counter()
server_args = case.server_args server_args = case.server_args
_print_case_log_separator(case.id, "BEGIN diffusion testcase")
# Skip ring attention tests on AMD/ROCm - Ring Attention requires Flash Attention # Skip ring attention tests on AMD/ROCm - Ring Attention requires Flash Attention
# which is not available on AMD. Use Ulysses parallelism instead. # which is not available on AMD. Use Ulysses parallelism instead.
@@ -203,6 +214,7 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
f"is not available in the installed version. " f"is not available in the installed version. "
f"Upgrade diffusers to enable this test." f"Upgrade diffusers to enable this test."
) )
_print_case_log_separator(case.id, "FAILED during server startup")
raise raise
try: try:
@@ -241,6 +253,7 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
f" }}\n" f" }}\n"
f'{"=" * 60}\n' f'{"=" * 60}\n'
) )
_print_case_log_separator(case.id, "END diffusion testcase")
class DiffusionServerBase: class DiffusionServerBase:
@@ -1199,6 +1212,24 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
- test_diffusion_generation[qwen_image_edit] - test_diffusion_generation[qwen_image_edit]
- etc. - etc.
""" """
try:
self._test_diffusion_generation_impl(case, diffusion_server)
except pytest.skip.Exception:
_print_case_log_separator(case.id, "SKIPPED diffusion testcase")
raise
except (KeyboardInterrupt, SystemExit):
raise
except BaseException:
_print_case_log_separator(case.id, "FAILED diffusion testcase")
raise
else:
_print_case_log_separator(case.id, "PASSED diffusion testcase")
def _test_diffusion_generation_impl(
self,
case: DiffusionTestCase,
diffusion_server: ServerContext,
):
# Check if we're in GT generation mode # Check if we're in GT generation mode
is_gt_gen_mode = os.environ.get("SGLANG_GEN_GT", "0") == "1" is_gt_gen_mode = os.environ.get("SGLANG_GEN_GT", "0") == "1"