| 49 | |
| 50 | template <typename SharedDeviceTensor, typename MultipleDeviceTensorHolder> |
| 51 | void 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 |
nothing calls this directly
no test coverage detected