| 18 | } |
| 19 | |
| 20 | void BatchedMatrixMulForwardImpl::AlgoBlas::exec(const ExecArgs& args) const { |
| 21 | auto batch = args.layout_a.shape[0]; |
| 22 | auto m = args.layout_c.shape[1], n = args.layout_c.shape[2]; |
| 23 | auto k = args.layout_a.shape[args.opr->param().transposeA ? 1 : 2]; |
| 24 | auto&& handle = concrete_handle(args.opr->handle()); |
| 25 | auto rocblas_handle_ = handle->get_rocblas_handle(); |
| 26 | |
| 27 | auto sgemm = [&]() { |
| 28 | auto zero = handle->zero_device(); |
| 29 | auto one = handle->one_device(); |
| 30 | rocblas_check(rocblas_sgemm_strided_batched( |
| 31 | rocblas_handle_, |
| 32 | args.opr->param().transposeB ? rocblas_operation_transpose |
| 33 | : rocblas_operation_none, |
| 34 | args.opr->param().transposeA ? rocblas_operation_transpose |
| 35 | : rocblas_operation_none, |
| 36 | n, m, k, one, args.tensor_b.ptr<dt_float32>(), |
| 37 | (rocblas_int)(args.layout_b.stride[1]), |
| 38 | (rocblas_int)(args.layout_b.stride[0]), args.tensor_a.ptr<dt_float32>(), |
| 39 | (rocblas_int)(args.layout_a.stride[1]), |
| 40 | (rocblas_int)(args.layout_a.stride[0]), zero, |
| 41 | args.tensor_c.ptr<dt_float32>(), (rocblas_int)(args.layout_c.stride[1]), |
| 42 | (rocblas_int)(args.layout_c.stride[0]), (rocblas_int)(batch))); |
| 43 | }; |
| 44 | |
| 45 | #if !MEGDNN_DISABLE_FLOAT16 |
| 46 | //! used for FLOAT_IO16xC32, not tested |
| 47 | auto gemm_ex = [&]() { |
| 48 | auto zero = handle->zero_device(); |
| 49 | auto one = handle->one_device(); |
| 50 | //! These two arguments for future use, see |
| 51 | //! https://github.com/ROCmSoftwarePlatform/rocBLAS/blob/develop/library/src/blas_ex/rocblas_gemm_ex.cpp |
| 52 | int32_t solution_index = 0; |
| 53 | uint32_t flags = 1; |
| 54 | size_t ws_size = 0; |
| 55 | |
| 56 | rocblas_check(rocblas_gemm_strided_batched_ex( |
| 57 | rocblas_handle_, |
| 58 | args.opr->param().transposeB ? rocblas_operation_transpose |
| 59 | : rocblas_operation_none, |
| 60 | args.opr->param().transposeA ? rocblas_operation_transpose |
| 61 | : rocblas_operation_none, |
| 62 | n, m, k, one, args.tensor_b.raw_ptr(), rocblas_datatype_i8_r, |
| 63 | args.layout_b.stride[1], args.layout_b.stride[0], |
| 64 | args.tensor_a.raw_ptr(), rocblas_datatype_i8_r, args.layout_a.stride[1], |
| 65 | args.layout_a.stride[0], zero, args.tensor_c.raw_ptr(), |
| 66 | rocblas_datatype_i32_r, args.layout_c.stride[1], |
| 67 | args.layout_c.stride[0], args.tensor_c.raw_ptr(), |
| 68 | rocblas_datatype_i32_r, args.layout_c.stride[1], |
| 69 | args.layout_c.stride[0], batch, rocblas_datatype_i32_r, |
| 70 | rocblas_gemm_algo_standard, solution_index, flags, &ws_size, nullptr)); |
| 71 | |
| 72 | MEGDNN_MARK_USED_VAR(ws_size); |
| 73 | }; |
| 74 | |
| 75 | auto hgemm = [&]() { |
| 76 | auto one_half = handle->one_device_h(); |
| 77 | auto zero_half = handle->zero_device_h(); |
nothing calls this directly
no test coverage detected