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

Method exec

dnn/src/cuda/warp_perspective/backward_mat.cpp:31–109  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

29}
30
31void WarpPerspectiveBackwardMatImpl::exec(
32 _megdnn_tensor_in ssrc, _megdnn_tensor_in smat, _megdnn_tensor_in smat_idx,
33 _megdnn_tensor_in sdiff, _megdnn_tensor_out sgrad,
34 _megdnn_workspace sworkspace) {
35 check_exec(
36 ssrc.layout, smat.layout, smat_idx.layout, sdiff.layout, sgrad.layout,
37 sworkspace.size);
38 TensorND src = ssrc;
39 TensorND mat = smat;
40 TensorND diff = sdiff;
41 TensorND grad = sgrad;
42 TensorND mat_idx = smat_idx;
43 auto bundle = get_workspace_bundle(
44 sworkspace.raw_ptr, ssrc.layout, smat.layout, sdiff.layout, sgrad.layout);
45 auto ctypecvt = CompTypeCvter<dtype::BFloat16, dtype::Float32>(
46 concrete_handle(this->handle()), &bundle);
47 if (ssrc.layout.dtype.enumv() == DTypeTrait<dtype::BFloat16>::enumv) {
48 ctypecvt.src_to_comp_type(ssrc, src)
49 .src_to_comp_type(smat, mat)
50 .src_to_comp_type(sdiff, diff)
51 .src_to_comp_type(sgrad, grad);
52 }
53 {
54 auto stream = cuda_stream(this->handle());
55 auto N = src.layout.shape[0], C = src.layout.shape[1], IH = src.layout.shape[2],
56 IW = src.layout.shape[3], OH = diff.layout.shape[2],
57 OW = diff.layout.shape[3];
58 int* midx_ptr = nullptr;
59 if (mat_idx.raw_ptr()) {
60 megdnn_assert(mat_idx.layout.ndim == 1);
61 N = mat_idx.layout.shape[0];
62 midx_ptr = mat_idx.ptr<int>();
63 } else {
64 megdnn_assert(mat_idx.layout.ndim == 0);
65 }
66
67 auto bval = param().border_val;
68 auto bmode = warp_perspective::get_bmode(param().bmode);
69
70 size_t batch_x_channel_size = N * C;
71 size_t max_batch_x_channel = max_batch_x_channel_size();
72 if (batch_x_channel_size <= max_batch_x_channel) {
73 warp_perspective::backward_mat_proxy(
74 src.ptr<dt_float32>(), mat.ptr<dt_float32>(), midx_ptr,
75 diff.ptr<dt_float32>(), grad.ptr<dt_float32>(), N, C, IH, IW, OH,
76 OW, bval, bmode, stream);
77 } else {
78 dt_float32* src_ptr = src.ptr<dt_float32>();
79 dt_float32* mat_ptr = mat.ptr<dt_float32>();
80 dt_float32* diff_ptr = diff.ptr<dt_float32>();
81 dt_float32* grad_ptr = grad.ptr<dt_float32>();
82 size_t max_batch_size = max_batch_x_channel / C;
83 while (N > 0) {
84 size_t curr_batch_size = N > max_batch_size ? max_batch_size : N;
85 warp_perspective::backward_mat_proxy(
86 src_ptr, mat_ptr, midx_ptr, diff_ptr, grad_ptr, curr_batch_size,
87 C, IH, IW, OH, OW, bval, bmode, stream);
88

Callers

nothing calls this directly

Calls 8

cuda_streamFunction · 0.85
get_bmodeFunction · 0.70
get_workspace_bundleFunction · 0.50
concrete_handleFunction · 0.50
paramFunction · 0.50
handleMethod · 0.45
enumvMethod · 0.45
raw_ptrMethod · 0.45

Tested by

no test coverage detected