diff --git a/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h b/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h index fe1c57e571..1d2bb9d235 100644 --- a/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h +++ b/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h @@ -202,7 +202,7 @@ Status CheckInputs(const T* query, if (mask_dims.size() == 1) { if (mask_dims[0] == static_cast(batch_size)) { mask_type = AttentionMaskType::MASK_1D_KEY_SEQ_LEN; - } else if (mask_dims[0] == static_cast(3 * batch_size + 2)) { + } else if (mask_dims[0] == static_cast(3) * static_cast(batch_size) + static_cast(2)) { mask_type = AttentionMaskType::MASK_1D_KEY_SEQ_LEN_START; } } else if (mask_dims.size() == 2 && mask_dims[0] == static_cast(batch_size) && mask_dims[1] == static_cast(kv_sequence_length)) { diff --git a/onnxruntime/contrib_ops/cuda/bert/longformer_attention.cc b/onnxruntime/contrib_ops/cuda/bert/longformer_attention.cc index 68b60b6e02..e556ae4a49 100644 --- a/onnxruntime/contrib_ops/cuda/bert/longformer_attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/longformer_attention.cc @@ -207,25 +207,26 @@ Status LongformerAttention::ComputeInternal(OpKernelContext* context) const { input_data, k, &zero, global_q, n, device_prop)); } else { - CUBLAS_RETURN_IF_ERROR(cublasGemmStridedBatchedHelper(cublas, - CUBLAS_OP_N, - CUBLAS_OP_N, - hidden_size, // m - max_num_global, // n - hidden_size, // k - &one, // alpha - global_q_weight, // A - hidden_size, // lda - 0, // strideA - input_data, // B - hidden_size, // ldb - sequence_length * hidden_size, // strideB - &zero, // beta - global_q, // C - hidden_size, // ldc - max_num_global * hidden_size, // strideC - batch_size, // batch count - device_prop)); + CUBLAS_RETURN_IF_ERROR(cublasGemmStridedBatchedHelper( + cublas, + CUBLAS_OP_N, + CUBLAS_OP_N, + hidden_size, // m + max_num_global, // n + hidden_size, // k + &one, // alpha + global_q_weight, // A + hidden_size, // lda + 0, // strideA + input_data, // B + hidden_size, // ldb + static_cast(sequence_length) * hidden_size, // strideB + &zero, // beta + global_q, // C + hidden_size, // ldc + static_cast(max_num_global) * hidden_size, // strideC + batch_size, // batch count + device_prop)); } // global k const CudaT* global_k_weight = global_weights_data + static_cast(hidden_size) * hidden_size; diff --git a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc index 7c4b65b113..9e9ed2e512 100644 --- a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc @@ -220,12 +220,13 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const { data.mask_index = (nullptr == key_padding_mask) ? nullptr : key_padding_mask->Data(); data.mask_index_dims = (nullptr == key_padding_mask) ? gsl::span() : key_padding_mask->Shape().GetDims(); data.past = nullptr; - data.past_key = (parameters.pass_past_in_kv) ? reinterpret_cast(key->Data()) - : (nullptr == past_key) ? nullptr - : reinterpret_cast(past_key->Data()); - data.past_value = (parameters.pass_past_in_kv) ? reinterpret_cast(value->Data()) - : (nullptr == past_value) ? nullptr - : reinterpret_cast(past_value->Data()); + const bool pass_key_value_as_past = (parameters.pass_past_in_kv && nullptr != key && nullptr != value); + data.past_key = pass_key_value_as_past ? reinterpret_cast(key->Data()) + : (nullptr == past_key) ? nullptr + : reinterpret_cast(past_key->Data()); + data.past_value = pass_key_value_as_past ? reinterpret_cast(value->Data()) + : (nullptr == past_value) ? nullptr + : reinterpret_cast(past_value->Data()); data.relative_position_bias = (nullptr == relative_position_bias) ? nullptr : reinterpret_cast(relative_position_bias->Data()); data.has_qkv_workspace = !no_qkv_workspace; data.workspace = reinterpret_cast(work_space.get()); diff --git a/onnxruntime/contrib_ops/cuda/bert/tensorrt_fused_multihead_attention/cross_attention/fmha_cross_attention.h b/onnxruntime/contrib_ops/cuda/bert/tensorrt_fused_multihead_attention/cross_attention/fmha_cross_attention.h index b0ed275b4a..d693c75bdf 100644 --- a/onnxruntime/contrib_ops/cuda/bert/tensorrt_fused_multihead_attention/cross_attention/fmha_cross_attention.h +++ b/onnxruntime/contrib_ops/cuda/bert/tensorrt_fused_multihead_attention/cross_attention/fmha_cross_attention.h @@ -217,7 +217,7 @@ static Fused_multihead_attention_params_mhca getMHCAParams( params.gmem_q_params.cu_seqlens = static_cast(cu_seqlens_q_d); params.gmem_kv_params.ptr = const_cast(kv_packed_d); - params.gmem_kv_params.stride_in_bytes = h * 2 * d * sizeof(half); + params.gmem_kv_params.stride_in_bytes = static_cast(h) * 2 * d * static_cast(sizeof(half)); params.gmem_kv_params.h = h; params.gmem_kv_params.d = d; params.gmem_kv_params.cu_seqlens = static_cast(cu_seqlens_kv_d);