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

Method concat_and_prepare

src/core/impl/graph/cg_impl_partial.cpp:584–694  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

582}
583
584const OprNodeArray* ComputingGraphImpl::MultiPartCompiler::concat_and_prepare() {
585 // no callback in out_spec_concat, so CallbackCaller would not be inserted
586 OutputSpec out_spec_concat;
587 std::vector<bool> out_spec_concat_from_original;
588
589 // init out_spec_concat
590 {
591 SymbolVarArray part_vars;
592 ExtraDependencyMerger dep_merger;
593 for (size_t part = 0; part < m_out_specs.size(); ++part) {
594 part_vars.clear();
595 for (auto&& i : m_out_specs[part]) {
596 part_vars.push_back(i.first);
597 }
598 auto&& dest_vars = dep_merger.add(part_vars);
599 for (size_t i = 0; i < dest_vars.size(); ++i) {
600 out_spec_concat.push_back({dest_vars[i].node(), {}});
601 out_spec_concat_from_original.push_back(i < part_vars.size());
602 }
603 dest_vars.clear();
604 }
605 }
606
607 auto remap_priority = [&](const VarNodeArray& dest_vars,
608 const TopoSorter::PriorityItem* items, size_t nr_item) {
609 const size_t nr_part = m_out_specs.size();
610 mgb_assert(
611 nr_item <= static_cast<size_t>(std::numeric_limits<int>::max()),
612 "too many oprs");
613
614 ThinHashMap<const OperatorNodeBase*, size_t> endpoint2part;
615
616 // remap optimized vars to specs and init endpoint2part
617 mgb_assert(dest_vars.size() == out_spec_concat.size());
618 size_t dest_var_idx = 0;
619 auto skip_non_orig_vars = [&](size_t part) {
620 while (!out_spec_concat_from_original[dest_var_idx]) {
621 ++dest_var_idx;
622 if (dest_var_idx == dest_vars.size()) {
623 mgb_assert(part == m_out_specs.size() - 1);
624 break;
625 } else {
626 // add extra vars to out spec so we do not need to handle
627 // extra_vardeps in graph copy
628 m_out_specs[part].push_back({dest_vars[dest_var_idx - 1], {}});
629 }
630 }
631 };
632 for (size_t part = 0; part < nr_part; ++part) {
633 int begin = dest_var_idx;
634 for (auto&& i : m_out_specs[part]) {
635 mgb_assert(out_spec_concat_from_original[dest_var_idx]);
636 i.first = dest_vars[dest_var_idx++];
637 }
638
639 if (dest_var_idx < dest_vars.size()) {
640 skip_non_orig_vars(part);
641 }

Callers

nothing calls this directly

Calls 15

maxFunction · 0.85
update_maxFunction · 0.85
sortFunction · 0.85
backMethod · 0.80
set_priority_remapperMethod · 0.80
compile_prepareMethod · 0.80
sizeMethod · 0.45
clearMethod · 0.45
push_backMethod · 0.45
addMethod · 0.45
nodeMethod · 0.45
insertMethod · 0.45

Tested by

no test coverage detected