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

Function TEST

src/opr/test/dnn/region_restricted_convolution.cpp:24–95  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

22using namespace mgb;
23
24TEST(TestOprDNN, REGIONCONV_FWD_CPU_WRAPPER) {
25 using Checker = AutoOprChecker<4, 1>;
26 megdnn::RegionRestrictedConvolution::Param param;
27 param.sparse = opr::RegionRestrictedConvolution::Param::Sparse::DENSE;
28
29 auto make_graph = [&](const Checker::SymInpArray& inputs) -> Checker::SymOutArray {
30 return {opr::RegionRestrictedConvolutionForward::make(
31 inputs[0], inputs[1], inputs[2], inputs[3], param)};
32 };
33
34 Checker::RunOptions option;
35 option.numdiff_eps = 0.1;
36 option.numdiff_max_err = 1e-2;
37
38 auto mask_gen = [&](HostTensorND& src) {
39 HostTensorGenerator<dtype::Int32, RandomDistribution::CONSTANT> gen(1);
40 src = *gen(src.shape(), src.comp_node());
41 };
42 auto float_gen = [&](HostTensorND& src) {
43 HostTensorGenerator<dtype::Float32, RandomDistribution::GAUSSIAN> gen;
44 src = *gen(src.shape(), src.comp_node());
45 };
46
47 auto fwd = [&](Checker::NumOutArray& dest, Checker::NumInpArray inp) {
48 auto opr =
49 megdnn_naive_handle()
50 ->create_operator<megdnn::RegionRestrictedConvolutionForward>();
51 opr->param() = param;
52 TensorLayout dest_layout;
53 opr->deduce_layout(
54 inp[0]->layout(), inp[1]->layout(), inp[2]->layout(), inp[3]->layout(),
55 dest_layout);
56 std::vector<dt_byte> workspace(opr->get_workspace_in_bytes(
57 inp[0]->layout(), inp[1]->layout(), inp[2]->layout(), inp[3]->layout(),
58 dest_layout));
59 dest[0].dtype(inp[0]->dtype())
60 .comp_node(inp[0]->comp_node())
61 .resize(dest_layout);
62 opr->exec(
63 inp[0]->as_megdnn(), inp[1]->as_megdnn(), inp[2]->as_megdnn(),
64 inp[3]->as_megdnn(), dest[0].as_megdnn(),
65 {workspace.data(), workspace.size()});
66 };
67
68 Checker(make_graph, fwd, CompNode::load("cpu0"))
69 .set_input_dtype(0, dtype::Float32())
70 .set_input_dtype(1, dtype::Float32())
71 .set_input_dtype(2, dtype::Int32())
72 .set_input_dtype(3, dtype::Int32())
73 .set_input_generator(0, float_gen)
74 .set_input_generator(1, float_gen)
75 .set_input_generator(2, mask_gen)
76 .set_input_generator(3, mask_gen)
77 .set_input_allow_grad(2, false)
78 .set_input_allow_grad(3, false)
79 // {n,ic,ih,iw}, {oc,ic,fh,fw}, {n,ih,iw}, {n,oh,ow}
80 .run({TensorShape{1, 2, 2, 2}, TensorShape{1, 2, 2, 2},
81 TensorShape{1, 2, 2}, TensorShape{1, 1, 1}},

Callers

nothing calls this directly

Calls 15

resizeMethod · 0.80
as_megdnnMethod · 0.80
makeFunction · 0.50
genFunction · 0.50
CheckerClass · 0.50
loadFunction · 0.50
shapeMethod · 0.45
comp_nodeMethod · 0.45
paramMethod · 0.45
deduce_layoutMethod · 0.45
layoutMethod · 0.45

Tested by

no test coverage detected