From 504b2c58cf94db16813745383dc71750a4a816aa Mon Sep 17 00:00:00 2001 From: triple-mu Date: Wed, 18 Feb 2026 00:47:38 +0800 Subject: [PATCH] [diffusion] improve: improve torch.compile for MOVA (#18914) Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- .../stages/model_specific_stages/mova.py | 24 ++++++++++--------- 1 file changed, 13 insertions(+), 11 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py index e83336766..da3e8bc3a 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py @@ -17,6 +17,7 @@ import os from collections.abc import Iterable import torch +import torch.nn as nn from diffusers.utils.torch_utils import randn_tensor from tqdm.auto import tqdm @@ -200,11 +201,13 @@ class MOVADenoisingStage(PipelineStage): partial = (1 - guidance_scale) * neg return cfg_model_parallel_all_reduce(partial) - def compile_module_with_torch_compile(self, module, server_args: ServerArgs): - if not server_args.enable_torch_compile or module is None: - return module - if not hasattr(module, "forward"): - return module + def _maybe_enable_torch_compile(self, module: nn.Module, server_args: ServerArgs): + """ + 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. + """ + if not server_args.enable_torch_compile or not isinstance(module, nn.Module): + return try: import torch._inductor.config as _inductor_cfg @@ -213,15 +216,14 @@ class MOVADenoisingStage(PipelineStage): pass mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs") logger.info("Compiling %s with mode: %s", module.__class__.__name__, mode) - compiled_forward = torch.compile(getattr(module, "forward"), mode=mode) - setattr(module, "forward", compiled_forward) - return module + # TODO(triple-mu): support customized fullgraph and dynamic in the future + module.compile(mode=mode, fullgraph=False, dynamic=None) def _maybe_compile_dits(self, server_args: ServerArgs): if self._torch_compiled or not server_args.enable_torch_compile: return for module in filter(None, [self.video_dit, self.video_dit_2, self.audio_dit]): - self.compile_module_with_torch_compile(module, server_args) + self._maybe_enable_torch_compile(module, server_args) self._torch_compiled = True def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult: @@ -310,8 +312,8 @@ class MOVADenoisingStage(PipelineStage): def _manage_device_placement( self, - model_to_use: torch.nn.Module | None, - model_to_offload: torch.nn.Module | None, + model_to_use: nn.Module | None, + model_to_offload: nn.Module | None, server_args: ServerArgs, ): if not server_args.dit_cpu_offload: