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

Method exec

dnn/src/cuda/local/backward_filter.cpp:22–46  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

20namespace cuda {
21
22void LocalBackwardFilterImpl::exec(
23 _megdnn_tensor_in src, _megdnn_tensor_in diff, _megdnn_tensor_out grad,
24 _megdnn_workspace workspace) {
25 check_exec(src.layout, diff.layout, grad.layout, workspace.size);
26 megdnn_assert(param().mode == Mode::CROSS_CORRELATION);
27 auto N = src.layout.shape[0], IC = src.layout.shape[1], IH = src.layout.shape[2],
28 IW = src.layout.shape[3];
29 auto OC = diff.layout.shape[1], OH = diff.layout.shape[2],
30 OW = diff.layout.shape[3];
31 auto FH = grad.layout.shape[3], FW = grad.layout.shape[4];
32 auto handle = concrete_handle(this->handle());
33 auto stream = cuda_stream(this->handle());
34 auto cublas = cublas_handle(this->handle());
35 auto one = handle->one_device();
36 auto zero = handle->zero_device();
37 if (use_cuda_convnet(src.layout, diff.layout, grad.layout)) {
38 local::backward_filter_proxy_convnet(
39 src.ptr<dt_float32>(), diff.ptr<dt_float32>(), grad.ptr<dt_float32>(),
40 reinterpret_cast<float*>(workspace.raw_ptr), N, IC, IH, IW, OC, OH, OW,
41 FH, FW, IC * IH * IW, OC * OH * OW, param().pad_h, param().pad_w,
42 param().stride_h, param().stride_w, cublas, stream, one, zero);
43 } else {
44 local::boom_backward_filter();
45 }
46}
47
48size_t LocalBackwardFilterImpl::get_workspace_in_bytes(
49 const TensorLayout& src, const TensorLayout& diff, const TensorLayout& grad) {

Callers

nothing calls this directly

Calls 8

cuda_streamFunction · 0.85
cublas_handleFunction · 0.85
boom_backward_filterFunction · 0.85
paramFunction · 0.50
concrete_handleFunction · 0.50
handleMethod · 0.45
one_deviceMethod · 0.45
zero_deviceMethod · 0.45

Tested by

no test coverage detected