[CPU] Add 4D input support for ROPE in sgl-kernel (#9337)
This commit is contained in:
@@ -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
@@ -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__":
|
||||||
|
|||||||
Reference in New Issue
Block a user