diff --git a/sgl-kernel/CMakeLists.txt b/sgl-kernel/CMakeLists.txt index 40ec8d11b..079eb2f45 100644 --- a/sgl-kernel/CMakeLists.txt +++ b/sgl-kernel/CMakeLists.txt @@ -278,6 +278,7 @@ set(SOURCES "csrc/elementwise/copy.cu" "csrc/elementwise/fused_add_rms_norm_kernel.cu" "csrc/elementwise/rope.cu" + "csrc/elementwise/pos_enc.cu" "csrc/elementwise/topk.cu" "csrc/expert_specialization/es_fp8_blockwise.cu" "csrc/expert_specialization/es_sm100_mxfp8_blockscaled.cu" diff --git a/sgl-kernel/csrc/common_extension.cc b/sgl-kernel/csrc/common_extension.cc index eaa2a2623..0145676e6 100644 --- a/sgl-kernel/csrc/common_extension.cc +++ b/sgl-kernel/csrc/common_extension.cc @@ -90,6 +90,12 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "Tensor? v, Tensor!? k_buffer, Tensor!? v_buffer, Tensor? kv_cache_loc) -> ()"); m.impl("apply_rope_pos_ids_cos_sin_cache", torch::kCUDA, &apply_rope_pos_ids_cos_sin_cache); + m.def( + "rotary_embedding(Tensor positions, Tensor! query," + " Tensor!? key, int head_size," + " Tensor cos_sin_cache, bool is_neox) -> ()"); + m.impl("rotary_embedding", torch::kCUDA, &rotary_embedding); + m.def( "downcast_fp8(Tensor k, Tensor v, Tensor k_out, Tensor v_out, Tensor k_scale, Tensor v_scale, Tensor loc, " "int mult, int offset) -> ()"); diff --git a/sgl-kernel/csrc/elementwise/pos_enc.cu b/sgl-kernel/csrc/elementwise/pos_enc.cu new file mode 100644 index 000000000..9b8bbd1d5 --- /dev/null +++ b/sgl-kernel/csrc/elementwise/pos_enc.cu @@ -0,0 +1,208 @@ +// Adapted from +// https://github.com/vllm-project/vllm/blob/014ece97c7aa49084a1119dca792af081a18dbc1/csrc/pos_encoding_kernels.cu + +#include +#include +#include + +#include "utils.h" + +template +inline __device__ void apply_token_rotary_embedding( + scalar_t* __restrict__ arr, + const scalar_t* __restrict__ cos_ptr, + const scalar_t* __restrict__ sin_ptr, + int rot_offset, + int embed_dim) { + int x_index, y_index; + scalar_t cos, sin; + if (IS_NEOX) { + // GPT-NeoX style rotary embedding. + x_index = rot_offset; + y_index = embed_dim + rot_offset; + cos = SGLANG_LDG(cos_ptr + x_index); + sin = SGLANG_LDG(sin_ptr + x_index); + } else { + // GPT-J style rotary embedding. + x_index = 2 * rot_offset; + y_index = 2 * rot_offset + 1; + cos = SGLANG_LDG(cos_ptr + x_index / 2); + sin = SGLANG_LDG(sin_ptr + x_index / 2); + } + + const scalar_t x = arr[x_index]; + const scalar_t y = arr[y_index]; + arr[x_index] = x * cos - y * sin; + arr[y_index] = y * cos + x * sin; +} + +template +inline __device__ void apply_rotary_embedding( + scalar_t* __restrict__ query, // [batch_size, seq_len, num_heads, + // head_size] or [num_tokens, num_heads, + // head_size] + scalar_t* __restrict__ key, // nullptr or + // [batch_size, seq_len, num_kv_heads, + // head_size] or [num_tokens, num_kv_heads, + // head_size] + const scalar_t* cache_ptr, + const int head_size, + const int num_heads, + const int num_kv_heads, + const int rot_dim, + const int token_idx, + const int64_t query_stride, + const int64_t key_stride, + const int64_t head_stride) { + const int embed_dim = rot_dim / 2; + const scalar_t* cos_ptr = cache_ptr; + const scalar_t* sin_ptr = cache_ptr + embed_dim; + + const int nq = num_heads * embed_dim; + for (int i = threadIdx.x; i < nq; i += blockDim.x) { + const int head_idx = i / embed_dim; + const int64_t token_head = token_idx * query_stride + head_idx * head_stride; + const int rot_offset = i % embed_dim; + apply_token_rotary_embedding(query + token_head, cos_ptr, sin_ptr, rot_offset, embed_dim); + } + + if (key != nullptr) { + const int nk = num_kv_heads * embed_dim; + for (int i = threadIdx.x; i < nk; i += blockDim.x) { + const int head_idx = i / embed_dim; + const int64_t token_head = token_idx * key_stride + head_idx * head_stride; + const int rot_offset = i % embed_dim; + apply_token_rotary_embedding(key + token_head, cos_ptr, sin_ptr, rot_offset, embed_dim); + } + } +} + +template +__global__ void rotary_embedding_kernel( + const int64_t* __restrict__ positions, // [batch_size, seq_len] or + // [num_tokens] + scalar_t* __restrict__ query, // [batch_size, seq_len, num_heads, + // head_size] or [num_tokens, num_heads, + // head_size] + scalar_t* __restrict__ key, // nullptr or + // [batch_size, seq_len, num_kv_heads, + // head_size] or [num_tokens, num_kv_heads, + // head_size] + const scalar_t* __restrict__ cos_sin_cache, // [max_position, 2, rot_dim // + // 2] + const int rot_dim, + const int64_t query_stride, + const int64_t key_stride, + const int64_t head_stride, + const int num_heads, + const int num_kv_heads, + const int head_size) { + // Each thread block is responsible for one token. + const int token_idx = blockIdx.x; + int64_t pos = positions[token_idx]; + const scalar_t* cache_ptr = cos_sin_cache + pos * rot_dim; + + apply_rotary_embedding( + query, + key, + cache_ptr, + head_size, + num_heads, + num_kv_heads, + rot_dim, + token_idx, + query_stride, + key_stride, + head_stride); +} + +void rotary_embedding( + torch::Tensor& positions, // [batch_size, seq_len] or [num_tokens] + torch::Tensor& query, // [batch_size, seq_len, num_heads * head_size] or + // [num_tokens, num_heads * head_size] or + // [batch_size, seq_len, num_heads, head_size] or + // [num_tokens, num_heads, head_size] + std::optional key, + // null or + // [batch_size, seq_len, num_kv_heads * head_size] or + // [num_tokens, num_kv_heads * head_size] or + // [batch_size, seq_len, num_heads, head_size] or + // [num_tokens, num_heads, head_size] + int64_t head_size, + torch::Tensor& cos_sin_cache, // [max_position, rot_dim] + bool is_neox) { + // num_tokens = batch_size * seq_len + int64_t num_tokens = positions.numel(); + int positions_ndim = positions.dim(); + + // Make sure num_tokens dim is consistent across positions, query, and key + TORCH_CHECK( + positions_ndim == 1 || positions_ndim == 2, "positions must have shape [num_tokens] or [batch_size, seq_len]"); + if (positions_ndim == 1) { + TORCH_CHECK( + query.size(0) == positions.size(0) && (!key.has_value() || key->size(0) == positions.size(0)), + "query, key and positions must have the same number of tokens"); + } + if (positions_ndim == 2) { + TORCH_CHECK( + query.size(0) == positions.size(0) && (!key.has_value() || key->size(0) == positions.size(0)) && + query.size(1) == positions.size(1) && (!key.has_value() || key->size(1) == positions.size(1)), + "query, key and positions must have the same batch_size and seq_len"); + } + + // Make sure head_size is valid for query and key + // hidden_size = num_heads * head_size + int query_hidden_size = query.numel() / num_tokens; + int key_hidden_size = key.has_value() ? key->numel() / num_tokens : 0; + TORCH_CHECK(query_hidden_size % head_size == 0); + TORCH_CHECK(key_hidden_size % head_size == 0); + + // Make sure query and key have consistent number of heads + int num_heads = query_hidden_size / head_size; + int num_kv_heads = key.has_value() ? key_hidden_size / head_size : num_heads; + TORCH_CHECK(num_heads % num_kv_heads == 0); + + int rot_dim = cos_sin_cache.size(1); + int seq_dim_idx = positions_ndim - 1; + int64_t query_stride = query.stride(seq_dim_idx); + int64_t key_stride = key.has_value() ? key->stride(seq_dim_idx) : 0; + // Determine head stride: for [*, heads, head_size] use stride of last dim; + // for flat [*, heads*head_size], heads blocks are contiguous of size + // head_size + int query_ndim = query.dim(); + int64_t head_stride = (query_ndim == positions_ndim + 2) ? query.stride(-2) : head_size; + + dim3 grid(num_tokens); + dim3 block(std::min(num_heads * rot_dim / 2, 512)); + const at::cuda::OptionalCUDAGuard device_guard(device_of(query)); + const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + DISPATCH_FLOAT_TYPES(query.scalar_type(), "rotary_embedding", [&] { + if (is_neox) { + rotary_embedding_kernel<<>>( + positions.data_ptr(), + query.data_ptr(), + key.has_value() ? key->data_ptr() : nullptr, + cos_sin_cache.data_ptr(), + rot_dim, + query_stride, + key_stride, + head_stride, + num_heads, + num_kv_heads, + head_size); + } else { + rotary_embedding_kernel<<>>( + positions.data_ptr(), + query.data_ptr(), + key.has_value() ? key->data_ptr() : nullptr, + cos_sin_cache.data_ptr(), + rot_dim, + query_stride, + key_stride, + head_stride, + num_heads, + num_kv_heads, + head_size); + } + }); +} diff --git a/sgl-kernel/include/sgl_kernel_ops.h b/sgl-kernel/include/sgl_kernel_ops.h index b5e519ac5..e171ab316 100644 --- a/sgl-kernel/include/sgl_kernel_ops.h +++ b/sgl-kernel/include/sgl_kernel_ops.h @@ -149,6 +149,14 @@ void apply_rope_pos_ids_cos_sin_cache( const std::optional& v_buffer, const std::optional& kv_cache_loc); +void rotary_embedding( + torch::Tensor& positions, + torch::Tensor& query, + std::optional key, + int64_t head_size, + torch::Tensor& cos_sin_cache, + bool is_neox); + void downcast_fp8( at::Tensor& k, at::Tensor& v, diff --git a/sgl-kernel/python/sgl_kernel/__init__.py b/sgl-kernel/python/sgl_kernel/__init__.py index 58436b17a..8e8994e04 100644 --- a/sgl-kernel/python/sgl_kernel/__init__.py +++ b/sgl-kernel/python/sgl_kernel/__init__.py @@ -30,6 +30,7 @@ from sgl_kernel.elementwise import ( gemma_fused_add_rmsnorm, gemma_rmsnorm, rmsnorm, + rotary_embedding, silu_and_mul, ) from sgl_kernel.expert_specialization import ( diff --git a/sgl-kernel/python/sgl_kernel/elementwise.py b/sgl-kernel/python/sgl_kernel/elementwise.py index 5684800be..c7a6a0ed1 100644 --- a/sgl-kernel/python/sgl_kernel/elementwise.py +++ b/sgl-kernel/python/sgl_kernel/elementwise.py @@ -353,6 +353,19 @@ def apply_rope_with_cos_sin_cache_inplace( ) +def rotary_embedding( + positions: torch.Tensor, + query: torch.Tensor, + key: torch.Tensor, + head_size: int, + cos_sin_cache: torch.Tensor, + is_neox: bool = True, +): + torch.ops.sgl_kernel.rotary_embedding.default( + positions, query, key, head_size, cos_sin_cache, is_neox + ) + + def downcast_fp8( k: torch.Tensor, v: torch.Tensor, diff --git a/sgl-kernel/python/sgl_kernel/testing/rotary_embedding.py b/sgl-kernel/python/sgl_kernel/testing/rotary_embedding.py index f1506479b..109778f29 100644 --- a/sgl-kernel/python/sgl_kernel/testing/rotary_embedding.py +++ b/sgl-kernel/python/sgl_kernel/testing/rotary_embedding.py @@ -97,10 +97,6 @@ class RotaryEmbedding(torch.nn.Module): num_tokens = positions.shape[0] cos_sin = self.cos_sin_cache.index_select(0, positions) - # Modification: float32 is required for the rotary embedding to work correctly - query = query.to(torch.float32) - key = key.to(torch.float32) - cos, sin = cos_sin.chunk(2, dim=-1) query_shape = query.shape @@ -146,6 +142,31 @@ class FlashInferRotaryEmbedding(RotaryEmbedding): return query, key +class SglKernelRotaryEmbedding(RotaryEmbedding): + def forward_cuda( + self, + positions: torch.Tensor, + query: torch.Tensor, + key: torch.Tensor, + offsets: Optional[torch.Tensor] = None, + fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + assert ( + fused_set_kv_buffer_arg is None + ), "fused_set_kv_buffer_arg is not supported for sgl-kernel implementation" + if self.cos_sin_cache.dtype != query.dtype: + self.cos_sin_cache = self.cos_sin_cache.to(query.dtype) + torch.ops.sgl_kernel.rotary_embedding( + positions, + query, + key, + self.head_size, + self.cos_sin_cache, + self.is_neox_style, + ) + return query, key + + class MHATokenToKVPool: KV_POOL_SIZE = 16384 diff --git a/sgl-kernel/tests/test_rotary_embedding.py b/sgl-kernel/tests/test_rotary_embedding.py index cc5374dbe..ed2cac32f 100644 --- a/sgl-kernel/tests/test_rotary_embedding.py +++ b/sgl-kernel/tests/test_rotary_embedding.py @@ -7,6 +7,7 @@ from sgl_kernel.testing.rotary_embedding import ( FlashInferRotaryEmbedding, MHATokenToKVPool, RotaryEmbedding, + SglKernelRotaryEmbedding, create_inputs, ) @@ -80,7 +81,7 @@ def test_correctness( rope_ref = RotaryEmbedding(**config).to(device) rope_flashinfer = FlashInferRotaryEmbedding(**config).to(device) - + rope_sglkernel = SglKernelRotaryEmbedding(**config).to(device) inputs = create_inputs( head_size=head_size, batch_size=batch_size, @@ -92,19 +93,27 @@ def test_correctness( ) if save_kv_cache: - pool_ref = MHATokenToKVPool(head_num=num_kv_heads, head_dim=head_size) + pool_ref_for_flashinfer = MHATokenToKVPool( + head_num=num_kv_heads, head_dim=head_size + ) pool_flashinfer = MHATokenToKVPool(head_num=num_kv_heads, head_dim=head_size) query_ref, key_ref = inputs["query"].clone(), inputs["key"].clone() query_flashinfer, key_flashinfer = inputs["query"].clone(), inputs["key"].clone() + query_sglkernel, key_sglkernel = inputs["query"].clone(), inputs["key"].clone() - query_ref_out, key_ref_out = rope_ref.forward_native( + # This is to align with the flashinfer implementation, flashinfer uses float32 cos/sin cache + query_ref_for_flashinfer_out, key_ref_for_flashinfer_out = rope_ref.forward_native( + inputs["pos_ids"], query_ref.to(torch.float32), key_ref.to(torch.float32) + ) + + query_ref_for_sglkernel_out, key_ref_for_sglkernel_out = rope_ref.forward_native( inputs["pos_ids"], query_ref, key_ref ) if save_kv_cache: - pool_ref.set_kv_buffer( + pool_ref_for_flashinfer.set_kv_buffer( loc=inputs["out_cache_loc"], - cache_k=key_ref_out.view(-1, num_kv_heads, head_size), + cache_k=key_ref_for_flashinfer_out.view(-1, num_kv_heads, head_size), cache_v=inputs["value"].view(-1, num_kv_heads, head_size), ) @@ -126,13 +135,27 @@ def test_correctness( ), ) - torch.testing.assert_close( - query_ref_out, query_flashinfer_out, atol=1e-2, rtol=1e-2 + query_sglkernel_out, key_sglkernel_out = rope_sglkernel.forward_cuda( + inputs["pos_ids"], + query_sglkernel, + key_sglkernel, + ) + + torch.testing.assert_close( + query_ref_for_flashinfer_out, query_flashinfer_out, atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + key_ref_for_flashinfer_out, key_flashinfer_out, atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + query_ref_for_sglkernel_out, query_sglkernel_out, atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + key_ref_for_sglkernel_out, key_sglkernel_out, atol=1e-2, rtol=1e-2 ) - torch.testing.assert_close(key_ref_out, key_flashinfer_out, atol=1e-2, rtol=1e-2) if save_kv_cache: for field in ["k_buffer", "v_buffer"]: - x_ref = getattr(pool_ref, field)[0] + x_ref = getattr(pool_ref_for_flashinfer, field)[0] x_flashinfer = getattr(pool_flashinfer, field)[0] torch.testing.assert_close(x_ref, x_flashinfer, atol=1e-2, rtol=1e-2) nonzero_ref = x_ref != 0