| 130 | } |
| 131 | |
| 132 | void ConvBiasForwardImpl::AlgoGroupConvGeneral::exec(const ExecArgs& args) const { |
| 133 | auto bundle = get_workspace_bundle(args.workspace.raw_ptr, args); |
| 134 | TensorND conv_dst_tensor = *args.dst_tensor; |
| 135 | if (args.dst_layout->dtype.enumv() != args.bias_layout->dtype.enumv()) { |
| 136 | conv_dst_tensor = TensorND{ |
| 137 | bundle.get(bundle.nr_workspace() - 1), args.dst_tensor->layout}; |
| 138 | conv_dst_tensor.layout.dtype = DType(); |
| 139 | args.opr->check_or_deduce_dtype_fwd( |
| 140 | args.src_layout->dtype, args.filter_layout->dtype, |
| 141 | conv_dst_tensor.layout.dtype); |
| 142 | } |
| 143 | { |
| 144 | auto sub_args = args; |
| 145 | sub_args.dst_tensor = &conv_dst_tensor; |
| 146 | sub_args.dst_layout = &conv_dst_tensor.layout; |
| 147 | |
| 148 | auto config = prepare_sub_opr(sub_args); |
| 149 | TensorND tsrc{args.src_tensor->raw_ptr(), config.first[0]}; |
| 150 | TensorND tfilter{args.filter_tensor->raw_ptr(), config.first[1]}; |
| 151 | TensorND tbias{args.bias_tensor->raw_ptr(), config.first[2]}; |
| 152 | TensorND tz{args.z_tensor->raw_ptr(), config.first[3]}; |
| 153 | TensorND tdst{conv_dst_tensor.raw_ptr(), config.first[4]}; |
| 154 | |
| 155 | size_t c_pos; |
| 156 | if (args.filter_meta.format == Param::Format::NCHW || |
| 157 | args.filter_meta.format == Param::Format::NCHW4) { |
| 158 | c_pos = 1; |
| 159 | } else { |
| 160 | megdnn_assert( |
| 161 | args.filter_meta.format == Param::Format::NHWC, |
| 162 | "invalid conv format"); |
| 163 | c_pos = 3; |
| 164 | } |
| 165 | |
| 166 | auto grp = args.filter_meta.group; |
| 167 | |
| 168 | auto&& fm = args.filter_meta; |
| 169 | auto strd_src = tsrc.layout.stride[c_pos] * fm.icpg * tsrc.layout.dtype.size(), |
| 170 | strd_dst = tdst.layout.stride[c_pos] * fm.ocpg * tdst.layout.dtype.size(), |
| 171 | strd_flt = fm.icpg * fm.ocpg * fm.spatial[0] * fm.spatial[1] * |
| 172 | tfilter.layout.dtype.size(); |
| 173 | if (args.filter_meta.format == Param::Format::NCHW4) { |
| 174 | strd_src >>= 2; |
| 175 | strd_dst >>= 2; |
| 176 | } |
| 177 | for (uint32_t g = 0; g < grp; ++g) { |
| 178 | config.second->exec( |
| 179 | tsrc, tfilter, tbias, tz, tdst, nullptr, bundle.get_workspace(0)); |
| 180 | incr_refp(tsrc.get_ref_ptr(), strd_src); |
| 181 | incr_refp(tdst.get_ref_ptr(), strd_dst); |
| 182 | incr_refp(tfilter.get_ref_ptr(), strd_flt); |
| 183 | } |
| 184 | } |
| 185 | handle_bias_and_nonlinear( |
| 186 | args.handle, args.nonlinear_mode, &conv_dst_tensor, args.dst_tensor, |
| 187 | args.bias_tensor); |
| 188 | } |
| 189 |
nothing calls this directly
no test coverage detected