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 19c566c2d..501bad456 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 @@ -205,6 +205,16 @@ class LayerwiseOffloadManager: except Exception: return None + def _managed_parameter_bytes(self) -> int: + total_bytes = 0 + for name, tensor in self.model.named_parameters(): + layer_idx = self._match_layer_idx(name) + if layer_idx is None or layer_idx >= self.num_layers: + continue + local_tensor = self._to_local_tensor(tensor) + total_bytes += local_tensor.numel() * local_tensor.element_size() + return total_bytes + def _get_shared_empty_tensor(self, dtype: torch.dtype) -> torch.Tensor: placeholder = self._offload_placeholders.get(dtype) if placeholder is None: @@ -252,13 +262,26 @@ class LayerwiseOffloadManager: if not self.enabled: return - self._named_parameters = dict(self.model.named_parameters()) - self._named_buffers = dict(self.model.named_buffers()) - if self._synchronous_mps: + self._named_parameters = dict(self.model.named_parameters()) + self._named_buffers = dict(self.model.named_buffers()) self._initialize_mps_cpu_weights() return + self._initialize_layer_weights() + + # Keep non-layer parameters resident on GPU. Layer tensors have already + # been replaced by tiny device placeholders, so this does not reload the + # offloaded layer weights. + if not self._has_dtensor_weights: + self.model.to(self.device) + + self._finalize_initialization() + + def _initialize_layer_weights(self) -> None: + self._named_parameters = dict(self.model.named_parameters()) + self._named_buffers = dict(self.model.named_buffers()) + # 1. collect and group layer parameters by dtype. Keep buffers resident: # shared buffers such as RoPE caches may be referenced by many layers. layer_groups: Dict[int, Dict[torch.dtype, List[Tuple[str, torch.Tensor]]]] = {} @@ -353,12 +376,7 @@ class LayerwiseOffloadManager: self._consolidated_cpu_weights[layer_idx][dtype] = cpu_buffer - # Keep non-layer parameters resident on GPU. Layer tensors have already - # been replaced by tiny device placeholders, so this does not reload the - # offloaded layer weights. - if not self._has_dtensor_weights: - self.model.to(self.device) - + def _finalize_initialization(self) -> None: # prefetch the head of the stream for warm-up; residency is not armed # yet, so this is layer 0 regardless of policy self.prepare_for_next_req(non_blocking=False) @@ -1012,7 +1030,7 @@ class LayerwiseOffloadableModuleMixin: pin_cpu_memory=server_args.pin_cpu_memory, prefetch_size=prefetch_size, resident_layers=resident_layers, - initialize=not current_platform.is_mps(), + initialize=False, residency_policy=( server_args.dit_layerwise_residency_policy if dit_tuning_enabled @@ -1026,6 +1044,33 @@ class LayerwiseOffloadableModuleMixin: for manager in self.layerwise_offload_managers: manager.initialize() self._capture_mps_cpu_non_layer_weights() + else: + enabled_managers = [ + manager + for manager in self.layerwise_offload_managers + if manager.enabled + ] + initialization_order = sorted( + enabled_managers, + key=lambda manager: manager._managed_parameter_bytes(), + reverse=True, + ) + # release the largest managed groups first when checkpoint loading + # already placed weights on the accelerator; keep the stored manager + # order unchanged for prefetch and forward lifecycle semantics + for manager in initialization_order: + manager._initialize_layer_weights() + + # Every managed layer group must be replaced before moving the + # remaining parameters, otherwise an earlier manager transiently + # moves later groups to the device. + if enabled_managers and not any( + manager._has_dtensor_weights for manager in enabled_managers + ): + self.to(enabled_managers[0].device) + + for manager in enabled_managers: + manager._finalize_initialization() if configured_layer_names: logger.debug( 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 bde610dd2..53b0d9a06 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -749,6 +749,25 @@ class _AuxiliaryResidentComponent(_ResidentComponent): layerwise_offload_dit_group_enabled = False +class _MultiGroupComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin): + layer_names = ["small_blocks", "large_blocks"] + + def __init__(self) -> None: + super().__init__() + self.small_blocks = torch.nn.ModuleList([_DummyBlock()]) + self.large_blocks = torch.nn.ModuleList( + [_DummyBlock(), _DummyBlock(), _DummyBlock()] + ) + self.non_layer = torch.nn.Parameter(torch.ones(2)) + self.to_parameter_shapes = [] + + def to(self, *args, **kwargs): + self.to_parameter_shapes.append( + {name: tuple(param.shape) for name, param in self.named_parameters()} + ) + return super().to(*args, **kwargs) + + def _patch_fake_device(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule @@ -916,6 +935,36 @@ def test_configure_resolves_residency_policy(monkeypatch): ) +def test_configure_offloads_all_layer_groups_before_moving_non_layers(monkeypatch): + _patch_fake_device(monkeypatch) + model = _MultiGroupComponent() + initialization_order = [] + initialize_layer_weights = LayerwiseOffloadManager._initialize_layer_weights + + def record_initialization(manager): + initialization_order.append(manager.layers_attr_str) + initialize_layer_weights(manager) + + monkeypatch.setattr( + LayerwiseOffloadManager, + "_initialize_layer_weights", + record_initialization, + ) + + model.configure_layerwise_offload(_server_args()) + + assert initialization_order == ["large_blocks", "small_blocks"] + assert [ + manager.layers_attr_str for manager in model.layerwise_offload_managers + ] == ["small_blocks", "large_blocks"] + assert len(model.to_parameter_shapes) == 1 + shapes_at_move = model.to_parameter_shapes[0] + assert shapes_at_move["non_layer"] == (2,) + for name, shape in shapes_at_move.items(): + if name != "non_layer": + assert shape == (1,), name + + def test_holds_residents_reflects_configuration(monkeypatch): _patch_fake_device(monkeypatch) resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)