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

Method exec

dnn/src/cuda/warp_perspective/backward_data.cpp:30–103  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 9

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
workspaceMethod · 0.45
raw_ptrMethod · 0.45

Tested by

no test coverage detected