[CPU] Add 4D input support for ROPE in sgl-kernel (#9337)

This commit is contained in:
blzheng
2025-12-16 17:27:39 +08:00
committed by GitHub
parent 272c5fe43e
commit e27635a02d
2 changed files with 102 additions and 56 deletions
+83 -42
View File
@@ -69,18 +69,23 @@ void rotary_embedding_3D_kernel_impl(
} }
template <typename scalar_t> template <typename scalar_t>
void rotary_embedding_neox_2D_kernel_impl( void rotary_embedding_neox_4D_kernel_impl(
int64_t* __restrict__ positions, int64_t* __restrict__ positions,
scalar_t* __restrict__ query, scalar_t* __restrict__ query,
scalar_t* __restrict__ key, scalar_t* __restrict__ key,
scalar_t* __restrict__ cos_sin_cache, scalar_t* __restrict__ cos_sin_cache,
int64_t rotary_dim, int64_t rotary_dim,
int64_t query_stride_b,
int64_t query_stride_s, int64_t query_stride_s,
int64_t query_stride_h,
int64_t key_stride_b,
int64_t key_stride_s, int64_t key_stride_s,
int64_t key_stride_h,
int64_t num_heads, int64_t num_heads,
int64_t num_kv_heads, int64_t num_kv_heads,
int64_t head_size, int64_t head_size,
int64_t num_tokens) { int64_t batch_size,
int64_t seq_len) {
using bVec = at::vec::Vectorized<scalar_t>; using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>; using fVec = at::vec::Vectorized<float>;
constexpr int64_t bVecSize = bVec::size(); constexpr int64_t bVecSize = bVec::size();
@@ -143,50 +148,57 @@ void rotary_embedding_neox_2D_kernel_impl(
} }
}; };
#pragma omp parallel for #pragma omp parallel for collapse(2)
for (int64_t token_idx = 0; token_idx < num_tokens; ++token_idx) { for (int64_t bs = 0; bs < batch_size; ++bs) {
int64_t pos = positions[token_idx]; for (int64_t seq = 0; seq < seq_len; ++seq) {
scalar_t* cache_ptr = cos_sin_cache + pos * rotary_dim; 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) { for (int64_t i = 0; i < num_heads; ++i) {
int64_t head_idx = i; 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;
compute_loop(token_head, cache_ptr, query); compute_loop(token_head, cache_ptr, query);
} }
for (int64_t i = 0; i < num_kv_heads; ++i) { for (int64_t i = 0; i < num_kv_heads; ++i) {
int64_t head_idx = i; 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 * key_stride_h;
compute_loop(token_head, cache_ptr, key); compute_loop(token_head, cache_ptr, key);
}
} }
} }
} }
template <typename scalar_t> template <typename scalar_t>
void rotary_embedding_2D_kernel_impl( void rotary_embedding_4D_kernel_impl(
int64_t* __restrict__ positions, int64_t* __restrict__ positions,
scalar_t* __restrict__ query, scalar_t* __restrict__ query,
scalar_t* __restrict__ key, scalar_t* __restrict__ key,
scalar_t* __restrict__ cos_sin_cache, scalar_t* __restrict__ cos_sin_cache,
int64_t rotary_dim, int64_t rotary_dim,
int64_t query_stride_b,
int64_t query_stride_s, int64_t query_stride_s,
int64_t query_stride_h,
int64_t key_stride_b,
int64_t key_stride_s, int64_t key_stride_s,
int64_t key_stride_h,
int64_t num_heads, int64_t num_heads,
int64_t num_kv_heads, int64_t num_kv_heads,
int64_t head_size, int64_t head_size,
int64_t num_tokens) { int64_t batch_size,
int64_t seq_len) {
int64_t embed_dim = rotary_dim / 2; 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) { at::parallel_for(0, batch_size * seq_len * num_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) {
int64_t token_idx = {0}, i = {0}; int64_t bs = {0}, seq = {0}, i = {0};
data_index_init(begin, token_idx, num_tokens, i, num_heads); data_index_init(begin, bs, batch_size, seq, seq_len, i, num_heads);
for ([[maybe_unused]] auto z : c10::irange(begin, end)) { 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* cache_ptr = cos_sin_cache + pos * rotary_dim;
scalar_t* cos_cache_ptr = cache_ptr; scalar_t* cos_cache_ptr = cache_ptr;
scalar_t* sin_cache_ptr = cache_ptr + embed_dim; scalar_t* sin_cache_ptr = cache_ptr + embed_dim;
int64_t head_idx = i; 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; scalar_t* head_query = token_head + query;
for (int64_t j = 0; j < embed_dim; j += 1) { for (int64_t j = 0; j < embed_dim; j += 1) {
int64_t rot_offset = j; 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[x_index] = x * cos - y * sin;
head_query[y_index] = y * cos + x * 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) { at::parallel_for(0, batch_size * seq_len * num_kv_heads, GRAIN_SIZE / rotary_dim, [&](int64_t begin, int64_t end) {
int64_t token_idx{0}, i = {0}; int64_t bs = {0}, seq = {0}, i = {0};
data_index_init(begin, token_idx, num_tokens, i, num_kv_heads); data_index_init(begin, bs, batch_size, seq, seq_len, i, num_kv_heads);
for ([[maybe_unused]] auto z : c10::irange(begin, end)) { 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* cache_ptr = cos_sin_cache + pos * rotary_dim;
scalar_t* cos_cache_ptr = cache_ptr; scalar_t* cos_cache_ptr = cache_ptr;
scalar_t* sin_cache_ptr = cache_ptr + embed_dim; scalar_t* sin_cache_ptr = cache_ptr + embed_dim;
int64_t head_idx = i; 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; scalar_t* head_key = key + token_head;
for (int64_t j = 0; j < embed_dim; j += 1) { for (int64_t j = 0; j < embed_dim; j += 1) {
int64_t rot_offset = j; 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[x_index] = x * cos - y * sin;
head_key[y_index] = y * cos + x * 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<at::Tensor, at::Tensor> rotary_embedding_cpu(
const auto input_dim = query.dim(); const auto input_dim = query.dim();
const auto input_dtype = query.scalar_type(); const auto input_dtype = query.scalar_type();
TORCH_CHECK( TORCH_CHECK(
input_dim == 2 || input_dim == 3, 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] tensor"); " 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_DIM(2, cos_sin_cache);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(query); CHECK_LAST_DIM_CONTIGUOUS_INPUT(query);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(key); CHECK_LAST_DIM_CONTIGUOUS_INPUT(key);
@@ -265,55 +278,83 @@ std::tuple<at::Tensor, at::Tensor> rotary_embedding_cpu(
} }
int64_t num_tokens = positions.numel(); int64_t num_tokens = positions.numel();
CHECK_EQ(key.size(0), num_tokens); if (input_dim <= 3) {
CHECK_EQ(query.size(0), num_tokens); 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(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 == 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"); 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_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(1); 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 key_stride_s = key.stride(0);
int64_t query_stride_s = query.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 == 2 ? head_size : query.stride(-2);
int64_t query_stride_h = input_dim == 3 ? query.stride(1) : -1; int64_t key_stride_h = input_dim == 2 ? head_size : key.stride(-2);
at::Tensor query_out = at::empty_like(query); at::Tensor query_out = at::empty_like(query);
at::Tensor key_out = at::empty_like(key); at::Tensor key_out = at::empty_like(key);
int64_t query_out_stride_s = query_out.stride(0); int64_t query_out_stride_s = query_out.stride(0);
int64_t key_out_stride_s = key_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 // 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 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", [&] { AT_DISPATCH_REDUCED_FLOATING_TYPES(input_dtype, "rotary_embedding_cpu", [&] {
if (input_dim == 2) { if (input_dim == 2 || input_dim == 4) {
if (is_neox) { if (is_neox) {
rotary_embedding_neox_2D_kernel_impl<scalar_t>( rotary_embedding_neox_4D_kernel_impl<scalar_t>(
positions.data_ptr<int64_t>(), positions.data_ptr<int64_t>(),
query.data_ptr<scalar_t>(), query.data_ptr<scalar_t>(),
key.data_ptr<scalar_t>(), key.data_ptr<scalar_t>(),
cos_sin_cache.data_ptr<scalar_t>(), cos_sin_cache.data_ptr<scalar_t>(),
rotary_dim, rotary_dim,
query_stride_b,
query_stride_s, query_stride_s,
query_stride_h,
key_stride_b,
key_stride_s, key_stride_s,
key_stride_h,
num_heads, num_heads,
num_kv_heads, num_kv_heads,
head_size, head_size,
num_tokens); batch_size,
seq_len);
} else { } else {
rotary_embedding_2D_kernel_impl<scalar_t>( rotary_embedding_4D_kernel_impl<scalar_t>(
positions.data_ptr<int64_t>(), positions.data_ptr<int64_t>(),
query.data_ptr<scalar_t>(), query.data_ptr<scalar_t>(),
key.data_ptr<scalar_t>(), key.data_ptr<scalar_t>(),
cos_sin_cache.data_ptr<scalar_t>(), cos_sin_cache.data_ptr<scalar_t>(),
rotary_dim, rotary_dim,
query_stride_b,
query_stride_s, query_stride_s,
query_stride_h,
key_stride_b,
key_stride_s, key_stride_s,
key_stride_h,
num_heads, num_heads,
num_kv_heads, num_kv_heads,
head_size, head_size,
num_tokens); batch_size,
seq_len);
} }
query_out = query; query_out = query;
key_out = key; key_out = key;
+19 -14
View File
@@ -88,6 +88,7 @@ class TestROPE(CustomTestCase):
rotary_dim: int, rotary_dim: int,
max_position_embeddings: int, max_position_embeddings: int,
base: int, base: int,
dims: int,
is_neox_style: bool, is_neox_style: bool,
dtype: torch.dtype, dtype: torch.dtype,
device: str, device: str,
@@ -119,7 +120,9 @@ class TestROPE(CustomTestCase):
dtype=dtype, dtype=dtype,
device=device, 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_ref, key_ref = query.clone(), key.clone()
query_cpu, key_cpu = query.clone(), key.clone() query_cpu, key_cpu = query.clone(), key.clone()
@@ -161,19 +164,21 @@ class TestROPE(CustomTestCase):
num_q_heads, num_q_heads,
num_kv_heads, num_kv_heads,
) in test_config: ) in test_config:
single_test( for dim in [2, 4]:
head_size, single_test(
rotary_dim, head_size,
max_position_embeddings, rotary_dim,
base, max_position_embeddings,
is_neox_style, base,
dtype, dim,
device, is_neox_style,
batch_size, dtype,
seq_len, device,
num_q_heads, batch_size,
num_kv_heads, seq_len,
) num_q_heads,
num_kv_heads,
)
if __name__ == "__main__": if __name__ == "__main__":