[CPU] Fix issues when running llama3.2-11B vision model with image tasks (#8666)
Co-authored-by: JieXin Liang <Alcanderian@users.noreply.github.com> Co-authored-by: Yineng Zhang <me@zhyncs.com> Co-authored-by: jianan-gu <jianan.gu@intel.com>
This commit is contained in:
co-authored by
JieXin Liang
Yineng Zhang
jianan-gu
parent
79b937aefb
commit
84ea47eb22
@@ -1038,6 +1038,7 @@ void decode_attention_kernel_impl(
|
||||
const index_t* __restrict__ req_to_token,
|
||||
const int64_t* __restrict__ req_pool_indices,
|
||||
const int64_t* __restrict__ seq_lens,
|
||||
const int64_t* __restrict__ encoder_lens,
|
||||
int64_t batches,
|
||||
int64_t num_heads,
|
||||
int64_t head_size,
|
||||
@@ -1053,7 +1054,9 @@ void decode_attention_kernel_impl(
|
||||
float logit_cap,
|
||||
int64_t max_num_reqs,
|
||||
int64_t max_context_len,
|
||||
int64_t max_total_num_tokens) {
|
||||
int64_t max_total_num_tokens,
|
||||
bool is_cross_attn,
|
||||
bool has_encoder_lens) {
|
||||
using Vec = at::vec::Vectorized<float>;
|
||||
|
||||
// strides
|
||||
@@ -1077,8 +1080,9 @@ void decode_attention_kernel_impl(
|
||||
const scalar_t* __restrict__ q_ptr = query + bs * q_strideM + head_id * q_strideH;
|
||||
|
||||
// get key/value
|
||||
int64_t seq_len_kv = seq_lens[bs];
|
||||
int64_t seq_len_kv = is_cross_attn ? encoder_lens[bs] : seq_lens[bs];
|
||||
int64_t req_pool_id = req_pool_indices[bs];
|
||||
int64_t kv_offset = (has_encoder_lens && (!is_cross_attn)) ? encoder_lens[bs] : 0;
|
||||
TORCH_CHECK(seq_len_kv <= max_context_len, "seq_len_kv out of scope!");
|
||||
TORCH_CHECK(req_pool_id < max_num_reqs, "req_pool_id out of scope!");
|
||||
|
||||
@@ -1102,7 +1106,7 @@ void decode_attention_kernel_impl(
|
||||
/* A */ q_ptr,
|
||||
/* B */ k_buffer + head_id * k_strideH,
|
||||
/* C */ s_i,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
|
||||
/* scl */ sm_scale,
|
||||
/* M */ 1,
|
||||
/* N */ n_size,
|
||||
@@ -1142,7 +1146,7 @@ void decode_attention_kernel_impl(
|
||||
/* A */ s_delta,
|
||||
/* B */ v_buffer + head_id * v_strideH,
|
||||
/* C */ v_prime,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
|
||||
/* scl */ &m_delta,
|
||||
/* M */ 1,
|
||||
/* N */ head_size_v,
|
||||
@@ -1159,6 +1163,8 @@ void decode_attention_kernel_impl(
|
||||
at::vec::map<float>([s](Vec out) { return out * Vec(s); }, v_prime, v_prime, head_size_v);
|
||||
|
||||
v_prime[head_size_v] = m_prime + std::log(s_prime);
|
||||
} else {
|
||||
v_prime[head_size_v] = -std::numeric_limits<float>::infinity();
|
||||
}
|
||||
|
||||
// move to the next index
|
||||
@@ -1350,6 +1356,10 @@ void decode_attention_mla_kernel_impl(
|
||||
[s](Vec out) { return out * Vec(s); }, v_prime + h * l_stride1, v_prime + h * l_stride1, head_size_v);
|
||||
(v_prime + h * l_stride1)[head_size_v] = m_prime[h] + std::log(s_prime[h]);
|
||||
}
|
||||
} else {
|
||||
for (int64_t h = 0; h < h_size; ++h) {
|
||||
(v_prime + h * l_stride1)[head_size_v] = -std::numeric_limits<float>::infinity();
|
||||
}
|
||||
}
|
||||
|
||||
// move to the next index
|
||||
@@ -1372,6 +1382,7 @@ void decode_attention_grouped_kernel_impl(
|
||||
const index_t* __restrict__ req_to_token,
|
||||
const int64_t* __restrict__ req_pool_indices,
|
||||
const int64_t* __restrict__ seq_lens,
|
||||
const int64_t* __restrict__ encoder_lens,
|
||||
int64_t batches,
|
||||
int64_t num_heads,
|
||||
int64_t num_heads_kv,
|
||||
@@ -1388,7 +1399,9 @@ void decode_attention_grouped_kernel_impl(
|
||||
float logit_cap,
|
||||
int64_t max_num_reqs,
|
||||
int64_t max_context_len,
|
||||
int64_t max_total_num_tokens) {
|
||||
int64_t max_total_num_tokens,
|
||||
bool is_cross_attn,
|
||||
bool has_encoder_lens) {
|
||||
using Vec = at::vec::Vectorized<float>;
|
||||
|
||||
// block length for heads
|
||||
@@ -1429,8 +1442,9 @@ void decode_attention_grouped_kernel_impl(
|
||||
// get query
|
||||
const scalar_t* __restrict__ q_ptr = query + bs * q_strideM + h_start * q_strideH;
|
||||
|
||||
int64_t seq_len_kv = seq_lens[bs];
|
||||
int64_t seq_len_kv = is_cross_attn ? encoder_lens[bs] : seq_lens[bs];
|
||||
int64_t req_pool_id = req_pool_indices[bs];
|
||||
int64_t kv_offset = (has_encoder_lens && (!is_cross_attn)) ? encoder_lens[bs] : 0;
|
||||
TORCH_CHECK(seq_len_kv <= max_context_len, "seq_len_kv out of scope!");
|
||||
TORCH_CHECK(req_pool_id < max_num_reqs, "req_pool_id out of scope!");
|
||||
|
||||
@@ -1456,7 +1470,7 @@ void decode_attention_grouped_kernel_impl(
|
||||
/* A */ q_ptr,
|
||||
/* B */ k_buffer + head_kv_id * k_strideH,
|
||||
/* C */ s_i,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
|
||||
/* scl */ sm_scale,
|
||||
/* M */ h_size,
|
||||
/* N */ n_size,
|
||||
@@ -1500,7 +1514,7 @@ void decode_attention_grouped_kernel_impl(
|
||||
/* A */ s_delta,
|
||||
/* B */ v_buffer + head_kv_id * v_strideH,
|
||||
/* C */ v_prime,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n,
|
||||
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
|
||||
/* scl */ m_delta,
|
||||
/* M */ h_size,
|
||||
/* N */ head_size_v,
|
||||
@@ -1519,6 +1533,10 @@ void decode_attention_grouped_kernel_impl(
|
||||
[s](Vec out) { return out * Vec(s); }, v_prime + h * l_stride1, v_prime + h * l_stride1, head_size_v);
|
||||
(v_prime + h * l_stride1)[head_size_v] = m_prime[h] + std::log(s_prime[h]);
|
||||
}
|
||||
} else {
|
||||
for (int64_t h = 0; h < h_size; ++h) {
|
||||
(v_prime + h * l_stride1)[head_size_v] = -std::numeric_limits<float>::infinity();
|
||||
}
|
||||
}
|
||||
|
||||
// move to the next index
|
||||
@@ -1540,32 +1558,30 @@ void decode_attention_grouped_kernel_impl(
|
||||
// req_to_token: [max_num_reqs, max_context_len] int32 or int64
|
||||
// req_pool_indices: [num_seqs] int64
|
||||
// seq_lens: [num_seqs] int64
|
||||
// encoder_lens: [num_seqs] int64 or None
|
||||
//
|
||||
void decode_attention_cpu(
|
||||
at::Tensor& query,
|
||||
at::Tensor& k_buffer,
|
||||
at::Tensor& v_buffer,
|
||||
at::Tensor& output,
|
||||
at::Tensor& key,
|
||||
at::Tensor& value,
|
||||
const std::optional<at::Tensor>& key,
|
||||
const std::optional<at::Tensor>& value,
|
||||
at::Tensor& loc,
|
||||
at::Tensor& attn_logits,
|
||||
at::Tensor& req_to_token,
|
||||
at::Tensor& req_pool_indices,
|
||||
at::Tensor& seq_lens,
|
||||
double sm_scale,
|
||||
double logit_cap) {
|
||||
double logit_cap,
|
||||
bool is_cross_attn,
|
||||
std::optional<at::Tensor> encoder_lens) {
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(query);
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(k_buffer);
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(v_buffer);
|
||||
// for MLA, key and value shares the same storage and value could be non-contiguous
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(key);
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(value);
|
||||
CHECK_DIM(3, query);
|
||||
CHECK_DIM(3, k_buffer);
|
||||
CHECK_DIM(3, v_buffer);
|
||||
CHECK_DIM(3, key);
|
||||
CHECK_DIM(3, value);
|
||||
CHECK_DIM(1, loc);
|
||||
|
||||
int64_t num_seqs = seq_lens.size(0);
|
||||
@@ -1580,7 +1596,6 @@ void decode_attention_cpu(
|
||||
|
||||
int64_t num_kv_splits = attn_logits.size(2);
|
||||
|
||||
CHECK_EQ(loc.numel(), num_seqs);
|
||||
CHECK_EQ(attn_logits.size(0), num_seqs);
|
||||
CHECK_EQ(attn_logits.size(1), num_heads);
|
||||
CHECK_EQ(attn_logits.size(3), head_size_v + 1);
|
||||
@@ -1595,11 +1610,6 @@ void decode_attention_cpu(
|
||||
int64_t k_strideH = k_buffer.stride(1);
|
||||
int64_t v_strideN = v_buffer.stride(0);
|
||||
int64_t v_strideH = v_buffer.stride(1);
|
||||
// strides for new key and value
|
||||
int64_t nk_strideN = key.stride(0);
|
||||
int64_t nk_strideH = key.stride(1);
|
||||
int64_t nv_strideN = value.stride(0);
|
||||
int64_t nv_strideH = value.stride(1);
|
||||
|
||||
// check index data types
|
||||
const auto index_dtype = req_to_token.scalar_type();
|
||||
@@ -1625,29 +1635,51 @@ void decode_attention_cpu(
|
||||
int num_threads = at::get_num_threads();
|
||||
int64_t size_per_thread = is_mla ? BLOCK_N * head_size + BLOCK_N * head_size_v : 0;
|
||||
auto buffer = at::empty({num_threads, size_per_thread}, k_buffer.options());
|
||||
|
||||
bool has_encoder_lens = encoder_lens.has_value();
|
||||
// Since encoder_lens is not used when it is None, encoder_lens_t can be initialized as any tensor of int64_t dtype.
|
||||
at::Tensor encoder_lens_t = seq_lens;
|
||||
if (has_encoder_lens) {
|
||||
encoder_lens_t = encoder_lens.value();
|
||||
CHECK_EQ(encoder_lens_t.size(0), num_seqs);
|
||||
}
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(query.scalar_type(), "decode_attention_kernel", [&] {
|
||||
AT_DISPATCH_INDEX_TYPES(index_dtype, "decode_attention_indices", [&] {
|
||||
// update the kv buffer
|
||||
decode_set_kv_buffer(
|
||||
(scalar_t*)k_buffer_data,
|
||||
(scalar_t*)v_buffer_data,
|
||||
key.data_ptr<scalar_t>(),
|
||||
value.data_ptr<scalar_t>(),
|
||||
loc.data_ptr<int64_t>(),
|
||||
num_seqs,
|
||||
num_heads_kv,
|
||||
head_size,
|
||||
head_size_v,
|
||||
k_strideN,
|
||||
k_strideH,
|
||||
v_strideN,
|
||||
v_strideH,
|
||||
nk_strideN,
|
||||
nk_strideH,
|
||||
nv_strideN,
|
||||
nv_strideH,
|
||||
is_mla);
|
||||
if (key.has_value()) {
|
||||
TORCH_CHECK(value.has_value(), "key and value should have values at the same time")
|
||||
CHECK_EQ(loc.numel(), num_seqs);
|
||||
auto key_tensor = key.value();
|
||||
auto value_tensor = value.value();
|
||||
// for MLA, key and value shares the same storage and value could be non-contiguous
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(key_tensor);
|
||||
CHECK_LAST_DIM_CONTIGUOUS_INPUT(value_tensor);
|
||||
CHECK_DIM(3, key_tensor);
|
||||
CHECK_DIM(3, value_tensor);
|
||||
// strides for new key and value
|
||||
int64_t nk_strideN = key_tensor.stride(0);
|
||||
int64_t nk_strideH = key_tensor.stride(1);
|
||||
int64_t nv_strideN = value_tensor.stride(0);
|
||||
int64_t nv_strideH = value_tensor.stride(1);
|
||||
// update the kv buffer
|
||||
decode_set_kv_buffer(
|
||||
(scalar_t*)k_buffer_data,
|
||||
(scalar_t*)v_buffer_data,
|
||||
key_tensor.data_ptr<scalar_t>(),
|
||||
value_tensor.data_ptr<scalar_t>(),
|
||||
loc.data_ptr<int64_t>(),
|
||||
num_seqs,
|
||||
num_heads_kv,
|
||||
head_size,
|
||||
head_size_v,
|
||||
k_strideN,
|
||||
k_strideH,
|
||||
v_strideN,
|
||||
v_strideH,
|
||||
nk_strideN,
|
||||
nk_strideH,
|
||||
nv_strideN,
|
||||
nv_strideH,
|
||||
is_mla);
|
||||
}
|
||||
|
||||
if (num_heads == num_heads_kv) {
|
||||
// MHA
|
||||
@@ -1660,6 +1692,7 @@ void decode_attention_cpu(
|
||||
req_to_token.data_ptr<index_t>(),
|
||||
req_pool_indices.data_ptr<int64_t>(),
|
||||
seq_lens.data_ptr<int64_t>(),
|
||||
encoder_lens_t.data_ptr<int64_t>(),
|
||||
num_seqs,
|
||||
num_heads,
|
||||
head_size,
|
||||
@@ -1675,7 +1708,9 @@ void decode_attention_cpu(
|
||||
logit_cap,
|
||||
max_num_reqs,
|
||||
max_context_len,
|
||||
max_total_num_tokens);
|
||||
max_total_num_tokens,
|
||||
is_cross_attn,
|
||||
has_encoder_lens);
|
||||
} else if (is_mla) {
|
||||
// MLA
|
||||
decode_attention_mla_kernel_impl<scalar_t, index_t, BLOCK_N>(
|
||||
@@ -1716,6 +1751,7 @@ void decode_attention_cpu(
|
||||
req_to_token.data_ptr<index_t>(),
|
||||
req_pool_indices.data_ptr<int64_t>(),
|
||||
seq_lens.data_ptr<int64_t>(),
|
||||
encoder_lens_t.data_ptr<int64_t>(),
|
||||
num_seqs,
|
||||
num_heads,
|
||||
num_heads_kv,
|
||||
@@ -1732,7 +1768,9 @@ void decode_attention_cpu(
|
||||
logit_cap,
|
||||
max_num_reqs,
|
||||
max_context_len,
|
||||
max_total_num_tokens);
|
||||
max_total_num_tokens,
|
||||
is_cross_attn,
|
||||
has_encoder_lens);
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user