[diffusion] feat: support resident layers for DiT (#31538)

This commit is contained in:
Yihao Wang
2026-07-29 16:52:55 +08:00
committed by GitHub
parent 0caf0fc01d
commit 227dadd79a
4 changed files with 240 additions and 5 deletions
@@ -16,6 +16,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_s
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
is_layerwise_offloaded_module,
is_resident_layerwise_module,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import (
is_dit_component_name,
@@ -419,6 +420,11 @@ class ComponentResidencyManager:
# Avoid making two vanilla-offloaded heavy components resident before
# a budget-aware planner can prove the overlap is safe.
return
if is_resident_layerwise_module(module):
# A layerwise DiT holding a large resident set must not be prefetched
# during a prior peer stage (e.g. text encoding): co-residing can lead
# to OOMs. Pin it lazily at the DiT's own use-site.
return
self._uses_seen[use.component_name] = use
if strategy.prefetch_for_use(module, use, self.state):
@@ -458,7 +464,9 @@ class ComponentResidencyManager:
if self.state.batch_is_warmup and use.keep_ready_after_warmup:
continue
preferred = component_name in preferred_uses
if not preferred and self._should_keep_single_dit(component_name):
if is_resident_layerwise_module(module):
preferred = False
elif not preferred and self._should_keep_single_dit(component_name):
continue
strategy = self.strategy_for(component_name, module)
if preferred and not self.state.batch_is_warmup:
@@ -537,6 +545,10 @@ class ComponentResidencyManager:
if use.component_name in future_component_names:
return True
if self._should_keep_single_dit(use.component_name):
module = self.get_module(use.component_name)
if module is not None and is_resident_layerwise_module(module):
# don't keep a layerwise DiT resident across the request to avoid OOMs
return False
return True
return False
@@ -42,12 +42,20 @@ class LayerwiseOffloadManager:
enabled: bool,
pin_cpu_memory: bool = True,
prefetch_size: int = 1,
resident_layers: int = 0,
) -> None:
self.model = model
self.layers_attr_str = layers_attr_str
self.num_layers = num_layers
self.pin_cpu_memory = pin_cpu_memory
self.prefetch_size = min(max(1, prefetch_size), self.num_layers)
# Leading layers held on GPU across denoise steps, instead of being
# re-streamed every step like the tail.
self.resident_layers = min(max(0, int(resident_layers)), self.num_layers)
# Armed on the first denoise forward, so that the load-time prefetch below
# does not pin the whole resident set before the DiT is the active component.
self._residency_active = False
self.enabled = bool(enabled and torch.get_device_module().is_available())
if not self.enabled:
return
@@ -247,18 +255,36 @@ class LayerwiseOffloadManager:
self.register_forward_hooks()
logger.info(
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, total num layers: {self.num_layers}"
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, num resident layers: {self.resident_layers}, total num layers: {self.num_layers}"
)
def prepare_for_next_req(self, non_blocking=True):
"""
Prepare for the next round of denoising loop with prefetching the necessary layers
"""
for i in range(self.prefetch_size):
num_prefetch_layers = max(self.prefetch_size, self._retained_layers)
for i in range(num_prefetch_layers):
self.prefetch_layer(i, non_blocking=non_blocking)
if not non_blocking and self.copy_stream is not None:
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
@property
def holds_residents(self) -> bool:
"""True if this manager keeps a resident leading-layer set beyond the
streaming prefetch window, so it must be denoise-stage-scoped."""
return self.enabled and self.resident_layers > 0
@property
def _retained_layers(self) -> int:
"""Leading layers currently held across denoise steps; 0 until armed."""
return self.resident_layers if self._residency_active else 0
@torch.compiler.disable
def _activate_residency(self) -> None:
"""Arm the resident set on the first denoise forward. The pinning itself is
done by the ``prepare_for_next_req`` that follows in the same hook."""
self._residency_active = True
def get_target_with_name(self, name: str) -> torch.Tensor:
"""get the target model weight/buffer to be replaced"""
if name in self._named_parameters:
@@ -333,14 +359,19 @@ class LayerwiseOffloadManager:
self._gpu_layers.add(layer_idx)
@torch.compiler.disable
def release_layer(self, layer_idx: int) -> None:
def release_layer(self, layer_idx: int, force: bool = False) -> None:
"""
lightweight release layer weights
Basically set the reference count to the gpu weight tensor to zero. The weights on cpu is untouched
Leading resident layers are kept across denoise steps
"""
if not self.enabled or self.device is None:
return
if not force and layer_idx < self._retained_layers:
return
# clear prefetch event, since it's useless and needs to be reset
self._prefetch_events.pop(layer_idx, None)
@@ -359,13 +390,15 @@ class LayerwiseOffloadManager:
@torch.compiler.disable
def release_all(self) -> None:
"""Release every layer, including the resident ones: this ends the
denoise stage that the resident set is scoped to."""
if not self.enabled or self.device is None:
return
if self.copy_stream is not None:
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
for layer_idx in list(self._gpu_layers):
self.release_layer(layer_idx)
self.release_layer(layer_idx, force=True)
@torch.compiler.disable
def load_all_layers(self) -> None:
@@ -521,6 +554,7 @@ class LayerwiseOffloadManager:
def make_pre_hook(i):
def hook(module, input):
if i == 0:
self._activate_residency()
self.prepare_for_next_req(non_blocking=False)
if i not in self._gpu_layers:
# LTX audio VAE traverses decoder.up in reverse order
@@ -589,6 +623,14 @@ class LayerwiseOffloadableModuleMixin:
else:
prefetch_size = int(server_args.dit_offload_prefetch_size)
resident_value = server_args.dit_layerwise_resident_layers
if resident_value <= 0:
resident_layers = 0
elif resident_value < 1.0:
resident_layers = max(1, int(round(resident_value * num_layers)))
else:
resident_layers = min(num_layers, int(resident_value))
manager = LayerwiseOffloadManager(
model=self,
layers_attr_str=layer_name,
@@ -596,6 +638,7 @@ class LayerwiseOffloadableModuleMixin:
enabled=True,
pin_cpu_memory=server_args.pin_cpu_memory,
prefetch_size=prefetch_size,
resident_layers=resident_layers,
)
self.layerwise_offload_managers.append(manager)
configured_layer_names.append(layer_name)
@@ -674,6 +717,15 @@ def is_layerwise_offloaded_module(module: torch.nn.Module) -> bool:
)
def is_resident_layerwise_module(module: torch.nn.Module) -> bool:
"""True if the module keeps leading DiT layers resident beyond the streaming
prefetch window.
"""
return isinstance(module, LayerwiseOffloadableModuleMixin) and any(
manager.holds_residents for manager in module.layerwise_offload_managers
)
def get_layerwise_offload_component_names_for_pipeline(
modules: Mapping[str, object],
component_names: Sequence[str] | None = None,
@@ -268,6 +268,8 @@ class ServerArgs(DisaggServerArgsMixin):
dit_layerwise_offload: bool | None = None
layerwise_offload_components: list[str] | None = None
dit_offload_prefetch_size: float = 0.0
# If set, keep this many leading DiT layers resident on GPU
dit_layerwise_resident_layers: float = 0.0
offload_during_compile: bool = True
text_encoder_cpu_offload: bool | None = None
image_encoder_cpu_offload: bool | None = None
@@ -1638,6 +1640,18 @@ class ServerArgs(DisaggServerArgsMixin):
default=ServerArgs.dit_offload_prefetch_size,
help="The size of prefetch for dit-layerwise-offload. If the value is between 0.0 and 1.0, it is treated as a ratio of the total number of layers. If the value is >= 1, it is treated as the absolute number of layers. 0.0 means prefetch 1 layer (lowest memory). Values above 0.5 might have peak memory close to no offload but worse performance.",
)
parser.add_argument(
"--dit-layerwise-resident-layers",
type=float,
default=ServerArgs.dit_layerwise_resident_layers,
help="With --dit-layerwise-offload, keep this many leading DiT layers "
"permanently resident on GPU (retained across denoise steps) and stream "
"only the tail with --dit-offload-prefetch-size. 0.0 = off (pure "
"streaming). Between 0.0 and 1.0 = ratio of layers; >= 1 = absolute "
"count. Unlike raising the prefetch size, resident layers are transferred "
"once (not re-streamed every step), so this trades VRAM for lower denoise "
"latency when memory is available.",
)
# offload flags
parser.add_argument(
@@ -2267,6 +2281,30 @@ class ServerArgs(DisaggServerArgsMixin):
"We do not recommend --dit-offload-prefetch-size to be between 0.5 and 1.0"
)
# validate dit_layerwise_resident_layers (same ratio/absolute convention)
if self.dit_layerwise_resident_layers < 0.0:
raise ValueError("dit_layerwise_resident_layers must be non-negative")
if self.dit_layerwise_resident_layers >= 1 and (
isinstance(self.dit_layerwise_resident_layers, float)
and not self.dit_layerwise_resident_layers.is_integer()
):
self.dit_layerwise_resident_layers = int(
math.floor(self.dit_layerwise_resident_layers)
)
logger.info(
"Invalid --dit-layerwise-resident-layers value passed, truncated to: "
f"{self.dit_layerwise_resident_layers}"
)
if (
self.dit_layerwise_resident_layers > 0
and not self.is_dit_layerwise_offload_selected
):
logger.warning(
"--dit-layerwise-resident-layers has no effect because the DiT is not "
"layerwise-offloaded. It only applies together with "
"--dit-layerwise-offload (or 'dit' in --layerwise-offload-components)."
)
# validate layerwise offload conflicts
if envs.SGLANG_CACHE_DIT_ENABLED and self.use_fsdp_inference:
if self.is_arg_explicitly_set("use_fsdp_inference"):
@@ -31,6 +31,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
configure_layerwise_offload_modules,
get_layerwise_offload_component_names_for_pipeline,
is_layerwise_offloaded_module,
is_resident_layerwise_module,
)
@@ -161,6 +162,7 @@ def _server_args(**kwargs):
image_encoder_cpu_offload=False,
vae_cpu_offload=False,
dit_offload_prefetch_size=1,
dit_layerwise_resident_layers=0.0,
pin_cpu_memory=False,
)
defaults.update(kwargs)
@@ -556,3 +558,134 @@ def test_layerwise_offload_aligns_contiguous_tensor_offsets(monkeypatch):
assert restored_bias.data_ptr() % 32 == 0
assert torch.equal(restored_weight, original_weight)
assert torch.equal(restored_bias, original_bias)
# ---------------------------------------------------------------------------
# --dit-layerwise-resident-layers: keep N leading layers resident (retained
# across denoise steps), streaming only the tail with the prefetch window.
# ---------------------------------------------------------------------------
class _MultiBlockModel(torch.nn.Module):
def __init__(self, n: int) -> None:
super().__init__()
self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)])
class _ResidentComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin):
layer_names = ["blocks"]
def __init__(self, n: int) -> None:
super().__init__()
self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)])
def _patch_fake_device(monkeypatch):
monkeypatch.setattr(
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
)
monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu")
def _resident_manager(model, *, num_layers, prefetch_size=1, resident_layers=0):
return LayerwiseOffloadManager(
model=model,
layers_attr_str="blocks",
num_layers=num_layers,
enabled=True,
pin_cpu_memory=False,
prefetch_size=prefetch_size,
resident_layers=resident_layers,
)
def _arm_residency(manager):
"""Mimic the first-layer pre-hook: arm the resident set, then pin it."""
manager._activate_residency()
manager.prepare_for_next_req(non_blocking=False)
def test_resident_layers_stay_pinned_until_stage_teardown(monkeypatch):
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(4), num_layers=4, prefetch_size=1, resident_layers=2
)
# The resident set is armed on the first forward, not at construction.
assert manager._retained_layers == 0
_arm_residency(manager)
assert manager._retained_layers == 2
assert {0, 1} <= manager._gpu_layers
# A non-force release keeps the leading resident layers pinned across steps.
manager.release_layer(0)
manager.release_layer(1)
assert {0, 1} <= manager._gpu_layers
# force=True (teardown) overrides the retention.
manager.release_layer(0, force=True)
assert 0 not in manager._gpu_layers
manager.release_all() # ends the denoise stage: residents go too
assert not manager._gpu_layers
def test_resident_layers_off_by_default_streams_everything(monkeypatch):
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(4), num_layers=4, prefetch_size=1, resident_layers=0
)
_arm_residency(manager)
assert manager._retained_layers == 0
assert manager.holds_residents is False
manager.prefetch_layer(2, non_blocking=False)
manager.release_layer(2) # no residents -> released like plain streaming
assert 2 not in manager._gpu_layers
def test_prepare_for_next_req_repins_residents(monkeypatch):
_patch_fake_device(monkeypatch)
manager = _resident_manager(
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=3
)
_arm_residency(manager)
manager.release_all()
assert not manager._gpu_layers
# The next denoise re-pins the resident set (union of prefetch window + residents).
manager.prepare_for_next_req(non_blocking=False)
assert {0, 1, 2} <= manager._gpu_layers
def test_holds_residents_reflects_configuration(monkeypatch):
_patch_fake_device(monkeypatch)
resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)
streaming = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=0)
assert resident.holds_residents is True
assert streaming.holds_residents is False
def test_is_resident_layerwise_module_detector():
class _Comp(torch.nn.Module, LayerwiseOffloadableModuleMixin):
pass
comp = _Comp()
comp.layerwise_offload_managers = [SimpleNamespace(holds_residents=True)]
assert is_resident_layerwise_module(comp) is True
comp.layerwise_offload_managers = [SimpleNamespace(holds_residents=False)]
assert is_resident_layerwise_module(comp) is False
def test_configure_resolves_resident_layers_absolute(monkeypatch):
_patch_fake_device(monkeypatch)
comp = _ResidentComponent(8)
comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=3))
assert comp.layerwise_offload_managers[0].resident_layers == 3
def test_configure_resolves_resident_layers_ratio(monkeypatch):
_patch_fake_device(monkeypatch)
comp = _ResidentComponent(8)
comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=0.5))
# 0.5 * 8 = 4 leading layers resident
assert comp.layerwise_offload_managers[0].resident_layers == 4