[diffusion] feat: support dynamically cpu offload components (#34391)
This commit is contained in:
@@ -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
|
||||
|
||||
+10
-3
@@ -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,
|
||||
)
|
||||
|
||||
+3
-7
@@ -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)
|
||||
|
||||
+13
-3
@@ -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
|
||||
|
||||
+30
-34
@@ -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
|
||||
|
||||
+50
@@ -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
|
||||
|
||||
|
||||
-1
@@ -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"],
|
||||
|
||||
Reference in New Issue
Block a user