From a9a355774a797a6f0377b0030f9662035a9f22e7 Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 12 Aug 2026 11:38:27 +0800 Subject: [PATCH] [diffusion] feat: support dynamically cpu offload components (#34391) --- docs/docs/sglang-diffusion/api/cli.mdx | 15 +- .../sglang-diffusion/deployment_cookbook.mdx | 2 +- .../configs/pipeline_configs/hunyuan.py | 3 +- .../configs/pipeline_configs/lingbot_world.py | 7 + .../model_deployment_config.py | 6 +- .../configs/pipeline_configs/wan.py | 9 +- .../configs/pipeline_configs/zimage.py | 5 +- .../component_loaders/adapter_loader.py | 5 +- .../loader/component_loaders/bridge_loader.py | 7 +- .../component_loaders/component_loader.py | 21 +- .../component_loaders/image_encoder_loader.py | 13 +- .../sound_tokenizer_loader.py | 10 +- .../component_loaders/text_encoder_loader.py | 16 +- .../component_loaders/transformer_loader.py | 5 + .../component_loaders/upsampler_loader.py | 5 +- .../loader/component_loaders/vae_loader.py | 8 +- .../component_loaders/vl_encoder_loader.py | 6 +- .../component_loaders/vocoder_loader.py | 8 +- .../runtime/managers/gpu_worker.py | 88 +++--- .../memory_managers/component_manager.py | 64 ++-- .../layerwise_offload_components.py | 50 ++++ .../pipelines_core/stages/causal_denoising.py | 2 + .../pipelines_core/stages/denoising.py | 6 +- .../stages/model_specific_stages/longlive2.py | 1 - .../runtime/server_args/auto_tune.py | 9 +- .../runtime/server_args/server_args.py | 81 ++++- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../unit/realtime/test_causal_denoising.py | 56 ++++ .../test/unit/test_layerwise_offload.py | 8 +- .../test/unit/test_nvtx_pytorch_hooks.py | 14 + .../test/unit/test_server_args.py | 278 ++++++++++++++++-- 31 files changed, 634 insertions(+), 176 deletions(-) diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index a8be23b96..ba14864bd 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -81,7 +81,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis - `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter - `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks. - `--num-gpus {N}`: number of GPUs to use -- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes. +- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and dispatches residency from selected-GPU headroom and workload type: image DiTs stay resident above the 45 GiB threshold, while video DiT placement remains model-specific. It uses FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes. - `--direct-gpu-weight-loading {true|false}`: opt into direct GPU loading for an unquantized, GPU-resident, TP=1 DiT by materializing its complete checkpoint state dict on GPU. Startup impact is model-dependent, so benchmark the target model before deployment. Disabled by default because checkpoint and model weights coexist temporarily, substantially increasing peak GPU memory. It is incompatible with DiT CPU/layerwise offload and FSDP. - `--tp-size {N}`: tensor parallelism size. Depending on the pipeline, it can shard the DiT, one or more encoders, or both. - `--sp-degree {N}`: sequence parallelism size @@ -197,6 +197,19 @@ For supported native pipelines, set `SGLANG_CACHE_DIT_ENABLED=true` to enable Ca For supported image pipelines, breakable CUDA graph can be enabled with `--enable-breakable-cuda-graph`, but you must declare every served resolution in `--warmup-resolutions` so warmup captures matching graph signatures. +### Component CPU Offload + +Use `--cpu-offload-components` to explicitly select components for coarse CPU offload: + +```bash Command +sglang generate \ + --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --cpu-offload-components dit text_encoder \ + --prompt "A quiet city street after rain" +``` + +Component names are matched against the loaded pipeline's component keys from `model_index.json`, so names such as `transformer_2`, `audio_vae`, and `connectors` are supported without adding them to a registry. The group aliases `dit`, `text_encoder`, `image_encoder`, and `vae` remain available. Use `all` to offload every loaded `torch.nn.Module`, or `none` to disable coarse component CPU offload. Non-module components such as tokenizers and schedulers are unaffected. This unified option cannot be combined with the legacy per-component CPU offload flags. Layerwise offload remains independently controlled by `--layerwise-offload-components`. + ### Layerwise Offload Use layerwise offload when a large component does not fit comfortably in GPU memory. By default, `--dit-layerwise-offload` only applies to legacy DiT components. Use `--layerwise-offload-components` to select pipeline component names explicitly (`--layerwise-offload-modules` is accepted as an alias): diff --git a/docs/docs/sglang-diffusion/deployment_cookbook.mdx b/docs/docs/sglang-diffusion/deployment_cookbook.mdx index 38f410705..0d4649c07 100644 --- a/docs/docs/sglang-diffusion/deployment_cookbook.mdx +++ b/docs/docs/sglang-diffusion/deployment_cookbook.mdx @@ -128,7 +128,7 @@ See [OpenAI API: Served model name](/docs/sglang-diffusion/api/openai_api#served -`auto` checks selected GPU memory before applying FSDP. In multi-GPU runs it uses the least available memory across selected GPUs, and only turns on FSDP automatically when doing so can replace DiT offload. Text encoder, image encoder, and other component residency still follow the offload policy unless the model marks a high-memory resident path as safe. When the model default uses CFG and the user did not set a parallelism policy, `auto` may also enable CFG parallelism. `speed` intentionally does not check memory; it is the mode for users who prefer latency/throughput and accept OOM risk. It keeps `torch.compile` disabled by default because its effect varies by model and workload. A model-specific deployment config may enable a validated compile path, and `--enable-torch-compile true` always opts in explicitly. +`auto` checks selected GPU memory before applying FSDP. In multi-GPU runs it uses the least available memory across selected GPUs, and only turns on FSDP automatically when doing so can replace DiT offload. For image workloads with at least 45 GiB available per selected GPU, it keeps the repeatedly reused DiT resident and uses layerwise offload for large auxiliary encoders; below that threshold it keeps the DiT offloaded. Video DiT residency remains model- and workload-specific because frame count and resolution change its peak memory substantially. When the model default uses CFG and the user did not set a parallelism policy, `auto` may also enable CFG parallelism. `speed` intentionally does not check memory; it is the mode for users who prefer latency/throughput and accept OOM risk. It keeps `torch.compile` disabled by default because its effect varies by model and workload. A model-specific deployment config may enable a validated compile path, and `--enable-torch-compile true` always opts in explicitly. The modes tune residency for native pipeline components declared to the component residency manager. Today this covers the major DiT, text/image encoder, VAE, vocoder, and upsampler components; DiT can use layerwise offload when supported, while text encoders use either resident execution or component CPU offload. Do not assume text-encoder layerwise offload unless a model implements and validates it. diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/hunyuan.py b/python/sglang/multimodal_gen/configs/pipeline_configs/hunyuan.py index 9fd4dcbd4..edeaaf09c 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/hunyuan.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/hunyuan.py @@ -164,5 +164,6 @@ class FastHunyuanConfig(HunyuanConfig): def get_model_deployment_config(self) -> ModelDeploymentConfig: return ModelDeploymentConfig( - keep_resident_components=("vae",), + keep_resident_min_available_gb=60, + keep_resident_components=("dit", "vae"), ) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py b/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py index 2585be683..f150dd6a3 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py @@ -338,6 +338,13 @@ class LingBotWorldCausalDMDConfig(LingBotWorldI2VConfig): interactive_kv_still_chunks: int = 2 lazy_vae_encode_black_frames: int = 0 + def get_model_deployment_config(self) -> ModelDeploymentConfig: + return ModelDeploymentConfig( + dit_layerwise_offload_modes=("memory",), + keep_resident_min_available_gb=70, + keep_resident_components=("dit",), + ) + def preprocess_vae_encode(self, image, vae): image = super().preprocess_vae_encode(image, vae) lazy_black_frames = envs.SGLANG_LINGBOT_LAZY_VAE_ENCODE_BLACK_FRAMES diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/model_deployment_config.py b/python/sglang/multimodal_gen/configs/pipeline_configs/model_deployment_config.py index ca37a4765..255a49af6 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/model_deployment_config.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/model_deployment_config.py @@ -14,9 +14,9 @@ class ModelDeploymentConfig: dit_layerwise_offload_modes: tuple[Literal["auto", "memory"], ...] = () auto_dit_offload_prefetch_size: float | None = None keep_resident_min_available_gb: float | None = None - # only vae -- it is tiny so keeping it resident barely shifts memory; large - # encoders stay offloaded and dit placement stays with the FSDP/dit-layerwise - # policy + # Per-model resident defaults. Auto mode additionally keeps an image DiT + # resident above the image workload memory threshold; video DiT placement + # stays with the model's FSDP/layerwise policy. keep_resident_components: tuple[OffloadComponentName, ...] = ("vae",) fsdp_auto_min_available_memory_gb: float | None = None fsdp_auto_requires_cfg: bool = True diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py index f3fc5d087..ed37e062c 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py @@ -99,6 +99,8 @@ class WanT2V480PConfig(PipelineConfig): def get_model_deployment_config(self) -> ModelDeploymentConfig: return ModelDeploymentConfig( dit_layerwise_offload_modes=("memory",), + keep_resident_min_available_gb=60, + keep_resident_components=("dit",), ) def expand_conditioning_to_sample_batch(self, batch): @@ -141,6 +143,7 @@ class TurboWanT2V1_3B480PConfig(TurboWanT2V480PConfig): dit_layerwise_offload_modes=("memory",), keep_resident_min_available_gb=60, keep_resident_components=( + "dit", "text_encoder", "image_encoder", "vae", @@ -182,11 +185,6 @@ class WanI2V480PConfig(WanT2V480PConfig, WanI2VCommonConfig): self.vae_config.load_encoder = True self.vae_config.load_decoder = True - def get_model_deployment_config(self) -> ModelDeploymentConfig: - return ModelDeploymentConfig( - dit_layerwise_offload_modes=("memory",), - ) - @dataclass class WanI2V720PConfig(WanI2V480PConfig): @@ -228,6 +226,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig): dit_layerwise_offload_modes=("memory",), keep_resident_min_available_gb=60, keep_resident_components=( + "dit", "text_encoder", "image_encoder", "vae", diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py index 2f8e061c0..ecee1a99e 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py @@ -82,7 +82,10 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): F_PATCH_SIZE: int = 1 def get_model_deployment_config(self) -> ModelDeploymentConfig: - return ModelDeploymentConfig(fsdp_auto_min_available_memory_gb=40) + return ModelDeploymentConfig( + keep_resident_min_available_gb=30, + fsdp_auto_min_available_memory_gb=40, + ) def prepare_sigmas(self, sigmas, num_inference_steps): return self._prepare_sigmas(sigmas, num_inference_steps) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py index 294e7b44d..af18bfc23 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py @@ -3,7 +3,6 @@ from safetensors.torch import load_file as safetensors_load_file from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import ( LTX2ConnectorConfig, ) -from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, ) @@ -50,7 +49,9 @@ class AdapterLoader(ComponentLoader): model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) - target_device = get_local_torch_device() + target_device = self.target_device( + server_args.should_cpu_offload_component("connectors") + ) default_dtype = resolve_precision( server_args, "connectors", precision_attr="dit_precision" ) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py index 2ff2dc8c4..57dfee124 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py @@ -74,6 +74,8 @@ class BridgeLoader(ComponentLoader): default_dtype, ) + component_cpu_offload = server_args.should_cpu_offload_component(component_name) + # Use the FSDP loader when FSDP is requested or shard rules are declared. fsdp_shard_conditions = getattr(model_cls, "_fsdp_shard_conditions", None) if server_args.use_fsdp_inference or ( @@ -88,7 +90,7 @@ class BridgeLoader(ComponentLoader): device=local_torch_device, hsdp_replicate_dim=server_args.hsdp_replicate_dim, hsdp_shard_dim=server_args.hsdp_shard_dim, - cpu_offload=server_args.dit_cpu_offload, + cpu_offload=component_cpu_offload, pin_cpu_memory=server_args.pin_cpu_memory, fsdp_inference=server_args.use_fsdp_inference, param_dtype=default_dtype, @@ -104,7 +106,8 @@ class BridgeLoader(ComponentLoader): model = model_cls.from_pretrained( component_model_path, torch_dtype=default_dtype ) - model = model.to(device=get_local_torch_device(), dtype=default_dtype) + target_device = self.target_device(component_cpu_offload) + model = model.to(device=target_device, dtype=default_dtype) total_params = sum(p.numel() for p in model.parameters()) logger.info("Loaded bridge model with %.2fM parameters", total_params / 1e6) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index 64cb290da..b1dbf3766 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -25,6 +25,9 @@ from sglang.multimodal_gen.runtime.loader.utils import ( component_name_to_loader_cls, get_memory_usage_of_component, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_strategies import ( + is_fsdp_managed_module, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( configure_layerwise_offload_modules, is_layerwise_offloaded_module, @@ -92,10 +95,14 @@ class ComponentLoader(ABC): self.component_architecture: str | None = None def should_offload( - self, server_args: ServerArgs, model_config: ModelConfig | None = None + self, + server_args: ServerArgs, + model_config: ModelConfig | None = None, + component_name: str | None = None, ): - # not offload by default - return False + return component_name is not None and server_args.should_cpu_offload_component( + component_name + ) def target_device(self, should_offload): if should_offload: @@ -215,7 +222,7 @@ class ComponentLoader(ABC): transformers_or_diffusers, component_name, ) - should_offload = self.should_offload(server_args) + should_offload = self.should_offload(server_args, component_name=component_name) target_device = self.target_device(should_offload) return component.to(device=target_device) @@ -301,6 +308,12 @@ class ComponentLoader(ABC): else: if isinstance(component, nn.Module): component = component.eval() + if ( + server_args.cpu_offload_components is not None + and server_args.should_cpu_offload_component(component_name) + and not is_fsdp_managed_module(component) + ): + component = component.to("cpu") current_gpu_mem = current_platform.get_available_gpu_memory() model_size = get_memory_usage_of_component(component) or "NA" consumed = gpu_mem_before_loading - current_gpu_mem diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py index 5818f6b37..dbe1367c8 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py @@ -16,8 +16,14 @@ class ImageEncoderLoader(TextEncoderLoader): component_names = ["image_encoder"] expected_library = "transformers" - def should_offload(self, server_args, model_config: ModelConfig | None = None): - should_offload = server_args.image_encoder_cpu_offload + def should_offload( + self, + server_args, + model_config: ModelConfig | None = None, + component_name: str | None = None, + ): + component_name = component_name or "image_encoder" + should_offload = server_args.should_cpu_offload_component(component_name) if not should_offload: return False # _fsdp_shard_conditions is in arch_config, not directly on model_config @@ -66,6 +72,7 @@ class ImageEncoderLoader(TextEncoderLoader): cpu_offload_flag=( cpu_offload_flag if cpu_offload_flag is not None - else server_args.image_encoder_cpu_offload + else server_args.should_cpu_offload_component(component_name) ), + component_name=component_name, ) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py index 8db15950d..32d42f108 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 from safetensors.torch import load_file as safetensors_load_file -from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, ) @@ -25,11 +24,6 @@ class SoundTokenizerLoader(ComponentLoader): component_names = ["sound_tokenizer"] expected_library = "diffusers" - def should_offload( - self, server_args: ServerArgs, model_config: ModelConfig | None = None - ) -> bool: - return server_args.vae_cpu_offload - def load_customized( self, component_model_path: str, server_args: ServerArgs, component_name: str ): @@ -46,7 +40,9 @@ class SoundTokenizerLoader(ComponentLoader): except AttributeError: precision = "bf16" dtype = PRECISION_TO_TYPE[precision] - target_device = self.target_device(self.should_offload(server_args)) + target_device = self.target_device( + server_args.should_cpu_offload_component(component_name) + ) with set_default_torch_dtype(dtype), skip_init_modules(): model_cls, _ = ModelRegistry.resolve_model_cls(class_name) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 8c9e85e81..0c8861a9f 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -80,8 +80,14 @@ class TextEncoderLoader(ComponentLoader): allow_patterns_overrides: list[str] | None = None """If defined, weights will load exclusively using these patterns.""" - def should_offload(self, server_args, model_config: ModelConfig | None = None): - should_offload = server_args.text_encoder_cpu_offload + def should_offload( + self, + server_args, + model_config: ModelConfig | None = None, + component_name: str | None = None, + ): + component_name = component_name or "text_encoder" + should_offload = server_args.should_cpu_offload_component(component_name) if not should_offload: return False # _fsdp_shard_conditions is in arch_config, not directly on model_config @@ -369,6 +375,7 @@ class TextEncoderLoader(ComponentLoader): server_args, encoder_dtype, cpu_offload_flag=cpu_offload_flag, + component_name=component_name, ) @staticmethod @@ -400,13 +407,16 @@ class TextEncoderLoader(ComponentLoader): server_args: ServerArgs, dtype: str = "fp16", cpu_offload_flag: bool | None = None, + component_name: str = "text_encoder", ): # Determine CPU offload behavior and target device local_torch_device = get_local_torch_device() if not current_platform.is_cpu(): - fsdp_cpu_offload = self.should_offload(server_args, model_config) + fsdp_cpu_offload = self.should_offload( + server_args, model_config, component_name + ) should_offload = ( cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload ) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 3e21e8c2a..ba2216eb2 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -152,6 +152,11 @@ class TransformerLoader(ComponentLoader): component_server_args = _server_args_for_transformer_component( server_args, component_name ) + if server_args.cpu_offload_components is not None: + component_server_args = copy.copy(component_server_args) + component_server_args.dit_cpu_offload = ( + server_args.should_cpu_offload_component(component_name) + ) # 1. hf config config = get_diffusers_component_config(component_path=component_model_path) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py index 6ad78e066..1fa6751e3 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py @@ -195,9 +195,6 @@ class UpsamplerLoader(ComponentLoader): component_names = ["spatial_upsampler"] expected_library = "diffusers" - def should_offload(self, server_args: ServerArgs, model_config=None): - return server_args.vae_cpu_offload - def load_customized( self, component_model_path: str, @@ -210,7 +207,7 @@ class UpsamplerLoader(ComponentLoader): logger.info("Loading LatentUpsampler with config: %s", config) - should_offload = self.should_offload(server_args) + should_offload = server_args.should_cpu_offload_component(component_name) target_device = self.target_device(should_offload) with torch.device("meta"): diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py index e59c87477..32fe0c39a 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py @@ -5,7 +5,6 @@ import torch import torch.nn as nn from safetensors.torch import load_file as safetensors_load_file -from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImagePipelineConfig, @@ -99,11 +98,6 @@ class VAELoader(ComponentLoader): component_names = ["vae", "audio_vae", "video_vae"] expected_library = "diffusers" - def should_offload( - self, server_args: ServerArgs, model_config: ModelConfig | None = None - ): - return server_args.vae_cpu_offload - def load_customized( self, component_model_path: str, server_args: ServerArgs, component_name: str ): @@ -139,7 +133,7 @@ class VAELoader(ComponentLoader): # NOTE: some post init logics are only available after updated with config vae_config.post_init() - should_offload = self.should_offload(server_args) + should_offload = server_args.should_cpu_offload_component(component_name) target_device = self.target_device(should_offload) native_only = component_name in getattr( diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py index d8b9891a2..970d19c83 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py @@ -3,7 +3,6 @@ from typing import Any import requests -from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, ) @@ -59,12 +58,15 @@ class VisionLanguageEncoderLoader(ComponentLoader): trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, ) + target_device = self.target_device( + server_args.should_cpu_offload_component("vision_language_encoder") + ) model = GlmImageForConditionalGeneration.from_pretrained( component_model_path, config=config, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, - ).to(get_local_torch_device()) + ).to(target_device) return model else: raise ValueError( diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py index 7e7ef4874..80d18675b 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py @@ -1,6 +1,5 @@ from safetensors.torch import load_file as safetensors_load_file -from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, ) @@ -25,11 +24,6 @@ class VocoderLoader(ComponentLoader): component_names = ["vocoder"] expected_library = "diffusers" - def should_offload( - self, server_args: ServerArgs, model_config: ModelConfig | None = None - ): - return server_args.vae_cpu_offload - def load_customized( self, component_model_path: str, server_args: ServerArgs, component_name: str ): @@ -55,7 +49,7 @@ class VocoderLoader(ComponentLoader): else PRECISION_TO_TYPE["fp32"] ) - should_offload = self.should_offload(server_args) + should_offload = server_args.should_cpu_offload_component(component_name) target_device = self.target_device(should_offload) with set_default_torch_dtype(vocoder_dtype), skip_init_modules(): diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 8b665082f..1063d6a83 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -317,46 +317,24 @@ class GPUWorker(GPUWorkerPostTrainingMixin): 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_str = self._format_offload_disable_suggestions(can_stay_resident) pool_overhead_gb = peak_reserved_gb - peak_allocated_gb - - logger.debug( - f"Peak GPU memory: {peak_reserved_gb:.2f} GB, " - f"Peak allocated: {peak_allocated_gb:.2f} GB, " - f"Memory pool overhead: {pool_overhead_gb:.2f} GB ({pool_overhead_gb / peak_reserved_gb * 100:.1f}%), " - f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. " - f"Components that could stay resident (based on the last request workload): {can_stay_resident}. " - f"Related offload server args to disable: {suggested_args_str}" + pool_overhead_pct = ( + pool_overhead_gb / peak_reserved_gb * 100 if peak_reserved_gb else 0.0 ) - 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.is_dit_layerwise_offload_selected: - 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" + logger.debug( + "GPU memory: peak=%.2f GB, allocated=%.2f GB, pool=%.2f GB (%.1f%%), " + "headroom=%.2f GB. Components that can remain on GPU: %s. " + "Adjust --cpu-offload-components or --layerwise-offload-components " + "to change residency.", + peak_reserved_gb, + peak_allocated_gb, + pool_overhead_gb, + pool_overhead_pct, + remaining_gpu_mem_gb, + can_stay_resident, + ) def execute_forward( self, batch: List[Req], return_req: bool = False @@ -974,28 +952,28 @@ class GPUWorker(GPUWorkerPostTrainingMixin): 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 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.is_dit_layerwise_offload_selected, - "vae": self.server_args.vae_cpu_offload, - "text_encoder": self.server_args.text_encoder_cpu_offload, - "text_encoder_2": self.server_args.text_encoder_cpu_offload, - "image_encoder": self.server_args.image_encoder_cpu_offload, - } - - 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: + memory_usages = self.pipeline.memory_usages + ordered_names = [ + name + for name in OFFLOAD_DISABLE_RECOMMENDATION_ORDER + if name in memory_usages + ] + ordered_names.extend( + name + for name in memory_usages + if name not in OFFLOAD_DISABLE_RECOMMENDATION_ORDER + ) + for name in ordered_names: + usage = memory_usages[name] + if not ( + self.server_args.should_cpu_offload_component(name) + or self.server_args.should_configure_layerwise_offload_for_lazy_component( + name + ) + ): 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/managers/memory_managers/component_manager.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py index 9d23a5b08..15c7ca14e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py @@ -18,12 +18,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im is_layerwise_offloaded_module, is_resident_layerwise_module, ) -from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( - is_dit_component_name, - is_image_encoder_component_name, - is_text_encoder_component_name, - is_vae_component_name, -) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -93,24 +87,6 @@ class ComponentResidencyPipeline(Protocol): component_residency_strategies: MutableMapping[str, "ComponentResidencyStrategy"] -def should_cpu_offload_component( - component_name: str, module: nn.Module, server_args: ServerArgs -) -> bool: - if current_platform.is_mps(): - return False - if server_args.use_fsdp_inference or is_fsdp_managed_module(module): - return False - if is_dit_component_name(component_name): - return bool(server_args.dit_cpu_offload) - if is_text_encoder_component_name(component_name): - return bool(server_args.text_encoder_cpu_offload) - if is_image_encoder_component_name(component_name): - return bool(server_args.image_encoder_cpu_offload) - if is_vae_component_name(component_name): - return bool(server_args.vae_cpu_offload) - return False - - def build_component_residency_strategy( component_name: str, module: nn.Module, @@ -118,7 +94,12 @@ def build_component_residency_strategy( ) -> ComponentResidencyStrategy: if is_layerwise_offloaded_module(module): return LayerwiseOffloadStrategy() - if should_cpu_offload_component(component_name, module, server_args): + if ( + not current_platform.is_mps() + and not server_args.use_fsdp_inference + and not is_fsdp_managed_module(module) + and server_args.should_cpu_offload_component(component_name) + ): return VanillaD2HStrategy() return ResidentStrategy() @@ -164,7 +145,6 @@ class ComponentResidencyManager: if pipeline is not self.pipeline: self._remove_nvtx_hooks() self.strategy_for.cache_clear() - self._should_keep_single_dit.cache_clear() self._active_use = None self._active_use_module = None self._uses_seen.clear() @@ -466,8 +446,15 @@ class ComponentResidencyManager: preferred = component_name in preferred_uses if is_resident_layerwise_module(module): preferred = False - elif not preferred and self._should_keep_single_dit(component_name): + keep_single_dit = self._should_keep_single_dit(component_name, module) + if not preferred and keep_single_dit: continue + # A preferred component is normally prefetched for the next request. + # Do not let that performance hint override CPU/layerwise offload for + # a single DiT, which must obey the selected memory policy. + preferred = preferred and ( + not self._is_single_dit_component(component_name) or keep_single_dit + ) strategy = self.strategy_for(component_name, module) if preferred and not self.state.batch_is_warmup: strategy.prepare_after_request(module, use, self.state) @@ -544,16 +531,25 @@ class ComponentResidencyManager: } if use.component_name in future_component_names: return True - if self._should_keep_single_dit(use.component_name): - module = self.get_module(use.component_name) - if module is not None and is_resident_layerwise_module(module): - # don't keep a layerwise DiT resident across the request to avoid OOMs - return False + module = self.get_module(use.component_name) + if module is not None and self._should_keep_single_dit( + use.component_name, module + ): return True return False - @lru_cache(maxsize=None) - def _should_keep_single_dit(self, component_name: str) -> bool: + def _should_keep_single_dit(self, component_name: str, module: nn.Module) -> bool: + """Keep a single DiT resident only when its effective strategy is resident. + + The single-DiT fast path is a performance optimization, not a memory + policy. In particular, it must not override explicit or auto-selected + CPU/layerwise offload. + """ + if not self._is_single_dit_component(component_name): + return False + return isinstance(self.strategy_for(component_name, module), ResidentStrategy) + + def _is_single_dit_component(self, component_name: str) -> bool: modules = self.pipeline.modules return (component_name == "transformer" and "transformer_2" not in modules) or ( component_name == "video_dit" and "video_dit_2" not in modules diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py index 71e3bf54e..a2ac0b6fc 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py @@ -47,12 +47,62 @@ CPU_OFFLOAD_FLAG_NAMES = ( "image_encoder_cpu_offload", "vae_cpu_offload", ) +CPU_OFFLOAD_ALL_COMPONENTS = "all" def is_dit_component_name(component_name: str) -> bool: return component_name in DIT_COMPONENT_NAMES +def normalize_cpu_offload_components( + component_names: str | Sequence[str] | None, +) -> list[str] | None: + """Normalize component keys accepted by ``--cpu-offload-components``.""" + if component_names is None: + return None + + raw_components = ( + [component_names] if isinstance(component_names, str) else component_names + ) + normalized_components: list[str] = [] + for raw_component in raw_components: + if not isinstance(raw_component, str): + raise ValueError(f"Invalid CPU offload component name: {raw_component}.") + normalized_components.extend( + component_name + for value in raw_component.split(",") + if (component_name := value.strip().replace("-", "_").lower()) + ) + + unique_components = list(dict.fromkeys(normalized_components)) + if "none" in unique_components: + if len(unique_components) != 1: + raise ValueError("'none' cannot be combined with other components.") + return [] + return unique_components or None + + +def cpu_offload_component_matches( + component_name: str, + selected_component_names: Collection[str] | None, +) -> bool: + if selected_component_names is None: + return False + if CPU_OFFLOAD_ALL_COMPONENTS in selected_component_names: + return True + if component_name in selected_component_names: + return True + if LAYERWISE_OFFLOAD_DIT_GROUP in selected_component_names: + return is_dit_component_name(component_name) + if LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP in selected_component_names: + return is_text_encoder_component_name(component_name) + if LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP in selected_component_names: + return component_name in ("image_encoder", "condition_image_encoder") + if LAYERWISE_OFFLOAD_VAE_GROUP in selected_component_names: + return is_vae_component_name(component_name) + return False + + def is_text_encoder_component_name(component_name: str) -> bool: return component_name.startswith("text_encoder") or component_name.endswith( "text_encoder" diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py index 44408e550..5a225e731 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py @@ -654,6 +654,8 @@ class CausalDMDDenoisingStage(DenoisingStage): target_dtype: torch.dtype, autocast_enabled: bool, ) -> torch.Tensor: + if self._component_residency_manager is not None: + self._manage_dit_use_site(self.transformer, "transformer", batch) with ( precision_autocast_context( target_dtype, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 891a9c533..04711febc 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -1531,12 +1531,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): batch: Req, ) -> None: """ - manage dit's residency by reporting the active sequential use - - only applicable for dual-dit architecture like Wan + manage dit residency by reporting the active sequential use Args: - current_model: the next active dit, transformer_1 or transformer_2 + current_model: the next active dit """ manager = self._component_residency_manager diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py index 9f06b4b7e..3716d5320 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py @@ -534,7 +534,6 @@ class LongLive2CausalDenoisingStage(CausalDMDDenoisingStage): target_dtype: torch.dtype, autocast_enabled: bool, ) -> torch.Tensor: - self._manage_dit_use_site(self.transformer, "transformer", batch) rope_start_frame = start_frame if self._rope_temporal_offset != 0.0: rope_start_frame = start_frame + self._rope_temporal_offset diff --git a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py index 0e6a993d3..62e2a82d1 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py @@ -140,7 +140,9 @@ class ServerArgsAutoTuner: and min_available_gb >= disable_threshold_gb ): changed = [] - components = deployment_config.keep_resident_components + components = set(deployment_config.keep_resident_components) + if args.pipeline_config.task_type.is_image_gen(): + components.add(LAYERWISE_OFFLOAD_DIT_GROUP) if ( args.layerwise_offload_components is not None and not args.is_arg_explicitly_set("layerwise_offload_components") @@ -253,6 +255,7 @@ class ServerArgsAutoTuner: if ( args.layerwise_offload_components is not None or args.dit_layerwise_offload is True + or args.is_arg_explicitly_set("cpu_offload_components") ): return if not current_platform.is_cuda(): @@ -410,6 +413,7 @@ class ServerArgsAutoTuner: if ( args.is_arg_explicitly_set("layerwise_offload_components") or args.dit_layerwise_offload is True + or args.is_arg_explicitly_set("cpu_offload_components") ): # The legacy --dit-layerwise-offload flag is a DiT-only selector. # Do not merge implicit defaults into that explicit mode. @@ -475,6 +479,7 @@ class ServerArgsAutoTuner: or envs.SGLANG_CACHE_DIT_ENABLED or args.use_fsdp_inference or args.is_arg_explicitly_set("dit_cpu_offload") + or args.is_arg_explicitly_set("cpu_offload_components") ): return False @@ -523,6 +528,7 @@ class ServerArgsAutoTuner: "dit_cpu_offload", "dit_layerwise_offload", "layerwise_offload_components", + "cpu_offload_components", ) ) @@ -533,6 +539,7 @@ class ServerArgsAutoTuner: for arg_name in ( "dit_layerwise_offload", "layerwise_offload_components", + "cpu_offload_components", ) ) diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index d1091c99a..93e5a79c5 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -35,8 +35,14 @@ from sglang.multimodal_gen.runtime.loader.utils import BYTES_PER_GB from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( LAYERWISE_OFFLOAD_ALL_COMPONENTS, LAYERWISE_OFFLOAD_DIT_GROUP, + cpu_offload_component_matches, cpu_offload_flags_for_layerwise_components, + is_dit_component_name, + is_image_encoder_component_name, + is_text_encoder_component_name, + is_vae_component_name, layerwise_component_matches_any_selection, + normalize_cpu_offload_components, normalize_layerwise_offload_components, ) from sglang.multimodal_gen.runtime.platforms import ( @@ -301,6 +307,8 @@ class ServerArgs(DisaggServerArgsMixin): lora_target_modules: list[str] | None = None # CPU offload parameters + # Exact component keys from model_index.json, or a legacy component group. + cpu_offload_components: list[str] | None = None dit_cpu_offload: bool | None = None # trade checkpoint-loading peak memory for faster ordinary DiT startup direct_gpu_weight_loading: bool = False @@ -486,13 +494,14 @@ class ServerArgs(DisaggServerArgsMixin): """set defaults and normalize values.""" auto_tuner = ServerArgsAutoTuner(self) auto_tuner.adjust_based_on_performance_mode() + self._adjust_cpu_offload_components() if auto_tuner.could_override_server_args(): self._adjust_offload() auto_tuner.maybe_adjust_auto_default_layerwise_offload() self._adjust_ltx2_two_stage_device_mode() if auto_tuner.could_override_server_args(): - auto_tuner.maybe_adjust_auto_component_residency_after_offload() auto_tuner.maybe_adjust_auto_fsdp_with_offload_enabled() + auto_tuner.maybe_adjust_auto_component_residency_after_offload() auto_tuner.maybe_replace_cpu_offloaded_components_with_layerwise() self._adjust_path() if self.served_model_name is None: @@ -739,6 +748,44 @@ class ServerArgs(DisaggServerArgsMixin): if self.image_encoder_cpu_offload is None: self.image_encoder_cpu_offload = True + def _adjust_cpu_offload_components(self) -> None: + """Apply the unified CPU offload component selector, when provided.""" + if self.cpu_offload_components is None: + return + + legacy_flags = ( + "dit_cpu_offload", + "text_encoder_cpu_offload", + "image_encoder_cpu_offload", + "vae_cpu_offload", + ) + conflicting_flags = [ + flag_name + for flag_name in legacy_flags + if self.is_arg_explicitly_set(flag_name) + ] + if conflicting_flags: + formatted_flags = ", ".join( + "--" + flag_name.replace("_", "-") for flag_name in conflicting_flags + ) + raise ValueError( + "--cpu-offload-components cannot be combined with the legacy " + f"CPU offload flags: {formatted_flags}" + ) + + selected_components = ( + normalize_cpu_offload_components(self.cpu_offload_components) or [] + ) + self.cpu_offload_components = selected_components + self.dit_cpu_offload = self.should_cpu_offload_component("transformer") + self.text_encoder_cpu_offload = self.should_cpu_offload_component( + "text_encoder" + ) + self.image_encoder_cpu_offload = self.should_cpu_offload_component( + "image_encoder" + ) + self.vae_cpu_offload = self.should_cpu_offload_component("vae") + def _adjust_ltx2_two_stage_device_mode(self): if not self._is_ltx23_two_stage_pipeline(): return @@ -1240,6 +1287,25 @@ class ServerArgs(DisaggServerArgsMixin): def is_arg_explicitly_set(self, arg_name: str) -> bool: return arg_name in self._explicit_arg_names + def should_cpu_offload_component(self, component_name: str) -> bool: + if self.cpu_offload_components is not None: + return cpu_offload_component_matches( + component_name, self.cpu_offload_components + ) + if is_dit_component_name(component_name) or component_name in ( + "connectors", + "unconditional_transformer", + "vision_language_encoder", + ): + return bool(self.dit_cpu_offload) + if is_text_encoder_component_name(component_name): + return bool(self.text_encoder_cpu_offload) + if is_image_encoder_component_name(component_name): + return bool(self.image_encoder_cpu_offload) + if is_vae_component_name(component_name) or component_name == "sound_tokenizer": + return bool(self.vae_cpu_offload) + return False + def should_configure_layerwise_offload_for_lazy_component( self, component_name: str ) -> bool: @@ -1797,6 +1863,19 @@ class ServerArgs(DisaggServerArgsMixin): "time depending on the model, but temporarily requires checkpoint " "weights and model weights to coexist on GPU. Disabled by default.", ) + parser.add_argument( + "--cpu-offload-components", + type=str, + nargs="+", + default=ServerArgs.cpu_offload_components, + help=( + "Select component keys from model_index.json for coarse CPU offload. " + "Use dit, text_encoder, image_encoder, or vae as group aliases; " + "all selects every loaded module and none disables component offload. " + "This unified option cannot be combined with the legacy " + "per-component CPU offload flags." + ), + ) parser.add_argument( "--dit-layerwise-offload", action=StoreBoolean, diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index da5d86cfa..204c984ad 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -39,7 +39,7 @@ logger = init_logger(__name__) # NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows # publish. SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion" -SGL_TEST_FILES_CI_DATA_REVISION = "eccf85dcebaaded92df8b0fce3064ebea910c6d4" +SGL_TEST_FILES_CI_DATA_REVISION = "cc3f27fd2d1b4d8e1a7d5eec1247a215a502b9c1" # The NPU pin is kept as a separate branch so ascend GT can be bumped independently # when it's regenerated on its own cadence. diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py b/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py index 59689ec18..ad22e4bfb 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py @@ -1,5 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 +from contextlib import nullcontext from types import MethodType, SimpleNamespace import torch @@ -21,6 +22,61 @@ class _Progress: self.count += 1 +def test_causal_transformer_prepares_dit_before_forward(monkeypatch): + from sglang.multimodal_gen.runtime.pipelines_core.stages import causal_denoising + + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage._component_residency_manager = object() + calls = [] + + def manage_dit(self, model, phase, batch): + del self + calls.append(("prepare", model, phase, batch)) + + def transformer(latents, *args, **kwargs): + del args, kwargs + calls.append(("forward",)) + return latents + + stage._manage_dit_use_site = MethodType(manage_dit, stage) + stage.transformer = transformer + monkeypatch.setattr( + causal_denoising, + "precision_autocast_context", + lambda *args, **kwargs: nullcontext(), + ) + monkeypatch.setattr( + causal_denoising, + "set_forward_context", + lambda **kwargs: nullcontext(), + ) + + batch = SimpleNamespace() + latents = torch.zeros(1) + result = stage._forward_causal_transformer( + batch, + latent_model_input=latents, + prompt_embeds=None, + timestep=torch.zeros(1), + kv_cache=[], + crossattn_cache=[], + current_start_tokens=0, + start_frame=0, + image_kwargs={}, + pos_cond_kwargs={}, + current_timestep=0, + attn_metadata=None, + target_dtype=torch.float16, + autocast_enabled=False, + ) + + assert result is latents + assert calls == [ + ("prepare", transformer, "transformer", batch), + ("forward",), + ] + + def test_causal_dmd_chunk_loop_uses_model_input_builder(): stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) predict_calls = [] diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py index 542e7a77f..49e0c363d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im is_layerwise_offloaded_module, is_resident_layerwise_module, ) +from sglang.multimodal_gen.runtime.server_args import ServerArgs class _FakeStream: @@ -161,8 +162,13 @@ class _LayerwiseComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin): self.layerwise_offload_managers = [SimpleNamespace(enabled=enabled)] +class _TestServerArgs(SimpleNamespace): + should_cpu_offload_component = ServerArgs.should_cpu_offload_component + + def _server_args(**kwargs): defaults = dict( + cpu_offload_components=None, use_fsdp_inference=False, dit_cpu_offload=False, text_encoder_cpu_offload=False, @@ -173,7 +179,7 @@ def _server_args(**kwargs): pin_cpu_memory=False, ) defaults.update(kwargs) - return SimpleNamespace(**defaults) + return _TestServerArgs(**defaults) def test_layerwise_offload_preserves_non_contiguous_stride(monkeypatch): diff --git a/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py b/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py index 5483ceef5..f3405afdb 100644 --- a/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py +++ b/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py @@ -16,6 +16,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im ComponentResidencyManager, ComponentUse, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_strategies import ( + ResidentStrategy, + VanillaD2HStrategy, +) from sglang.multimodal_gen.runtime.utils import nvtx_pytorch_hooks from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import ( DiffusionNvtxHooks, @@ -228,6 +232,16 @@ def _test_manager( class TestComponentResidencyNvtxHooks(unittest.TestCase): + def test_single_dit_residency_does_not_override_offload_strategy(self) -> None: + module = torch.nn.Linear(2, 2) + manager = _test_manager({"transformer": module}) + + manager.strategy_for = lambda _component_name, _module: ResidentStrategy() + self.assertTrue(manager._should_keep_single_dit("transformer", module)) + + manager.strategy_for = lambda _component_name, _module: VanillaD2HStrategy() + self.assertFalse(manager._should_keep_single_dit("transformer", module)) + def test_disabled_flag_is_noop(self) -> None: module = torch.nn.Linear(2, 2) manager = _test_manager({"linear": module}, enable_flag=False) 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 06ba57208..1b61d2562 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -17,6 +17,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import ( PipelineConfig, ) from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanConfig +from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( + LingBotWorldCausalDMDConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( LTX2PipelineConfig, LTX23PipelineConfig, @@ -385,6 +388,32 @@ class TestServerArgsPathExpansion(unittest.TestCase): server_args.layerwise_offload_components, ["transformer", "text_encoder"] ) + def test_cpu_offload_components_cli_args(self): + parser = FlexibleArgumentParser() + ServerArgs.add_cli_args(parser) + argv = [ + "--model-path", + "/fake", + "--performance-mode", + "manual", + "--cpu-offload-components", + "transformer", + "vae", + ] + + with patch.object(sys, "argv", ["sglang"] + argv): + args, unknown_args = parser.parse_known_args(argv) + with patch.object( + PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig() + ): + server_args = ServerArgs.from_cli_args(args, unknown_args) + + self.assertEqual(server_args.cpu_offload_components, ["transformer", "vae"]) + self.assertTrue(server_args.dit_cpu_offload) + self.assertTrue(server_args.vae_cpu_offload) + self.assertFalse(server_args.text_encoder_cpu_offload) + self.assertFalse(server_args.image_encoder_cpu_offload) + def test_serve_cli_preserves_config_and_dynamic_unknown_args(self): from sglang.multimodal_gen.runtime.entrypoints.cli.serve import ( add_multimodal_gen_serve_args, @@ -814,6 +843,74 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.vae_cpu_offload) + def test_cpu_offload_components_preserves_model_index_names(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "performance_mode": "manual", + "cpu_offload_components": [ + "transformer_2", + "audio_vae", + "connectors", + ], + }, + ) + + self.assertEqual( + args.cpu_offload_components, + ["transformer_2", "audio_vae", "connectors"], + ) + self.assertTrue(args.should_cpu_offload_component("transformer_2")) + self.assertTrue(args.should_cpu_offload_component("audio_vae")) + self.assertTrue(args.should_cpu_offload_component("connectors")) + self.assertFalse(args.dit_cpu_offload) + self.assertFalse(args.vae_cpu_offload) + + def test_cpu_offload_components_all_matches_dynamic_components(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "performance_mode": "manual", + "cpu_offload_components": ["all"], + }, + ) + + self.assertEqual(args.cpu_offload_components, ["all"]) + self.assertTrue(args.should_cpu_offload_component("transformer_2")) + self.assertTrue(args.should_cpu_offload_component("connectors")) + + def test_cpu_offload_components_none_disables_all_legacy_flags(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "performance_mode": "manual", + "cpu_offload_components": ["none"], + }, + ) + + self.assertEqual(args.cpu_offload_components, []) + self.assertFalse(args.dit_cpu_offload) + self.assertFalse(args.text_encoder_cpu_offload) + self.assertFalse(args.image_encoder_cpu_offload) + self.assertFalse(args.vae_cpu_offload) + + with self.assertRaisesRegex(ValueError, "cannot be combined"): + self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={"cpu_offload_components": ["none", "vae"]}, + ) + + def test_cpu_offload_components_rejects_legacy_flag_conflict(self): + with self.assertRaisesRegex(ValueError, "cannot be combined"): + self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "performance_mode": "manual", + "cpu_offload_components": ["dit"], + "dit_cpu_offload": True, + }, + ) + def test_vae_cpu_offload_defaults_false_on_low_memory_gpu(self): args = self._from_dict_with_task_type( ModelTaskType.T2V, @@ -906,6 +1003,7 @@ class TestOffloadDefaults(unittest.TestCase): wan_deployment = WanT2V480PConfig().get_model_deployment_config() mova_deployment = MOVAPipelineConfig().get_model_deployment_config() zimage_deployment = ZImagePipelineConfig().get_model_deployment_config() + lingbot_deployment = LingBotWorldCausalDMDConfig().get_model_deployment_config() ltx_deployment = LTX2PipelineConfig().get_model_deployment_config() ltx23_config = LTX23PipelineConfig() sana_wm_deployment = SanaWMPipelineConfig().get_model_deployment_config() @@ -915,6 +1013,8 @@ class TestOffloadDefaults(unittest.TestCase): self.assertIsNone(wan_deployment.fsdp_auto_min_available_memory_gb) self.assertEqual(wan_deployment.dit_layerwise_offload_modes, ("memory",)) + self.assertEqual(wan_deployment.keep_resident_min_available_gb, 60) + self.assertEqual(wan_deployment.keep_resident_components, ("dit",)) self.assertIsNone(mova_deployment.fsdp_auto_min_available_memory_gb) self.assertEqual( @@ -924,9 +1024,14 @@ class TestOffloadDefaults(unittest.TestCase): self.assertEqual(mova_deployment.keep_resident_components, ("dit", "vae")) self.assertEqual(zimage_deployment.fsdp_auto_min_available_memory_gb, 40) + self.assertEqual(zimage_deployment.keep_resident_min_available_gb, 30) self.assertTrue(zimage_deployment.fsdp_auto_requires_cfg) self.assertEqual(zimage_deployment.dit_layerwise_offload_modes, ()) + self.assertEqual(lingbot_deployment.dit_layerwise_offload_modes, ("memory",)) + self.assertEqual(lingbot_deployment.keep_resident_min_available_gb, 70) + self.assertEqual(lingbot_deployment.keep_resident_components, ("dit",)) + self.assertEqual(ltx_deployment.keep_resident_min_available_gb, 70) self.assertEqual(ltx_deployment.keep_resident_components, ("dit",)) self.assertEqual( @@ -945,10 +1050,23 @@ class TestOffloadDefaults(unittest.TestCase): self.assertEqual(sana_wm_deployment.fsdp_auto_min_available_memory_gb, 60) self.assertEqual(sana_wm_deployment.dit_layerwise_offload_modes, ("memory",)) - # 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",)) + self.assertEqual(fast_hunyuan_deployment.keep_resident_min_available_gb, 60) + self.assertEqual( + fast_hunyuan_deployment.keep_resident_components, ("dit", "vae") + ) + + fast_wan_deployment = FastWan2_2_TI2V_5B_Config().get_model_deployment_config() + self.assertEqual(fast_wan_deployment.keep_resident_min_available_gb, 60) + self.assertEqual(fast_wan_deployment.keep_resident_components, ("dit",)) + + for dual_dit_config in ( + Wan2_2_T2V_A14B_Config(), + Wan2_2_I2V_A14B_Config(), + ): + dual_dit_deployment = dual_dit_config.get_model_deployment_config() + self.assertIsNone(dual_dit_deployment.keep_resident_min_available_gb) + self.assertEqual(dual_dit_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",)) @@ -1051,9 +1169,9 @@ class TestOffloadDefaults(unittest.TestCase): self.assertEqual(args.performance_mode, "auto") 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) + # 80gb > image threshold (45gb): vae and dit stay resident, while the + # large encoders use layerwise offload. + self.assertFalse(args.dit_cpu_offload) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder"], @@ -1075,6 +1193,53 @@ class TestOffloadDefaults(unittest.TestCase): ["text_encoder", "image_encoder", "vae"], ) + def test_auto_zimage_keeps_dit_resident_on_5090(self): + args = self._from_dict_with_pipeline_config( + ZImagePipelineConfig(), + memory_gb=32, + available_memory_gb=31, + kwargs={ + "model_path": "Tongyi-MAI/Z-Image-Turbo", + "performance_mode": "auto", + }, + ) + + self.assertFalse(args.dit_cpu_offload) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder"], + ) + + def test_auto_lingbot_keeps_dit_resident_on_h100(self): + args = self._from_dict_with_pipeline_config( + LingBotWorldCausalDMDConfig(), + memory_gb=80, + available_memory_gb=72, + kwargs={ + "model_path": "robbyant/lingbot-world-fast-diffusers", + "performance_mode": "auto", + "text_encoder_cpu_offload": True, + }, + ) + + self.assertFalse(args.dit_cpu_offload) + self.assertTrue(args.text_encoder_cpu_offload) + self.assertEqual( + args.layerwise_offload_components, + ["image_encoder", "vae"], + ) + + def test_auto_image_preserves_explicit_dit_cpu_offload(self): + args = self._from_dict_with_pipeline_config( + QwenImagePipelineConfig(), + kwargs={ + "model_path": "Qwen/Qwen-Image", + "dit_cpu_offload": True, + }, + ) + + self.assertTrue(args.dit_cpu_offload) + def test_auto_ltx_original_replaces_component_cpu_offload( self, ): @@ -1098,7 +1263,7 @@ class TestOffloadDefaults(unittest.TestCase): ["text_encoder", "image_encoder", "vae"], ) - def test_auto_wan_layerwise_offload_is_enabled_without_fsdp(self): + def test_auto_wan_keeps_single_dit_resident_on_h100(self): args = self._from_dict_with_pipeline_config( WanT2V480PConfig(), kwargs={"performance_mode": "auto"}, @@ -1106,7 +1271,7 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.use_fsdp_inference) - self.assertTrue(args.dit_cpu_offload) + self.assertFalse(args.dit_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( @@ -1114,6 +1279,19 @@ class TestOffloadDefaults(unittest.TestCase): ["text_encoder", "image_encoder", "vae"], ) + def test_auto_wan_offloads_single_dit_below_resident_threshold(self): + args = self._from_dict_with_pipeline_config( + WanT2V480PConfig(), + memory_gb=48, + kwargs={"performance_mode": "auto"}, + ) + + self.assertTrue(args.dit_cpu_offload) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) + def test_auto_wan2_2_a14b_layerwise_offload_adds_dit(self): for pipeline_config, model_path in ( (Wan2_2_T2V_A14B_Config(), "Wan-AI/Wan2.2-T2V-A14B-Diffusers"), @@ -1142,7 +1320,7 @@ class TestOffloadDefaults(unittest.TestCase): ["dit", "text_encoder", "image_encoder", "vae"], ) - def test_auto_wan2_1_14b_layerwise_offload_uses_non_dit_default(self): + def test_auto_wan2_1_14b_keeps_dit_resident_on_h100(self): for pipeline_config, model_path in ( (WanT2V720PConfig(), "Wan-AI/Wan2.1-T2V-14B-Diffusers"), (WanI2V480PConfig(), "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"), @@ -1158,13 +1336,25 @@ class TestOffloadDefaults(unittest.TestCase): ) self.assertTrue(args.layerwise_offload_components) - self.assertTrue(args.dit_cpu_offload) + self.assertFalse(args.dit_cpu_offload) self.assertEqual(args.dit_offload_prefetch_size, 0.0) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder", "vae"], ) + def test_auto_wan2_1_14b_offloads_dit_below_resident_threshold(self): + args = self._from_dict_with_pipeline_config( + WanI2V720PConfig(), + memory_gb=48, + kwargs={ + "model_path": "Wan-AI/Wan2.1-I2V-14B-720P-Diffusers", + "performance_mode": "auto", + }, + ) + + self.assertTrue(args.dit_cpu_offload) + def test_memory_wan_layerwise_offload_is_enabled_without_fsdp(self): args = self._from_dict_with_pipeline_config( WanT2V480PConfig(), @@ -1259,9 +1449,26 @@ class TestOffloadDefaults(unittest.TestCase): ["dit", "text_encoder", "image_encoder", "vae"], ) - def test_auto_fastwan_layerwise_offload_does_not_implicitly_add_dit(self): + def test_auto_fastwan_keeps_dit_resident_on_h100(self): args = self._from_dict_with_pipeline_config( FastWan2_2_TI2V_5B_Config(), + available_memory_gb=72, + kwargs={ + "model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", + "performance_mode": "auto", + }, + ) + + self.assertFalse(args.dit_cpu_offload) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) + + def test_auto_fastwan_offloads_dit_below_resident_threshold(self): + args = self._from_dict_with_pipeline_config( + FastWan2_2_TI2V_5B_Config(), + memory_gb=48, kwargs={ "model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", "performance_mode": "auto", @@ -1269,12 +1476,33 @@ class TestOffloadDefaults(unittest.TestCase): ) self.assertTrue(args.dit_cpu_offload) - self.assertEqual( - args.layerwise_offload_components, - ["text_encoder", "image_encoder", "vae"], + + def test_auto_fast_hunyuan_keeps_dit_resident_on_h100(self): + args = self._from_dict_with_pipeline_config( + FastHunyuanConfig(), + available_memory_gb=72, + kwargs={ + "model_path": "FastVideo/FastHunyuan-diffusers", + "performance_mode": "auto", + }, ) - def test_auto_turbo_wan_layerwise_offload_does_not_implicitly_add_dit(self): + self.assertFalse(args.dit_cpu_offload) + self.assertFalse(args.vae_cpu_offload) + + def test_auto_fast_hunyuan_offloads_dit_below_resident_threshold(self): + args = self._from_dict_with_pipeline_config( + FastHunyuanConfig(), + memory_gb=48, + kwargs={ + "model_path": "FastVideo/FastHunyuan-diffusers", + "performance_mode": "auto", + }, + ) + + self.assertTrue(args.dit_cpu_offload) + + def test_auto_turbo_wan_keeps_dit_resident_on_h100(self): args = self._from_dict_with_pipeline_config( TurboWanT2V480PConfig(), kwargs={ @@ -1283,7 +1511,7 @@ class TestOffloadDefaults(unittest.TestCase): }, ) - self.assertTrue(args.dit_cpu_offload) + self.assertFalse(args.dit_cpu_offload) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder", "vae"], @@ -1316,7 +1544,7 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.enable_cfg_parallel) - self.assertTrue(args.dit_cpu_offload) + self.assertFalse(args.dit_cpu_offload) self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) @@ -1448,9 +1676,9 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) self.assertTrue(args.enable_cfg_parallel) - # 80gb > image threshold (45gb): only vae resident, encoders offloaded; - # cfg/dit unchanged - self.assertTrue(args.dit_cpu_offload) + # 80gb > image threshold (45gb): vae and dit stay resident, while the + # large encoders use layerwise offload. + self.assertFalse(args.dit_cpu_offload) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder"], @@ -1518,9 +1746,9 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) 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) + # 50gb still > image threshold (45gb): vae and dit stay resident, while + # the encoders remain offloaded; qwen does not opt into auto fsdp. + self.assertFalse(args.dit_cpu_offload) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder"], @@ -1557,8 +1785,8 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) 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) + # vae and dit stay resident, while the encoders remain offloaded. + self.assertFalse(args.dit_cpu_offload) self.assertEqual( args.layerwise_offload_components, ["text_encoder", "image_encoder"],