diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py new file mode 100644 index 000000000..269bbee5e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/host_memory_budget.py @@ -0,0 +1,171 @@ +# SPDX-License-Identifier: Apache-2.0 +"""How much host memory this process may still commit, and to what. + +Offloaded weights live in pinned host memory, which the kernel can neither swap +nor drop. That is what makes the asynchronous host-to-device copies work, and +also what turns "using a lot of RAM" into "the container gets OOM-killed": +ordinary pages would have been reclaimed instead. + +`psutil.virtual_memory()` cannot be the whole answer here, because it reads +/proc/meminfo, which is host-wide and blind to a container's limit. Measured on +a rented 4-GPU box: psutil reports 2015.7 GiB total while the cgroup caps the +container at 1117.2 GiB, a 900 GiB over-report. Serving runs in containers, so +the cap is read directly from whichever cgroup version is mounted. +""" + +import psutil + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +GIB_BYTES = 1024**3 + +_CGROUP_V2 = ("/sys/fs/cgroup/memory.max", "/sys/fs/cgroup/memory.current") +_CGROUP_V1 = ( + "/sys/fs/cgroup/memory/memory.limit_in_bytes", + "/sys/fs/cgroup/memory/memory.usage_in_bytes", +) + +# An unlimited v1 cgroup reports a sentinel near 2**63 rather than omitting the +# file, so treat anything implausibly large as "no cap". +_UNLIMITED_ABOVE = 1 << 62 + +# Left unpinned so the process can still allocate activations, staging buffers +# and whatever the allocator needs mid-request. A share of the cap rather than a +# flat number, because the same absolute headroom is generous on a desktop and +# nothing on a serving host. +HOST_RESERVE_FRACTION = 0.05 +MIN_HOST_RESERVE_BYTES = 2 * GIB_BYTES + + +def _read_int(path: str) -> int | None: + try: + with open(path) as handle: + text = handle.read().strip() + except OSError: + return None + if text == "max": + return None + try: + return int(text) + except ValueError: + return None + + +def cgroup_memory_limit_bytes() -> tuple[int, int] | None: + """This process's (cap, usage) under its cgroup, or None when uncapped.""" + for limit_path, usage_path in (_CGROUP_V2, _CGROUP_V1): + limit = _read_int(limit_path) + if limit is None or limit >= _UNLIMITED_ABOVE: + continue + usage = _read_int(usage_path) or 0 + return limit, usage + return None + + +def host_memory_available_bytes() -> int: + """Bytes this process can still commit without hitting a wall. + + The smaller of what the kernel reports free and what the cgroup still + allows, so a container does not plan against the whole machine. + """ + available = int(psutil.virtual_memory().available) + capped = cgroup_memory_limit_bytes() + if capped is None: + return available + limit, usage = capped + return min(available, max(0, limit - usage)) + + +class HostPinBudget: + """Hands out pinned-host-memory allowances until the headroom runs out. + + Pinning is not all-or-nothing per process: a component whose weights stream + once per request gains far less from pinning than one re-streamed on every + denoise step. So the hot components are offered the budget first, and a cold + component that no longer fits falls back to pageable host memory. + + That fallback is a last resort, not a cheap safety net. Measured on an + RTX 4090 with Wan2.1-1.3B, dropping the text encoder to pageable left the + denoise loop untouched but doubled its own stage (3.04 s -> 6.26 s at best, + and up to 8x when the host's memory bandwidth was contended). It is still + the right trade against exhausting host memory -- slower is not dead -- but + it only fires when the bytes genuinely do not fit. + """ + + def __init__(self, available_bytes: int | None = None) -> None: + if available_bytes is None: + available_bytes = host_memory_available_bytes() + self.available_bytes = available_bytes + self.reserve_bytes = max( + int(available_bytes * HOST_RESERVE_FRACTION), MIN_HOST_RESERVE_BYTES + ) + self.committed_bytes = 0 + + @property + def spendable_bytes(self) -> int: + return max(0, self.available_bytes - self.reserve_bytes - self.committed_bytes) + + def request(self, *, component_name: str, weight_bytes: int) -> bool: + """Whether `component_name` may pin `weight_bytes`, and book it if so. + + The cap is hard even for the hot components. Granting past it does not + buy a smaller footprint, it just moves the failure: the pinned + allocation itself starts failing, or the box begins swapping. Priority + is expressed by asking in hot-first order, not by overrunning. + """ + if weight_bytes <= 0: + return True + if weight_bytes <= self.spendable_bytes: + self.committed_bytes += weight_bytes + return True + logger.info( + "Host pin budget: %s stays pageable (%.2f GB of weights, %.2f GB " + "spendable of %.2f GB available). Its host-to-device copies fall " + "back to staged transfers -- measured at roughly 2x the time for " + "the stage that uses it, and more under memory-bandwidth " + "contention. Nothing is re-read from disk.", + component_name, + weight_bytes / GIB_BYTES, + self.spendable_bytes / GIB_BYTES, + self.available_bytes / GIB_BYTES, + ) + return False + + +def pin_benefit_bytes(*, weight_bytes: int, uses_per_request: int) -> int: + """Host-to-device bytes a pin would cover for one request. + + Ranking on this product rather than on "is it the DiT" matters for few-step + models: a 20 GB text encoder used once moves more per request than a 1 GB + DiT stepped four times, so it is the one that should claim the budget. + """ + return max(0, weight_bytes) * max(1, uses_per_request) + + +def module_weight_bytes(module) -> int: + """Bytes of parameters and buffers a module would hand to the host.""" + seen: set[int] = set() + total = 0 + for tensor in list(module.parameters()) + list(module.buffers()): + storage = tensor.untyped_storage() + pointer = storage.data_ptr() + if pointer == 0 or pointer in seen: + continue + seen.add(pointer) + total += storage.nbytes() + return total + + +def describe_host_memory() -> str: + """One line for startup logs: what the cap is and where it came from.""" + capped = cgroup_memory_limit_bytes() + available = host_memory_available_bytes() + if capped is None: + return f"host memory available: {available / GIB_BYTES:.1f} GiB (no cgroup cap)" + limit, usage = capped + return ( + f"host memory available: {available / GIB_BYTES:.1f} GiB " + f"(cgroup cap {limit / GIB_BYTES:.1f} GiB, in use {usage / GIB_BYTES:.1f} GiB)" + ) 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 9ce5f169c..3bcf32398 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 @@ -13,6 +13,12 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency LAYERWISE_OFFLOAD, ComponentResidencyError, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import ( + HostPinBudget, + describe_host_memory, + module_weight_bytes, + pin_benefit_bytes, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( LAYERWISE_OFFLOAD_ALL_COMPONENTS, LAYERWISE_OFFLOAD_DIT_GROUP, @@ -1084,7 +1090,13 @@ class LayerwiseOffloadableModuleMixin: for name, tensor in self._mps_cpu_buffers.items(): buffers[name].data = tensor - def configure_layerwise_offload(self, server_args: ServerArgs): + def configure_layerwise_offload( + self, + server_args: ServerArgs, + *, + pin_budget: HostPinBudget | None = None, + component_name: str | None = None, + ): self.layerwise_offload_managers = [] named_modules = dict(self.named_modules()) configured_layer_names = [] @@ -1120,12 +1132,22 @@ class LayerwiseOffloadableModuleMixin: else: resident_layers = min(num_layers, int(resident_value)) + # Pinning these weights is what lets the copy stream run ahead of + # compute, but pinned pages are the ones the kernel cannot reclaim, + # so a component only gets them while the budget lasts. + pin_cpu_memory = server_args.pin_cpu_memory + if pin_cpu_memory and pin_budget is not None: + pin_cpu_memory = pin_budget.request( + component_name=f"{component_name or type(self).__name__}.{layer_name}", + weight_bytes=module_weight_bytes(module_list), + ) + manager = LayerwiseOffloadManager( model=self, layers_attr_str=layer_name, num_layers=num_layers, enabled=True, - pin_cpu_memory=server_args.pin_cpu_memory, + pin_cpu_memory=pin_cpu_memory, prefetch_size=prefetch_size, resident_layers=resident_layers, initialize=False, @@ -1436,6 +1458,49 @@ def configure_layerwise_offload_modules( sorted(unsupported_component_names), ) + def _default_num_inference_steps() -> int: + from sglang.multimodal_gen.registry import get_pipeline_config_classes + + pipeline_class_name = server_args.pipeline_class_name + if not pipeline_class_name: + return 1 + config_classes = get_pipeline_config_classes(pipeline_class_name) + if config_classes is None: + return 1 + return max(1, int(config_classes[1]().num_inference_steps)) + + default_steps = _default_num_inference_steps() + + def _h2d_bytes_a_pin_would_save(name: str) -> int: + """What pinning this component is worth, in bytes moved per request. + + A DiT under layerwise offload re-streams its layers on every denoise + step; everything else transfers once. Ranking on the product rather + than on "is it the DiT" matters for few-step models, where a large + one-shot text encoder can move more bytes per request than a small DiT + stepped four times. + """ + module = modules[name] + if not isinstance(module, LayerwiseOffloadableModuleMixin): + return 0 + return pin_benefit_bytes( + weight_bytes=module_weight_bytes(module), + uses_per_request=( + default_steps if module.layerwise_offload_dit_group_enabled else 1 + ), + ) + + # Offer the budget in descending order of what a pin saves, so the bytes + # that would move most often claim it first. sorted() is stable, so equal + # rankings keep their original order. + selected_pipeline_component_names = sorted( + selected_pipeline_component_names, + key=_h2d_bytes_a_pin_would_save, + reverse=True, + ) + pin_budget = HostPinBudget() + logger.info("Layerwise offload: %s", describe_host_memory()) + for component_name in selected_pipeline_component_names: module = modules[component_name] if not isinstance(module, LayerwiseOffloadableModuleMixin): @@ -1451,7 +1516,9 @@ def configure_layerwise_offload_modules( configured_module_ids.add(module_id) if not is_layerwise_offloaded_module(module): - module.configure_layerwise_offload(server_args) + module.configure_layerwise_offload( + server_args, pin_budget=pin_budget, component_name=component_name + ) if not is_layerwise_offloaded_module(module): raise ComponentResidencyError( f"Component {component_name!r} did not enable layerwise offload" diff --git a/python/sglang/multimodal_gen/test/unit/test_host_memory_budget.py b/python/sglang/multimodal_gen/test/unit/test_host_memory_budget.py new file mode 100644 index 000000000..100db4d33 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_host_memory_budget.py @@ -0,0 +1,156 @@ +"""Pinned host memory is planned against the cgroup cap, not the whole machine.""" + +import torch +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.host_memory_budget import ( + GIB_BYTES, + HostPinBudget, + cgroup_memory_limit_bytes, + host_memory_available_bytes, + module_weight_bytes, + pin_benefit_bytes, +) + + +def _point_at(monkeypatch, tmp_path, *, v2=None, v1=None): + """Redirect the cgroup lookups at files under tmp_path.""" + + def write(name, value): + path = tmp_path / name + path.write_text(str(value)) + return str(path) + + missing = str(tmp_path / "absent") + v2_paths = ( + (write("memory.max", v2[0]), write("memory.current", v2[1])) + if v2 + else (missing, missing) + ) + v1_paths = ( + (write("limit_in_bytes", v1[0]), write("usage_in_bytes", v1[1])) + if v1 + else (missing, missing) + ) + monkeypatch.setattr(host_memory_budget, "_CGROUP_V2", v2_paths) + monkeypatch.setattr(host_memory_budget, "_CGROUP_V1", v1_paths) + + +class TestCgroupLimit: + def test_v2_cap_is_read(self, monkeypatch, tmp_path): + _point_at(monkeypatch, tmp_path, v2=(32 * GIB_BYTES, 4 * GIB_BYTES)) + assert cgroup_memory_limit_bytes() == (32 * GIB_BYTES, 4 * GIB_BYTES) + + def test_v1_cap_is_read_when_v2_is_absent(self, monkeypatch, tmp_path): + _point_at(monkeypatch, tmp_path, v1=(64 * GIB_BYTES, 8 * GIB_BYTES)) + assert cgroup_memory_limit_bytes() == (64 * GIB_BYTES, 8 * GIB_BYTES) + + def test_no_cgroup_reports_uncapped(self, monkeypatch, tmp_path): + _point_at(monkeypatch, tmp_path) + assert cgroup_memory_limit_bytes() is None + + def test_v2_max_keyword_is_uncapped(self, monkeypatch, tmp_path): + _point_at(monkeypatch, tmp_path, v2=("max", 4 * GIB_BYTES)) + assert cgroup_memory_limit_bytes() is None + + def test_v1_sentinel_is_uncapped(self, monkeypatch, tmp_path): + # an unlimited v1 cgroup reports a number near 2**63 rather than "max" + _point_at(monkeypatch, tmp_path, v1=(2**63 - 4096, 8 * GIB_BYTES)) + assert cgroup_memory_limit_bytes() is None + + def test_the_cap_wins_over_what_the_kernel_reports_free( + self, monkeypatch, tmp_path + ): + # the case measured on a rented box: psutil sees the whole machine + _point_at(monkeypatch, tmp_path, v2=(32 * GIB_BYTES, 8 * GIB_BYTES)) + monkeypatch.setattr( + host_memory_budget.psutil, + "virtual_memory", + lambda: type("VM", (), {"available": 900 * GIB_BYTES})(), + ) + assert host_memory_available_bytes() == 24 * GIB_BYTES + + def test_free_memory_wins_when_it_is_the_smaller_number( + self, monkeypatch, tmp_path + ): + _point_at(monkeypatch, tmp_path, v2=(900 * GIB_BYTES, 0)) + monkeypatch.setattr( + host_memory_budget.psutil, + "virtual_memory", + lambda: type("VM", (), {"available": 12 * GIB_BYTES})(), + ) + assert host_memory_available_bytes() == 12 * GIB_BYTES + + +class TestHostPinBudget: + def test_a_component_that_fits_is_granted(self): + budget = HostPinBudget(available_bytes=40 * GIB_BYTES) + assert budget.request(component_name="dit", weight_bytes=20 * GIB_BYTES) + + def test_the_reserve_is_not_spendable(self): + budget = HostPinBudget(available_bytes=40 * GIB_BYTES) + # 5% of 40 GiB is 2 GiB, so 38 GiB is spendable and 39 GiB is not + assert not budget.request(component_name="dit", weight_bytes=39 * GIB_BYTES) + assert budget.request(component_name="dit", weight_bytes=38 * GIB_BYTES) + + def test_the_reserve_has_a_floor_on_a_small_host(self): + budget = HostPinBudget(available_bytes=8 * GIB_BYTES) + # 5% of 8 GiB is well under the 2 GiB floor + assert budget.reserve_bytes == 2 * GIB_BYTES + + def test_a_later_component_is_denied_once_the_budget_is_spent(self): + budget = HostPinBudget(available_bytes=40 * GIB_BYTES) + assert budget.request(component_name="dit", weight_bytes=30 * GIB_BYTES) + assert not budget.request( + component_name="text_encoder", weight_bytes=20 * GIB_BYTES + ) + + def test_the_cap_binds_even_for_the_first_component(self): + # priority comes from asking first, not from being allowed to overrun + budget = HostPinBudget(available_bytes=8 * GIB_BYTES) + assert not budget.request(component_name="dit", weight_bytes=20 * GIB_BYTES) + + def test_a_weightless_component_needs_no_budget(self): + budget = HostPinBudget(available_bytes=0) + assert budget.request(component_name="scheduler", weight_bytes=0) + + +class TestModuleWeightBytes: + def test_parameters_and_buffers_are_counted(self): + module = nn.Linear(64, 64, bias=False) + assert module_weight_bytes(module) == 64 * 64 * module.weight.element_size() + + def test_shared_storage_is_counted_once(self): + module = nn.Module() + backing = torch.empty(1024, dtype=torch.float32) + module.register_buffer("a", backing[:512]) + module.register_buffer("b", backing[512:]) + assert module_weight_bytes(module) == 4096 + + +class TestPinBenefit: + def test_a_stepped_component_counts_every_step(self): + assert pin_benefit_bytes(weight_bytes=1000, uses_per_request=50) == 50_000 + + def test_a_one_shot_component_counts_once(self): + assert pin_benefit_bytes(weight_bytes=1000, uses_per_request=1) == 1000 + + def test_a_few_step_model_inverts_the_obvious_order(self): + # 1 GB DiT over 4 steps against a 20 GB one-shot text encoder: ranking + # by "is it the DiT" would hand the budget to the wrong one + dit = pin_benefit_bytes(weight_bytes=1 * GIB_BYTES, uses_per_request=4) + text_encoder = pin_benefit_bytes( + weight_bytes=20 * GIB_BYTES, uses_per_request=1 + ) + assert text_encoder > dit + + def test_a_many_step_model_keeps_the_dit_first(self): + dit = pin_benefit_bytes(weight_bytes=3 * GIB_BYTES, uses_per_request=50) + text_encoder = pin_benefit_bytes( + weight_bytes=21 * GIB_BYTES, uses_per_request=1 + ) + assert dit > text_encoder + + def test_missing_step_count_is_treated_as_one_use(self): + assert pin_benefit_bytes(weight_bytes=1000, uses_per_request=0) == 1000 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 53b0d9a06..e9bd41974 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -211,6 +211,9 @@ def _server_args(**kwargs): dit_layerwise_resident_layers=0.0, dit_layerwise_residency_policy=RESIDENCY_POLICY_LEADING, pin_cpu_memory=False, + # the pin budget ranks candidates by bytes x steps, and reads the step + # count off the pipeline's sampling defaults + pipeline_class_name=None, ) defaults.update(kwargs) return _TestServerArgs(**defaults)