Skip to content
Merged
3 changes: 2 additions & 1 deletion docs/OperatorKernels.md
Original file line number Diff line number Diff line change
Expand Up @@ -810,7 +810,8 @@ The **OpSet Version** column uses the following notation:
|||[9, 12]|**T1** = tensor(double), tensor(float), tensor(float16)<br/> **T2** = tensor(bool)|
|LRN|*in* X:**T**<br> *out* Y:**T**|13+|**T** = tensor(double), tensor(float), tensor(float16)|
|||[1, 12]|**T** = tensor(double), tensor(float), tensor(float16)|
|LSTM|*in* X:**T**<br> *in* W:**T**<br> *in* R:**T**<br> *in* B:**T**<br> *in* sequence_lens:**T1**<br> *in* initial_h:**T**<br> *in* initial_c:**T**<br> *in* P:**T**<br> *out* Y:**T**<br> *out* Y_h:**T**<br> *out* Y_c:**T**|14+|**T** = tensor(double), tensor(float), tensor(float16)<br/> **T1** = tensor(int32)|
|LSTM|*in* X:**T**<br> *in* W:**T**<br> *in* R:**T**<br> *in* B:**T**<br> *in* sequence_lens:**T1**<br> *in* initial_h:**T**<br> *in* initial_c:**T**<br> *in* P:**T**<br> *out* Y:**T**<br> *out* Y_h:**T**<br> *out* Y_c:**T**|22+|**T** = tensor(double), tensor(float), tensor(float16)<br/> **T1** = tensor(int32)|
|||[14, 21]|**T** = tensor(double), tensor(float), tensor(float16)<br/> **T1** = tensor(int32)|
|||[7, 13]|**T** = tensor(double), tensor(float), tensor(float16)<br/> **T1** = tensor(int32)|
|LayerNormalization|*in* X:**T**<br> *in* Scale:**T**<br> *in* B:**T**<br> *out* Y:**T**<br> *out* Mean:**U**<br> *out* InvStdDev:**U**<br><br>or<br><br>*in* X:**T**<br> *in* Scale:**V**<br> *in* B:**V**<br> *out* Y:**V**<br> *out* Mean:**U**<br> *out* InvStdDev:**U**|17+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)<br/> **U** = tensor(float)|
|||[1, 16]|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)<br/> **U** = tensor(double), tensor(float)<br/> **V** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)|
Expand Down
18 changes: 12 additions & 6 deletions onnxruntime/core/providers/cuda/cuda_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1409,9 +1409,9 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kO
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, GRU);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, MLFloat16, GRU);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 18, Identity);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, float, LSTM);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, double, LSTM);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, MLFloat16, LSTM);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, float, LSTM);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, LSTM);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, MLFloat16, LSTM);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 18, Reshape);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, float, RNN);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, RNN);
Expand Down Expand Up @@ -1702,6 +1702,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, double, HardSwish);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, HardSwish);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, BFloat16, HardSwish);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, float, LSTM);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, double, LSTM);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, LSTM);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomNormal);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomNormalLike);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomUniform);
Expand Down Expand Up @@ -2681,9 +2684,9 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, float, GRU)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, GRU)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, MLFloat16, GRU)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, float, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, double, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, MLFloat16, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, float, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, MLFloat16, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 18, Reshape)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, float, RNN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, 21, double, RNN)>,
Expand Down Expand Up @@ -2975,6 +2978,9 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, double, HardSwish)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, HardSwish)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, BFloat16, HardSwish)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, float, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, double, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, MLFloat16, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomNormal)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomNormalLike)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 22, RandomUniform)>,
Expand Down
80 changes: 49 additions & 31 deletions onnxruntime/core/providers/cuda/rnn/cudnn_rnn_base.cc
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
#include "core/providers/shared_library/provider_api.h"
#include "cudnn_rnn_base.h"
#include "rnn_impl.h"
#include "core/common/safeint.h"

namespace onnxruntime {
namespace cuda {
Expand All @@ -16,7 +17,7 @@ Status CudnnRnnBase<T>::SetWeightBias(const cudnnHandle_t handle,
const void* reorganized_w_data,
const int lin_layer_id,
const T* pos,
int& offset,
size_t& offset,
bool is_matrix,
cudaStream_t cuda_stream) const {
int numDims;
Expand All @@ -37,12 +38,12 @@ Status CudnnRnnBase<T>::SetWeightBias(const cudnnHandle_t handle,
is_matrix ? tensor_desc_matrix : tensor_desc_bias, 3, &dt, &numDims, matDims.data(), strideA.data()));

mem_offset = is_matrix ? mem_offset_matrix : mem_offset_bias;
int count = matDims[0] * matDims[1] * matDims[2];
size_t count = SafeInt<size_t>(matDims[0]) * matDims[1] * matDims[2];

if (strideA[0] != count) {
if (static_cast<size_t>(strideA[0]) != count) {
return ORT_MAKE_STATUS(ONNXRUNTIME, StatusCode::INVALID_ARGUMENT, "Stride is not packed");
}
CUDA_CALL_THROW(cudaMemcpyAsync(mem_offset, pos + offset, count * sizeof(T), cudaMemcpyDeviceToDevice, cuda_stream));
CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(mem_offset, pos + offset, count * sizeof(T), cudaMemcpyDeviceToDevice, cuda_stream));

offset += count;

Expand All @@ -57,9 +58,9 @@ Status CudnnRnnBase<T>::SetCudnnRnnWeightBias(const cudnnHandle_t cudnn_handle,
const T* R_data,
const T* B_data,
cudaStream_t cuda_stream) const {
int w_offset = 0;
int r_offset = 0;
int bias_offset = 0;
size_t w_offset = 0;
size_t r_offset = 0;
size_t bias_offset = 0;
for (int layer = 0; layer < RNN_NUM_LAYERS * num_directions_; ++layer) {
for (size_t idx = 0; idx < W_lin_layer_id_.size(); ++idx) {
ORT_RETURN_IF_ERROR(SetWeightBias(
Expand Down Expand Up @@ -92,6 +93,12 @@ Status CudnnRnnBase<T>::ReorganizeWeights(const Tensor* W, const Tensor* R, cons
void* alloc_stream, cudaStream_t cuda_stream,
cudnnHandle_t cudnn_handle) const {
typedef typename ToCudaType<T>::MappedType CudaT;
ORT_RETURN_IF(W->Shape().NumDimensions() != 3,
"Weight W must be 3-D [num_directions, hidden_size, input_size], got rank ",
W->Shape().NumDimensions());
ORT_RETURN_IF(R->Shape().NumDimensions() != 3,
"Recurrence R must be 3-D [num_directions, hidden_size, hidden_size], got rank ",
R->Shape().NumDimensions());
int64_t input_size = W->Shape()[2];
// RNN W[num_directions_, hidden_size_, input_size]
// RNN R[num_directions_, hidden_size_, hidden_size_]
Expand All @@ -103,12 +110,12 @@ Status CudnnRnnBase<T>::ReorganizeWeights(const Tensor* W, const Tensor* R, cons
// LSTM R[num_directions_, 4*hidden_size_, hidden_size_]
// LSTM B[num_directions_, 8*hidden_size_]
size_t number = W_lin_layer_id_.size();
int64_t w_size = num_directions_ * (number * hidden_size_ * (input_size + hidden_size_ + 2));
int64_t w_size = SafeInt<int64_t>(num_directions_) * number * hidden_size_ * (input_size + hidden_size_ + 2);
TensorShapeVector dims_w({w_size, 1, 1});
ORT_RETURN_IF_ERROR(target_w_desc.Set(dims_w, CudnnTensor::GetDataType<CudaT>()));

// Prepare the weight data
reorganized_w_data_size_in_bytes = w_size * sizeof(T);
reorganized_w_data_size_in_bytes = SafeInt<size_t>(w_size) * sizeof(T);
reorganized_w_data = GetScratchBuffer<void>(reorganized_w_data_size_in_bytes, alloc_stream);

// In many cases, this allocation is bigger than needed, leaving part of
Expand Down Expand Up @@ -142,6 +149,8 @@ Status CudnnRnnBase<T>::CacheCudnnRnnWeights(const OpKernelInfo& info) {
bool has_bias = B != nullptr;

if (get_W && get_R) {
ORT_RETURN_IF(W->Shape().NumDimensions() != 3,
"Constant W must be 3-D, got rank ", W->Shape().NumDimensions());
CudnnRNN tmp_rnn_desc;
auto proj_size = hidden_size_;
ORT_RETURN_IF_ERROR(tmp_rnn_desc.Set(W->Shape()[2], // input_size
Expand All @@ -163,7 +172,7 @@ Status CudnnRnnBase<T>::CacheCudnnRnnWeights(const OpKernelInfo& info) {
w_data_cache_size_in_bytes_, w_data_cache_, w_desc_cache_,
tmp_rnn_desc, nullptr, nullptr, DefaultCudnnHandle()));
}
cudaStreamSynchronize(nullptr);
CUDA_RETURN_IF_ERROR(cudaStreamSynchronize(nullptr));

weight_cached_ = true;
}
Expand All @@ -177,6 +186,9 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
// inputs
const Tensor* X = ctx->Input<Tensor>(RNN_Input_Index::X); // inputs. [seq_length, batch_size, input_size]
ORT_ENFORCE(nullptr != X);
ORT_RETURN_IF(X->Shape().NumDimensions() != 3,
"Input X must be 3-D [seq_length, batch_size, input_size], got rank ",
X->Shape().NumDimensions());

// optional inputs
// [batch_size]
Expand All @@ -187,6 +199,10 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
if (rnn_mode_ == CUDNN_LSTM) {
// initial cell. [num_directions_, batch_size, hidden_size_]
initial_c = ctx->Input<Tensor>(RNN_Input_Index::initial_c);
// cuDNN LSTM does not support peephole weights (ONNX input P at index 7)
const Tensor* P = ctx->Input<Tensor>(7);
ORT_RETURN_IF(P != nullptr,
"CUDA LSTM does not support peephole weights (input P). Use CPU EP instead.");
}

size_t proj_size = hidden_size_;
Expand All @@ -197,7 +213,7 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
// we thread a single input as sequence_lens of length 1, require to expand to [batch_size]?
std::vector<int32_t> sequence_lengths_temp;
if (!sequence_lens) {
sequence_lengths_temp.resize(batch_size, gsl::narrow_cast<int32_t>(seq_length));
sequence_lengths_temp.resize(batch_size, gsl::narrow<int32_t>(seq_length));
}

const int32_t* sequence_lens_data = (sequence_lens == nullptr)
Expand All @@ -214,10 +230,10 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {

// 0-len sequences are not supported by cuDNN.
// Replace them by sequences of len 1 and mask them out with SetZeroSequences
for (int i = 0; i < batch_size; ++i) {
for (int64_t i = 0; i < batch_size; ++i) {
if (0 == sequence_lens_data[i]) {
seq_len_array[i] = 1;
zero_seq_index_cache[zero_seq_count] = i;
zero_seq_index_cache[zero_seq_count] = gsl::narrow<int32_t>(i);
++zero_seq_count;
} else {
seq_len_array[i] = sequence_lens_data[i];
Expand Down Expand Up @@ -257,15 +273,15 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
const T* x_data = X->Data<T>();
if (reverse_) {
// reverse input data
x_reversed_data = GetScratchBuffer<T>(seq_length * batch_size * input_size, GetComputeStream(ctx));
x_reversed_data = GetScratchBuffer<T>(SafeInt<int64_t>(seq_length) * batch_size * input_size, GetComputeStream(ctx));
ReverseBySequence(Stream(ctx),
gsl::narrow_cast<int32_t>(seq_length),
gsl::narrow<int32_t>(seq_length),
sequence_lens_buffer.GpuPtr(),
gsl::narrow_cast<int32_t>(batch_size),
gsl::narrow_cast<int32_t>(input_size),
gsl::narrow<int32_t>(batch_size),
gsl::narrow<int32_t>(input_size),
reinterpret_cast<const CudaT*>(x_data),
reinterpret_cast<CudaT*>(x_reversed_data.get()),
seq_length * batch_size * input_size);
SafeInt<int64_t>(seq_length) * batch_size * input_size);
}

const T* x_data_input = reverse_ ? x_reversed_data.get() : x_data;
Expand All @@ -274,7 +290,7 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
const T* cx_data = (initial_c == nullptr) ? nullptr : initial_c->Data<T>();
T* y_h_data = (Y_h == nullptr) ? nullptr : Y_h->MutableData<T>();
T* y_c_data = (Y_c == nullptr) ? nullptr : Y_c->MutableData<T>();
int64_t output_size = seq_length * num_directions_ * batch_size * hidden_size_;
int64_t output_size = SafeInt<int64_t>(seq_length) * num_directions_ * batch_size * hidden_size_;
T* y_data = nullptr;
IAllocatorUniquePtr<T> y_alloc_data;
if (Y != nullptr) {
Expand Down Expand Up @@ -357,7 +373,8 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
// Mask on output for 0 sequence batches
if (zero_seq_count > 0) {
// Mask on output for 0 sequence batches
SetZeroSequences(zero_seq_count, zero_seq_index_cache, y_data, y_h_data, y_c_data, GetComputeStream(ctx), Stream(ctx));
SetZeroSequences(gsl::span<const int32_t>(zero_seq_index_cache.data(), zero_seq_count),
y_data, y_h_data, y_c_data, GetComputeStream(ctx), Stream(ctx));
}
return Status::OK();
}
Expand All @@ -369,26 +386,26 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {
if (reverse_) {
// reverse output data
ReverseBySequence(Stream(ctx),
gsl::narrow_cast<int32_t>(seq_length),
gsl::narrow<int32_t>(seq_length),
sequence_lens_buffer.GpuPtr(),
gsl::narrow_cast<int32_t>(batch_size),
gsl::narrow_cast<int32_t>(hidden_size_),
gsl::narrow<int32_t>(batch_size),
gsl::narrow<int32_t>(hidden_size_),
reinterpret_cast<CudaT*>(y_data),
reinterpret_cast<CudaT*>(y_reorganized_data.get()),
output_size);
} else {
ReorderBidirectionalDataInSequence(Stream(ctx),
gsl::narrow_cast<int32_t>(seq_length),
gsl::narrow_cast<int32_t>(batch_size),
gsl::narrow_cast<int32_t>(hidden_size_),
gsl::narrow<int32_t>(seq_length),
gsl::narrow<int32_t>(batch_size),
gsl::narrow<int32_t>(hidden_size_),
reinterpret_cast<CudaT*>(y_data),
reinterpret_cast<CudaT*>(y_reorganized_data.get()),
output_size);
}

if (Y != nullptr) {
// User specified this optional output, so need to copy the reversed data to original place
CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(y_data, y_reorganized_data.get(), output_size * sizeof(T),
CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(y_data, y_reorganized_data.get(), SafeInt<size_t>(output_size) * sizeof(T),
cudaMemcpyDeviceToDevice, Stream(ctx)));
} else {
y_data = y_reorganized_data.get();
Expand All @@ -397,26 +414,27 @@ Status CudnnRnnBase<T>::ComputeInternal(OpKernelContext* ctx) const {

// Mask on output for 0 sequence batches
if (zero_seq_count > 0) {
SetZeroSequences(zero_seq_count, zero_seq_index_cache, y_data, y_h_data, y_c_data, GetComputeStream(ctx), Stream(ctx));
SetZeroSequences(gsl::span<const int32_t>(zero_seq_index_cache.data(), zero_seq_count),
y_data, y_h_data, y_c_data, GetComputeStream(ctx), Stream(ctx));
}

return Status::OK();
}

template <typename T>
void CudnnRnnBase<T>::SetZeroSequences(const int64_t zero_seq_index_cache_size,
const std::vector<int32_t> zero_seq_index_cache,
void CudnnRnnBase<T>::SetZeroSequences(gsl::span<const int32_t> zero_seq_index_cache,
T* y_data,
T* y_h_data,
T* y_c_data,
void* alloc_stream, cudaStream_t cuda_stream) const {
typedef typename ToCudaType<T>::MappedType CudaT;
const int64_t zero_seq_index_cache_size = static_cast<int64_t>(zero_seq_index_cache.size());
CudaAsyncBuffer<int32_t> zero_seq_index_cache_async_buffer(this, zero_seq_index_cache_size);
memcpy(zero_seq_index_cache_async_buffer.CpuPtr(), zero_seq_index_cache.data(),
zero_seq_index_cache_size * sizeof(int32_t));
ORT_THROW_IF_ERROR(zero_seq_index_cache_async_buffer.CopyToGpu(alloc_stream));
MaskZeroSequences(cuda_stream,
gsl::narrow_cast<int32_t>(hidden_size_),
gsl::narrow<int32_t>(hidden_size_),
reinterpret_cast<CudaT*>(y_data),
reinterpret_cast<CudaT*>(y_h_data),
reinterpret_cast<CudaT*>(y_c_data),
Expand Down
Loading
Loading