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

Function gen_dct_constriant

dnn/test/common/dct_ref.cpp:19–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17}
18
19CheckerHelper::TensorsConstriant gen_dct_constriant(
20 const size_t /* n */, const size_t ic, const size_t ih, const size_t iw,
21 const size_t oc, Param param) {
22 auto constraint = [=](CheckerHelper::TensorValueArray& tensors_orig) {
23 const size_t block = param.dct_block_size;
24 const int block_c = param.format == Param::Format::NCHW4 ? 4 : 1;
25 megdnn_assert(oc % block_c == 0, "oc mod block_c must == 0");
26 std::shared_ptr<DctTestcase> test_case_ptr = DctTestcase::make();
27 DctTestcase& test_case = *test_case_ptr.get();
28 UniformIntRNG rng(0, 255);
29 UniformIntRNG mask_rng(0, 64 / block_c - 1);
30 const size_t no_mask_oc = ic * block * block;
31 megdnn_assert(ih % block == 0, "%zu mod %zu == 0", ih, block);
32 megdnn_assert(iw % block == 0, "%zu mod %zu == 0", iw, block);
33
34 TensorND mask_offset;
35 TensorND mask_val;
36 std::vector<int>& mask_offset_vec = test_case.mask_offset_vec;
37 std::vector<int>& mask_val_vec = test_case.mask_val_vec;
38 UniformIntRNG rng_oc(0, oc);
39 if (param.fastImpl == Param::FastImpl::FIX_32_MASK) {
40 auto fix_32_mask = get_fix_mask(Param::FastImpl::FIX_32_MASK);
41 mask_offset_vec = fix_32_mask.mask_offset;
42 mask_val_vec = fix_32_mask.mask_val;
43 megdnn_assert(oc == 32, "oc must eq 32");
44 } else if (no_mask_oc > oc) {
45 size_t remain_oc = oc;
46 mask_offset_vec.resize(ic + 1);
47 mask_val_vec.resize(oc);
48 mask_offset_vec[0] = 0;
49 for (size_t ic_idx = 0; ic_idx < ic; ++ic_idx) {
50 size_t random_len = (int)rng_oc.gen_single_val() * block_c;
51 size_t mask_len = (ic_idx == ic - 1) || (remain_oc == 0)
52 ? remain_oc
53 : random_len % remain_oc;
54 megdnn_assert(
55 mask_len % block_c == 0,
56 "mask_len mod block_c == 0, but %zu mod %d ", mask_len,
57 block_c);
58 const size_t oc_idx = mask_offset_vec[ic_idx];
59 remain_oc -= mask_len;
60 mask_offset_vec[ic_idx + 1] = oc_idx + mask_len;
61 for (size_t mask_idx = 0; mask_idx < mask_len; ++mask_idx) {
62 mask_val_vec[oc_idx + mask_idx] = (int)mask_rng.gen_single_val();
63 }
64 }
65 }
66 mask_offset = TensorND(
67 mask_offset_vec.data(), {{mask_offset_vec.size()}, dtype::Int32()});
68 mask_val =
69 TensorND(mask_val_vec.data(), {{mask_val_vec.size()}, dtype::Int32()});
70 if (tensors_orig.size() > 1) {
71 megdnn_assert(tensors_orig.size() == 4, "tensors_orig.size() == 4");
72 megdnn_assert(mask_offset_vec.size() >= 2, "mask_offset_vec.size() >= 2");
73 megdnn_assert(
74 tensors_orig[1].layout == mask_offset.layout,
75 "tensors_orig[1].layout == mask_offset.layout");
76 megdnn_assert(

Callers 1

TEST_FFunction · 0.85

Calls 11

get_fix_maskFunction · 0.85
resizeMethod · 0.80
dist_byteMethod · 0.80
spanMethod · 0.80
makeFunction · 0.50
TensorNDClass · 0.50
getMethod · 0.45
gen_single_valMethod · 0.45
dataMethod · 0.45
sizeMethod · 0.45
raw_ptrMethod · 0.45

Tested by

no test coverage detected