[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-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):
@@ -128,7 +128,7 @@ See [OpenAI API: Served model name](/docs/sglang-diffusion/api/openai_api#served
</tbody>
</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.
@@ -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"),
)
@@ -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
@@ -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
@@ -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",
@@ -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)
@@ -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"
)
@@ -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)
@@ -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
@@ -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,
)
@@ -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)
@@ -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
)
@@ -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)
@@ -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"):
@@ -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(
@@ -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(
@@ -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():
@@ -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
@@ -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
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
@@ -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"
@@ -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,
@@ -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
@@ -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
@@ -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",
)
)
@@ -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,
@@ -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.
@@ -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 = []
@@ -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):
@@ -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)
@@ -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"],