forked from Karylab-cklius/vllm
[KV Offload] Unified memory layout for offloading workers (#37206)
Signed-off-by: omerpaz95 <omerpaz95@gmail.com> Co-authored-by: Or Ozeri <oro@il.ibm.com>
This commit is contained in:
@@ -2,12 +2,14 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import random
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
from vllm.v1.kv_offload.cpu.shared_offload_region import SharedOffloadRegion
|
||||
from vllm.v1.kv_offload.mediums import CPULoadStoreSpec, GPULoadStoreSpec
|
||||
from vllm.v1.kv_offload.spec import (
|
||||
CanonicalKVCacheRef,
|
||||
@@ -36,6 +38,7 @@ NUM_MAPPINGS = [3]
|
||||
@pytest.mark.parametrize("num_tensors", NUM_TENSORS)
|
||||
@pytest.mark.parametrize("seed", SEEDS)
|
||||
@pytest.mark.parametrize("device", DEVICES)
|
||||
@pytest.mark.parametrize("use_shared_memory", [False, True])
|
||||
@torch.inference_mode()
|
||||
def test_transfer(
|
||||
default_vllm_config,
|
||||
@@ -48,6 +51,7 @@ def test_transfer(
|
||||
num_tensors: int,
|
||||
seed: int,
|
||||
device: str,
|
||||
use_shared_memory: bool,
|
||||
) -> None:
|
||||
set_random_seed(seed)
|
||||
|
||||
@@ -83,10 +87,24 @@ def test_transfer(
|
||||
tensors=kv_cache_tensors,
|
||||
group_data_refs=kv_cache_groups_data_refs,
|
||||
)
|
||||
|
||||
mmap_region: SharedOffloadRegion | None = None
|
||||
if use_shared_memory:
|
||||
cpu_page_size = gpu_page_size_bytes * num_tensors * block_size_factor
|
||||
mmap_region = SharedOffloadRegion(
|
||||
instance_id=str(uuid.uuid4()),
|
||||
total_size_bytes=num_cpu_blocks * cpu_page_size,
|
||||
num_blocks=num_cpu_blocks,
|
||||
rank=0,
|
||||
num_workers=1,
|
||||
cpu_page_size=cpu_page_size,
|
||||
)
|
||||
|
||||
handlers = CpuGpuOffloadingHandlers(
|
||||
kv_caches=kv_caches,
|
||||
block_size_factor=block_size_factor,
|
||||
num_cpu_blocks=num_cpu_blocks,
|
||||
mmap_region=mmap_region,
|
||||
)
|
||||
|
||||
# select block mappings
|
||||
@@ -137,10 +155,8 @@ def test_transfer(
|
||||
if finished:
|
||||
assert finished[0].job_id == 1
|
||||
assert finished[0].success
|
||||
assert (
|
||||
finished[0].transfer_type == ("GPU", "CPU")
|
||||
if gpu_to_cpu
|
||||
else ("CPU", "GPU")
|
||||
assert finished[0].transfer_type == (
|
||||
("GPU", "CPU") if gpu_to_cpu else ("CPU", "GPU")
|
||||
)
|
||||
assert finished[0].transfer_size == (
|
||||
len(gpu_blocks) * handler.group_block_size_in_bytes[0]
|
||||
@@ -161,9 +177,9 @@ def test_transfer(
|
||||
orig_dst_tensors,
|
||||
):
|
||||
# view both GPU and CPU tensors as (n, gpu_page_size_bytes) for comparison.
|
||||
src_view = src_tensor.view(-1, gpu_page_size_bytes)
|
||||
dst_view = dst_tensor.view(-1, gpu_page_size_bytes)
|
||||
orig_dst_view = orig_dst_tensor.view(-1, gpu_page_size_bytes)
|
||||
src_view = src_tensor.reshape(-1, gpu_page_size_bytes)
|
||||
dst_view = dst_tensor.reshape(-1, gpu_page_size_bytes)
|
||||
orig_dst_view = orig_dst_tensor.reshape(-1, gpu_page_size_bytes)
|
||||
for dst_sub_block in range(num_dst_sub_blocks):
|
||||
src_sub_block = dst_to_src.get(dst_sub_block)
|
||||
if src_sub_block is not None:
|
||||
@@ -171,3 +187,12 @@ def test_transfer(
|
||||
else:
|
||||
expected = orig_dst_view[dst_sub_block]
|
||||
torch.testing.assert_close(dst_view[dst_sub_block].cpu(), expected.cpu())
|
||||
|
||||
# Drop loop-variable refs so mmap_obj has no exported buffers at cleanup.
|
||||
del orig_tensor, tensor, src_tensor, dst_tensor, orig_dst_tensor
|
||||
del src_view, dst_view, orig_dst_view, expected
|
||||
|
||||
handlers.cpu_to_gpu_handler.shutdown()
|
||||
handlers.gpu_to_cpu_handler.shutdown()
|
||||
if mmap_region:
|
||||
mmap_region.cleanup()
|
||||
|
||||
@@ -0,0 +1,625 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Unit tests for SharedOffloadRegion."""
|
||||
|
||||
import contextlib
|
||||
import mmap
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm.utils.system_utils import get_mp_context
|
||||
from vllm.v1.kv_offload.cpu.shared_offload_region import (
|
||||
SharedOffloadRegion,
|
||||
_wait_for_file_size,
|
||||
)
|
||||
|
||||
PAGE_SIZE = mmap.PAGESIZE
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers / fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _set_spawn_method(monkeypatch):
|
||||
# On WSL, NVML is not compatible with fork so vLLM auto-overrides the
|
||||
# multiprocessing start method to 'spawn' with a warning. Set it explicitly
|
||||
# here so the override is a no-op and the warning is suppressed.
|
||||
monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn")
|
||||
|
||||
|
||||
def _make_region(
|
||||
instance_id: str,
|
||||
num_blocks: int = 4,
|
||||
cpu_page_size: int = PAGE_SIZE,
|
||||
num_workers: int = 1,
|
||||
rank: int = 0,
|
||||
) -> SharedOffloadRegion:
|
||||
total_size_bytes = num_blocks * num_workers * cpu_page_size
|
||||
assert total_size_bytes % PAGE_SIZE == 0
|
||||
return SharedOffloadRegion(
|
||||
instance_id=instance_id,
|
||||
total_size_bytes=total_size_bytes,
|
||||
num_blocks=num_blocks,
|
||||
rank=rank,
|
||||
num_workers=num_workers,
|
||||
cpu_page_size=cpu_page_size,
|
||||
)
|
||||
|
||||
|
||||
def _cleanup_file(path: str) -> None:
|
||||
"""Best-effort file removal for test teardown."""
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _region(instance_id: str, **kwargs):
|
||||
"""Context manager: create one region, clean up on exit."""
|
||||
r = _make_region(instance_id, **kwargs)
|
||||
try:
|
||||
yield r
|
||||
finally:
|
||||
r.cleanup()
|
||||
_cleanup_file(r.mmap_path)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _multi_region(
|
||||
instance_id: str,
|
||||
num_workers: int,
|
||||
num_blocks: int = 4,
|
||||
cpu_page_size: int = PAGE_SIZE,
|
||||
):
|
||||
"""Context manager: create one SharedOffloadRegion per rank, clean up on exit."""
|
||||
total = num_blocks * num_workers * cpu_page_size
|
||||
regions = [
|
||||
SharedOffloadRegion(
|
||||
instance_id=instance_id,
|
||||
total_size_bytes=total,
|
||||
num_blocks=num_blocks,
|
||||
rank=rank,
|
||||
num_workers=num_workers,
|
||||
cpu_page_size=cpu_page_size,
|
||||
)
|
||||
for rank in range(num_workers)
|
||||
]
|
||||
try:
|
||||
yield regions
|
||||
finally:
|
||||
for r in regions:
|
||||
r.cleanup()
|
||||
_cleanup_file(regions[0].mmap_path)
|
||||
|
||||
|
||||
def _race_construct(
|
||||
instance_id: str,
|
||||
num_workers: int,
|
||||
num_blocks: int = 4,
|
||||
cpu_page_size: int = PAGE_SIZE,
|
||||
) -> tuple[list[SharedOffloadRegion], list[Exception]]:
|
||||
"""Spawn num_workers threads that all race to construct SharedOffloadRegion."""
|
||||
total = num_blocks * num_workers * cpu_page_size
|
||||
regions: list[SharedOffloadRegion | None] = [None] * num_workers
|
||||
errors: list[Exception] = []
|
||||
barrier = threading.Barrier(num_workers)
|
||||
|
||||
def worker(rank: int) -> None:
|
||||
barrier.wait() # all threads start at the same instant
|
||||
try:
|
||||
regions[rank] = SharedOffloadRegion(
|
||||
instance_id=instance_id,
|
||||
total_size_bytes=total,
|
||||
num_blocks=num_blocks,
|
||||
rank=rank,
|
||||
num_workers=num_workers,
|
||||
cpu_page_size=cpu_page_size,
|
||||
)
|
||||
except Exception as e:
|
||||
errors.append(e)
|
||||
|
||||
threads = [threading.Thread(target=worker, args=(i,)) for i in range(num_workers)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
return [r for r in regions if r is not None], errors
|
||||
|
||||
|
||||
def _mp_race_construct_and_write(
|
||||
instance_id: str,
|
||||
total_bytes: int,
|
||||
num_blocks: int,
|
||||
rank: int,
|
||||
num_workers: int,
|
||||
cpu_page_size: int,
|
||||
fill_value: int,
|
||||
done_queue,
|
||||
cleanup_queue,
|
||||
) -> None:
|
||||
"""Race to construct a SharedOffloadRegion, write fill_value, then wait
|
||||
for the parent's cleanup signal before tearing down. The wait gives the
|
||||
parent a window to read the raw mmap before the creator removes the file."""
|
||||
try:
|
||||
region = SharedOffloadRegion(
|
||||
instance_id=instance_id,
|
||||
total_size_bytes=total_bytes,
|
||||
num_blocks=num_blocks,
|
||||
rank=rank,
|
||||
num_workers=num_workers,
|
||||
cpu_page_size=cpu_page_size,
|
||||
)
|
||||
t = region.create_next_view(cpu_page_size)
|
||||
t[:, :] = fill_value
|
||||
done_queue.put({"rank": rank, "error": None})
|
||||
cleanup_queue.get() # wait for parent's verification to finish
|
||||
del t # release view before cleanup to avoid BufferError
|
||||
region.cleanup()
|
||||
except Exception as e:
|
||||
done_queue.put({"rank": rank, "error": repr(e)})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def iid():
|
||||
"""Fresh instance ID for each test."""
|
||||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_next_view — shape, stride and storage offset
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_next_view_shape_and_stride(iid):
|
||||
"""Returned tensor must have shape (num_blocks, tensor_page_size) and
|
||||
stride (row_stride, 1) where row_stride = cpu_page_size * num_workers."""
|
||||
with _region(iid, num_blocks=4, cpu_page_size=2 * PAGE_SIZE) as r:
|
||||
t = r.create_next_view(PAGE_SIZE)
|
||||
assert t.shape == (4, PAGE_SIZE)
|
||||
# num_workers=1 → row_stride = cpu_page_size
|
||||
assert t.stride() == (2 * PAGE_SIZE, 1)
|
||||
del t
|
||||
|
||||
|
||||
def test_create_next_view_storage_offset_rank0(iid):
|
||||
"""rank=0 worker's first tensor must start at byte 0 of the mmap."""
|
||||
with _region(iid, cpu_page_size=PAGE_SIZE, num_workers=2, rank=0) as r:
|
||||
t = r.create_next_view(PAGE_SIZE)
|
||||
assert t.data_ptr() == r._base.data_ptr() # storage_offset == 0
|
||||
del t
|
||||
|
||||
|
||||
def test_create_next_view_storage_offset_rank1(iid):
|
||||
"""rank=1 worker's first tensor must start cpu_page_size bytes into the mmap."""
|
||||
with _multi_region(iid, num_workers=2, num_blocks=4) as (r0, r1):
|
||||
t1 = r1.create_next_view(PAGE_SIZE)
|
||||
assert t1.data_ptr() == r1._base.data_ptr() + PAGE_SIZE
|
||||
del t1
|
||||
|
||||
|
||||
def test_create_next_view_row_stride_with_multiple_workers(iid):
|
||||
"""With num_workers=4, row_stride must be 4 * cpu_page_size."""
|
||||
with _region(iid, num_blocks=2, num_workers=4) as r:
|
||||
t = r.create_next_view(PAGE_SIZE)
|
||||
assert t.stride(0) == 4 * PAGE_SIZE
|
||||
del t
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_next_view — cursor advancement
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_next_view_cursor_advances(iid):
|
||||
"""Each call to create_next_view must advance _worker_offset by tensor_page_size."""
|
||||
with _region(iid, cpu_page_size=3 * PAGE_SIZE) as r:
|
||||
assert r._worker_offset == 0
|
||||
r.create_next_view(PAGE_SIZE)
|
||||
assert r._worker_offset == PAGE_SIZE
|
||||
r.create_next_view(PAGE_SIZE)
|
||||
assert r._worker_offset == 2 * PAGE_SIZE
|
||||
r.create_next_view(PAGE_SIZE)
|
||||
assert r._worker_offset == 3 * PAGE_SIZE # exactly at area end
|
||||
|
||||
|
||||
def test_create_next_view_exact_fill_succeeds(iid):
|
||||
"""Allocations whose total exactly equals cpu_page_size must all succeed."""
|
||||
with _region(iid, cpu_page_size=2 * PAGE_SIZE) as r:
|
||||
r.create_next_view(PAGE_SIZE) # first half
|
||||
r.create_next_view(PAGE_SIZE) # fills to area end — must not raise
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_next_view — overflow guard
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_next_view_single_overflow_raises(iid):
|
||||
"""A single allocation larger than cpu_page_size must raise AssertionError."""
|
||||
with (
|
||||
_region(iid) as r,
|
||||
pytest.raises(AssertionError, match="exceeds worker area end"),
|
||||
):
|
||||
r.create_next_view(PAGE_SIZE + 1)
|
||||
|
||||
|
||||
def test_create_next_view_cumulative_overflow_raises(iid):
|
||||
"""Successive allocations that cumulatively exceed cpu_page_size must raise."""
|
||||
with _region(iid, cpu_page_size=2 * PAGE_SIZE) as r:
|
||||
r.create_next_view(PAGE_SIZE) # ok — half used
|
||||
r.create_next_view(PAGE_SIZE) # ok — full
|
||||
with pytest.raises(AssertionError, match="exceeds worker area end"):
|
||||
r.create_next_view(1) # one byte too many
|
||||
|
||||
|
||||
def test_create_next_view_overflow_does_not_mutate_cursor(iid):
|
||||
"""A failed create_next_view must leave _worker_offset unchanged."""
|
||||
with _region(iid) as r:
|
||||
offset_before = r._worker_offset
|
||||
with pytest.raises(AssertionError):
|
||||
r.create_next_view(PAGE_SIZE + 1)
|
||||
assert r._worker_offset == offset_before
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_next_view — data correctness and layout
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_next_view_write_visible_in_raw_mmap(iid):
|
||||
"""Writes into a create_next_view view must appear at the correct raw mmap offset"""
|
||||
with _region(iid, num_blocks=4) as r:
|
||||
t = r.create_next_view(PAGE_SIZE)
|
||||
t[2, :] = 42 # write to block row 2
|
||||
|
||||
raw = memoryview(r.mmap_obj)
|
||||
# num_workers=1 → row_stride = PAGE_SIZE; block 2 starts at byte 2*PAGE_SIZE
|
||||
chunk = bytes(raw[2 * PAGE_SIZE : 3 * PAGE_SIZE])
|
||||
assert all(b == 42 for b in chunk)
|
||||
del raw, t
|
||||
|
||||
|
||||
def test_create_next_view_multi_tensor_layout(iid):
|
||||
"""Two tensors from the same worker land at consecutive byte offsets per row."""
|
||||
with _region(iid, num_blocks=2, cpu_page_size=2 * PAGE_SIZE) as r:
|
||||
ta = r.create_next_view(PAGE_SIZE)
|
||||
tb = r.create_next_view(PAGE_SIZE)
|
||||
|
||||
ta[:, :] = 1
|
||||
tb[:, :] = 2
|
||||
|
||||
raw = memoryview(r.mmap_obj)
|
||||
for blk in range(2):
|
||||
row_offset = blk * 2 * PAGE_SIZE # num_workers=1
|
||||
assert all(b == 1 for b in raw[row_offset : row_offset + PAGE_SIZE])
|
||||
assert all(
|
||||
b == 2 for b in raw[row_offset + PAGE_SIZE : row_offset + 2 * PAGE_SIZE]
|
||||
)
|
||||
del raw, ta, tb
|
||||
|
||||
|
||||
def test_create_next_view_multiprocess_slots(iid):
|
||||
"""Each worker process calls create_next_view and writes distinct data;
|
||||
the parent verifies each slot lands at the correct interleaved offset."""
|
||||
num_workers = 2
|
||||
num_blocks = 4
|
||||
total_bytes = num_blocks * num_workers * PAGE_SIZE
|
||||
|
||||
ctx = get_mp_context()
|
||||
done_queue = ctx.Queue()
|
||||
cleanup_queue = ctx.Queue()
|
||||
|
||||
# Parent is rank 0 (creator); child is rank 1 (joiner).
|
||||
region = SharedOffloadRegion(
|
||||
instance_id=iid,
|
||||
total_size_bytes=total_bytes,
|
||||
num_blocks=num_blocks,
|
||||
rank=0,
|
||||
num_workers=num_workers,
|
||||
cpu_page_size=PAGE_SIZE,
|
||||
)
|
||||
try:
|
||||
child = ctx.Process(
|
||||
target=_mp_race_construct_and_write,
|
||||
args=(
|
||||
iid,
|
||||
total_bytes,
|
||||
num_blocks,
|
||||
1,
|
||||
num_workers,
|
||||
PAGE_SIZE,
|
||||
22,
|
||||
done_queue,
|
||||
cleanup_queue,
|
||||
),
|
||||
)
|
||||
child.start()
|
||||
|
||||
t0 = region.create_next_view(PAGE_SIZE)
|
||||
t0[:, :] = 11
|
||||
|
||||
result = done_queue.get(timeout=30)
|
||||
assert result["error"] is None, result["error"]
|
||||
|
||||
raw = memoryview(region.mmap_obj)
|
||||
for blk in range(num_blocks):
|
||||
row_start = blk * num_workers * PAGE_SIZE
|
||||
w0 = bytes(raw[row_start : row_start + PAGE_SIZE])
|
||||
w1 = bytes(raw[row_start + PAGE_SIZE : row_start + 2 * PAGE_SIZE])
|
||||
assert all(b == 11 for b in w0), f"block {blk}: rank0 slot wrong"
|
||||
assert all(b == 22 for b in w1), f"block {blk}: rank1 slot wrong"
|
||||
|
||||
del raw, t0 # release before finally triggers cleanup
|
||||
cleanup_queue.put(True)
|
||||
child.join(timeout=10)
|
||||
assert child.exitcode == 0
|
||||
finally:
|
||||
region.cleanup()
|
||||
_cleanup_file(region.mmap_path)
|
||||
|
||||
|
||||
def test_create_next_view_worker_isolation(iid):
|
||||
"""Writes by worker 0 must not affect worker 1's slot and vice versa."""
|
||||
num_workers = 2
|
||||
num_blocks = 4
|
||||
with _multi_region(iid, num_workers=num_workers, num_blocks=num_blocks) as regions:
|
||||
t0 = regions[0].create_next_view(PAGE_SIZE)
|
||||
t1 = regions[1].create_next_view(PAGE_SIZE)
|
||||
|
||||
t0[:, :] = 11
|
||||
t1[:, :] = 22
|
||||
|
||||
raw = memoryview(regions[0].mmap_obj)
|
||||
for blk in range(num_blocks):
|
||||
row_start = blk * num_workers * PAGE_SIZE
|
||||
w0 = bytes(raw[row_start : row_start + PAGE_SIZE])
|
||||
w1 = bytes(raw[row_start + PAGE_SIZE : row_start + 2 * PAGE_SIZE])
|
||||
assert all(b == 11 for b in w0), f"block {blk}: worker0 slot corrupted"
|
||||
assert all(b == 22 for b in w1), f"block {blk}: worker1 slot corrupted"
|
||||
del raw, t0, t1 # release before finally triggers cleanup
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constructor — creator vs joiner semantics
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_creator_flag_set_on_first_open(iid):
|
||||
"""The first worker to open the file must have _creator == True."""
|
||||
with _region(iid) as r:
|
||||
assert r._creator is True
|
||||
|
||||
|
||||
def test_joiner_flag_not_set(iid):
|
||||
"""A second worker opening the same file must have _creator == False."""
|
||||
with _multi_region(iid, num_workers=2) as (r0, r1):
|
||||
assert r0._creator is True
|
||||
assert r1._creator is False
|
||||
|
||||
|
||||
def test_file_exists_after_construction(iid):
|
||||
"""The mmap file must be present on disk after __init__ completes."""
|
||||
with _region(iid) as r:
|
||||
assert os.path.exists(r.mmap_path)
|
||||
|
||||
|
||||
def test_file_has_correct_size(iid):
|
||||
"""The mmap file size on disk must equal total_size_bytes."""
|
||||
with _region(iid, num_blocks=4) as r:
|
||||
assert os.path.getsize(r.mmap_path) == 4 * PAGE_SIZE
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Multi-worker race — concurrent construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_multi_worker_race_exactly_one_creator(iid):
|
||||
"""When N threads race to create the same region, exactly one becomes creator."""
|
||||
num_workers = 8
|
||||
regions, errors = _race_construct(iid, num_workers=num_workers)
|
||||
try:
|
||||
assert not errors, f"Workers raised: {errors}"
|
||||
assert len(regions) == num_workers, "Some workers failed to construct"
|
||||
|
||||
creators = [r for r in regions if r._creator]
|
||||
assert len(creators) == 1, f"Expected 1 creator, got {len(creators)}"
|
||||
assert sum(1 for r in regions if not r._creator) == num_workers - 1, (
|
||||
f"Expected {num_workers - 1} non-creators, got "
|
||||
f"{sum(1 for r in regions if not r._creator)}"
|
||||
)
|
||||
|
||||
for r in regions:
|
||||
assert not r.mmap_obj.closed
|
||||
assert r.total_size_bytes == 4 * num_workers * PAGE_SIZE
|
||||
finally:
|
||||
for r in regions:
|
||||
r.cleanup()
|
||||
_cleanup_file(regions[0].mmap_path)
|
||||
|
||||
|
||||
def test_multi_worker_race_shared_memory_visible(iid):
|
||||
"""After a concurrent construction race, MAP_SHARED is intact across all workers."""
|
||||
num_workers = 4
|
||||
regions, errors = _race_construct(iid, num_workers=num_workers)
|
||||
assert not errors
|
||||
try:
|
||||
regions[0].mmap_obj[0:1] = b"\xab"
|
||||
for r in regions[1:]:
|
||||
assert memoryview(r.mmap_obj)[0:1] == b"\xab"
|
||||
finally:
|
||||
for r in regions:
|
||||
r.cleanup()
|
||||
_cleanup_file(regions[0].mmap_path)
|
||||
|
||||
|
||||
def test_multiprocess_race_construct_and_write(iid):
|
||||
"""N processes race to construct the same SharedOffloadRegion, each writes
|
||||
fill_value = rank+1 into their slot; parent verifies interleaved layout."""
|
||||
num_workers = 4
|
||||
num_blocks = 3
|
||||
total_bytes = num_blocks * num_workers * PAGE_SIZE
|
||||
|
||||
ctx = get_mp_context()
|
||||
done_queue = ctx.Queue()
|
||||
cleanup_queue = ctx.Queue()
|
||||
|
||||
procs = [
|
||||
ctx.Process(
|
||||
target=_mp_race_construct_and_write,
|
||||
args=(
|
||||
iid,
|
||||
total_bytes,
|
||||
num_blocks,
|
||||
rank,
|
||||
num_workers,
|
||||
PAGE_SIZE,
|
||||
rank + 1,
|
||||
done_queue,
|
||||
cleanup_queue,
|
||||
),
|
||||
)
|
||||
for rank in range(num_workers)
|
||||
]
|
||||
for p in procs:
|
||||
p.start()
|
||||
|
||||
results = {}
|
||||
for _ in range(num_workers):
|
||||
r = done_queue.get(timeout=30)
|
||||
results[r["rank"]] = r
|
||||
|
||||
for rank, r in results.items():
|
||||
assert r["error"] is None, f"rank {rank}: {r['error']}"
|
||||
|
||||
# Read the raw file while all workers still hold it open.
|
||||
mmap_path = f"/dev/shm/vllm_offload_{iid}.mmap"
|
||||
with open(mmap_path, "rb") as f:
|
||||
raw = f.read()
|
||||
|
||||
for blk in range(num_blocks):
|
||||
for w in range(num_workers):
|
||||
slot_start = (blk * num_workers + w) * PAGE_SIZE
|
||||
slot = raw[slot_start : slot_start + PAGE_SIZE]
|
||||
expected = w + 1 # fill_value = rank + 1
|
||||
assert all(b == expected for b in slot), (
|
||||
f"block {blk}, worker {w}: expected {expected} but got wrong bytes"
|
||||
)
|
||||
|
||||
# Unblock all workers to clean up.
|
||||
for _ in range(num_workers):
|
||||
cleanup_queue.put(True)
|
||||
for p in procs:
|
||||
p.join(timeout=10)
|
||||
assert p.exitcode == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cleanup
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_cleanup_creator_all_effects(iid):
|
||||
"""cleanup() on the creator closes mmap, closes fd, and removes the file."""
|
||||
r = _make_region(iid)
|
||||
path = r.mmap_path
|
||||
fd = r.fd
|
||||
mmap_obj = r.mmap_obj
|
||||
|
||||
r.cleanup()
|
||||
|
||||
assert mmap_obj.closed, "mmap should be closed after cleanup"
|
||||
assert not os.path.exists(path), "creator should remove the file"
|
||||
with pytest.raises(OSError):
|
||||
os.fstat(fd) # fd should be closed
|
||||
|
||||
|
||||
def test_cleanup_non_creator_all_effects(iid):
|
||||
"""cleanup() on a non-creator closes mmap and fd, but leaves the file on disk."""
|
||||
r0 = _make_region(iid) # creator
|
||||
r1 = _make_region(iid) # joiner
|
||||
path = r0.mmap_path
|
||||
fd1 = r1.fd
|
||||
mmap_obj1 = r1.mmap_obj
|
||||
try:
|
||||
r1.cleanup()
|
||||
|
||||
assert mmap_obj1.closed, "mmap should be closed after cleanup"
|
||||
assert os.path.exists(path), "non-creator must not remove the file"
|
||||
with pytest.raises(OSError):
|
||||
os.fstat(fd1) # fd should be closed
|
||||
finally:
|
||||
r0.cleanup()
|
||||
_cleanup_file(path)
|
||||
|
||||
|
||||
def test_cleanup_idempotent(iid):
|
||||
"""Calling cleanup() twice must not raise any exception."""
|
||||
r = _make_region(iid)
|
||||
r.cleanup()
|
||||
r.cleanup() # must be a no-op
|
||||
|
||||
|
||||
def test_cleanup_after_create_next_view_releases_mmap(iid):
|
||||
"""cleanup() must close the mmap even after create_next_view was called.
|
||||
create_next_view returns a view that shares storage with _base; both must be
|
||||
released before mmap.close() can succeed."""
|
||||
r = _make_region(iid)
|
||||
mmap_obj = r.mmap_obj
|
||||
|
||||
t = r.create_next_view(PAGE_SIZE)
|
||||
del t
|
||||
|
||||
r.cleanup()
|
||||
|
||||
assert mmap_obj.closed, "mmap should be closed after releasing the tensor"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _wait_for_file_size
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_wait_for_file_size_already_large_enough(tmp_path):
|
||||
"""_wait_for_file_size must return immediately when file is already big enough."""
|
||||
fd = os.open(str(tmp_path / "ready.mmap"), os.O_CREAT | os.O_RDWR, 0o600)
|
||||
try:
|
||||
os.ftruncate(fd, PAGE_SIZE)
|
||||
start = time.monotonic()
|
||||
_wait_for_file_size(fd, PAGE_SIZE, timeout=5.0)
|
||||
assert time.monotonic() - start < 0.5
|
||||
finally:
|
||||
os.close(fd)
|
||||
|
||||
|
||||
def test_wait_for_file_size_waits_for_grow(tmp_path):
|
||||
"""_wait_for_file_size must return once a background thread grows the file."""
|
||||
fd = os.open(str(tmp_path / "grow.mmap"), os.O_CREAT | os.O_RDWR, 0o600)
|
||||
try:
|
||||
|
||||
def grow():
|
||||
time.sleep(0.05)
|
||||
os.ftruncate(fd, PAGE_SIZE)
|
||||
|
||||
t = threading.Thread(target=grow)
|
||||
t.start()
|
||||
_wait_for_file_size(fd, PAGE_SIZE, timeout=5.0) # must not raise
|
||||
t.join()
|
||||
finally:
|
||||
os.close(fd)
|
||||
|
||||
|
||||
def test_wait_for_file_size_timeout(tmp_path):
|
||||
"""_wait_for_file_size must raise TimeoutError when the file never grows."""
|
||||
fd = os.open(str(tmp_path / "stuck.mmap"), os.O_CREAT | os.O_RDWR, 0o600)
|
||||
try:
|
||||
with pytest.raises(TimeoutError):
|
||||
_wait_for_file_size(fd, PAGE_SIZE, timeout=0.1)
|
||||
finally:
|
||||
os.close(fd)
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import mmap
|
||||
import os
|
||||
import time
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _wait_for_file_size(fd: int, expected_size: int, timeout: float = 30.0) -> None:
|
||||
"""Spin-wait until the file reaches expected_size (creator truncated it)."""
|
||||
deadline = time.monotonic() + timeout
|
||||
while True:
|
||||
if os.fstat(fd).st_size >= expected_size:
|
||||
return
|
||||
if time.monotonic() > deadline:
|
||||
raise TimeoutError(
|
||||
f"Timed out waiting for mmap file to reach {expected_size} bytes"
|
||||
)
|
||||
time.sleep(0.005)
|
||||
|
||||
|
||||
class SharedOffloadRegion:
|
||||
"""
|
||||
Single mmap-backed memory region shared across all workers for a
|
||||
vLLM instance. Workers coordinate via the filesystem: the first worker
|
||||
to open the file with O_EXCL becomes the creator and calls ftruncate;
|
||||
the rest open the existing file and wait until it reaches the expected
|
||||
size. Each worker then mmap()s the full file.
|
||||
|
||||
File path: /dev/shm/vllm_offload_{instance_id}.mmap
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
instance_id: str,
|
||||
total_size_bytes: int,
|
||||
num_blocks: int,
|
||||
rank: int | None,
|
||||
num_workers: int,
|
||||
cpu_page_size: int,
|
||||
) -> None:
|
||||
self.page_size = mmap.PAGESIZE
|
||||
|
||||
self.total_size_bytes = total_size_bytes
|
||||
self.mmap_path = f"/dev/shm/vllm_offload_{instance_id}.mmap"
|
||||
self._creator = False # set True only if this worker creates the file
|
||||
self.num_blocks = num_blocks
|
||||
self.rank = rank
|
||||
# interleaved-layout stride: one row = all workers' data for one block
|
||||
self._row_stride = cpu_page_size * num_workers
|
||||
if rank is not None:
|
||||
# byte offset to this worker's first slot within each block row
|
||||
self._worker_offset = rank * cpu_page_size
|
||||
# exclusive upper bound for this worker's area within each row
|
||||
self._worker_area_end = (rank + 1) * cpu_page_size
|
||||
try:
|
||||
# Exclusive create — only one worker succeeds
|
||||
self.fd: int | None = os.open(
|
||||
self.mmap_path, os.O_CREAT | os.O_EXCL | os.O_RDWR, 0o600
|
||||
)
|
||||
os.ftruncate(self.fd, self.total_size_bytes)
|
||||
self._creator = True
|
||||
logger.info(
|
||||
"Created mmap file %s (%.2f GB)",
|
||||
self.mmap_path,
|
||||
self.total_size_bytes / 1e9,
|
||||
)
|
||||
except FileExistsError:
|
||||
self.fd = os.open(self.mmap_path, os.O_RDWR)
|
||||
_wait_for_file_size(self.fd, self.total_size_bytes)
|
||||
logger.info("Opened existing mmap file %s", self.mmap_path)
|
||||
|
||||
self.mmap_obj: mmap.mmap | None = mmap.mmap(
|
||||
self.fd,
|
||||
self.total_size_bytes,
|
||||
flags=mmap.MAP_SHARED,
|
||||
prot=mmap.PROT_READ | mmap.PROT_WRITE,
|
||||
)
|
||||
|
||||
# MADV_POPULATE_WRITE was added in Linux 5.14 (value 23).
|
||||
_MADV_POPULATE_WRITE = getattr(mmap, "MADV_POPULATE_WRITE", 23)
|
||||
if rank is not None:
|
||||
# Populate only this worker's pages (one slot per block row).
|
||||
worker_offset = rank * cpu_page_size
|
||||
_t0 = time.perf_counter()
|
||||
page_size = self.page_size
|
||||
for block in range(num_blocks):
|
||||
raw_offset = block * self._row_stride + worker_offset
|
||||
aligned_offset = (raw_offset // page_size) * page_size
|
||||
end = raw_offset + cpu_page_size
|
||||
aligned_length = end - aligned_offset
|
||||
self.mmap_obj.madvise(
|
||||
_MADV_POPULATE_WRITE, aligned_offset, aligned_length
|
||||
)
|
||||
logger.debug(
|
||||
"MADV_POPULATE_WRITE loop: %d blocks in %.3f s",
|
||||
num_blocks,
|
||||
time.perf_counter() - _t0,
|
||||
)
|
||||
else:
|
||||
# No rank — populate the entire shared region in one call.
|
||||
_t0 = time.perf_counter()
|
||||
self.mmap_obj.madvise(_MADV_POPULATE_WRITE, 0, self.total_size_bytes)
|
||||
logger.debug(
|
||||
"MADV_POPULATE_WRITE entire region: %.3f s", time.perf_counter() - _t0
|
||||
)
|
||||
|
||||
self._base = torch.frombuffer(memoryview(self.mmap_obj), dtype=torch.int8)
|
||||
self._views: list[torch.Tensor] = []
|
||||
self.is_pinned: bool = False
|
||||
|
||||
def create_next_view(self, tensor_page_size: int) -> torch.Tensor:
|
||||
"""Allocate a strided int8 view for this worker, one canonical tensor.
|
||||
|
||||
Must be called once per canonical tensor. The full mmap layout is:
|
||||
|
||||
worker0_block0 | worker1_block0 | ... | worker{M-1}_block0
|
||||
worker0_block1 | worker1_block1 | ... | worker{M-1}_block1
|
||||
...
|
||||
|
||||
Each worker_block cell is cpu_page_size bytes and holds all canonical
|
||||
tensors for that worker and block concatenated:
|
||||
[ tensor0_data | tensor1_data | ... | tensor{L-1}_data ]
|
||||
|
||||
Consecutive rows are separated by row_stride = cpu_page_size * M.
|
||||
|
||||
Returns an int8 tensor of shape (num_blocks, tensor_page_size) with stride
|
||||
(row_stride, 1). Using int8 keeps stride == bytes, so swap_blocks
|
||||
address arithmetic works without any dtype conversion.
|
||||
|
||||
Args:
|
||||
tensor_page_size: Bytes per block for this tensor.
|
||||
"""
|
||||
assert self.rank is not None
|
||||
new_offset = self._worker_offset + tensor_page_size
|
||||
assert new_offset <= self._worker_area_end, (
|
||||
f"Worker offset {new_offset} exceeds worker area end "
|
||||
f"{self._worker_area_end} (overflowed by "
|
||||
f"{new_offset - self._worker_area_end} bytes)"
|
||||
)
|
||||
worker_layer_view = torch.as_strided(
|
||||
self._base,
|
||||
size=(self.num_blocks, tensor_page_size),
|
||||
stride=(self._row_stride, 1),
|
||||
storage_offset=self._worker_offset,
|
||||
)
|
||||
self._worker_offset = new_offset
|
||||
self._views.append(worker_layer_view)
|
||||
return worker_layer_view
|
||||
|
||||
def cleanup(self) -> None:
|
||||
if self.is_pinned and self._base is not None:
|
||||
base_ptr = self._base.data_ptr()
|
||||
result = torch.cuda.cudart().cudaHostUnregister(base_ptr)
|
||||
if result.value != 0:
|
||||
logger.warning(
|
||||
"cudaHostUnregister failed for rank=%d (code=%d)", self.rank, result
|
||||
)
|
||||
self.is_pinned = False
|
||||
# Release views before _base: each view holds a _base reference and a
|
||||
# direct StorageImpl reference. Freeing views first lets both refcounts
|
||||
# drop so the storage (which holds the mmap_obj buffer export) is freed
|
||||
# before mmap_obj.close() is called below.
|
||||
if self._views is not None:
|
||||
self._views.clear()
|
||||
self._base = None
|
||||
if self.mmap_obj:
|
||||
try:
|
||||
self.mmap_obj.close()
|
||||
except Exception:
|
||||
logger.warning("Failed to close mmap_obj", exc_info=True)
|
||||
self.mmap_obj = None
|
||||
if self.fd is not None:
|
||||
try:
|
||||
os.close(self.fd)
|
||||
except Exception:
|
||||
logger.warning("Failed to close fd %s", self.fd, exc_info=True)
|
||||
self.fd = None
|
||||
if self._creator and getattr(self, "mmap_path", None):
|
||||
try:
|
||||
os.unlink(self.mmap_path)
|
||||
logger.info("Removed mmap file %s", self.mmap_path)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to unlink path %s", self.mmap_path, exc_info=True
|
||||
)
|
||||
self._creator = False
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import time
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
|
||||
@@ -9,6 +10,7 @@ import torch
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm.logger import init_logger
|
||||
from vllm.utils.platform_utils import is_pin_memory_available
|
||||
from vllm.v1.kv_offload.cpu.shared_offload_region import SharedOffloadRegion
|
||||
from vllm.v1.kv_offload.mediums import BlockIDsLoadStoreSpec
|
||||
from vllm.v1.kv_offload.spec import CanonicalKVCacheRef, CanonicalKVCaches
|
||||
from vllm.v1.kv_offload.worker.worker import (
|
||||
@@ -29,37 +31,76 @@ class Transfer:
|
||||
num_bytes: int
|
||||
|
||||
|
||||
def expand_block_ids(
|
||||
def compute_sub_block_ptrs(
|
||||
block_ids: np.ndarray,
|
||||
block_size_factor: int,
|
||||
output: np.ndarray,
|
||||
tensor: torch.Tensor,
|
||||
skip_count: int = 0,
|
||||
):
|
||||
"""
|
||||
Convert a list of block IDs to a list of matching block ids,
|
||||
assuming each block is composed of actual block_size_factor blocks.
|
||||
Outputs to output tensor.
|
||||
The first skip_count blocks will be skipped.
|
||||
Note that skip_count must be less than block_size_factor.
|
||||
Compute byte pointers for sub-blocks of the given block IDs.
|
||||
|
||||
For example, if block_ids = [0, 1, 3] and block_size_factor = 4,
|
||||
then it yields [0, 1, 2, 3, 4, 5, 6, 7, 12, 13, 14, 15]
|
||||
since 0 maps to [0, 1, 2, 3]
|
||||
1 maps to [4, 5, 6, 7]
|
||||
and 3 maps to [12, 13, 14, 15]
|
||||
Each block in block_ids contains block_size_factor sub-blocks.
|
||||
The pointer for sub-block j of block b is:
|
||||
base_ptr + b * row_stride + j * sub_block_size
|
||||
|
||||
where sub_block_size = tensor.shape[1] // block_size_factor (gpu page size).
|
||||
|
||||
This handles tensors where row_stride != block_size_factor * sub_block_size
|
||||
(e.g. non-contiguous CPU tensors).
|
||||
|
||||
Args:
|
||||
block_ids: array of block IDs at the tensor's native granularity.
|
||||
block_size_factor: number of sub-blocks per block.
|
||||
output: pre-allocated int64 array to write pointers into.
|
||||
tensor: the source or destination tensor.
|
||||
skip_count: sub-blocks to skip in the first block.
|
||||
"""
|
||||
assert skip_count < block_size_factor
|
||||
|
||||
first_range = np.arange(skip_count, block_size_factor)
|
||||
full_range = np.arange(0, block_size_factor)
|
||||
num_sub_blocks = len(output)
|
||||
base_ptr = tensor.data_ptr()
|
||||
row_stride = tensor.stride(0)
|
||||
|
||||
output_idx = 0
|
||||
for i, block_id in enumerate(block_ids):
|
||||
base_block_id = block_id * block_size_factor
|
||||
indices = first_range if i == 0 else full_range
|
||||
output_end_idx = output_idx + len(indices)
|
||||
output[output_idx:output_end_idx] = base_block_id + indices
|
||||
output_idx = output_end_idx
|
||||
if block_size_factor == 1:
|
||||
# Fast path: 1:1 mapping, no sub-block expansion needed.
|
||||
output[:] = base_ptr + block_ids[:num_sub_blocks] * row_stride
|
||||
return
|
||||
|
||||
# Vectorized expansion for block_size_factor > 1.
|
||||
assert tensor.shape[1] % block_size_factor == 0
|
||||
sub_block_size = tensor.shape[1] // block_size_factor
|
||||
sub_offsets = np.arange(block_size_factor, dtype=np.int64) * sub_block_size
|
||||
# (num_blocks, 1) + (1, block_size_factor) -> (num_blocks, block_size_factor)
|
||||
all_ptrs = (
|
||||
base_ptr + block_ids.astype(np.int64)[:, np.newaxis] * row_stride
|
||||
) + sub_offsets[np.newaxis, :]
|
||||
# Flatten and apply skip_count / truncation
|
||||
flat = all_ptrs.ravel()
|
||||
output[:] = flat[skip_count : skip_count + num_sub_blocks]
|
||||
|
||||
|
||||
def pin_mmap_region(region: SharedOffloadRegion) -> None:
|
||||
"""Register the entire mmap as CUDA pinned memory via cudaHostRegister."""
|
||||
rank = region.rank
|
||||
|
||||
base_ptr = region._base.data_ptr()
|
||||
result = torch.cuda.cudart().cudaHostRegister(base_ptr, region.total_size_bytes, 0)
|
||||
if result.value != 0:
|
||||
logger.warning(
|
||||
"cudaHostRegister failed for rank=%d (code=%d) — "
|
||||
"transfers will still work but may be slower (unpinned DMA)",
|
||||
rank,
|
||||
result,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"cudaHostRegister rank=%d %.2f GB",
|
||||
rank,
|
||||
region.total_size_bytes / 1e9,
|
||||
)
|
||||
region.is_pinned = True
|
||||
|
||||
|
||||
class SingleDirectionOffloadingHandler(OffloadingHandler):
|
||||
@@ -149,13 +190,7 @@ class SingleDirectionOffloadingHandler(OffloadingHandler):
|
||||
# list of CUDA events available for re-use
|
||||
self._event_pool: list[torch.Event] = []
|
||||
|
||||
# Pre-compute base pointers and block sizes for batch copies.
|
||||
self._src_base_ptrs = np.array(
|
||||
[t.data_ptr() for t in self.src_tensors], dtype=np.int64
|
||||
)
|
||||
self._dst_base_ptrs = np.array(
|
||||
[t.data_ptr() for t in self.dst_tensors], dtype=np.int64
|
||||
)
|
||||
# Pre-compute block sizes for batch copies.
|
||||
self._block_size_in_bytes_arr = np.array(
|
||||
self.tensor_block_size_in_bytes, dtype=np.int64
|
||||
)
|
||||
@@ -176,17 +211,6 @@ class SingleDirectionOffloadingHandler(OffloadingHandler):
|
||||
|
||||
assert dst_sub_block_count == src_sub_block_count - src_sub_blocks_to_skip
|
||||
|
||||
src_block_ids = np.empty(dst_sub_block_count, dtype=np.int64)
|
||||
dst_block_ids = np.empty(dst_sub_block_count, dtype=np.int64)
|
||||
expand_block_ids(
|
||||
src_blocks,
|
||||
self.src_block_size_factor,
|
||||
src_block_ids,
|
||||
skip_count=src_sub_blocks_to_skip,
|
||||
)
|
||||
expand_block_ids(dst_blocks, self.dst_block_size_factor, dst_block_ids)
|
||||
|
||||
# Build flat pointer arrays for all tensors × all block pairs.
|
||||
num_pairs = dst_sub_block_count
|
||||
num_tensors = len(self.src_tensors)
|
||||
total = num_pairs * num_tensors
|
||||
@@ -198,8 +222,19 @@ class SingleDirectionOffloadingHandler(OffloadingHandler):
|
||||
for t_idx, bsz in enumerate(self._block_size_in_bytes_arr):
|
||||
start = t_idx * num_pairs
|
||||
end = start + num_pairs
|
||||
all_src[start:end] = self._src_base_ptrs[t_idx] + src_block_ids * bsz
|
||||
all_dst[start:end] = self._dst_base_ptrs[t_idx] + dst_block_ids * bsz
|
||||
compute_sub_block_ptrs(
|
||||
block_ids=src_blocks,
|
||||
block_size_factor=self.src_block_size_factor,
|
||||
output=all_src[start:end],
|
||||
tensor=self.src_tensors[t_idx],
|
||||
skip_count=src_sub_blocks_to_skip,
|
||||
)
|
||||
compute_sub_block_ptrs(
|
||||
block_ids=dst_blocks,
|
||||
block_size_factor=self.dst_block_size_factor,
|
||||
output=all_dst[start:end],
|
||||
tensor=self.dst_tensors[t_idx],
|
||||
)
|
||||
all_sizes[start:end] = bsz
|
||||
|
||||
batch_src = torch.from_numpy(all_src)
|
||||
@@ -281,6 +316,8 @@ class SingleDirectionOffloadingHandler(OffloadingHandler):
|
||||
self._transfer_events.clear()
|
||||
self._stream_pool.clear()
|
||||
self._event_pool.clear()
|
||||
self.src_tensors.clear()
|
||||
self.dst_tensors.clear()
|
||||
|
||||
|
||||
class CpuGpuOffloadingHandlers:
|
||||
@@ -289,9 +326,14 @@ class CpuGpuOffloadingHandlers:
|
||||
kv_caches: CanonicalKVCaches,
|
||||
block_size_factor: int,
|
||||
num_cpu_blocks: int,
|
||||
mmap_region: SharedOffloadRegion | None = None,
|
||||
):
|
||||
pin_memory = is_pin_memory_available()
|
||||
logger.info("Allocating %d CPU tensors...", len(kv_caches.tensors))
|
||||
self._mmap_region = mmap_region
|
||||
if mmap_region is not None and pin_memory:
|
||||
pin_mmap_region(mmap_region)
|
||||
|
||||
gpu_tensors: list[torch.Tensor] = []
|
||||
cpu_tensors: list[torch.Tensor] = []
|
||||
for kv_cache_tensor in kv_caches.tensors:
|
||||
@@ -300,12 +342,24 @@ class CpuGpuOffloadingHandlers:
|
||||
(-1, gpu_page_size_bytes)
|
||||
)
|
||||
cpu_page_size_bytes = gpu_page_size_bytes * block_size_factor
|
||||
cpu_tensor = torch.zeros(
|
||||
(num_cpu_blocks, cpu_page_size_bytes),
|
||||
dtype=torch.int8,
|
||||
device="cpu",
|
||||
pin_memory=pin_memory,
|
||||
)
|
||||
|
||||
if mmap_region is not None:
|
||||
cpu_tensor = mmap_region.create_next_view(cpu_page_size_bytes)
|
||||
else:
|
||||
t0 = time.monotonic()
|
||||
cpu_tensor = torch.zeros(
|
||||
(num_cpu_blocks, cpu_page_size_bytes),
|
||||
dtype=torch.int8,
|
||||
device="cpu",
|
||||
pin_memory=pin_memory,
|
||||
)
|
||||
logger.debug(
|
||||
"torch.zeros pinned tensor %d×%d (%.2f GB): %.3f s",
|
||||
num_cpu_blocks,
|
||||
cpu_page_size_bytes,
|
||||
num_cpu_blocks * cpu_page_size_bytes / 1e9,
|
||||
time.monotonic() - t0,
|
||||
)
|
||||
|
||||
gpu_tensors.append(gpu_tensor)
|
||||
cpu_tensors.append(cpu_tensor)
|
||||
|
||||
Reference in New Issue
Block a user