diff --git a/docker/xeon.Dockerfile b/docker/xeon.Dockerfile index 98e443a1f..c29e2ad45 100644 --- a/docker/xeon.Dockerfile +++ b/docker/xeon.Dockerfile @@ -42,7 +42,8 @@ RUN source $HOME/.local/bin/env && \ uv pip install . && \ cd ../sgl-kernel && \ cp pyproject_cpu.toml pyproject.toml && \ - uv pip install . + uv pip install . && \ + uv pip install pytest ENV SGLANG_USE_CPU_ENGINE=1 ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc.so.4:/usr/lib/x86_64-linux-gnu/libtbbmalloc.so:/opt/.venv/lib/libiomp5.so diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 78127c05f..0d520a5ba 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -111,6 +111,16 @@ def _set_kv_buffer_impl( row_bytes=row_bytes, ) + if _is_cpu and _cpu_has_amx_support: + return torch.ops.sgl_kernel.store_cache_cpu( + k, + v, + k_cache, + v_cache, + indices, + row_dim, + ) + from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode if get_is_capture_mode() and alt_stream is not None: diff --git a/sgl-kernel/csrc/cpu/common.h b/sgl-kernel/csrc/cpu/common.h index 139121859..0da10f0bc 100644 --- a/sgl-kernel/csrc/cpu/common.h +++ b/sgl-kernel/csrc/cpu/common.h @@ -1,6 +1,7 @@ #pragma once #include +#include #include #if defined(_OPENMP) @@ -44,6 +45,14 @@ namespace { } \ }() +// Half + BFloat16, plus one extra scalar type +#define AT_DISPATCH_CASE_REDUCED_FLOATING_TYPES_AND(SCALARTYPE, ...) \ + AT_DISPATCH_CASE_REDUCED_FLOATING_TYPES(__VA_ARGS__) \ + AT_DISPATCH_CASE(SCALARTYPE, __VA_ARGS__) + +#define AT_DISPATCH_REDUCED_FLOATING_TYPES_AND(SCALARTYPE, TYPE, NAME, ...) \ + AT_DISPATCH_SWITCH(TYPE, NAME, AT_DISPATCH_CASE_REDUCED_FLOATING_TYPES_AND(SCALARTYPE, __VA_ARGS__)) + // dispatch: bfloat16, float16, int8_t, fp8_e4m3, uint8_t(mxfp4/int4) #define CPU_DISPATCH_PACKED_TYPES(TYPE, ...) \ [&] { \ diff --git a/sgl-kernel/csrc/cpu/kvcache.cpp b/sgl-kernel/csrc/cpu/kvcache.cpp new file mode 100644 index 000000000..ca2dc81b8 --- /dev/null +++ b/sgl-kernel/csrc/cpu/kvcache.cpp @@ -0,0 +1,130 @@ +#include "common.h" +#include "vec.h" + +namespace { + +template +inline void copy_stub(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int size) { + int d = 0; +#if defined(CPU_CAPABILITY_AVX512) + using Vec = at::vec::Vectorized; + constexpr int kVecSize = Vec::size(); + + for (; d <= size - kVecSize; d += kVecSize) { + Vec data = Vec::loadu(src + d); + data.store(dst + d); + } +#endif + for (; d < size; ++d) { + dst[d] = src[d]; + } +} + +template +void store_cache_kernel_impl( + const scalar_t* __restrict__ k, + const scalar_t* __restrict__ v, + scalar_t* __restrict__ k_cache, + scalar_t* __restrict__ v_cache, + const index_t* __restrict__ indices, + int64_t batch_size, + int64_t num_pages, + int64_t row_dim, + int64_t k_stride, + int64_t v_stride, + int64_t kc_stride, + int64_t vc_stride) { + at::parallel_for(0, batch_size, 0, [&](int64_t begin, int64_t end) { + for (int64_t bs = begin; bs < end; ++bs) { + const int64_t idx = static_cast(indices[bs]); + const scalar_t* k_ptr = k + bs * k_stride; + const scalar_t* v_ptr = v + bs * v_stride; + scalar_t* kc_ptr = k_cache + idx * kc_stride; + scalar_t* vc_ptr = v_cache + idx * vc_stride; + copy_stub(kc_ptr, k_ptr, row_dim); + copy_stub(vc_ptr, v_ptr, row_dim); + } + }); +} + +} // anonymous namespace + +// check tensor last two dimensions are contiguous +#define CHECK_LAST2_DIM_CONTIGUOUS(x, ndim) \ + do { \ + const auto& _x = (x); \ + const auto _ndim = _x.dim(); \ + const auto _strides = _x.strides(); \ + const auto _sizes = _x.sizes(); \ + TORCH_CHECK(_ndim == ndim, #x " must have " #ndim " dimensions"); \ + TORCH_CHECK( \ + _ndim >= 2 && _strides[_ndim - 1] == 1 && _strides[_ndim - 2] == _sizes[_ndim - 1], \ + #x " must be contiguous at the last two dimensions"); \ + } while (0) + +// [NB]: store_cache takes 3 dimension tensors, +// This is to avoid the overhead of creating a new TensorImpl +// from .view(-1, row_dim) +// +// k : [batch_size, num_heads, head_size] -> [batch_size, row_dim] +// v : [batch_size, num_heads, head_size] -> [batch_size, row_dim] +// k_cache : [num_pages, num_heads, head_size] -> [num_pages, row_dim] +// v_cache : [num_pages, num_heads, head_size] -> [num_pages, row_dim] +// indices : [batch_size] +// +void store_cache_cpu( + const at::Tensor& k, + const at::Tensor& v, + const at::Tensor& k_cache, + const at::Tensor& v_cache, + const at::Tensor& indices, + std::optional row_dim) { + CHECK_LAST2_DIM_CONTIGUOUS(k, 3); + CHECK_LAST2_DIM_CONTIGUOUS(v, 3); + CHECK_LAST2_DIM_CONTIGUOUS(k_cache, 3); + CHECK_LAST2_DIM_CONTIGUOUS(v_cache, 3); + CHECK_INPUT(indices); + + int64_t batch_size = k.size(0); + int64_t num_heads = k.size(1); + int64_t head_size = k.size(2); + int64_t num_pages = k_cache.size(0); + int64_t row_dim_value = num_heads * head_size; + if (row_dim.has_value()) { + CHECK_EQ(row_dim.value(), row_dim_value); + } + CHECK_EQ(indices.size(0), batch_size); + + // strides: batch dimension (dim 0) stride in elements + int64_t k_stride = k.stride(0); + int64_t v_stride = v.stride(0); + int64_t kc_stride = k_cache.stride(0); + int64_t vc_stride = v_cache.stride(0); + + const auto dtype = k.scalar_type(); + TORCH_CHECK( + dtype == v.scalar_type() && dtype == k_cache.scalar_type() && dtype == v_cache.scalar_type(), + "store_cache_cpu: input tensors must have the same dtype"); + const auto index_dtype = indices.scalar_type(); + TORCH_CHECK(index_dtype == at::kLong || index_dtype == at::kInt, "indices must be int64 or int32"); + + // dtype : [bfloat16, float16, uint8] for fp8 KV stored as uint8 + // index_dtype : [int64, int32] + AT_DISPATCH_REDUCED_FLOATING_TYPES_AND(at::ScalarType::Byte, dtype, "store_cache_cpu", [&] { + AT_DISPATCH_INDEX_TYPES(index_dtype, "store_cache_cpu_index", [&] { + store_cache_kernel_impl( + k.data_ptr(), + v.data_ptr(), + k_cache.data_ptr(), + v_cache.data_ptr(), + indices.data_ptr(), + batch_size, + num_pages, + row_dim_value, + k_stride, + v_stride, + kc_stride, + vc_stride); + }); + }); +} diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index e86d249da..5ed291986 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -410,6 +410,15 @@ std::tuple image_preprocess_cpu( bool disable_grouping, at::ScalarType out_dtype); +// kvcache +void store_cache_cpu( + const at::Tensor& k, + const at::Tensor& v, + const at::Tensor& k_cache, + const at::Tensor& v_cache, + const at::Tensor& indices, + std::optional row_dim); + // [NOTE] When registering kernels, we should accurately describe the in-place information. // Taking fused_add_rmsnorm_cpu as an example, add `Tensor(a!)` modifier to all tensors that // will be modified in-place to avoid incorrect fusing and execution order on graph mode. @@ -658,6 +667,12 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "image_std, int patch_size, int temporal_patch_size, int merge_size, bool disable_grouping, ScalarType " "out_dtype) -> (Tensor, Tensor)"); m.impl("image_preprocess_cpu", torch::kCPU, &image_preprocess_cpu); + + // kvcache + m.def( + "store_cache_cpu(Tensor k, Tensor v, Tensor(a!) k_cache, Tensor(a!) v_cache, Tensor indices, int? row_dim) -> " + "()"); + m.impl("store_cache_cpu", torch::kCPU, &store_cache_cpu); } TORCH_LIBRARY_IMPL(sgl_kernel, CatchAll, m) { diff --git a/test/registered/cpu/test_store_cache.py b/test/registered/cpu/test_store_cache.py new file mode 100644 index 000000000..80bd94e06 --- /dev/null +++ b/test/registered/cpu/test_store_cache.py @@ -0,0 +1,83 @@ +import sys + +import pytest +import torch + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=25, suite="base-b-test-cpu") + +torch.manual_seed(42) + +DEVICE = "cpu" +CACHE_SIZE = 4096 + +# for fp8 KV stored as uint8, e.g. float8_e4m3fn and float8_e5m2 +DTYPES = [torch.float16, torch.bfloat16, torch.uint8] +DTYPE_IDS = ["float16", "bfloat16", "uint8"] + + +def _store_cache_cpu(k, v, k_cache, v_cache, indices): + row_dim = k.size(1) * k.size(2) + torch.ops.sgl_kernel.store_cache_cpu(k, v, k_cache, v_cache, indices, row_dim) + + +def _random_tensor(shape, dtype): + """FP8 KV is stored as uint8; randn is not implemented for Byte.""" + if dtype == torch.uint8: + return torch.randint(0, 256, shape, dtype=torch.uint8, device=DEVICE) + return torch.randn(shape, dtype=dtype, device=DEVICE) + + +@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS) +@pytest.mark.parametrize("head_dim", [64, 128]) +@pytest.mark.parametrize("num_heads", [1, 8, 16, 32]) +@pytest.mark.parametrize("batch_size", [1, 7, 133]) +def test_store_cache(batch_size, num_heads, head_dim, dtype): + shape = (batch_size, num_heads, head_dim) + cache_shape = (CACHE_SIZE, num_heads, head_dim) + k = _random_tensor(shape, dtype) + v = _random_tensor(shape, dtype) + k_cache = _random_tensor(cache_shape, dtype) + v_cache = _random_tensor(cache_shape, dtype) + indices = torch.randperm(CACHE_SIZE, device=DEVICE, dtype=torch.int64)[:batch_size] + + k_cache_ref = k_cache.clone() + v_cache_ref = v_cache.clone() + k_cache_ref[indices] = k + v_cache_ref[indices] = v + + _store_cache_cpu(k, v, k_cache, v_cache, indices) + + assert torch.equal(k_cache, k_cache_ref) + assert torch.equal(v_cache, v_cache_ref) + + +@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS) +@pytest.mark.parametrize("head_dim", [64, 128]) +@pytest.mark.parametrize("num_heads", [1, 8]) +@pytest.mark.parametrize("batch_size", [11]) +def test_store_cache_int32_indices(batch_size, num_heads, head_dim, dtype): + shape = (batch_size, num_heads, head_dim) + cache_shape = (CACHE_SIZE, num_heads, head_dim) + k = _random_tensor(shape, dtype) + v = _random_tensor(shape, dtype) + k_cache = _random_tensor(cache_shape, dtype) + v_cache = _random_tensor(cache_shape, dtype) + indices = torch.randperm(CACHE_SIZE, device=DEVICE, dtype=torch.int64)[ + :batch_size + ].to(torch.int32) + + k_cache_ref = k_cache.clone() + v_cache_ref = v_cache.clone() + k_cache_ref[indices.long()] = k + v_cache_ref[indices.long()] = v + + _store_cache_cpu(k, v, k_cache, v_cache, indices) + + assert torch.equal(k_cache, k_cache_ref) + assert torch.equal(v_cache, v_cache_ref) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__]))