From 39dff5ff39a8f45caeb81d345d7a3d70c2723e40 Mon Sep 17 00:00:00 2001 From: pschlan-amd Date: Tue, 12 May 2026 01:11:49 +0200 Subject: [PATCH] Add VLLM_USE_SPINLOOP_EXT to use more efficient busy polling (#36517) Signed-off-by: Patrick Schlangen --- CMakeLists.txt | 18 ++ csrc/spinloop.cpp | 204 ++++++++++++++++++ setup.py | 3 + .../device_communicators/shm_broadcast.py | 45 ++-- vllm/envs.py | 3 + 5 files changed, 259 insertions(+), 14 deletions(-) create mode 100644 csrc/spinloop.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 38af64851ed..fd6c7eeffd0 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -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. # diff --git a/csrc/spinloop.cpp b/csrc/spinloop.cpp new file mode 100644 index 00000000000..c29e48a5f0e --- /dev/null +++ b/csrc/spinloop.cpp @@ -0,0 +1,204 @@ +#include + +extern "C" { + +#include +#include + +#if defined(__i386__) || defined(__x86_64__) + #include + #include +#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" diff --git a/setup.py b/setup.py index 7c226a72425..dc963c9e789 100644 --- a/setup.py +++ b/setup.py @@ -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")) diff --git a/vllm/distributed/device_communicators/shm_broadcast.py b/vllm/distributed/device_communicators/shm_broadcast.py index 9c8bf3ad165..dc7e6d151a4 100644 --- a/vllm/distributed/device_communicators/shm_broadcast.py +++ b/vllm/distributed/device_communicators/shm_broadcast.py @@ -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 diff --git a/vllm/envs.py b/vllm/envs.py index 66e2a33bc1b..03230eed068 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -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"))), }