MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / get_channel_wise_args

Function get_channel_wise_args

dnn/test/arm_common/conv_bias_multi_thread.cpp:139–215  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

137}
138
139std::vector<conv_bias::TestArg> get_channel_wise_args(
140 std::vector<size_t> kernel, size_t stride, bool no_bias, bool no_nonlinemode,
141 bool no_full_bias, bool support_relu) {
142 using namespace conv_bias;
143 using Param = param::ConvBias;
144 using NLMode = param::ConvBias::NonlineMode;
145 std::vector<TestArg> args;
146
147 auto pack = [&](size_t n, size_t group, size_t w, size_t h, size_t kernel,
148 size_t stride, NLMode nlmode, bool pad) {
149 Param param;
150 param.stride_h = stride;
151 param.stride_w = stride;
152 if (pad) {
153 param.pad_h = kernel / 2;
154 param.pad_w = kernel / 2;
155 } else {
156 param.pad_h = 0;
157 param.pad_w = 0;
158 }
159 param.nonlineMode = nlmode;
160 param.format = param::ConvBias::Format::NCHW;
161 param.sparse = param::ConvBias::Sparse::GROUP;
162
163 args.emplace_back(
164 param, TensorShape{n, group, h, w},
165 TensorShape{group, 1, 1, kernel, kernel}, TensorShape{});
166 if (!no_bias) {
167 args.emplace_back(
168 param, TensorShape{n, group, h, w},
169 TensorShape{group, 1, 1, kernel, kernel},
170 TensorShape{1, group, 1, 1});
171 }
172 if (!no_full_bias) {
173 args.emplace_back(
174 param, TensorShape{n, group, h, w},
175 TensorShape{group, 1, 1, kernel, kernel},
176 TensorShape{
177 n, group, (h + 2 * param.pad_w - kernel) / stride + 1,
178 (w + 2 * param.pad_w - kernel) / stride + 1});
179 }
180 };
181
182 std::vector<NLMode> nonlinemode = {NLMode::IDENTITY};
183 if (!no_nonlinemode) {
184 nonlinemode.emplace_back(NLMode::RELU);
185 nonlinemode.emplace_back(NLMode::H_SWISH);
186 } else if (support_relu) {
187 nonlinemode.emplace_back(NLMode::RELU);
188 }
189
190 for (size_t n : {1, 2}) {
191 for (auto nlmode : nonlinemode) {
192 for (bool pad : {true}) {
193 for (size_t group : {1, 3, 7}) {
194 for (size_t size : {4, 6, 7, 9, 16, 20, 32, 55}) {
195 for (size_t kern : kernel) {
196 pack(n, group, size, size, kern, stride, nlmode, pad);

Callers 1

TEST_FFunction · 0.85

Calls 2

packFunction · 0.85
emplace_backMethod · 0.80

Tested by

no test coverage detected