[diffusion] feat: support resident layers for DiT (#31538)
This commit is contained in:
@@ -16,6 +16,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_s
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
is_layerwise_offloaded_module,
|
is_layerwise_offloaded_module,
|
||||||
|
is_resident_layerwise_module,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import (
|
||||||
is_dit_component_name,
|
is_dit_component_name,
|
||||||
@@ -419,6 +420,11 @@ class ComponentResidencyManager:
|
|||||||
# Avoid making two vanilla-offloaded heavy components resident before
|
# Avoid making two vanilla-offloaded heavy components resident before
|
||||||
# a budget-aware planner can prove the overlap is safe.
|
# a budget-aware planner can prove the overlap is safe.
|
||||||
return
|
return
|
||||||
|
if is_resident_layerwise_module(module):
|
||||||
|
# A layerwise DiT holding a large resident set must not be prefetched
|
||||||
|
# during a prior peer stage (e.g. text encoding): co-residing can lead
|
||||||
|
# to OOMs. Pin it lazily at the DiT's own use-site.
|
||||||
|
return
|
||||||
|
|
||||||
self._uses_seen[use.component_name] = use
|
self._uses_seen[use.component_name] = use
|
||||||
if strategy.prefetch_for_use(module, use, self.state):
|
if strategy.prefetch_for_use(module, use, self.state):
|
||||||
@@ -458,7 +464,9 @@ class ComponentResidencyManager:
|
|||||||
if self.state.batch_is_warmup and use.keep_ready_after_warmup:
|
if self.state.batch_is_warmup and use.keep_ready_after_warmup:
|
||||||
continue
|
continue
|
||||||
preferred = component_name in preferred_uses
|
preferred = component_name in preferred_uses
|
||||||
if not preferred and self._should_keep_single_dit(component_name):
|
if is_resident_layerwise_module(module):
|
||||||
|
preferred = False
|
||||||
|
elif not preferred and self._should_keep_single_dit(component_name):
|
||||||
continue
|
continue
|
||||||
strategy = self.strategy_for(component_name, module)
|
strategy = self.strategy_for(component_name, module)
|
||||||
if preferred and not self.state.batch_is_warmup:
|
if preferred and not self.state.batch_is_warmup:
|
||||||
@@ -537,6 +545,10 @@ class ComponentResidencyManager:
|
|||||||
if use.component_name in future_component_names:
|
if use.component_name in future_component_names:
|
||||||
return True
|
return True
|
||||||
if self._should_keep_single_dit(use.component_name):
|
if self._should_keep_single_dit(use.component_name):
|
||||||
|
module = self.get_module(use.component_name)
|
||||||
|
if module is not None and is_resident_layerwise_module(module):
|
||||||
|
# don't keep a layerwise DiT resident across the request to avoid OOMs
|
||||||
|
return False
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -42,12 +42,20 @@ class LayerwiseOffloadManager:
|
|||||||
enabled: bool,
|
enabled: bool,
|
||||||
pin_cpu_memory: bool = True,
|
pin_cpu_memory: bool = True,
|
||||||
prefetch_size: int = 1,
|
prefetch_size: int = 1,
|
||||||
|
resident_layers: int = 0,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.model = model
|
self.model = model
|
||||||
self.layers_attr_str = layers_attr_str
|
self.layers_attr_str = layers_attr_str
|
||||||
self.num_layers = num_layers
|
self.num_layers = num_layers
|
||||||
self.pin_cpu_memory = pin_cpu_memory
|
self.pin_cpu_memory = pin_cpu_memory
|
||||||
self.prefetch_size = min(max(1, prefetch_size), self.num_layers)
|
self.prefetch_size = min(max(1, prefetch_size), self.num_layers)
|
||||||
|
# Leading layers held on GPU across denoise steps, instead of being
|
||||||
|
# re-streamed every step like the tail.
|
||||||
|
self.resident_layers = min(max(0, int(resident_layers)), self.num_layers)
|
||||||
|
# Armed on the first denoise forward, so that the load-time prefetch below
|
||||||
|
# does not pin the whole resident set before the DiT is the active component.
|
||||||
|
self._residency_active = False
|
||||||
|
|
||||||
self.enabled = bool(enabled and torch.get_device_module().is_available())
|
self.enabled = bool(enabled and torch.get_device_module().is_available())
|
||||||
if not self.enabled:
|
if not self.enabled:
|
||||||
return
|
return
|
||||||
@@ -247,18 +255,36 @@ class LayerwiseOffloadManager:
|
|||||||
|
|
||||||
self.register_forward_hooks()
|
self.register_forward_hooks()
|
||||||
logger.info(
|
logger.info(
|
||||||
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, total num layers: {self.num_layers}"
|
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, num resident layers: {self.resident_layers}, total num layers: {self.num_layers}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def prepare_for_next_req(self, non_blocking=True):
|
def prepare_for_next_req(self, non_blocking=True):
|
||||||
"""
|
"""
|
||||||
Prepare for the next round of denoising loop with prefetching the necessary layers
|
Prepare for the next round of denoising loop with prefetching the necessary layers
|
||||||
"""
|
"""
|
||||||
for i in range(self.prefetch_size):
|
num_prefetch_layers = max(self.prefetch_size, self._retained_layers)
|
||||||
|
for i in range(num_prefetch_layers):
|
||||||
self.prefetch_layer(i, non_blocking=non_blocking)
|
self.prefetch_layer(i, non_blocking=non_blocking)
|
||||||
if not non_blocking and self.copy_stream is not None:
|
if not non_blocking and self.copy_stream is not None:
|
||||||
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
|
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def holds_residents(self) -> bool:
|
||||||
|
"""True if this manager keeps a resident leading-layer set beyond the
|
||||||
|
streaming prefetch window, so it must be denoise-stage-scoped."""
|
||||||
|
return self.enabled and self.resident_layers > 0
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _retained_layers(self) -> int:
|
||||||
|
"""Leading layers currently held across denoise steps; 0 until armed."""
|
||||||
|
return self.resident_layers if self._residency_active else 0
|
||||||
|
|
||||||
|
@torch.compiler.disable
|
||||||
|
def _activate_residency(self) -> None:
|
||||||
|
"""Arm the resident set on the first denoise forward. The pinning itself is
|
||||||
|
done by the ``prepare_for_next_req`` that follows in the same hook."""
|
||||||
|
self._residency_active = True
|
||||||
|
|
||||||
def get_target_with_name(self, name: str) -> torch.Tensor:
|
def get_target_with_name(self, name: str) -> torch.Tensor:
|
||||||
"""get the target model weight/buffer to be replaced"""
|
"""get the target model weight/buffer to be replaced"""
|
||||||
if name in self._named_parameters:
|
if name in self._named_parameters:
|
||||||
@@ -333,14 +359,19 @@ class LayerwiseOffloadManager:
|
|||||||
self._gpu_layers.add(layer_idx)
|
self._gpu_layers.add(layer_idx)
|
||||||
|
|
||||||
@torch.compiler.disable
|
@torch.compiler.disable
|
||||||
def release_layer(self, layer_idx: int) -> None:
|
def release_layer(self, layer_idx: int, force: bool = False) -> None:
|
||||||
"""
|
"""
|
||||||
lightweight release layer weights
|
lightweight release layer weights
|
||||||
Basically set the reference count to the gpu weight tensor to zero. The weights on cpu is untouched
|
Basically set the reference count to the gpu weight tensor to zero. The weights on cpu is untouched
|
||||||
|
|
||||||
|
Leading resident layers are kept across denoise steps
|
||||||
"""
|
"""
|
||||||
if not self.enabled or self.device is None:
|
if not self.enabled or self.device is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if not force and layer_idx < self._retained_layers:
|
||||||
|
return
|
||||||
|
|
||||||
# clear prefetch event, since it's useless and needs to be reset
|
# clear prefetch event, since it's useless and needs to be reset
|
||||||
self._prefetch_events.pop(layer_idx, None)
|
self._prefetch_events.pop(layer_idx, None)
|
||||||
|
|
||||||
@@ -359,13 +390,15 @@ class LayerwiseOffloadManager:
|
|||||||
|
|
||||||
@torch.compiler.disable
|
@torch.compiler.disable
|
||||||
def release_all(self) -> None:
|
def release_all(self) -> None:
|
||||||
|
"""Release every layer, including the resident ones: this ends the
|
||||||
|
denoise stage that the resident set is scoped to."""
|
||||||
if not self.enabled or self.device is None:
|
if not self.enabled or self.device is None:
|
||||||
return
|
return
|
||||||
if self.copy_stream is not None:
|
if self.copy_stream is not None:
|
||||||
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
|
torch.get_device_module().current_stream().wait_stream(self.copy_stream)
|
||||||
|
|
||||||
for layer_idx in list(self._gpu_layers):
|
for layer_idx in list(self._gpu_layers):
|
||||||
self.release_layer(layer_idx)
|
self.release_layer(layer_idx, force=True)
|
||||||
|
|
||||||
@torch.compiler.disable
|
@torch.compiler.disable
|
||||||
def load_all_layers(self) -> None:
|
def load_all_layers(self) -> None:
|
||||||
@@ -521,6 +554,7 @@ class LayerwiseOffloadManager:
|
|||||||
def make_pre_hook(i):
|
def make_pre_hook(i):
|
||||||
def hook(module, input):
|
def hook(module, input):
|
||||||
if i == 0:
|
if i == 0:
|
||||||
|
self._activate_residency()
|
||||||
self.prepare_for_next_req(non_blocking=False)
|
self.prepare_for_next_req(non_blocking=False)
|
||||||
if i not in self._gpu_layers:
|
if i not in self._gpu_layers:
|
||||||
# LTX audio VAE traverses decoder.up in reverse order
|
# LTX audio VAE traverses decoder.up in reverse order
|
||||||
@@ -589,6 +623,14 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
else:
|
else:
|
||||||
prefetch_size = int(server_args.dit_offload_prefetch_size)
|
prefetch_size = int(server_args.dit_offload_prefetch_size)
|
||||||
|
|
||||||
|
resident_value = server_args.dit_layerwise_resident_layers
|
||||||
|
if resident_value <= 0:
|
||||||
|
resident_layers = 0
|
||||||
|
elif resident_value < 1.0:
|
||||||
|
resident_layers = max(1, int(round(resident_value * num_layers)))
|
||||||
|
else:
|
||||||
|
resident_layers = min(num_layers, int(resident_value))
|
||||||
|
|
||||||
manager = LayerwiseOffloadManager(
|
manager = LayerwiseOffloadManager(
|
||||||
model=self,
|
model=self,
|
||||||
layers_attr_str=layer_name,
|
layers_attr_str=layer_name,
|
||||||
@@ -596,6 +638,7 @@ class LayerwiseOffloadableModuleMixin:
|
|||||||
enabled=True,
|
enabled=True,
|
||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||||
prefetch_size=prefetch_size,
|
prefetch_size=prefetch_size,
|
||||||
|
resident_layers=resident_layers,
|
||||||
)
|
)
|
||||||
self.layerwise_offload_managers.append(manager)
|
self.layerwise_offload_managers.append(manager)
|
||||||
configured_layer_names.append(layer_name)
|
configured_layer_names.append(layer_name)
|
||||||
@@ -674,6 +717,15 @@ def is_layerwise_offloaded_module(module: torch.nn.Module) -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_resident_layerwise_module(module: torch.nn.Module) -> bool:
|
||||||
|
"""True if the module keeps leading DiT layers resident beyond the streaming
|
||||||
|
prefetch window.
|
||||||
|
"""
|
||||||
|
return isinstance(module, LayerwiseOffloadableModuleMixin) and any(
|
||||||
|
manager.holds_residents for manager in module.layerwise_offload_managers
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_layerwise_offload_component_names_for_pipeline(
|
def get_layerwise_offload_component_names_for_pipeline(
|
||||||
modules: Mapping[str, object],
|
modules: Mapping[str, object],
|
||||||
component_names: Sequence[str] | None = None,
|
component_names: Sequence[str] | None = None,
|
||||||
|
|||||||
@@ -268,6 +268,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
dit_layerwise_offload: bool | None = None
|
dit_layerwise_offload: bool | None = None
|
||||||
layerwise_offload_components: list[str] | None = None
|
layerwise_offload_components: list[str] | None = None
|
||||||
dit_offload_prefetch_size: float = 0.0
|
dit_offload_prefetch_size: float = 0.0
|
||||||
|
# If set, keep this many leading DiT layers resident on GPU
|
||||||
|
dit_layerwise_resident_layers: float = 0.0
|
||||||
offload_during_compile: bool = True
|
offload_during_compile: bool = True
|
||||||
text_encoder_cpu_offload: bool | None = None
|
text_encoder_cpu_offload: bool | None = None
|
||||||
image_encoder_cpu_offload: bool | None = None
|
image_encoder_cpu_offload: bool | None = None
|
||||||
@@ -1638,6 +1640,18 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
default=ServerArgs.dit_offload_prefetch_size,
|
default=ServerArgs.dit_offload_prefetch_size,
|
||||||
help="The size of prefetch for dit-layerwise-offload. If the value is between 0.0 and 1.0, it is treated as a ratio of the total number of layers. If the value is >= 1, it is treated as the absolute number of layers. 0.0 means prefetch 1 layer (lowest memory). Values above 0.5 might have peak memory close to no offload but worse performance.",
|
help="The size of prefetch for dit-layerwise-offload. If the value is between 0.0 and 1.0, it is treated as a ratio of the total number of layers. If the value is >= 1, it is treated as the absolute number of layers. 0.0 means prefetch 1 layer (lowest memory). Values above 0.5 might have peak memory close to no offload but worse performance.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dit-layerwise-resident-layers",
|
||||||
|
type=float,
|
||||||
|
default=ServerArgs.dit_layerwise_resident_layers,
|
||||||
|
help="With --dit-layerwise-offload, keep this many leading DiT layers "
|
||||||
|
"permanently resident on GPU (retained across denoise steps) and stream "
|
||||||
|
"only the tail with --dit-offload-prefetch-size. 0.0 = off (pure "
|
||||||
|
"streaming). Between 0.0 and 1.0 = ratio of layers; >= 1 = absolute "
|
||||||
|
"count. Unlike raising the prefetch size, resident layers are transferred "
|
||||||
|
"once (not re-streamed every step), so this trades VRAM for lower denoise "
|
||||||
|
"latency when memory is available.",
|
||||||
|
)
|
||||||
|
|
||||||
# offload flags
|
# offload flags
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -2267,6 +2281,30 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
"We do not recommend --dit-offload-prefetch-size to be between 0.5 and 1.0"
|
"We do not recommend --dit-offload-prefetch-size to be between 0.5 and 1.0"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# validate dit_layerwise_resident_layers (same ratio/absolute convention)
|
||||||
|
if self.dit_layerwise_resident_layers < 0.0:
|
||||||
|
raise ValueError("dit_layerwise_resident_layers must be non-negative")
|
||||||
|
if self.dit_layerwise_resident_layers >= 1 and (
|
||||||
|
isinstance(self.dit_layerwise_resident_layers, float)
|
||||||
|
and not self.dit_layerwise_resident_layers.is_integer()
|
||||||
|
):
|
||||||
|
self.dit_layerwise_resident_layers = int(
|
||||||
|
math.floor(self.dit_layerwise_resident_layers)
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Invalid --dit-layerwise-resident-layers value passed, truncated to: "
|
||||||
|
f"{self.dit_layerwise_resident_layers}"
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
self.dit_layerwise_resident_layers > 0
|
||||||
|
and not self.is_dit_layerwise_offload_selected
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"--dit-layerwise-resident-layers has no effect because the DiT is not "
|
||||||
|
"layerwise-offloaded. It only applies together with "
|
||||||
|
"--dit-layerwise-offload (or 'dit' in --layerwise-offload-components)."
|
||||||
|
)
|
||||||
|
|
||||||
# validate layerwise offload conflicts
|
# validate layerwise offload conflicts
|
||||||
if envs.SGLANG_CACHE_DIT_ENABLED and self.use_fsdp_inference:
|
if envs.SGLANG_CACHE_DIT_ENABLED and self.use_fsdp_inference:
|
||||||
if self.is_arg_explicitly_set("use_fsdp_inference"):
|
if self.is_arg_explicitly_set("use_fsdp_inference"):
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
|||||||
configure_layerwise_offload_modules,
|
configure_layerwise_offload_modules,
|
||||||
get_layerwise_offload_component_names_for_pipeline,
|
get_layerwise_offload_component_names_for_pipeline,
|
||||||
is_layerwise_offloaded_module,
|
is_layerwise_offloaded_module,
|
||||||
|
is_resident_layerwise_module,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -161,6 +162,7 @@ def _server_args(**kwargs):
|
|||||||
image_encoder_cpu_offload=False,
|
image_encoder_cpu_offload=False,
|
||||||
vae_cpu_offload=False,
|
vae_cpu_offload=False,
|
||||||
dit_offload_prefetch_size=1,
|
dit_offload_prefetch_size=1,
|
||||||
|
dit_layerwise_resident_layers=0.0,
|
||||||
pin_cpu_memory=False,
|
pin_cpu_memory=False,
|
||||||
)
|
)
|
||||||
defaults.update(kwargs)
|
defaults.update(kwargs)
|
||||||
@@ -556,3 +558,134 @@ def test_layerwise_offload_aligns_contiguous_tensor_offsets(monkeypatch):
|
|||||||
assert restored_bias.data_ptr() % 32 == 0
|
assert restored_bias.data_ptr() % 32 == 0
|
||||||
assert torch.equal(restored_weight, original_weight)
|
assert torch.equal(restored_weight, original_weight)
|
||||||
assert torch.equal(restored_bias, original_bias)
|
assert torch.equal(restored_bias, original_bias)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# --dit-layerwise-resident-layers: keep N leading layers resident (retained
|
||||||
|
# across denoise steps), streaming only the tail with the prefetch window.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
class _MultiBlockModel(torch.nn.Module):
|
||||||
|
def __init__(self, n: int) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)])
|
||||||
|
|
||||||
|
|
||||||
|
class _ResidentComponent(torch.nn.Module, LayerwiseOffloadableModuleMixin):
|
||||||
|
layer_names = ["blocks"]
|
||||||
|
|
||||||
|
def __init__(self, n: int) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.blocks = torch.nn.ModuleList([_DummyBlock() for _ in range(n)])
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_fake_device(monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _resident_manager(model, *, num_layers, prefetch_size=1, resident_layers=0):
|
||||||
|
return LayerwiseOffloadManager(
|
||||||
|
model=model,
|
||||||
|
layers_attr_str="blocks",
|
||||||
|
num_layers=num_layers,
|
||||||
|
enabled=True,
|
||||||
|
pin_cpu_memory=False,
|
||||||
|
prefetch_size=prefetch_size,
|
||||||
|
resident_layers=resident_layers,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _arm_residency(manager):
|
||||||
|
"""Mimic the first-layer pre-hook: arm the resident set, then pin it."""
|
||||||
|
manager._activate_residency()
|
||||||
|
manager.prepare_for_next_req(non_blocking=False)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resident_layers_stay_pinned_until_stage_teardown(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
manager = _resident_manager(
|
||||||
|
_MultiBlockModel(4), num_layers=4, prefetch_size=1, resident_layers=2
|
||||||
|
)
|
||||||
|
# The resident set is armed on the first forward, not at construction.
|
||||||
|
assert manager._retained_layers == 0
|
||||||
|
|
||||||
|
_arm_residency(manager)
|
||||||
|
assert manager._retained_layers == 2
|
||||||
|
assert {0, 1} <= manager._gpu_layers
|
||||||
|
|
||||||
|
# A non-force release keeps the leading resident layers pinned across steps.
|
||||||
|
manager.release_layer(0)
|
||||||
|
manager.release_layer(1)
|
||||||
|
assert {0, 1} <= manager._gpu_layers
|
||||||
|
|
||||||
|
# force=True (teardown) overrides the retention.
|
||||||
|
manager.release_layer(0, force=True)
|
||||||
|
assert 0 not in manager._gpu_layers
|
||||||
|
manager.release_all() # ends the denoise stage: residents go too
|
||||||
|
assert not manager._gpu_layers
|
||||||
|
|
||||||
|
|
||||||
|
def test_resident_layers_off_by_default_streams_everything(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
manager = _resident_manager(
|
||||||
|
_MultiBlockModel(4), num_layers=4, prefetch_size=1, resident_layers=0
|
||||||
|
)
|
||||||
|
_arm_residency(manager)
|
||||||
|
|
||||||
|
assert manager._retained_layers == 0
|
||||||
|
assert manager.holds_residents is False
|
||||||
|
|
||||||
|
manager.prefetch_layer(2, non_blocking=False)
|
||||||
|
manager.release_layer(2) # no residents -> released like plain streaming
|
||||||
|
assert 2 not in manager._gpu_layers
|
||||||
|
|
||||||
|
|
||||||
|
def test_prepare_for_next_req_repins_residents(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
manager = _resident_manager(
|
||||||
|
_MultiBlockModel(6), num_layers=6, prefetch_size=1, resident_layers=3
|
||||||
|
)
|
||||||
|
_arm_residency(manager)
|
||||||
|
manager.release_all()
|
||||||
|
assert not manager._gpu_layers
|
||||||
|
|
||||||
|
# The next denoise re-pins the resident set (union of prefetch window + residents).
|
||||||
|
manager.prepare_for_next_req(non_blocking=False)
|
||||||
|
assert {0, 1, 2} <= manager._gpu_layers
|
||||||
|
|
||||||
|
|
||||||
|
def test_holds_residents_reflects_configuration(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
resident = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=2)
|
||||||
|
streaming = _resident_manager(_MultiBlockModel(3), num_layers=3, resident_layers=0)
|
||||||
|
assert resident.holds_residents is True
|
||||||
|
assert streaming.holds_residents is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_is_resident_layerwise_module_detector():
|
||||||
|
class _Comp(torch.nn.Module, LayerwiseOffloadableModuleMixin):
|
||||||
|
pass
|
||||||
|
|
||||||
|
comp = _Comp()
|
||||||
|
comp.layerwise_offload_managers = [SimpleNamespace(holds_residents=True)]
|
||||||
|
assert is_resident_layerwise_module(comp) is True
|
||||||
|
|
||||||
|
comp.layerwise_offload_managers = [SimpleNamespace(holds_residents=False)]
|
||||||
|
assert is_resident_layerwise_module(comp) is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_configure_resolves_resident_layers_absolute(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
comp = _ResidentComponent(8)
|
||||||
|
comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=3))
|
||||||
|
assert comp.layerwise_offload_managers[0].resident_layers == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_configure_resolves_resident_layers_ratio(monkeypatch):
|
||||||
|
_patch_fake_device(monkeypatch)
|
||||||
|
comp = _ResidentComponent(8)
|
||||||
|
comp.configure_layerwise_offload(_server_args(dit_layerwise_resident_layers=0.5))
|
||||||
|
# 0.5 * 8 = 4 leading layers resident
|
||||||
|
assert comp.layerwise_offload_managers[0].resident_layers == 4
|
||||||
|
|||||||
Reference in New Issue
Block a user