mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-24 19:43:35 +00:00
Add Attention op for multi-head self attention in BERT (#1984)
* Add Attention op for multi head self attention in BERT * Add test cases * Move op from kOnnxDomain to kMSDomain. Limit test to run by CUDA provider only. * fix test * Add float16 test * fix cpu build error * handle cuda error * get last cuda error when failed
This commit is contained in:
parent
7d2f0c79bd
commit
7b39f5090c
10 changed files with 906 additions and 1 deletions
155
onnxruntime/contrib_ops/cuda/bert/attention.cc
Normal file
155
onnxruntime/contrib_ops/cuda/bert/attention.cc
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "attention.h"
|
||||
#include "core/framework/tensorprotoutils.h"
|
||||
#include "core/providers/cuda/cudnn_common.h"
|
||||
#include "core/providers/cuda/cuda_common.h"
|
||||
#include "core/providers/cuda/shared_inc/fpgeneric.h"
|
||||
#include "attention_impl.h"
|
||||
|
||||
using namespace onnxruntime::cuda;
|
||||
using namespace ::onnxruntime::common;
|
||||
using namespace ONNX_NAMESPACE;
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
|
||||
#define REGISTER_KERNEL_TYPED(T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
Attention, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Attention<T>);
|
||||
|
||||
REGISTER_KERNEL_TYPED(float)
|
||||
REGISTER_KERNEL_TYPED(MLFloat16)
|
||||
|
||||
template <typename T>
|
||||
Attention<T>::Attention(const OpKernelInfo& info) : CudaKernel(info) {
|
||||
int64_t num_heads = 0;
|
||||
ORT_ENFORCE(info.GetAttr("num_heads", &num_heads).IsOK() && num_heads > 0);
|
||||
num_heads_ = static_cast<int>(num_heads);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
Status Attention<T>::ComputeInternal(OpKernelContext* context) const {
|
||||
// Input and output shapes:
|
||||
// Input 0 - input : (batch_size, sequence_length, hidden_size)
|
||||
// Input 1 - weights : (hidden_size, 3 * hidden_size)
|
||||
// Input 2 - bias : (3 * hidden_size)
|
||||
// Input 3 - mask_index : (batch_size)
|
||||
// Output : (batch_size, sequence_length, hidden_size)
|
||||
|
||||
const Tensor* input = context->Input<Tensor>(0);
|
||||
const auto dims = input->Shape().GetDims();
|
||||
if (dims.size() != 3) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 0 is expected to have 3 dimensions, got ", dims.size());
|
||||
}
|
||||
int batch_size = static_cast<int>(dims[0]);
|
||||
int sequence_length = static_cast<int>(dims[1]);
|
||||
int hidden_size = static_cast<int>(dims[2]);
|
||||
if (hidden_size % num_heads_ != 0) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 0 dimension 2 should be divisiable by value of the num_heads attribute.");
|
||||
}
|
||||
int head_size = hidden_size / num_heads_;
|
||||
|
||||
const Tensor* weights = context->Input<Tensor>(1);
|
||||
const auto weights_dims = weights->Shape().GetDims();
|
||||
if (weights_dims.size() != 2) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 1 is expected to have 2 dimensions, got ", weights_dims.size());
|
||||
}
|
||||
if (weights_dims[0] != dims[2]) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 1 dimension 0 should have same length as dimension 2 of input 0");
|
||||
}
|
||||
if (weights_dims[1] != 3 * weights_dims[0]) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 1 dimension 1 should be 3 times of dimension 0");
|
||||
}
|
||||
|
||||
const Tensor* bias = context->Input<Tensor>(2);
|
||||
const auto bias_dims = bias->Shape().GetDims();
|
||||
if (bias_dims.size() != 1) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 2 is expected to have 1 dimension, got ", bias_dims.size());
|
||||
}
|
||||
if (bias_dims[0] != weights_dims[1]) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 2 dimension 0 should have same length as dimension 1 of input 1");
|
||||
}
|
||||
|
||||
const Tensor* mask_index = context->Input<Tensor>(3);
|
||||
const auto mask_dims = mask_index->Shape().GetDims();
|
||||
if (mask_dims.size() != 1) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Input 3 is expected to have 1 dimension, got ", mask_dims.size());
|
||||
}
|
||||
if (static_cast<int>(mask_dims[0]) != batch_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"Inputs 3 and 0 shall have same length at dimension 0");
|
||||
}
|
||||
|
||||
TensorShape output_shape(dims);
|
||||
Tensor* output = context->Output(0, output_shape);
|
||||
|
||||
cublasHandle_t cublas = CublasHandle();
|
||||
const size_t element_size = sizeof(T);
|
||||
|
||||
// Use GEMM for fully connection.
|
||||
int m = batch_size * sequence_length;
|
||||
int n = 3 * hidden_size;
|
||||
int k = hidden_size;
|
||||
auto gemm_buffer = GetScratchBuffer<T>(batch_size * sequence_length * 3 * hidden_size * element_size);
|
||||
|
||||
typedef typename ToCudaType<T>::MappedType CudaT;
|
||||
CudaT one = ToCudaType<T>::FromFloat(1.0f);
|
||||
CudaT zero = ToCudaType<T>::FromFloat(0.0f);
|
||||
|
||||
// Bias shape is (N), broadcast using B(N, M) = 1 * bias(N, 1) x ones(1, M) + 0 * B.
|
||||
// TODO: use custom kernel of expand to improve the performance.
|
||||
CUBLAS_RETURN_IF_ERROR(cublasGemmHelper(
|
||||
cublas, CUBLAS_OP_N, CUBLAS_OP_N, n, m, 1, &one,
|
||||
reinterpret_cast<const CudaT*>(bias->template Data<T>()), n,
|
||||
GetConstOnes<CudaT>(m), 1,
|
||||
&zero, reinterpret_cast<CudaT*>(gemm_buffer.get()), n));
|
||||
|
||||
// Gemm, note that CUDA assumes col-major, so result(N, M) = 1 * weights x input + 1 x B.
|
||||
CUBLAS_RETURN_IF_ERROR(cublasGemmHelper(
|
||||
cublas, CUBLAS_OP_N, CUBLAS_OP_N, n, m, k, &one,
|
||||
reinterpret_cast<const CudaT*>(weights->template Data<T>()), n,
|
||||
reinterpret_cast<const CudaT*>(input->template Data<T>()), k,
|
||||
&one, reinterpret_cast<CudaT*>(gemm_buffer.get()), n));
|
||||
|
||||
size_t workSpaceSize = GetAttentionWorkspaceSize(element_size, batch_size, num_heads_, head_size, sequence_length);
|
||||
auto temp_buffer = GetScratchBuffer<void>(workSpaceSize);
|
||||
if (!LaunchAttentionKernel(
|
||||
reinterpret_cast<const CudaT*>(gemm_buffer.get()),
|
||||
mask_index->template Data<int>(),
|
||||
output->template MutableData<T>(),
|
||||
batch_size,
|
||||
sequence_length,
|
||||
num_heads_,
|
||||
head_size,
|
||||
temp_buffer.get(),
|
||||
cublas,
|
||||
element_size)) {
|
||||
// Get last error to reset it to cudaSuccess.
|
||||
CUDA_CALL(cudaGetLastError());
|
||||
return Status(common::ONNXRUNTIME, common::FAIL);
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
28
onnxruntime/contrib_ops/cuda/bert/attention.h
Normal file
28
onnxruntime/contrib_ops/cuda/bert/attention.h
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "core/common/common.h"
|
||||
#include "core/framework/op_kernel.h"
|
||||
#include "core/providers/cuda/cudnn_common.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
|
||||
using namespace onnxruntime::cuda;
|
||||
|
||||
template <typename T>
|
||||
class Attention final : public CudaKernel {
|
||||
public:
|
||||
Attention(const OpKernelInfo& info);
|
||||
Status ComputeInternal(OpKernelContext* context) const override;
|
||||
|
||||
private:
|
||||
int num_heads_; // number of attention heads
|
||||
};
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
413
onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
Normal file
413
onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
Normal file
|
|
@ -0,0 +1,413 @@
|
|||
/*
|
||||
The implementation of this file is based on qkvToContext plugin in TensorRT demo:
|
||||
https://github.com/NVIDIA/TensorRT/tree/release/5.1/demo/BERT/
|
||||
|
||||
Copyright 2019 NVIDIA Corporation
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
// Modifications: scaling is moved from masked softmax to the gemm before that.
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include <cub/cub.cuh>
|
||||
#include <cublas_v2.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include "core/providers/cuda/cu_inc/common.cuh"
|
||||
#include "core/providers/cuda/cuda_common.h"
|
||||
#include "attention_impl.h"
|
||||
|
||||
using namespace onnxruntime::cuda;
|
||||
using namespace cub;
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
|
||||
static size_t AlignTo(size_t a, size_t b) {
|
||||
return CeilDiv(a, b) * b;
|
||||
}
|
||||
|
||||
size_t ScratchSize(size_t element_size, int batch_size, int num_heads, int sequence_length) {
|
||||
const size_t len = batch_size * num_heads * sequence_length * sequence_length;
|
||||
const size_t bytes = len * element_size;
|
||||
|
||||
const size_t alignment = 256;
|
||||
const size_t bytesAligned = AlignTo(bytes, alignment);
|
||||
return bytesAligned;
|
||||
}
|
||||
|
||||
size_t GetAttentionWorkspaceSize(size_t element_size, int batch_size, int num_heads, int head_size, int sequence_length) {
|
||||
size_t qkv_size = 3 * batch_size * sequence_length * num_heads * head_size * element_size;
|
||||
return qkv_size + 2 * ScratchSize(element_size, batch_size, num_heads, sequence_length);
|
||||
}
|
||||
|
||||
template <typename T, unsigned TPB>
|
||||
__device__ inline void Softmax(const int ld, const int last_valid, const T* input, T* output) {
|
||||
using BlockReduce = cub::BlockReduce<float, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmp_storage;
|
||||
|
||||
__shared__ float reverse_z;
|
||||
|
||||
float thread_data(0);
|
||||
const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * ld;
|
||||
for (int i = threadIdx.x; i < last_valid; i += TPB) {
|
||||
const int index = offset + i;
|
||||
const float val = input[index];
|
||||
thread_data += expf(val);
|
||||
}
|
||||
|
||||
cub::Sum sum;
|
||||
const auto z = BlockReduce(tmp_storage).Reduce(thread_data, sum);
|
||||
if (threadIdx.x == 0) {
|
||||
reverse_z = 1.f / z;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int i = threadIdx.x; i < ld; i += TPB) {
|
||||
const int index = offset + i;
|
||||
const float val = (i < last_valid) ? expf(float(input[index])) * reverse_z : 0.f;
|
||||
output[index] = T(val);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, unsigned TPB>
|
||||
__device__ inline void SoftmaxSmall(const int ld, const int last_valid, const T* input, T* output) {
|
||||
using BlockReduce = cub::BlockReduce<float, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmp_storage;
|
||||
|
||||
__shared__ float reverse_z;
|
||||
|
||||
float thread_data(0);
|
||||
const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * ld;
|
||||
const int index = offset + threadIdx.x;
|
||||
if (threadIdx.x < last_valid) {
|
||||
const float val = input[index];
|
||||
thread_data = expf(val);
|
||||
}
|
||||
|
||||
cub::Sum sum;
|
||||
const auto z = BlockReduce(tmp_storage).Reduce(thread_data, sum);
|
||||
if (threadIdx.x == 0) {
|
||||
reverse_z = (1.f) / z;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x < ld) {
|
||||
// this will be 0 for threadIdx.x >= last_valid
|
||||
output[index] = T(thread_data * reverse_z);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, unsigned TPB>
|
||||
__global__ void MaskedSoftmaxKernelSmall(const int sequence_length, const int* mask_index, const T* input, T* output) {
|
||||
__shared__ int last_valid;
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
last_valid = min(sequence_length, mask_index[blockIdx.y]);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
SoftmaxSmall<T, TPB>(sequence_length, last_valid, input, output);
|
||||
}
|
||||
|
||||
template <typename T, unsigned TPB>
|
||||
__global__ void MaskedSoftmaxKernel(const int sequence_length, const int* mask_index, const T* input, T* output) {
|
||||
__shared__ int last_valid;
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
last_valid = min(sequence_length, mask_index[blockIdx.y]);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
Softmax<T, TPB>(sequence_length, last_valid, input, output);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
bool ComputeMaskedSoftmax(cudaStream_t stream, const int sequence_length, const int batch_size, const int num_heads,
|
||||
const int* mask_index, const T* input, T* output) {
|
||||
// Mask is of length batch_size and assumes the valid region is contiguous starting
|
||||
// from the beginning of the sequence
|
||||
|
||||
const dim3 grid(sequence_length * num_heads, batch_size, 1);
|
||||
|
||||
if (sequence_length <= 32) {
|
||||
const int blockSize = 32;
|
||||
MaskedSoftmaxKernelSmall<T, blockSize>
|
||||
<<<grid, blockSize, 0, stream>>>(sequence_length, mask_index, input, output);
|
||||
} else if (sequence_length <= 128) {
|
||||
const int blockSize = 128;
|
||||
MaskedSoftmaxKernelSmall<T, blockSize>
|
||||
<<<grid, blockSize, 0, stream>>>(sequence_length, mask_index, input, output);
|
||||
} else if (sequence_length == 384) {
|
||||
const int blockSize = 384;
|
||||
MaskedSoftmaxKernelSmall<T, blockSize>
|
||||
<<<grid, blockSize, 0, stream>>>(sequence_length, mask_index, input, output);
|
||||
} else {
|
||||
const int blockSize = 256;
|
||||
MaskedSoftmaxKernel<T, blockSize>
|
||||
<<<grid, blockSize, 0, stream>>>(sequence_length, mask_index, input, output);
|
||||
}
|
||||
|
||||
return CUDA_CALL(cudaPeekAtLastError());
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__global__ void TransposeCtx(const int H, const T* input, T* output) {
|
||||
// Input: BxNxSxH
|
||||
// Output: BxSxNxH
|
||||
|
||||
int n = threadIdx.y;
|
||||
int s = blockIdx.x;
|
||||
int b = blockIdx.y;
|
||||
|
||||
int num_heads = blockDim.y;
|
||||
int sequence_length = gridDim.x;
|
||||
|
||||
const int NH = num_heads * H;
|
||||
const int NHS = NH * sequence_length;
|
||||
const int in_offset = s * H + n * sequence_length * H + b * NHS;
|
||||
const int out_offset = n * H + s * NH + b * NHS;
|
||||
|
||||
const int i = threadIdx.x;
|
||||
if (i < H) {
|
||||
output[out_offset + i] = input[in_offset + i];
|
||||
}
|
||||
}
|
||||
|
||||
bool LaunchTransCtx(cudaStream_t stream,
|
||||
const int sequence_length, const int batch_size, const int head_size, const int num_heads,
|
||||
const float* input, float* output) {
|
||||
const dim3 grid(sequence_length, batch_size, 1);
|
||||
if (0 == (head_size & 1)) {
|
||||
const int H = head_size / 2;
|
||||
const float2* input2 = reinterpret_cast<const float2*>(input);
|
||||
float2* output2 = reinterpret_cast<float2*>(output);
|
||||
const dim3 block(H, num_heads, 1);
|
||||
TransposeCtx<float2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else {
|
||||
const dim3 block(head_size, num_heads, 1);
|
||||
TransposeCtx<float><<<grid, block, 0, stream>>>(head_size, input, output);
|
||||
}
|
||||
return CUDA_CALL(cudaPeekAtLastError());
|
||||
}
|
||||
|
||||
bool LaunchTransCtx(cudaStream_t stream,
|
||||
const int sequence_length, const int batch_size, const int head_size, const int num_heads,
|
||||
const half* input, half* output) {
|
||||
const dim3 grid(sequence_length, batch_size, 1);
|
||||
if (0 == (head_size % 4)) {
|
||||
const int H = head_size / 4;
|
||||
const dim3 block(H, num_heads, 1);
|
||||
const float2* input2 = reinterpret_cast<const float2*>(input);
|
||||
float2* output2 = reinterpret_cast<float2*>(output);
|
||||
TransposeCtx<float2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else if (0 == (head_size & 1)) {
|
||||
const int H = head_size / 2;
|
||||
const dim3 block(H, num_heads, 1);
|
||||
const half2* input2 = reinterpret_cast<const half2*>(input);
|
||||
half2* output2 = reinterpret_cast<half2*>(output);
|
||||
TransposeCtx<half2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else { // this should be an "odd" case. probably not worth catching it in the half2 kernel.
|
||||
const dim3 block(head_size, num_heads, 1);
|
||||
TransposeCtx<half><<<grid, block, 0, stream>>>(head_size, input, output);
|
||||
}
|
||||
|
||||
return CUDA_CALL(cudaPeekAtLastError());
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__global__ void TransposeQKV(const int H, const T* input, T* output) {
|
||||
// Input: BxSx3xNxH
|
||||
// Output: 3xBxNxSxH
|
||||
|
||||
int n = threadIdx.y;
|
||||
int s = blockIdx.x;
|
||||
int b = blockIdx.y;
|
||||
int m = blockIdx.z; // matrix id
|
||||
|
||||
const int num_heads = blockDim.y;
|
||||
|
||||
const int sequence_length = gridDim.x;
|
||||
const int batch_size = gridDim.y;
|
||||
const int NH = num_heads * H;
|
||||
const int NHS = NH * sequence_length;
|
||||
const int in_offset = n * H + m * NH + s * 3 * NH + b * NHS * 3;
|
||||
const int out_offset = s * H + n * sequence_length * H + b * NHS + m * NHS * batch_size;
|
||||
|
||||
const int i = threadIdx.x;
|
||||
if (i < H) {
|
||||
output[out_offset + i] = input[in_offset + i];
|
||||
}
|
||||
}
|
||||
|
||||
bool LaunchTransQkv(cudaStream_t stream,
|
||||
const int sequence_length, const int batch_size, const int head_size, const int num_heads,
|
||||
const float* input, float* output) {
|
||||
const dim3 grid(sequence_length, batch_size, 3);
|
||||
if (0 == (head_size & 1)) {
|
||||
const int H = head_size / 2;
|
||||
const float2* input2 = reinterpret_cast<const float2*>(input);
|
||||
float2* output2 = reinterpret_cast<float2*>(output);
|
||||
const dim3 block(H, num_heads, 1);
|
||||
TransposeQKV<float2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else {
|
||||
const dim3 block(head_size, num_heads, 1);
|
||||
TransposeQKV<float><<<grid, block, 0, stream>>>(head_size, input, output);
|
||||
}
|
||||
return CUDA_CALL(cudaPeekAtLastError());
|
||||
}
|
||||
|
||||
bool LaunchTransQkv(cudaStream_t stream,
|
||||
const int sequence_length, const int batch_size, const int head_size, const int num_heads,
|
||||
const half* input, half* output) {
|
||||
const dim3 grid(sequence_length, batch_size, 3);
|
||||
if (0 == (head_size % 4)) {
|
||||
const int H = head_size / 4;
|
||||
const dim3 block(H, num_heads, 1);
|
||||
const float2* input2 = reinterpret_cast<const float2*>(input);
|
||||
float2* output2 = reinterpret_cast<float2*>(output);
|
||||
TransposeQKV<float2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else if (0 == (head_size & 1)) {
|
||||
const int H = head_size / 2;
|
||||
const dim3 block(H, num_heads, 1);
|
||||
const half2* input2 = reinterpret_cast<const half2*>(input);
|
||||
half2* output2 = reinterpret_cast<half2*>(output);
|
||||
TransposeQKV<half2><<<grid, block, 0, stream>>>(H, input2, output2);
|
||||
} else { // this should be an "odd" case. probably not worth catching it in the half2 kernel..
|
||||
const dim3 block(head_size, num_heads, 1);
|
||||
TransposeQKV<half><<<grid, block, 0, stream>>>(head_size, input, output);
|
||||
}
|
||||
return CUDA_CALL(cudaPeekAtLastError());
|
||||
}
|
||||
|
||||
cublasStatus_t inline CublasGemmStridedBatched(
|
||||
cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb,
|
||||
int m, int n, int k, const float alpha,
|
||||
const float* A, int lda, long long int strideA, const float* B, int ldb, long long int strideB,
|
||||
const float beta, float* C, int ldc, long long int strideC, int batchCount) {
|
||||
return cublasSgemmStridedBatched(
|
||||
handle, transa, transb, m, n, k, &alpha, A, lda, strideA, B, ldb, strideB, &beta, C, ldc, strideC, batchCount);
|
||||
}
|
||||
|
||||
cublasStatus_t inline CublasGemmStridedBatched(
|
||||
cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb,
|
||||
int m, int n, int k, const half alpha,
|
||||
const half* A, int lda, long long int strideA, const half* B, int ldb, long long int strideB,
|
||||
const half beta, half* C, int ldc, long long int strideC, int batchCount) {
|
||||
return cublasHgemmStridedBatched(
|
||||
handle, transa, transb, m, n, k, &alpha, A, lda, strideA, B, ldb, strideB, &beta, C, ldc, strideC, batchCount);
|
||||
}
|
||||
|
||||
struct CublasConfigHelper {
|
||||
cublasPointerMode_t pointer_mode_;
|
||||
cublasMath_t math_mode_;
|
||||
cublasHandle_t cublas_;
|
||||
CublasConfigHelper(cublasHandle_t cublas)
|
||||
: cublas_(cublas) {
|
||||
cublasGetPointerMode(cublas_, &pointer_mode_);
|
||||
cublasGetMathMode(cublas_, &math_mode_);
|
||||
cublasSetPointerMode(cublas_, CUBLAS_POINTER_MODE_HOST);
|
||||
cublasSetMathMode(cublas_, CUBLAS_TENSOR_OP_MATH);
|
||||
}
|
||||
~CublasConfigHelper() {
|
||||
cublasSetMathMode(cublas_, math_mode_);
|
||||
cublasSetPointerMode(cublas_, pointer_mode_);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
bool QkvToContext(
|
||||
cublasHandle_t& cublas, cudaStream_t stream,
|
||||
const int batch_size, const int sequence_length, const int num_heads, const int head_size, const size_t element_size,
|
||||
const T* input, T* output, T* workspace,
|
||||
const int* mask_index) {
|
||||
const size_t bytes = ScratchSize(element_size, batch_size, num_heads, sequence_length);
|
||||
T* scratch1 = workspace;
|
||||
T* scratch2 = scratch1 + (bytes / element_size);
|
||||
T* scratch3 = scratch2 + (bytes / element_size);
|
||||
|
||||
// input should be BxSx3xNxH => scratch3: 3xBxNxSxH
|
||||
if (!LaunchTransQkv(stream, sequence_length, batch_size, head_size, num_heads, input, scratch3)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// now scratch3 has Q, K, V: each has size BxNxSxH
|
||||
const int batches = batch_size * num_heads;
|
||||
const int size_per_batch = sequence_length * head_size;
|
||||
const int total_size = batches * size_per_batch;
|
||||
const int temp_matrix_size = sequence_length * sequence_length;
|
||||
|
||||
const T* q = scratch3;
|
||||
const T* k = q + total_size;
|
||||
const T* v = k + total_size;
|
||||
|
||||
cublasSetStream(cublas, stream);
|
||||
CublasConfigHelper helper(cublas);
|
||||
|
||||
// compute Q*K' (as K'*Q), scaled by 1/sqrt(H) and store in scratch1: BxNxSxS
|
||||
const float rsqrt_head_size = 1.f / sqrt(static_cast<float>(head_size));
|
||||
if (!CUBLAS_CALL(CublasGemmStridedBatched(
|
||||
cublas, CUBLAS_OP_T, CUBLAS_OP_N, sequence_length, sequence_length, head_size, rsqrt_head_size, k, head_size, size_per_batch,
|
||||
q, head_size, size_per_batch, 0.f, scratch1, sequence_length, temp_matrix_size, batches))) {
|
||||
return false;
|
||||
}
|
||||
|
||||
|
||||
// apply softmax and store result P to scratch2: BxNxSxS
|
||||
if (!ComputeMaskedSoftmax<T>(stream, sequence_length, batch_size, num_heads, mask_index, scratch1, scratch2)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// compute P*V (as V*P), and store in scratch3: BxNxSxH
|
||||
if (!CUBLAS_CALL(CublasGemmStridedBatched(
|
||||
cublas, CUBLAS_OP_N, CUBLAS_OP_N, head_size, sequence_length, sequence_length, 1.f, v, head_size, size_per_batch,
|
||||
scratch2, sequence_length, temp_matrix_size, 0.f, scratch3, head_size, size_per_batch, batches))) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// scratch3 is BxNxSxH, transpose to output BxSxNxH
|
||||
return LaunchTransCtx(stream, sequence_length, batch_size, head_size, num_heads, scratch3, output);
|
||||
}
|
||||
|
||||
bool LaunchAttentionKernel(
|
||||
const void* input,
|
||||
const int* mask_index,
|
||||
void* output,
|
||||
const int batch_size,
|
||||
const int sequence_length,
|
||||
const int num_heads,
|
||||
const int head_size,
|
||||
void* workspace,
|
||||
cublasHandle_t& cublas,
|
||||
const size_t element_size) {
|
||||
// use default stream
|
||||
const cudaStream_t stream = nullptr;
|
||||
|
||||
if (element_size == 2) {
|
||||
return QkvToContext(cublas, stream,
|
||||
batch_size, sequence_length, num_heads, head_size, element_size,
|
||||
reinterpret_cast<const half*>(input), reinterpret_cast<half*>(output), reinterpret_cast<half*>(workspace),
|
||||
mask_index);
|
||||
} else {
|
||||
return QkvToContext(cublas, stream,
|
||||
batch_size, sequence_length, num_heads, head_size, element_size,
|
||||
reinterpret_cast<const float*>(input), reinterpret_cast<float*>(output), reinterpret_cast<float*>(workspace),
|
||||
mask_index);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
27
onnxruntime/contrib_ops/cuda/bert/attention_impl.h
Normal file
27
onnxruntime/contrib_ops/cuda/bert/attention_impl.h
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
#include "core/providers/cuda/shared_inc/cuda_utils.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
size_t GetAttentionWorkspaceSize(size_t element_size, int batchsize, int num_heads, int head_size, int sequence_length);
|
||||
|
||||
bool LaunchAttentionKernel(
|
||||
const void* input, // Input tensor
|
||||
const int* mask_index, // Nask index where each element is length of a sequence
|
||||
void* output, // Output tensor
|
||||
int batch_size, // Batch size (B)
|
||||
int sequence_length, // Sequence length (S)
|
||||
int num_heads, // Number of attention heads (N)
|
||||
int head_size, // Hidden layer size per head (H)
|
||||
void* workspace, // Temporary buffer
|
||||
cublasHandle_t& cublas, // Cublas handle
|
||||
const size_t element_size // Element size of input tensor
|
||||
);
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -18,6 +18,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Affine);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Affine);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16, Affine);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, Attention);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, Attention);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ConvTransposeWithDynamicPads);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Crop);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Crop);
|
||||
|
|
@ -51,6 +53,8 @@ void RegisterCudaContribKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Affine)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Affine)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16, Affine)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, Attention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, Attention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, ConvTransposeWithDynamicPads)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Crop)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Crop)>,
|
||||
|
|
|
|||
|
|
@ -192,6 +192,25 @@ void RegisterNchwcSchemas() {
|
|||
.FillUsing(NchwcGlobalPoolOpSchemaGenerator);
|
||||
}
|
||||
|
||||
void RegisterBertSchemas() {
|
||||
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(Attention)
|
||||
.SetDomain(kMSDomain)
|
||||
.SinceVersion(1)
|
||||
.SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL)
|
||||
.SetDoc("Multi-Head Self Attention")
|
||||
.Attr("num_heads", "Number of attention heads", AttributeProto::INT)
|
||||
.Input(0, "input", "3D input tensor with shape (batch_size, sequence_length, hidden_size), hidden_size = num_heads * head_size", "T")
|
||||
.Input(1, "weight", "2D input tensor with shape (hidden_size, 3 * hidden_size)", "T")
|
||||
.Input(2, "bias", "1D input tensor with shape (3 * hidden_size)", "T")
|
||||
.Input(3, "mask_index", "Attention mask index with shape (batch_size)", "M")
|
||||
.Output(0, "output", "3D output tensor with shape (batch_size, sequence_length, hidden_size)", "T")
|
||||
.TypeConstraint("T", {"tensor(float)", "tensor(float16)"}, "Constrain input and output types to float tensors.")
|
||||
.TypeConstraint("M", {"tensor(int32)"}, "Constrain mask index to integer types")
|
||||
.TypeAndShapeInferenceFunction(ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput);
|
||||
|
||||
}
|
||||
|
||||
void RegisterContribSchemas() {
|
||||
// Register removed experimental ops for backward compatibility.
|
||||
// Experimental operators do not have version history. However, RS5 takes bunch of experimental operators
|
||||
|
|
@ -1773,6 +1792,8 @@ Example 4:
|
|||
}
|
||||
});
|
||||
|
||||
RegisterBertSchemas();
|
||||
|
||||
// Register the NCHWc schemas if supported by the platform.
|
||||
if (MlasNchwcGetBlockSize() > 1) {
|
||||
RegisterNchwcSchemas();
|
||||
|
|
|
|||
41
onnxruntime/test/common/cuda_op_test_utils.h
Normal file
41
onnxruntime/test/common/cuda_op_test_utils.h
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "test/util/include/default_providers.h"
|
||||
#ifdef USE_CUDA
|
||||
#include "cuda_runtime_api.h"
|
||||
#endif
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace test {
|
||||
|
||||
inline bool HasCudaEnvironment(int min_cuda_architecture) {
|
||||
if (DefaultCudaExecutionProvider().get() == nullptr) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (min_cuda_architecture == 0) {
|
||||
return true;
|
||||
}
|
||||
|
||||
int cuda_architecture = 0;
|
||||
|
||||
#ifdef USE_CUDA
|
||||
int currentCudaDevice = 0;
|
||||
cudaGetDevice(¤tCudaDevice);
|
||||
cudaDeviceSynchronize();
|
||||
cudaDeviceProp prop;
|
||||
if (cudaSuccess != cudaGetDeviceProperties(&prop, currentCudaDevice)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
cuda_architecture = prop.major * 100 + prop.minor * 10;
|
||||
#endif
|
||||
|
||||
return cuda_architecture >= min_cuda_architecture;
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -43,5 +43,15 @@ inline void Normalize(std::vector<float>& v,
|
|||
std::bind(std::divides<float>(), std::placeholders::_1, stdev));
|
||||
}
|
||||
}
|
||||
|
||||
inline std::vector<MLFloat16> ToFloat16(const std::vector<float>& data) {
|
||||
std::vector<MLFloat16> result;
|
||||
result.reserve(data.size());
|
||||
for (size_t i = 0; i < data.size(); i++) {
|
||||
result.push_back(MLFloat16(math::floatToHalf(data[i])));
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
206
onnxruntime/test/contrib_ops/attention_op_test.cc
Normal file
206
onnxruntime/test/contrib_ops/attention_op_test.cc
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
#include "test/common/tensor_op_test_utils.h"
|
||||
#include "test/common/cuda_op_test_utils.h"
|
||||
#include "test/providers/provider_test_utils.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace test {
|
||||
|
||||
static void RunAttentionTest(
|
||||
const std::vector<float>& input_data, // input: [batch_size, sequence_length, hidden_size]
|
||||
const std::vector<float>& weights_data, // weights: [hidden_size, 3 * hidden_size]
|
||||
const std::vector<float>& bias_data, // bias: [3 * hidden_size]
|
||||
const std::vector<int32_t>& mask_index_data, // mask_index: [batch_size]
|
||||
const std::vector<float>& output_data, // output: [batch_size, sequence_length, hidden_size]
|
||||
int batch_size,
|
||||
int sequence_length,
|
||||
int hidden_size,
|
||||
int number_of_heads,
|
||||
bool use_float16 = false) {
|
||||
int min_cuda_architecture = use_float16 ? 530 : 0;
|
||||
if (HasCudaEnvironment(min_cuda_architecture)) {
|
||||
OpTester tester("Attention", 1, onnxruntime::kMSDomain);
|
||||
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(number_of_heads));
|
||||
|
||||
std::vector<int64_t> input_dims = {batch_size, sequence_length, hidden_size};
|
||||
std::vector<int64_t> weights_dims = {hidden_size, 3 * hidden_size};
|
||||
std::vector<int64_t> bias_dims = {3 * hidden_size};
|
||||
std::vector<int64_t> mask_index_dims = {batch_size};
|
||||
std::vector<int64_t> output_dims = input_dims;
|
||||
|
||||
if (use_float16) {
|
||||
tester.AddInput<MLFloat16>("input", input_dims, ToFloat16(input_data));
|
||||
tester.AddInput<MLFloat16>("weight", weights_dims, ToFloat16(weights_data));
|
||||
tester.AddInput<MLFloat16>("bias", bias_dims, ToFloat16(bias_data));
|
||||
tester.AddInput<int32_t>("mask_index", mask_index_dims, mask_index_data);
|
||||
tester.AddOutput<MLFloat16>("output", output_dims, ToFloat16(output_data));
|
||||
} else {
|
||||
tester.AddInput<float>("input", input_dims, input_data);
|
||||
tester.AddInput<float>("weight", weights_dims, weights_data);
|
||||
tester.AddInput<float>("bias", bias_dims, bias_data);
|
||||
tester.AddInput<int32_t>("mask_index", mask_index_dims, mask_index_data);
|
||||
tester.AddOutput<float>("output", output_dims, output_data);
|
||||
}
|
||||
|
||||
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
|
||||
execution_providers.push_back(DefaultCudaExecutionProvider());
|
||||
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
|
||||
}
|
||||
}
|
||||
|
||||
TEST(AttentionTest, AttentionBatch1) {
|
||||
int batch_size = 1;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
int number_of_heads = 2;
|
||||
|
||||
std::vector<float> input_data = {
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f};
|
||||
|
||||
std::vector<float> weight_data = {
|
||||
0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
|
||||
0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
|
||||
0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
|
||||
0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
|
||||
|
||||
std::vector<float> bias_data = {
|
||||
-0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
|
||||
|
||||
std::vector<int32_t> mask_index_data = {2L};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f,
|
||||
3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f};
|
||||
|
||||
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
|
||||
batch_size, sequence_length, hidden_size, number_of_heads);
|
||||
}
|
||||
|
||||
TEST(AttentionTest, AttentionBatch1_Float16) {
|
||||
int batch_size = 1;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
int number_of_heads = 2;
|
||||
|
||||
std::vector<float> input_data = {
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f};
|
||||
|
||||
std::vector<float> weight_data = {
|
||||
0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
|
||||
0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
|
||||
0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
|
||||
0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
|
||||
|
||||
std::vector<float> bias_data = {
|
||||
-0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
|
||||
|
||||
std::vector<int32_t> mask_index_data = {2L};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
3.154296875, 0.1082763671875, 4.25, 5.6484375,
|
||||
3.970703125, 0.072998046875, 4.25, 5.6484375};
|
||||
|
||||
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
|
||||
batch_size, sequence_length, hidden_size, number_of_heads, true);
|
||||
}
|
||||
|
||||
TEST(AttentionTest, AttentionBatch2) {
|
||||
int batch_size = 2;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
int number_of_heads = 2;
|
||||
|
||||
std::vector<float> input_data = {
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f,
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f};
|
||||
|
||||
std::vector<float> weight_data = {
|
||||
0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
|
||||
0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
|
||||
0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
|
||||
0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
|
||||
|
||||
std::vector<float> bias_data = {
|
||||
-0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
|
||||
|
||||
std::vector<int32_t> mask_index_data = {2L, 2L};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f,
|
||||
3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f,
|
||||
3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f,
|
||||
3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f};
|
||||
|
||||
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
|
||||
batch_size, sequence_length, hidden_size, number_of_heads);
|
||||
}
|
||||
|
||||
TEST(AttentionTest, AttentionMaskPartialSequence) {
|
||||
int batch_size = 1;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
int number_of_heads = 2;
|
||||
|
||||
std::vector<float> input_data = {
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f};
|
||||
|
||||
std::vector<float> weight_data = {
|
||||
0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
|
||||
0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
|
||||
0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
|
||||
0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
|
||||
|
||||
std::vector<float> bias_data = {
|
||||
-0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
|
||||
|
||||
// Test mask_index < sequence_length
|
||||
std::vector<int32_t> mask_index_data = {1L};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
8.6899995803833008f, -0.13000002503395081f, 4.25f, 5.6499996185302734f,
|
||||
8.6899995803833008f, -0.13000002503395081f, 4.2499995231628418f, 5.6499991416931152f};
|
||||
|
||||
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
|
||||
batch_size, sequence_length, hidden_size, number_of_heads);
|
||||
}
|
||||
|
||||
TEST(AttentionTest, AttentionMaskExceedSequence) {
|
||||
int batch_size = 1;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
int number_of_heads = 2;
|
||||
|
||||
std::vector<float> input_data = {
|
||||
0.8f, -0.5f, 0.0f, 1.f,
|
||||
0.5f, 0.2f, 0.3f, -0.6f};
|
||||
|
||||
std::vector<float> weight_data = {
|
||||
0.1f, -0.2f, 0.3f, 1.0f, 1.1f, 0.3f, 0.5f, 0.2f, 0.3f, -0.6f, 1.5f, 2.0f,
|
||||
0.5f, 0.1f, 0.4f, 1.6f, 1.0f, 2.0f, 0.4f, 0.8f, 0.9f, 0.1f, -1.3f, 0.7f,
|
||||
0.3f, 0.2f, 4.0f, 2.2f, 1.6f, 1.1f, 0.7f, 0.2f, 0.4f, 1.0f, 1.2f, 0.5f,
|
||||
0.2f, 0.1f, 0.4f, 1.6f, 2.4f, 3.3f, 2.1f, 4.2f, 8.4f, 0.0f, 2.1f, 3.2f};
|
||||
|
||||
std::vector<float> bias_data = {
|
||||
-0.5f, 0.6f, 1.2f, 2.1f, 0.5f, 0.7f, 0.2f, 1.2f, 0.5f, 0.4f, 0.3f, 1.2f};
|
||||
|
||||
// Test mask_index > sequence_length
|
||||
std::vector<int32_t> mask_index_data = {3L};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
3.1495983600616455f, 0.10843668878078461f, 4.25f, 5.6499996185302734f,
|
||||
3.9696791172027588f, 0.073143675923347473f, 4.2499995231628418f, 5.6499991416931152f};
|
||||
|
||||
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
|
||||
batch_size, sequence_length, hidden_size, number_of_heads);
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
#include "core/framework/execution_provider.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
|
|
|
|||
Loading…
Reference in a new issue