diff --git a/onnxscript/_internal/converter.py b/onnxscript/_internal/converter.py index c2f6b0fb63..3789bb5f47 100644 --- a/onnxscript/_internal/converter.py +++ b/onnxscript/_internal/converter.py @@ -1108,7 +1108,12 @@ def ret(exp, i, suffix): 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) + 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..ea5eb1f339 100644 --- a/onnxscript/_internal/converter_test.py +++ b/onnxscript/_internal/converter_test.py @@ -607,6 +607,23 @@ 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_bool_attr_promotion(self): @script() def if_then_else(flag: bool, Y, Z):