| 8 | using namespace mgb; |
| 9 | |
| 10 | TEST(TestOprDNN, SlidingWindowTranspose) { |
| 11 | using Checker = AutoOprChecker<1, 1>; |
| 12 | |
| 13 | opr::SlidingWindowTranspose::Param param; |
| 14 | param.pad_h = 1; |
| 15 | param.pad_w = 2; |
| 16 | param.stride_w = 2; |
| 17 | param.window_h = 4; |
| 18 | param.dilate_h = 2; |
| 19 | unsigned long ih = 16, iw = 15; |
| 20 | unsigned long oh = (ih + 2 * param.pad_h - param.dilate_h * (param.window_h - 1) - |
| 21 | 1) / param.stride_h + |
| 22 | 1; |
| 23 | unsigned long ow = (iw + 2 * param.pad_w - param.dilate_w * (param.window_w - 1) - |
| 24 | 1) / param.stride_w + |
| 25 | 1; |
| 26 | param.out_h = ih; |
| 27 | param.out_w = iw; |
| 28 | |
| 29 | auto make_graph = [&](const Checker::SymInpArray& inputs) -> Checker::SymOutArray { |
| 30 | return {opr::SlidingWindowTranspose::make(inputs[0], param)}; |
| 31 | }; |
| 32 | |
| 33 | auto fwd = [&](Checker::NumOutArray& dest, Checker::NumInpArray inp) { |
| 34 | auto opr = megdnn_naive_handle() |
| 35 | ->create_operator<megdnn::SlidingWindowTranspose>(); |
| 36 | opr->param() = param; |
| 37 | TensorLayout dest_layout; |
| 38 | opr->deduce_layout(inp[0]->layout(), dest_layout); |
| 39 | std::vector<dt_byte> workspace( |
| 40 | opr->get_workspace_in_bytes(inp[0]->layout(), dest_layout)); |
| 41 | dest[0].dtype(dtype::Float32()) |
| 42 | .comp_node(inp[0]->comp_node()) |
| 43 | .resize(dest_layout); |
| 44 | opr->exec( |
| 45 | inp[0]->as_megdnn(), dest[0].as_megdnn(), |
| 46 | {workspace.data(), workspace.size()}); |
| 47 | }; |
| 48 | Checker::RunOptions opt; |
| 49 | opt.numdiff_eps = 1; |
| 50 | Checker checker{make_graph, fwd}; |
| 51 | |
| 52 | checker.run({TensorShape{2, 3, oh, ow, param.window_h, param.window_w}}, opt) |
| 53 | .run({TensorShape{4, 5, oh, ow, param.window_h, param.window_w}}, opt) |
| 54 | .run({TensorShape{3, 2, oh, ow, param.window_h, param.window_w}}, opt); |
| 55 | } |
| 56 | |
| 57 | // vim: syntax=cpp.doxygen foldmethod=marker foldmarker=f{{{,f}}} |
nothing calls this directly
no test coverage detected