| 582 | } |
| 583 | |
| 584 | const 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 | } |
nothing calls this directly
no test coverage detected