[diffusion] feat: support dynamically cpu offload components (#34391)

This commit is contained in:
Mick
2026-08-12 11:38:27 +08:00
committed by GitHub
parent 5899674504
commit a9a355774a
31 changed files with 634 additions and 176 deletions
+14 -1
View File
@@ -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-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. - `--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 - `--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. - `--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. - `--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 - `--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. 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 ### 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): 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):
@@ -128,7 +128,7 @@ See [OpenAI API: Served model name](/docs/sglang-diffusion/api/openai_api#served
</tbody> </tbody>
</table> </table>
`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. 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.
@@ -164,5 +164,6 @@ class FastHunyuanConfig(HunyuanConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
keep_resident_components=("vae",), keep_resident_min_available_gb=60,
keep_resident_components=("dit", "vae"),
) )
@@ -338,6 +338,13 @@ class LingBotWorldCausalDMDConfig(LingBotWorldI2VConfig):
interactive_kv_still_chunks: int = 2 interactive_kv_still_chunks: int = 2
lazy_vae_encode_black_frames: int = 0 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): def preprocess_vae_encode(self, image, vae):
image = super().preprocess_vae_encode(image, vae) image = super().preprocess_vae_encode(image, vae)
lazy_black_frames = envs.SGLANG_LINGBOT_LAZY_VAE_ENCODE_BLACK_FRAMES lazy_black_frames = envs.SGLANG_LINGBOT_LAZY_VAE_ENCODE_BLACK_FRAMES
@@ -14,9 +14,9 @@ class ModelDeploymentConfig:
dit_layerwise_offload_modes: tuple[Literal["auto", "memory"], ...] = () dit_layerwise_offload_modes: tuple[Literal["auto", "memory"], ...] = ()
auto_dit_offload_prefetch_size: float | None = None auto_dit_offload_prefetch_size: float | None = None
keep_resident_min_available_gb: float | None = None keep_resident_min_available_gb: float | None = None
# only vae -- it is tiny so keeping it resident barely shifts memory; large # Per-model resident defaults. Auto mode additionally keeps an image DiT
# encoders stay offloaded and dit placement stays with the FSDP/dit-layerwise # resident above the image workload memory threshold; video DiT placement
# policy # stays with the model's FSDP/layerwise policy.
keep_resident_components: tuple[OffloadComponentName, ...] = ("vae",) keep_resident_components: tuple[OffloadComponentName, ...] = ("vae",)
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
@@ -99,6 +99,8 @@ class WanT2V480PConfig(PipelineConfig):
def get_model_deployment_config(self) -> ModelDeploymentConfig: def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig( return ModelDeploymentConfig(
dit_layerwise_offload_modes=("memory",), dit_layerwise_offload_modes=("memory",),
keep_resident_min_available_gb=60,
keep_resident_components=("dit",),
) )
def expand_conditioning_to_sample_batch(self, batch): def expand_conditioning_to_sample_batch(self, batch):
@@ -141,6 +143,7 @@ class TurboWanT2V1_3B480PConfig(TurboWanT2V480PConfig):
dit_layerwise_offload_modes=("memory",), dit_layerwise_offload_modes=("memory",),
keep_resident_min_available_gb=60, keep_resident_min_available_gb=60,
keep_resident_components=( keep_resident_components=(
"dit",
"text_encoder", "text_encoder",
"image_encoder", "image_encoder",
"vae", "vae",
@@ -182,11 +185,6 @@ class WanI2V480PConfig(WanT2V480PConfig, WanI2VCommonConfig):
self.vae_config.load_encoder = True self.vae_config.load_encoder = True
self.vae_config.load_decoder = True self.vae_config.load_decoder = True
def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig(
dit_layerwise_offload_modes=("memory",),
)
@dataclass @dataclass
class WanI2V720PConfig(WanI2V480PConfig): class WanI2V720PConfig(WanI2V480PConfig):
@@ -228,6 +226,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
dit_layerwise_offload_modes=("memory",), dit_layerwise_offload_modes=("memory",),
keep_resident_min_available_gb=60, keep_resident_min_available_gb=60,
keep_resident_components=( keep_resident_components=(
"dit",
"text_encoder", "text_encoder",
"image_encoder", "image_encoder",
"vae", "vae",
@@ -82,7 +82,10 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig):
F_PATCH_SIZE: int = 1 F_PATCH_SIZE: int = 1
def get_model_deployment_config(self) -> ModelDeploymentConfig: 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): def prepare_sigmas(self, sigmas, num_inference_steps):
return self._prepare_sigmas(sigmas, num_inference_steps) return self._prepare_sigmas(sigmas, num_inference_steps)
@@ -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 ( from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import (
LTX2ConnectorConfig, LTX2ConnectorConfig,
) )
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
) )
@@ -50,7 +49,9 @@ class AdapterLoader(ComponentLoader):
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) 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( default_dtype = resolve_precision(
server_args, "connectors", precision_attr="dit_precision" server_args, "connectors", precision_attr="dit_precision"
) )
@@ -74,6 +74,8 @@ class BridgeLoader(ComponentLoader):
default_dtype, 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. # Use the FSDP loader when FSDP is requested or shard rules are declared.
fsdp_shard_conditions = getattr(model_cls, "_fsdp_shard_conditions", None) fsdp_shard_conditions = getattr(model_cls, "_fsdp_shard_conditions", None)
if server_args.use_fsdp_inference or ( if server_args.use_fsdp_inference or (
@@ -88,7 +90,7 @@ class BridgeLoader(ComponentLoader):
device=local_torch_device, device=local_torch_device,
hsdp_replicate_dim=server_args.hsdp_replicate_dim, hsdp_replicate_dim=server_args.hsdp_replicate_dim,
hsdp_shard_dim=server_args.hsdp_shard_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, pin_cpu_memory=server_args.pin_cpu_memory,
fsdp_inference=server_args.use_fsdp_inference, fsdp_inference=server_args.use_fsdp_inference,
param_dtype=default_dtype, param_dtype=default_dtype,
@@ -104,7 +106,8 @@ class BridgeLoader(ComponentLoader):
model = model_cls.from_pretrained( model = model_cls.from_pretrained(
component_model_path, torch_dtype=default_dtype 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()) total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded bridge model with %.2fM parameters", total_params / 1e6) logger.info("Loaded bridge model with %.2fM parameters", total_params / 1e6)
@@ -25,6 +25,9 @@ from sglang.multimodal_gen.runtime.loader.utils import (
component_name_to_loader_cls, component_name_to_loader_cls,
get_memory_usage_of_component, 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 ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
configure_layerwise_offload_modules, configure_layerwise_offload_modules,
is_layerwise_offloaded_module, is_layerwise_offloaded_module,
@@ -92,10 +95,14 @@ class ComponentLoader(ABC):
self.component_architecture: str | None = None self.component_architecture: str | None = None
def should_offload( 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 component_name is not None and server_args.should_cpu_offload_component(
return False component_name
)
def target_device(self, should_offload): def target_device(self, should_offload):
if should_offload: if should_offload:
@@ -215,7 +222,7 @@ class ComponentLoader(ABC):
transformers_or_diffusers, transformers_or_diffusers,
component_name, 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) target_device = self.target_device(should_offload)
return component.to(device=target_device) return component.to(device=target_device)
@@ -301,6 +308,12 @@ class ComponentLoader(ABC):
else: else:
if isinstance(component, nn.Module): if isinstance(component, nn.Module):
component = component.eval() 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() current_gpu_mem = current_platform.get_available_gpu_memory()
model_size = get_memory_usage_of_component(component) or "NA" model_size = get_memory_usage_of_component(component) or "NA"
consumed = gpu_mem_before_loading - current_gpu_mem consumed = gpu_mem_before_loading - current_gpu_mem
@@ -16,8 +16,14 @@ class ImageEncoderLoader(TextEncoderLoader):
component_names = ["image_encoder"] component_names = ["image_encoder"]
expected_library = "transformers" expected_library = "transformers"
def should_offload(self, server_args, model_config: ModelConfig | None = None): def should_offload(
should_offload = server_args.image_encoder_cpu_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: if not should_offload:
return False return False
# _fsdp_shard_conditions is in arch_config, not directly on model_config # _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=(
cpu_offload_flag cpu_offload_flag
if cpu_offload_flag is not None 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,
) )
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from safetensors.torch import load_file as safetensors_load_file 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 ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
) )
@@ -25,11 +24,6 @@ class SoundTokenizerLoader(ComponentLoader):
component_names = ["sound_tokenizer"] component_names = ["sound_tokenizer"]
expected_library = "diffusers" 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( def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str self, component_model_path: str, server_args: ServerArgs, component_name: str
): ):
@@ -46,7 +40,9 @@ class SoundTokenizerLoader(ComponentLoader):
except AttributeError: except AttributeError:
precision = "bf16" precision = "bf16"
dtype = PRECISION_TO_TYPE[precision] 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(): with set_default_torch_dtype(dtype), skip_init_modules():
model_cls, _ = ModelRegistry.resolve_model_cls(class_name) model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
@@ -80,8 +80,14 @@ class TextEncoderLoader(ComponentLoader):
allow_patterns_overrides: list[str] | None = None allow_patterns_overrides: list[str] | None = None
"""If defined, weights will load exclusively using these patterns.""" """If defined, weights will load exclusively using these patterns."""
def should_offload(self, server_args, model_config: ModelConfig | None = None): def should_offload(
should_offload = server_args.text_encoder_cpu_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: if not should_offload:
return False return False
# _fsdp_shard_conditions is in arch_config, not directly on model_config # _fsdp_shard_conditions is in arch_config, not directly on model_config
@@ -369,6 +375,7 @@ class TextEncoderLoader(ComponentLoader):
server_args, server_args,
encoder_dtype, encoder_dtype,
cpu_offload_flag=cpu_offload_flag, cpu_offload_flag=cpu_offload_flag,
component_name=component_name,
) )
@staticmethod @staticmethod
@@ -400,13 +407,16 @@ class TextEncoderLoader(ComponentLoader):
server_args: ServerArgs, server_args: ServerArgs,
dtype: str = "fp16", dtype: str = "fp16",
cpu_offload_flag: bool | None = None, cpu_offload_flag: bool | None = None,
component_name: str = "text_encoder",
): ):
# Determine CPU offload behavior and target device # Determine CPU offload behavior and target device
local_torch_device = get_local_torch_device() local_torch_device = get_local_torch_device()
if not current_platform.is_cpu(): 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 = ( should_offload = (
cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload
) )
@@ -152,6 +152,11 @@ class TransformerLoader(ComponentLoader):
component_server_args = _server_args_for_transformer_component( component_server_args = _server_args_for_transformer_component(
server_args, component_name 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 # 1. hf config
config = get_diffusers_component_config(component_path=component_model_path) config = get_diffusers_component_config(component_path=component_model_path)
@@ -195,9 +195,6 @@ class UpsamplerLoader(ComponentLoader):
component_names = ["spatial_upsampler"] component_names = ["spatial_upsampler"]
expected_library = "diffusers" expected_library = "diffusers"
def should_offload(self, server_args: ServerArgs, model_config=None):
return server_args.vae_cpu_offload
def load_customized( def load_customized(
self, self,
component_model_path: str, component_model_path: str,
@@ -210,7 +207,7 @@ class UpsamplerLoader(ComponentLoader):
logger.info("Loading LatentUpsampler with config: %s", config) 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) target_device = self.target_device(should_offload)
with torch.device("meta"): with torch.device("meta"):
@@ -5,7 +5,6 @@ import torch
import torch.nn as nn import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file 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.ltx_2 import LTX2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImagePipelineConfig, QwenImagePipelineConfig,
@@ -99,11 +98,6 @@ class VAELoader(ComponentLoader):
component_names = ["vae", "audio_vae", "video_vae"] component_names = ["vae", "audio_vae", "video_vae"]
expected_library = "diffusers" expected_library = "diffusers"
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
):
return server_args.vae_cpu_offload
def load_customized( def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str 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 # NOTE: some post init logics are only available after updated with config
vae_config.post_init() 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) target_device = self.target_device(should_offload)
native_only = component_name in getattr( native_only = component_name in getattr(
@@ -3,7 +3,6 @@ from typing import Any
import requests import requests
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
) )
@@ -59,12 +58,15 @@ class VisionLanguageEncoderLoader(ComponentLoader):
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, revision=server_args.revision,
) )
target_device = self.target_device(
server_args.should_cpu_offload_component("vision_language_encoder")
)
model = GlmImageForConditionalGeneration.from_pretrained( model = GlmImageForConditionalGeneration.from_pretrained(
component_model_path, component_model_path,
config=config, config=config,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, revision=server_args.revision,
).to(get_local_torch_device()) ).to(target_device)
return model return model
else: else:
raise ValueError( raise ValueError(
@@ -1,6 +1,5 @@
from safetensors.torch import load_file as safetensors_load_file 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 ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
) )
@@ -25,11 +24,6 @@ class VocoderLoader(ComponentLoader):
component_names = ["vocoder"] component_names = ["vocoder"]
expected_library = "diffusers" expected_library = "diffusers"
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
):
return server_args.vae_cpu_offload
def load_customized( def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str self, component_model_path: str, server_args: ServerArgs, component_name: str
): ):
@@ -55,7 +49,7 @@ class VocoderLoader(ComponentLoader):
else PRECISION_TO_TYPE["fp32"] 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) target_device = self.target_device(should_offload)
with set_default_torch_dtype(vocoder_dtype), skip_init_modules(): with set_default_torch_dtype(vocoder_dtype), skip_init_modules():
@@ -317,46 +317,24 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
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_str = self._format_offload_disable_suggestions(can_stay_resident)
pool_overhead_gb = peak_reserved_gb - peak_allocated_gb pool_overhead_gb = peak_reserved_gb - peak_allocated_gb
pool_overhead_pct = (
logger.debug( pool_overhead_gb / peak_reserved_gb * 100 if peak_reserved_gb else 0.0
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}"
) )
def _format_offload_disable_suggestions(self, components: List[str]) -> str: logger.debug(
component_set = set(components) "GPU memory: peak=%.2f GB, allocated=%.2f GB, pool=%.2f GB (%.1f%%), "
suggestions = [] "headroom=%.2f GB. Components that can remain on GPU: %s. "
seen_args = set() "Adjust --cpu-offload-components or --layerwise-offload-components "
"to change residency.",
for component in OFFLOAD_DISABLE_RECOMMENDATION_ORDER: peak_reserved_gb,
if component not in component_set: peak_allocated_gb,
continue pool_overhead_gb,
pool_overhead_pct,
arg = None remaining_gpu_mem_gb,
if component == "vae": can_stay_resident,
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"
def execute_forward( def execute_forward(
self, batch: List[Req], return_req: bool = False self, batch: List[Req], return_req: bool = False
@@ -974,28 +952,28 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
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. memory_usages = self.pipeline.memory_usages
# If the flag is False, the component is already resident, so we do not suggest it. ordered_names = [
# If the flag is True, it is currently offloaded, so it is a candidate to stay resident. name
offload_flags = { for name in OFFLOAD_DISABLE_RECOMMENDATION_ORDER
"transformer": self.server_args.dit_cpu_offload if name in memory_usages
or self.server_args.is_dit_layerwise_offload_selected, ]
"vae": self.server_args.vae_cpu_offload, ordered_names.extend(
"text_encoder": self.server_args.text_encoder_cpu_offload, name
"text_encoder_2": self.server_args.text_encoder_cpu_offload, for name in memory_usages
"image_encoder": self.server_args.image_encoder_cpu_offload, if name not in OFFLOAD_DISABLE_RECOMMENDATION_ORDER
} )
for name in ordered_names:
for name in OFFLOAD_DISABLE_RECOMMENDATION_ORDER: usage = memory_usages[name]
# Only consider components that are currently configured to be offloaded if not (
is_offload_configured = offload_flags.get(name, False) self.server_args.should_cpu_offload_component(name)
if not is_offload_configured: or self.server_args.should_configure_layerwise_offload_for_lazy_component(
name
)
):
continue continue
usage = self.pipeline.memory_usages.get(name)
if usage is None: if usage is None:
continue 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
@@ -18,12 +18,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
is_layerwise_offloaded_module, is_layerwise_offloaded_module,
is_resident_layerwise_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.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -93,24 +87,6 @@ class ComponentResidencyPipeline(Protocol):
component_residency_strategies: MutableMapping[str, "ComponentResidencyStrategy"] 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( def build_component_residency_strategy(
component_name: str, component_name: str,
module: nn.Module, module: nn.Module,
@@ -118,7 +94,12 @@ def build_component_residency_strategy(
) -> ComponentResidencyStrategy: ) -> ComponentResidencyStrategy:
if is_layerwise_offloaded_module(module): if is_layerwise_offloaded_module(module):
return LayerwiseOffloadStrategy() 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 VanillaD2HStrategy()
return ResidentStrategy() return ResidentStrategy()
@@ -164,7 +145,6 @@ class ComponentResidencyManager:
if pipeline is not self.pipeline: if pipeline is not self.pipeline:
self._remove_nvtx_hooks() self._remove_nvtx_hooks()
self.strategy_for.cache_clear() self.strategy_for.cache_clear()
self._should_keep_single_dit.cache_clear()
self._active_use = None self._active_use = None
self._active_use_module = None self._active_use_module = None
self._uses_seen.clear() self._uses_seen.clear()
@@ -466,8 +446,15 @@ class ComponentResidencyManager:
preferred = component_name in preferred_uses preferred = component_name in preferred_uses
if is_resident_layerwise_module(module): if is_resident_layerwise_module(module):
preferred = False 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 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) strategy = self.strategy_for(component_name, module)
if preferred and not self.state.batch_is_warmup: if preferred and not self.state.batch_is_warmup:
strategy.prepare_after_request(module, use, self.state) strategy.prepare_after_request(module, use, self.state)
@@ -544,16 +531,25 @@ class ComponentResidencyManager:
} }
if use.component_name in future_component_names: if use.component_name in future_component_names:
return True return True
if self._should_keep_single_dit(use.component_name): module = self.get_module(use.component_name)
module = self.get_module(use.component_name) if module is not None and self._should_keep_single_dit(
if module is not None and is_resident_layerwise_module(module): use.component_name, module
# don't keep a layerwise DiT resident across the request to avoid OOMs ):
return False
return True return True
return False return False
@lru_cache(maxsize=None) def _should_keep_single_dit(self, component_name: str, module: nn.Module) -> bool:
def _should_keep_single_dit(self, component_name: str) -> 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 modules = self.pipeline.modules
return (component_name == "transformer" and "transformer_2" not in modules) or ( return (component_name == "transformer" and "transformer_2" not in modules) or (
component_name == "video_dit" and "video_dit_2" not in modules component_name == "video_dit" and "video_dit_2" not in modules
@@ -47,12 +47,62 @@ CPU_OFFLOAD_FLAG_NAMES = (
"image_encoder_cpu_offload", "image_encoder_cpu_offload",
"vae_cpu_offload", "vae_cpu_offload",
) )
CPU_OFFLOAD_ALL_COMPONENTS = "all"
def is_dit_component_name(component_name: str) -> bool: def is_dit_component_name(component_name: str) -> bool:
return component_name in DIT_COMPONENT_NAMES 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: def is_text_encoder_component_name(component_name: str) -> bool:
return component_name.startswith("text_encoder") or component_name.endswith( return component_name.startswith("text_encoder") or component_name.endswith(
"text_encoder" "text_encoder"
@@ -654,6 +654,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
target_dtype: torch.dtype, target_dtype: torch.dtype,
autocast_enabled: bool, autocast_enabled: bool,
) -> torch.Tensor: ) -> torch.Tensor:
if self._component_residency_manager is not None:
self._manage_dit_use_site(self.transformer, "transformer", batch)
with ( with (
precision_autocast_context( precision_autocast_context(
target_dtype, target_dtype,
@@ -1531,12 +1531,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
batch: Req, batch: Req,
) -> None: ) -> None:
""" """
manage dit's residency by reporting the active sequential use manage dit residency by reporting the active sequential use
only applicable for dual-dit architecture like Wan
Args: Args:
current_model: the next active dit, transformer_1 or transformer_2 current_model: the next active dit
""" """
manager = self._component_residency_manager manager = self._component_residency_manager
@@ -534,7 +534,6 @@ class LongLive2CausalDenoisingStage(CausalDMDDenoisingStage):
target_dtype: torch.dtype, target_dtype: torch.dtype,
autocast_enabled: bool, autocast_enabled: bool,
) -> torch.Tensor: ) -> torch.Tensor:
self._manage_dit_use_site(self.transformer, "transformer", batch)
rope_start_frame = start_frame rope_start_frame = start_frame
if self._rope_temporal_offset != 0.0: if self._rope_temporal_offset != 0.0:
rope_start_frame = start_frame + self._rope_temporal_offset rope_start_frame = start_frame + self._rope_temporal_offset
@@ -140,7 +140,9 @@ class ServerArgsAutoTuner:
and min_available_gb >= disable_threshold_gb and min_available_gb >= disable_threshold_gb
): ):
changed = [] 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 ( 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")
@@ -253,6 +255,7 @@ class ServerArgsAutoTuner:
if ( if (
args.layerwise_offload_components is not None args.layerwise_offload_components is not None
or args.dit_layerwise_offload is True or args.dit_layerwise_offload is True
or args.is_arg_explicitly_set("cpu_offload_components")
): ):
return return
if not current_platform.is_cuda(): if not current_platform.is_cuda():
@@ -410,6 +413,7 @@ class ServerArgsAutoTuner:
if ( if (
args.is_arg_explicitly_set("layerwise_offload_components") args.is_arg_explicitly_set("layerwise_offload_components")
or args.dit_layerwise_offload is True 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. # The legacy --dit-layerwise-offload flag is a DiT-only selector.
# Do not merge implicit defaults into that explicit mode. # Do not merge implicit defaults into that explicit mode.
@@ -475,6 +479,7 @@ class ServerArgsAutoTuner:
or envs.SGLANG_CACHE_DIT_ENABLED or envs.SGLANG_CACHE_DIT_ENABLED
or args.use_fsdp_inference or args.use_fsdp_inference
or args.is_arg_explicitly_set("dit_cpu_offload") or args.is_arg_explicitly_set("dit_cpu_offload")
or args.is_arg_explicitly_set("cpu_offload_components")
): ):
return False return False
@@ -523,6 +528,7 @@ class ServerArgsAutoTuner:
"dit_cpu_offload", "dit_cpu_offload",
"dit_layerwise_offload", "dit_layerwise_offload",
"layerwise_offload_components", "layerwise_offload_components",
"cpu_offload_components",
) )
) )
@@ -533,6 +539,7 @@ class ServerArgsAutoTuner:
for arg_name in ( for arg_name in (
"dit_layerwise_offload", "dit_layerwise_offload",
"layerwise_offload_components", "layerwise_offload_components",
"cpu_offload_components",
) )
) )
@@ -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 ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import (
LAYERWISE_OFFLOAD_ALL_COMPONENTS, LAYERWISE_OFFLOAD_ALL_COMPONENTS,
LAYERWISE_OFFLOAD_DIT_GROUP, LAYERWISE_OFFLOAD_DIT_GROUP,
cpu_offload_component_matches,
cpu_offload_flags_for_layerwise_components, 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, layerwise_component_matches_any_selection,
normalize_cpu_offload_components,
normalize_layerwise_offload_components, normalize_layerwise_offload_components,
) )
from sglang.multimodal_gen.runtime.platforms import ( from sglang.multimodal_gen.runtime.platforms import (
@@ -301,6 +307,8 @@ class ServerArgs(DisaggServerArgsMixin):
lora_target_modules: list[str] | None = None lora_target_modules: list[str] | None = None
# CPU offload parameters # 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 dit_cpu_offload: bool | None = None
# trade checkpoint-loading peak memory for faster ordinary DiT startup # trade checkpoint-loading peak memory for faster ordinary DiT startup
direct_gpu_weight_loading: bool = False direct_gpu_weight_loading: bool = False
@@ -486,13 +494,14 @@ class ServerArgs(DisaggServerArgsMixin):
"""set defaults and normalize values.""" """set defaults and normalize values."""
auto_tuner = ServerArgsAutoTuner(self) auto_tuner = ServerArgsAutoTuner(self)
auto_tuner.adjust_based_on_performance_mode() auto_tuner.adjust_based_on_performance_mode()
self._adjust_cpu_offload_components()
if auto_tuner.could_override_server_args(): if auto_tuner.could_override_server_args():
self._adjust_offload() self._adjust_offload()
auto_tuner.maybe_adjust_auto_default_layerwise_offload() auto_tuner.maybe_adjust_auto_default_layerwise_offload()
self._adjust_ltx2_two_stage_device_mode() self._adjust_ltx2_two_stage_device_mode()
if auto_tuner.could_override_server_args(): 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_fsdp_with_offload_enabled()
auto_tuner.maybe_adjust_auto_component_residency_after_offload()
auto_tuner.maybe_replace_cpu_offloaded_components_with_layerwise() auto_tuner.maybe_replace_cpu_offloaded_components_with_layerwise()
self._adjust_path() self._adjust_path()
if self.served_model_name is None: if self.served_model_name is None:
@@ -739,6 +748,44 @@ class ServerArgs(DisaggServerArgsMixin):
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
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): def _adjust_ltx2_two_stage_device_mode(self):
if not self._is_ltx23_two_stage_pipeline(): if not self._is_ltx23_two_stage_pipeline():
return return
@@ -1240,6 +1287,25 @@ class ServerArgs(DisaggServerArgsMixin):
def is_arg_explicitly_set(self, arg_name: str) -> bool: def is_arg_explicitly_set(self, arg_name: str) -> bool:
return arg_name in self._explicit_arg_names 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( def should_configure_layerwise_offload_for_lazy_component(
self, component_name: str self, component_name: str
) -> bool: ) -> bool:
@@ -1797,6 +1863,19 @@ class ServerArgs(DisaggServerArgsMixin):
"time depending on the model, but temporarily requires checkpoint " "time depending on the model, but temporarily requires checkpoint "
"weights and model weights to coexist on GPU. Disabled by default.", "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( parser.add_argument(
"--dit-layerwise-offload", "--dit-layerwise-offload",
action=StoreBoolean, action=StoreBoolean,
@@ -39,7 +39,7 @@ logger = init_logger(__name__)
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows # NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
# publish. # publish.
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion" 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 # The NPU pin is kept as a separate branch so ascend GT can be bumped independently
# when it's regenerated on its own cadence. # when it's regenerated on its own cadence.
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from contextlib import nullcontext
from types import MethodType, SimpleNamespace from types import MethodType, SimpleNamespace
import torch import torch
@@ -21,6 +22,61 @@ class _Progress:
self.count += 1 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(): def test_causal_dmd_chunk_loop_uses_model_input_builder():
stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage)
predict_calls = [] predict_calls = []
@@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
is_layerwise_offloaded_module, is_layerwise_offloaded_module,
is_resident_layerwise_module, is_resident_layerwise_module,
) )
from sglang.multimodal_gen.runtime.server_args import ServerArgs
class _FakeStream: class _FakeStream:
@@ -161,8 +162,13 @@ class _LayerwiseComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin):
self.layerwise_offload_managers = [SimpleNamespace(enabled=enabled)] self.layerwise_offload_managers = [SimpleNamespace(enabled=enabled)]
class _TestServerArgs(SimpleNamespace):
should_cpu_offload_component = ServerArgs.should_cpu_offload_component
def _server_args(**kwargs): def _server_args(**kwargs):
defaults = dict( defaults = dict(
cpu_offload_components=None,
use_fsdp_inference=False, use_fsdp_inference=False,
dit_cpu_offload=False, dit_cpu_offload=False,
text_encoder_cpu_offload=False, text_encoder_cpu_offload=False,
@@ -173,7 +179,7 @@ def _server_args(**kwargs):
pin_cpu_memory=False, pin_cpu_memory=False,
) )
defaults.update(kwargs) defaults.update(kwargs)
return SimpleNamespace(**defaults) return _TestServerArgs(**defaults)
def test_layerwise_offload_preserves_non_contiguous_stride(monkeypatch): def test_layerwise_offload_preserves_non_contiguous_stride(monkeypatch):
@@ -16,6 +16,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im
ComponentResidencyManager, ComponentResidencyManager,
ComponentUse, 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 import nvtx_pytorch_hooks
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import ( from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import (
DiffusionNvtxHooks, DiffusionNvtxHooks,
@@ -228,6 +232,16 @@ def _test_manager(
class TestComponentResidencyNvtxHooks(unittest.TestCase): 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: def test_disabled_flag_is_noop(self) -> None:
module = torch.nn.Linear(2, 2) module = torch.nn.Linear(2, 2)
manager = _test_manager({"linear": module}, enable_flag=False) manager = _test_manager({"linear": module}, enable_flag=False)
@@ -17,6 +17,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
PipelineConfig, PipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanConfig 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 ( from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig, LTX2PipelineConfig,
LTX23PipelineConfig, LTX23PipelineConfig,
@@ -385,6 +388,32 @@ class TestServerArgsPathExpansion(unittest.TestCase):
server_args.layerwise_offload_components, ["transformer", "text_encoder"] 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): def test_serve_cli_preserves_config_and_dynamic_unknown_args(self):
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import ( from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args, add_multimodal_gen_serve_args,
@@ -814,6 +843,74 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.vae_cpu_offload) 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): def test_vae_cpu_offload_defaults_false_on_low_memory_gpu(self):
args = self._from_dict_with_task_type( args = self._from_dict_with_task_type(
ModelTaskType.T2V, ModelTaskType.T2V,
@@ -906,6 +1003,7 @@ class TestOffloadDefaults(unittest.TestCase):
wan_deployment = WanT2V480PConfig().get_model_deployment_config() wan_deployment = WanT2V480PConfig().get_model_deployment_config()
mova_deployment = MOVAPipelineConfig().get_model_deployment_config() mova_deployment = MOVAPipelineConfig().get_model_deployment_config()
zimage_deployment = ZImagePipelineConfig().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() ltx_deployment = LTX2PipelineConfig().get_model_deployment_config()
ltx23_config = LTX23PipelineConfig() ltx23_config = LTX23PipelineConfig()
sana_wm_deployment = SanaWMPipelineConfig().get_model_deployment_config() 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.assertIsNone(wan_deployment.fsdp_auto_min_available_memory_gb)
self.assertEqual(wan_deployment.dit_layerwise_offload_modes, ("memory",)) 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.assertIsNone(mova_deployment.fsdp_auto_min_available_memory_gb)
self.assertEqual( self.assertEqual(
@@ -924,9 +1024,14 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertEqual(mova_deployment.keep_resident_components, ("dit", "vae")) self.assertEqual(mova_deployment.keep_resident_components, ("dit", "vae"))
self.assertEqual(zimage_deployment.fsdp_auto_min_available_memory_gb, 40) 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.assertTrue(zimage_deployment.fsdp_auto_requires_cfg)
self.assertEqual(zimage_deployment.dit_layerwise_offload_modes, ()) 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_min_available_gb, 70)
self.assertEqual(ltx_deployment.keep_resident_components, ("dit",)) self.assertEqual(ltx_deployment.keep_resident_components, ("dit",))
self.assertEqual( 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.fsdp_auto_min_available_memory_gb, 60)
self.assertEqual(sana_wm_deployment.dit_layerwise_offload_modes, ("memory",)) 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() 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_min_available_gb, 60)
self.assertEqual(fast_hunyuan_deployment.keep_resident_components, ("vae",)) 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) # default keeps only vae resident (encoders are large, dit owned by FSDP)
self.assertEqual(qwen_deployment.keep_resident_components, ("vae",)) self.assertEqual(qwen_deployment.keep_resident_components, ("vae",))
@@ -1051,9 +1169,9 @@ 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 # 80gb > image threshold (45gb): vae and dit stay resident, while the
# offloaded layerwise, dit unchanged # large encoders use layerwise offload.
self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.dit_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder"], ["text_encoder", "image_encoder"],
@@ -1075,6 +1193,53 @@ class TestOffloadDefaults(unittest.TestCase):
["text_encoder", "image_encoder", "vae"], ["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( def test_auto_ltx_original_replaces_component_cpu_offload(
self, self,
): ):
@@ -1098,7 +1263,7 @@ class TestOffloadDefaults(unittest.TestCase):
["text_encoder", "image_encoder", "vae"], ["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( args = self._from_dict_with_pipeline_config(
WanT2V480PConfig(), WanT2V480PConfig(),
kwargs={"performance_mode": "auto"}, kwargs={"performance_mode": "auto"},
@@ -1106,7 +1271,7 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertTrue(args.layerwise_offload_components) self.assertTrue(args.layerwise_offload_components)
self.assertFalse(args.use_fsdp_inference) 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.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload)
self.assertEqual( self.assertEqual(
@@ -1114,6 +1279,19 @@ class TestOffloadDefaults(unittest.TestCase):
["text_encoder", "image_encoder", "vae"], ["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): def test_auto_wan2_2_a14b_layerwise_offload_adds_dit(self):
for pipeline_config, model_path in ( for pipeline_config, model_path in (
(Wan2_2_T2V_A14B_Config(), "Wan-AI/Wan2.2-T2V-A14B-Diffusers"), (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"], ["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 ( for pipeline_config, model_path in (
(WanT2V720PConfig(), "Wan-AI/Wan2.1-T2V-14B-Diffusers"), (WanT2V720PConfig(), "Wan-AI/Wan2.1-T2V-14B-Diffusers"),
(WanI2V480PConfig(), "Wan-AI/Wan2.1-I2V-14B-480P-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.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.dit_offload_prefetch_size, 0.0)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["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): def test_memory_wan_layerwise_offload_is_enabled_without_fsdp(self):
args = self._from_dict_with_pipeline_config( args = self._from_dict_with_pipeline_config(
WanT2V480PConfig(), WanT2V480PConfig(),
@@ -1259,9 +1449,26 @@ class TestOffloadDefaults(unittest.TestCase):
["dit", "text_encoder", "image_encoder", "vae"], ["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( args = self._from_dict_with_pipeline_config(
FastWan2_2_TI2V_5B_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={ kwargs={
"model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", "model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"performance_mode": "auto", "performance_mode": "auto",
@@ -1269,12 +1476,33 @@ class TestOffloadDefaults(unittest.TestCase):
) )
self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.dit_cpu_offload)
self.assertEqual(
args.layerwise_offload_components, def test_auto_fast_hunyuan_keeps_dit_resident_on_h100(self):
["text_encoder", "image_encoder", "vae"], 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( args = self._from_dict_with_pipeline_config(
TurboWanT2V480PConfig(), TurboWanT2V480PConfig(),
kwargs={ kwargs={
@@ -1283,7 +1511,7 @@ class TestOffloadDefaults(unittest.TestCase):
}, },
) )
self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.dit_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder", "vae"], ["text_encoder", "image_encoder", "vae"],
@@ -1316,7 +1544,7 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.use_fsdp_inference)
self.assertFalse(args.enable_cfg_parallel) 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.assertTrue(args.layerwise_offload_components)
self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_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.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; # 80gb > image threshold (45gb): vae and dit stay resident, while the
# cfg/dit unchanged # large encoders use layerwise offload.
self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.dit_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder"], ["text_encoder", "image_encoder"],
@@ -1518,9 +1746,9 @@ 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; # 50gb still > image threshold (45gb): vae and dit stay resident, while
# fsdp skipped (qwen does not opt into auto fsdp) # the encoders remain offloaded; qwen does not opt into auto fsdp.
self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.dit_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder"], ["text_encoder", "image_encoder"],
@@ -1557,8 +1785,8 @@ 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): # min available across selected gpus is 72gb > image threshold (45gb):
# vae resident, encoders offloaded # vae and dit stay resident, while the encoders remain offloaded.
self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.dit_cpu_offload)
self.assertEqual( self.assertEqual(
args.layerwise_offload_components, args.layerwise_offload_components,
["text_encoder", "image_encoder"], ["text_encoder", "image_encoder"],