From c611a3fb782f6635b641e3a9988ab949e8b64509 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 4 May 2026 08:24:51 +0800 Subject: [PATCH] [diffusion] chore: disable VAE cpu offload by default (#24315) --- .../runtime/managers/gpu_worker.py | 69 ++++++++++++------- .../multimodal_gen/runtime/server_args.py | 12 ++-- .../test/unit/test_server_args.py | 50 +++++++++++++- 3 files changed, 98 insertions(+), 33 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 5f667d892..1a224d4f5 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -66,6 +66,14 @@ from sglang.srt.utils.network import NetworkAddress logger = init_logger(__name__) +OFFLOAD_DISABLE_RECOMMENDATION_ORDER = ( + "vae", + "image_encoder", + "text_encoder", + "text_encoder_2", + "transformer", +) + @dataclass class _ExpandedOutputParts: @@ -192,26 +200,7 @@ class GPUWorker: current_platform.get_device_total_memory() / (1024**3) - peak_reserved_gb ) can_stay_resident = self.get_can_stay_resident_components(remaining_gpu_mem_gb) - suggested_args = set() - component_to_arg = { - "vae": "--vae-cpu-offload", - "text_encoder": "--text-encoder-cpu-offload", - "text_encoder_2": "--text-encoder-cpu-offload", - "image_encoder": "--image-encoder-cpu-offload", - } - - for component in can_stay_resident: - if component == "transformer": - if self.server_args.dit_layerwise_offload: - suggested_args.add("--dit-layerwise-offload") - elif self.server_args.dit_cpu_offload: - suggested_args.add("--dit-cpu-offload") - elif component in component_to_arg: - suggested_args.add(component_to_arg[component]) - - suggested_args_str = ( - ", ".join(sorted(suggested_args)) if suggested_args else "None" - ) + suggested_args_str = self._format_offload_disable_suggestions(can_stay_resident) pool_overhead_gb = peak_reserved_gb - peak_allocated_gb @@ -224,6 +213,34 @@ class GPUWorker: f"Related offload server args to disable: {suggested_args_str}" ) + def _format_offload_disable_suggestions(self, components: List[str]) -> str: + component_set = set(components) + suggestions = [] + seen_args = set() + + for component in OFFLOAD_DISABLE_RECOMMENDATION_ORDER: + if component not in component_set: + continue + + arg = None + if component == "vae": + arg = "--vae-cpu-offload" + elif component == "image_encoder": + arg = "--image-encoder-cpu-offload" + elif component in ("text_encoder", "text_encoder_2"): + arg = "--text-encoder-cpu-offload" + elif component == "transformer": + if self.server_args.dit_layerwise_offload: + arg = "--dit-layerwise-offload" + elif self.server_args.dit_cpu_offload: + arg = "--dit-cpu-offload" + + if arg is not None and arg not in seen_args: + suggestions.append(arg) + seen_args.add(arg) + + return ", ".join(suggestions) if suggestions else "None" + def execute_forward( self, batch: List[Req], return_req: bool = False ) -> OutputBatch | Req: @@ -641,9 +658,9 @@ class GPUWorker: if not self.pipeline: return can_stay_resident - # Map memory_usage keys to server_args offload flags - # If the flag is False, the component is ALREADY resident, so we don't suggest it. - # If the flag is True, it is currently offloaded, so it's a candidate to "stay resident". + # Map memory_usage keys to server_args offload flags. + # If the flag is False, the component is already resident, so we do not suggest it. + # If the flag is True, it is currently offloaded, so it is a candidate to stay resident. offload_flags = { "transformer": self.server_args.dit_cpu_offload or self.server_args.dit_layerwise_offload, @@ -653,12 +670,16 @@ class GPUWorker: "image_encoder": self.server_args.image_encoder_cpu_offload, } - for name, usage in self.pipeline.memory_usages.items(): + for name in OFFLOAD_DISABLE_RECOMMENDATION_ORDER: # Only consider components that are currently configured to be offloaded is_offload_configured = offload_flags.get(name, False) if not is_offload_configured: continue + usage = self.pipeline.memory_usages.get(name) + if usage is None: + continue + if usage <= remaining_gpu_mem_gb: can_stay_resident.append(name) remaining_gpu_mem_gb -= usage diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 8b83a02c6..6aad57647 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -190,7 +190,7 @@ class ServerArgs(DisaggArgsMixin): dit_offload_prefetch_size: float = 0.0 text_encoder_cpu_offload: bool | None = None image_encoder_cpu_offload: bool | None = None - vae_cpu_offload: bool | None = None + vae_cpu_offload: bool | None = False use_fsdp_inference: bool = False pin_cpu_memory: bool = True ltx2_two_stage_device_mode: str | None = None @@ -382,15 +382,15 @@ class ServerArgs(DisaggArgsMixin): # TODO: to be handled by each platform if current_platform.get_device_total_memory() / BYTES_PER_GB < 30: - logger.info("Enabling all offloading for GPU with low device memory") + logger.info( + "Enabling large component offloading for GPU with low device memory" + ) if self.dit_cpu_offload is None: self.dit_cpu_offload = True if self.text_encoder_cpu_offload is None: self.text_encoder_cpu_offload = True if self.image_encoder_cpu_offload is None: self.image_encoder_cpu_offload = True - if self.vae_cpu_offload is None: - self.vae_cpu_offload = True elif self.pipeline_config.task_type.is_image_gen(): logger.info( "Disabling some offloading (except dit, text_encoder) for image generation model" @@ -401,8 +401,6 @@ class ServerArgs(DisaggArgsMixin): self.text_encoder_cpu_offload = True if self.image_encoder_cpu_offload is None: self.image_encoder_cpu_offload = False - if self.vae_cpu_offload is None: - self.vae_cpu_offload = False else: if self.dit_cpu_offload is None: self.dit_cpu_offload = True @@ -410,8 +408,6 @@ class ServerArgs(DisaggArgsMixin): self.text_encoder_cpu_offload = True if self.image_encoder_cpu_offload is None: self.image_encoder_cpu_offload = True - if self.vae_cpu_offload is None: - self.vae_cpu_offload = True def _adjust_ltx2_two_stage_device_mode(self): if not self._is_ltx23_two_stage_pipeline(): diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 17b6c84ee..a1a525c35 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -3,7 +3,10 @@ import sys import unittest from unittest.mock import patch -from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ModelTaskType, + PipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImagePipelineConfig, ) @@ -46,6 +49,51 @@ class TestServerArgsPathExpansion(unittest.TestCase): ) +class TestOffloadDefaults(unittest.TestCase): + def _from_dict_with_task_type( + self, + task_type, + *, + memory_gb=80, + kwargs=None, + ): + pipeline_config = PipelineConfig() + pipeline_config.task_type = task_type + with ( + patch.object(PipelineConfig, "from_kwargs", return_value=pipeline_config), + patch( + "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + return_value=False, + ), + patch( + "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + return_value=memory_gb * 1024**3, + ), + ): + return ServerArgs.from_dict({"model_path": "/fake", **(kwargs or {})}) + + def test_vae_cpu_offload_defaults_false_for_video_generation(self): + args = self._from_dict_with_task_type(ModelTaskType.T2V) + + self.assertFalse(args.vae_cpu_offload) + + def test_vae_cpu_offload_defaults_false_on_low_memory_gpu(self): + args = self._from_dict_with_task_type(ModelTaskType.T2V, memory_gb=16) + + self.assertFalse(args.vae_cpu_offload) + self.assertTrue(args.dit_cpu_offload) + self.assertTrue(args.text_encoder_cpu_offload) + self.assertTrue(args.image_encoder_cpu_offload) + + def test_explicit_vae_cpu_offload_true_is_preserved(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={"vae_cpu_offload": True}, + ) + + self.assertTrue(args.vae_cpu_offload) + + class TestModelIdResolution(unittest.TestCase): def setUp(self): _get_config_info.cache_clear()