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

Function TEST

runtime/test/instruction/dimshuffle.cpp:8–78  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6using namespace test;
7
8TEST(INSTRUCTION, Dimshuffle) {
9 //! input shape =[10, 20, 10, 10]
10 std::shared_ptr<Tensor> output;
11 std::vector<int> data0(20 * 10 * 10 * 10);
12 std::vector<int> data1(20 * 10 * 10 * 10);
13 for (size_t i = 0; i < 20 * 10 * 10 * 10; i++) {
14 data0[i] = i + 1;
15 data1[i] = i + 1;
16 }
17
18 auto create_dimshuffle = [&](Tensor* input, std::vector<uint32_t> pattern) {
19 auto dimshuffle = std::make_shared<Dimshuffle>();
20 dimshuffle->pattern_dim = pattern.size();
21 for (int i = 0; i < pattern.size(); i++) {
22 dimshuffle->pattern[i] = pattern[i];
23 }
24
25 output = std::make_shared<Tensor>();
26 output->is_dynamic = true;
27
28 dimshuffle->input = input;
29 dimshuffle->output = output.get();
30 return dimshuffle;
31 };
32 VM* vm = create_vm();
33 auto test_dimshuffle = [&](Tensor* input, std::vector<uint32_t> pattern,
34 const Tensor& expect) {
35 auto dimshuffle = create_dimshuffle(input, pattern);
36 Instruction inst;
37 inst.tag = TinyNN_INST_DIMSHUFFLE;
38 inst.workload.dimshuffle = *dimshuffle;
39 vm_instruction_call(vm, &inst);
40 check_tensor(*dimshuffle->output, expect);
41 vm->model->host_dev.free(dimshuffle->output->ptr);
42 };
43
44 auto test_case = [&](std::vector<uint32_t> pattern) {
45 std::vector<uint32_t> shape{10, 20, 10, 10};
46 auto src = create_tensor(shape, TinyNNDType::TinyNN_INT, data0.data());
47 Tensor src_copy = *src;
48 Layout src_layout = src->layout, out_layout = src->layout;
49 for (int i = 0; i < pattern.size(); i++) {
50 src_layout.dims[i] = src->layout.dims[pattern[i]];
51 out_layout.dims[i] = src->layout.dims[pattern[i]];
52 src_layout.stride[i] = src->layout.stride[pattern[i]];
53 }
54 out_layout.stride[out_layout.nr_dim - 1] = 1;
55 for (int index = out_layout.nr_dim - 2; index >= 0; index--) {
56 out_layout.stride[index] =
57 out_layout.dims[index + 1] * out_layout.stride[index + 1];
58 }
59
60 auto trueth = create_tensor(shape, TinyNNDType::TinyNN_INT, data1.data());
61 trueth->layout = out_layout;
62 NoconIter src_iter = init_iter(src_layout);
63 NoconIter dst_iter = init_iter(out_layout);
64 int* dst_data = static_cast<int*>(trueth->ptr);
65 int* src_data = static_cast<int*>(src->ptr);

Callers

nothing calls this directly

Calls 5

vm_instruction_callFunction · 0.85
init_iterFunction · 0.85
inc_iterFunction · 0.85
getMethod · 0.45
freeMethod · 0.45

Tested by

no test coverage detected