| 10 | |
| 11 | namespace { |
| 12 | std::pair<TensorLayoutArray, ConvBiasForward::Param> sub_opr_config( |
| 13 | const TensorLayout& src, const TensorLayout& filter, const TensorLayout& dst, |
| 14 | const ConvolutionForwardImpl* opr) { |
| 15 | auto conv_param = opr->param(); |
| 16 | megdnn_assert(src.dtype.category() == DTypeCategory::FLOAT); |
| 17 | // todo check dtype |
| 18 | DType bias_type = src.dtype; |
| 19 | |
| 20 | std::pair<TensorLayoutArray, ConvBiasForward::Param> ret; |
| 21 | ret.second = { |
| 22 | param::ConvBias::NonlineMode::IDENTITY, |
| 23 | conv_param.mode, |
| 24 | conv_param.sparse, |
| 25 | conv_param.format, |
| 26 | conv_param.pad_h, |
| 27 | conv_param.pad_w, |
| 28 | conv_param.stride_h, |
| 29 | conv_param.stride_w, |
| 30 | conv_param.dilate_h, |
| 31 | conv_param.dilate_w, |
| 32 | conv_param.compute_mode}; |
| 33 | // bias cant be null |
| 34 | size_t out_channel = dst.shape[3]; |
| 35 | ret.first.push_back(TensorLayout({1, 1, 1, out_channel}, bias_type)); // bias |
| 36 | ret.first.push_back(TensorLayout({}, dst.dtype)); // z |
| 37 | return ret; |
| 38 | } |
| 39 | |
| 40 | std::pair<TensorLayoutArray, std::unique_ptr<ConvBiasForward>> prepare_sub_opr( |
| 41 | const ConvolutionForwardImpl::AlgoBase::SizeArgs& args) { |
no test coverage detected