[diffusion] feat: keep image-model auxiliary components resident under auto memory policy (#29649)

This commit is contained in:
Mick
2026-06-30 01:25:37 +08:00
committed by GitHub
parent e6c15f76f3
commit b0be644133
7 changed files with 85 additions and 51 deletions
@@ -164,6 +164,5 @@ class FastHunyuanConfig(HunyuanConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
auto_disable_component_offload_min_available_memory_gb=150, keep_resident_components=("vae",),
auto_disable_component_offload_components=("vae",),
) )
@@ -200,8 +200,8 @@ class LTX2PipelineConfig(PipelineConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
auto_disable_component_offload_min_available_memory_gb=70, keep_resident_min_available_gb=70,
auto_disable_component_offload_components=("dit",), keep_resident_components=("dit",),
auto_cfg_parallel_degree_by_num_gpus=((4, 1), (8, 1)), auto_cfg_parallel_degree_by_num_gpus=((4, 1), (8, 1)),
) )
@@ -14,13 +14,11 @@ class ModelDeploymentConfig:
auto_dit_layerwise_offload: bool = False auto_dit_layerwise_offload: bool = False
# if the available memory is bigger than this value, keep dit resident instead of apply layerwise-offload # if the available memory is bigger than this value, keep dit resident instead of apply layerwise-offload
auto_dit_layerwise_offload_high_memory_disable_gb: float | None = None auto_dit_layerwise_offload_high_memory_disable_gb: float | None = None
auto_disable_component_offload_min_available_memory_gb: float | None = None keep_resident_min_available_gb: float | None = None
# keep this explicit because large encoders can OOM even when DiT fits resident # only vae -- it is tiny so keeping it resident barely shifts memory; large
auto_disable_component_offload_components: tuple[OffloadComponentName, ...] = ( # encoders stay offloaded and dit placement stays with the FSDP/dit-layerwise
"dit", # policy
"text_encoder", keep_resident_components: tuple[OffloadComponentName, ...] = ("vae",)
"image_encoder",
)
fsdp_auto_min_available_memory_gb: float | None = None fsdp_auto_min_available_memory_gb: float | None = None
fsdp_auto_requires_cfg: bool = True fsdp_auto_requires_cfg: bool = True
fsdp_auto_requires_default_parallelism: bool = True fsdp_auto_requires_default_parallelism: bool = True
@@ -329,7 +329,7 @@ class SanaWMRealtimeConfig(SanaWMPipelineConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
auto_dit_layerwise_offload=True, auto_dit_layerwise_offload=True,
auto_disable_component_offload_min_available_memory_gb=120, keep_resident_min_available_gb=120,
auto_disable_component_offload_components=("dit",), keep_resident_components=("dit",),
auto_enable_cfg_parallel=False, auto_enable_cfg_parallel=False,
) )
@@ -115,8 +115,8 @@ class TurboWanT2V1_3B480PConfig(TurboWanT2V480PConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
auto_dit_layerwise_offload=True, auto_dit_layerwise_offload=True,
auto_disable_component_offload_min_available_memory_gb=60, keep_resident_min_available_gb=60,
auto_disable_component_offload_components=( keep_resident_components=(
"text_encoder", "text_encoder",
"image_encoder", "image_encoder",
"vae", "vae",
@@ -202,8 +202,8 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
auto_dit_layerwise_offload=True, auto_dit_layerwise_offload=True,
auto_disable_component_offload_min_available_memory_gb=60, keep_resident_min_available_gb=60,
auto_disable_component_offload_components=( keep_resident_components=(
"text_encoder", "text_encoder",
"image_encoder", "image_encoder",
"vae", "vae",
@@ -32,6 +32,12 @@ DEFAULT_LAYERWISE_COMPONENT_ARG_NAMES = (
(LAYERWISE_OFFLOAD_VAE_GROUP, "vae_cpu_offload"), (LAYERWISE_OFFLOAD_VAE_GROUP, "vae_cpu_offload"),
) )
# task-type defaults for keep_resident_min_available_gb when a model does not pin
# one: image vae is tiny so any datacenter gpu keeps it resident, video vae is
# larger so it only stays resident on very-high-memory gpus
IMAGE_GEN_KEEP_RESIDENT_MIN_AVAILABLE_GB = 45.0
DEFAULT_KEEP_RESIDENT_MIN_AVAILABLE_GB = 120.0
class ServerArgsAutoTuner: class ServerArgsAutoTuner:
"""Auto-tunes the server-arg for the given performance-mode, based on practical deployment experience with different model architectures""" """Auto-tunes the server-arg for the given performance-mode, based on practical deployment experience with different model architectures"""
@@ -46,6 +52,17 @@ class ServerArgsAutoTuner:
def _deployment_config(self) -> ModelDeploymentConfig: def _deployment_config(self) -> ModelDeploymentConfig:
return self.server_args.pipeline_config.get_model_deployment_config() return self.server_args.pipeline_config.get_model_deployment_config()
def _resolve_keep_resident_min_available_gb(
self, deployment_config: ModelDeploymentConfig
) -> float | None:
# explicit per-model > task-type default > global default
explicit = deployment_config.keep_resident_min_available_gb
if explicit is not None:
return explicit
if self.server_args.pipeline_config.task_type.is_image_gen():
return IMAGE_GEN_KEEP_RESIDENT_MIN_AVAILABLE_GB
return DEFAULT_KEEP_RESIDENT_MIN_AVAILABLE_GB
def adjust_based_on_performance_mode(self) -> None: def adjust_based_on_performance_mode(self) -> None:
"""Adjust the server args based on the performance mode""" """Adjust the server args based on the performance mode"""
args = self.server_args args = self.server_args
@@ -96,8 +113,8 @@ class ServerArgsAutoTuner:
min_available_gb = self._get_min_available_device_memory_gb() min_available_gb = self._get_min_available_device_memory_gb()
deployment_config = self._deployment_config() deployment_config = self._deployment_config()
disable_threshold_gb = ( disable_threshold_gb = self._resolve_keep_resident_min_available_gb(
deployment_config.auto_disable_component_offload_min_available_memory_gb deployment_config
) )
if ( if (
min_available_gb is not None min_available_gb is not None
@@ -105,7 +122,7 @@ class ServerArgsAutoTuner:
and min_available_gb >= disable_threshold_gb and min_available_gb >= disable_threshold_gb
): ):
changed = [] changed = []
components = deployment_config.auto_disable_component_offload_components components = deployment_config.keep_resident_components
if ( if (
args.layerwise_offload_components is not None args.layerwise_offload_components is not None
and not args.is_arg_explicitly_set("layerwise_offload_components") and not args.is_arg_explicitly_set("layerwise_offload_components")
@@ -400,9 +417,7 @@ class ServerArgsAutoTuner:
return components return components
deployment_config = self._deployment_config() deployment_config = self._deployment_config()
threshold_gb = ( threshold_gb = self._resolve_keep_resident_min_available_gb(deployment_config)
deployment_config.auto_disable_component_offload_min_available_memory_gb
)
if threshold_gb is None: if threshold_gb is None:
return components return components
@@ -410,9 +425,7 @@ class ServerArgsAutoTuner:
if min_available_gb is None or min_available_gb < threshold_gb: if min_available_gb is None or min_available_gb < threshold_gb:
return components return components
resident_components = set( resident_components = set(deployment_config.keep_resident_components)
deployment_config.auto_disable_component_offload_components
)
filtered_components = [ filtered_components = [
component component
for component in components for component in components
@@ -15,6 +15,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType, ModelTaskType,
PipelineConfig, PipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanConfig
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig, LTX2PipelineConfig,
LTX23PipelineConfig, LTX23PipelineConfig,
@@ -837,12 +838,8 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertTrue(zimage_deployment.fsdp_auto_requires_cfg) self.assertTrue(zimage_deployment.fsdp_auto_requires_cfg)
self.assertFalse(zimage_deployment.auto_dit_layerwise_offload) self.assertFalse(zimage_deployment.auto_dit_layerwise_offload)
self.assertEqual( self.assertEqual(ltx_deployment.keep_resident_min_available_gb, 70)
ltx_deployment.auto_disable_component_offload_min_available_memory_gb, 70 self.assertEqual(ltx_deployment.keep_resident_components, ("dit",))
)
self.assertEqual(
ltx_deployment.auto_disable_component_offload_components, ("dit",)
)
self.assertEqual( self.assertEqual(
ltx_deployment.auto_cfg_parallel_degree_by_num_gpus, ((4, 1), (8, 1)) ltx_deployment.auto_cfg_parallel_degree_by_num_gpus, ((4, 1), (8, 1))
) )
@@ -859,6 +856,15 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertEqual(sana_wm_deployment.fsdp_auto_min_available_memory_gb, 60) self.assertEqual(sana_wm_deployment.fsdp_auto_min_available_memory_gb, 60)
self.assertTrue(sana_wm_deployment.auto_dit_layerwise_offload) self.assertTrue(sana_wm_deployment.auto_dit_layerwise_offload)
# fasthunyuan no longer pins 150gb -- falls back to the global video default
fast_hunyuan_deployment = FastHunyuanConfig().get_model_deployment_config()
self.assertIsNone(fast_hunyuan_deployment.keep_resident_min_available_gb)
self.assertEqual(fast_hunyuan_deployment.keep_resident_components, ("vae",))
# default keeps only vae resident (encoders are large, dit owned by FSDP)
self.assertEqual(qwen_deployment.keep_resident_components, ("vae",))
self.assertIsNone(qwen_deployment.keep_resident_min_available_gb)
def test_auto_multi_gpu_sana_wm_prefers_fsdp_and_cfg_parallel(self): def test_auto_multi_gpu_sana_wm_prefers_fsdp_and_cfg_parallel(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
SanaWMPipelineConfig(), SanaWMPipelineConfig(),
@@ -948,7 +954,7 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertIsNone(args.image_encoder_cpu_offload) self.assertIsNone(args.image_encoder_cpu_offload)
self.assertFalse(args.enable_cfg_parallel) self.assertFalse(args.enable_cfg_parallel)
def test_default_auto_replaces_text_encoder_cpu_offload_with_layerwise(self): def test_default_auto_keeps_image_vae_resident_when_memory_allows(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
QwenImagePipelineConfig(), QwenImagePipelineConfig(),
kwargs={"model_path": "Qwen/Qwen-Image"}, kwargs={"model_path": "Qwen/Qwen-Image"},
@@ -956,10 +962,25 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertEqual(args.performance_mode, "auto") self.assertEqual(args.performance_mode, "auto")
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
# 80gb > image threshold (45gb): only vae kept resident, encoders stay
# offloaded layerwise, dit unchanged
self.assertTrue(args.dit_cpu_offload)
self.assertEqual(
args.layerwise_offload_components,
["text_encoder", "image_encoder"],
)
self.assertFalse(args.vae_cpu_offload)
def test_auto_image_offloads_aux_below_resident_threshold(self):
# 40gb < image threshold (45gb): aux incl. vae still offloaded to save vram
args = self._from_dict_with_pipeline_config(
QwenImagePipelineConfig(),
memory_gb=40,
kwargs={"model_path": "Qwen/Qwen-Image"},
)
self.assertEqual(args.performance_mode, "auto")
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertTrue(args.layerwise_offload_components)
self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder", "vae"],
@@ -1296,7 +1317,7 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertEqual(args.ltx2_two_stage_device_mode, "resident") self.assertEqual(args.ltx2_two_stage_device_mode, "resident")
self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) self.assertEqual(args.layerwise_offload_components, ["text_encoder"])
def test_auto_multi_gpu_qwen_replaces_text_encoder_offload_with_cfg(self): def test_auto_multi_gpu_qwen_keeps_vae_resident_with_cfg(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
QwenImagePipelineConfig(), QwenImagePipelineConfig(),
kwargs={ kwargs={
@@ -1308,14 +1329,14 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertTrue(args.enable_cfg_parallel) self.assertTrue(args.enable_cfg_parallel)
# 80gb > image threshold (45gb): only vae resident, encoders offloaded;
# cfg/dit unchanged
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertTrue(args.layerwise_offload_components)
self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder"],
) )
self.assertFalse(args.vae_cpu_offload)
def test_auto_multi_gpu_zimage_base_prefers_fsdp(self): def test_auto_multi_gpu_zimage_base_prefers_fsdp(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
@@ -1357,11 +1378,12 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertTrue(args.enable_cfg_parallel) self.assertTrue(args.enable_cfg_parallel)
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.vae_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload) # explicit use_fsdp_inference skips the residency pass, but the layerwise
# filter still drops vae (kept resident); encoders stay offloaded
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder"],
) )
def test_auto_multi_gpu_qwen_skips_fsdp_when_available_memory_is_low(self): def test_auto_multi_gpu_qwen_skips_fsdp_when_available_memory_is_low(self):
@@ -1377,13 +1399,14 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertTrue(args.enable_cfg_parallel) self.assertTrue(args.enable_cfg_parallel)
# 50gb still > image threshold (45gb): vae resident, encoders offloaded;
# fsdp skipped (qwen does not opt into auto fsdp)
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder"],
) )
self.assertFalse(args.vae_cpu_offload)
def test_auto_multi_gpu_qwen_uses_selected_gpu_min_available_memory(self): def test_auto_multi_gpu_qwen_uses_selected_gpu_min_available_memory(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
@@ -1400,7 +1423,7 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertTrue(args.enable_cfg_parallel) self.assertTrue(args.enable_cfg_parallel)
def test_auto_multi_gpu_qwen_replaces_text_encoder_offload_with_headroom(self): def test_auto_multi_gpu_qwen_keeps_vae_resident_with_headroom(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
QwenImagePipelineConfig(), QwenImagePipelineConfig(),
available_memory_gb={1: 72, 2: 80}, available_memory_gb={1: 72, 2: 80},
@@ -1414,13 +1437,14 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertTrue(args.enable_cfg_parallel) self.assertTrue(args.enable_cfg_parallel)
# min available across selected gpus is 72gb > image threshold (45gb):
# vae resident, encoders offloaded
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder"],
) )
self.assertFalse(args.vae_cpu_offload)
def test_speed_mode_single_gpu_disables_offload(self): def test_speed_mode_single_gpu_disables_offload(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(