| 54 | } |
| 55 | |
| 56 | void LocalShareBackwardDataImpl::AlgoImplicitGemm::exec(const ExecArgs& args) const { |
| 57 | local_share::Param kern_param; |
| 58 | auto&& param = args.opr->param(); |
| 59 | unpack_local_share_params( |
| 60 | args.grad_layout, args.filter_layout, args.diff_layout, param); |
| 61 | kern_param.n = n, kern_param.co = co, kern_param.ci = ci, kern_param.hi = hi, |
| 62 | kern_param.wi = wi, kern_param.ph = ph, kern_param.pw = pw, |
| 63 | kern_param.grp_ho = ho / sgh, kern_param.grp_wo = wo / sgw, kern_param.sgh = sgh, |
| 64 | kern_param.sgw = sgw; |
| 65 | auto&& handle = concrete_handle(args.opr->handle()); |
| 66 | auto&& cublas_hdl = cublas_handle(args.opr->handle()); |
| 67 | auto&& stream = cuda_stream(args.opr->handle()); |
| 68 | |
| 69 | auto one = handle->one_device(); |
| 70 | auto zero = handle->zero_device(); |
| 71 | |
| 72 | local_share_bwd_data::_do_local_share_bwd_data_implicit_gemm( |
| 73 | args.filter_tensor->ptr<dt_float32>(), args.diff_tensor->ptr<dt_float32>(), |
| 74 | args.grad_tensor->ptr<dt_float32>(), |
| 75 | reinterpret_cast<float*>(args.workspace.raw_ptr), fh, fw, sh, sw, |
| 76 | kern_param, cublas_hdl, stream, one, zero); |
| 77 | } |
| 78 | |
| 79 | // vim: syntax=cpp.doxygen |
nothing calls this directly
no test coverage detected