| 539 | } |
| 540 | |
| 541 | void convolution::test_conv_config_combinations( |
| 542 | int k_size, Handle* handle, bool test_int8, bool test_backward, bool is_cuda, |
| 543 | ConvEPSGetter eps_getter, bool use_io16xc32) { |
| 544 | Checker<Convolution> checker(handle); |
| 545 | std::unique_ptr<Checker<ConvolutionBackwardData>> checker_bwd_data_ptr; |
| 546 | std::unique_ptr<Checker<ConvolutionBackwardFilter>> checker_bwd_filter_ptr; |
| 547 | if (test_backward) { |
| 548 | checker_bwd_data_ptr.reset( |
| 549 | new std::remove_reference<decltype(*checker_bwd_data_ptr)>::type( |
| 550 | handle)); |
| 551 | checker_bwd_filter_ptr.reset( |
| 552 | new std::remove_reference<decltype(*checker_bwd_filter_ptr)>::type( |
| 553 | handle)); |
| 554 | } |
| 555 | auto&& checker_bwd_data = *checker_bwd_data_ptr; |
| 556 | auto&& checker_bwd_filter = *checker_bwd_filter_ptr; |
| 557 | |
| 558 | #define CONF_BOOL(var) for (int var : {0, 1}) |
| 559 | |
| 560 | std::unordered_set<Convolution::AlgorithmDesc> used_algos; |
| 561 | std::unordered_set<ConvolutionBackwardData::AlgorithmDesc> used_algos_bwd_data; |
| 562 | std::unordered_set<ConvolutionBackwardFilter::AlgorithmDesc> used_algos_bwd_flt; |
| 563 | |
| 564 | using Param = Convolution::Param; |
| 565 | CONF_BOOL(conv) |
| 566 | CONF_BOOL(padding) |
| 567 | CONF_BOOL(stride) |
| 568 | CONF_BOOL(group) |
| 569 | CONF_BOOL(non_square) |
| 570 | CONF_BOOL(dilation) |
| 571 | CONF_BOOL(format) |
| 572 | // dtype: 0: f32; 1: f16; 2: i8x8x16 3: i8x8x32 |
| 573 | for (int dtype = 0; dtype < (test_int8 ? 4 : 2); ++dtype) |
| 574 | for (int ksize : {1, k_size}) { |
| 575 | // When is_cuda is on, test cases where format is NHWC and |
| 576 | // data type is not INT8x8x32 are disabled. |
| 577 | if (is_cuda) { |
| 578 | if (format && dtype != 3) |
| 579 | continue; |
| 580 | } |
| 581 | auto config2str = [&]() -> std::string { |
| 582 | std::ostringstream ostr; |
| 583 | ostr << conv << padding << stride << group << non_square << dilation |
| 584 | << format << dtype << ksize; |
| 585 | return ostr.str(); |
| 586 | }; |
| 587 | auto errmsg = [&](const char* name) { |
| 588 | std::string ret; |
| 589 | ret += "checker failed for algorithm "; |
| 590 | ret += name; |
| 591 | ret += " with conv,padding,stride,group,non_square,dilation,format," |
| 592 | "dtype,ksize="; |
| 593 | ret += config2str(); |
| 594 | return ret; |
| 595 | }; |
| 596 | MEGDNN_MARK_USED_VAR(errmsg); |
| 597 | Param param; |
| 598 | param.mode = |
nothing calls this directly
no test coverage detected