[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
@@ -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
module = self.get_module(use.component_name)
if module is not None and self._should_keep_single_dit(
use.component_name, module
):
return True
return False
@lru_cache(maxsize=None)
def _should_keep_single_dit(self, component_name: str) -> bool:
def _should_keep_single_dit(self, component_name: str, module: nn.Module) -> bool:
"""Keep a single DiT resident only when its effective strategy is resident.
The single-DiT fast path is a performance optimization, not a memory
policy. In particular, it must not override explicit or auto-selected
CPU/layerwise offload.
"""
if not self._is_single_dit_component(component_name):
return False
return isinstance(self.strategy_for(component_name, module), ResidentStrategy)
def _is_single_dit_component(self, component_name: str) -> bool:
modules = self.pipeline.modules
return (component_name == "transformer" and "transformer_2" not in modules) or (
component_name == "video_dit" and "video_dit_2" not in modules
@@ -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"],