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