[diffusion] feat: add explicit snapshot-offload component residency (#38535)

Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
Mick
2026-09-09 13:44:54 +08:00
committed by GitHub
co-authored by Mick Qian
parent 0ee8e41a4e
commit 00a9028e87
20 changed files with 835 additions and 29 deletions
+29
View File
@@ -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:
@@ -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 \
@@ -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"
)
@@ -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:
@@ -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
@@ -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
@@ -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
@@ -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."""
@@ -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.
@@ -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:
@@ -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)
@@ -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
@@ -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
@@ -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
@@ -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)
@@ -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."
@@ -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()
@@ -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
@@ -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(
@@ -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