diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index ddb186cb1..2e9cf9e4e 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -31,8 +31,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( LAYERWISE_OFFLOAD_ALL_COMPONENTS, - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS, - layerwise_component_matches_selection, + LAYERWISE_OFFLOAD_DIT_GROUP, + layerwise_component_matches_any_selection, + normalize_layerwise_offload_components, ) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -115,18 +116,19 @@ class ComponentLoader(ABC): server_args: ServerArgs, component_name: str ) -> bool: """if a component should be loaded in a layerwise-fashion""" - selected_component_names = server_args.layerwise_offload_components + selected_component_names = normalize_layerwise_offload_components( + server_args.layerwise_offload_components + ) if selected_component_names is None: return False selected_component_names = set(selected_component_names) if LAYERWISE_OFFLOAD_ALL_COMPONENTS in selected_component_names: return True explicit_component_names = selected_component_names - { - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS + LAYERWISE_OFFLOAD_DIT_GROUP } - return any( - layerwise_component_matches_selection(component_name, selected_component) - for selected_component in explicit_component_names + return layerwise_component_matches_any_selection( + component_name, explicit_component_names ) def _maybe_configure_layerwise_after_startup_cpu_staging( diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index ad6294bb7..a832ca725 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -141,9 +141,6 @@ class TransformerLoader(ComponentLoader): for post_load_hook in quant_spec.post_load_hooks: post_load_hook(model) - total_params = sum(p.numel() for p in model.parameters()) - logger.info("Loaded model with %.2fB parameters", total_params / 1e9) - # considering the existent of mixed-precision models (e.g., nunchaku) if ( next(model.parameters()).dtype != quant_spec.param_dtype diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index b07f210ee..ee4f1cd07 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -170,6 +170,12 @@ class GPUWorker: self.pipeline.modules, self.server_args, component_names=self.server_args.layerwise_offload_components, + warn_missing=( + self.server_args.is_arg_explicitly_set( + "layerwise_offload_components" + ) + or self.server_args.is_arg_explicitly_set("dit_layerwise_offload") + ), ) logger.info( diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py index 646e04568..99db87030 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py @@ -6,8 +6,10 @@ import torch from torch.distributed.tensor import DTensor from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS, - layerwise_component_matches_selection, + LAYERWISE_OFFLOAD_ALL_COMPONENTS, + LAYERWISE_OFFLOAD_DIT_GROUP, + layerwise_component_matches_any_selection, + normalize_layerwise_offload_components, ) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -513,9 +515,11 @@ class LayerwiseOffloadManager: def make_pre_hook(i): def hook(module, input): - # wait only for the current layer if it's being prefetched if i == 0: self.prepare_for_next_req(non_blocking=False) + if i not in self._gpu_layers: + # LTX audio VAE traverses decoder.up in reverse order + self.prefetch_layer(i, non_blocking=False) if i in self._prefetch_events: torch.get_device_module().current_stream().wait_event( self._prefetch_events[i] @@ -554,8 +558,9 @@ class LayerwiseOffloadManager: class LayerwiseOffloadableModuleMixin: """A mixin that registers forward hooks to enable layerwise offload.""" - # Legacy --dit-layerwise-offload configures these modules when no component is named. - layerwise_offload_default_enabled: bool = True + # whether the current module is selected by the `dit` group + layerwise_offload_dit_group_enabled: bool = True + # The list of names of this module's layer/block ModuleList or Sequential attributes. layer_names: List[str] = [] layerwise_offload_managers: list[LayerwiseOffloadManager] = [] @@ -664,6 +669,54 @@ def is_layerwise_offloaded_module(module: torch.nn.Module) -> bool: ) +def get_layerwise_offload_component_names_for_pipeline( + modules: Mapping[str, object], + component_names: Sequence[str] | None = None, +) -> list[str]: + """Resolve layerwise selectors against the current pipeline modules. + + Explicit unsupported component names are kept so callers can report them. + """ + normalized_component_names = normalize_layerwise_offload_components(component_names) + selected_component_names = ( + set(normalized_component_names) + if normalized_component_names is not None + else None + ) + + if selected_component_names is None: + return [ + component_name + for component_name, module in modules.items() + if isinstance(module, LayerwiseOffloadableModuleMixin) + and module.layerwise_offload_dit_group_enabled + ] + + if LAYERWISE_OFFLOAD_ALL_COMPONENTS in selected_component_names: + return [ + component_name + for component_name, module in modules.items() + if isinstance(module, LayerwiseOffloadableModuleMixin) + ] + + explicit_component_names = selected_component_names - {LAYERWISE_OFFLOAD_DIT_GROUP} + select_dit_group = LAYERWISE_OFFLOAD_DIT_GROUP in selected_component_names + selected_pipeline_component_names: list[str] = [] + for component_name, module in modules.items(): + if layerwise_component_matches_any_selection( + component_name, explicit_component_names + ): + selected_pipeline_component_names.append(component_name) + continue + if ( + select_dit_group + and isinstance(module, LayerwiseOffloadableModuleMixin) + and module.layerwise_offload_dit_group_enabled + ): + selected_pipeline_component_names.append(component_name) + return selected_pipeline_component_names + + def configure_layerwise_offload_modules( modules: Mapping[str, object], server_args: ServerArgs, @@ -682,27 +735,33 @@ def configure_layerwise_offload_modules( # components which has already been configured to be layerwise-offload configured_component_names: list[str] = [] configured_module_ids: set[int] = set() + normalized_component_names = normalize_layerwise_offload_components(component_names) selected_component_names = ( - set(component_names) if component_names is not None else None + set(normalized_component_names) + if normalized_component_names is not None + else None ) select_all = ( - selected_component_names is not None and "all" in selected_component_names - ) - select_default = ( selected_component_names is not None - and LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS in selected_component_names + and LAYERWISE_OFFLOAD_ALL_COMPONENTS in selected_component_names + ) + selected_pipeline_component_names = ( + get_layerwise_offload_component_names_for_pipeline( + modules, + normalized_component_names, + ) ) if warn_missing and selected_component_names is not None and not select_all: explicit_component_names = selected_component_names - { - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS + LAYERWISE_OFFLOAD_DIT_GROUP } missing_component_names = [ selected_component_name for selected_component_name in explicit_component_names if not any( - layerwise_component_matches_selection( - component_name, selected_component_name + layerwise_component_matches_any_selection( + component_name, [selected_component_name] ) for component_name in modules ) @@ -717,13 +776,7 @@ def configure_layerwise_offload_modules( unsupported_component_names = [ component_name - for component_name in modules - if any( - layerwise_component_matches_selection( - component_name, selected_component_name - ) - for selected_component_name in explicit_component_names - ) + for component_name in selected_pipeline_component_names if not isinstance(modules[component_name], LayerwiseOffloadableModuleMixin) ] if unsupported_component_names: @@ -732,27 +785,13 @@ def configure_layerwise_offload_modules( sorted(unsupported_component_names), ) - for component_name, module in modules.items(): + for component_name in selected_pipeline_component_names: + module = modules[component_name] if not isinstance(module, LayerwiseOffloadableModuleMixin): continue - if selected_component_names is None: - if not module.layerwise_offload_default_enabled: - continue - elif ( - not select_all - and not any( - layerwise_component_matches_selection( - component_name, selected_component_name - ) - for selected_component_name in selected_component_names - ) - and not (select_default and module.layerwise_offload_default_enabled) - ): - # if the current component is not selected to be layerwise-offload, skip - continue module_id = id(module) if module_id in configured_module_ids: - # avoid multiple configures on a same module + # avoid duplicated configures on a same module continue configured_module_ids.add(module_id) diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py index 48297d547..71e3bf54e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py @@ -1,7 +1,19 @@ -from collections.abc import Sequence +from collections.abc import Collection, Sequence LAYERWISE_OFFLOAD_ALL_COMPONENTS = "all" -LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS = "default" +LAYERWISE_OFFLOAD_DIT_GROUP = "dit" +LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP = "text_encoder" +LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP = "image_encoder" +LAYERWISE_OFFLOAD_VAE_GROUP = "vae" +LAYERWISE_OFFLOAD_DEFAULT_GROUP = "default" + +# Components whose layerwise policy has been validated as a better default than +# component-level CPU offload when the user has not pinned their placement. +LAYERWISE_OFFLOAD_DEFAULT_GROUP_COMPONENTS = ( + LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP, + LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP, + LAYERWISE_OFFLOAD_VAE_GROUP, +) DIT_COMPONENT_NAMES = frozenset( { "transformer", @@ -22,6 +34,13 @@ VAE_COMPONENT_NAMES = frozenset( "condition_image_encoder", } ) +DEFAULT_LAYERWISE_VAE_COMPONENT_NAMES = frozenset( + { + "vae", + "video_vae", + "condition_image_encoder", + } +) CPU_OFFLOAD_FLAG_NAMES = ( "dit_cpu_offload", "text_encoder_cpu_offload", @@ -53,25 +72,37 @@ def layerwise_component_matches_selection( selected_component_name: str, ) -> bool: """if the provided component_name (unnormalized, e.g., text_encoder_2) matches with the selected_component_name (normalized)""" - if selected_component_name == "text_encoder": + if selected_component_name == LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP: return is_text_encoder_component_name(component_name) - if selected_component_name == "vae": - return is_vae_component_name(component_name) + if selected_component_name == LAYERWISE_OFFLOAD_VAE_GROUP: + # `vae` is a default-policy selector; AV-side decoders remain explicit-only + return component_name in DEFAULT_LAYERWISE_VAE_COMPONENT_NAMES return component_name == selected_component_name +def layerwise_component_matches_any_selection( + component_name: str, + selected_component_names: Collection[str], +) -> bool: + return any( + layerwise_component_matches_selection(component_name, selected_component_name) + for selected_component_name in selected_component_names + ) + + def cpu_offload_flags_for_layerwise_components( component_names: Sequence[str], ) -> tuple[str, ...]: + component_names = normalize_layerwise_offload_components(component_names) or [] if LAYERWISE_OFFLOAD_ALL_COMPONENTS in component_names: return CPU_OFFLOAD_FLAG_NAMES flag_names: list[str] = [] - if LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS in component_names: + if LAYERWISE_OFFLOAD_DIT_GROUP in component_names: flag_names.append("dit_cpu_offload") for component_name in component_names: - if component_name == LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS: + if component_name == LAYERWISE_OFFLOAD_DIT_GROUP: continue if is_dit_component_name(component_name): flag_name = "dit_cpu_offload" @@ -90,6 +121,12 @@ def cpu_offload_flags_for_layerwise_components( return tuple(flag_names) +def expand_layerwise_offload_component_group(component_name: str) -> tuple[str, ...]: + if component_name == LAYERWISE_OFFLOAD_DEFAULT_GROUP: + return LAYERWISE_OFFLOAD_DEFAULT_GROUP_COMPONENTS + return (component_name,) + + def normalize_layerwise_offload_components( component_names: str | Sequence[str] | None, ) -> list[str] | None: @@ -109,9 +146,12 @@ def normalize_layerwise_offload_components( component_name = component_name.strip().replace("-", "_").lower() if not component_name: continue - if component_name == LAYERWISE_OFFLOAD_ALL_COMPONENTS: - return [LAYERWISE_OFFLOAD_ALL_COMPONENTS] - if component_name not in normalized_components: - normalized_components.append(component_name) + for expanded_component_name in expand_layerwise_offload_component_group( + component_name + ): + if expanded_component_name == LAYERWISE_OFFLOAD_ALL_COMPONENTS: + return [LAYERWISE_OFFLOAD_ALL_COMPONENTS] + if expanded_component_name not in normalized_components: + normalized_components.append(expanded_component_name) return normalized_components or None diff --git a/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py b/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py index 4277bd24b..3607dfc68 100644 --- a/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py +++ b/python/sglang/multimodal_gen/runtime/models/bridges/mova_dual_tower.py @@ -409,7 +409,7 @@ class DualTowerConditionalBridge( 3. Cross-attention interaction between the hidden states of the two DiTs. """ - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False _fsdp_shard_conditions = MOVADualTowerConfig()._fsdp_shard_conditions _compile_conditions = MOVADualTowerConfig()._compile_conditions diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/base.py b/python/sglang/multimodal_gen/runtime/models/encoders/base.py index ebaa924ec..70d2cb906 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/base.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/base.py @@ -19,7 +19,7 @@ from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "layers", "encoder.block", @@ -60,7 +60,7 @@ class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin): class ImageEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "layers", "vision_model.encoder.layers", diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py b/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py index 2d0eaa27a..a3bbce1af 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py @@ -287,7 +287,7 @@ class Gemma2Model(nn.Module, LayerwiseOffloadableModuleMixin): """Gemma2 text encoder model for SANA pipeline.""" _fsdp_shard_conditions = [] - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["layers"] def __init__(self, config: Gemma2Config, **kwargs): diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py b/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py index 0bc0ae964..e9c639d83 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/gemma_3.py @@ -941,7 +941,7 @@ class Gemma3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin) # transformers 5.6.0 flattened SiglipVisionModel, dropping the # `vision_model` intermediate wrapper. Our reimpl keeps it, so remap # HF source keys back into our nested namespace when transferring weights. - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["language_model.layers"] param_names_mapping = { diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/hunyuan3d.py b/python/sglang/multimodal_gen/runtime/models/encoders/hunyuan3d.py index 2f4347009..dc355f46e 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/hunyuan3d.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/hunyuan3d.py @@ -33,7 +33,7 @@ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): class ImageEncoder(nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "model.encoder.layer", "model.vision_model.encoder.layers", @@ -213,7 +213,7 @@ def build_image_encoder(config): class DualImageEncoder(nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "main_image_encoder.model.encoder.layer", "main_image_encoder.model.vision_model.encoder.layers", @@ -250,7 +250,7 @@ class DualImageEncoder(nn.Module, LayerwiseOffloadableModuleMixin): class SingleImageEncoder(nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "main_image_encoder.model.encoder.layer", "main_image_encoder.model.vision_model.encoder.layers", diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py index b6b5066cc..72a7d9587 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py @@ -381,7 +381,7 @@ class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixi } _tied_weights_keys = ["lm_head.weight"] uses_sglang_forward_context = False - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["model.language_model.layers"] def __init__(self, config: LlavaConfig): diff --git a/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py b/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py index ab1cbd693..35573b478 100644 --- a/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py +++ b/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py @@ -165,7 +165,7 @@ class LatentUpsampler(torch.nn.Module, LayerwiseOffloadableModuleMixin): rational_resampler: Whether to use rational resampler for spatial upsampling. """ - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["res_blocks", "post_upsample_res_blocks"] def __init__( diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py index ec028423d..ecb780458 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py @@ -62,7 +62,7 @@ class AutoencoderKL(nn.Module, LayerwiseOffloadableModuleMixin): mid_block will only have resnet blocks """ - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False _supports_gradient_checkpointing = True _no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"] layer_names = ["encoder.down_blocks", "decoder.up_blocks"] diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py index ac913e62d..8259a5b7c 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py @@ -17,7 +17,7 @@ logger = init_logger(__name__) class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin): """Deep Compression Autoencoder wrapper with 32x spatial compression.""" - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["_inner_model.encoder.down_blocks", "_inner_model.decoder.up_blocks"] def __init__(self, config: SanaVAEConfig = None, **kwargs): diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/common.py b/python/sglang/multimodal_gen/runtime/models/vaes/common.py index 57fc5b3b3..d83d417d1 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/common.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/common.py @@ -24,7 +24,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im class ParallelTiledVAE(ABC, nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "encoder.down_blocks", "decoder.up_blocks", diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/dac.py b/python/sglang/multimodal_gen/runtime/models/vaes/dac.py index cd3797f2b..ee4750f3d 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/dac.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/dac.py @@ -417,7 +417,7 @@ class Decoder(nn.Module): class DAC(nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["encoder.block", "decoder.model"] def __init__( diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuan3d_vae.py b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuan3d_vae.py index 692998f16..c9a54fb63 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuan3d_vae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuan3d_vae.py @@ -1105,7 +1105,7 @@ SurfaceExtractors = { class VectsetVAE(nn.Module, LayerwiseOffloadableModuleMixin): """Base VAE class for vector set encoding.""" - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["transformer.resblocks"] def __init__(self, volume_decoder=None, surface_extractor=None): diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_3_condition_encoder.py b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_3_condition_encoder.py index cba5764b9..75871f10a 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_3_condition_encoder.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_3_condition_encoder.py @@ -114,7 +114,7 @@ def _make_ltx23_encoder_block( class LTX23VideoConditionEncoder(nn.Module, LayerwiseOffloadableModuleMixin): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = ["down_blocks"] def __init__(self, config: dict[str, Any]) -> None: diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_audio.py b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_audio.py index bf1613e5a..0c7097225 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_audio.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_audio.py @@ -851,6 +851,19 @@ class AutoencoderKLLTX2Audio(ParallelTiledVAE): # TODO: confirm whether the mel compression ratio below is correct self.mel_compression_ratio = LATENT_DOWNSAMPLE_FACTOR self.use_slicing = False + # stage containers are not called directly, so hooks attach to called lists + self.layer_names = [ + layer_name + for prefix, num_resolutions in ( + ("encoder.down", self.encoder.num_resolutions), + ("decoder.up", self.decoder.num_resolutions), + ) + for level in range(num_resolutions) + for layer_name in ( + f"{prefix}.{level}.block", + f"{prefix}.{level}.attn", + ) + ] def _encode(self, x: torch.Tensor) -> torch.Tensor: return self.encoder(x) diff --git a/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py b/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py index 48985e8bd..6432db263 100644 --- a/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py +++ b/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py @@ -539,7 +539,7 @@ class LTX2Vocoder(ABC, nn.Module, LayerwiseOffloadableModuleMixin): LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms. """ - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False layer_names = [ "upsamplers", "resnets", diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py index dedacb7d7..a309f304a 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py @@ -490,7 +490,9 @@ class LTX2ImageEncodingStage(PipelineStage): safetensors_load_file(weights_path), strict=True ) self._condition_image_encoder_dir = encoder_dir - if server_args.should_configure_layerwise_offload_for_lazy_component(): + if server_args.should_configure_layerwise_offload_for_lazy_component( + "condition_image_encoder" + ): modules = {"condition_image_encoder": self._condition_image_encoder} configure_layerwise_offload_modules( modules, diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 8479ad6ae..f747a0eaf 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -38,8 +38,10 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i ) 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_DEFAULT_COMPONENTS, + LAYERWISE_OFFLOAD_ALL_COMPONENTS, + LAYERWISE_OFFLOAD_DIT_GROUP, cpu_offload_flags_for_layerwise_components, + layerwise_component_matches_any_selection, normalize_layerwise_offload_components, ) from sglang.multimodal_gen.runtime.platforms import ( @@ -199,7 +201,7 @@ class ServerArgs(DisaggArgsMixin): # CPU offload parameters dit_cpu_offload: bool | None = None - # if true, add the legacy default DiT components + # if true, select the DiT layerwise group dit_layerwise_offload: bool | None = None layerwise_offload_components: list[str] | None = None dit_offload_prefetch_size: float = 0.0 @@ -209,6 +211,7 @@ class ServerArgs(DisaggArgsMixin): use_fsdp_inference: bool | None = None pin_cpu_memory: bool = True ltx2_two_stage_device_mode: str | None = None + _explicit_arg_names: set[str] = field(default_factory=set, repr=False) # ComfyUI integration comfyui_mode: bool = False @@ -831,14 +834,28 @@ class ServerArgs(DisaggArgsMixin): self.dit_layerwise_offload = False self.layerwise_offload_components = None - def should_configure_layerwise_offload_for_lazy_component(self) -> bool: + def is_arg_explicitly_set(self, arg_name: str) -> bool: + return arg_name in self._explicit_arg_names + + def should_configure_layerwise_offload_for_lazy_component( + self, component_name: str + ) -> bool: """Return whether a lazy-loaded component should try layerwise offload. Lazy components are loaded after the normal pipeline-wide configuration - pass, so they should only attempt layerwise configuration when a - component scope is present. + pass, so they should only attempt layerwise configuration when their + component name is covered by the selected layerwise scope. """ - return bool(self.layerwise_offload_components) + component_names = normalize_layerwise_offload_components( + self.layerwise_offload_components + ) + if not component_names: + return False + if LAYERWISE_OFFLOAD_ALL_COMPONENTS in component_names: + return True + return layerwise_component_matches_any_selection( + component_name, component_names + ) @property def is_dit_layerwise_offload_selected(self) -> bool: @@ -856,13 +873,10 @@ class ServerArgs(DisaggArgsMixin): ) if self.dit_layerwise_offload: if explicitly_set_component_names is None: - explicitly_set_component_names = [LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS] - elif ( - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS - not in explicitly_set_component_names - ): + explicitly_set_component_names = [LAYERWISE_OFFLOAD_DIT_GROUP] + elif LAYERWISE_OFFLOAD_DIT_GROUP not in explicitly_set_component_names: explicitly_set_component_names = [ - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS, + LAYERWISE_OFFLOAD_DIT_GROUP, *explicitly_set_component_names, ] @@ -1197,7 +1211,7 @@ class ServerArgs(DisaggArgsMixin): action=StoreBoolean, default=ServerArgs.dit_layerwise_offload, help="Enable layerwise CPU offload with async H2D prefetch overlap for DiTs. " - "It only selects the legacy default DiT components. Cannot be used together with cache-dit " + "It selects only the DiT layerwise group. Cannot be used together with cache-dit " "(SGLANG_CACHE_DIT_ENABLED), dit_cpu_offload, or use_fsdp_inference.", ) parser.add_argument( @@ -1207,10 +1221,11 @@ class ServerArgs(DisaggArgsMixin): nargs="+", default=ServerArgs.layerwise_offload_components, help="Select pipeline components for layerwise offload. " - "Use default to select the legacy default DiT components, " + "Use dit to select the DiT layerwise group, default for the default group " + "(currently text_encoder, image_encoder, and vae), " "or all to select every layerwise-offloadable component. " "This option does not imply --dit-layerwise-offload. Example: " - "--layerwise-offload-components text_encoder image_encoder.", + "--layerwise-offload-components text_encoder image_encoder vae.", ) parser.add_argument( "--dit-offload-prefetch-size", @@ -1647,6 +1662,7 @@ class ServerArgs(DisaggArgsMixin): component_paths = dict(kwargs.get("component_paths") or {}) if component_paths: server_args_kwargs["component_paths"] = component_paths + server_args_kwargs["_explicit_arg_names"] = set(kwargs) for attr in attrs: if attr == "pipeline_config": @@ -1682,6 +1698,8 @@ class ServerArgs(DisaggArgsMixin): @classmethod def from_kwargs(cls, **kwargs: Any) -> "ServerArgs": + explicit_arg_names = set(kwargs) + # Convert backend string to enum if necessary if "backend" in kwargs and isinstance(kwargs["backend"], str): kwargs["backend"] = Backend.from_string(kwargs["backend"]) @@ -1690,6 +1708,7 @@ class ServerArgs(DisaggArgsMixin): convert_disagg_role_string(kwargs) kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs) + kwargs["_explicit_arg_names"] = explicit_arg_names return cls(**kwargs) @staticmethod @@ -1712,6 +1731,8 @@ class ServerArgs(DisaggArgsMixin): provided_arg_names.add(arg_name) if "mode" in provided_arg_names: provided_arg_names.add("performance_mode") + if "layerwise_offload_modules" in provided_arg_names: + provided_arg_names.add("layerwise_offload_components") # Populate provided_args if the argument from the namespace was on the command line. for k, v in vars(args).items(): @@ -1749,13 +1770,13 @@ class ServerArgs(DisaggArgsMixin): if self.dit_offload_prefetch_size < 0.0: raise ValueError("dit_offload_prefetch_size must be non-negative") - if self.use_fsdp_inference: + should_disable_dit_cpu_offload = self.is_dit_layerwise_offload_selected + if self.use_fsdp_inference and should_disable_dit_cpu_offload: logger.warning( - "layerwise offload components are selected, automatically disabling use_fsdp_inference." + "layerwise offload is selected for DiT components, automatically disabling use_fsdp_inference." ) self.use_fsdp_inference = False - should_disable_dit_cpu_offload = self.is_dit_layerwise_offload_selected if should_disable_dit_cpu_offload and self.dit_cpu_offload is not False: logger.warning( "layerwise offload is selected for DiT components, automatically disabling dit_cpu_offload." diff --git a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py index d9e3811d9..21a647c65 100644 --- a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py @@ -11,7 +11,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config impo ModelDeploymentConfig, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( - LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS, + LAYERWISE_OFFLOAD_DIT_GROUP, + LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP, + LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP, + LAYERWISE_OFFLOAD_VAE_GROUP, ) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -23,6 +26,12 @@ logger = init_logger(__name__) PERFORMANCE_MODES = ("manual", "auto", "speed", "memory") +DEFAULT_LAYERWISE_COMPONENT_ARG_NAMES = ( + (LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP, "text_encoder_cpu_offload"), + (LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP, "image_encoder_cpu_offload"), + (LAYERWISE_OFFLOAD_VAE_GROUP, "vae_cpu_offload"), +) + class ServerArgsAutoTuner: """Auto-tunes the server-arg for the given performance-mode, based on practical deployment experience with different model architectures""" @@ -60,6 +69,13 @@ class ServerArgsAutoTuner: logger.info("Applying performance_mode=memory") if args.use_fsdp_inference: self._set_gpu_resident_defaults(use_fsdp=True) + if ( + args.layerwise_offload_components is None + and self._can_apply_default_layerwise_offload_policy() + ): + args.layerwise_offload_components = ( + self._default_layerwise_components_for_unset_placement() or None + ) return args.use_fsdp_inference = False if self._can_apply_default_layerwise_offload_policy(): @@ -96,15 +112,34 @@ class ServerArgsAutoTuner: components = tuple( component for component in components if component != "dit" ) - if args.dit_cpu_offload and "dit" in components: + if ( + args.dit_cpu_offload + and "dit" in components + and not args.is_arg_explicitly_set("dit_cpu_offload") + ): args.dit_cpu_offload = False changed.append("dit_cpu_offload=False") - if args.text_encoder_cpu_offload and "text_encoder" in components: + if ( + args.text_encoder_cpu_offload + and LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP in components + and not args.is_arg_explicitly_set("text_encoder_cpu_offload") + ): args.text_encoder_cpu_offload = False changed.append("text_encoder_cpu_offload=False") - if args.image_encoder_cpu_offload and "image_encoder" in components: + if ( + args.image_encoder_cpu_offload + and LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP in components + and not args.is_arg_explicitly_set("image_encoder_cpu_offload") + ): args.image_encoder_cpu_offload = False changed.append("image_encoder_cpu_offload=False") + if ( + args.vae_cpu_offload + and LAYERWISE_OFFLOAD_VAE_GROUP in components + and not args.is_arg_explicitly_set("vae_cpu_offload") + ): + args.vae_cpu_offload = False + changed.append("vae_cpu_offload=False") if changed: logger.info( "Disabling component offload for %s because minimum available memory on selected GPUs is %.2f GiB: %s", @@ -134,56 +169,30 @@ class ServerArgsAutoTuner: self._enable_cfg_parallel_if_supported() def maybe_adjust_auto_default_layerwise_offload(self) -> None: - """adjust the default layerwise offload policy""" + """Enable verified non-DiT layerwise defaults for unset component placement.""" args = self.server_args + if args.performance_mode != "auto": + return if not self.could_override_server_args(): return - if self._explicit_memory_policy: - return - deployment_config = self._deployment_config() - if envs.SGLANG_CACHE_DIT_ENABLED: - return if ( - not deployment_config.auto_dit_layerwise_offload - or args.dit_layerwise_offload is not None + args.layerwise_offload_components is not None + or args.dit_layerwise_offload is True ): return - if args.use_fsdp_inference: - # if fsdp is enabled, layerwise-offload is weakened since the parameter has already been sharded - args.dit_layerwise_offload = False + if not current_platform.is_cuda(): return - auto_enable_layerwise_offload = ( - current_platform.enable_dit_layerwise_offload_for_wan_by_default() - ) - disable_threshold_gb = ( - deployment_config.auto_dit_layerwise_offload_high_memory_disable_gb - ) - if ( - auto_enable_layerwise_offload - and current_platform.is_cuda() - and disable_threshold_gb is not None - ): - # auto turn off layerwise-offload if we have sufficient VRAM headroom - device_total_memory_gb = current_platform.get_device_total_memory() / ( - 1 << 30 - ) - if device_total_memory_gb >= disable_threshold_gb: - logger.info( - "Skipping automatic dit_layerwise_offload for %s on a high-memory CUDA GPU (e.g. H200/B200/B300-class, %.2f GiB total)", - args.pipeline_config.__class__.__name__, - device_total_memory_gb, - ) - auto_enable_layerwise_offload = False - args.dit_layerwise_offload = False + layerwise_components = self._default_layerwise_components_for_unset_placement() + if not layerwise_components: + return - if auto_enable_layerwise_offload: - logger.info( - "Automatically enable dit_layerwise_offload for %s for low memory and performance balance", - args.pipeline_config.__class__.__name__, - ) - args.dit_layerwise_offload = True - args.dit_cpu_offload = False + logger.info( + "Automatically enable default non-DiT layerwise offload for %s: %s", + args.pipeline_config.__class__.__name__, + layerwise_components, + ) + args.layerwise_offload_components = layerwise_components def maybe_replace_cpu_offloaded_components_with_layerwise(self) -> None: args = self.server_args @@ -200,18 +209,22 @@ class ServerArgsAutoTuner: layerwise_components: list[str] = [] if args.dit_layerwise_offload: - layerwise_components.append(LAYERWISE_OFFLOAD_DEFAULT_COMPONENTS) + layerwise_components.append(LAYERWISE_OFFLOAD_DIT_GROUP) changed: list[str] = [] - if args.text_encoder_cpu_offload: - layerwise_components.append("text_encoder") - changed.append("text_encoder") - if args.image_encoder_cpu_offload: - layerwise_components.append("image_encoder") - changed.append("image_encoder") - if args.vae_cpu_offload: - layerwise_components.append("vae") - changed.append("vae") + if args.text_encoder_cpu_offload and not args.is_arg_explicitly_set( + "text_encoder_cpu_offload" + ): + layerwise_components.append(LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP) + changed.append(LAYERWISE_OFFLOAD_TEXT_ENCODER_GROUP) + if args.image_encoder_cpu_offload and not args.is_arg_explicitly_set( + "image_encoder_cpu_offload" + ): + layerwise_components.append(LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP) + changed.append(LAYERWISE_OFFLOAD_IMAGE_ENCODER_GROUP) + if args.vae_cpu_offload and not args.is_arg_explicitly_set("vae_cpu_offload"): + layerwise_components.append(LAYERWISE_OFFLOAD_VAE_GROUP) + changed.append(LAYERWISE_OFFLOAD_VAE_GROUP) if not changed: return @@ -301,25 +314,46 @@ class ServerArgsAutoTuner: def _set_layerwise_offload_defaults(self) -> None: args = self.server_args - if args.dit_layerwise_offload is None: - args.dit_layerwise_offload = True + if args.layerwise_offload_components is None: + args.layerwise_offload_components = ( + self._default_layerwise_components_for_unset_placement() or None + ) if args.dit_cpu_offload is None: - args.dit_cpu_offload = False + args.dit_cpu_offload = True if args.text_encoder_cpu_offload is None: - args.text_encoder_cpu_offload = True + args.text_encoder_cpu_offload = False if args.image_encoder_cpu_offload is None: - args.image_encoder_cpu_offload = True + args.image_encoder_cpu_offload = False def _can_apply_default_layerwise_offload_policy(self) -> bool: - return ( - self._deployment_config().auto_dit_layerwise_offload - and not envs.SGLANG_CACHE_DIT_ENABLED - and current_platform.enable_dit_layerwise_offload_for_wan_by_default() - ) + return current_platform.is_cuda() + + def _default_layerwise_components_for_unset_placement(self) -> list[str]: + args = self.server_args + if ( + args.is_arg_explicitly_set("layerwise_offload_components") + or args.dit_layerwise_offload is True + ): + # The legacy --dit-layerwise-offload flag is a DiT-only selector. + # Do not merge implicit non-DiT defaults into that explicit mode. + return [] + + # `*_cpu_offload` is the component placement knob. If a user explicitly + # set it to either true or false, keep that component out of default + # layerwise selection. + return [ + component_name + for component_name, arg_name in DEFAULT_LAYERWISE_COMPONENT_ARG_NAMES + if not args.is_arg_explicitly_set(arg_name) + ] def _auto_uses_dit_offload(self) -> bool: args = self.server_args - return bool(args.dit_cpu_offload or args.dit_layerwise_offload) + return bool( + args.dit_cpu_offload + or args.dit_layerwise_offload + or args.is_dit_layerwise_offload_selected + ) def _get_min_available_device_memory_gb(self) -> float | None: args = self.server_args @@ -339,23 +373,24 @@ class ServerArgsAutoTuner: def _has_explicit_memory_policy(self) -> bool: args = self.server_args - return ( - args.use_fsdp_inference is not None - or args.dit_cpu_offload is not None - or args.dit_layerwise_offload is not None - or args.layerwise_offload_components is not None - or args.text_encoder_cpu_offload is not None - or args.image_encoder_cpu_offload is not None + return any( + args.is_arg_explicitly_set(arg_name) + for arg_name in ( + "use_fsdp_inference", + "dit_cpu_offload", + "dit_layerwise_offload", + "layerwise_offload_components", + ) ) def _has_explicit_layerwise_replacement_policy(self) -> bool: args = self.server_args - return ( - args.dit_layerwise_offload is not None - or args.layerwise_offload_components is not None - or args.text_encoder_cpu_offload is not None - or args.image_encoder_cpu_offload is not None - or args.vae_cpu_offload is True + return any( + args.is_arg_explicitly_set(arg_name) + for arg_name in ( + "dit_layerwise_offload", + "layerwise_offload_components", + ) ) def _has_explicit_parallel_policy(self) -> bool: diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 8463871dd..cf224c3a8 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -37,7 +37,7 @@ "scenarios": { "qwen_image_t2i": { "stages_ms": { - "TextEncodingStage": 588.48, + "TextEncodingStage": 520.0, "DenoisingStage": 12404.23, "InputValidationStage": 0.05, "LatentPreparationStage": 0.23, @@ -104,7 +104,7 @@ "qwen_image_t2i_2_gpus": { "stages_ms": { "InputValidationStage": 0.13, - "TextEncodingStage": 1114.31, + "TextEncodingStage": 1050.0, "LatentPreparationStage": 0.26, "TimestepPreparationStage": 20.44, "DenoisingStage": 10159.11, @@ -239,7 +239,7 @@ "DenoisingStage": 24036.73, "DecodingStage": 11.76, "LatentPreparationStage": 1.17, - "TextEncodingStage": 500.54, + "TextEncodingStage": 430.0, "InputValidationStage": 0.05, "ImageVAEEncodingStage": 0.01 }, @@ -389,7 +389,7 @@ }, "flux_2_ti2i": { "stages_ms": { - "TextEncodingStage": 500.3, + "TextEncodingStage": 430.0, "DenoisingStage": 47133.15, "InputValidationStage": 43.89, "LatentPreparationStage": 1.19, @@ -592,7 +592,7 @@ "DecodingStage": 8.86, "InputValidationStage": 0.05, "DenoisingStage": 675.8, - "TextEncodingStage": 173.34, + "TextEncodingStage": 155.0, "LatentPreparationStage": 0.14, "TimestepPreparationStage": 36.26 }, @@ -614,7 +614,7 @@ }, "zimage_image_t2i_fp8": { "stages_ms": { - "TextEncodingStage": 176.72, + "TextEncodingStage": 155.0, "DenoisingStage": 634.42, "InputValidationStage": 0.04, "LatentPreparationStage": 0.11, @@ -643,7 +643,7 @@ "DenoisingStage": 673.95, "DecodingStage": 8.43, "LatentPreparationStage": 0.11, - "TextEncodingStage": 175.95, + "TextEncodingStage": 155.0, "InputValidationStage": 0.04 }, "denoise_step_ms": { @@ -812,7 +812,7 @@ "qwen_image_t2i_cache_dit_enabled": { "stages_ms": { "InputValidationStage": 0.06, - "TextEncodingStage": 633.32, + "TextEncodingStage": 540.0, "LatentPreparationStage": 0.24, "TimestepPreparationStage": 17.53, "DenoisingStage": 4279.45, @@ -870,7 +870,7 @@ "48": 189.41, "49": 177.01 }, - "expected_e2e_ms": 4956.71, + "expected_e2e_ms": 4800.0, "expected_avg_denoise_ms": 85.38, "expected_median_denoise_ms": 57.6, "estimated_full_test_time_s": 124.9 @@ -1014,7 +1014,7 @@ "DenoisingStage": 4934.97, "DecodingStage": 631.04, "TimestepPreparationStage": 3.66, - "TextEncodingStage": 1707.94 + "TextEncodingStage": 1450.0 }, "denoise_step_ms": { "0": 58.84, @@ -1068,7 +1068,7 @@ "48": 91.15, "49": 91.1 }, - "expected_e2e_ms": 7617.99, + "expected_e2e_ms": 6800.0, "expected_avg_denoise_ms": 98.56, "expected_median_denoise_ms": 100.05, "estimated_full_test_time_s": 127.6 @@ -1163,12 +1163,12 @@ }, "wan2_2_ti2v_5b": { "stages_ms": { - "InputValidationStage": 23.99, - "TextEncodingStage": 1152.06, + "InputValidationStage": 380.0, + "TextEncodingStage": 750.0, "LatentPreparationStage": 0.13, "TimestepPreparationStage": 2.32, "DenoisingStage": 18728.92, - "DecodingStage": 4110.73 + "DecodingStage": 1550.0 }, "denoise_step_ms": { "0": 225.06, @@ -1222,7 +1222,7 @@ "48": 373.94, "49": 371.58 }, - "expected_e2e_ms": 24821.68, + "expected_e2e_ms": 18500.0, "expected_avg_denoise_ms": 364.69, "expected_median_denoise_ms": 367.69, "estimated_full_test_time_s": 141.7 @@ -1350,19 +1350,19 @@ }, "fastwan2_2_ti2v_5b": { "stages_ms": { - "InputValidationStage": 26.57, - "TextEncodingStage": 1207.93, + "InputValidationStage": 380.0, + "TextEncodingStage": 700.0, "TimestepPreparationStage": 41.12, "LatentPreparationStage": 0.15, "DmdDenoisingStage": 428.28, - "DecodingStage": 2501.21 + "DecodingStage": 1550.0 }, "denoise_step_ms": { "0": 49.59, "1": 171.65, "2": 200.43 }, - "expected_e2e_ms": 5014.57, + "expected_e2e_ms": 3150.0, "expected_avg_denoise_ms": 140.56, "expected_median_denoise_ms": 171.65, "estimated_full_test_time_s": 125.2 @@ -1370,7 +1370,7 @@ "fast_hunyuan_video": { "stages_ms": { "InputValidationStage": 0.06, - "TextEncodingStage": 321.95, + "TextEncodingStage": 300.0, "TimestepPreparationStage": 28.98, "LatentPreparationStage": 0.13, "DenoisingStage": 5898.72, @@ -1586,7 +1586,7 @@ "wan2_2_t2v_a14b_2gpu": { "stages_ms": { "InputValidationStage": 0.05, - "TextEncodingStage": 1012.25, + "TextEncodingStage": 800.0, "LatentPreparationStage": 0.21, "TimestepPreparationStage": 1.89, "DenoisingStage": 82060.15, @@ -1641,7 +1641,7 @@ }, "wan2_1_t2v_14b_2gpu": { "stages_ms": { - "TextEncodingStage": 1693.7, + "TextEncodingStage": 1450.0, "DecodingStage": 637.43, "TimestepPreparationStage": 3.4, "InputValidationStage": 0.05, @@ -1708,7 +1708,7 @@ "wan2_2_t2v_a14b_lora_2gpu": { "stages_ms": { "InputValidationStage": 0.06, - "TextEncodingStage": 1693.53, + "TextEncodingStage": 1500.0, "LatentPreparationStage": 0.15, "TimestepPreparationStage": 4.13, "DenoisingStage": 57638.76, @@ -1898,7 +1898,7 @@ "flux_2_image_t2i_2_gpus": { "stages_ms": { "InputValidationStage": 0.28, - "TextEncodingStage": 952.67, + "TextEncodingStage": 820.0, "ImageVAEEncodingStage": 0.01, "LatentPreparationStage": 0.83, "TimestepPreparationStage": 46.57, @@ -2023,7 +2023,7 @@ "stages_ms": { "InputValidationStage": 0.06, "LatentPreparationStage": 0.16, - "TextEncodingStage": 305.97, + "TextEncodingStage": 280.0, "TimestepPreparationStage": 57.19, "DecodingStage": 16.88, "DenoisingStage": 2422.53 @@ -2179,7 +2179,7 @@ }, "flux_2_image_t2i_upscaling_4x": { "stages_ms": { - "TextEncodingStage": 494.65, + "TextEncodingStage": 430.0, "DenoisingStage": 23822.05, "InputValidationStage": 0.06, "LatentPreparationStage": 1.2, @@ -2541,7 +2541,7 @@ "ltx_2_3_hq_pipeline": { "stages_ms": { "InputValidationStage": 0.09, - "TextEncodingStage": 987.02, + "TextEncodingStage": 900.0, "LTX2TextConnectorStage": 31.22, "LTX2HalveResolutionStage": 0.12, "LTX2LoRASwitchStage": 0.03, diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py index 22c62311f..6c11501af 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -28,6 +28,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im LayerwiseOffloadableModuleMixin, LayerwiseOffloadManager, configure_layerwise_offload_modules, + get_layerwise_offload_component_names_for_pipeline, is_layerwise_offloaded_module, ) @@ -112,8 +113,34 @@ class _SharedBufferModel(torch.nn.Module): ) +class _OrderedLinearLayer(torch.nn.Module): + def __init__(self, scale: float) -> None: + super().__init__() + self.weight = torch.nn.Parameter(torch.eye(2, dtype=torch.float32) * scale) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return x @ self.weight + + +class _ReverseLayerwiseModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList( + [ + _OrderedLinearLayer(2.0), + _OrderedLinearLayer(3.0), + _OrderedLinearLayer(5.0), + ] + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + for block in reversed(self.blocks): + x = block(x) + return x + + class _NestedEncoderDummyModel(_NestedDummyModel): - layerwise_offload_default_enabled = False + layerwise_offload_dit_group_enabled = False class _LayerwiseComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin): @@ -207,6 +234,28 @@ def test_layerwise_offload_keeps_shared_buffers_resident(monkeypatch): assert torch.equal(cache.index_select(0, torch.tensor([2])), original_cache[2:3]) +def test_layerwise_offload_loads_current_layer_for_reverse_execution(monkeypatch): + monkeypatch.setattr( + layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule + ) + monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu") + + model = _ReverseLayerwiseModel() + x = torch.ones(1, 2, dtype=torch.float32) + expected = model(x) + + LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=3, + enabled=True, + pin_cpu_memory=False, + prefetch_size=1, + ) + + assert torch.equal(model(x), expected) + + def test_modelopt_fp8_adapter_keeps_layerwise_offload_enabled(): server_args = SimpleNamespace( dit_cpu_offload=True, @@ -234,7 +283,7 @@ def test_layerwise_capability_selects_layerwise_strategy_for_any_component(): assert isinstance(strategy, LayerwiseOffloadStrategy) -def test_layerwise_configuration_uses_legacy_default_components(monkeypatch): +def test_layerwise_pipeline_selection_uses_dit_group(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule ) @@ -246,8 +295,10 @@ def test_layerwise_configuration_uses_legacy_default_components(monkeypatch): "scheduler": object(), } + selected = get_layerwise_offload_component_names_for_pipeline(modules) configured = configure_layerwise_offload_modules(modules, _server_args()) + assert selected == ["text_encoder", "text_encoder_alias"] assert configured == ["text_encoder"] assert is_layerwise_offloaded_module(layerwise_module) @@ -276,7 +327,7 @@ def test_layerwise_configuration_filters_by_component_name(monkeypatch): assert not is_layerwise_offloaded_module(vae) -def test_layerwise_configuration_default_marker_extends_legacy_defaults(monkeypatch): +def test_layerwise_configuration_default_group_selects_non_dit_defaults(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule ) @@ -284,37 +335,63 @@ def test_layerwise_configuration_default_marker_extends_legacy_defaults(monkeypa text_encoder = _NestedEncoderDummyModel() text_encoder_2 = _NestedEncoderDummyModel() transformer = _NestedDummyModel() + image_encoder = _NestedEncoderDummyModel() vae = _NestedEncoderDummyModel() audio_vae = _NestedEncoderDummyModel() + vocoder = _NestedEncoderDummyModel() + spatial_upsampler = _NestedEncoderDummyModel() condition_image_encoder = _NestedEncoderDummyModel() modules = { "text_encoder": text_encoder, "text_encoder_2": text_encoder_2, "transformer": transformer, + "image_encoder": image_encoder, "vae": vae, "audio_vae": audio_vae, + "vocoder": vocoder, + "spatial_upsampler": spatial_upsampler, "condition_image_encoder": condition_image_encoder, } configured = configure_layerwise_offload_modules( - modules, _server_args(), component_names=["default", "text_encoder", "vae"] + modules, _server_args(), component_names=["default"] ) + assert get_layerwise_offload_component_names_for_pipeline(modules, ["default"]) == [ + "text_encoder", + "text_encoder_2", + "image_encoder", + "vae", + "condition_image_encoder", + ] assert configured == [ "text_encoder", "text_encoder_2", - "transformer", + "image_encoder", "vae", - "audio_vae", "condition_image_encoder", ] assert is_layerwise_offloaded_module(text_encoder) assert is_layerwise_offloaded_module(text_encoder_2) - assert is_layerwise_offloaded_module(transformer) + assert not is_layerwise_offloaded_module(transformer) + assert is_layerwise_offloaded_module(image_encoder) assert is_layerwise_offloaded_module(vae) - assert is_layerwise_offloaded_module(audio_vae) + assert not is_layerwise_offloaded_module(audio_vae) + assert not is_layerwise_offloaded_module(vocoder) + assert not is_layerwise_offloaded_module(spatial_upsampler) assert is_layerwise_offloaded_module(condition_image_encoder) + for component_name, module in ( + ("audio_vae", audio_vae), + ("vocoder", vocoder), + ("spatial_upsampler", spatial_upsampler), + ): + configured = configure_layerwise_offload_modules( + modules, _server_args(), component_names=[component_name] + ) + assert configured == [component_name] + assert is_layerwise_offloaded_module(module) + def test_layerwise_configuration_all_selects_every_capable_component(monkeypatch): monkeypatch.setattr( diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 40dcbb455..c2507ff38 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -162,7 +162,7 @@ class TestServerArgsPathExpansion(unittest.TestCase): args.layerwise_offload_components, ["text_encoder", "transformer"] ) - def test_dit_layerwise_offload_extends_default_components(self): + def test_dit_layerwise_offload_selects_dit_group(self): args = self._from_dict_without_model_resolution( { "model_path": "/data/my-model", @@ -172,7 +172,7 @@ class TestServerArgsPathExpansion(unittest.TestCase): ) self.assertTrue(args.layerwise_offload_components) - self.assertEqual(args.layerwise_offload_components, ["default"]) + self.assertEqual(args.layerwise_offload_components, ["dit"]) def test_dit_layerwise_offload_from_kwargs(self): with patch.object( @@ -185,7 +185,7 @@ class TestServerArgsPathExpansion(unittest.TestCase): ) self.assertTrue(args.layerwise_offload_components) - self.assertEqual(args.layerwise_offload_components, ["default"]) + self.assertEqual(args.layerwise_offload_components, ["dit"]) def test_layerwise_offload_components_normalize_commas(self): args = self._from_dict_without_model_resolution( @@ -201,6 +201,21 @@ class TestServerArgsPathExpansion(unittest.TestCase): args.layerwise_offload_components, ["text_encoder", "transformer"] ) + def test_layerwise_offload_components_normalize_default_group(self): + args = self._from_dict_without_model_resolution( + { + "model_path": "/data/my-model", + "performance_mode": "manual", + } + ) + args.layerwise_offload_components = ["default", "text_encoder"] + args._adjust_layerwise_offload_components() + + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) + def test_dit_layerwise_offload_cli_arg(self): parser = FlexibleArgumentParser() ServerArgs.add_cli_args(parser) @@ -221,7 +236,7 @@ class TestServerArgsPathExpansion(unittest.TestCase): server_args = ServerArgs.from_cli_args(args, unknown_args) self.assertTrue(server_args.layerwise_offload_components) - self.assertEqual(server_args.layerwise_offload_components, ["default"]) + self.assertEqual(server_args.layerwise_offload_components, ["dit"]) def test_layerwise_offload_components_cli_args(self): parser = FlexibleArgumentParser() @@ -340,10 +355,11 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( - args.layerwise_offload_components, ["text_encoder", "image_encoder"] + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], ) - def test_explicit_vae_cpu_offload_true_is_preserved_without_component_selection( + def test_explicit_vae_cpu_offload_true_is_preserved_by_default_layerwise( self, ): args = self._from_dict_with_task_type( @@ -352,7 +368,18 @@ class TestOffloadDefaults(unittest.TestCase): ) self.assertTrue(args.vae_cpu_offload) - self.assertFalse(args.layerwise_offload_components) + self.assertEqual( + args.layerwise_offload_components, ["text_encoder", "image_encoder"] + ) + + def test_explicit_component_resident_is_preserved_by_default_layerwise(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={"text_encoder_cpu_offload": False}, + ) + + self.assertFalse(args.text_encoder_cpu_offload) + self.assertEqual(args.layerwise_offload_components, ["image_encoder", "vae"]) def test_layerwise_components_disable_matching_cpu_offloads(self): args = self._from_dict_with_task_type( @@ -438,7 +465,10 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) def test_auto_ltx_snapshot_keeps_dit_offload_and_replaces_encoder_cpu_offload( self, @@ -460,7 +490,8 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( - args.layerwise_offload_components, ["text_encoder", "image_encoder"] + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], ) def test_auto_wan_layerwise_offload_is_enabled_without_fsdp(self): @@ -471,11 +502,12 @@ 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.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( args.layerwise_offload_components, - ["default", "text_encoder", "image_encoder"], + ["text_encoder", "image_encoder", "vae"], ) def test_memory_wan_layerwise_offload_is_enabled_without_fsdp(self): @@ -486,11 +518,12 @@ 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.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( args.layerwise_offload_components, - ["default", "text_encoder", "image_encoder"], + ["text_encoder", "image_encoder", "vae"], ) def test_auto_wan_layerwise_offload_does_not_disable_explicit_fsdp(self): @@ -504,7 +537,10 @@ class TestOffloadDefaults(unittest.TestCase): }, ) - self.assertFalse(args.layerwise_offload_components) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) self.assertTrue(args.use_fsdp_inference) def test_auto_multi_gpu_wan_uses_layerwise_offload_without_cfg(self): @@ -520,16 +556,16 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) self.assertFalse(args.enable_cfg_parallel) - self.assertFalse(args.dit_cpu_offload) + self.assertTrue(args.dit_cpu_offload) self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( args.layerwise_offload_components, - ["default", "text_encoder", "image_encoder"], + ["text_encoder", "image_encoder", "vae"], ) - def test_explicit_multi_gpu_dit_layerwise_only_selects_default_component(self): + def test_explicit_multi_gpu_dit_layerwise_only_selects_dit_group(self): args = self._from_dict_with_pipeline_config( MOVAPipelineConfig(), kwargs={ @@ -544,7 +580,7 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.layerwise_offload_components) self.assertTrue(args.text_encoder_cpu_offload) self.assertTrue(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["default"]) + self.assertEqual(args.layerwise_offload_components, ["dit"]) def test_auto_multi_gpu_ltx_replaces_component_cpu_offload_with_resident_dit(self): args = self._from_dict_with_pipeline_config( @@ -563,7 +599,8 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( - args.layerwise_offload_components, ["text_encoder", "image_encoder"] + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], ) def test_auto_multi_gpu_qwen_replaces_text_encoder_offload_with_cfg(self): @@ -582,7 +619,10 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) def test_auto_multi_gpu_zimage_base_prefers_fsdp(self): args = self._from_dict_with_pipeline_config( @@ -626,7 +666,10 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) def test_auto_multi_gpu_qwen_skips_fsdp_when_available_memory_is_low(self): args = self._from_dict_with_pipeline_config( @@ -644,7 +687,10 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) def test_auto_multi_gpu_qwen_uses_selected_gpu_min_available_memory(self): args = self._from_dict_with_pipeline_config( @@ -678,7 +724,10 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) - self.assertEqual(args.layerwise_offload_components, ["text_encoder"]) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) def test_speed_mode_single_gpu_disables_offload(self): args = self._from_dict_with_pipeline_config( @@ -722,12 +771,12 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(args.use_fsdp_inference) self.assertTrue(args.layerwise_offload_components) - self.assertFalse(args.dit_cpu_offload) + self.assertTrue(args.dit_cpu_offload) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) self.assertEqual( args.layerwise_offload_components, - ["default", "text_encoder", "image_encoder"], + ["text_encoder", "image_encoder", "vae"], ) def test_memory_mode_preserves_explicit_fsdp(self): @@ -742,7 +791,10 @@ class TestOffloadDefaults(unittest.TestCase): ) self.assertTrue(args.use_fsdp_inference) - self.assertFalse(args.layerwise_offload_components) + self.assertEqual( + args.layerwise_offload_components, + ["text_encoder", "image_encoder", "vae"], + ) self.assertFalse(args.dit_cpu_offload) def test_invalid_performance_mode_raises(self):