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

Function param_merge

src/gopt/impl/inference.cpp:51–95  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

49
50template <typename SharedDeviceTensor, typename MultipleDeviceTensorHolder>
51void param_merge(OptState& opt_state) {
52 auto rewriter = opt_state.graph().make_rewriter();
53 ThinHashMap<OperatorNodeBase*, size_t> opr2idx;
54 std::vector<OperatorNodeBase*> all_oprs;
55 typename MultipleDeviceTensorHolder::ValueArray all_values;
56
57 auto cb_find_opr = [&](cg::OperatorNodeBase* opr) {
58 if (opr->same_type<SharedDeviceTensor>()) {
59 auto p = &opr->cast_final<SharedDeviceTensor>();
60 // ShredD may be manu
61 opr2idx[p] = all_values.size();
62 all_values.push_back(p->dev_data());
63 all_oprs.push_back(p);
64 }
65 };
66 opt_state.graph().iter(cb_find_opr);
67 SymbolVarArray new_vars;
68 auto cb_replace = [&](cg::OperatorNodeBase* opr) {
69 auto iter = opr2idx.find(opr);
70 if (iter == opr2idx.end()) {
71 rewriter.auto_replace_outputs(opr);
72 } else {
73 if (new_vars.empty()) {
74 // new oprs must be created in iter callback; so we populate
75 // new_vars lazily
76 new_vars = MultipleDeviceTensorHolder::make(
77 *opt_state.graph().comp_graph(), std::move(all_values),
78 {ssprintf("merged%zu", all_values.size())});
79 for (size_t i = 0; i < new_vars.size(); ++i) {
80 auto src = all_oprs[i]->output(0);
81 if (src->has_name_set()) {
82 new_vars[i].rename(src->name());
83 }
84 }
85 }
86 rewriter.replace_var(
87 opr->output(0), new_vars.at(iter->second).node(),
88 mgb_cstr_log("replace multi SharedDeviceTensor(Format) to "
89 "MultipleDeviceTensorHolder(Format)"));
90 }
91 };
92 opt_state.graph().iter(cb_replace);
93
94 rewriter.apply_inplace();
95}
96
97} // namespace
98

Callers

nothing calls this directly

Calls 15

make_rewriterMethod · 0.80
comp_graphMethod · 0.80
has_name_setMethod · 0.80
renameMethod · 0.80
apply_inplaceMethod · 0.80
makeFunction · 0.50
graphMethod · 0.45
sizeMethod · 0.45
push_backMethod · 0.45
iterMethod · 0.45
findMethod · 0.45
endMethod · 0.45

Tested by

no test coverage detected