diff --git a/onnxscript/optimizer/_optimizer.py b/onnxscript/optimizer/_optimizer.py index 36beb7f848..1f7f342c35 100644 --- a/onnxscript/optimizer/_optimizer.py +++ b/onnxscript/optimizer/_optimizer.py @@ -66,6 +66,7 @@ def optimize_ir( common_passes.CommonSubexpressionEliminationPass(), common_passes.OutputFixPass(), common_passes.NameFixPass(), + common_passes.RemoveUnusedNodesPass(), ] if inline: # Inline all functions first before optimizing diff --git a/onnxscript/optimizer/_optimizer_test.py b/onnxscript/optimizer/_optimizer_test.py index 05064fbc70..4385bd7e61 100644 --- a/onnxscript/optimizer/_optimizer_test.py +++ b/onnxscript/optimizer/_optimizer_test.py @@ -84,6 +84,32 @@ def test_static_split_to_sequence_with_uneven_split_ir(self): self.assertEqual(len(model_ir.graph.node(0).outputs), 2) self.assertEqual(model_ir.graph.node(0).op_type, "Split") + def test_name_fix_preserves_unnamed_unused_outputs(self): + model_proto = onnx.parser.parse_model( + """ + + main_graph ( + float[1, 2, 3, 3] x, + float[2] scale, + float[2] bias, + float[2] mean, + float[2] variance + ) => (float[1, 2, 3, 3] y) { + y, running_mean, running_var = BatchNormalization + (x, scale, bias, mean, variance) + } + """ + ) + model_ir = ir.serde.deserialize_model(model_proto) + model_ir.graph.inputs[1].name = "x" + + optimizer.optimize_ir(model_ir, num_iterations=1, onnx_shape_inference=False) + + self.assertEqual([input.name for input in model_ir.graph.inputs[:2]], ["x", "x_1"]) + self.assertEqual( + [output.name for output in model_ir.graph.node(0).outputs], ["y", "", ""] + ) + if __name__ == "__main__": unittest.main()