From 00a9028e8789df6f928f54eda306b099d1e7c3c6 Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 9 Sep 2026 13:44:54 +0800 Subject: [PATCH] [diffusion] feat: add explicit snapshot-offload component residency (#38535) Co-authored-by: Mick Qian --- docs/docs/sglang-diffusion/api/cli.mdx | 29 +++ .../sglang-diffusion/deployment_cookbook.mdx | 2 +- .../runtime/loader/transformer_load_utils.py | 11 +- .../multimodal_gen/runtime/loader/utils.py | 5 + .../runtime/managers/gpu_worker.py | 4 + .../memory_managers/component_manager.py | 34 +++- .../memory_managers/component_residency.py | 6 + .../component_residency_strategies.py | 43 +++- .../memory_managers/host_memory_budget.py | 10 + .../memory_managers/layerwise_offload.py | 58 ++++-- .../memory_occupation_controller.py | 6 +- .../memory_managers/weight_snapshot.py | 78 ++++++++ .../runtime/pipelines_core/lora/pipeline.py | 11 ++ .../pipelines_core/stages/image_encoding.py | 7 + .../runtime/post_training/weights_updater.py | 4 + .../runtime/server_args/server_args.py | 30 ++- .../test/unit/test_component_residency.py | 149 ++++++++++++++ .../test/unit/test_layerwise_offload.py | 144 ++++++++++++++ .../test/unit/test_server_args.py | 47 +++++ .../test/unit/test_weight_snapshot.py | 186 ++++++++++++++++++ 20 files changed, 835 insertions(+), 29 deletions(-) create mode 100644 python/sglang/multimodal_gen/runtime/managers/memory_managers/weight_snapshot.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_weight_snapshot.py diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index b5bc8c1c9..5cf41faa6 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -313,6 +313,7 @@ The available modes are: - `resident`: keep the complete component on the accelerator. - `component-offload`: keep the complete component on CPU between uses, moving it to the accelerator before each declared use and back to CPU afterward. +- `snapshot-offload`: retain CPU parameter storage while a complete component runs on GPU. At the end of a use, restore the CPU parameters instead of copying GPU weights back. Mutable buffers still move back to CPU. - `layerwise-offload`: keep component weights on CPU and stream its declared layers during execution. Selectors match exact loaded component keys from `model_index.json`, including names such as `transformer_2`, `audio_vae`, and `connectors`. The group selectors `dit`, `text_encoder`, `image_encoder`, and `vae` are also available, together with `all`. An exact key overrides a matching group, and a group overrides `all`. Components without a matching canonical selector retain their explicit legacy setting or automatic/model default. @@ -321,6 +322,34 @@ The existing `--dit-cpu-offload`, `--text-encoder-cpu-offload`, `--image-encoder Layerwise selection is strict. A native weighted component selected for `layerwise-offload` must declare its layer structure; otherwise startup fails with the unsupported component name instead of silently changing modes. Explicit non-resident placement also requires a request-time component-use declaration, so it cannot silently select a module that the pipeline does not manage. FSDP applies only to resident components. The Diffusers backend supports only pipeline-wide `all=resident` and `all=component-offload`. +#### Snapshot offload + +On NVIDIA CUDA GPUs with separate host and device memory, explicitly select +`--component-residency vae=snapshot-offload` to avoid repeated weight D2H transfers. +Existing defaults and `component-offload` behavior are unchanged. + +Snapshot offload retains a full host copy during GPU execution. With +`--pin-cpu-memory` enabled (the default), it shares a per-worker host pin budget +with layerwise offload. Layerwise initialization claims its allowance first; +snapshot weights use the remainder. Shared parameter storage is pinned once and +reused across requests; its allowance is returned when the storage is released. +Pinning also checks current host headroom, including the cgroup limit. + +Weights that do not fit the budget, or have pinning disabled, retain their +existing CPU storage. Checkpoint mmap stays file-backed unless pinning, dtype +conversion, or weight mutation materializes a copy. Pageable H2D can be slower, +and pinning temporarily needs space for both the original and pinned storage. +This budget limits new pinned weight allocations, not total process RAM or the +CUDA host allocator's cache. Benchmark repeated requests with the actual host +storage and workload; avoiding D2H does not guarantee an end-to-end speedup. + +The whole component must still fit on GPU. Unlike layerwise offload with all layers +retained across denoising steps, snapshot offload does not require a layer-structure +declaration and also covers parameters outside transformer blocks. Weight updates, +LoRA merge/unmerge, and sleep restore the CPU weights before mutation or release. +FSDP-managed components, shared-memory GPUs, and snapshot-offloaded DiTs with +breakable CUDA graphs are not supported. + ### Layerwise Offload Tuning Use layerwise offload when a component does not fit comfortably in GPU memory. The compatibility options `--dit-layerwise-offload` and `--layerwise-offload-components` remain available (`--layerwise-offload-modules` is an alias), while new deployments can select the mode directly: diff --git a/docs/docs/sglang-diffusion/deployment_cookbook.mdx b/docs/docs/sglang-diffusion/deployment_cookbook.mdx index 17dc36634..87538d287 100644 --- a/docs/docs/sglang-diffusion/deployment_cookbook.mdx +++ b/docs/docs/sglang-diffusion/deployment_cookbook.mdx @@ -132,7 +132,7 @@ See [OpenAI API: Served model name](/docs/sglang-diffusion/api/openai_api#served The modes tune native pipeline components declared to the component residency manager. DiTs, text/image encoders, VAEs, vocoders, adapters, and upsamplers can use layerwise offload when their native module declares its executable layer structure. Explicitly selecting an unsupported component fails at startup instead of falling back to another residency mode. -For direct control, assign one of `resident`, `component-offload`, or `layerwise-offload` with `--component-residency COMPONENT=MODE`: +For direct control, assign one of `resident`, `component-offload`, `snapshot-offload`, or `layerwise-offload` with `--component-residency COMPONENT=MODE`: ```bash sglang generate \ diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 78b6df23d..d6f13135e 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -39,6 +39,7 @@ from sglang.multimodal_gen.runtime.loader.gguf_weights import ( from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( COMPONENT_OFFLOAD, + SNAPSHOT_OFFLOAD, ComponentResidencyError, ) from sglang.multimodal_gen.runtime.platforms import current_platform @@ -232,7 +233,7 @@ def _uses_component_offload( ) -> bool: if component_name is None: return legacy_enabled - return server_args.residency_mode(component_name) == COMPONENT_OFFLOAD + return server_args.should_cpu_offload_component(component_name) def _reject_explicit_component_selector( @@ -243,12 +244,12 @@ def _reject_explicit_component_selector( ) -> None: if component_name is None: return - selected_by_component_residency = ( - server_args.canonical_residency_mode(component_name) == COMPONENT_OFFLOAD - ) + selected_by_component_residency = server_args.canonical_residency_mode( + component_name + ) in (COMPONENT_OFFLOAD, SNAPSHOT_OFFLOAD) if selected_by_component_residency: raise ComponentResidencyError( - f"{feature_name} does not support component-offload for " + f"{feature_name} does not support {server_args.canonical_residency_mode(component_name)} for " f"{component_name!r}; select resident or layerwise-offload" ) diff --git a/python/sglang/multimodal_gen/runtime/loader/utils.py b/python/sglang/multimodal_gen/runtime/loader/utils.py index 4ec8a6f65..c96dcce28 100644 --- a/python/sglang/multimodal_gen/runtime/loader/utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/utils.py @@ -19,6 +19,9 @@ from safetensors.torch import load_file as safetensors_load_file from torch import nn from torch.nn.utils import parametrize +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + weight_snapshot, +) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.weights.source import ( filter_duplicate_precision_variant_safetensors, @@ -691,6 +694,8 @@ def component_residency_bytes(module) -> Dict[str, int]: add(tensor) for tensor in module.buffers(): add(tensor) + for tensor in (weight_snapshot(module) or {}).values(): + add(tensor) for manager in getattr(module, "layerwise_offload_managers", None) or []: iter_cpu_weights = getattr(manager, "iter_cpu_weights", None) if iter_cpu_weights is None: diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index de32a17a7..8856f670b 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency impor resolve_default_workload, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + get_global_component_residency_manager, peek_global_component_residency_manager, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( @@ -417,6 +418,9 @@ class GPUWorker(GPUWorkerPostTrainingMixin): configure_layerwise_offload_modules( self.pipeline.modules, self.server_args, + pin_budget=get_global_component_residency_manager( + self.pipeline, self.server_args + ).host_pin_budget, component_names=( None if self.server_args.component_residency is not None diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py index 2a26d6f44..8d94d5141 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py @@ -10,6 +10,7 @@ from sglang.multimodal_gen import envs from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD, + SNAPSHOT_OFFLOAD, ComponentResidencyError, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import ( @@ -17,8 +18,12 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_ ComponentResidencyStrategy, LayerwiseOffloadStrategy, ResidentStrategy, + SnapshotOffloadStrategy, is_fsdp_managed_module, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import ( + HostPinBudget, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, is_resident_layerwise_module, @@ -87,6 +92,8 @@ def build_component_residency_strategy( component_name: str, module: nn.Module, server_args: ServerArgs, + *, + pin_budget: HostPinBudget | None = None, ) -> ComponentResidencyStrategy: residency_mode = server_args.residency_mode(component_name) if is_layerwise_offloaded_module(module): @@ -96,11 +103,16 @@ def build_component_residency_strategy( f"Component {component_name!r} resolved to layerwise-offload, but its " "loaded module did not enable layerwise offload" ) - if residency_mode == COMPONENT_OFFLOAD and is_fsdp_managed_module(module): + if residency_mode in ( + COMPONENT_OFFLOAD, + SNAPSHOT_OFFLOAD, + ) and is_fsdp_managed_module(module): raise ComponentResidencyError( - f"Component {component_name!r} resolved to component-offload, but it " + f"Component {component_name!r} resolved to {residency_mode}, but it " "was loaded as an FSDP-managed module" ) + if residency_mode == SNAPSHOT_OFFLOAD: + return SnapshotOffloadStrategy(pin_budget=pin_budget) if ( not current_platform.is_mps() and not is_fsdp_managed_module(module) @@ -119,6 +131,7 @@ class ComponentResidencyManager: self.pipeline = pipeline self.server_args = server_args self.state = ResidencyState() + self._host_pin_budget: HostPinBudget | None = None self._stage_names_by_id: dict[int, str] = {} self._stage_uses_by_index: list[tuple[ComponentUse, ...]] = [] self._ordered_uses: tuple[ComponentUse, ...] = () @@ -146,9 +159,17 @@ class ComponentResidencyManager: self._warmup_phase_peaks: dict[str, WarmupPhasePeak] = {} self._completed_warmup_phase_peaks: dict[str, WarmupPhasePeak] = {} + @property + def host_pin_budget(self) -> HostPinBudget: + # measure headroom after loading, when the first offload path needs it + if self._host_pin_budget is None: + self._host_pin_budget = HostPinBudget() + return self._host_pin_budget + def refresh_pipeline(self, pipeline: ComponentResidencyPipeline) -> None: custom_strategies = dict(pipeline.component_residency_strategies) if pipeline is not self.pipeline: + self._host_pin_budget = None self._remove_nvtx_hooks() self._strategy_cache.clear() self._active_use = None @@ -231,7 +252,7 @@ class ComponentResidencyManager: for component_name, module in self.pipeline.modules.items() if isinstance(module, nn.Module) and self.server_args.explicit_residency_mode(component_name) - in (COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD) + in (COMPONENT_OFFLOAD, SNAPSHOT_OFFLOAD, LAYERWISE_OFFLOAD) and component_name not in declared_components ) if unmanaged_components: @@ -821,6 +842,13 @@ class ComponentResidencyManager: component_name, module, self.server_args, + pin_budget=( + self.host_pin_budget + if self.server_args.residency_mode(component_name) + == SNAPSHOT_OFFLOAD + and self.server_args.pin_cpu_memory + else None + ), ) else: strategy = custom_strategy diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency.py index 0835d331a..1ab8611c3 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency.py @@ -18,11 +18,13 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_co RESIDENT = "resident" COMPONENT_OFFLOAD = "component-offload" +SNAPSHOT_OFFLOAD = "snapshot-offload" LAYERWISE_OFFLOAD = "layerwise-offload" COMPONENT_RESIDENCY_MODES = frozenset( ( RESIDENT, COMPONENT_OFFLOAD, + SNAPSHOT_OFFLOAD, LAYERWISE_OFFLOAD, ) ) @@ -151,6 +153,10 @@ def resolve_diffusers_pipeline_offload( "--component-residency layerwise-offload requires the native SGLang backend" ) + if SNAPSHOT_OFFLOAD in assignments.values(): + raise ComponentResidencyError( + "--component-residency snapshot-offload requires the native SGLang backend" + ) pipeline_mode = assignments.get(LAYERWISE_OFFLOAD_ALL_COMPONENTS) if len(assignments) == 1 and pipeline_mode is not None: return pipeline_mode == COMPONENT_OFFLOAD diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py index 326a14a90..9267c7b84 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py @@ -10,11 +10,17 @@ from torch.distributed.fsdp import FSDPModule from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import ( + HostPinBudget, shared_pool_available_bytes, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + capture_weight_snapshot, + restore_weight_snapshot, + weight_snapshot, +) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -134,13 +140,16 @@ class ComponentOffloadStrategy(ComponentResidencyStrategy): self._prefetch_stream: object | None = None self._ready_events: dict[str, object] = {} + def _load_component(self, module: nn.Module, use: ComponentUse) -> None: + _module_to_local_device(module, dtype=use.target_dtype) + def prepare_for_use( self, module: nn.Module, use: ComponentUse, state: ResidencyState, ) -> None: - _module_to_local_device(module, dtype=use.target_dtype) + self._load_component(module, use) def wait_for_use( self, @@ -169,7 +178,7 @@ class ComponentOffloadStrategy(ComponentResidencyStrategy): device=get_local_torch_device() ) with torch.get_device_module().stream(self._prefetch_stream): - _module_to_local_device(module, dtype=use.target_dtype) + self._load_component(module, use) event = torch.get_device_module().Event() event.record(self._prefetch_stream) self._ready_events[use.component_name] = event @@ -210,6 +219,36 @@ class ComponentOffloadStrategy(ComponentResidencyStrategy): self.finish_use(module, use, state) +class SnapshotOffloadStrategy(ComponentOffloadStrategy): + """Keep CPU weights during device use; restore them without weight D2H.""" + + def __init__(self, *, pin_budget: HostPinBudget | None = None) -> None: + super().__init__() + self._pin_budget = pin_budget + + def _load_component(self, module: nn.Module, use: ComponentUse) -> None: + if weight_snapshot(module) is not None and not _module_ready_on_local_device( + module, dtype=use.target_dtype + ): + restore_weight_snapshot(module) + if weight_snapshot(module) is None: + if use.target_dtype is not None: + module.to(dtype=use.target_dtype) + capture_weight_snapshot( + module, pin_budget=self._pin_budget, component_name=use.component_name + ) + super()._load_component(module, use) + + def finish_use( + self, module: nn.Module, use: ComponentUse, state: ResidencyState + ) -> None: + self.wait_for_use(module, use, state) + if restore_weight_snapshot(module): + self._ready_events.pop(use.component_name, None) + else: + super().finish_use(module, use, state) + + class LayerwiseOffloadStrategy(ComponentResidencyStrategy): """Run the lifecycle of an already configured layerwise component.""" diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py index e5d9a1f88..f851966cc 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py @@ -14,8 +14,10 @@ the cap is read directly from whichever cgroup version is mounted. """ import os +import weakref import psutil +import torch from sglang.multimodal_gen import envs from sglang.multimodal_gen.runtime.platforms import current_platform @@ -283,6 +285,14 @@ class HostPinBudget: ) return False + def release(self, weight_bytes: int) -> None: + """Return an allowance when its pinned storage is no longer owned.""" + self.committed_bytes -= weight_bytes + + def track_storage(self, storage: torch.UntypedStorage) -> None: + """Tie an already booked allowance to the storage's last owner.""" + weakref.finalize(storage, self.release, storage.nbytes()) + def pin_benefit_bytes(*, weight_bytes: int, uses_per_request: int) -> int: """Host-to-device bytes a pin would cover for one request. 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 7a7600c94..44c2a0e66 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 @@ -8,7 +8,7 @@ import re import sys import threading import time -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from contextlib import nullcontext from time import perf_counter from typing import ( @@ -1073,24 +1073,40 @@ class LayerwiseOffloadManager: def _layer_byte_totals( self, layer_groups: Dict ) -> Tuple[Dict[int, int], Dict[int, int]]: - """Per layer: (all weight bytes, the subset that are checkpoint views).""" + """Per layer: (host allocation bytes, checkpoint-view bytes).""" totals: Dict[int, int] = {} mapped: Dict[int, int] = {} for layer_idx, dtype_to_params in layer_groups.items(): total = 0 from_mapping = 0 - for weights in dtype_to_params.values(): + for dtype, weights in dtype_to_params.items(): + offset = 0 for _, weight in weights: tensor = self._to_local_tensor(weight) - nbytes = tensor.untyped_storage().nbytes() - total += nbytes + if tensor.is_contiguous(): + offset = ( + self._align_numel_offset(offset, dtype) + tensor.numel() + ) + else: + # match empty_strided's allocation, including view holes + total += ( + torch.empty_strided( + tensor.shape, + tensor.stride(), + dtype=dtype, + device="meta", + ) + .untyped_storage() + .nbytes() + ) if self._mapped_regions.holds(tensor): - from_mapping += nbytes + from_mapping += tensor.untyped_storage().nbytes() + total += offset * dtype.itemsize totals[layer_idx] = total mapped[layer_idx] = from_mapping return totals, mapped - def _plan_layer_hosting(self, layer_groups: Dict) -> Dict[int, str]: + def _plan_layer_hosting(self, layer_groups: Dict) -> Tuple[Dict[int, str], int]: """Where each layer's weights live on the host: pinned, pageable or mapped. Pinning is what lets the copy stream run ahead of compute; a pageable @@ -1128,7 +1144,7 @@ class LayerwiseOffloadManager: len(totals), sum(1 for where in hosting.values() if where == "pageable"), ) - return hosting + return hosting, 0 pinned_bytes = 0 hosting: Dict[int, str] = {} pin_order: List[int] = [] @@ -1203,7 +1219,7 @@ class LayerwiseOffloadManager: counts["mapped"], sum(totals.values()) / 1024**3, ) - return hosting + return hosting, pinned_bytes def _initialize_layer_weights(self) -> None: self._named_parameters = dict(self.model.named_parameters()) @@ -1224,8 +1240,19 @@ class LayerwiseOffloadManager: local_tensor.dtype, [] ).append((name, tensor)) - layer_hosting = self._plan_layer_hosting(layer_groups) + layer_hosting, untracked_bytes = self._plan_layer_hosting(layer_groups) + try: + for storage in self._initialize_host_stores(layer_groups, layer_hosting): + self._pin_budget.track_storage(storage) + untracked_bytes -= storage.nbytes() + finally: + # failed allocations have no storage finalizer to return their allowance + self._pin_budget.release(untracked_bytes) + def _initialize_host_stores( + self, layer_groups: Dict, layer_hosting: Dict[int, str] + ) -> Iterator[torch.UntypedStorage]: + """Yield each pinned allocation before copying weights to transfer its lease.""" # 2. concat and offload (in pinned memory) for layer_idx, dtype_to_params in layer_groups.items(): self._consolidated_cpu_weights[layer_idx] = {} @@ -1283,6 +1310,8 @@ class LayerwiseOffloadManager: dtype=dtype, pin_memory=pin_this_layer, ) + if pin_this_layer: + yield cpu_tensor.untyped_storage() cpu_tensor.copy_(local_weight) self._strided_cpu_weights[layer_idx][name] = cpu_tensor self._weight_metadata[layer_idx][name] = { @@ -1315,6 +1344,8 @@ class LayerwiseOffloadManager: cpu_buffer = torch.empty( total_numel, dtype=dtype, pin_memory=pin_this_layer ) + if pin_this_layer: + yield cpu_buffer.untyped_storage() # offload weights to the buffer for name, weight, local_weight in contiguous_weights: @@ -1942,13 +1973,11 @@ class LayerwiseOffloadManager: "cannot release host stores with mapped copies in flight" ) - self._pin_budget.release(self.pinned_host_weight_bytes()) self._consolidated_cpu_weights.clear() self._strided_cpu_weights.clear() self._mapped_cpu_weights.clear() self._mps_cpu_weights.clear() self._weight_metadata.clear() - self._layer_hosting.clear() self._prefetch_events.clear() self._mapped_bytes = 0 self._configured = False @@ -2730,6 +2759,8 @@ def configure_layerwise_offload_modules( server_args: ServerArgs, component_names: Sequence[str] | None = None, warn_missing: bool = True, + *, + pin_budget: HostPinBudget | None = None, ) -> list[str]: """Configure layerwise offload for the given modules, from the given component_names @@ -2907,7 +2938,8 @@ def configure_layerwise_offload_modules( key=_h2d_bytes_a_pin_would_save, reverse=True, ) - pin_budget = HostPinBudget() + if pin_budget is None: + pin_budget = HostPinBudget() logger.info("Layerwise offload host memory: %s", describe_host_memory()) for component_name in selected_pipeline_component_names: diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/memory_occupation_controller.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/memory_occupation_controller.py index 612445b8e..008caaf09 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/memory_occupation_controller.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/memory_occupation_controller.py @@ -7,6 +7,9 @@ import torch from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + restore_weight_snapshot, +) from sglang.multimodal_gen.runtime.pipelines_core import ComposedPipelineBase from sglang.multimodal_gen.runtime.post_training.weights_updater import ( get_updatable_modules, @@ -132,7 +135,8 @@ class MemoryOccupationController: module = modules[name] src_device_map[name] = _get_module_device(module) if device.startswith("cpu"): - _module_to_pinned_cpu(module) + if not restore_weight_snapshot(module): + _module_to_pinned_cpu(module) else: module.to(device, non_blocking=True) moved.append(name) diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/weight_snapshot.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/weight_snapshot.py new file mode 100644 index 000000000..880004950 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/weight_snapshot.py @@ -0,0 +1,78 @@ +"""CPU weight ownership while a snapshot-offloaded component runs on device.""" + +import torch +from torch import nn + +from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import ( + HostPinBudget, + host_copies_would_not_fit, +) + + +def weight_snapshot(module: nn.Module) -> dict[str, torch.Tensor] | None: + return module.__dict__.get("_offload_weight_snapshot") + + +def capture_weight_snapshot( + module: nn.Module, + *, + pin_budget: HostPinBudget | None = None, + component_name: str = "", +) -> None: + parameters = dict(module.named_parameters()) + snapshot = { + name: parameter.detach().to("cpu") for name, parameter in parameters.items() + } + if pin_budget is not None: + # pin each storage once, preserving tied views and releasing the original + # CPU storage as we go rather than staging another complete model copy + storage_names: dict[int, list[str]] = {} + for name, tensor in snapshot.items(): + storage_names.setdefault(tensor.untyped_storage().data_ptr(), []).append( + name + ) + for names in storage_names.values(): + storage = snapshot[names[0]].untyped_storage() + size = storage.nbytes() + if storage.is_pinned() or not size or size > pin_budget.spendable_bytes: + continue + if host_copies_would_not_fit(size): + continue + if not pin_budget.request(component_name=component_name, weight_bytes=size): + continue + try: + pinned_storage = storage.pin_memory() + except Exception: + pin_budget.release(size) + raise + # the lease outlives strategy rebuilds, snapshots and LoRA backups; + # only the last tensor releasing this storage returns its allowance + pin_budget.track_storage(pinned_storage) + for name in names: + tensor = snapshot[name] + pinned = torch.empty(0, dtype=tensor.dtype, device="cpu").set_( + pinned_storage, + tensor.storage_offset(), + tensor.shape, + tensor.stride(), + ) + snapshot[name] = pinned + if parameters[name].device.type == "cpu": + parameters[name].data = pinned + module._offload_weight_snapshot = snapshot + + +def restore_weight_snapshot(module: nn.Module) -> bool: + """Restore CPU parameters before offload or mutation, preserving live buffers.""" + snapshot = weight_snapshot(module) + if snapshot is None: + return False + # drain prefetch reads of the host weights before a writer can mutate them + torch.get_device_module().synchronize() + with torch.no_grad(): + for name, parameter in module.named_parameters(): + parameter.data = snapshot[name] + # buffers may change during forward and must not be restored from a snapshot + module.to("cpu") + del module._offload_weight_snapshot + return True diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py index 462307d7b..5f6620b8b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py @@ -23,6 +23,10 @@ from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + restore_weight_snapshot, + weight_snapshot, +) from sglang.multimodal_gen.runtime.models.dits.base import BaseDiT from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( ComposedPipelineBase, @@ -297,6 +301,8 @@ class LoRAPipeline(ComposedPipelineBase): offload_disabled_modules = [] for module_name in module_names: module = self.modules.get(module_name) + if isinstance(module, torch.nn.Module): + restore_weight_snapshot(module) if module is not None and is_layerwise_offloaded_module(module): module.disable_offload() offload_disabled_modules.append(module) @@ -320,6 +326,11 @@ class LoRAPipeline(ComposedPipelineBase): if any(layer.merged for layer in lora_layers_dict.values()): return True module = self.modules.get(module_name) + if ( + isinstance(module, torch.nn.Module) + and weight_snapshot(module) is not None + ): + return True if module is not None and is_layerwise_offloaded_module(module): return True return False 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 ba14eb7eb..6d59329cf 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 @@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentUse, + peek_global_component_residency_manager, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( configure_layerwise_offload_modules, @@ -544,9 +545,15 @@ class LTX2ImageEncodingStage(PipelineStage): "condition_image_encoder" ): modules = {"condition_image_encoder": self._condition_image_encoder} + residency_manager = peek_global_component_residency_manager() configure_layerwise_offload_modules( modules, server_args, + pin_budget=( + residency_manager.host_pin_budget + if residency_manager is not None + else None + ), component_names=( None if server_args.component_residency is not None diff --git a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py index 2111beb2c..75359a82c 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py +++ b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py @@ -57,6 +57,9 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + restore_weight_snapshot, +) from sglang.multimodal_gen.runtime.models.dits.base import BaseDiT from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import DiffusersPipeline from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import ( @@ -185,6 +188,7 @@ def _load_weights_into_module(module: torch.nn.Module, weights_iter) -> None: and returns an HTTP error. """ with torch.inference_mode(): + restore_weight_snapshot(module) model_params = dict(module.named_parameters()) weights_iter = _iter_module_weight_updates(module, weights_iter, model_params) diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index c44474b99..cd9a9a49a 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -36,6 +36,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD, RESIDENT, + SNAPSHOT_OFFLOAD, normalize_component_residency, resolve_component_residency_mode, resolve_diffusers_pipeline_offload, @@ -915,6 +916,23 @@ class ServerArgs(DisaggServerArgsMixin): self.component_residency = normalize_component_residency( self.component_residency ) + if SNAPSHOT_OFFLOAD in (self.component_residency or {}).values(): + if ( + not current_platform.is_cuda() + or current_platform.device_shares_host_memory() + ): + raise ValueError( + "snapshot-offload requires CUDA with separate host and device memory; " + "use component-offload or layerwise-offload on this platform" + ) + if self.enable_breakable_cuda_graph and any( + self.canonical_residency_mode(name) == SNAPSHOT_OFFLOAD + for name in ("transformer", "transformer_2") + ): + raise ValueError( + "snapshot-offload for DiT is incompatible with " + "--enable-breakable-cuda-graph because weight addresses change" + ) def _adjust_ltx2_two_stage_device_mode(self): if not self._is_ltx23_two_stage_pipeline(): @@ -942,7 +960,7 @@ class ServerArgs(DisaggServerArgsMixin): component_name: residency_mode for component_name in ("transformer", "transformer_2") if (residency_mode := self.explicit_residency_mode(component_name)) - in (COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD) + in (COMPONENT_OFFLOAD, SNAPSHOT_OFFLOAD, LAYERWISE_OFFLOAD) } if mode == "resident" and explicit_nonresident_dits: configured = ", ".join( @@ -1644,11 +1662,15 @@ class ServerArgs(DisaggServerArgsMixin): return RESIDENT def should_cpu_offload_component(self, component_name: str) -> bool: - return self.residency_mode(component_name) == COMPONENT_OFFLOAD + return self.residency_mode(component_name) in ( + COMPONENT_OFFLOAD, + SNAPSHOT_OFFLOAD, + ) def should_start_component_on_cpu(self, component_name: str) -> bool: return self.residency_mode(component_name) in ( COMPONENT_OFFLOAD, + SNAPSHOT_OFFLOAD, LAYERWISE_OFFLOAD, ) @@ -1745,7 +1767,7 @@ class ServerArgs(DisaggServerArgsMixin): has_explicit_dit_offload = bool( self.canonical_residency_mode("transformer") - in (COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD) + in (COMPONENT_OFFLOAD, SNAPSHOT_OFFLOAD, LAYERWISE_OFFLOAD) or self.is_explicit_layerwise_offload_component("transformer") or ( self.is_arg_explicitly_set("cpu_offload_components") @@ -2408,7 +2430,7 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.component_residency, metavar="COMPONENT=MODE", help=( - "Select resident, component-offload, or layerwise-offload for " + "Select resident, component-offload, snapshot-offload, or layerwise-offload for " "pipeline components. Exact model_index.json component keys override " "the dit, text_encoder, image_encoder, vae, and all groups. " "Components without an assignment keep their automatic placement." diff --git a/python/sglang/multimodal_gen/test/unit/test_component_residency.py b/python/sglang/multimodal_gen/test/unit/test_component_residency.py index 3c51e8418..3af44ea18 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_residency.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_residency.py @@ -3,12 +3,15 @@ from unittest.mock import Mock import pytest import torch +from safetensors.torch import load_file, save_file +from sglang.multimodal_gen.runtime.loader.utils import component_residency_bytes from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentResidencyManager, ComponentUse, ResidencyState, WarmupPhasePeak, + build_component_residency_strategy, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( ComponentResidencyError, @@ -16,7 +19,15 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import ( ComponentOffloadStrategy, ResidentStrategy, + SnapshotOffloadStrategy, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.memory_occupation_controller import ( + MemoryOccupationController, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + weight_snapshot, +) +from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( ImageEncodingStage, ) @@ -24,6 +35,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime.text_encoding RealtimeTextEncodingStage, ) from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.post_training.weights_updater import ( + _load_weights_into_module, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -52,6 +66,141 @@ def test_component_offload_releases_preferred_component_after_request(): strategy.finish_use.assert_called_once_with(module, use, state) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("host_storage", ["pageable", "pinned", "mmap"]) +@pytest.mark.parametrize("prefetch", [False, True]) +def test_snapshot_offload_preserves_host_storage_and_live_buffers( + tmp_path, monkeypatch, host_storage, prefetch +): + module = torch.nn.Linear(16, 16, bias=False) + if host_storage == "pinned": + module.weight.data = module.weight.detach().pin_memory() + elif host_storage == "mmap": + checkpoint = str(tmp_path / "model.safetensors") + save_file(module.state_dict(), checkpoint) + module.load_state_dict(load_file(checkpoint), assign=True) + module.register_buffer("counter", torch.zeros((), dtype=torch.int64)) + original_host = module.weight.detach() + pointer = original_host.data_ptr() + x = torch.randn(2, 16, device="cuda") + expected = torch.nn.functional.linear(x, original_host.to("cuda")) + args = SimpleNamespace(residency_mode=lambda _: "snapshot-offload") + strategy = build_component_residency_strategy("vae", module, args) + assert isinstance(strategy, SnapshotOffloadStrategy) + use = ComponentUse("decode", "vae") + state = ResidencyState() + + original_to = torch.Tensor.to + weight_d2h = [] + + def tracked_to(tensor, *args, **kwargs): + device = kwargs.get("device", args[0] if args else None) + if ( + tensor.device.type == "cuda" + and isinstance(device, (str, torch.device)) + and torch.device(device).type == "cpu" + and tensor.numel() == 256 + ): + weight_d2h.append(tensor.numel()) + return original_to(tensor, *args, **kwargs) + + monkeypatch.setattr(torch.Tensor, "to", tracked_to) + for iteration in range(3): + if prefetch: + strategy.prefetch_for_use(module, use, state) + else: + strategy.prepare_for_use(module, use, state) + strategy.wait_for_use(module, use, state) + assert weight_snapshot(module)["weight"].data_ptr() == pointer + totals = component_residency_bytes(module) + assert sum(totals[k] for k in ("host", "host_pinned", "host_mapped")) == 1024 + torch.testing.assert_close(module(x), expected, rtol=0, atol=0) + module.counter.add_(1) + strategy.finish_use(module, use, state) + assert module.weight.device.type == "cpu" + assert module.weight.data_ptr() == pointer + assert module.counter.item() == iteration + 1 + assert weight_snapshot(module) is None + assert not weight_d2h + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_snapshot_offload_refit_and_lora_context_refresh_weights(): + module = torch.nn.Linear(4, 4, bias=False) + strategy = SnapshotOffloadStrategy() + use = ComponentUse("denoise", "transformer") + state = ResidencyState() + strategy.prefetch_for_use(module, use, state) + strategy.wait_for_use(module, use, state) + _load_weights_into_module(module, [("weight", torch.full((4, 4), 2.0))]) + assert module.weight.device.type == "cpu" + strategy.prepare_for_use(module, use, state) + torch.testing.assert_close(module.weight, torch.full((4, 4), 2.0, device="cuda")) + + pipeline = SimpleNamespace(modules={"transformer": module}) + with LoRAPipeline._temporarily_disable_offload( + pipeline, target="transformer", use_module_names_only=True + ): + # exercise the same weight-mutation boundary as merge and layer replacement + module.weight = torch.nn.Parameter(torch.full((4, 4), 3.0)) + strategy.prepare_for_use(module, use, state) + torch.testing.assert_close(module.weight, torch.full((4, 4), 3.0, device="cuda")) + strategy.finish_use(module, use, state) + with LoRAPipeline._temporarily_disable_offload( + pipeline, target="transformer", use_module_names_only=True + ): + with torch.no_grad(): + module.weight.sub_(1) + strategy.prepare_for_use(module, use, state) + torch.testing.assert_close(module.weight, torch.full((4, 4), 2.0, device="cuda")) + strategy.finish_use(module, use, state) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_snapshot_offload_dtype_tied_storage_and_strategy_rebuild(): + module = torch.nn.Module() + storage = torch.randn(32, dtype=torch.bfloat16) + module.register_parameter("a", torch.nn.Parameter(storage[:16])) + module.register_parameter("b", torch.nn.Parameter(storage[16:])) + module.register_parameter("tied", module.a) + strategy = SnapshotOffloadStrategy() + use = ComponentUse("decode", "vae", target_dtype=torch.bfloat16) + state = ResidencyState() + strategy.prepare_for_use(module, use, state) + assert module.a is module.tied + rebuilt = SnapshotOffloadStrategy() + rebuilt.finish_use(module, use, state) + assert module.a is module.tied + assert ( + module.a.untyped_storage().data_ptr() == module.b.untyped_storage().data_ptr() + ) + assert module.a.data_ptr() == storage.data_ptr() + rebuilt.prepare_for_use( + module, ComponentUse("decode", "vae", target_dtype=torch.float32), state + ) + rebuilt.finish_use(module, use, state) + assert module.a.dtype == torch.float32 + assert module.a is module.tied + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_snapshot_offload_sleep_uses_existing_host_storage(): + module = torch.nn.Linear(4, 4, bias=False) + host_pointer = module.weight.data_ptr() + strategy = SnapshotOffloadStrategy() + use = ComponentUse("denoise", "transformer") + state = ResidencyState() + strategy.prepare_for_use(module, use, state) + pipeline = SimpleNamespace(modules={"transformer": module}) + controller = MemoryOccupationController(pipeline, rank=0, use_fsdp_inference=False) + controller._move_modules(["transformer"], "cpu") + assert module.weight.data_ptr() == host_pointer + assert weight_snapshot(module) is None + strategy.prepare_for_use(module, use, state) + strategy.finish_use(module, use, state) + assert module.weight.data_ptr() == host_pointer + + def test_component_offload_keeps_preferred_component_after_warmup(): strategy = ComponentOffloadStrategy() strategy.prepare_for_use = Mock() 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 64db72b0a..ad35ddae1 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -1,3 +1,4 @@ +import gc import pathlib from contextlib import nullcontext from types import SimpleNamespace @@ -23,7 +24,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers import ( layerwise_offload as layerwise_offload_mod, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + ComponentResidencyManager, ComponentUse, + ResidencyState, build_component_residency_strategy, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( @@ -470,6 +473,147 @@ def test_pin_budget_ranks_by_steps_resolved_from_model_index(monkeypatch): ) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_snapshot_and_layerwise_share_the_residency_managers_pin_budget(): + transformer = _NestedDummyModel() + vae = torch.nn.Linear(4, 4, bias=False) + encoder = torch.nn.Linear(32, 32, bias=False) + modules = {"transformer": transformer, "vae": vae, "text_encoder": encoder} + pipeline = SimpleNamespace( + modules=modules, _stage_name_mapping={}, component_residency_strategies={} + ) + args = _server_args( + component_residency={ + "transformer": "layerwise-offload", + "vae": "snapshot-offload", + "text_encoder": "snapshot-offload", + }, + pin_cpu_memory=True, + ) + manager = ComponentResidencyManager(pipeline, args) + budget = manager.host_pin_budget + budget.available_bytes = host_memory_budget.MIN_HOST_RESERVE_BYTES + 1024 + budget.reserve_bytes = host_memory_budget.MIN_HOST_RESERVE_BYTES + configured = configure_layerwise_offload_modules(modules, args, pin_budget=budget) + assert configured == ["transformer"] + layerwise = transformer.layerwise_offload_managers[0] + assert layerwise._pin_budget is budget + booked = budget.committed_bytes + assert 0 < booked < 1024 - 64 + for name in ("vae", "text_encoder"): + module = modules[name] + strategy = manager.strategy_for(name, module) + use = ComponentUse("encode", name) + strategy.prepare_for_use(module, use, ResidencyState()) + strategy.finish_use(module, use, ResidencyState()) + assert budget.committed_bytes == booked + 64 + assert module.weight.is_pinned() == (name == "vae") + transformer.disable_offload() + layerwise.release_host_stores() + assert budget.committed_bytes == 64 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_layerwise_pin_lease_includes_alignment_and_survives_host_aliases(): + budget = host_memory_budget.HostPinBudget( + available_bytes=host_memory_budget.MIN_HOST_RESERVE_BYTES + 76 + ) + for _ in range(2): + model = torch.nn.Module() + model.blocks = torch.nn.ModuleList([torch.nn.Linear(3, 3)]) + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=1, + enabled=True, + pin_cpu_memory=True, + pin_budget=budget, + ) + # 36 bytes of weights, 28 bytes of alignment, then a 12-byte bias + assert budget.committed_bytes == 76 + host_alias = manager._consolidated_cpu_weights[0][torch.float32].detach() + manager.remove_forward_hooks() + manager.load_all_layers() + torch.cuda.synchronize() + manager.enabled = False + manager.release_host_stores() + assert budget.committed_bytes == 76 + del host_alias, manager, model + gc.collect() + assert budget.committed_bytes == 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("stride,allocation_bytes", [(1, 16), (3, 40)]) +def test_layerwise_budget_uses_view_allocation_not_the_backing_storage( + stride, allocation_bytes +): + model = torch.nn.Module() + block = torch.nn.Module() + source = torch.arange(32, dtype=torch.float32)[1 : 1 + 4 * stride : stride] + expected = source.clone() + block.weight = torch.nn.Parameter(source) + model.blocks = torch.nn.ModuleList([block]) + budget = host_memory_budget.HostPinBudget( + available_bytes=host_memory_budget.MIN_HOST_RESERVE_BYTES + allocation_bytes + ) + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=1, + enabled=True, + pin_cpu_memory=True, + pin_budget=budget, + ) + assert budget.committed_bytes == allocation_bytes + manager.load_all_layers() + torch.cuda.synchronize() + torch.testing.assert_close(block.weight.cpu(), expected, rtol=0, atol=0) + manager.remove_forward_hooks() + manager.enabled = False + manager.release_host_stores() + assert budget.committed_bytes == 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_failed_layerwise_allocation_refunds_only_unallocated_allowance(monkeypatch): + budget = host_memory_budget.HostPinBudget( + available_bytes=host_memory_budget.MIN_HOST_RESERVE_BYTES + 1024 + ) + assert budget.request(component_name="other", weight_bytes=64) + model = torch.nn.Module() + model.blocks = torch.nn.ModuleList([torch.nn.Linear(3, 3) for _ in range(2)]) + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=2, + enabled=True, + initialize=False, + pin_cpu_memory=True, + pin_budget=budget, + ) + empty = torch.empty + allocations = 0 + + def fail_second_pin(*args, **kwargs): + nonlocal allocations + if kwargs.get("pin_memory"): + allocations += 1 + if allocations == 2: + raise RuntimeError("pin allocation failed") + return empty(*args, **kwargs) + + monkeypatch.setattr(torch, "empty", fail_second_pin) + with pytest.raises(RuntimeError, match="pin allocation failed"): + manager.initialize() + assert allocations == 2 + gc.collect() + assert budget.committed_bytes == 64 + 76 + del manager, model + gc.collect() + assert budget.committed_bytes == 64 + + def test_layerwise_configuration_filters_by_component_name(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule 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 00106d08d..aa481a729 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -69,6 +69,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency COMPONENT_OFFLOAD, LAYERWISE_OFFLOAD, RESIDENT, + SNAPSHOT_OFFLOAD, normalize_component_residency, resolve_component_residency_mode, resolve_diffusers_pipeline_offload, @@ -1108,6 +1109,50 @@ class TestOffloadDefaults(unittest.TestCase): with self.assertRaisesRegex(ValueError, "Invalid component residency mode"): normalize_component_residency(["dit=cpu"]) + def test_snapshot_offload_is_explicit_and_uses_cpu_load_policy(self): + args = self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "performance_mode": "manual", + "component_residency": ["vae=snapshot_offload", "dit=resident"], + "vae_cpu_offload": False, + "use_fsdp_inference": True, + }, + ) + self.assertEqual(args.residency_mode("video_vae"), SNAPSHOT_OFFLOAD) + self.assertTrue(args.should_cpu_offload_component("video_vae")) + self.assertTrue(args.should_start_component_on_cpu("video_vae")) + self.assertFalse(args.should_use_fsdp_for_component("video_vae")) + self.assertEqual(args.residency_mode("transformer"), RESIDENT) + self.assertTrue(args.should_use_fsdp_for_component("transformer")) + self.assertEqual( + resolve_component_residency_mode( + "video_vae", + normalize_component_residency( + "vae=snapshot-offload,video_vae=resident" + ), + ), + RESIDENT, + ) + + def test_snapshot_offload_rejects_shared_memory_and_captured_dit(self): + with patch.object( + current_platform, "device_shares_host_memory", return_value=True + ): + with self.assertRaisesRegex(ValueError, "separate host and device memory"): + self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={"component_residency": ["vae=snapshot-offload"]}, + ) + with self.assertRaisesRegex(ValueError, "weight addresses change"): + self._from_dict_with_task_type( + ModelTaskType.T2V, + kwargs={ + "component_residency": ["dit=snapshot-offload"], + "enable_breakable_cuda_graph": True, + }, + ) + def test_component_residency_resolves_exact_group_and_all_precedence(self): assignments = normalize_component_residency( [ @@ -1325,6 +1370,8 @@ class TestOffloadDefaults(unittest.TestCase): resolve_diffusers_pipeline_offload({"dit": COMPONENT_OFFLOAD}) with self.assertRaisesRegex(ValueError, "native SGLang backend"): resolve_diffusers_pipeline_offload({"all": LAYERWISE_OFFLOAD}) + with self.assertRaisesRegex(ValueError, "native SGLang backend"): + resolve_diffusers_pipeline_offload({"all": SNAPSHOT_OFFLOAD}) def test_memory_mode_layerwise_offloads_vae_on_low_memory_gpu(self): args = self._from_dict_with_task_type( diff --git a/python/sglang/multimodal_gen/test/unit/test_weight_snapshot.py b/python/sglang/multimodal_gen/test/unit/test_weight_snapshot.py new file mode 100644 index 000000000..10c660b19 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_weight_snapshot.py @@ -0,0 +1,186 @@ +"""Snapshot pin allowances follow storage ownership, not request lifetimes.""" + +import gc +from types import SimpleNamespace + +import pytest +import torch +from safetensors.torch import load_file, save_file + +from sglang.multimodal_gen.runtime.managers.memory_managers import host_memory_budget +from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + ComponentResidencyManager, + ComponentUse, + ResidencyState, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import ( + MIN_HOST_RESERVE_BYTES, + HostPinBudget, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.memory_occupation_controller import ( + MemoryOccupationController, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.weight_snapshot import ( + capture_weight_snapshot, + restore_weight_snapshot, + weight_snapshot, +) +from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") + + +@pytest.mark.parametrize("pin", [False, True]) +@pytest.mark.parametrize("mapped", [False, True]) +def test_manager_pins_once_across_requests_and_strategy_rebuilds(tmp_path, pin, mapped): + module = torch.nn.Linear(4, 4, bias=False) + if mapped: + checkpoint = str(tmp_path / "model.safetensors") + save_file(module.state_dict(), checkpoint) + module.load_state_dict(load_file(checkpoint), assign=True) + expected = module.weight.detach().clone() + pipeline = SimpleNamespace( + modules={"vae": module}, + _stage_name_mapping={}, + component_residency_strategies={}, + ) + args = SimpleNamespace( + residency_mode=lambda _: "snapshot-offload", pin_cpu_memory=pin + ) + manager = ComponentResidencyManager(pipeline, args) + budget = manager.host_pin_budget + use = ComponentUse("decode", "vae") + state = ResidencyState() + pointer = None + for _ in range(3): + strategy = manager.strategy_for("vae", module) + strategy.prefetch_for_use(module, use, state) + strategy.wait_for_use(module, use, state) + snapshot = weight_snapshot(module) + assert snapshot["weight"].is_pinned() == pin + assert budget.committed_bytes == (64 if pin else 0) + torch.testing.assert_close(module.weight, expected.cuda(), rtol=0, atol=0) + strategy.finish_use(module, use, state) + if pointer is not None: + assert module.weight.data_ptr() == pointer + pointer = module.weight.data_ptr() + # real server-args replacement invalidates the cached strategy + manager.refresh_server_args(SimpleNamespace(**vars(args))) + assert manager.host_pin_budget is budget + del snapshot + module.weight = torch.nn.Parameter(expected) + gc.collect() + assert budget.committed_bytes == 0 + + +def test_partial_budget_preserves_shared_storage_and_strided_views(): + storage = torch.arange(32, dtype=torch.float32) + module = torch.nn.Module() + module.a = torch.nn.Parameter(storage[:16].reshape(4, 4).T) + module.b = torch.nn.Parameter(storage[16:]) + module.tied = module.a + module.other = torch.nn.Parameter(torch.ones(32)) + budget = HostPinBudget(available_bytes=MIN_HOST_RESERVE_BYTES + 128) + capture_weight_snapshot(module, pin_budget=budget, component_name="vae") + assert budget.committed_bytes == 128 + assert module.a.is_pinned() and module.b.is_pinned() + assert not module.other.is_pinned() + assert module.a is module.tied + assert module.a.stride() == (1, 4) + assert module.b.storage_offset() == 16 + assert ( + module.a.untyped_storage().data_ptr() == module.b.untyped_storage().data_ptr() + ) + torch.testing.assert_close(module.a, storage[:16].reshape(4, 4).T, rtol=0, atol=0) + torch.testing.assert_close(module.b, storage[16:], rtol=0, atol=0) + restore_weight_snapshot(module) + backup = module.b.detach() + del module.a, module.b, module.tied + gc.collect() + assert budget.committed_bytes == 128 + del backup + gc.collect() + assert budget.committed_bytes == 0 + + +def test_lora_replacement_and_dtype_change_return_pin_allowance(): + module = torch.nn.Linear(4, 4, bias=False) + pipeline = SimpleNamespace(modules={"transformer": module}) + budget = HostPinBudget(available_bytes=MIN_HOST_RESERVE_BYTES + 128) + for value in (2.0, 3.0): + capture_weight_snapshot(module, pin_budget=budget) + module.cuda() + assert budget.committed_bytes == 64 + with LoRAPipeline._temporarily_disable_offload( + pipeline, target="transformer", use_module_names_only=True + ): + module.weight = torch.nn.Parameter(torch.full((4, 4), value)) + gc.collect() + assert budget.committed_bytes == 0 + capture_weight_snapshot(module, pin_budget=budget) + module.cuda() + torch.testing.assert_close( + module.weight, torch.full((4, 4), value, device="cuda") + ) + restore_weight_snapshot(module) + module.to(dtype=torch.bfloat16) + gc.collect() + assert budget.committed_bytes == 0 + capture_weight_snapshot(module, pin_budget=budget) + assert budget.committed_bytes == 32 + restore_weight_snapshot(module) + + +def test_live_headroom_prevents_pin_copy(monkeypatch): + module = torch.nn.Linear(4, 4, bias=False) + pointer = module.weight.data_ptr() + budget = HostPinBudget(available_bytes=16 * 1024**3) + monkeypatch.setattr(host_memory_budget, "host_memory_available_bytes", lambda: 0) + capture_weight_snapshot(module, pin_budget=budget) + assert module.weight.data_ptr() == pointer + assert not module.weight.is_pinned() + assert budget.committed_bytes == 0 + restore_weight_snapshot(module) + + +def test_sleep_keeps_the_pin_lease_and_existing_pins_are_reused(): + module = torch.nn.Linear(4, 4, bias=False) + budget = HostPinBudget(available_bytes=MIN_HOST_RESERVE_BYTES + 64) + capture_weight_snapshot(module, pin_budget=budget) + pointer = module.weight.data_ptr() + module.cuda() + pipeline = SimpleNamespace(modules={"transformer": module}) + controller = MemoryOccupationController(pipeline, rank=0, use_fsdp_inference=False) + controller._move_modules(["transformer"], "cpu") + assert module.weight.data_ptr() == pointer + assert budget.committed_bytes == 64 + capture_weight_snapshot(module, pin_budget=budget) + assert budget.committed_bytes == 64 + restore_weight_snapshot(module) + del module, pipeline, controller + gc.collect() + assert budget.committed_bytes == 0 + + module = torch.nn.Linear(4, 4, bias=False) + module.weight.data = module.weight.detach().pin_memory() + pointer = module.weight.data_ptr() + capture_weight_snapshot(module, pin_budget=budget) + assert module.weight.data_ptr() == pointer + assert budget.committed_bytes == 0 + restore_weight_snapshot(module) + + +def test_failed_pin_allocation_returns_allowance(monkeypatch): + module = torch.nn.Linear(4, 4, bias=False) + pointer = module.weight.data_ptr() + budget = HostPinBudget(available_bytes=MIN_HOST_RESERVE_BYTES + 64) + + def fail_pin(storage, device="cuda"): + raise RuntimeError("pin allocation failed") + + monkeypatch.setattr(torch.UntypedStorage, "pin_memory", fail_pin) + with pytest.raises(RuntimeError, match="pin allocation failed"): + capture_weight_snapshot(module, pin_budget=budget) + assert budget.committed_bytes == 0 + assert module.weight.data_ptr() == pointer + assert weight_snapshot(module) is None