| 155 | } |
| 156 | |
| 157 | ReformatEmitter::UnderlyingBuilders ReformatEmitter::analyze() const { |
| 158 | struct Dim { |
| 159 | Dimension dim; |
| 160 | int index; |
| 161 | Dim(Dimension dim_, int index_) : dim{dim_}, index{index_} {} |
| 162 | }; |
| 163 | SmallVector<Dim> src_dims; |
| 164 | SmallVector<Dim> dest_dims; |
| 165 | for (size_t i = 0; i < m_src.ndim; ++i) |
| 166 | src_dims.emplace_back(Dim(m_src[i], i)); |
| 167 | for (size_t i = 0; i < m_dest.ndim; ++i) |
| 168 | dest_dims.emplace_back(Dim(m_dest[i], i)); |
| 169 | auto compare = [](const Dim& lhs, const Dim& rhs) { return lhs.dim < rhs.dim; }; |
| 170 | std::sort(src_dims.begin(), src_dims.end(), compare); |
| 171 | std::sort(dest_dims.begin(), dest_dims.end(), compare); |
| 172 | auto src_iter = src_dims.begin(); |
| 173 | auto dest_iter = dest_dims.begin(); |
| 174 | for (; src_iter != src_dims.end() && dest_iter != dest_dims.end();) { |
| 175 | if (src_iter->dim == dest_iter->dim) { |
| 176 | src_iter++; |
| 177 | dest_iter++; |
| 178 | } else if (src_iter->dim < dest_iter->dim) { |
| 179 | auto split = dest_iter->dim / src_iter->dim; |
| 180 | int dim_idx = dest_iter->index; |
| 181 | dest_iter = dest_dims.insert(dest_iter, Dim(src_iter->dim, dim_idx)); |
| 182 | dest_iter++; |
| 183 | dest_iter->dim = split; |
| 184 | dest_iter->index = dim_idx; |
| 185 | src_iter++; |
| 186 | } else { |
| 187 | auto split = src_iter->dim / dest_iter->dim; |
| 188 | int dim_idx = src_iter->index; |
| 189 | src_iter = src_dims.insert(src_iter, Dim(dest_iter->dim, dim_idx)); |
| 190 | src_iter++; |
| 191 | src_iter->dim = split; |
| 192 | src_iter->index = dim_idx; |
| 193 | dest_iter++; |
| 194 | } |
| 195 | } |
| 196 | mgb_assert(src_dims.size() == dest_dims.size()); |
| 197 | std::vector<int> src_perm(src_dims.size()); |
| 198 | std::vector<int> permute(dest_dims.size()); |
| 199 | std::iota(src_perm.begin(), src_perm.end(), 0); |
| 200 | std::iota(permute.begin(), permute.end(), 0); |
| 201 | std::sort(src_perm.begin(), src_perm.end(), [&](const int a, const int b) { |
| 202 | if (src_dims[a].index != src_dims[b].index) |
| 203 | return src_dims[a].index < src_dims[b].index; |
| 204 | return src_dims[a].dim < src_dims[b].dim; |
| 205 | }); |
| 206 | std::sort(permute.begin(), permute.end(), [&](const int a, const int b) { |
| 207 | int perm_a = src_perm[a]; |
| 208 | int perm_b = src_perm[b]; |
| 209 | if (dest_dims[perm_a].index != dest_dims[perm_b].index) |
| 210 | return dest_dims[perm_a].index < dest_dims[perm_b].index; |
| 211 | return dest_dims[perm_a].dim < dest_dims[perm_b].dim; |
| 212 | }); |
| 213 | NamedTensorShape i1, i2; |
| 214 | i1.ndim = src_dims.size(), i2.ndim = dest_dims.size(); |
nothing calls this directly
no test coverage detected