Fix KV-canary workspace accounting after graph capture (#38596)
This commit is contained in:
@@ -197,6 +197,10 @@ class VerifyPlan:
|
|||||||
verify_num_valid: torch.Tensor
|
verify_num_valid: torch.Tensor
|
||||||
enable: torch.Tensor
|
enable: torch.Tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def allocation_bytes(verify_capacity: int) -> int:
|
||||||
|
return 4 * verify_capacity * torch.int64.itemsize + 2 * torch.int32.itemsize
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def allocate(cls, *, verify_capacity: int, device: torch.device) -> VerifyPlan:
|
def allocate(cls, *, verify_capacity: int, device: torch.device) -> VerifyPlan:
|
||||||
if verify_capacity <= 0:
|
if verify_capacity <= 0:
|
||||||
|
|||||||
@@ -46,6 +46,12 @@ class WritePlan:
|
|||||||
write_seed_slot_indices: torch.Tensor
|
write_seed_slot_indices: torch.Tensor
|
||||||
write_num_valid_reqs: torch.Tensor
|
write_num_valid_reqs: torch.Tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def allocation_bytes(write_req_capacity: int) -> int:
|
||||||
|
return (
|
||||||
|
2 * write_req_capacity + 1
|
||||||
|
) * torch.int64.itemsize + torch.int32.itemsize
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def allocate(
|
def allocate(
|
||||||
cls,
|
cls,
|
||||||
|
|||||||
@@ -3,6 +3,10 @@ from __future__ import annotations
|
|||||||
import math
|
import math
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from sglang.kernels.ops.kv_canary.verify import VerifyPlan
|
||||||
|
from sglang.kernels.ops.kv_canary.write import WritePlan
|
||||||
|
from sglang.srt.kv_canary.expected_inputs import ExpectedInputs
|
||||||
|
from sglang.srt.kv_canary.plan_input import PlanInput
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_exec,
|
get_exec,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
@@ -116,3 +120,14 @@ class CanaryLaunchCapacities:
|
|||||||
per_forward_write_req_capacity=max_bs,
|
per_forward_write_req_capacity=max_bs,
|
||||||
per_forward_write_entry_capacity=write_entry_capacity,
|
per_forward_write_entry_capacity=write_entry_capacity,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def per_forward_workspace_bytes(self, *, num_buffer_groups: int) -> int:
|
||||||
|
return (
|
||||||
|
num_buffer_groups
|
||||||
|
* (
|
||||||
|
VerifyPlan.allocation_bytes(self.per_forward_verify_capacity)
|
||||||
|
+ WritePlan.allocation_bytes(self.per_forward_write_req_capacity)
|
||||||
|
)
|
||||||
|
+ ExpectedInputs.allocation_bytes(self.per_forward_write_entry_capacity)
|
||||||
|
+ PlanInput.allocation_bytes(self.per_forward_write_req_capacity)
|
||||||
|
)
|
||||||
|
|||||||
@@ -10,6 +10,10 @@ class ExpectedInputs:
|
|||||||
tokens: torch.Tensor
|
tokens: torch.Tensor
|
||||||
positions: torch.Tensor
|
positions: torch.Tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def allocation_bytes(capacity: int) -> int:
|
||||||
|
return 2 * capacity * torch.int64.itemsize
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def allocate(cls, *, capacity: int, device: torch.device) -> ExpectedInputs:
|
def allocate(cls, *, capacity: int, device: torch.device) -> ExpectedInputs:
|
||||||
return cls(
|
return cls(
|
||||||
|
|||||||
@@ -41,6 +41,10 @@ class PlanInput:
|
|||||||
extend_seq_lens: torch.Tensor
|
extend_seq_lens: torch.Tensor
|
||||||
req_to_verify_expected_tokens_valid_lens: torch.Tensor
|
req_to_verify_expected_tokens_valid_lens: torch.Tensor
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def allocation_bytes(bs_capacity: int) -> int:
|
||||||
|
return 4 * bs_capacity * torch.int64.itemsize
|
||||||
|
|
||||||
def zero_(self) -> None:
|
def zero_(self) -> None:
|
||||||
self.req_pool_indices.zero_()
|
self.req_pool_indices.zero_()
|
||||||
self.prefix_lens.zero_()
|
self.prefix_lens.zero_()
|
||||||
|
|||||||
@@ -64,6 +64,7 @@ class CanaryManager:
|
|||||||
self._model_forward_bracket_depth: int = 0
|
self._model_forward_bracket_depth: int = 0
|
||||||
|
|
||||||
self._buffer_groups: tuple[CanaryBufferGroup, ...] = tuple(buffer_groups)
|
self._buffer_groups: tuple[CanaryBufferGroup, ...] = tuple(buffer_groups)
|
||||||
|
self._launch_capacities = launch_capacities
|
||||||
|
|
||||||
self._device_state = CanaryDeviceState.allocate(
|
self._device_state = CanaryDeviceState.allocate(
|
||||||
config=config,
|
config=config,
|
||||||
@@ -167,6 +168,11 @@ class CanaryManager:
|
|||||||
for _ in range(num_sfms)
|
for _ in range(num_sfms)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def per_forward_workspace_bytes(self) -> int:
|
||||||
|
return self._launch_capacities.per_forward_workspace_bytes(
|
||||||
|
num_buffer_groups=len(self._buffer_groups)
|
||||||
|
)
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def with_active_single_forward_manager(self, index: int) -> Iterator[None]:
|
def with_active_single_forward_manager(self, index: int) -> Iterator[None]:
|
||||||
assert self._active_single_forward_manager_index is None, (
|
assert self._active_single_forward_manager_index is None, (
|
||||||
|
|||||||
@@ -1115,7 +1115,13 @@ class Scheduler(
|
|||||||
self.draft_worker.prewarm_sampling()
|
self.draft_worker.prewarm_sampling()
|
||||||
if model_runner.token_to_kv_pool.post_capture_active:
|
if model_runner.token_to_kv_pool.post_capture_active:
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
model_runner.post_capture_resize_kv_pool()
|
model_runner.post_capture_resize_kv_pool(
|
||||||
|
draft_runners=(
|
||||||
|
self.draft_worker._draft_model_runners()
|
||||||
|
if self.draft_worker is not None
|
||||||
|
else ()
|
||||||
|
)
|
||||||
|
)
|
||||||
self.kv_cache_allocation_time += time.perf_counter() - tic
|
self.kv_cache_allocation_time += time.perf_counter() - tic
|
||||||
|
|
||||||
if get_model().is_startup_weight_load_overlap:
|
if get_model().is_startup_weight_load_overlap:
|
||||||
|
|||||||
@@ -966,8 +966,8 @@ class ModelRunner:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def post_capture_resize_kv_pool(self):
|
def post_capture_resize_kv_pool(self, *, draft_runners=()):
|
||||||
resize = compute_post_capture_kv_resize(self)
|
resize = compute_post_capture_kv_resize(self, draft_runners=draft_runners)
|
||||||
self.max_total_num_tokens = resize.max_total_num_tokens
|
self.max_total_num_tokens = resize.max_total_num_tokens
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
self.full_max_total_num_tokens = resize.full_max_total_num_tokens
|
self.full_max_total_num_tokens = resize.full_max_total_num_tokens
|
||||||
|
|||||||
@@ -48,6 +48,8 @@ class PostCaptureKVResize(msgspec.Struct, frozen=True, kw_only=True):
|
|||||||
|
|
||||||
def compute_post_capture_kv_resize(
|
def compute_post_capture_kv_resize(
|
||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
|
*,
|
||||||
|
draft_runners: tuple[ModelRunner, ...] = (),
|
||||||
) -> PostCaptureKVResize:
|
) -> PostCaptureKVResize:
|
||||||
"""Resize the KV pool after capture and return the new sizes for the
|
"""Resize the KV pool after capture and return the new sizes for the
|
||||||
orchestrator to assign. Takes the live ModelRunner because it reads
|
orchestrator to assign. Takes the live ModelRunner because it reads
|
||||||
@@ -97,9 +99,19 @@ def compute_post_capture_kv_resize(
|
|||||||
is_multimodal=model_runner.model_config.is_multimodal,
|
is_multimodal=model_runner.model_config.is_multimodal,
|
||||||
mm_feature_transport=get_mm().mm_feature_transport,
|
mm_feature_transport=get_mm().mm_feature_transport,
|
||||||
)
|
)
|
||||||
|
# Sequential target/draft forwards reuse workspace at unchanged capacities.
|
||||||
|
canary_workspace_bytes = max(
|
||||||
|
(
|
||||||
|
runner.canary_manager.per_forward_workspace_bytes()
|
||||||
|
for runner in (model_runner, *draft_runners)
|
||||||
|
if runner.canary_manager is not None
|
||||||
|
),
|
||||||
|
default=0,
|
||||||
|
)
|
||||||
budget_bytes = (
|
budget_bytes = (
|
||||||
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
|
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
|
||||||
+ pool.post_capture_backed_bytes
|
+ pool.post_capture_backed_bytes
|
||||||
|
- canary_workspace_bytes
|
||||||
)
|
)
|
||||||
config = model_runner.kv_cache_configurator.config_from_budget(
|
config = model_runner.kv_cache_configurator.config_from_budget(
|
||||||
budget_bytes, cap_tokens=model_runner.max_total_num_tokens
|
budget_bytes, cap_tokens=model_runner.max_total_num_tokens
|
||||||
@@ -107,6 +119,11 @@ def compute_post_capture_kv_resize(
|
|||||||
pool.finalize_backing(config)
|
pool.finalize_backing(config)
|
||||||
model_runner.token_to_kv_pool_allocator.resize(config)
|
model_runner.token_to_kv_pool_allocator.resize(config)
|
||||||
model_runner.req_to_token_pool.reset_aux_cache_allocator()
|
model_runner.req_to_token_pool.reset_aux_cache_allocator()
|
||||||
|
if canary_workspace_bytes:
|
||||||
|
logger.info(
|
||||||
|
"Post-capture KV sizing: KV-canary per-forward workspace %.2f GB",
|
||||||
|
canary_workspace_bytes / (1 << 30),
|
||||||
|
)
|
||||||
|
|
||||||
capped_max_running_requests = None
|
capped_max_running_requests = None
|
||||||
if model_runner.max_running_requests is not None:
|
if model_runner.max_running_requests is not None:
|
||||||
|
|||||||
@@ -2,7 +2,13 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.kv_canary.verify import VerifyPlan
|
||||||
|
from sglang.kernels.ops.kv_canary.write import WritePlan
|
||||||
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
||||||
|
from sglang.srt.kv_canary.expected_inputs import ExpectedInputs
|
||||||
|
from sglang.srt.kv_canary.plan_input import PlanInput
|
||||||
from sglang.srt.model_executor.cuda_graph_config import (
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
Backend,
|
Backend,
|
||||||
CudaGraphConfig,
|
CudaGraphConfig,
|
||||||
@@ -88,6 +94,59 @@ class TestComputeLaunchCapacities(CustomTestCase):
|
|||||||
with self.assertRaisesRegex(ValueError, "pool_slot_count"):
|
with self.assertRaisesRegex(ValueError, "pool_slot_count"):
|
||||||
self._from_args(max_bs=1, max_seq_len=1, max_total_num_tokens=0)
|
self._from_args(max_bs=1, max_seq_len=1, max_total_num_tokens=0)
|
||||||
|
|
||||||
|
def test_workspace_matches_allocated_tensors(self) -> None:
|
||||||
|
device = torch.device("cpu")
|
||||||
|
for slots, requests, entries, groups in ((1, 1, 1, 1), (1024, 8, 128, 3)):
|
||||||
|
with self.subTest(slots=slots, groups=groups):
|
||||||
|
capacities = CanaryLaunchCapacities(
|
||||||
|
per_forward_verify_capacity=3 * slots,
|
||||||
|
per_forward_write_req_capacity=requests,
|
||||||
|
per_forward_write_entry_capacity=entries,
|
||||||
|
)
|
||||||
|
verify = VerifyPlan.allocate(verify_capacity=3 * slots, device=device)
|
||||||
|
write = WritePlan.allocate(write_req_capacity=requests, device=device)
|
||||||
|
expected = ExpectedInputs.allocate(capacity=entries, device=device)
|
||||||
|
plan = PlanInput.allocate(bs_capacity=requests, device=device)
|
||||||
|
group_tensors = (
|
||||||
|
verify.verify_slot_indices,
|
||||||
|
verify.verify_expected_tokens,
|
||||||
|
verify.verify_expected_positions,
|
||||||
|
verify.verify_prev_slot_indices,
|
||||||
|
verify.verify_num_valid,
|
||||||
|
verify.enable,
|
||||||
|
write.write_offsets,
|
||||||
|
write.write_seed_slot_indices,
|
||||||
|
write.write_num_valid_reqs,
|
||||||
|
)
|
||||||
|
shared_tensors = (
|
||||||
|
expected.tokens,
|
||||||
|
expected.positions,
|
||||||
|
plan.req_pool_indices,
|
||||||
|
plan.prefix_lens,
|
||||||
|
plan.extend_seq_lens,
|
||||||
|
plan.req_to_verify_expected_tokens_valid_lens,
|
||||||
|
)
|
||||||
|
actual_bytes = groups * sum(t.nbytes for t in group_tensors) + sum(
|
||||||
|
t.nbytes for t in shared_tensors
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
capacities.per_forward_workspace_bytes(num_buffer_groups=groups),
|
||||||
|
actual_bytes,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_workspace_scales_with_pool_slots(self) -> None:
|
||||||
|
for groups in (1, 4):
|
||||||
|
with self.subTest(groups=groups):
|
||||||
|
small, large = (
|
||||||
|
self._from_args(max_bs=8, max_seq_len=64, max_total_num_tokens=n)
|
||||||
|
for n in (1024, 4096)
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
large.per_forward_workspace_bytes(num_buffer_groups=groups)
|
||||||
|
- small.per_forward_workspace_bytes(num_buffer_groups=groups),
|
||||||
|
(4096 - 1024) * 96 * groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
+11
-3
@@ -4,7 +4,7 @@ import dataclasses
|
|||||||
import re
|
import re
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import call, patch
|
from unittest.mock import Mock, call, patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
@@ -711,7 +711,9 @@ class _SchedulerWorker:
|
|||||||
forward_stream=object(),
|
forward_stream=object(),
|
||||||
prewarm_sampling=lambda: trace.append("prewarm"),
|
prewarm_sampling=lambda: trace.append("prewarm"),
|
||||||
token_to_kv_pool=SimpleNamespace(post_capture_active=post_capture_active),
|
token_to_kv_pool=SimpleNamespace(post_capture_active=post_capture_active),
|
||||||
post_capture_resize_kv_pool=lambda: trace.append("resize"),
|
post_capture_resize_kv_pool=Mock(
|
||||||
|
side_effect=lambda *, draft_runners: trace.append("resize")
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def start_startup_weight_load(self):
|
def start_startup_weight_load(self):
|
||||||
@@ -750,7 +752,10 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
trace = []
|
trace = []
|
||||||
worker = _SchedulerWorker(trace, post_capture_active=True)
|
worker = _SchedulerWorker(trace, post_capture_active=True)
|
||||||
draft_worker = (
|
draft_worker = (
|
||||||
SimpleNamespace(prewarm_sampling=lambda: trace.append("draft_prewarm"))
|
SimpleNamespace(
|
||||||
|
prewarm_sampling=lambda: trace.append("draft_prewarm"),
|
||||||
|
_draft_model_runners=lambda: (worker.model_runner,),
|
||||||
|
)
|
||||||
if use_draft_worker
|
if use_draft_worker
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -795,6 +800,9 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
):
|
):
|
||||||
scheduler.init_model_worker()
|
scheduler.init_model_worker()
|
||||||
|
|
||||||
|
worker.model_runner.post_capture_resize_kv_pool.assert_called_once_with(
|
||||||
|
draft_runners=(worker.model_runner,) if use_draft_worker else ()
|
||||||
|
)
|
||||||
return trace
|
return trace
|
||||||
|
|
||||||
def test_serial_path_skips_overlap_hooks(self):
|
def test_serial_path_skips_overlap_hooks(self):
|
||||||
|
|||||||
@@ -0,0 +1,139 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
||||||
|
from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
||||||
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
|
Backend,
|
||||||
|
CudaGraphConfig,
|
||||||
|
PhaseConfig,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.model_runner_components import kv_pool_runtime
|
||||||
|
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, stage="base-a", runner_config="cpu")
|
||||||
|
|
||||||
|
_GIB = 1 << 30
|
||||||
|
|
||||||
|
|
||||||
|
def _manager(slots, groups):
|
||||||
|
manager = CanaryManager.__new__(CanaryManager)
|
||||||
|
manager._launch_capacities = CanaryLaunchCapacities(
|
||||||
|
per_forward_verify_capacity=3 * slots,
|
||||||
|
per_forward_write_req_capacity=128,
|
||||||
|
per_forward_write_entry_capacity=4096,
|
||||||
|
)
|
||||||
|
manager._buffer_groups = (None,) * groups
|
||||||
|
return manager
|
||||||
|
|
||||||
|
|
||||||
|
class TestCanaryHeadroom(CustomTestCase):
|
||||||
|
def _resize(self, target, drafts=(), *, graph_borrow=False, eager_gap=False):
|
||||||
|
config = MemoryPoolConfig(max_total_num_tokens=1024)
|
||||||
|
pool = Mock(post_capture_backed_bytes=2 * _GIB, dtype="bfloat16")
|
||||||
|
runner = SimpleNamespace(
|
||||||
|
token_to_kv_pool=pool,
|
||||||
|
device="cuda",
|
||||||
|
gpu_id=0,
|
||||||
|
pre_model_load_memory=32,
|
||||||
|
mem_fraction_static=0.875,
|
||||||
|
max_running_requests=16 if eager_gap else None,
|
||||||
|
model_config=SimpleNamespace(is_multimodal=False),
|
||||||
|
sampling_prewarm_result=SimpleNamespace(sampling_headroom_bytes=6 * _GIB),
|
||||||
|
canary_manager=target,
|
||||||
|
kv_cache_configurator=Mock(),
|
||||||
|
max_total_num_tokens=1_000_000,
|
||||||
|
token_to_kv_pool_allocator=Mock(),
|
||||||
|
req_to_token_pool=Mock(),
|
||||||
|
)
|
||||||
|
runner.kv_cache_configurator.config_from_budget.return_value = config
|
||||||
|
runner.kv_cache_configurator.resolve_max_num_reqs.return_value = 16
|
||||||
|
with (
|
||||||
|
get_context().override_server_args(
|
||||||
|
disaggregation_mode="null",
|
||||||
|
cuda_graph_config=CudaGraphConfig(
|
||||||
|
decode=PhaseConfig(backend=Backend.FULL, max_bs=8)
|
||||||
|
),
|
||||||
|
),
|
||||||
|
patch.object(kv_pool_runtime.torch.cuda, "synchronize"),
|
||||||
|
patch.object(
|
||||||
|
kv_pool_runtime,
|
||||||
|
"get_world_group",
|
||||||
|
return_value=SimpleNamespace(world_size=1, cpu_group=None),
|
||||||
|
),
|
||||||
|
patch.object(kv_pool_runtime, "get_available_gpu_memory", return_value=20),
|
||||||
|
patch.object(kv_pool_runtime, "mambaish_config", return_value=None),
|
||||||
|
patch.object(
|
||||||
|
kv_pool_runtime, "get_device_memory_capacity", return_value=32
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
kv_pool_runtime,
|
||||||
|
"pre_capture_activation_reserve_mb",
|
||||||
|
return_value=8 * 1024,
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
kv_pool_runtime, "graph_pool_borrow_enabled", return_value=graph_borrow
|
||||||
|
),
|
||||||
|
patch.object(kv_pool_runtime, "mm_runtime_reservation_gb", return_value=1),
|
||||||
|
):
|
||||||
|
resize = kv_pool_runtime.compute_post_capture_kv_resize(
|
||||||
|
runner,
|
||||||
|
draft_runners=tuple(SimpleNamespace(canary_manager=m) for m in drafts),
|
||||||
|
)
|
||||||
|
pool.finalize_backing.assert_called_once_with(config)
|
||||||
|
runner.token_to_kv_pool_allocator.resize.assert_called_once_with(config)
|
||||||
|
runner.req_to_token_pool.reset_aux_cache_allocator.assert_called_once_with()
|
||||||
|
self.assertEqual(resize.max_total_num_tokens, config.max_total_num_tokens)
|
||||||
|
args, kwargs = runner.kv_cache_configurator.config_from_budget.call_args
|
||||||
|
self.assertEqual(kwargs, {"cap_tokens": 1_000_000})
|
||||||
|
return args[0]
|
||||||
|
|
||||||
|
def test_canary_off_preserves_every_budget_byte(self):
|
||||||
|
for graph_borrow, eager_gap, headroom in (
|
||||||
|
(False, False, 6),
|
||||||
|
(True, False, 4),
|
||||||
|
(True, True, 8),
|
||||||
|
):
|
||||||
|
with self.subTest(graph_borrow=graph_borrow, eager_gap=eager_gap):
|
||||||
|
self.assertEqual(
|
||||||
|
self._resize(
|
||||||
|
None, (None,), graph_borrow=graph_borrow, eager_gap=eager_gap
|
||||||
|
),
|
||||||
|
(20 - headroom - 1 + 2) * _GIB,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_workspace_is_added_to_other_headroom_at_installed_capacity(self):
|
||||||
|
manager = _manager(1_000_000, 3)
|
||||||
|
workspace = manager.per_forward_workspace_bytes()
|
||||||
|
for graph_borrow, eager_gap in ((False, False), (True, False), (True, True)):
|
||||||
|
with self.subTest(graph_borrow=graph_borrow, eager_gap=eager_gap):
|
||||||
|
baseline = self._resize(
|
||||||
|
None, graph_borrow=graph_borrow, eager_gap=eager_gap
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
self._resize(
|
||||||
|
manager, graph_borrow=graph_borrow, eager_gap=eager_gap
|
||||||
|
),
|
||||||
|
baseline - workspace,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_sequential_target_and_drafts_reserve_largest_workspace(self):
|
||||||
|
small, large = _manager(1024, 1), _manager(4096, 4)
|
||||||
|
for target, drafts in (
|
||||||
|
(small, (large, None)),
|
||||||
|
(large, (small,)),
|
||||||
|
(None, (small, large)),
|
||||||
|
):
|
||||||
|
with self.subTest(target=target, drafts=drafts):
|
||||||
|
self.assertEqual(
|
||||||
|
self._resize(target, drafts),
|
||||||
|
self._resize(None) - large.per_forward_workspace_bytes(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user