From c5b9106c1aaf82648d9ea000b44ed470bb020715 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Tue, 16 Jun 2026 15:45:00 +0800 Subject: [PATCH] [perf] Use default torch compile mode for Wan2.2 T2V A14B (#28304) --- .../configs/models/dits/base.py | 8 +++++++ .../configs/models/dits/ltx_2.py | 1 + .../configs/pipeline_configs/__init__.py | 6 ++++- .../configs/pipeline_configs/base.py | 3 --- .../configs/pipeline_configs/ltx_2.py | 5 ++++ .../configs/pipeline_configs/wan.py | 1 + python/sglang/multimodal_gen/registry.py | 7 ++++-- .../pipelines_core/stages/denoising.py | 7 ++++-- .../model_specific_stages/hunyuan3d/paint.py | 7 ++++-- .../stages/model_specific_stages/mova.py | 23 +++++++++++++++---- 10 files changed, 53 insertions(+), 15 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/models/dits/base.py b/python/sglang/multimodal_gen/configs/models/dits/base.py index 5e77f0d7a..ea494a3f5 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/base.py +++ b/python/sglang/multimodal_gen/configs/models/dits/base.py @@ -59,6 +59,7 @@ class DiTConfig(ModelConfig): # sglang-diffusion DiT-specific parameters prefix: str = "" quant_config: QuantizationConfig | None = None + torch_compile_mode: str = "max-autotune-no-cudagraphs" @staticmethod def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any: @@ -78,5 +79,12 @@ class DiTConfig(ModelConfig): default=None, help="Quantization configuration for the DiT model", ) + parser.add_argument( + f"--{prefix}.torch-compile-mode", + type=str, + dest=f"{prefix.replace('-', '_')}.torch_compile_mode", + default=DiTConfig.torch_compile_mode, + help="torch.compile mode for the DiT model", + ) return parser diff --git a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py index f0318a559..22b036f0c 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py @@ -186,3 +186,4 @@ class LTX2Config(DiTConfig): arch_config: LTX2ArchConfig = field(default_factory=LTX2ArchConfig) prefix: str = "ltx2" + torch_compile_mode: str = "default" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py index 60de1090e..2f48ced11 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py @@ -34,7 +34,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.ideogram import ( from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( LingBotWorldCausalDMDConfig, ) -from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( + LTX2PipelineConfig, + LTX23PipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import ( @@ -75,5 +78,6 @@ __all__ = [ "SelfForcingWanT2V480PConfig", "ZImagePipelineConfig", "LTX2PipelineConfig", + "LTX23PipelineConfig", "LingBotWorldCausalDMDConfig", ] diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index c8b255064..039e95655 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -274,9 +274,6 @@ class PipelineConfig: # Wan2.2 TI2V parameters boundary_ratio: float | None = None - # Compilation - # enable_torch_compile: bool = False - # calculate the adjust size for condition image # width: original condition image width # height: original condition image height diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py index edde43fb7..20e7d6399 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py @@ -710,6 +710,11 @@ class LTX2PipelineConfig(PipelineConfig): return latents, audio_latents +@dataclasses.dataclass +class LTX23PipelineConfig(LTX2PipelineConfig): + """Configuration overrides for LTX-2.3.""" + + @dataclasses.dataclass class LTX2I2VPipelineConfig(LTX2PipelineConfig): task_type: ModelTaskType = ModelTaskType.TI2V diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py index 07d01f064..bf41e59f3 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py @@ -222,6 +222,7 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig): def __post_init__(self) -> None: self.dit_config.boundary_ratio = self.boundary_ratio + self.dit_config.torch_compile_mode = "default" @dataclass diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index 22f7d7cf6..11d7fb743 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -64,7 +64,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.ideogram import ( from sglang.multimodal_gen.configs.pipeline_configs.joy_image import ( JoyImageEditPipelineConfig, ) -from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( + LTX2PipelineConfig, + LTX23PipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.mova import ( MOVA360PConfig, MOVA720PConfig, @@ -647,7 +650,7 @@ def _register_configs(): ) register_configs( sampling_param_cls=LTX23SamplingParams, - pipeline_config_cls=LTX2PipelineConfig, + pipeline_config_cls=LTX23PipelineConfig, hf_model_paths=["Lightricks/LTX-2.3"], model_detectors=[ lambda path: "ltx-2.3" in path.lower(), diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 920412adb..10f739084 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -321,8 +321,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): _inductor_cfg.reorder_for_compute_comm_overlap = True except ImportError: pass - mode = os.environ.get( - "SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs" + dit_config = getattr(self.server_args.pipeline_config, "dit_config", None) + mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr( + dit_config, + "torch_compile_mode", + "max-autotune-no-cudagraphs", ) compile_kwargs["mode"] = mode logger.info(f"Compiling transformer with mode: {mode}") diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/paint.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/paint.py index 2c6851743..a031da099 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/paint.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/paint.py @@ -616,8 +616,11 @@ class Hunyuan3DPaintTexGenStage(PipelineStage): ddim_timesteps=30, ).to(self.device) if server_args.enable_torch_compile: - compile_mode = os.environ.get( - "SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs" + dit_config = getattr(server_args.pipeline_config, "dit_config", None) + compile_mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr( + dit_config, + "torch_compile_mode", + "max-autotune-no-cudagraphs", ) logger.info("Compiling paint transformer with mode: %s", compile_mode) self.transformer.compile(mode=compile_mode, fullgraph=False, dynamic=None) 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 8be9a042e..44cbcef63 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 @@ -235,7 +235,12 @@ class MOVADenoisingStage(PipelineStage): partial = (1 - guidance_scale) * neg return cfg_model_parallel_all_reduce(partial) - def _maybe_enable_torch_compile(self, module: nn.Module, server_args: ServerArgs): + def _maybe_enable_torch_compile( + self, + module: nn.Module, + server_args: ServerArgs, + model_config: object | None = None, + ): """ 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. @@ -266,8 +271,10 @@ class MOVADenoisingStage(PipelineStage): _inductor_cfg.reorder_for_compute_comm_overlap = True except ImportError: pass - mode = os.environ.get( - "SGLANG_TORCH_COMPILE_MODE", "max-autotune-no-cudagraphs" + mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr( + model_config, + "torch_compile_mode", + "max-autotune-no-cudagraphs", ) compile_kwargs["mode"] = mode logger.info("Compiling %s with mode: %s", module.__class__.__name__, mode) @@ -278,8 +285,14 @@ class MOVADenoisingStage(PipelineStage): 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._maybe_enable_torch_compile(module, server_args) + module_configs = [ + (self.video_dit, server_args.pipeline_config.dit_config), + (self.video_dit_2, server_args.pipeline_config.dit_config), + (self.audio_dit, server_args.pipeline_config.audio_dit_config), + ] + for module, model_config in module_configs: + if module is not None: + self._maybe_enable_torch_compile(module, server_args, model_config) self._torch_compiled = True def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult: