| 4 | using namespace mgb; |
| 5 | |
| 6 | SymbolVar Network::add_conv( |
| 7 | SymbolVar f, size_t output_channels, KernSize kern_size, DType out_dtype, |
| 8 | bool has_relu, Stride stride, Padding padding) { |
| 9 | static int weight_idx = 0; |
| 10 | static int bias_idx = 0; |
| 11 | |
| 12 | size_t input_channels = f.node()->shape()[1]; |
| 13 | auto weight = add_cvar( |
| 14 | ssprintf("w%d", weight_idx).c_str(), |
| 15 | {output_channels, input_channels, kern_size[0], kern_size[1]}); |
| 16 | auto bias = add_cvar(ssprintf("b%d", bias_idx).c_str(), {1, output_channels, 1, 1}); |
| 17 | if (out_dtype.category() == DTypeCategory::QUANTIZED) { |
| 18 | weight = add_type_cvt(weight, out_dtype); |
| 19 | bias = add_type_cvt(bias, dtype::QuantizedS32{1.f}); |
| 20 | } |
| 21 | opr::ConvBias::Param param; |
| 22 | param.stride_h = stride[0], param.stride_w = stride[1]; |
| 23 | param.pad_h = padding[0], param.pad_w = padding[1]; |
| 24 | if (has_relu) { |
| 25 | param.nonlineMode = opr::ConvBias::Param::NonlineMode::RELU; |
| 26 | } else { |
| 27 | param.nonlineMode = opr::ConvBias::Param::NonlineMode::IDENTITY; |
| 28 | } |
| 29 | |
| 30 | SymbolVar conv; |
| 31 | if (out_dtype.category() == DTypeCategory::QUANTIZED) { |
| 32 | conv = opr::ConvBias::make( |
| 33 | f, weight, bias, param, {}, OperatorNodeConfig{out_dtype}); |
| 34 | } else { |
| 35 | conv = opr::ConvBias::make(f, weight, bias, param, {}); |
| 36 | } |
| 37 | weight_idx++; |
| 38 | bias_idx++; |
| 39 | return conv; |
| 40 | } |
| 41 | |
| 42 | SymbolVar Network::add_group_conv( |
| 43 | SymbolVar f, size_t output_channels, size_t groups, KernSize kern_size, |
no test coverage detected