diff --git a/sgl-kernel/csrc/cpu/rope.cpp b/sgl-kernel/csrc/cpu/rope.cpp index 1c6249466..a6ff0cdf2 100644 --- a/sgl-kernel/csrc/cpu/rope.cpp +++ b/sgl-kernel/csrc/cpu/rope.cpp @@ -69,18 +69,23 @@ void rotary_embedding_3D_kernel_impl( } template -void rotary_embedding_neox_2D_kernel_impl( +void rotary_embedding_neox_4D_kernel_impl( int64_t* __restrict__ positions, scalar_t* __restrict__ query, scalar_t* __restrict__ key, scalar_t* __restrict__ cos_sin_cache, int64_t rotary_dim, + int64_t query_stride_b, int64_t query_stride_s, + int64_t query_stride_h, + int64_t key_stride_b, int64_t key_stride_s, + int64_t key_stride_h, int64_t num_heads, int64_t num_kv_heads, int64_t head_size, - int64_t num_tokens) { + int64_t batch_size, + int64_t seq_len) { using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; constexpr int64_t bVecSize = bVec::size(); @@ -143,50 +148,57 @@ void rotary_embedding_neox_2D_kernel_impl( } }; -#pragma omp parallel for - for (int64_t token_idx = 0; token_idx < num_tokens; ++token_idx) { - int64_t pos = positions[token_idx]; - scalar_t* cache_ptr = cos_sin_cache + pos * rotary_dim; +#pragma omp parallel for collapse(2) + for (int64_t bs = 0; bs < batch_size; ++bs) { + for (int64_t seq = 0; seq < seq_len; ++seq) { + int64_t pos = positions[bs * seq_len + seq]; + scalar_t* cache_ptr = cos_sin_cache + pos * rotary_dim; - for (int64_t i = 0; i < num_heads; ++i) { - int64_t head_idx = i; - int64_t token_head = token_idx * query_stride_s + head_idx * head_size; - compute_loop(token_head, cache_ptr, query); - } + for (int64_t i = 0; i < num_heads; ++i) { + int64_t head_idx = i; + int64_t token_head = bs * query_stride_b + seq * query_stride_s + head_idx * query_stride_h; + compute_loop(token_head, cache_ptr, query); + } - for (int64_t i = 0; i < num_kv_heads; ++i) { - int64_t head_idx = i; - int64_t token_head = token_idx * key_stride_s + head_idx * head_size; - compute_loop(token_head, cache_ptr, key); + for (int64_t i = 0; i < num_kv_heads; ++i) { + int64_t head_idx = i; + int64_t token_head = bs * key_stride_b + seq * key_stride_s + head_idx * key_stride_h; + compute_loop(token_head, cache_ptr, key); + } } } } template -void rotary_embedding_2D_kernel_impl( +void rotary_embedding_4D_kernel_impl( int64_t* __restrict__ positions, scalar_t* __restrict__ query, scalar_t* __restrict__ key, scalar_t* __restrict__ cos_sin_cache, int64_t rotary_dim, + int64_t query_stride_b, int64_t query_stride_s, + int64_t query_stride_h, + int64_t key_stride_b, int64_t key_stride_s, + int64_t key_stride_h, int64_t num_heads, int64_t num_kv_heads, int64_t head_size, - int64_t num_tokens) { + int64_t batch_size, + int64_t seq_len) { int64_t embed_dim = rotary_dim / 2; - at::parallel_for(0, num_tokens * num_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) { - int64_t token_idx = {0}, i = {0}; - data_index_init(begin, token_idx, num_tokens, i, num_heads); + at::parallel_for(0, batch_size * seq_len * num_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) { + int64_t bs = {0}, seq = {0}, i = {0}; + data_index_init(begin, bs, batch_size, seq, seq_len, i, num_heads); for ([[maybe_unused]] auto z : c10::irange(begin, end)) { - int64_t pos = positions[token_idx]; + int64_t pos = positions[bs * seq_len + seq]; scalar_t* cache_ptr = cos_sin_cache + pos * rotary_dim; scalar_t* cos_cache_ptr = cache_ptr; scalar_t* sin_cache_ptr = cache_ptr + embed_dim; int64_t head_idx = i; - int64_t token_head = token_idx * query_stride_s + head_idx * head_size; + int64_t token_head = bs * query_stride_b + seq * query_stride_s + head_idx * query_stride_h; scalar_t* head_query = token_head + query; for (int64_t j = 0; j < embed_dim; j += 1) { int64_t rot_offset = j; @@ -202,20 +214,20 @@ void rotary_embedding_2D_kernel_impl( head_query[x_index] = x * cos - y * sin; head_query[y_index] = y * cos + x * sin; } - data_index_step(token_idx, num_tokens, i, num_heads); + data_index_step(bs, batch_size, seq, seq_len, i, num_heads); } }); - at::parallel_for(0, num_tokens * num_kv_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) { - int64_t token_idx{0}, i = {0}; - data_index_init(begin, token_idx, num_tokens, i, num_kv_heads); + at::parallel_for(0, batch_size * seq_len * num_kv_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) { + int64_t bs = {0}, seq = {0}, i = {0}; + data_index_init(begin, bs, batch_size, seq, seq_len, i, num_kv_heads); for ([[maybe_unused]] auto z : c10::irange(begin, end)) { - int64_t pos = positions[token_idx]; + int64_t pos = positions[bs * seq_len + seq]; scalar_t* cache_ptr = cos_sin_cache + pos * rotary_dim; scalar_t* cos_cache_ptr = cache_ptr; scalar_t* sin_cache_ptr = cache_ptr + embed_dim; int64_t head_idx = i; - int64_t token_head = token_idx * key_stride_s + head_idx * head_size; + int64_t token_head = bs * key_stride_b + seq * key_stride_s + head_idx * head_size; scalar_t* head_key = key + token_head; for (int64_t j = 0; j < embed_dim; j += 1) { int64_t rot_offset = j; @@ -231,7 +243,7 @@ void rotary_embedding_2D_kernel_impl( head_key[x_index] = x * cos - y * sin; head_key[y_index] = y * cos + x * sin; } - data_index_step(token_idx, num_tokens, i, num_kv_heads); + data_index_step(bs, batch_size, seq, seq_len, i, num_kv_heads); } }); } @@ -250,8 +262,9 @@ std::tuple rotary_embedding_cpu( const auto input_dim = query.dim(); const auto input_dtype = query.scalar_type(); TORCH_CHECK( - input_dim == 2 || input_dim == 3, - " Query/Key must be 2D [num_tokens, num_heads*head_size] or 3D [num_tokens, num_heads, head_size] tensor"); + input_dim == 2 || input_dim == 3 || input_dim == 4, + " Query/Key must be 2D [num_tokens, num_heads*head_size] or 3D [num_tokens, num_heads, head_size] or 4D " + "[batch_size, seq_len, num_heads, head_size] tensor"); CHECK_DIM(2, cos_sin_cache); CHECK_LAST_DIM_CONTIGUOUS_INPUT(query); CHECK_LAST_DIM_CONTIGUOUS_INPUT(key); @@ -265,55 +278,83 @@ std::tuple rotary_embedding_cpu( } int64_t num_tokens = positions.numel(); - CHECK_EQ(key.size(0), num_tokens); - CHECK_EQ(query.size(0), num_tokens); + if (input_dim <= 3) { + CHECK_EQ(key.size(0), num_tokens); + CHECK_EQ(query.size(0), num_tokens); + } TORCH_CHECK(positions.scalar_type() == at::kLong, "expect positions to be int64, got ", positions.scalar_type()); TORCH_CHECK(input_dtype == key.scalar_type(), "query and key must have the same data type"); TORCH_CHECK(input_dtype == cos_sin_cache.scalar_type(), "query and cos_sin_cache must have the same data type"); - int64_t num_heads = input_dim == 2 ? query.size(-1) / head_size : query.size(1); - int64_t num_kv_heads = input_dim == 2 ? key.size(-1) / head_size : key.size(1); + int64_t num_heads = input_dim == 2 ? query.size(-1) / head_size : query.size(-2); + int64_t num_kv_heads = input_dim == 2 ? key.size(-1) / head_size : key.size(-2); int64_t key_stride_s = key.stride(0); int64_t query_stride_s = query.stride(0); - // input stride of num head dim is meaningful only when input dim = 3 - int64_t query_stride_h = input_dim == 3 ? query.stride(1) : -1; + int64_t query_stride_h = input_dim == 2 ? head_size : query.stride(-2); + int64_t key_stride_h = input_dim == 2 ? head_size : key.stride(-2); at::Tensor query_out = at::empty_like(query); at::Tensor key_out = at::empty_like(key); int64_t query_out_stride_s = query_out.stride(0); int64_t key_out_stride_s = key_out.stride(0); // output stride of num head dim is meaningful only when input dim = 3 int64_t query_out_stride_h = input_dim == 3 ? query_out.stride(1) : -1; + int64_t batch_size = 1; + int64_t seq_len = num_tokens; + int64_t query_stride_b = 0; + int64_t key_stride_b = 0; + if (input_dim == 4) { + batch_size = query.size(0); + seq_len = query.size(1); + query_stride_b = query.stride(0); + key_stride_b = key.stride(0); + query_stride_s = query.stride(1); + key_stride_s = key.stride(1); + CHECK_EQ(batch_size, key.size(0)); + CHECK_EQ(seq_len, key.size(1)); + CHECK_EQ(key.size(0) * key.size(1), num_tokens); + CHECK_EQ(query.size(0) * query.size(1), num_tokens); + } AT_DISPATCH_REDUCED_FLOATING_TYPES(input_dtype, "rotary_embedding_cpu", [&] { - if (input_dim == 2) { + if (input_dim == 2 || input_dim == 4) { if (is_neox) { - rotary_embedding_neox_2D_kernel_impl( + rotary_embedding_neox_4D_kernel_impl( positions.data_ptr(), query.data_ptr(), key.data_ptr(), cos_sin_cache.data_ptr(), rotary_dim, + query_stride_b, query_stride_s, + query_stride_h, + key_stride_b, key_stride_s, + key_stride_h, num_heads, num_kv_heads, head_size, - num_tokens); + batch_size, + seq_len); } else { - rotary_embedding_2D_kernel_impl( + rotary_embedding_4D_kernel_impl( positions.data_ptr(), query.data_ptr(), key.data_ptr(), cos_sin_cache.data_ptr(), rotary_dim, + query_stride_b, query_stride_s, + query_stride_h, + key_stride_b, key_stride_s, + key_stride_h, num_heads, num_kv_heads, head_size, - num_tokens); + batch_size, + seq_len); } query_out = query; key_out = key; diff --git a/test/srt/cpu/test_rope.py b/test/srt/cpu/test_rope.py index 8c1dfe9aa..00aace4ed 100644 --- a/test/srt/cpu/test_rope.py +++ b/test/srt/cpu/test_rope.py @@ -88,6 +88,7 @@ class TestROPE(CustomTestCase): rotary_dim: int, max_position_embeddings: int, base: int, + dims: int, is_neox_style: bool, dtype: torch.dtype, device: str, @@ -119,7 +120,9 @@ class TestROPE(CustomTestCase): dtype=dtype, device=device, ) - + if dims == 4: + query = query.view(batch_size, seq_len, num_q_heads, head_size) + key = key.view(batch_size, seq_len, num_kv_heads, head_size) query_ref, key_ref = query.clone(), key.clone() query_cpu, key_cpu = query.clone(), key.clone() @@ -161,19 +164,21 @@ class TestROPE(CustomTestCase): num_q_heads, num_kv_heads, ) in test_config: - single_test( - head_size, - rotary_dim, - max_position_embeddings, - base, - is_neox_style, - dtype, - device, - batch_size, - seq_len, - num_q_heads, - num_kv_heads, - ) + for dim in [2, 4]: + single_test( + head_size, + rotary_dim, + max_position_embeddings, + base, + dim, + is_neox_style, + dtype, + device, + batch_size, + seq_len, + num_q_heads, + num_kv_heads, + ) if __name__ == "__main__":