diff --git a/src/onnx/onnx_parser.cpp b/src/onnx/onnx_parser.cpp index a091d4f7086..5545be47069 100644 --- a/src/onnx/onnx_parser.cpp +++ b/src/onnx/onnx_parser.cpp @@ -620,6 +620,9 @@ static void log_node_parse_exception(const onnx::NodeProto& node, std::vector onnx_parser::parse_graph(module* mod, const onnx::GraphProto& graph, bool inlining) { + // Save/restore parent_input_nodes so sibling subgraphs can't see each other's inputs. + auto saved_parent_inputs = parent_input_nodes; + std::vector node_indices(graph.node_size()); if(check_sorted(graph, parent_input_nodes)) @@ -741,6 +744,7 @@ onnx_parser::parse_graph(module* mod, const onnx::GraphProto& graph, bool inlini erase_if(instructions, [&](auto&& p) { return mod->has_instruction(p.second); }); } + parent_input_nodes = std::move(saved_parent_inputs); return output_ins; } diff --git a/src/replace_allocate.cpp b/src/replace_allocate.cpp index a8573f94ea0..737aeac9378 100644 --- a/src/replace_allocate.cpp +++ b/src/replace_allocate.cpp @@ -115,17 +115,24 @@ get_output_debug_symbols(const module& mod) return mod_output_debug_symbols; } -void insert_copy(module& m, const allocation_model& model) +void insert_copy(module& m, const allocation_model& model, bool is_root) { auto returns = m.get_returns(); std::unordered_set returns_set(returns.begin(), returns.end()); for(auto ins : returns_set) { + if(is_root and (ins->name() == "@param" or ins->name() == "@literal")) + continue; if(ins->get_shape().any_of_dynamic()) continue; auto aliases = instruction::get_output_alias(ins); if(std::any_of(aliases.begin(), aliases.end(), [&](instruction_ref alias) { - return alias->get_shape() == ins->get_shape(); + if(alias->get_shape() != ins->get_shape()) + return false; + if(alias->name() == "allocate" or alias->name() == model.name()) + return true; + return alias->name() == "@param" and alias != ins and + not ins->get_operator().is_context_free(); })) continue; auto insert_ins = std::next(ins); @@ -176,7 +183,7 @@ void replace_allocate::apply(module_pass_manager& mpm) const insert_submod_allocations(ins, m, model); } if(not root_offload_copy and model.needs_out_params()) - insert_copy(m, model); + insert_copy(m, model, is_root); auto mod_output_names = create_output_names(m); auto mod_output_debug_symbols = get_output_debug_symbols(m); for(auto ins : iterator_for(m))