Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions onnxruntime/contrib_ops/cpu/sparse/sparse_attention_helper.h
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,10 @@

#pragma once

#include <initializer_list>

#include "core/common/common.h"
#include "core/common/safeint.h"
#include "core/providers/common.h"
#include "contrib_ops/cpu/bert/attention_common.h"
#include "contrib_ops/cpu/bert/attention_parameters.h"
Expand All @@ -12,6 +15,16 @@ namespace onnxruntime {
namespace contrib {
namespace sparse_attention_helper {

inline bool ProductExceedsInt32Max(std::initializer_list<int64_t> factors) {
int32_t product = 1;
for (int64_t factor : factors) {
int32_t next_product;
if (factor < 0 || !SafeMultiply(product, factor, next_product)) return true;
product = next_product;
}
return false;
}

Status CheckInputs(void* params,
const Tensor* query,
const Tensor* key,
Expand Down Expand Up @@ -264,6 +277,22 @@ Status CheckInputs(void* params,
parameters->stride_row_indices = static_cast<int>(block_row_indices_dim[1]);
parameters->stride_col_indices = static_cast<int>(block_col_indices_dim[1]);

// Buffer sizes and kernel strides are computed as products of these dimensions, and the Triton kernels
// take 32-bit strides. Reject shapes whose products do not fit in int32 so that the sizes and offsets
// derived from them cannot wrap.
if (ProductExceedsInt32Max({batch_size, sequence_length,
static_cast<int64_t>(num_heads) + 2 * static_cast<int64_t>(kv_num_heads), head_size})) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"batch_size * sequence_length * (num_heads + 2 * kv_num_heads) * head_size shall not "
"exceed int32 max");
}

if (ProductExceedsInt32Max({batch_size, kv_num_heads, max_cache_sequence_length, head_size})) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"batch_size * kv_num_heads * max_cache_sequence_length * head_size shall not exceed "
"int32 max");
}

return Status::OK();
}

Expand Down
11 changes: 6 additions & 5 deletions onnxruntime/contrib_ops/cuda/sparse/sparse_attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#include "contrib_ops/cpu/sparse/sparse_attention_helper.h"
#include "contrib_ops/cuda/sparse/sparse_attention_v1/sparse_attention_v1_api.h"
#include "contrib_ops/cuda/sparse/sparse_attention_v2/sparse_attention_v2_api.h"
#include "core/common/safeint.h"
#include "core/platform/env_var_utils.h"
#include "contrib_ops/cuda/bert/transformer_cuda_common.h"

Expand Down Expand Up @@ -238,16 +239,16 @@ Status SparseAttention<T>::ComputeInternal(OpKernelContext* context) const {

size_t rotary_buffer_bytes = 0;
if (do_rotary_) {
rotary_buffer_bytes = 2 * sizeof(T) * parameters.batch_size * parameters.num_heads *
rotary_buffer_bytes = 2 * sizeof(T) * SafeInt<size_t>(parameters.batch_size) * parameters.num_heads *
parameters.sequence_length * parameters.head_size;
rotary_buffer_bytes += sizeof(int64_t) * parameters.batch_size * parameters.sequence_length;
rotary_buffer_bytes += sizeof(int64_t) * SafeInt<size_t>(parameters.batch_size) * parameters.sequence_length;
}
auto rotary_buffer = GetScratchBuffer<void>(rotary_buffer_bytes, GetComputeStream(context));
data.rotary_buffer = reinterpret_cast<CudaT*>(rotary_buffer.get());

size_t transposed_q_bytes = 0;
if (!parameters.is_packed_qkv) {
transposed_q_bytes = parameters.batch_size * parameters.sequence_length *
transposed_q_bytes = SafeInt<size_t>(parameters.batch_size) * parameters.sequence_length *
parameters.num_heads * parameters.head_size * sizeof(T);
}
auto transposed_q_buffer = GetScratchBuffer<void>(transposed_q_bytes, GetComputeStream(context));
Expand All @@ -257,8 +258,8 @@ Status SparseAttention<T>::ComputeInternal(OpKernelContext* context) const {

size_t unpacked_qkv_bytes = 0;
if (parameters.is_packed_qkv) {
unpacked_qkv_bytes = (parameters.batch_size * parameters.sequence_length *
(parameters.num_heads + 2 * parameters.kv_num_heads) *
unpacked_qkv_bytes = (SafeInt<size_t>(parameters.batch_size) * parameters.sequence_length *
(SafeInt<size_t>(parameters.num_heads) + 2 * parameters.kv_num_heads) *
parameters.head_size * sizeof(T));
}
auto unpacked_qkv_buffer = GetScratchBuffer<void>(unpacked_qkv_bytes, GetComputeStream(context));
Expand Down
7 changes: 4 additions & 3 deletions onnxruntime/contrib_ops/cuda/sparse/sparse_attention_impl.cu
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
// Licensed under the MIT License.

#include "contrib_ops/cuda/sparse/sparse_attention_impl.h"
#include "core/common/safeint.h"
#include "contrib_ops/cuda/utils/dump_cuda_tensor.h"
#include "contrib_ops/cuda/bert/rotary_embedding_impl.h"
#include "contrib_ops/cuda/bert/group_query_attention_impl.h"
Expand Down Expand Up @@ -132,8 +133,8 @@ Status QkvToContext(
key = reinterpret_cast<const void*>(data.key);
value = reinterpret_cast<const void*>(data.value);
} else {
size_t q_size = static_cast<size_t>(batch_size * sequence_length * num_heads * head_size);
size_t k_size = static_cast<size_t>(batch_size * sequence_length * kv_num_heads * head_size);
size_t q_size = SafeInt<size_t>(batch_size) * sequence_length * num_heads * head_size;
size_t k_size = SafeInt<size_t>(batch_size) * sequence_length * kv_num_heads * head_size;
auto q = reinterpret_cast<T*>(data.unpacked_qkv_buffer);
auto k = reinterpret_cast<T*>(data.unpacked_qkv_buffer + q_size);
auto v = reinterpret_cast<T*>(data.unpacked_qkv_buffer + q_size + k_size);
Expand Down Expand Up @@ -165,7 +166,7 @@ Status QkvToContext(
#endif

if (parameters.do_rotary) {
size_t bsh = static_cast<size_t>(parameters.batch_size * parameters.sequence_length * parameters.head_size);
size_t bsh = SafeInt<size_t>(parameters.batch_size) * parameters.sequence_length * parameters.head_size;
size_t q_size = bsh * static_cast<size_t>(parameters.num_heads);
size_t k_size = bsh * static_cast<size_t>(parameters.kv_num_heads);
auto q_buffer = reinterpret_cast<T*>(data.rotary_buffer);
Expand Down
Loading