| 1104 | } |
| 1105 | |
| 1106 | std::vector<instruction_ref> |
| 1107 | module::fuse(const std::vector<instruction_ref>& inss, |
| 1108 | std::unordered_map<instruction_ref, instruction_ref>* map_ins, |
| 1109 | module::inserter insert, |
| 1110 | const std::function<shape(const shape&)>& shape_transform) |
| 1111 | { |
| 1112 | std::unordered_map<instruction_ref, instruction_ref> default_map_ins; |
| 1113 | if(map_ins == nullptr) |
| 1114 | map_ins = &default_map_ins; |
| 1115 | std::vector<instruction_ref> inputs; |
| 1116 | for(auto ins : inss) |
| 1117 | { |
| 1118 | for(auto input : ins->inputs()) |
| 1119 | { |
| 1120 | if(contains(inss, input)) |
| 1121 | continue; |
| 1122 | if(contains(inputs, input)) |
| 1123 | continue; |
| 1124 | inputs.push_back(input); |
| 1125 | } |
| 1126 | } |
| 1127 | insert_params(*this, inputs, *map_ins, shape_transform); |
| 1128 | return this->add_instructions(inss, map_ins, std::move(insert)); |
| 1129 | } |
| 1130 | |
| 1131 | std::vector<instruction_ref> |
| 1132 | module::fuse(const module& m, |