| 31 | } |
| 32 | |
| 33 | void MatrixMulForwardImpl::AlgoBlas::exec(const ExecArgs& args) const { |
| 34 | auto m = args.layout_c.shape[0], n = args.layout_c.shape[1]; |
| 35 | auto k = args.layout_a.shape[args.opr->param().transposeA ? 0 : 1]; |
| 36 | auto&& handle = concrete_handle(args.opr->handle()); |
| 37 | auto rocblas_handle_ = handle->get_rocblas_handle(); |
| 38 | |
| 39 | auto sgemm = [&]() { |
| 40 | auto zero = handle->zero_device(); |
| 41 | auto one = handle->one_device(); |
| 42 | rocblas_check(rocblas_sgemm( |
| 43 | rocblas_handle_, |
| 44 | args.opr->param().transposeB ? rocblas_operation_transpose |
| 45 | : rocblas_operation_none, |
| 46 | args.opr->param().transposeA ? rocblas_operation_transpose |
| 47 | : rocblas_operation_none, |
| 48 | n, m, k, one, args.tensor_b.ptr<dt_float32>(), args.layout_b.stride[0], |
| 49 | args.tensor_a.ptr<dt_float32>(), args.layout_a.stride[0], zero, |
| 50 | args.tensor_c.ptr<dt_float32>(), args.layout_c.stride[0])); |
| 51 | }; |
| 52 | |
| 53 | #if !MEGDNN_DISABLE_FLOAT16 |
| 54 | //! used for FLOAT_IO16xC32, not tested |
| 55 | auto gemm_ex = [&]() { |
| 56 | auto zero = handle->zero_device(); |
| 57 | auto one = handle->one_device(); |
| 58 | //! These two arguments for future use, see |
| 59 | //! https://github.com/ROCmSoftwarePlatform/rocBLAS/blob/develop/library/src/blas_ex/rocblas_gemm_ex.cpp |
| 60 | int32_t solution_index = 0; |
| 61 | uint32_t flags = 1; |
| 62 | size_t ws_size = 0; |
| 63 | auto gemm_ex_err = rocblas_gemm_ex( |
| 64 | rocblas_handle_, |
| 65 | args.opr->param().transposeB ? rocblas_operation_transpose |
| 66 | : rocblas_operation_none, |
| 67 | args.opr->param().transposeA ? rocblas_operation_transpose |
| 68 | : rocblas_operation_none, |
| 69 | n, m, k, one, args.tensor_b.raw_ptr(), rocblas_datatype_f16_r, |
| 70 | args.layout_b.stride[0], args.tensor_a.raw_ptr(), |
| 71 | rocblas_datatype_f16_r, args.layout_a.stride[0], zero, |
| 72 | args.tensor_c.raw_ptr(), rocblas_datatype_f16_r, |
| 73 | args.layout_c.stride[0], args.tensor_c.raw_ptr(), |
| 74 | rocblas_datatype_f16_r, args.layout_c.stride[0], rocblas_datatype_f32_r, |
| 75 | rocblas_gemm_algo_standard, solution_index, flags, &ws_size, nullptr); |
| 76 | rocblas_check(gemm_ex_err); |
| 77 | MEGDNN_MARK_USED_VAR(ws_size); |
| 78 | }; |
| 79 | |
| 80 | auto hgemm = [&]() { |
| 81 | auto one_half = handle->one_device_h(); |
| 82 | auto zero_half = handle->zero_device_h(); |
| 83 | auto hgemm_err = rocblas_hgemm( |
| 84 | rocblas_handle_, |
| 85 | args.opr->param().transposeB ? rocblas_operation_transpose |
| 86 | : rocblas_operation_none, |
| 87 | args.opr->param().transposeA ? rocblas_operation_transpose |
| 88 | : rocblas_operation_none, |
| 89 | n, m, k, reinterpret_cast<const rocblas_half*>(one_half), |
| 90 | static_cast<const rocblas_half*>(args.tensor_b.raw_ptr()), |
nothing calls this directly
no test coverage detected