[diffusion] feat: plan pinned host memory against the cgroup cap not the machine (#35641)
This commit is contained in:
@@ -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)"
|
||||||
|
)
|
||||||
@@ -13,6 +13,12 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency
|
|||||||
LAYERWISE_OFFLOAD,
|
LAYERWISE_OFFLOAD,
|
||||||
ComponentResidencyError,
|
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 (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import (
|
||||||
LAYERWISE_OFFLOAD_ALL_COMPONENTS,
|
LAYERWISE_OFFLOAD_ALL_COMPONENTS,
|
||||||
LAYERWISE_OFFLOAD_DIT_GROUP,
|
LAYERWISE_OFFLOAD_DIT_GROUP,
|
||||||
@@ -1084,7 +1090,13 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
for name, tensor in self._mps_cpu_buffers.items():
|
for name, tensor in self._mps_cpu_buffers.items():
|
||||||
buffers[name].data = tensor
|
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 = []
|
self.layerwise_offload_managers = []
|
||||||
named_modules = dict(self.named_modules())
|
named_modules = dict(self.named_modules())
|
||||||
configured_layer_names = []
|
configured_layer_names = []
|
||||||
@@ -1120,12 +1132,22 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
else:
|
else:
|
||||||
resident_layers = min(num_layers, int(resident_value))
|
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(
|
manager = LayerwiseOffloadManager(
|
||||||
model=self,
|
model=self,
|
||||||
layers_attr_str=layer_name,
|
layers_attr_str=layer_name,
|
||||||
num_layers=num_layers,
|
num_layers=num_layers,
|
||||||
enabled=True,
|
enabled=True,
|
||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=pin_cpu_memory,
|
||||||
prefetch_size=prefetch_size,
|
prefetch_size=prefetch_size,
|
||||||
resident_layers=resident_layers,
|
resident_layers=resident_layers,
|
||||||
initialize=False,
|
initialize=False,
|
||||||
@@ -1436,6 +1458,49 @@ def configure_layerwise_offload_modules(
|
|||||||
sorted(unsupported_component_names),
|
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:
|
for component_name in selected_pipeline_component_names:
|
||||||
module = modules[component_name]
|
module = modules[component_name]
|
||||||
if not isinstance(module, LayerwiseOffloadableModuleMixin):
|
if not isinstance(module, LayerwiseOffloadableModuleMixin):
|
||||||
@@ -1451,7 +1516,9 @@ def configure_layerwise_offload_modules(
|
|||||||
configured_module_ids.add(module_id)
|
configured_module_ids.add(module_id)
|
||||||
|
|
||||||
if not is_layerwise_offloaded_module(module):
|
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):
|
if not is_layerwise_offloaded_module(module):
|
||||||
raise ComponentResidencyError(
|
raise ComponentResidencyError(
|
||||||
f"Component {component_name!r} did not enable layerwise offload"
|
f"Component {component_name!r} did not enable layerwise offload"
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -211,6 +211,9 @@ def _server_args(**kwargs):
|
|||||||
dit_layerwise_resident_layers=0.0,
|
dit_layerwise_resident_layers=0.0,
|
||||||
dit_layerwise_residency_policy=RESIDENCY_POLICY_LEADING,
|
dit_layerwise_residency_policy=RESIDENCY_POLICY_LEADING,
|
||||||
pin_cpu_memory=False,
|
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)
|
defaults.update(kwargs)
|
||||||
return _TestServerArgs(**defaults)
|
return _TestServerArgs(**defaults)
|
||||||
|
|||||||
Reference in New Issue
Block a user