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

Method analyze

src/gopt/impl/global_layout_transform/reformat_emitter.cpp:157–230  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

155}
156
157ReformatEmitter::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();

Callers

nothing calls this directly

Calls 13

DimClass · 0.85
sortFunction · 0.85
MakeShapeEmitterClass · 0.85
ReshapeEmitterClass · 0.85
DimshuffleEmitterClass · 0.85
emplace_backMethod · 0.80
iotaFunction · 0.50
beginMethod · 0.45
endMethod · 0.45
insertMethod · 0.45
sizeMethod · 0.45
eq_shapeMethod · 0.45

Tested by

no test coverage detected