| 120 | } |
| 121 | |
| 122 | void Algo::exec(const ExecArgs& args) const { |
| 123 | auto&& opr = args.opr; |
| 124 | auto&& handle = concrete_handle(opr->handle()); |
| 125 | auto&& param = opr->param(); |
| 126 | auto p = create_param(args, param, handle->cublas_handle(), handle->stream()); |
| 127 | auto bundle = get_bundle(args); |
| 128 | bundle.set(args.workspace.raw_ptr); |
| 129 | |
| 130 | float* dev_im = args.im_tensor.ptr<float>(); |
| 131 | float* dev_filter = args.filter_tensor.ptr<float>(); |
| 132 | float* dev_offset = args.offset_tensor.ptr<float>(); |
| 133 | float* dev_mask = args.mask_tensor.ptr<float>(); |
| 134 | float* dev_out_grad = args.out_grad_tensor.ptr<float>(); |
| 135 | |
| 136 | float* dev_im_grad = args.im_grad_tensor.ptr<float>(); |
| 137 | float* dev_offset_grad = args.offset_grad_tensor.ptr<float>(); |
| 138 | float* dev_mask_grad = args.mask_grad_tensor.ptr<float>(); |
| 139 | |
| 140 | void* bmm_ws = bundle.get(0); |
| 141 | float* result_ws = static_cast<float*>(bundle.get(1)); |
| 142 | float* relayout_ws1 = static_cast<float*>(bundle.get(2)); |
| 143 | |
| 144 | // clear out grad |
| 145 | { |
| 146 | size_t im_sz = p.batch_sz * p.IC * p.IH * p.IW * sizeof(float); |
| 147 | size_t offset_sz = p.batch_sz * 2 * p.deformable_group * p.FH * p.FW * p.OH * |
| 148 | p.OW * sizeof(float); |
| 149 | size_t mask_sz = p.batch_sz * p.deformable_group * p.FH * p.FW * p.OH * p.OW * |
| 150 | sizeof(float); |
| 151 | |
| 152 | cudaMemsetAsync(dev_im_grad, 0, im_sz, p.stream); |
| 153 | cudaMemsetAsync(dev_offset_grad, 0, offset_sz, p.stream); |
| 154 | cudaMemsetAsync(dev_mask_grad, 0, mask_sz, p.stream); |
| 155 | } |
| 156 | |
| 157 | // relayout out_grad to [oc, N, OH, OW] |
| 158 | { |
| 159 | auto&& dt = args.im_layout.dtype; |
| 160 | size_t dim0 = p.batch_sz, dim1 = p.OC, dim2 = p.OH * p.OW; |
| 161 | TensorLayout C2l({dim0, dim1, dim2}, dt), C3l = C2l; |
| 162 | C3l.stride[0] = dim2; |
| 163 | C3l.stride[1] = dim0 * dim2; |
| 164 | C3l.stride[2] = 1; |
| 165 | TensorND C2(dev_out_grad, C2l); |
| 166 | TensorND C3(relayout_ws1, C3l); |
| 167 | |
| 168 | args.handle->relayout_opr()->exec(C2, C3); |
| 169 | } |
| 170 | // matmul [g, icpg, FH, FW, ocpg] * [g, ocpg, N, OH, OW] => |
| 171 | // => [g, icpg, FH, FW, N, OH, OW] |
| 172 | { |
| 173 | auto config = prepare_sub_opr(args); |
| 174 | |
| 175 | TensorND A(static_cast<void*>(dev_filter), config.first[0]), |
| 176 | B(static_cast<void*>(relayout_ws1), config.first[1]), |
| 177 | C(static_cast<void*>(result_ws), config.first[2]); |
| 178 | |
| 179 | size_t bmm_ws_size = bundle.get_size(0); |
nothing calls this directly
no test coverage detected