MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / exec

Method exec

dnn/src/rocm/batched_matrix_mul/blas.cpp:20–122  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

18}
19
20void 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();

Callers

nothing calls this directly

Calls 9

get_rocblas_handleMethod · 0.80
concrete_handleFunction · 0.50
paramMethod · 0.45
handleMethod · 0.45
zero_deviceMethod · 0.45
one_deviceMethod · 0.45
raw_ptrMethod · 0.45
one_device_hMethod · 0.45
zero_device_hMethod · 0.45

Tested by

no test coverage detected