| 17 | } |
| 18 | |
| 19 | CheckerHelper::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( |