From be0c855ebd6ed4a8bface6c7df6d99150fdb5d90 Mon Sep 17 00:00:00 2001 From: omerpaz95 <73347585+omerpaz95@users.noreply.github.com> Date: Tue, 14 Apr 2026 21:33:33 +0300 Subject: [PATCH] [KV Offload] Unified memory layout for offloading workers (#37206) Signed-off-by: omerpaz95 Co-authored-by: Or Ozeri --- tests/v1/kv_offload/test_cpu_gpu.py | 39 +- .../kv_offload/test_shared_offload_region.py | 625 ++++++++++++++++++ .../kv_offload/cpu/shared_offload_region.py | 192 ++++++ vllm/v1/kv_offload/worker/cpu_gpu.py | 146 ++-- 4 files changed, 949 insertions(+), 53 deletions(-) create mode 100644 tests/v1/kv_offload/test_shared_offload_region.py create mode 100644 vllm/v1/kv_offload/cpu/shared_offload_region.py diff --git a/tests/v1/kv_offload/test_cpu_gpu.py b/tests/v1/kv_offload/test_cpu_gpu.py index 2da3a5e56b1..de482aec4a4 100644 --- a/tests/v1/kv_offload/test_cpu_gpu.py +++ b/tests/v1/kv_offload/test_cpu_gpu.py @@ -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() diff --git a/tests/v1/kv_offload/test_shared_offload_region.py b/tests/v1/kv_offload/test_shared_offload_region.py new file mode 100644 index 00000000000..b33a27ca645 --- /dev/null +++ b/tests/v1/kv_offload/test_shared_offload_region.py @@ -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) diff --git a/vllm/v1/kv_offload/cpu/shared_offload_region.py b/vllm/v1/kv_offload/cpu/shared_offload_region.py new file mode 100644 index 00000000000..b2e21d06c9b --- /dev/null +++ b/vllm/v1/kv_offload/cpu/shared_offload_region.py @@ -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 diff --git a/vllm/v1/kv_offload/worker/cpu_gpu.py b/vllm/v1/kv_offload/worker/cpu_gpu.py index e14c85549b2..dd12a533ede 100644 --- a/vllm/v1/kv_offload/worker/cpu_gpu.py +++ b/vllm/v1/kv_offload/worker/cpu_gpu.py @@ -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)