| 7 | using Mode = PoolingForward::Param::Mode; |
| 8 | |
| 9 | TEST(ARMCOMMON, PoolingNchw44Int8) { |
| 10 | #ifdef __aarch64__ |
| 11 | Checker<PoolingForward> checker(Arch::ARM64, 0); |
| 12 | #else |
| 13 | Checker<PoolingForward> checker(Arch::ARMV7, 0); |
| 14 | #endif |
| 15 | PoolingForward::Param param; |
| 16 | UniformIntRNG rng(-127, 127); |
| 17 | checker.set_rng(0, &rng); |
| 18 | checker.set_kernel_symbol("ArmCommon_FilterX_modeX_.*"); |
| 19 | param.format = param::Pooling::Format::NCHW44; |
| 20 | |
| 21 | auto run = [&](std::list<megdnn::DType> dtypes, std::list<Mode> modes) { |
| 22 | for (auto dtype : dtypes) |
| 23 | for (auto mode : modes) |
| 24 | for (size_t window : {2, 3, 4, 5}) |
| 25 | for (size_t stride : {1, 2}) |
| 26 | for (size_t pad : {size_t(0), size_t(window / 2)}) { |
| 27 | param.mode = mode; |
| 28 | checker.set_dtype(0, dtype).set_dtype(1, dtype); |
| 29 | param.pad_h = pad; |
| 30 | param.pad_w = pad; |
| 31 | param.window_h = window; |
| 32 | param.window_w = window; |
| 33 | param.stride_h = stride; |
| 34 | param.stride_w = stride; |
| 35 | checker.set_param(param); |
| 36 | checker.set_before_exec_callback( |
| 37 | megdnn::test::AlgoChecker<PoolingForward>( |
| 38 | ("ARM_POOLING_FILTER" + |
| 39 | std::to_string(window) + |
| 40 | "_MODEX_STRIDEX_NCHW44") |
| 41 | .c_str())); |
| 42 | checker.execs({{2, 3, 5, 5, 4}, {}}); |
| 43 | checker.execs({{1, 2, 7, 7, 4}, {}}); |
| 44 | } |
| 45 | }; |
| 46 | |
| 47 | run({dtype::Int8()}, {Mode::MAX}); |
| 48 | run({dtype::QuantizedS8(0.35f), dtype::QuantizedS8(1.6f)}, |
| 49 | {Mode::AVERAGE, Mode::MAX}); |
| 50 | } |