From e0dc140a461c6fc85fe68adfdb06e7bb8f20ea98 Mon Sep 17 00:00:00 2001 From: Mohammed Alkindi Date: Wed, 7 Oct 2026 21:52:57 -0700 Subject: [PATCH] Keep the saturate attribute when constant folding rewrites CastLike to Cast --- onnxscript/optimizer/_constant_folding.py | 3 +- .../optimizer/_constant_folding_test.py | 81 +++++++++++++++++++ 2 files changed, 83 insertions(+), 1 deletion(-) diff --git a/onnxscript/optimizer/_constant_folding.py b/onnxscript/optimizer/_constant_folding.py index 66179e8efa..60d439781f 100644 --- a/onnxscript/optimizer/_constant_folding.py +++ b/onnxscript/optimizer/_constant_folding.py @@ -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") diff --git a/onnxscript/optimizer/_constant_folding_test.py b/onnxscript/optimizer/_constant_folding_test.py index 4cac9debd1..d89f68ff05 100644 --- a/onnxscript/optimizer/_constant_folding_test.py +++ b/onnxscript/optimizer/_constant_folding_test.py @@ -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 = """ + + agraph (float[3] x) => (float[3] z) { + like = Constant () + x_f8 = CastLike (x, like) + z = Cast (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 = """ + + agraph () => (float[3] z) { + x = Constant () + like = Constant () + x_f8 = CastLike (x, like) + z = Cast (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 = """ + + agraph (float[2] x) => (float8e4m3fn[2] z) { + like = Constant () + z = CastLike (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 = """ + + agraph (float[3] x) => (float8e4m3fn[3] z) { + z = this.function (x) + } + + function (x) => (z) { + like = Constant () + z = CastLike (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 = """