Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions onnxscript/version_converter/_version_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,11 +160,14 @@ def dft_19_20(node: ir.Node, op):
dft_length = node.inputs[1] if len(node.inputs) > 1 else None
inverse = _get_int_attribute(node, "inverse", 0)
onesided = _get_int_attribute(node, "onesided", 0)
axis = _get_int_attribute(node, "axis", None)
if axis is not None:
axis_value = op.Constant(value_int=axis)
return op.DFT(input, dft_length, axis_value, inverse=inverse, onesided=onesided)
return None
# In opset 19 `axis` is an attribute defaulting to 1; in opset 20 it became an
# input defaulting to -2. An omitted attribute must therefore be materialized
# explicitly, or the converted node silently transforms a different axis.
axis = _get_int_attribute(node, "axis", 1)
if axis is None:
return None
axis_value = op.Constant(value_int=axis)
return op.DFT(input, dft_length, axis_value, inverse=inverse, onesided=onesided)


@register("GridSample", node_version=19, up_conversion=True)
Expand Down
22 changes: 22 additions & 0 deletions onnxscript/version_converter/_version_converter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,28 @@ def test_version_convert_compatible(self):
self.assertEqual(model.graph.node(3).version, 20)
self.assertEqual(len(model.graph.node(3).inputs), 3)

def test_version_convert_dft_without_axis_attribute(self):
# Opset 19 defines DFT axis as an attribute defaulting to 1. Opset 20 moved
# it to an input defaulting to -2, so an omitted attribute has to be
# materialized or the converted node transforms a different axis.
model = ir.from_onnx_text(
"""
<ir_version: 7, opset_import: [ "" : 19]>
agraph (float[2, 3, 4, 2] input_x) => (float[2, 3, 4, 2] output)
{
output = DFT (input_x)
}
"""
)
version_converter.convert_version(model, target_version=20)
self.assertEqual(model.opset_imports[""], 20)

dft_node = next(node for node in model.graph if node.op_type == "DFT")
self.assertEqual(len(dft_node.inputs), 3)
axis_input = dft_node.inputs[2]
self.assertIsNotNone(axis_input)
self.assertEqual(axis_input.producer().attributes["value_int"].value, 1)

def test_version_convert_gridsample_linear(self):
model = ir.from_onnx_text(
"""
Expand Down
Loading