Skip to content
Open
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
3 changes: 2 additions & 1 deletion onnxscript/optimizer/_constant_folding.py
Original file line number Diff line number Diff line change
Expand Up @@ -526,7 +526,8 @@ def cast_like(node: ir.Node, op, state: OptimizerState) -> ReturnValue:
return None
if source_element_type == target_element_type:
return op.Identity(input0)
return op.Cast(input0, to=target_element_type)
# Keep saturate, which is what ONNX's CastLike function body passes on to Cast.
return op.Cast(input0, to=target_element_type, saturate=node.attributes.get("saturate"))


@register("Shape")
Expand Down
81 changes: 81 additions & 0 deletions onnxscript/optimizer/_constant_folding_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,6 +250,87 @@ def test_fold_redundant_cast2(self):
self.assertEqual(optimized.graph[0].outputs[0].name, "z")
self.assertEqual(optimized.graph[0].inputs[0].name, "x")

def test_cast_like_to_cast_preserves_saturate(self):
model = """
<ir_version: 10, opset_import: [ "" : 19]>
agraph (float[3] x) => (float[3] z) {
like = Constant <value = float8e4m3fn[1] {0}> ()
x_f8 = CastLike <saturate = 0> (x, like)
z = Cast <to = 1> (x_f8)
}
"""
original = ir.from_onnx_text(model)
data = np.array([1000.0, 1.0, -1000.0], dtype=np.float32)
# With saturate=0, out-of-range values become NaN instead of +/-448
expected = np.array([np.nan, 1.0, np.nan], dtype=np.float32)
np.testing.assert_array_equal(
ReferenceEvaluator(ir.serde.serialize_model(original)).run(None, {"x": data})[0],
expected,
)

optimized = self._fold(original)
self.assertEqual(len(optimized.graph), 2)
cast = optimized.graph[0]
self.assertEqual(cast.op_type, "Cast")
self.assertEqual(cast.attributes["to"].as_int(), ir.DataType.FLOAT8E4M3FN)
self.assertEqual(cast.attributes["saturate"].as_int(), 0)
np.testing.assert_array_equal(
ReferenceEvaluator(ir.serde.serialize_model(optimized)).run(None, {"x": data})[0],
expected,
)

def test_fold_cast_like_preserves_saturate(self):
model = """
<ir_version: 10, opset_import: [ "" : 19]>
agraph () => (float[3] z) {
x = Constant <value = float[3] {1000.0, 1.0, -1000.0}> ()
like = Constant <value = float8e4m3fn[1] {0}> ()
x_f8 = CastLike <saturate = 0> (x, like)
z = Cast <to = 1> (x_f8)
}
"""

optimized = self._fold(model)
self.assertEqual(len(optimized.graph), 0)
np.testing.assert_array_equal(
optimized.graph.initializers["z"].const_value.numpy(),
np.array([np.nan, 1.0, np.nan], dtype=np.float32),
)

def test_cast_like_to_cast_does_not_add_round_mode(self):
# ONNX's CastLike function body passes only saturate on to Cast.
model = """
<ir_version: 11, opset_import: [ "" : 24]>
agraph (float[2] x) => (float8e4m3fn[2] z) {
like = Constant <value = float8e4m3fn[1] {0}> ()
z = CastLike <round_mode = "down"> (x, like)
}
"""

optimized = self._fold(model)
self.assertEqual([n.op_type for n in optimized.graph], ["Cast"])
self.assertNotIn("round_mode", optimized.graph[0].attributes)

def test_cast_like_with_attribute_reference_keeps_it(self):
model = """
<ir_version: 10, opset_import: ["this" : 1, "" : 19]>
agraph (float[3] x) => (float8e4m3fn[3] z) {
z = this.function <sat = 0> (x)
}
<domain: "this", opset_import: ["" : 19]>
function <sat> (x) => (z) {
like = Constant <value = float8e4m3fn[1] {0}> ()
z = CastLike <saturate: int = @sat> (x, like)
}
"""

optimized = self._fold(model)
function = next(iter(optimized.functions.values()))
self.assertEqual([n.op_type for n in function], ["Cast"])
saturate = function[0].attributes["saturate"]
self.assertTrue(saturate.is_ref())
self.assertEqual(saturate.ref_attr_name, "sat")

def test_shape_inference(self):
model = """
<ir_version: 7, opset_import: [ "" : 17]>
Expand Down