Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,17 @@ endif()
# cumem_allocator extension
#

if(VLLM_GPU_LANG STREQUAL "CUDA")
define_extension_target(
_ple_memops
DESTINATION vllm
LANGUAGE CXX
SOURCES "csrc/ple_memops.cpp"
LIBRARIES CUDA::cuda_driver
USE_SABI 3.8
WITH_SOABI)
endif()

set(VLLM_CUMEM_EXT_SRC
"csrc/cumem_allocator.cpp")

Expand Down
67 changes: 67 additions & 0 deletions csrc/ple_memops.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
// Adapted flag protocol from FreeToken's ple_store_ext.cpp at
// af71ba43206e124f5ff6419b47ee36c6e9981078 (Apache-2.0).
// This implementation uses the CUDA driver directly and translates pinned
// host addresses rather than assuming identical host/device virtual addresses.
#include <Python.h>
#include <cuda.h>
#include <cstdint>

static CUresult device_flag(unsigned long long host, CUdeviceptr* device) {
return cuMemHostGetDevicePointer(device, reinterpret_cast<void*>(host), 0);
}

static PyObject* signal_flag(PyObject*, PyObject* args) {
unsigned long long flag;
if (!PyArg_ParseTuple(args, "K", &flag)) return nullptr;
__atomic_store_n(reinterpret_cast<uint64_t*>(flag), uint64_t{1},
__ATOMIC_RELEASE);
Py_RETURN_NONE;
}

static PyObject* memop_write(PyObject*, PyObject* args) {
unsigned long long stream, host, value;
if (!PyArg_ParseTuple(args, "KKK", &stream, &host, &value)) return nullptr;
CUdeviceptr flag;
CUresult status = device_flag(host, &flag);
if (status == CUDA_SUCCESS)
status = cuStreamWriteValue64(reinterpret_cast<CUstream>(stream), flag,
value, CU_STREAM_WRITE_VALUE_DEFAULT);
return PyLong_FromLong(status);
}

static PyObject* memop_wait_geq(PyObject*, PyObject* args) {
unsigned long long stream, host, value;
if (!PyArg_ParseTuple(args, "KKK", &stream, &host, &value)) return nullptr;
CUdeviceptr flag;
CUresult status = device_flag(host, &flag);
if (status == CUDA_SUCCESS)
status = cuStreamWaitValue64(reinterpret_cast<CUstream>(stream), flag,
value, CU_STREAM_WAIT_VALUE_GEQ);
return PyLong_FromLong(status);
}

static PyObject* memop_wait_reset(PyObject*, PyObject* args) {
unsigned long long stream, host;
if (!PyArg_ParseTuple(args, "KK", &stream, &host)) return nullptr;
CUdeviceptr flag;
CUresult status = device_flag(host, &flag);
auto s = reinterpret_cast<CUstream>(stream);
if (status == CUDA_SUCCESS)
status = cuStreamWaitValue64(s, flag, 1, CU_STREAM_WAIT_VALUE_GEQ);
if (status == CUDA_SUCCESS)
status = cuStreamWriteValue64(s, flag, 0, CU_STREAM_WRITE_VALUE_DEFAULT);
return PyLong_FromLong(status);
}

static PyMethodDef methods[] = {
{"signal_flag", signal_flag, METH_VARARGS, "Release-store the host flag."},
{"memop_write", memop_write, METH_VARARGS, "Queue a 64-bit flag write."},
{"memop_wait_geq", memop_wait_geq, METH_VARARGS, "Queue a flag wait."},
{"memop_wait_reset", memop_wait_reset, METH_VARARGS,
"Queue wait and reset."},
{nullptr, nullptr, 0, nullptr}};
static PyModuleDef module = {PyModuleDef_HEAD_INIT, "_ple_memops", nullptr, -1,
methods};
PyMODINIT_FUNC PyInit__ple_memops() { return PyModule_Create(&module); }
1 change: 1 addition & 0 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -1360,6 +1360,7 @@ def _read_requirements(filename: str) -> list[str]:
ext_modules.append(CMakeExtension(name="vllm._rocm_C"))

if _is_cuda():
ext_modules.append(CMakeExtension(name="vllm._ple_memops"))
ext_modules.append(CMakeExtension(name="vllm.vllm_flash_attn._vllm_fa2_C"))
if USE_PRECOMPILED_EXTENSIONS or (
CUDA_HOME and get_nvcc_cuda_version() >= Version("12.3")
Expand Down
94 changes: 94 additions & 0 deletions tests/models/qwen4_exp/test_ple_deferred_cuda.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Real memop/copy capture smoke; checkpoint integration is a separate gate."""

from types import SimpleNamespace
from unittest.mock import Mock, patch

import numpy as np
import pytest
import torch

from vllm.config import CUDAGraphMode
from vllm.models.qwen4_exp.nvidia import ple_layer
from vllm.models.qwen4_exp.nvidia.ple_wait import (
DeferredRows,
StreamMemopsUnavailable,
)


class _Table:
def __init__(self, source):
self.source = source

def gather(self, ids):
# Production mmap gather consumes a flattened ID array.
assert ids.ndim == 1
return np.stack(
[self.source[int(ids[h]), h].view(torch.uint8).numpy() for h in range(2)]
)


def test_complete_passes_flat_ids_to_cuda_smoke_table():
source = torch.arange(64, dtype=torch.bfloat16).reshape(8, 2, 4)
helper = object.__new__(DeferredRows)
helper.table = _Table(source)
helper.ids = torch.tensor([[5, 3]], dtype=torch.int64)
helper.rows = torch.empty((1, 2, 4), dtype=torch.bfloat16)
helper._poisoned = False
helper._pending = True
helper._readback_event = Mock()
helper._ext = Mock()
helper.flag = torch.zeros(1, dtype=torch.int64)

helper.complete()

expected = torch.stack([source[5, 0], source[3, 1]])[None]
assert torch.equal(helper.rows.view(torch.uint8), expected.view(torch.uint8))
helper._readback_event.synchronize.assert_called_once_with()
helper._ext.signal_flag.assert_called_once_with(helper.flag.data_ptr())
assert not helper.pending


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_runtime_none_capture_replays_fresh_rows_and_resets_flag():
# Distinct raw BF16 rows catch stale data, zeros and ID mixups without
# relying on floating point comparisons or a model's generated text.
source = torch.arange(64, dtype=torch.bfloat16).reshape(8, 2, 4)

destination = torch.empty((1, 2, 4), dtype=torch.bfloat16, device="cuda")
try:
helper = DeferredRows(destination, _Table(source))
except StreamMemopsUnavailable as exc:
pytest.skip(str(exc))
helper.prepare_dummy()
graph = torch.cuda.CUDAGraph()
context = SimpleNamespace(
cudagraph_runtime_mode=CUDAGraphMode.NONE,
no_compile_layers={
"ple": SimpleNamespace(ple_embedding=SimpleNamespace(deferred_rows=helper))
},
)
with (
patch.object(ple_layer, "get_forward_context", return_value=context),
torch.cuda.graph(graph),
):
ple_layer.qwen4_exp_ple_deferred_rows(destination, "ple")

for row_ids in ([1, 2], [5, 3], [0, 7]):
ids = torch.tensor([row_ids], dtype=torch.int64, device="cuda")
helper.prepare(ids)
graph.replay()
# Never synchronize a replay waiting for the host before releasing it.
try:
helper.complete()
except BaseException:
helper.abort()
raise
torch.accelerator.synchronize()
expected = torch.stack([source[row_ids[h], h] for h in range(2)])[None]
assert torch.equal(
destination.cpu().view(torch.uint8), expected.view(torch.uint8)
)
assert int(helper.flag[0]) == 0
assert not helper.pending
161 changes: 161 additions & 0 deletions tests/models/qwen4_exp/test_ple_deferred_state.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""CPU lifecycle contracts for the model state."""

import unittest
from types import SimpleNamespace
from unittest.mock import Mock, patch

import torch

from vllm.config import CUDAGraphMode
from vllm.models.qwen4_exp.nvidia import ple_layer
from vllm.models.qwen4_exp.nvidia.model_state import Qwen4ExpModelState
from vllm.v1.worker.gpu.model_states.mamba_hybrid import MambaHybridModelState


class DeferredStateTests(unittest.TestCase):
def test_consume_records_during_full_capture_runtime_none(self):
for mode, capturing, rows, expected in (
(CUDAGraphMode.NONE, True, 1, True),
(CUDAGraphMode.NONE, False, 1, False),
(CUDAGraphMode.PIECEWISE, True, 1, False),
(CUDAGraphMode.FULL, False, 1, True),
(CUDAGraphMode.NONE, True, 2, False),
):
with self.subTest(mode=mode, capturing=capturing, rows=rows):
helper = Mock()
context = SimpleNamespace(
cudagraph_runtime_mode=mode,
no_compile_layers={
"ple": SimpleNamespace(
ple_embedding=SimpleNamespace(deferred_rows=helper)
)
},
)
output = torch.empty(rows, 2, 3)
with (
patch.object(
ple_layer, "get_forward_context", return_value=context
),
patch.object(
torch.cuda,
"is_current_stream_capturing",
return_value=capturing,
),
):
ple_layer.qwen4_exp_ple_deferred_rows(output, "ple")
if expected:
helper.consume.assert_called_once_with(destination=output)
else:
helper.consume.assert_not_called()

def make_state(self, count=3):
state = object.__new__(Qwen4ExpModelState)
state._mmap_ple_modules = tuple(
SimpleNamespace(deferred_rows=Mock()) for _ in range(count)
)
state._deferred_ple_step = False
state._deferred_ple_poisoned = False
return state

def test_mixed_capability_disables_all_helpers_before_capture(self):
state = self.make_state()
state.device = torch.device("cpu")
state.max_num_tokens = 8
modules = state._mmap_ple_modules
modules[1].deferred_rows = None
for module in modules:
module.mmap_staging_nbytes = Mock(return_value=16)
module.initialize_mmap_staging = Mock()
with patch(
"vllm.models.qwen4_exp.nvidia.model_state.MemorySnapshot",
return_value=SimpleNamespace(free_memory=1024),
):
state._initialize_mmap_staging(modules)
self.assertTrue(all(m.deferred_rows is None for m in modules))
for module in modules:
module.initialize_mmap_staging.assert_called_once_with(8, state.device)
state.set_deferred_ple_step(True)
self.assertFalse(state._deferred_ple_step)

def test_complete_all_layers_and_next_step(self):
state = self.make_state()
state.set_deferred_ple_step(True)
state.complete_deferred_ple()
for module in state._mmap_ple_modules:
module.deferred_rows.complete.assert_called_once()
module.deferred_rows.abort.assert_not_called()
self.assertFalse(state._deferred_ple_step)
state.set_deferred_ple_step(True)
self.assertTrue(state._deferred_ple_step)

def test_fill_failure_releases_unvisited_layers_and_poison_latches(self):
state = self.make_state()
state._mmap_ple_modules[0].deferred_rows.complete.side_effect = ValueError(
"disk"
)
state.set_deferred_ple_step(True)
with self.assertRaisesRegex(ValueError, "disk"):
state.complete_deferred_ple()
for module in state._mmap_ple_modules:
module.deferred_rows.abort.assert_called_once()
state._mmap_ple_modules[1].deferred_rows.complete.assert_not_called()
with self.assertRaisesRegex(RuntimeError, "poisoned"):
state.set_deferred_ple_step(False)

def test_failed_release_still_attempts_every_layer(self):
state = self.make_state()
state.set_deferred_ple_step(True)
state._mmap_ple_modules[0].deferred_rows.abort.side_effect = ValueError(
"release"
)
with self.assertRaisesRegex(ValueError, "release"):
state.abort_deferred_ple()
for module in state._mmap_ple_modules:
module.deferred_rows.abort.assert_called_once()
self.assertTrue(state._deferred_ple_poisoned)

def test_prepare_failure_releases_already_prepared_layers(self):
state = self.make_state()
state.uses_ngram_embedding = True
state.ple_query_start_loc = torch.zeros(2, dtype=torch.int32)
state._prepare_ngram_context = Mock(return_value=torch.zeros((1, 2)))
batch = SimpleNamespace(
num_reqs_after_padding=1,
num_tokens=1,
num_tokens_after_padding=1,
num_reqs=1,
input_ids=torch.tensor([3]),
query_start_loc=torch.tensor([0, 1], dtype=torch.int32),
)
for module in state._mmap_ple_modules:
module.prepare_deferred_mmap_rows = Mock()
state._mmap_ple_modules[1].prepare_deferred_mmap_rows.side_effect = ValueError(
"ids"
)
state.set_deferred_ple_step(True)
with (
patch.object(MambaHybridModelState, "prepare_inputs", return_value={}),
self.assertRaisesRegex(ValueError, "ids"),
):
state.prepare_inputs(batch, None)
state._mmap_ple_modules[0].prepare_deferred_mmap_rows.assert_called_once()
state._mmap_ple_modules[2].prepare_deferred_mmap_rows.assert_not_called()
for module in state._mmap_ple_modules:
module.deferred_rows.abort.assert_called_once()

def test_disabled_and_empty_have_no_effect(self):
for state in (self.make_state(), self.make_state(0)):
state.set_deferred_ple_step(False)
state.complete_deferred_ple()
state.abort_deferred_ple()
self.assertFalse(state._deferred_ple_poisoned)
state = self.make_state()
state._mmap_ple_modules[1].deferred_rows = None
state.set_deferred_ple_step(True)
self.assertFalse(state._deferred_ple_step)


if __name__ == "__main__":
unittest.main()
Loading