diff --git a/onnxscript/optimizer/_optimizer.py b/onnxscript/optimizer/_optimizer.py index 36beb7f848..253be790ba 100644 --- a/onnxscript/optimizer/_optimizer.py +++ b/onnxscript/optimizer/_optimizer.py @@ -14,6 +14,27 @@ logger = logging.getLogger(__name__) +class _RemoveUnusedBatchNormalizationOutputsPass(ir.passes.InPlacePass): + """Remove BatchNormalization output slots marked unused by earlier passes.""" + + def call(self, model: ir.Model) -> ir.passes.PassResult: + modified = False + for graph_like in (model.graph, *model.functions.values()): + for node in ir.traversal.RecursiveGraphIterator(graph_like): + if node.domain not in {"", "ai.onnx"} or node.op_type != "BatchNormalization": + continue + output_count = len(node.outputs) + while output_count: + output = node.outputs[output_count - 1] + if output.name or output.uses(): + break + output_count -= 1 + if output_count != len(node.outputs): + node.resize_outputs(output_count) + modified = True + return ir.passes.PassResult(model, modified=modified) + + def optimize_ir( model: ir.Model, num_iterations: int = 2, @@ -65,6 +86,7 @@ def optimize_ir( common_passes.DeduplicateInitializersPass(), common_passes.CommonSubexpressionEliminationPass(), common_passes.OutputFixPass(), + _RemoveUnusedBatchNormalizationOutputsPass(), common_passes.NameFixPass(), ] if inline: diff --git a/onnxscript/optimizer/_optimizer_test.py b/onnxscript/optimizer/_optimizer_test.py index 05064fbc70..8def3a5ccd 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_does_not_restore_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"]) + self.assertNotIn("training_mode", model_ir.graph.node(0).attributes) + onnx.checker.check_model(ir.serde.serialize_model(model_ir), full_check=True) + if __name__ == "__main__": unittest.main()