MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / DoBlasGemmStridedBatched

Method DoBlasGemmStridedBatched

tensorflow/stream_executor/cuda/cuda_blas.cc:2711–2772  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

2709}
2710
2711bool CUDABlas::DoBlasGemmStridedBatched(
2712 Stream *stream, blas::Transpose transa, blas::Transpose transb, uint64 m,
2713 uint64 n, uint64 k, float alpha, const DeviceMemory<Eigen::half> &a,
2714 int lda, int64 stride_a, const DeviceMemory<Eigen::half> &b, int ldb,
2715 int64 stride_b, float beta, DeviceMemory<Eigen::half> *c, int ldc,
2716 int64 stride_c, int batch_count) {
2717 bool use_tensor_ops = false;
2718#if CUDA_VERSION >= 9000
2719 int cc_major, cc_minor;
2720 if (stream->parent()->GetDeviceDescription().cuda_compute_capability(
2721 &cc_major, &cc_minor)) {
2722 // GPUs < sm_70 don't support tensor ops.
2723 if (cc_major >= 7) {
2724 use_tensor_ops = true;
2725 }
2726#if CUDA_VERSION >= 9010
2727 if (cc_major >= 5) {
2728 cublasGemmAlgo_t algo =
2729 (use_tensor_ops ? CUBLAS_GEMM_DFALT_TENSOR_OP : CUBLAS_GEMM_DFALT);
2730 cublasStatus_t (*bgemm) (cublasHandle_t, cublasOperation_t, cublasOperation_t,
2731 int, int, int, const void*, const void*, cudaDataType,
2732 int, long long int, const void*, cudaDataType, int,
2733 long long int, const void*, void*, cudaDataType, int,
2734 long long int, int, cudaDataType,
2735 cublasGemmAlgo_t algo) = cublasGemmStridedBatchedEx;
2736 bool ok = DoBlasInternalImpl<Eigen::half>(
2737 bgemm, stream, true /* = pointer_mode_host */,
2738 true /* = err_on_failure */, CUDABlasTranspose(transa),
2739 CUDABlasTranspose(transb), m, n, k, &alpha, GpuMemory(a), CUDA_R_16F,
2740 lda, stride_a, GpuMemory(b), CUDA_R_16F, ldb, stride_b, &beta,
2741 GpuMemoryMutable(c), CUDA_R_16F, ldc, stride_c, batch_count, CUDA_R_32F,
2742 algo);
2743 if (ok) {
2744 return true;
2745 }
2746 LOG(ERROR) << "failed BLAS call, see log for details";
2747 return false;
2748 }
2749#endif
2750 }
2751#endif
2752 // Either CUDA_VERSION < 9.1 or SM < 5.0. Fall back to a loop.
2753 for (int batch = 0; batch < batch_count; ++batch) {
2754 const auto *a_matrix =
2755 reinterpret_cast<const __half *>(GpuMemory(a) + batch * stride_a);
2756 const auto *b_matrix =
2757 reinterpret_cast<const __half *>(GpuMemory(b) + batch * stride_b);
2758 auto *c_matrix =
2759 reinterpret_cast<__half *>(GpuMemoryMutable(c) + batch * stride_c);
2760 bool ok = DoBlasInternalImpl<Eigen::half>(
2761 cublasSgemmEx, stream, true /* = pointer_mode_host */,
2762 true /* = err_on_failure= */, CUDABlasTranspose(transa),
2763 CUDABlasTranspose(transb), m, n, k, &alpha, a_matrix, SE_CUDA_DATA_HALF,
2764 lda, b_matrix, SE_CUDA_DATA_HALF, ldb, &beta, c_matrix,
2765 SE_CUDA_DATA_HALF, ldc);
2766 if (!ok) {
2767 LOG(ERROR) << "failed BLAS call, see log for details";
2768 return false;

Callers

nothing calls this directly

Calls 6

CUDABlasTransposeFunction · 0.85
GpuMemoryFunction · 0.85
GpuMemoryMutableFunction · 0.85
GpuComplexFunction · 0.85
parentMethod · 0.45

Tested by

no test coverage detected