[diffusion] feat: allow offloaded weights stay on the checkpoint mapping (#35701)
This commit is contained in:
+43
-5
@@ -39,6 +39,10 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import (
|
|||||||
pt_weights_iterator,
|
pt_weights_iterator,
|
||||||
safetensors_weights_iterator,
|
safetensors_weights_iterator,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
||||||
|
host_copies_would_not_fit,
|
||||||
|
host_memory_available_bytes,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.encoders.base import (
|
from sglang.multimodal_gen.runtime.models.encoders.base import (
|
||||||
EncoderTensorParallelMixin,
|
EncoderTensorParallelMixin,
|
||||||
TextEncoder,
|
TextEncoder,
|
||||||
@@ -205,6 +209,34 @@ def _process_quantized_encoder_weights(
|
|||||||
return processed_layers
|
return processed_layers
|
||||||
|
|
||||||
|
|
||||||
|
def _checkpoint_bytes(model_path: str) -> int:
|
||||||
|
"""On-disk size of a checkpoint, readable before any weight of it is."""
|
||||||
|
total = 0
|
||||||
|
for path in glob.glob(
|
||||||
|
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
total += os.path.getsize(path)
|
||||||
|
except OSError:
|
||||||
|
continue
|
||||||
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
def _keep_this_checkpoint_mapped(model_path: str) -> bool:
|
||||||
|
"""Whether this encoder's weights should stay on their file mapping."""
|
||||||
|
checkpoint_bytes = _checkpoint_bytes(model_path)
|
||||||
|
if not host_copies_would_not_fit(checkpoint_bytes):
|
||||||
|
return False
|
||||||
|
logger.info(
|
||||||
|
"Text encoder checkpoint is %.2f GiB against %.2f GiB of host memory, "
|
||||||
|
"so its compatible weights stay on the checkpoint mapping instead of "
|
||||||
|
"being copied in.",
|
||||||
|
checkpoint_bytes / 1024**3,
|
||||||
|
host_memory_available_bytes() / 1024**3,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
class TextEncoderLoader(ComponentLoader):
|
class TextEncoderLoader(ComponentLoader):
|
||||||
"""Loader for text encoders."""
|
"""Loader for text encoders."""
|
||||||
|
|
||||||
@@ -613,11 +645,17 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
)
|
)
|
||||||
model.bind_encoder_tp_group(encoder_tp_group)
|
model.bind_encoder_tp_group(encoder_tp_group)
|
||||||
|
|
||||||
if current_platform.is_mps() and component_starts_on_cpu:
|
if component_starts_on_cpu and (
|
||||||
# the h3 encoder is layered immediately after this loader returns
|
current_platform.is_mps() or _keep_this_checkpoint_mapped(model_path)
|
||||||
# compatible CPU safetensors stay mapped instead of copying the
|
):
|
||||||
# full Qwen checkpoint into unified memory
|
# The encoder is layered immediately after this loader returns,
|
||||||
model._mps_zero_copy_weight_loading = True
|
# so compatible CPU safetensors can stay mapped instead of being
|
||||||
|
# copied. On MPS that is always the right call -- the memory is
|
||||||
|
# unified. On any host it becomes the only call once the
|
||||||
|
# checkpoint is larger than host memory, because the copy is
|
||||||
|
# what does not fit: H3's encoder is 62.13 GiB against a 32 GiB
|
||||||
|
# target.
|
||||||
|
model._keep_checkpoint_mapping = True
|
||||||
|
|
||||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||||
loaded_weights = model.load_weights(
|
loaded_weights = model.load_weights(
|
||||||
|
|||||||
@@ -299,15 +299,15 @@ def maybe_load_fsdp_model(
|
|||||||
logger.info("Disabling FSDP for MPS platform as it's not compatible")
|
logger.info("Disabling FSDP for MPS platform as it's not compatible")
|
||||||
|
|
||||||
weight_load_plan = weight_load_plan or WeightLoadPlan(checkpoint_load_device=device)
|
weight_load_plan = weight_load_plan or WeightLoadPlan(checkpoint_load_device=device)
|
||||||
mps_zero_copy_weight_loading = bool(
|
keep_checkpoint_mapping = bool(
|
||||||
current_platform.is_mps()
|
current_platform.is_mps()
|
||||||
and weight_load_plan.mps_layerwise_cpu_staging
|
and weight_load_plan.mps_layerwise_cpu_staging
|
||||||
and weight_load_plan.checkpoint_load_device.type == "cpu"
|
and weight_load_plan.checkpoint_load_device.type == "cpu"
|
||||||
)
|
)
|
||||||
if mps_zero_copy_weight_loading:
|
if keep_checkpoint_mapping:
|
||||||
# layerwise offload replaces block parameters with mps placeholders after
|
# layerwise offload replaces block parameters with placeholders after
|
||||||
# load, so compatible checkpoint tensors stay file-backed on CPU
|
# load, so compatible checkpoint tensors stay file-backed on CPU
|
||||||
model._mps_zero_copy_weight_loading = True
|
model._keep_checkpoint_mapping = True
|
||||||
defer_cpu_placement = bool(
|
defer_cpu_placement = bool(
|
||||||
component_starts_on_cpu
|
component_starts_on_cpu
|
||||||
and weight_load_plan.defer_cpu_placement
|
and weight_load_plan.defer_cpu_placement
|
||||||
@@ -425,7 +425,7 @@ def maybe_load_fsdp_model(
|
|||||||
strict=strict,
|
strict=strict,
|
||||||
cpu_offload=load_on_cpu,
|
cpu_offload=load_on_cpu,
|
||||||
param_names_mapping=param_names_mapping_fn,
|
param_names_mapping=param_names_mapping_fn,
|
||||||
mps_zero_copy_weight_loading=mps_zero_copy_weight_loading,
|
keep_checkpoint_mapping=keep_checkpoint_mapping,
|
||||||
preconverted_state_dict=preconverted_state_dict,
|
preconverted_state_dict=preconverted_state_dict,
|
||||||
)
|
)
|
||||||
if bnb_quant_states:
|
if bnb_quant_states:
|
||||||
@@ -549,7 +549,7 @@ def load_model_from_full_model_state_dict(
|
|||||||
strict: bool = False,
|
strict: bool = False,
|
||||||
cpu_offload: bool = False,
|
cpu_offload: bool = False,
|
||||||
param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None,
|
param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None,
|
||||||
mps_zero_copy_weight_loading: bool = False,
|
keep_checkpoint_mapping: bool = False,
|
||||||
preconverted_state_dict: (
|
preconverted_state_dict: (
|
||||||
tuple[
|
tuple[
|
||||||
dict[
|
dict[
|
||||||
@@ -574,7 +574,7 @@ def load_model_from_full_model_state_dict(
|
|||||||
strict (bool): flag to check if to load the model in strict mode
|
strict (bool): flag to check if to load the model in strict mode
|
||||||
cpu_offload (bool): flag to check if FSDP offload is enabled
|
cpu_offload (bool): flag to check if FSDP offload is enabled
|
||||||
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
|
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
|
||||||
mps_zero_copy_weight_loading (bool): retain compatible CPU checkpoint tensors for MPS layerwise offload
|
keep_checkpoint_mapping (bool): retain compatible CPU checkpoint tensors instead of copying them
|
||||||
Returns:
|
Returns:
|
||||||
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
|
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
|
||||||
* **missing_keys** is a list of str containing the missing keys
|
* **missing_keys** is a list of str containing the missing keys
|
||||||
@@ -717,9 +717,9 @@ def load_model_from_full_model_state_dict(
|
|||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
use_checkpoint_tensor_directly = bool(
|
use_checkpoint_tensor_directly = bool(
|
||||||
mps_zero_copy_weight_loading
|
keep_checkpoint_mapping
|
||||||
and actual_param is not None
|
and actual_param is not None
|
||||||
and not getattr(actual_param, "mps_zero_copy_unsafe", False)
|
and not getattr(actual_param, "checkpoint_mapping_unsafe", False)
|
||||||
and tuple(meta_sharded_param.shape) == tuple(full_tensor.shape)
|
and tuple(meta_sharded_param.shape) == tuple(full_tensor.shape)
|
||||||
and full_tensor.device.type == "cpu"
|
and full_tensor.device.type == "cpu"
|
||||||
and full_tensor.dtype == target_dtype
|
and full_tensor.dtype == target_dtype
|
||||||
|
|||||||
@@ -357,6 +357,44 @@ def _read_process_mappings() -> tuple[list[int], list[int], list[bool]] | None:
|
|||||||
return [r[0] for r in rows], [r[1] for r in rows], [r[2] for r in rows]
|
return [r[0] for r in rows], [r[1] for r in rows], [r[2] for r in rows]
|
||||||
|
|
||||||
|
|
||||||
|
class MappedRegions:
|
||||||
|
"""Answers whether a tensor's bytes live in a file mapping.
|
||||||
|
|
||||||
|
Built once and reused. The lookup table comes from /proc/self/maps, so
|
||||||
|
rebuilding it per tensor would be quadratic over a checkpoint's worth of
|
||||||
|
weights -- H3's DiT alone has tens of thousands.
|
||||||
|
|
||||||
|
A snapshot, not a live view: mappings created after construction are
|
||||||
|
unknown to it. Callers that need to classify freshly loaded weights should
|
||||||
|
build one after loading, which is when the mappings exist.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._maps = _read_process_mappings()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def available(self) -> bool:
|
||||||
|
"""False where /proc is absent, in which case nothing is classified."""
|
||||||
|
return self._maps is not None
|
||||||
|
|
||||||
|
def holds_pointer(self, pointer: int) -> bool:
|
||||||
|
if self._maps is None or pointer == 0:
|
||||||
|
return False
|
||||||
|
starts, ends, backed = self._maps
|
||||||
|
index = bisect.bisect_right(starts, pointer) - 1
|
||||||
|
if index < 0 or pointer >= ends[index]:
|
||||||
|
return False
|
||||||
|
return backed[index]
|
||||||
|
|
||||||
|
def holds(self, tensor: torch.Tensor) -> bool:
|
||||||
|
if tensor.device.type != "cpu":
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
return self.holds_pointer(tensor.untyped_storage().data_ptr())
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def component_residency_bytes(module) -> Dict[str, int]:
|
def component_residency_bytes(module) -> Dict[str, int]:
|
||||||
"""Where a component's weights actually sit, in bytes.
|
"""Where a component's weights actually sit, in bytes.
|
||||||
|
|
||||||
@@ -380,16 +418,10 @@ def component_residency_bytes(module) -> Dict[str, int]:
|
|||||||
|
|
||||||
totals = {"vram": 0, "host_pinned": 0, "host_mapped": 0, "host": 0}
|
totals = {"vram": 0, "host_pinned": 0, "host_mapped": 0, "host": 0}
|
||||||
seen: set[int] = set()
|
seen: set[int] = set()
|
||||||
mappings = _read_process_mappings()
|
regions = MappedRegions()
|
||||||
|
|
||||||
def is_file_backed(pointer: int) -> bool:
|
def is_file_backed(pointer: int) -> bool:
|
||||||
if mappings is None:
|
return regions.holds_pointer(pointer)
|
||||||
return False
|
|
||||||
starts, ends, backed = mappings
|
|
||||||
index = bisect.bisect_right(starts, pointer) - 1
|
|
||||||
if index < 0 or pointer >= ends[index]:
|
|
||||||
return False
|
|
||||||
return backed[index]
|
|
||||||
|
|
||||||
def add(tensor: torch.Tensor) -> None:
|
def add(tensor: torch.Tensor) -> None:
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -41,6 +41,10 @@ _UNLIMITED_ABOVE = 1 << 62
|
|||||||
HOST_RESERVE_FRACTION = 0.05
|
HOST_RESERVE_FRACTION = 0.05
|
||||||
MIN_HOST_RESERVE_BYTES = 2 * GIB_BYTES
|
MIN_HOST_RESERVE_BYTES = 2 * GIB_BYTES
|
||||||
|
|
||||||
|
# Left free when weighing a checkpoint against host memory: activations, staging
|
||||||
|
# buffers and allocator slack are none of them in the weight total.
|
||||||
|
HOST_COPY_RESERVE_BYTES = 4 * GIB_BYTES
|
||||||
|
|
||||||
|
|
||||||
def _read_int(path: str) -> int | None:
|
def _read_int(path: str) -> int | None:
|
||||||
try:
|
try:
|
||||||
@@ -134,6 +138,19 @@ def host_memory_available_bytes() -> int:
|
|||||||
return min(available, max(0, limit - usage))
|
return min(available, max(0, limit - usage))
|
||||||
|
|
||||||
|
|
||||||
|
def host_copies_would_not_fit(weight_bytes: int) -> bool:
|
||||||
|
"""Whether copying `weight_bytes` into host memory would run the host out.
|
||||||
|
|
||||||
|
The alternative to a copy is leaving the weights on their file mapping,
|
||||||
|
which the kernel may drop under pressure and re-read from disk. That is
|
||||||
|
slower per byte but bounded, so it is the right answer exactly when the
|
||||||
|
copies do not fit -- and the wrong one when they do.
|
||||||
|
"""
|
||||||
|
if weight_bytes <= 0:
|
||||||
|
return False
|
||||||
|
return weight_bytes >= host_memory_available_bytes() - HOST_COPY_RESERVE_BYTES
|
||||||
|
|
||||||
|
|
||||||
class HostPinBudget:
|
class HostPinBudget:
|
||||||
"""Hands out pinned-host-memory allowances until the headroom runs out.
|
"""Hands out pinned-host-memory allowances until the headroom runs out.
|
||||||
|
|
||||||
|
|||||||
+102
-2
@@ -8,6 +8,7 @@ import torch
|
|||||||
from torch.distributed.tensor import DTensor
|
from torch.distributed.tensor import DTensor
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import MappedRegions
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||||
COMPONENT_RESIDENCY_GROUPS,
|
COMPONENT_RESIDENCY_GROUPS,
|
||||||
LAYERWISE_OFFLOAD,
|
LAYERWISE_OFFLOAD,
|
||||||
@@ -16,6 +17,8 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency
|
|||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
||||||
HostPinBudget,
|
HostPinBudget,
|
||||||
describe_host_memory,
|
describe_host_memory,
|
||||||
|
host_copies_would_not_fit,
|
||||||
|
host_memory_available_bytes,
|
||||||
module_weight_bytes,
|
module_weight_bytes,
|
||||||
pin_benefit_bytes,
|
pin_benefit_bytes,
|
||||||
)
|
)
|
||||||
@@ -270,6 +273,12 @@ class LayerwiseOffloadManager:
|
|||||||
# layer_idx -> {name: pinned_cpu_tensor_with_original_stride}
|
# layer_idx -> {name: pinned_cpu_tensor_with_original_stride}
|
||||||
# stores tensors whose original non-contiguous stride/layout must be preserved
|
# stores tensors whose original non-contiguous stride/layout must be preserved
|
||||||
self._strided_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
self._strided_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
||||||
|
# layer_idx -> {name: tensor still viewing the checkpoint file}
|
||||||
|
# Weights left on their mapping rather than copied into host memory, so
|
||||||
|
# the page cache decides what stays resident. Used when the copies would
|
||||||
|
# not fit; see _keep_weights_on_their_mapping.
|
||||||
|
self._mapped_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
||||||
|
self._mapped_bytes = 0
|
||||||
# mps keeps the original CPU tensor for each layer instead of building a
|
# mps keeps the original CPU tensor for each layer instead of building a
|
||||||
# second flattened host copy
|
# second flattened host copy
|
||||||
self._mps_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
self._mps_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {}
|
||||||
@@ -285,6 +294,9 @@ class LayerwiseOffloadManager:
|
|||||||
self._named_buffers: Dict[str, torch.Tensor] = {}
|
self._named_buffers: Dict[str, torch.Tensor] = {}
|
||||||
self._offload_placeholders: Dict[torch.dtype, torch.Tensor] = {}
|
self._offload_placeholders: Dict[torch.dtype, torch.Tensor] = {}
|
||||||
self._has_dtensor_weights = False
|
self._has_dtensor_weights = False
|
||||||
|
# A snapshot of this process's mappings, taken now because the weights
|
||||||
|
# have just been loaded and their mappings exist.
|
||||||
|
self._mapped_regions = MappedRegions()
|
||||||
# Store forward hooks for removal
|
# Store forward hooks for removal
|
||||||
self._forward_hooks: List[Any] = []
|
self._forward_hooks: List[Any] = []
|
||||||
|
|
||||||
@@ -378,6 +390,48 @@ class LayerwiseOffloadManager:
|
|||||||
|
|
||||||
self._finalize_initialization()
|
self._finalize_initialization()
|
||||||
|
|
||||||
|
def _keep_weights_on_their_mapping(self, layer_groups: Dict) -> bool:
|
||||||
|
"""Whether to leave file-backed weights on their mapping.
|
||||||
|
|
||||||
|
Copying them into host memory buys pinning, and pinning is what lets the
|
||||||
|
copy stream run ahead of compute -- worth 1.03 s against 1.90 s per step
|
||||||
|
on a measured Wan2.1 run, so it is not given up lightly.
|
||||||
|
|
||||||
|
It is given up when the copies do not fit. H3's DiT is 61.73 GiB of
|
||||||
|
weights, of which 50.53 GiB arrive as views into the checkpoint; on a
|
||||||
|
32 GiB host the copy cannot be made at all, and the choice is between a
|
||||||
|
mapping and not running. Above that, page-cache residency makes the read
|
||||||
|
nearly as fast as pinned -- 12.38 GB/s against 13.39 measured -- and what
|
||||||
|
is lost is the overlap, not the bandwidth.
|
||||||
|
"""
|
||||||
|
if not self._mapped_regions.available:
|
||||||
|
return False
|
||||||
|
mapped_bytes = sum(
|
||||||
|
tensor.untyped_storage().nbytes()
|
||||||
|
for dtype_to_params in layer_groups.values()
|
||||||
|
for weights in dtype_to_params.values()
|
||||||
|
for _, weight in weights
|
||||||
|
for tensor in (self._to_local_tensor(weight),)
|
||||||
|
if self._mapped_regions.holds(tensor)
|
||||||
|
)
|
||||||
|
if mapped_bytes <= 0:
|
||||||
|
return False
|
||||||
|
# The copies land in host memory on top of the page cache already holding
|
||||||
|
# these bytes, so the room needed is the copy itself plus a reserve.
|
||||||
|
if not host_copies_would_not_fit(mapped_bytes):
|
||||||
|
return False
|
||||||
|
available = host_memory_available_bytes()
|
||||||
|
logger.info(
|
||||||
|
"Layerwise offload: leaving %.2f GiB of weights on the checkpoint "
|
||||||
|
"mapping -- copying them into host memory needs more than the "
|
||||||
|
"%.2f GiB available. The page cache decides what stays resident, so "
|
||||||
|
"reads may come from disk; the copies cannot be pinned, so they run "
|
||||||
|
"on the compute stream instead of ahead of it.",
|
||||||
|
mapped_bytes / 1024**3,
|
||||||
|
available / 1024**3,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
def _initialize_layer_weights(self) -> None:
|
def _initialize_layer_weights(self) -> None:
|
||||||
self._named_parameters = dict(self.model.named_parameters())
|
self._named_parameters = dict(self.model.named_parameters())
|
||||||
self._named_buffers = dict(self.model.named_buffers())
|
self._named_buffers = dict(self.model.named_buffers())
|
||||||
@@ -397,16 +451,36 @@ class LayerwiseOffloadManager:
|
|||||||
local_tensor.dtype, []
|
local_tensor.dtype, []
|
||||||
).append((name, tensor))
|
).append((name, tensor))
|
||||||
|
|
||||||
|
keep_mapping = self._keep_weights_on_their_mapping(layer_groups)
|
||||||
|
|
||||||
# 2. concat and offload (in pinned memory)
|
# 2. concat and offload (in pinned memory)
|
||||||
for layer_idx, dtype_to_params in layer_groups.items():
|
for layer_idx, dtype_to_params in layer_groups.items():
|
||||||
self._consolidated_cpu_weights[layer_idx] = {}
|
self._consolidated_cpu_weights[layer_idx] = {}
|
||||||
self._strided_cpu_weights[layer_idx] = {}
|
self._strided_cpu_weights[layer_idx] = {}
|
||||||
|
self._mapped_cpu_weights[layer_idx] = {}
|
||||||
self._weight_metadata[layer_idx] = {}
|
self._weight_metadata[layer_idx] = {}
|
||||||
|
|
||||||
for dtype, weights in dtype_to_params.items():
|
for dtype, weights in dtype_to_params.items():
|
||||||
contiguous_weights: List[Tuple[str, torch.Tensor, torch.Tensor]] = []
|
contiguous_weights: List[Tuple[str, torch.Tensor, torch.Tensor]] = []
|
||||||
for name, weight in weights:
|
for name, weight in weights:
|
||||||
local_weight = self._to_local_tensor(weight)
|
local_weight = self._to_local_tensor(weight)
|
||||||
|
if keep_mapping and self._mapped_regions.holds(local_weight):
|
||||||
|
# Already a view into the checkpoint. Copying it would
|
||||||
|
# add a second copy of bytes the page cache holds
|
||||||
|
# anyway, and that copy is what does not fit.
|
||||||
|
self._mapped_cpu_weights[layer_idx][name] = local_weight
|
||||||
|
self._weight_metadata[layer_idx][name] = {
|
||||||
|
"dtype": local_weight.dtype,
|
||||||
|
"shape": tuple(local_weight.shape),
|
||||||
|
"stride": local_weight.stride(),
|
||||||
|
"preserve_strides": False,
|
||||||
|
"mapped": True,
|
||||||
|
}
|
||||||
|
self._mapped_bytes += local_weight.untyped_storage().nbytes()
|
||||||
|
weight.data = self._get_shared_empty_tensor_for_target(
|
||||||
|
weight, local_weight.dtype
|
||||||
|
)
|
||||||
|
continue
|
||||||
if local_weight.is_contiguous():
|
if local_weight.is_contiguous():
|
||||||
contiguous_weights.append((name, weight, local_weight))
|
contiguous_weights.append((name, weight, local_weight))
|
||||||
continue
|
continue
|
||||||
@@ -629,7 +703,9 @@ class LayerwiseOffloadManager:
|
|||||||
)
|
)
|
||||||
self._gpu_layers.add(layer_idx)
|
self._gpu_layers.add(layer_idx)
|
||||||
return
|
return
|
||||||
if layer_idx not in self._consolidated_cpu_weights:
|
if layer_idx not in self._consolidated_cpu_weights and not (
|
||||||
|
self._mapped_cpu_weights.get(layer_idx)
|
||||||
|
):
|
||||||
return
|
return
|
||||||
if self.copy_stream is not None:
|
if self.copy_stream is not None:
|
||||||
self.copy_stream.wait_stream(torch.get_device_module().current_stream())
|
self.copy_stream.wait_stream(torch.get_device_module().current_stream())
|
||||||
@@ -646,7 +722,9 @@ class LayerwiseOffloadManager:
|
|||||||
torch.no_grad(),
|
torch.no_grad(),
|
||||||
stream_context,
|
stream_context,
|
||||||
):
|
):
|
||||||
for dtype, cpu_buffer in self._consolidated_cpu_weights[layer_idx].items():
|
for dtype, cpu_buffer in self._consolidated_cpu_weights.get(
|
||||||
|
layer_idx, {}
|
||||||
|
).items():
|
||||||
gpu_buffer = torch.empty(
|
gpu_buffer = torch.empty(
|
||||||
cpu_buffer.shape, dtype=dtype, device=self.device
|
cpu_buffer.shape, dtype=dtype, device=self.device
|
||||||
)
|
)
|
||||||
@@ -657,6 +735,18 @@ class LayerwiseOffloadManager:
|
|||||||
# so the recorded event covers both flat-buffer and stride-preserving copies.
|
# so the recorded event covers both flat-buffer and stride-preserving copies.
|
||||||
for name, meta in self._weight_metadata[layer_idx].items():
|
for name, meta in self._weight_metadata[layer_idx].items():
|
||||||
target = self.get_target_with_name(name)
|
target = self.get_target_with_name(name)
|
||||||
|
if meta.get("mapped", False):
|
||||||
|
# Straight from the mapping. Not pinned, so this copy runs
|
||||||
|
# on the compute stream rather than ahead of it, and a page
|
||||||
|
# the kernel has reclaimed is faulted back in here.
|
||||||
|
cpu_tensor = self._mapped_cpu_weights[layer_idx][name]
|
||||||
|
gpu_tensor = torch.empty(
|
||||||
|
meta["shape"], dtype=meta["dtype"], device=self.device
|
||||||
|
)
|
||||||
|
gpu_tensor.copy_(cpu_tensor, non_blocking=False)
|
||||||
|
target.data = self._wrap_for_target(target, gpu_tensor)
|
||||||
|
continue
|
||||||
|
|
||||||
if meta.get("preserve_strides", False):
|
if meta.get("preserve_strides", False):
|
||||||
# Recreate the original view layout instead of flatten+view.
|
# Recreate the original view layout instead of flatten+view.
|
||||||
# ModelOpt FP8 relies on a transposed runtime weight layout,
|
# ModelOpt FP8 relies on a transposed runtime weight layout,
|
||||||
@@ -766,6 +856,12 @@ class LayerwiseOffloadManager:
|
|||||||
|
|
||||||
# Collect current GPU weights and write back to CPU buffer
|
# Collect current GPU weights and write back to CPU buffer
|
||||||
for name, meta in self._weight_metadata.get(layer_idx, {}).items():
|
for name, meta in self._weight_metadata.get(layer_idx, {}).items():
|
||||||
|
if meta.get("mapped", False):
|
||||||
|
# The store is the checkpoint file. Inference does not mutate
|
||||||
|
# weights, so there is nothing to write back; writing would
|
||||||
|
# copy-on-write the mapping into the anonymous memory this
|
||||||
|
# exists to avoid.
|
||||||
|
continue
|
||||||
target = self.get_target_with_name(name)
|
target = self.get_target_with_name(name)
|
||||||
target_local = self._to_local_tensor(target)
|
target_local = self._to_local_tensor(target)
|
||||||
if meta.get("preserve_strides", False):
|
if meta.get("preserve_strides", False):
|
||||||
@@ -903,6 +999,10 @@ class LayerwiseOffloadManager:
|
|||||||
|
|
||||||
for layer_idx in sorted(self._weight_metadata):
|
for layer_idx in sorted(self._weight_metadata):
|
||||||
for name, meta in self._weight_metadata[layer_idx].items():
|
for name, meta in self._weight_metadata[layer_idx].items():
|
||||||
|
if meta.get("mapped", False):
|
||||||
|
yield name, self._mapped_cpu_weights[layer_idx][name]
|
||||||
|
continue
|
||||||
|
|
||||||
if meta.get("preserve_strides", False):
|
if meta.get("preserve_strides", False):
|
||||||
# Some quantized weights rely on a non-contiguous layout.
|
# Some quantized weights rely on a non-contiguous layout.
|
||||||
# Yield the strided tensor directly instead of rebuilding it
|
# Yield the strided tensor directly instead of rebuilding it
|
||||||
|
|||||||
@@ -628,7 +628,7 @@ class MiniMaxH3Attention(nn.Module):
|
|||||||
weight = self.qkv_proj.weight
|
weight = self.qkv_proj.weight
|
||||||
# h3 checkpoints interleave each attention head's Q, K, and V rows
|
# h3 checkpoints interleave each attention head's Q, K, and V rows
|
||||||
# this parameter needs reordering before the native QKV projection
|
# this parameter needs reordering before the native QKV projection
|
||||||
weight.mps_zero_copy_unsafe = True
|
weight.checkpoint_mapping_unsafe = True
|
||||||
base_loader = weight.weight_loader
|
base_loader = weight.weight_loader
|
||||||
|
|
||||||
def _reorder_checkpoint_weight(loaded_weight: torch.Tensor) -> torch.Tensor:
|
def _reorder_checkpoint_weight(loaded_weight: torch.Tensor) -> torch.Tensor:
|
||||||
|
|||||||
@@ -200,7 +200,7 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
|
|||||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||||
try:
|
try:
|
||||||
can_keep_checkpoint_tensor = bool(
|
can_keep_checkpoint_tensor = bool(
|
||||||
getattr(self, "_mps_zero_copy_weight_loading", False)
|
getattr(self, "_keep_checkpoint_mapping", False)
|
||||||
and weight_loader is default_weight_loader
|
and weight_loader is default_weight_loader
|
||||||
and param.device.type == "cpu"
|
and param.device.type == "cpu"
|
||||||
and loaded_weight.device.type == "cpu"
|
and loaded_weight.device.type == "cpu"
|
||||||
|
|||||||
@@ -6,8 +6,10 @@ import torch.nn as nn
|
|||||||
from sglang.multimodal_gen.runtime.managers.memory_managers import host_memory_budget
|
from sglang.multimodal_gen.runtime.managers.memory_managers import host_memory_budget
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
||||||
GIB_BYTES,
|
GIB_BYTES,
|
||||||
|
HOST_COPY_RESERVE_BYTES,
|
||||||
HostPinBudget,
|
HostPinBudget,
|
||||||
cgroup_memory_limit_bytes,
|
cgroup_memory_limit_bytes,
|
||||||
|
host_copies_would_not_fit,
|
||||||
host_memory_available_bytes,
|
host_memory_available_bytes,
|
||||||
module_weight_bytes,
|
module_weight_bytes,
|
||||||
pin_benefit_bytes,
|
pin_benefit_bytes,
|
||||||
@@ -98,6 +100,49 @@ class TestCgroupLimit:
|
|||||||
assert host_memory_available_bytes() == 12 * GIB_BYTES
|
assert host_memory_available_bytes() == 12 * GIB_BYTES
|
||||||
|
|
||||||
|
|
||||||
|
def _host_has(monkeypatch, tmp_path, available_bytes):
|
||||||
|
_point_at(monkeypatch, tmp_path)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
host_memory_budget.psutil,
|
||||||
|
"virtual_memory",
|
||||||
|
lambda: type("VM", (), {"available": available_bytes})(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestHostCopiesWouldNotFit:
|
||||||
|
def test_a_checkpoint_larger_than_the_host_does_not_fit(
|
||||||
|
self, monkeypatch, tmp_path
|
||||||
|
):
|
||||||
|
# the measured H3 case: a 62 GiB encoder against a 32 GiB host
|
||||||
|
_host_has(monkeypatch, tmp_path, 32 * GIB_BYTES)
|
||||||
|
assert host_copies_would_not_fit(62 * GIB_BYTES)
|
||||||
|
|
||||||
|
def test_a_checkpoint_well_under_the_host_fits(self, monkeypatch, tmp_path):
|
||||||
|
_host_has(monkeypatch, tmp_path, 128 * GIB_BYTES)
|
||||||
|
assert not host_copies_would_not_fit(20 * GIB_BYTES)
|
||||||
|
|
||||||
|
def test_the_reserve_is_not_available_to_the_copy(self, monkeypatch, tmp_path):
|
||||||
|
# a copy that fits only by eating the reserve is treated as not fitting
|
||||||
|
_host_has(monkeypatch, tmp_path, 40 * GIB_BYTES)
|
||||||
|
spendable = 40 * GIB_BYTES - HOST_COPY_RESERVE_BYTES
|
||||||
|
assert host_copies_would_not_fit(spendable)
|
||||||
|
assert not host_copies_would_not_fit(spendable - 1)
|
||||||
|
|
||||||
|
def test_nothing_to_copy_always_fits(self, monkeypatch, tmp_path):
|
||||||
|
_host_has(monkeypatch, tmp_path, 0)
|
||||||
|
assert not host_copies_would_not_fit(0)
|
||||||
|
|
||||||
|
def test_the_cgroup_cap_decides_and_not_the_machine(self, monkeypatch, tmp_path):
|
||||||
|
# psutil sees the whole box; only the cap makes the copy impossible
|
||||||
|
_point_at(monkeypatch, tmp_path, v2=(32 * GIB_BYTES, 0))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
host_memory_budget.psutil,
|
||||||
|
"virtual_memory",
|
||||||
|
lambda: type("VM", (), {"available": 900 * GIB_BYTES})(),
|
||||||
|
)
|
||||||
|
assert host_copies_would_not_fit(62 * GIB_BYTES)
|
||||||
|
|
||||||
|
|
||||||
class TestNestedCgroup:
|
class TestNestedCgroup:
|
||||||
def test_a_tighter_nested_cap_wins_over_the_root(self, monkeypatch, tmp_path):
|
def test_a_tighter_nested_cap_wins_over_the_root(self, monkeypatch, tmp_path):
|
||||||
# a systemd scope with MemoryMax, or --cgroup-parent: planning against
|
# a systemd scope with MemoryMax, or --cgroup-parent: planning against
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import pathlib
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
@@ -15,6 +16,9 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
|||||||
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
||||||
component_residency_strategies as component_residency_strategies_mod,
|
component_residency_strategies as component_residency_strategies_mod,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
||||||
|
host_memory_budget,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
||||||
layerwise_offload as layerwise_offload_mod,
|
layerwise_offload as layerwise_offload_mod,
|
||||||
)
|
)
|
||||||
@@ -1228,6 +1232,82 @@ def test_strided_forward_leaves_exactly_the_resident_set(monkeypatch):
|
|||||||
assert len(manager._gpu_layers) <= len(resident) + manager.prefetch_size
|
assert len(manager._gpu_layers) <= len(resident) + manager.prefetch_size
|
||||||
|
|
||||||
|
|
||||||
|
class _FileBackedBlock(torch.nn.Module):
|
||||||
|
"""A block whose weight is a view into a file, as a loaded checkpoint is."""
|
||||||
|
|
||||||
|
def __init__(self, path: pathlib.Path) -> None:
|
||||||
|
super().__init__()
|
||||||
|
path.write_bytes(b"\x00" * (64 * 4))
|
||||||
|
mapped = torch.from_file(str(path), shared=True, size=64, dtype=torch.float32)
|
||||||
|
self.weight = torch.nn.Parameter(mapped.reshape(8, 8), requires_grad=False)
|
||||||
|
|
||||||
|
|
||||||
|
class _FileBackedModel(torch.nn.Module):
|
||||||
|
def __init__(self, path: pathlib.Path) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.blocks = torch.nn.ModuleList([_FileBackedBlock(path)])
|
||||||
|
|
||||||
|
|
||||||
|
def _mapped_manager(tmp_path, monkeypatch, *, available_gib):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu")
|
||||||
|
# Two bindings for one fact: the budget module's copy is what decides
|
||||||
|
# whether the copies fit, and layerwise_offload's own is what the log reports.
|
||||||
|
available_bytes = int(available_gib * 1024**3)
|
||||||
|
for module in (host_memory_budget, layerwise_offload_mod):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
module, "host_memory_available_bytes", lambda: available_bytes
|
||||||
|
)
|
||||||
|
model = _FileBackedModel(tmp_path / "weights.bin")
|
||||||
|
return LayerwiseOffloadManager(
|
||||||
|
model=model,
|
||||||
|
layers_attr_str="blocks",
|
||||||
|
num_layers=1,
|
||||||
|
enabled=True,
|
||||||
|
pin_cpu_memory=True,
|
||||||
|
prefetch_size=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_weights_stay_on_the_mapping_when_copies_do_not_fit(tmp_path, monkeypatch):
|
||||||
|
if not pathlib.Path("/proc/self/maps").exists():
|
||||||
|
pytest.skip("needs /proc to tell a mapping from anonymous memory")
|
||||||
|
# the reserve alone exceeds this, so no copy can be afforded
|
||||||
|
manager = _mapped_manager(tmp_path, monkeypatch, available_gib=0.001)
|
||||||
|
assert manager._mapped_cpu_weights[0], "expected the weight to stay mapped"
|
||||||
|
assert manager._weight_metadata[0]["blocks.0.weight"]["mapped"] is True
|
||||||
|
assert not manager._consolidated_cpu_weights.get(0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_weights_are_copied_when_they_fit(tmp_path, monkeypatch):
|
||||||
|
if not pathlib.Path("/proc/self/maps").exists():
|
||||||
|
pytest.skip("needs /proc to tell a mapping from anonymous memory")
|
||||||
|
manager = _mapped_manager(tmp_path, monkeypatch, available_gib=64)
|
||||||
|
assert not manager._mapped_cpu_weights[0], "a copy was affordable"
|
||||||
|
assert manager._consolidated_cpu_weights[0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_a_mapped_weight_is_not_written_back(tmp_path, monkeypatch):
|
||||||
|
if not pathlib.Path("/proc/self/maps").exists():
|
||||||
|
pytest.skip("needs /proc to tell a mapping from anonymous memory")
|
||||||
|
manager = _mapped_manager(tmp_path, monkeypatch, available_gib=0.001)
|
||||||
|
before = manager._mapped_cpu_weights[0]["blocks.0.weight"].clone()
|
||||||
|
manager._gpu_layers.add(0)
|
||||||
|
manager.sync_layer_to_cpu(0)
|
||||||
|
after = manager._mapped_cpu_weights[0]["blocks.0.weight"]
|
||||||
|
assert torch.equal(before, after), "writeback must not touch the checkpoint"
|
||||||
|
|
||||||
|
|
||||||
|
def test_mapped_weights_are_visible_to_checksums(tmp_path, monkeypatch):
|
||||||
|
if not pathlib.Path("/proc/self/maps").exists():
|
||||||
|
pytest.skip("needs /proc to tell a mapping from anonymous memory")
|
||||||
|
manager = _mapped_manager(tmp_path, monkeypatch, available_gib=0.001)
|
||||||
|
names = {name for name, _ in manager.iter_cpu_weights()}
|
||||||
|
assert "blocks.0.weight" in names
|
||||||
|
|
||||||
|
|
||||||
def test_layerwise_tuning_defaults_match_the_group():
|
def test_layerwise_tuning_defaults_match_the_group():
|
||||||
"""No per-component entry: the DiT group keeps its knobs, auxiliaries do not."""
|
"""No per-component entry: the DiT group keeps its knobs, auxiliaries do not."""
|
||||||
args = _server_args(
|
args = _server_args(
|
||||||
|
|||||||
Reference in New Issue
Block a user