| 2709 | } |
| 2710 | |
| 2711 | bool 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; |
nothing calls this directly
no test coverage detected