[diffusion] feat: support resident layers for DiT (#31538)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user