diff --git a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py index 020746b6f..75fdbf265 100644 --- a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py @@ -1,6 +1,7 @@ import re from contextlib import contextmanager -from typing import Dict, Set +from itertools import chain +from typing import Any, Dict, List, Set, Tuple import torch @@ -45,15 +46,20 @@ class LayerwiseOffloadManager: rf"(^|\.){re.escape(module_list_attr)}\.(\d+)(\.|$)" ) - self._cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {} - self._cpu_dtypes: Dict[int, Dict[str, torch.dtype]] = {} - self._gpu_layers: Dict[int, Set[str]] = {} + # layer_idx -> {dtype: consolidated_pinned_cpu_tensor} + # stores the consolidated weight from a same layer, of same dtype + self._consolidated_cpu_weights: Dict[int, Dict[torch.dtype, torch.Tensor]] = {} + # layer_idx -> {name: {dtype, offset, numel, shape}} + # stores the offset and numel of each weight from a same layer, of same dtype + self._weight_metadata: Dict[int, Dict[str, Dict[str, Any]]] = {} + # layer indices that are already in gpu + self._gpu_layers: Set[int] = set() self._named_parameters: Dict[str, torch.nn.Parameter] = {} self._named_buffers: Dict[str, torch.Tensor] = {} if auto_initialize: - self.initialize() + self._initialize() def _match_layer_idx(self, name: str) -> int | None: m = self._layer_name_re.search(name) @@ -64,40 +70,59 @@ class LayerwiseOffloadManager: except Exception: return None - def _offload_tensor(self, name: str, tensor: torch.Tensor, layer_idx: int) -> None: - if layer_idx not in self._cpu_weights: - self._cpu_weights[layer_idx] = {} - self._cpu_dtypes[layer_idx] = {} - - cpu_weight = tensor.detach().to("cpu") - if self.pin_cpu_memory: - cpu_weight = cpu_weight.pin_memory() - self._cpu_weights[layer_idx][name] = cpu_weight - self._cpu_dtypes[layer_idx][name] = tensor.dtype - - if self.device is not None: - tensor.data = torch.empty((1,), device=self.device, dtype=tensor.dtype) - @torch.compiler.disable - def initialize(self) -> None: + def _initialize(self) -> None: if not self.enabled: return self._named_parameters = dict(self.model.named_parameters()) self._named_buffers = dict(self.model.named_buffers()) - for name, param in self._named_parameters.items(): + # 1. collect and group tensors by layer and dtype + layer_groups: Dict[int, Dict[torch.dtype, List[Tuple[str, torch.Tensor]]]] = {} + all_tensors = chain(self._named_parameters.items(), self._named_buffers.items()) + for name, tensor in all_tensors: layer_idx = self._match_layer_idx(name) if layer_idx is None or layer_idx >= self.num_layers: continue - self._offload_tensor(name, param, layer_idx) + layer_groups.setdefault(layer_idx, {}).setdefault(tensor.dtype, []).append( + (name, tensor) + ) - for name, buf in self._named_buffers.items(): - layer_idx = self._match_layer_idx(name) - if layer_idx is None or layer_idx >= self.num_layers: - continue - self._offload_tensor(name, buf, layer_idx) + # 2. concat and offload (in pinned memory) + for layer_idx, dtype_to_params in layer_groups.items(): + self._consolidated_cpu_weights[layer_idx] = {} + self._weight_metadata[layer_idx] = {} + for dtype, weights in dtype_to_params.items(): + total_numel = sum(t.numel() for _, t in weights) + + # create concatenated CPU buffer (in pinned memory) + cpu_buffer = torch.empty( + total_numel, dtype=dtype, pin_memory=self.pin_cpu_memory + ) + + # offload weights to the buffer + current_offset = 0 + for name, weight in weights: + numel = weight.numel() + cpu_buffer[current_offset : current_offset + numel].copy_( + weight.flatten() + ) + self._weight_metadata[layer_idx][name] = { + "dtype": dtype, + "offset": current_offset, + "numel": numel, + "shape": weight.shape, + } + + weight.data = torch.empty((1,), device=self.device, dtype=dtype) + + current_offset += numel + + self._consolidated_cpu_weights[layer_idx][dtype] = cpu_buffer + + # prefetch the first layer for warm-up self.prepare_for_next_denoise(non_blocking=False) def prepare_for_next_denoise(self, non_blocking=True): @@ -105,6 +130,14 @@ class LayerwiseOffloadManager: if not non_blocking and self.copy_stream is not None: torch.cuda.current_stream().wait_stream(self.copy_stream) + 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: + target = self._named_parameters[name] + else: + target = self._named_buffers[name] + return target + @torch.compiler.disable def prefetch_layer(self, layer_idx: int, non_blocking: bool = True) -> None: if not self.enabled or self.device is None or self.copy_stream is None: @@ -113,30 +146,33 @@ class LayerwiseOffloadManager: return if layer_idx in self._gpu_layers: return - if layer_idx not in self._cpu_weights: + if layer_idx not in self._consolidated_cpu_weights: return self.copy_stream.wait_stream(torch.cuda.current_stream()) - param_names: Set[str] = set() + # create gpu buffer and load from CPU buffer + gpu_buffers: Dict[torch.dtype, torch.Tensor] = {} + with torch.cuda.stream(self.copy_stream): + for dtype, cpu_buffer in self._consolidated_cpu_weights[layer_idx].items(): + gpu_buffer = torch.empty( + cpu_buffer.shape, dtype=dtype, device=self.device + ) + gpu_buffer.copy_(cpu_buffer, non_blocking=non_blocking) + gpu_buffers[dtype] = gpu_buffer - for name, cpu_weight in self._cpu_weights[layer_idx].items(): - if name in self._named_parameters: - target = self._named_parameters[name] - else: - target = self._named_buffers[name] + # restore model's weights by their metadata using gpu buffer + for name, meta in self._weight_metadata[layer_idx].items(): + dtype = meta["dtype"] + gpu_buffer = gpu_buffers[dtype] - gpu_weight = torch.empty( - cpu_weight.shape, - dtype=self._cpu_dtypes[layer_idx][name], - device=self.device, - ) - with torch.cuda.stream(self.copy_stream): - gpu_weight.copy_(cpu_weight, non_blocking=non_blocking) - target.data = gpu_weight - param_names.add(name) + # map the parameter's data to the correct slice of the GPU buffer + target = self.get_target_with_name(name) + target.data = gpu_buffer[ + meta["offset"] : meta["offset"] + meta["numel"] + ].view(meta["shape"]) - self._gpu_layers[layer_idx] = param_names + self._gpu_layers.add(layer_idx) @contextmanager def layer_scope( @@ -169,16 +205,14 @@ class LayerwiseOffloadManager: if layer_idx <= 0: return - param_names = self._gpu_layers.pop(layer_idx, None) - if not param_names: + if layer_idx not in self._gpu_layers: return - for name in param_names: - if name in self._named_parameters: - target = self._named_parameters[name] - else: - target = self._named_buffers[name] - target.data = torch.empty((1,), device=self.device, dtype=target.dtype) + for name, meta in self._weight_metadata.get(layer_idx, {}).items(): + target = self.get_target_with_name(name) + target.data = torch.empty((1,), device=self.device, dtype=meta["dtype"]) + + self._gpu_layers.discard(layer_idx) @torch.compiler.disable def release_all(self) -> None: @@ -187,14 +221,5 @@ class LayerwiseOffloadManager: if self.copy_stream is not None: torch.cuda.current_stream().wait_stream(self.copy_stream) - layer_indices = list(self._gpu_layers.keys()) - for layer_idx in layer_indices: - param_names = self._gpu_layers.pop(layer_idx, None) - if not param_names: - continue - for name in param_names: - if name in self._named_parameters: - target = self._named_parameters[name] - else: - target = self._named_buffers[name] - target.data = torch.empty((1,), device=self.device, dtype=target.dtype) + for layer_idx in list(self._gpu_layers): + self.release_layer(layer_idx)