| 13 | |
| 14 | namespace { |
| 15 | deformable_conv::Param create_param( |
| 16 | const Algo::SizeArgs& args, const OprParam& opr_param, cublasHandle_t handle, |
| 17 | cudaStream_t stream) { |
| 18 | deformable_conv::Param p; |
| 19 | auto&& fm = args.filter_meta; |
| 20 | |
| 21 | p.handle = handle; |
| 22 | p.stream = stream; |
| 23 | p.group = fm.group; |
| 24 | p.deformable_group = fm.deformable_group; |
| 25 | p.batch_sz = args.im_layout[0]; |
| 26 | |
| 27 | p.IC = args.im_layout[1]; |
| 28 | p.IH = args.im_layout[2]; |
| 29 | p.IW = args.im_layout[3]; |
| 30 | p.OC = args.out_grad_layout[1]; |
| 31 | p.OH = args.out_grad_layout[2]; |
| 32 | p.OW = args.out_grad_layout[3]; |
| 33 | p.FH = fm.spatial[0]; |
| 34 | p.FW = fm.spatial[1]; |
| 35 | p.PH = opr_param.pad_h; |
| 36 | p.PW = opr_param.pad_w; |
| 37 | p.SH = opr_param.stride_h; |
| 38 | p.SW = opr_param.stride_w; |
| 39 | p.DH = opr_param.dilate_h; |
| 40 | p.DW = opr_param.dilate_w; |
| 41 | |
| 42 | p.icpg = p.IC / p.group; |
| 43 | p.icpdg = p.IC / p.deformable_group; |
| 44 | p.ocpg = p.OC / p.group; |
| 45 | p.ocpdg = p.OC / p.deformable_group; |
| 46 | |
| 47 | return p; |
| 48 | } |
| 49 | |
| 50 | std::pair<TensorLayoutArray, BatchedMatrixMulForward::Param> sub_opr_config( |
| 51 | const DeformableConvForwardImpl::CanonizedFilterMeta& fm, |