| 12 | namespace { |
| 13 | |
| 14 | std::unique_ptr<ConvBiasForward> prepare_conv_opr( |
| 15 | const MatrixMulForwardImpl::AlgoBase::SizeArgs& args) { |
| 16 | auto conv_bias_opr_ptr = args.opr->handle()->create_operator<ConvBiasForward>(); |
| 17 | |
| 18 | auto conv_param_computemode = |
| 19 | (args.opr->param().compute_mode == param::MatrixMul::ComputeMode::DEFAULT) |
| 20 | ? param::Convolution::ComputeMode::DEFAULT |
| 21 | : param::Convolution::ComputeMode::FLOAT32; |
| 22 | conv_bias_opr_ptr->param() = { |
| 23 | param::ConvBias::NonlineMode::IDENTITY, |
| 24 | param::Convolution::Mode::CROSS_CORRELATION, |
| 25 | param::Convolution::Sparse::DENSE, |
| 26 | param::Convolution::Format::NCHW, |
| 27 | 0, // pad_h |
| 28 | 0, // pad_w |
| 29 | 1, // stride_h |
| 30 | 1, // stride_w |
| 31 | 1, // dilate_h |
| 32 | 1, // dilate_w |
| 33 | conv_param_computemode}; |
| 34 | |
| 35 | return conv_bias_opr_ptr; |
| 36 | } |
| 37 | std::tuple<size_t, size_t, size_t> gen_matrixmul_shape( |
| 38 | const MatrixMulForwardImpl::AlgoBase::SizeArgs& args) { |
| 39 | size_t m, k, n; |
no test coverage detected