[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__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
OFFLOAD_DISABLE_RECOMMENDATION_ORDER = (
|
||||||
|
"vae",
|
||||||
|
"image_encoder",
|
||||||
|
"text_encoder",
|
||||||
|
"text_encoder_2",
|
||||||
|
"transformer",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class _ExpandedOutputParts:
|
class _ExpandedOutputParts:
|
||||||
@@ -192,26 +200,7 @@ class GPUWorker:
|
|||||||
current_platform.get_device_total_memory() / (1024**3) - peak_reserved_gb
|
current_platform.get_device_total_memory() / (1024**3) - peak_reserved_gb
|
||||||
)
|
)
|
||||||
can_stay_resident = self.get_can_stay_resident_components(remaining_gpu_mem_gb)
|
can_stay_resident = self.get_can_stay_resident_components(remaining_gpu_mem_gb)
|
||||||
suggested_args = set()
|
suggested_args_str = self._format_offload_disable_suggestions(can_stay_resident)
|
||||||
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"
|
|
||||||
)
|
|
||||||
|
|
||||||
pool_overhead_gb = peak_reserved_gb - peak_allocated_gb
|
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}"
|
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(
|
def execute_forward(
|
||||||
self, batch: List[Req], return_req: bool = False
|
self, batch: List[Req], return_req: bool = False
|
||||||
) -> OutputBatch | Req:
|
) -> OutputBatch | Req:
|
||||||
@@ -641,9 +658,9 @@ class GPUWorker:
|
|||||||
if not self.pipeline:
|
if not self.pipeline:
|
||||||
return can_stay_resident
|
return can_stay_resident
|
||||||
|
|
||||||
# Map memory_usage keys to server_args offload flags
|
# 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 False, the component is already resident, so we do not suggest it.
|
||||||
# If the flag is True, it is currently offloaded, so it's a candidate to "stay resident".
|
# If the flag is True, it is currently offloaded, so it is a candidate to stay resident.
|
||||||
offload_flags = {
|
offload_flags = {
|
||||||
"transformer": self.server_args.dit_cpu_offload
|
"transformer": self.server_args.dit_cpu_offload
|
||||||
or self.server_args.dit_layerwise_offload,
|
or self.server_args.dit_layerwise_offload,
|
||||||
@@ -653,12 +670,16 @@ class GPUWorker:
|
|||||||
"image_encoder": self.server_args.image_encoder_cpu_offload,
|
"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
|
# Only consider components that are currently configured to be offloaded
|
||||||
is_offload_configured = offload_flags.get(name, False)
|
is_offload_configured = offload_flags.get(name, False)
|
||||||
if not is_offload_configured:
|
if not is_offload_configured:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
usage = self.pipeline.memory_usages.get(name)
|
||||||
|
if usage is None:
|
||||||
|
continue
|
||||||
|
|
||||||
if usage <= remaining_gpu_mem_gb:
|
if usage <= remaining_gpu_mem_gb:
|
||||||
can_stay_resident.append(name)
|
can_stay_resident.append(name)
|
||||||
remaining_gpu_mem_gb -= usage
|
remaining_gpu_mem_gb -= usage
|
||||||
|
|||||||
@@ -190,7 +190,7 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
dit_offload_prefetch_size: float = 0.0
|
dit_offload_prefetch_size: float = 0.0
|
||||||
text_encoder_cpu_offload: bool | None = None
|
text_encoder_cpu_offload: bool | None = None
|
||||||
image_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
|
use_fsdp_inference: bool = False
|
||||||
pin_cpu_memory: bool = True
|
pin_cpu_memory: bool = True
|
||||||
ltx2_two_stage_device_mode: str | None = None
|
ltx2_two_stage_device_mode: str | None = None
|
||||||
@@ -382,15 +382,15 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
|
|
||||||
# TODO: to be handled by each platform
|
# TODO: to be handled by each platform
|
||||||
if current_platform.get_device_total_memory() / BYTES_PER_GB < 30:
|
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:
|
if self.dit_cpu_offload is None:
|
||||||
self.dit_cpu_offload = True
|
self.dit_cpu_offload = True
|
||||||
if self.text_encoder_cpu_offload is None:
|
if self.text_encoder_cpu_offload is None:
|
||||||
self.text_encoder_cpu_offload = True
|
self.text_encoder_cpu_offload = True
|
||||||
if self.image_encoder_cpu_offload is None:
|
if self.image_encoder_cpu_offload is None:
|
||||||
self.image_encoder_cpu_offload = True
|
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():
|
elif self.pipeline_config.task_type.is_image_gen():
|
||||||
logger.info(
|
logger.info(
|
||||||
"Disabling some offloading (except dit, text_encoder) for image generation model"
|
"Disabling some offloading (except dit, text_encoder) for image generation model"
|
||||||
@@ -401,8 +401,6 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
self.text_encoder_cpu_offload = True
|
self.text_encoder_cpu_offload = True
|
||||||
if self.image_encoder_cpu_offload is None:
|
if self.image_encoder_cpu_offload is None:
|
||||||
self.image_encoder_cpu_offload = False
|
self.image_encoder_cpu_offload = False
|
||||||
if self.vae_cpu_offload is None:
|
|
||||||
self.vae_cpu_offload = False
|
|
||||||
else:
|
else:
|
||||||
if self.dit_cpu_offload is None:
|
if self.dit_cpu_offload is None:
|
||||||
self.dit_cpu_offload = True
|
self.dit_cpu_offload = True
|
||||||
@@ -410,8 +408,6 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
self.text_encoder_cpu_offload = True
|
self.text_encoder_cpu_offload = True
|
||||||
if self.image_encoder_cpu_offload is None:
|
if self.image_encoder_cpu_offload is None:
|
||||||
self.image_encoder_cpu_offload = True
|
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):
|
def _adjust_ltx2_two_stage_device_mode(self):
|
||||||
if not self._is_ltx23_two_stage_pipeline():
|
if not self._is_ltx23_two_stage_pipeline():
|
||||||
|
|||||||
@@ -3,7 +3,10 @@ import sys
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
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 (
|
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||||
QwenImagePipelineConfig,
|
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):
|
class TestModelIdResolution(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
_get_config_info.cache_clear()
|
_get_config_info.cache_clear()
|
||||||
|
|||||||
Reference in New Issue
Block a user