From 227dadd79a17dc16c2a5e0c33ad7ea68f1e7093f Mon Sep 17 00:00:00 2001 From: Yihao Wang <42559837+AgainstEntropy@users.noreply.github.com> Date: Wed, 29 Jul 2026 01:52:55 -0700 Subject: [PATCH] [diffusion] feat: support resident layers for DiT (#31538) --- .../memory_managers/component_manager.py | 14 +- .../memory_managers/layerwise_offload.py | 60 +++++++- .../runtime/server_args/server_args.py | 38 +++++ .../test/unit/test_layerwise_offload.py | 133 ++++++++++++++++++ 4 files changed, 240 insertions(+), 5 deletions(-) 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 dd225ec0e..9d23a5b08 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 @@ -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 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 95b66392c..372a3934c 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 @@ -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, 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 987d3b149..4c3b4c922 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -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"): 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 35e4f1b94..0053793f0 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -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