Using cublasGemmBatchedEx/cublasGemmStridedBatchedEx for training (#4731)

* use cublas extenstion API for fp16

* Using cublasGemmBatchedEx/cublasGemmStridedBatchedEx for training

To avoid accuracy, the accumulation needs to be done in FP32 for training.

Co-authored-by: Weixing Zhang <wezhan@microsoft.com>
This commit is contained in:
Weixing Zhang 2020-08-14 02:12:14 -07:00 committed by GitHub
parent ec36c793e8
commit afa89566d7
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
8 changed files with 235 additions and 48 deletions

View file

@ -87,6 +87,7 @@ Status Attention<T>::ComputeInternal(OpKernelContext* context) const {
size_t workSpaceSize = GetAttentionWorkspaceSize(element_size, batch_size, num_heads_, head_size, sequence_length, past_sequence_length);
auto temp_buffer = GetScratchBuffer<void>(workSpaceSize);
if (!LaunchAttentionKernel(
device_prop,
reinterpret_cast<const CudaT*>(gemm_buffer.get()),
nullptr == mask_index ? nullptr : mask_index->template Data<int>(),
nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()),

View file

@ -635,7 +635,7 @@ cublasStatus_t inline CublasGemmStridedBatched(
template <typename T>
bool QkvToContext(
cublasHandle_t& cublas, cudaStream_t stream,
const cudaDeviceProp& prop, 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 std::vector<int64_t>* mask_index_dims,
@ -661,7 +661,7 @@ bool QkvToContext(
const T* v = k + total_size;
cublasSetStream(cublas, stream);
CublasMathModeSetter helper(cublas, CUBLAS_TENSOR_OP_MATH);
CublasMathModeSetter helper(prop, cublas, CUBLAS_TENSOR_OP_MATH);
// Concat past (2xBxNxS'xH) to present (2xBxNxS*xH):
// past_k (BxNxS'xH) + k (BxNxSxH) => present_k (BxNxS*xH)
@ -720,6 +720,7 @@ bool QkvToContext(
}
bool LaunchAttentionKernel(
const cudaDeviceProp& prop,
const void* input,
const int* mask_index,
const std::vector<int64_t>* mask_index_dims,
@ -739,13 +740,13 @@ bool LaunchAttentionKernel(
const cudaStream_t stream = nullptr;
if (element_size == 2) {
return QkvToContext(cublas, stream,
return QkvToContext(prop, 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, mask_index_dims, is_unidirectional,
past_sequence_length, reinterpret_cast<const half*>(past), reinterpret_cast<half*>(present));
} else {
return QkvToContext(cublas, stream,
return QkvToContext(prop, 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, mask_index_dims, is_unidirectional,

View file

@ -16,6 +16,7 @@ size_t GetAttentionWorkspaceSize(
int past_sequence_length);
bool LaunchAttentionKernel(
const cudaDeviceProp& prop, // Device Properties
const void* input, // Input tensor
const int* mask_index, // Attention mask raw data or index (end position of each sequence, or end positions and start positions). NULL means no mask.
const std::vector<int64_t>* mask_index_dims, // Mask index shape

View file

@ -171,6 +171,7 @@ Status QAttention<T, int8_t>::ComputeInternal(OpKernelContext* context) const {
size_t workSpaceSize = GetAttentionWorkspaceSize(element_size, batch_size, num_heads_, head_size, sequence_length, past_sequence_length);
auto temp_buffer = GetScratchBuffer<void>(workSpaceSize);
if (!LaunchAttentionKernel(
GetDeviceProp(),
reinterpret_cast<const CudaT*>(gemm_buffer.get()),
nullptr == mask_index ? nullptr : mask_index->template Data<int>(),
nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()),

View file

@ -225,15 +225,20 @@ inline bool CalculateFdmStrides(gsl::span<fast_divmod> p, const std::vector<int6
class CublasMathModeSetter {
public:
CublasMathModeSetter(cublasHandle_t handle, cublasMath_t mode) : handle_(handle) {
cublasGetMathMode(handle, &mode_);
cublasSetMathMode(handle, mode);
CublasMathModeSetter(const cudaDeviceProp& prop,cublasHandle_t handle, cublasMath_t mode) : prop_(prop), handle_(handle) {
if (prop_.major >= 7) {
cublasGetMathMode(handle, &mode_);
cublasSetMathMode(handle, mode);
}
}
~CublasMathModeSetter() {
cublasSetMathMode(handle_, mode_);
if (prop_.major >= 7) {
cublasSetMathMode(handle_, mode_);
}
}
private:
const cudaDeviceProp& prop_;
cublasHandle_t handle_;
cublasMath_t mode_;
};

View file

@ -60,7 +60,8 @@ Status MatMul(const T* input_1_data, const T* input_2_data, T* output_data,
reinterpret_cast<CudaT*>(output_data),
static_cast<int>(N),
static_cast<int>(output_stride),
static_cast<int>(num_batches)));
static_cast<int>(num_batches),
static_cast<EinsumCudaAssets*>(einsum_cuda_assets)->cuda_ep_->GetDeviceProp()));
return Status::OK();
}

View file

@ -143,7 +143,8 @@ Status MatMul<T>::ComputeInternal(OpKernelContext* ctx) const {
reinterpret_cast<CudaT*>(Y->template MutableData<T>()),
ldc,
stride_C,
static_cast<int>(batch_count)));
static_cast<int>(batch_count),
device_prop));
return Status::OK();
}
@ -175,7 +176,8 @@ Status MatMul<T>::ComputeInternal(OpKernelContext* ctx) const {
&zero,
output_arrays.GpuPtr(),
ldc,
static_cast<int>(helper.OutputOffsets().size())));
static_cast<int>(helper.OutputOffsets().size()),
device_prop));
return Status::OK();
}

View file

@ -18,49 +18,175 @@
// Generalize library calls to be use in template functions
// gemm
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb,
int m, int n, int k, const float* alpha, const float* A, int lda,
const float* B, int ldb, const float* beta, float* C, int ldc,
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const float* alpha,
const float* A, int lda,
const float* B, int ldb,
const float* beta,
float* C, int ldc,
const cudaDeviceProp& /*prop*/) {
return cublasSgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc);
return cublasSgemm(handle,
transa,
transb,
m, n, k,
alpha,
A, lda,
B, ldb,
beta,
C, ldc);
}
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb,
int m, int n, int k, const double* alpha, const double* A, int lda,
const double* B, int ldb, const double* beta, double* C, int ldc,
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const double* alpha,
const double* A, int lda,
const double* B, int ldb,
const double* beta,
double* C, int ldc,
const cudaDeviceProp& /*prop*/) {
return cublasDgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc);
return cublasDgemm(handle,
transa,
transb,
m, n, k,
alpha,
A, lda,
B, ldb,
beta,
C, ldc);
}
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb,
int m, int n, int k, const half* alpha, const half* A, int lda,
const half* B, int ldb, const half* beta, half* C, int ldc,
inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const half* alpha,
const half* A, int lda,
const half* B, int ldb,
const half* beta,
half* C, int ldc,
const cudaDeviceProp& prop) {
#ifndef ENABLE_TRAINING
// This does true FP16 computation which is slow for non-Volta GPUs
if (prop.major >= 7) {
onnxruntime::cuda::CublasMathModeSetter math_mode_setter(handle, CUBLAS_TENSOR_OP_MATH);
return cublasHgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc);
}
#else
ORT_UNUSED_PARAMETER(prop);
#endif
onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH);
//This does pseudo FP16 computation (input/output in fp16, computation in fp32)
#ifdef ENABLE_TRAINING
float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(alpha));
float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(beta));
cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH);
return cublasGemmEx(handle, transa, transb, m, n, k, &h_a, A, CUDA_R_16F, lda, B, CUDA_R_16F, ldb, &h_b, C, CUDA_R_16F, ldc, CUDA_R_32F, CUBLAS_GEMM_DFALT);
// accumulating in FP32
return cublasGemmEx(handle,
transa,
transb,
m, n, k,
&h_a,
A, CUDA_R_16F, lda,
B, CUDA_R_16F, ldb,
&h_b,
C, CUDA_R_16F, ldc,
CUDA_R_32F,
CUBLAS_GEMM_DEFAULT);
#else
// accumulating in FP16
return cublasHgemm(handle,
transa,
transb,
m, n, k,
alpha,
A, lda,
B, ldb,
beta,
C, ldc);
#endif
}
// batched gemm
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const float* alpha, const float* Aarray[], int lda, const float* Barray[], int ldb, const float* beta, float* Carray[], int ldc, int batch_count) {
return cublasSgemmBatched(handle, transa, transb, m, n, k, alpha, Aarray, lda, Barray, ldb, beta, Carray, ldc, batch_count);
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const float* alpha,
const float* Aarray[], int lda,
const float* Barray[], int ldb,
const float* beta,
float* Carray[], int ldc,
int batch_count,
const cudaDeviceProp& /*prop*/) {
return cublasSgemmBatched(handle,
transa,
transb,
m, n, k,
alpha,
Aarray, lda,
Barray, ldb,
beta,
Carray, ldc,
batch_count);
}
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const double* alpha, const double* Aarray[], int lda, const double* Barray[], int ldb, const double* beta, double* Carray[], int ldc, int batch_count) {
return cublasDgemmBatched(handle, transa, transb, m, n, k, alpha, Aarray, lda, Barray, ldb, beta, Carray, ldc, batch_count);
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const double* alpha,
const double* Aarray[], int lda,
const double* Barray[], int ldb,
const double* beta,
double* Carray[], int ldc,
int batch_count,
const cudaDeviceProp& /*prop*/) {
return cublasDgemmBatched(handle,
transa,
transb,
m, n, k,
alpha,
Aarray, lda,
Barray, ldb,
beta,
Carray, ldc,
batch_count);
}
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const half* alpha, const half* Aarray[], int lda, const half* Barray[], int ldb, const half* beta, half* Carray[], int ldc, int batch_count) {
cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH);
return cublasHgemmBatched(handle, transa, transb, m, n, k, alpha, (const __half**)Aarray, lda, (const __half**)Barray, ldb, beta, (__half**)Carray, ldc, batch_count);
inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int m, int n, int k,
const half* alpha,
const half* Aarray[], int lda,
const half* Barray[], int ldb,
const half* beta,
half* Carray[], int ldc,
int batch_count,
const cudaDeviceProp& prop) {
onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH);
#ifdef ENABLE_TRAINING
float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(alpha));
float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(beta));
// accumulating in FP32
return cublasGemmBatchedEx(handle,
transa,
transb,
m, n, k,
&h_a,
(const void**)Aarray, CUDA_R_16F, lda,
(const void**)Barray, CUDA_R_16F, ldb,
&h_b,
(void**)Carray, CUDA_R_16F, ldc,
batch_count,
CUDA_R_32F,
CUBLAS_GEMM_DEFAULT);
#else
// accumulating in FP16
return cublasHgemmBatched(handle,
transa,
transb,
m, n, k,
alpha,
(const __half**)Aarray, lda,
(const __half**)Barray, ldb,
beta,
(__half**)Carray, ldc,
batch_count);
#endif
}
// strided batched gemm
@ -76,8 +202,18 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle,
const float* beta,
float* C, int ldc,
long long int strideC,
int batch_count) {
return cublasSgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count);
int batch_count,
const cudaDeviceProp& /*prop*/) {
return cublasSgemmStridedBatched(handle,
transa,
transb,
m, n, k,
alpha,
A, lda, strideA,
B, ldb, strideB,
beta,
C, ldc, strideC,
batch_count);
}
inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle,
@ -92,8 +228,18 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle,
const double* beta,
double* C, int ldc,
long long int strideC,
int batch_count) {
return cublasDgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count);
int batch_count,
const cudaDeviceProp& /*prop*/) {
return cublasDgemmStridedBatched(handle,
transa,
transb,
m, n, k,
alpha,
A, lda, strideA,
B, ldb, strideB,
beta,
C, ldc, strideC,
batch_count);
}
inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle,
@ -108,9 +254,38 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle,
const __half* beta,
__half* C, int ldc,
long long int strideC,
int batch_count) {
cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH);
return cublasHgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count);
int batch_count,
const cudaDeviceProp& prop) {
onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH);
#ifdef ENABLE_TRAINING
float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(alpha));
float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast<const uint16_t*>(beta));
// accumulating in FP32
return cublasGemmStridedBatchedEx(handle,
transa,
transb,
m, n, k,
&h_a,
A, CUDA_R_16F, lda, strideA,
B, CUDA_R_16F, ldb, strideB,
&h_b,
C, CUDA_R_16F, ldc, strideC,
batch_count,
CUDA_R_32F,
CUBLAS_GEMM_DEFAULT);
#else
// accumulating in FP16
return cublasHgemmStridedBatched(handle,
transa,
transb,
m, n, k,
alpha,
A, lda, strideA,
B, ldb, strideB,
beta,
C, ldc, strideC,
batch_count);
#endif
}
// axpy