[diffusion] chore: disable VAE cpu offload by default (#24315)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user