[diffusion] chore: disable VAE cpu offload by default (#24315)

This commit is contained in:
Mick
2026-05-04 08:24:51 +08:00
committed by GitHub
parent 00d620b77d
commit c611a3fb78
3 changed files with 98 additions and 33 deletions
@@ -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
@@ -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():
@@ -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()