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

Method test_conv_config_combinations

dnn/test/common/convolution.cpp:541–741  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

539}
540
541void 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 =

Callers

nothing calls this directly

Calls 14

serialize_write_podFunction · 0.85
set_dtypeMethod · 0.80
prev_succMethod · 0.80
sqrtFunction · 0.50
resetMethod · 0.45
strMethod · 0.45
oprMethod · 0.45
paramMethod · 0.45
deduce_layoutMethod · 0.45
insertMethod · 0.45
handleMethod · 0.45

Tested by

no test coverage detected