| 2138 | |
| 2139 | template <size_t nr_out> |
| 2140 | void test_param_pack_split(const TensorShapeArray& shapes) { |
| 2141 | auto cn = CompNode::load("xpu0"); |
| 2142 | auto align = std::max<size_t>(cn.get_mem_addr_alignment() / 4, 1); |
| 2143 | size_t concat_size = 0; |
| 2144 | mgb_assert(shapes.size() == nr_out); |
| 2145 | for (auto&& i : shapes) { |
| 2146 | concat_size = get_aligned_power2(concat_size, align) + i.total_nr_elems(); |
| 2147 | } |
| 2148 | |
| 2149 | using Checker = AutoOprChecker<1, nr_out>; |
| 2150 | |
| 2151 | auto make_graph = [&](const typename Checker::SymInpArray& inputs) -> |
| 2152 | typename Checker::SymOutArray { |
| 2153 | auto offsets_val = megdnn::ParamPackConcat::gen_offsets( |
| 2154 | shapes, cn.get_mem_addr_alignment(), 4); |
| 2155 | HostTensorND offsets; |
| 2156 | std::copy_n( |
| 2157 | offsets_val.data(), offsets_val.size(), |
| 2158 | offsets.dtype(dtype::Int32{}) |
| 2159 | .comp_node(cn) |
| 2160 | .resize({offsets_val.size()}) |
| 2161 | .ptr<dt_int32>()); |
| 2162 | auto out = opr::ParamPackSplit::make(inputs[0], offsets_val, shapes); |
| 2163 | mgb_assert(out.size() == nr_out); |
| 2164 | typename Checker::SymOutArray ret; |
| 2165 | for (size_t i = 0; i < nr_out; ++i) { |
| 2166 | ret[i] = out[i]; |
| 2167 | } |
| 2168 | return ret; |
| 2169 | }; |
| 2170 | |
| 2171 | auto fwd = [&](typename Checker::NumOutArray& dest, |
| 2172 | typename Checker::NumInpArray inp) { |
| 2173 | size_t offset = 0; |
| 2174 | auto ptr = inp[0]->template ptr<float>(); |
| 2175 | for (size_t i = 0; i < nr_out; ++i) { |
| 2176 | dest[i].resize(shapes[i]); |
| 2177 | offset = get_aligned_power2(offset, align); |
| 2178 | auto nr_elem = shapes[i].total_nr_elems(); |
| 2179 | memcpy(dest[i].template ptr<float>(), ptr + offset, nr_elem * 4); |
| 2180 | offset += nr_elem; |
| 2181 | } |
| 2182 | }; |
| 2183 | |
| 2184 | Checker{make_graph, fwd} |
| 2185 | .run({TensorShape{concat_size}}) |
| 2186 | .run({TensorShape{concat_size}}) |
| 2187 | .run({TensorShape{concat_size}}); |
| 2188 | } |
| 2189 | |
| 2190 | } // anonymous namespace |
| 2191 |
nothing calls this directly
no test coverage detected