diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py index 9cccc78bc..ee60f09e1 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_residency_strategies.py @@ -210,6 +210,7 @@ class LayerwiseOffloadStrategy(ComponentResidencyStrategy): return _module_to_local_device(module, dtype=use.target_dtype) return + module.restore_non_layer_weights() module.prepare_for_next_req() def finish_use( @@ -222,6 +223,10 @@ class LayerwiseOffloadStrategy(ComponentResidencyStrategy): return for manager in module.layerwise_offload_managers: manager.release_all() + # The layers are gone; the rest of this component is dead weight on the + # device until it is used again, and the stage that follows may be the + # one that needs the room. + module.park_non_layer_weights() if current_platform.is_mps(): torch.mps.synchronize() module.restore_mps_cpu_non_layer_weights() 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 eb416647d..7234e1b31 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 @@ -97,6 +97,13 @@ def compute_streamed_layers( # of table size to rows actually read makes residency clearly wasteful. HOST_RESIDENT_TABLE_MIN_BYTES = 256 * 1024**2 +# Parking a component's non-layer weights frees device memory at the cost of two +# transfers per use and a host copy that competes with the page cache. It is +# worth that only when what it frees is a meaningful share of the headroom +# actually available; on a card with room it is pure loss. Below this share of +# free device memory, the component stays where it is. +PARK_SIGNIFICANCE = 0.1 + def _resolve_submodule(root: torch.nn.Module, path: str) -> torch.nn.Module | None: current: Any = root @@ -1430,6 +1437,106 @@ class LayerwiseOffloadableModuleMixin: host_resident_table_names: List[str] = [] layerwise_offload_managers: list[LayerwiseOffloadManager] = [] + # Whether to park non-layer parameters on the host between uses. Costs a + # transfer per request and is worth it only when device memory is the + # binding constraint, so it follows --performance-mode memory. + park_non_layer_weights_between_uses: bool = False + + def _managed_layer_parameter_names(self) -> set: + """Parameter names some layerwise manager already streams.""" + return { + name + for manager in self.layerwise_offload_managers + for names in manager._weight_metadata.values() + for name in names + } + + def park_non_layer_weights(self) -> None: + """Move the parameters no manager streams back to the host. + + A layerwise component holds its non-layer parameters on the device for + the whole request. That is right while it is the component being used + and pure cost afterwards. Measured on H3 at 864x480 / 124 frames: the + DiT keeps 2.09 GB and the text encoder 1.40 GB through a VAE decode + that touches neither, and the decode is exactly where the budget runs + out -- with the VAE's blocks held resident it needs 11.86 GiB against a + 12 GiB card, and fails for want of 20 MiB. + + Buffers are left where they are. Layerwise offload keeps them resident + on purpose, because a shared buffer such as a RoPE cache is referenced + by many layers. + """ + if not self.park_non_layer_weights_between_uses: + return + if current_platform.is_mps(): + # MPS parks its own non-layer weights, scoped to subphases + return + managed = self._managed_layer_parameter_names() + resident = [ + (name, parameter) + for name, parameter in self.named_parameters() + if name not in managed and parameter.device.type != "cpu" + ] + holds = sum(p.numel() * p.element_size() for _, p in resident) + if holds <= self._device_headroom_bytes() * PARK_SIGNIFICANCE: + # There is room. Give back any host copies rather than hold them. + self._parked_non_layer_weights.clear() + return + + parked = self._parked_non_layer_weights + with torch.inference_mode(False), torch.no_grad(): + for name, parameter in resident: + if name not in parked: + parked[name] = parameter.detach().to("cpu", copy=True) + parameter.data = self._park_placeholder(parameter) + + def _device_headroom_bytes(self) -> int: + """What an allocation could get without the allocator growing its pool. + + `get_available_gpu_memory` reports driver-level free memory, which + excludes blocks the caching allocator has already reserved and not + handed out. On a warm process that undercounts the real headroom badly, + so the allocator's own unused reserve is added back. + """ + free = int( + current_platform.get_available_gpu_memory(empty_cache=False) * (1 << 30) + ) + device_module = torch.get_device_module() + unused_reserve = ( + device_module.memory_reserved() - device_module.memory_allocated() + ) + return free + max(0, unused_reserve) + + def _park_placeholder(self, parameter: torch.Tensor) -> torch.Tensor: + """One shared stand-in per (device, dtype), not one per parked weight.""" + key = (parameter.device, parameter.dtype) + placeholder = self._park_placeholders.get(key) + if placeholder is None: + placeholder = torch.empty( + (1,), dtype=parameter.dtype, device=parameter.device + ) + self._park_placeholders[key] = placeholder + return placeholder + + def restore_non_layer_weights(self) -> None: + """Bring parked parameters back before this component is used again.""" + parked = self._parked_non_layer_weights + if not parked: + return + device = current_platform.get_local_torch_device() + parameters = dict(self.named_parameters()) + with torch.inference_mode(False), torch.no_grad(): + for name, host_tensor in parked.items(): + parameter = parameters.get(name) + if parameter is None: + continue + # The parked copy is pageable, so this transfer stages through + # the driver's own pinned buffer and is synchronous whatever is + # asked for. Pinning it instead would make the copy async, at + # the price of host memory the kernel can never reclaim -- the + # wrong trade on the hosts this path exists for. + parameter.data = host_tensor.to(device) + def _capture_mps_cpu_non_layer_weights(self) -> None: managed_names = { name @@ -1522,6 +1629,22 @@ class LayerwiseOffloadableModuleMixin: for name, tensor in self._mps_cpu_buffers.items(): buffers[name].data = tensor + @property + def _parked_non_layer_weights(self) -> dict: + store = self.__dict__.get("_parked_non_layer_weight_store") + if store is None: + store = {} + self.__dict__["_parked_non_layer_weight_store"] = store + return store + + @property + def _park_placeholders(self) -> dict: + store = self.__dict__.get("_park_placeholder_store") + if store is None: + store = {} + self.__dict__["_park_placeholder_store"] = store + return store + def configure_layerwise_offload( self, server_args: ServerArgs, @@ -1529,6 +1652,9 @@ class LayerwiseOffloadableModuleMixin: pin_budget: HostPinBudget | None = None, component_name: str | None = None, ): + self.park_non_layer_weights_between_uses = ( + server_args.performance_mode == "memory" + ) self.layerwise_offload_managers = [] named_modules = dict(self.named_modules()) configured_layer_names = [] 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 51dee69d0..7f76c6ed2 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -205,6 +205,7 @@ def _server_args(**kwargs): defaults = dict( component_residency=None, disagg_role=RoleType.MONOLITHIC, + performance_mode="auto", _required_resident_components=set(), _component_layerwise_capabilities={}, _explicit_arg_names=set(), @@ -797,6 +798,12 @@ class _ResidentComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin): self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)]) +class _ParkableResidentComponent(_ResidentComponent): + def __init__(self, n: int) -> None: + super().__init__(n) + self.non_layer = torch.nn.Parameter(torch.ones(2)) + + class _AuxiliaryResidentComponent(_ResidentComponent): layerwise_offload_dit_group_enabled = False @@ -1684,3 +1691,97 @@ def test_layerwise_tuning_accepts_json_and_pair_forms(): assert pair.layerwise_tuning_for("text_encoder", dit_group=False)[1] == 2.0 as_json = _server_args(layerwise_resident_layers='{"vae": 6}') assert as_json.layerwise_tuning_for("vae", dit_group=False)[1] == 6.0 + + +def test_non_layer_parking_follows_memory_performance_mode(monkeypatch): + """The extra transfer per request only pays for itself under memory mode.""" + _patch_fake_device(monkeypatch) + tight = _ResidentComponent(4) + tight.configure_layerwise_offload(_server_args(performance_mode="memory")) + assert tight.park_non_layer_weights_between_uses + + relaxed = _ResidentComponent(4) + relaxed.configure_layerwise_offload(_server_args(performance_mode="speed")) + assert not relaxed.park_non_layer_weights_between_uses + + +def test_parking_leaves_streamed_layer_weights_alone(monkeypatch): + """Only the parameters no manager streams are moved to the host.""" + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="memory")) + _headroom(monkeypatch, 0) + managed = comp._managed_layer_parameter_names() + assert managed, "the managers should own the block parameters" + + comp.park_non_layer_weights() + parked = comp._parked_non_layer_weights + assert not (set(parked) & managed), "a streamed layer weight was parked" + for name, host_tensor in parked.items(): + assert host_tensor.device.type == "cpu", name + + comp.restore_non_layer_weights() + restored = dict(comp.named_parameters()) + for name, host_tensor in parked.items(): + assert restored[name].shape == host_tensor.shape + + +def test_parking_is_a_no_op_outside_memory_mode(monkeypatch): + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="speed")) + comp.park_non_layer_weights() + assert not comp._parked_non_layer_weights + + +def _headroom(monkeypatch, gib): + monkeypatch.setattr( + layerwise_offload_mod.current_platform, + "get_available_gpu_memory", + lambda **_: float(gib), + ) + module = layerwise_offload_mod.torch.get_device_module() + monkeypatch.setattr(module, "memory_reserved", lambda *_: 0, raising=False) + monkeypatch.setattr(module, "memory_allocated", lambda *_: 0, raising=False) + + +def test_parking_is_skipped_when_the_card_has_room(monkeypatch): + """A component holding a sliver of a large headroom is left alone.""" + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="memory")) + _headroom(monkeypatch, 400) + comp.park_non_layer_weights() + assert not comp._parked_non_layer_weights + + +def test_parking_happens_when_the_headroom_is_small(monkeypatch): + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="memory")) + _headroom(monkeypatch, 0) + comp.park_non_layer_weights() + assert comp._parked_non_layer_weights + + +def test_host_copies_are_given_back_when_room_appears(monkeypatch): + """Skipping must not leave host memory held for a park that will not happen.""" + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="memory")) + _headroom(monkeypatch, 0) + comp.park_non_layer_weights() + assert comp._parked_non_layer_weights + comp.restore_non_layer_weights() + + _headroom(monkeypatch, 400) + comp.park_non_layer_weights() + assert not comp._parked_non_layer_weights, "host copies should be released" + + +def test_park_placeholders_are_shared(monkeypatch): + """One stand-in per (device, dtype), not one allocation per parked weight.""" + comp = _ParkableResidentComponent(4) + comp.configure_layerwise_offload(_server_args(performance_mode="memory")) + _headroom(monkeypatch, 0) + comp.park_non_layer_weights() + managed = comp._managed_layer_parameter_names() + stand_ins = { + id(p) for n, p in comp.named_parameters() if n not in managed and p.numel() == 1 + } + assert len(stand_ins) <= len(comp._park_placeholders)