diff --git a/onnxscript/_internal/converter.py b/onnxscript/_internal/converter.py index c2f6b0fb63..f584315a5f 100644 --- a/onnxscript/_internal/converter.py +++ b/onnxscript/_internal/converter.py @@ -1103,12 +1103,14 @@ def check_num_outputs(n): def ret(exp, i, suffix): preferred_name = f"return_val{suffix}" return_var = self._translate_expr(exp, preferred_name) - val = self._lookup(return_var.name, self._source_of(exp), raise_exception=False) - if isinstance(val, values.SymbolValue) and isinstance(val.value, ir.Value): - if val.value.is_graph_input(): - # In ONNX, a graph-input cannot be an output of the graph. - # We need to insert a copy. - return_var = self._emit_copy(return_var, preferred_name) + if return_var.is_graph_input(): + # Use the resolved ONNX value: the Python input name may have been rebound. + copy_name = ( + exp.id + if isinstance(exp, ast.Name) and exp.id != return_var.name + else preferred_name + ) + return_var = self._emit_copy(return_var, copy_name) for prev_output in self._current_fn.outputs: if prev_output.name == return_var.name: # ONNX does not allow duplicate output names. diff --git a/onnxscript/_internal/converter_test.py b/onnxscript/_internal/converter_test.py index c1fe276ef5..fba141b428 100644 --- a/onnxscript/_internal/converter_test.py +++ b/onnxscript/_internal/converter_test.py @@ -607,6 +607,58 @@ def duplicate_output(X): outputs = duplicate_output.to_function_proto().output self.assertNotEqual(outputs[0], outputs[1]) + def test_returned_input_alias_preserves_name(self): + @script(default_opset=op) + def returned_alias(X): + Y = X + return Y + + function_proto = returned_alias.to_function_proto() + self.assertEqual(function_proto.output[0], "Y") + self.assertEqual(function_proto.node[-1].op_type, "Identity") + self.assertEqual(function_proto.node[-1].output[0], "Y") + + @script(default_opset=op) + def returned_input(X): + return X + + self.assertEqual(returned_input.to_function_proto().output[0], "return_val") + + def test_returned_input_alias_after_rebinding_input(self): + @script(default_opset=op) + def returned_alias(X: FLOAT[2]) -> FLOAT[2]: + Y = X + X = op.Neg(X) + return Y + + model = returned_alias.to_model_proto() + onnx.checker.check_model(model, full_check=True) + self.assertEqual(model.graph.output[0].name, "Y") + self.assertEqual(model.graph.node[-1].op_type, "Identity") + self.assertEqual(list(model.graph.node[-1].input), ["X"]) + x = np.array([1.0, -2.0], dtype=np.float32) + actual = create_cpu_inference_session(model.SerializeToString()).run(None, {"X": x}) + np.testing.assert_array_equal(actual[0], x) + + def test_returned_input_alias_name_collision(self): + @script(default_opset=op) + def returned_alias(X: FLOAT[2], Y: FLOAT[2]) -> (FLOAT[2], FLOAT[2]): + Y = X + return Y, Y + + model = returned_alias.to_model_proto() + onnx.checker.check_model(model, full_check=True) + outputs = [value.name for value in model.graph.output] + self.assertEqual(len(set(outputs)), 2) + self.assertTrue(all(name.startswith("Y_") for name in outputs)) + self.assertTrue(set(outputs).isdisjoint({"X", "Y"})) + x = np.array([1.0, -2.0], dtype=np.float32) + actual = create_cpu_inference_session(model.SerializeToString()).run( + None, {"X": x, "Y": -x} + ) + for output in actual: + np.testing.assert_array_equal(output, x) + def test_bool_attr_promotion(self): @script() def if_then_else(flag: bool, Y, Z):