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

Method exec

dnn/src/rocm/matrix_mul/blas.cpp:33–146  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

31}
32
33void 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()),

Callers

nothing calls this directly

Calls 12

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
zero_device_i32Method · 0.45
one_device_i32Method · 0.45

Tested by

no test coverage detected