| 120 | } |
| 121 | |
| 122 | void Algo::exec(const ExecArgs& args) const { |
| 123 | auto&& opr = args.opr; |
| 124 | auto&& param = opr->param(); |
| 125 | auto&& handle = concrete_handle(opr->handle()); |
| 126 | |
| 127 | auto p = create_param(args, param, handle->cublas_handle(), handle->stream()); |
| 128 | |
| 129 | auto bundle = get_bundle(args); |
| 130 | bundle.set(args.workspace.raw_ptr); |
| 131 | |
| 132 | const float* dev_im = args.im_tensor.ptr<float>(); |
| 133 | const float* dev_offset = args.offset_tensor.ptr<float>(); |
| 134 | const float* dev_mask = args.mask_tensor.ptr<float>(); |
| 135 | float* dev_out_grad = args.out_grad_tensor.ptr<float>(); |
| 136 | float* dev_filter_grad = args.filter_grad_tensor.ptr<float>(); |
| 137 | |
| 138 | float* col_ws = static_cast<float*>(bundle.get(0)); |
| 139 | float* out_grad_ws = static_cast<float*>(bundle.get(1)); |
| 140 | void* bmm_ws = bundle.get(2); |
| 141 | |
| 142 | // im2col |
| 143 | deformable_conv::im2col(dev_im, dev_offset, dev_mask, col_ws, p); |
| 144 | // relayout |
| 145 | auto&& dt = args.im_layout.dtype; |
| 146 | size_t dim0 = p.batch_sz, dim1 = p.OC, dim2 = p.OH * p.OW; |
| 147 | TensorLayout C2l({dim0, dim1, dim2}, dt), C3l = C2l; |
| 148 | C3l.stride[0] = dim2; |
| 149 | C3l.stride[1] = dim0 * dim2; |
| 150 | C3l.stride[2] = 1; |
| 151 | TensorND C2(dev_out_grad, C2l); |
| 152 | TensorND C3(out_grad_ws, C3l); |
| 153 | |
| 154 | args.handle->relayout_opr()->exec(C2, C3); |
| 155 | // matmul |
| 156 | auto config = prepare_sub_opr(args); |
| 157 | |
| 158 | TensorND A(static_cast<void*>(out_grad_ws), config.first[0]), |
| 159 | B(static_cast<void*>(col_ws), config.first[1]), |
| 160 | C(static_cast<void*>(dev_filter_grad), config.first[2]); |
| 161 | |
| 162 | size_t bmm_ws_size = bundle.get_size(2); |
| 163 | config.second->exec( |
| 164 | A, B, C, Workspace(static_cast<megdnn::dt_byte*>(bmm_ws), bmm_ws_size)); |
| 165 | } |
| 166 | // vim: syntax=cpp.doxygen |
nothing calls this directly
no test coverage detected