mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-11 16:28:16 +00:00
Add VLLM_USE_SPINLOOP_EXT to use more efficient busy polling (#36517)
Signed-off-by: Patrick Schlangen <[email protected]>
This commit is contained in:
@@ -109,6 +109,24 @@ else()
|
||||
set(CUDA_SUPPORTED_ARCHS "7.0;7.5;8.0;8.6;8.7;8.9;9.0")
|
||||
endif()
|
||||
|
||||
#
|
||||
# spinloop extension (pure CXX; must stay above the non-CUDA device branch so
|
||||
# CPU builds define the target before the early return)
|
||||
#
|
||||
set(VLLM_SPINLOOP_EXT_SRC "csrc/spinloop.cpp")
|
||||
set(SPINLOOP_COMPILE_FLAGS "")
|
||||
if(CMAKE_SYSTEM_PROCESSOR MATCHES "x86_64|amd64")
|
||||
list(APPEND SPINLOOP_COMPILE_FLAGS "-mmwaitx")
|
||||
endif()
|
||||
define_extension_target(
|
||||
spinloop
|
||||
DESTINATION vllm
|
||||
LANGUAGE CXX
|
||||
SOURCES ${VLLM_SPINLOOP_EXT_SRC}
|
||||
COMPILE_FLAGS ${SPINLOOP_COMPILE_FLAGS}
|
||||
USE_SABI 3.11
|
||||
WITH_SOABI)
|
||||
|
||||
#
|
||||
# Forward the non-CUDA device extensions to external CMake scripts.
|
||||
#
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
#include <Python.h>
|
||||
|
||||
extern "C" {
|
||||
|
||||
#include <stdbool.h>
|
||||
#include <time.h>
|
||||
|
||||
#if defined(__i386__) || defined(__x86_64__)
|
||||
#include <cpuid.h>
|
||||
#include <mwaitxintrin.h>
|
||||
#endif
|
||||
|
||||
#if defined(CLOCK_MONOTONIC_RAW)
|
||||
#define TIMEOUT_CLOCK CLOCK_MONOTONIC_RAW
|
||||
#else
|
||||
#define TIMEOUT_CLOCK CLOCK_MONOTONIC
|
||||
#endif
|
||||
|
||||
#define CPU_SUPPORT_NONE 0
|
||||
#define CPU_SUPPORT_MONITORX 1
|
||||
|
||||
#define MWAITX_DEFAULT_TIMEOUT_CYCLES 1000000
|
||||
|
||||
typedef struct {
|
||||
unsigned int cpu_support;
|
||||
unsigned int max_monitor_line_size;
|
||||
} spinloop_state_t;
|
||||
|
||||
static void determine_cpu_support(spinloop_state_t* state) {
|
||||
state->cpu_support = CPU_SUPPORT_NONE;
|
||||
state->max_monitor_line_size = 0;
|
||||
|
||||
#if defined(__i386__) || defined(__x86_64__)
|
||||
unsigned int eax, ebx, ecx, edx;
|
||||
if (__get_cpuid(0, &eax, &ebx, &ecx, &edx) == 1) {
|
||||
// AMD CPU (possible monitorx/mwaitx support)
|
||||
if (ebx == 0x68747541 && edx == 0x69746e65 && ecx == 0x444d4163) {
|
||||
if (__get_cpuid(0x80000000, &eax, &ebx, &ecx, &edx) == 1 &&
|
||||
eax >= 0x80000001 &&
|
||||
__get_cpuid(0x80000001, &eax, &ebx, &ecx, &edx) == 1) {
|
||||
if ((ecx & (1 << 29)) != 0) {
|
||||
state->cpu_support = CPU_SUPPORT_MONITORX;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (state->cpu_support == CPU_SUPPORT_MONITORX) {
|
||||
if (__get_cpuid(5, &eax, &ebx, &ecx, &edx) == 1) {
|
||||
state->max_monitor_line_size = ebx & 0xff;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
static PyObject* method_spinloop(PyObject* self, PyObject* args,
|
||||
PyObject* kwargs) {
|
||||
Py_buffer buffer;
|
||||
PyObject* callback;
|
||||
double timeout = 0.;
|
||||
|
||||
spinloop_state_t* state = (spinloop_state_t*)PyModule_GetState(self);
|
||||
if (state == NULL) {
|
||||
PyErr_SetString(PyExc_TypeError, "Failed to retrieve module state!");
|
||||
return NULL;
|
||||
}
|
||||
|
||||
static const char* keywords[] = {"buffer", "callback", "timeout", NULL};
|
||||
if (!PyArg_ParseTupleAndKeywords(args, kwargs, "y*O|d", (char**)keywords,
|
||||
&buffer, &callback, &timeout)) {
|
||||
return NULL;
|
||||
}
|
||||
|
||||
if (!PyCallable_Check(callback)) {
|
||||
PyErr_SetString(PyExc_TypeError, "callback parameter must be callable!");
|
||||
PyBuffer_Release(&buffer);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
struct timespec t_start;
|
||||
if (clock_gettime(TIMEOUT_CLOCK, &t_start) != 0) {
|
||||
PyErr_SetString(PyExc_RuntimeError, "clock_gettime() failed!");
|
||||
PyBuffer_Release(&buffer);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
bool result = false;
|
||||
bool error = false;
|
||||
bool have_timeout = (timeout > 1e-9);
|
||||
unsigned int iteration = 0;
|
||||
const bool buffer_qualifies = (buffer.len <= state->max_monitor_line_size);
|
||||
|
||||
while (true) {
|
||||
PyObject* res = PyObject_CallNoArgs(callback);
|
||||
if (res == NULL) {
|
||||
error = true;
|
||||
break;
|
||||
}
|
||||
int ok = (res == Py_True);
|
||||
Py_DECREF(res);
|
||||
|
||||
if (ok) {
|
||||
result = true;
|
||||
break;
|
||||
}
|
||||
|
||||
// Check timeout at most every 16 iterations to avoid clock_gettime and
|
||||
// comparison cost
|
||||
if (have_timeout && (iteration & 15u) == 0) {
|
||||
struct timespec t_now;
|
||||
if (clock_gettime(TIMEOUT_CLOCK, &t_now) != 0) {
|
||||
PyErr_SetString(PyExc_RuntimeError, "clock_gettime() failed!");
|
||||
error = true;
|
||||
break;
|
||||
}
|
||||
|
||||
const double elapsed = (double)(t_now.tv_sec - t_start.tv_sec) +
|
||||
(t_now.tv_nsec - t_start.tv_nsec) * 1e-9;
|
||||
if (elapsed >= timeout) {
|
||||
result = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
++iteration;
|
||||
|
||||
#if defined(__i386__) || defined(__x86_64__)
|
||||
// monitorx + mwaitx with qualified buffer
|
||||
if (buffer_qualifies && state->cpu_support == CPU_SUPPORT_MONITORX) {
|
||||
_mm_monitorx(buffer.buf, 0, 0);
|
||||
|
||||
// Check once more in case the buffer has been modified while we were
|
||||
// arming the monitor hardware
|
||||
res = PyObject_CallNoArgs(callback);
|
||||
if (res == NULL) {
|
||||
error = true;
|
||||
break;
|
||||
}
|
||||
ok = (res == Py_True);
|
||||
Py_DECREF(res);
|
||||
|
||||
if (ok) {
|
||||
result = true;
|
||||
break;
|
||||
}
|
||||
|
||||
// Run mwaitx with enabled timeout (bit 1). The actual timeout value
|
||||
// is not very important, we just want to ensure we don't lock up
|
||||
// here for too long.
|
||||
Py_BEGIN_ALLOW_THREADS _mm_mwaitx((1 << 1), 0,
|
||||
MWAITX_DEFAULT_TIMEOUT_CYCLES);
|
||||
Py_END_ALLOW_THREADS
|
||||
}
|
||||
|
||||
// Fallback: Busy poll
|
||||
else {
|
||||
#endif
|
||||
// Give other threads a chance to be scheduled
|
||||
Py_BEGIN_ALLOW_THREADS
|
||||
#if defined(__i386__) || defined(__x86_64__)
|
||||
__builtin_ia32_pause();
|
||||
#elif defined(__aarch64__)
|
||||
__asm__ volatile("yield" :: : "memory");
|
||||
#endif
|
||||
Py_END_ALLOW_THREADS
|
||||
#if defined(__i386__) || defined(__x86_64__)
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
PyBuffer_Release(&buffer);
|
||||
|
||||
if (error) {
|
||||
return NULL;
|
||||
}
|
||||
|
||||
if (result) {
|
||||
Py_RETURN_TRUE;
|
||||
}
|
||||
|
||||
Py_RETURN_FALSE;
|
||||
}
|
||||
|
||||
static PyMethodDef spinloop_methods[] = {
|
||||
{"spinloop", (PyCFunction)method_spinloop, METH_VARARGS | METH_KEYWORDS,
|
||||
"Wait for store with callback"},
|
||||
{NULL, NULL, 0, NULL}};
|
||||
|
||||
static struct PyModuleDef spinloop_module = {
|
||||
PyModuleDef_HEAD_INIT, "spinloop",
|
||||
"Hardware-optimized spinloops for Python", sizeof(spinloop_state_t),
|
||||
spinloop_methods};
|
||||
|
||||
PyMODINIT_FUNC PyInit_spinloop(void) {
|
||||
PyObject* m = PyModule_Create(&spinloop_module);
|
||||
if (m != NULL) {
|
||||
spinloop_state_t* state = (spinloop_state_t*)PyModule_GetState(m);
|
||||
if (state != NULL) {
|
||||
determine_cpu_support(state);
|
||||
}
|
||||
}
|
||||
return m;
|
||||
}
|
||||
|
||||
} // extern "C"
|
||||
@@ -686,6 +686,7 @@ class precompiled_wheel_utils:
|
||||
"vllm/vllm_flash_attn/_vllm_fa2_C.abi3.so",
|
||||
"vllm/vllm_flash_attn/_vllm_fa3_C.abi3.so",
|
||||
"vllm/cumem_allocator.abi3.so",
|
||||
"vllm/spinloop.abi3.so",
|
||||
# ROCm-specific libraries
|
||||
"vllm/_rocm_C.abi3.so",
|
||||
]
|
||||
@@ -993,6 +994,8 @@ if _is_cuda() or _is_hip():
|
||||
# copying the relevant .py files from the source repository.
|
||||
ext_modules.append(CMakeExtension(name="vllm.triton_kernels", optional=True))
|
||||
|
||||
ext_modules.append(CMakeExtension(name="vllm.spinloop"))
|
||||
|
||||
if _is_hip():
|
||||
ext_modules.append(CMakeExtension(name="vllm._rocm_C"))
|
||||
|
||||
|
||||
@@ -38,6 +38,11 @@ from vllm.utils.network_utils import (
|
||||
is_valid_ipv6_address,
|
||||
)
|
||||
|
||||
if envs.VLLM_USE_SPINLOOP_EXT:
|
||||
from vllm.spinloop import spinloop
|
||||
|
||||
SPINLOOP_TIMEOUT_SECONDS = 0.1
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from _typeshed import SizedBuffer
|
||||
|
||||
@@ -540,13 +545,17 @@ class MessageQueue:
|
||||
n_warning = 1
|
||||
while True:
|
||||
with self.buffer.get_metadata(self.current_idx) as metadata_buffer:
|
||||
# Memory fence ensures we see the latest read flags from readers.
|
||||
# Without this, we may read stale flags from our CPU cache and
|
||||
# spin indefinitely even though readers have completed.
|
||||
memory_fence()
|
||||
read_count = sum(metadata_buffer[1:])
|
||||
written_flag = metadata_buffer[0]
|
||||
if written_flag and read_count != self.buffer.n_reader:
|
||||
|
||||
def check():
|
||||
memory_fence()
|
||||
read_count = sum(metadata_buffer[1:])
|
||||
written_flag = metadata_buffer[0]
|
||||
return not (written_flag and read_count != self.buffer.n_reader)
|
||||
|
||||
if envs.VLLM_USE_SPINLOOP_EXT and not check():
|
||||
spinloop(metadata_buffer, check, timeout=SPINLOOP_TIMEOUT_SECONDS)
|
||||
|
||||
if not check():
|
||||
# this block is written and not read by all readers
|
||||
# for writers, `self.current_idx` is the next block to write
|
||||
# if this block is not ready to write,
|
||||
@@ -657,13 +666,21 @@ class MessageQueue:
|
||||
)
|
||||
with self.buffer.get_metadata(self.current_idx) as metadata_buffer:
|
||||
while True:
|
||||
# Memory fence ensures we see the latest writes from the writer.
|
||||
# Without this, we may read stale flags from our CPU cache
|
||||
# and spin indefinitely even though writer has updated them.
|
||||
memory_fence()
|
||||
read_flag = metadata_buffer[self.local_reader_rank + 1]
|
||||
written_flag = metadata_buffer[0]
|
||||
if not written_flag or read_flag:
|
||||
|
||||
def check():
|
||||
memory_fence()
|
||||
read_flag = metadata_buffer[self.local_reader_rank + 1]
|
||||
written_flag = metadata_buffer[0]
|
||||
return not (not written_flag or read_flag)
|
||||
|
||||
if envs.VLLM_USE_SPINLOOP_EXT and not check():
|
||||
spinloop(
|
||||
metadata_buffer[0 : self.local_reader_rank + 1],
|
||||
check,
|
||||
timeout=SPINLOOP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if not check():
|
||||
# this block is either
|
||||
# (1) not written
|
||||
# (2) already read by this reader
|
||||
|
||||
@@ -1780,6 +1780,9 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"VLLM_LORA_ENABLE_DUAL_STREAM": lambda: bool(
|
||||
int(os.getenv("VLLM_LORA_ENABLE_DUAL_STREAM", "0"))
|
||||
),
|
||||
# If set to 1, use Python spinloop extension to poll in a more efficient
|
||||
# way when using the mp backend.
|
||||
"VLLM_USE_SPINLOOP_EXT": lambda: bool(int(os.getenv("VLLM_USE_SPINLOOP_EXT", "0"))),
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user