| 29 | } |
| 30 | |
| 31 | void 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 |
nothing calls this directly
no test coverage detected