From 41c5e6524bb0592b551b8fd43994f48e6fd50d0e Mon Sep 17 00:00:00 2001 From: Mohammed Alkindi Date: Wed, 7 Oct 2026 21:48:22 -0700 Subject: [PATCH] Skip constant folding for non-deterministic Bernoulli --- onnxscript/optimizer/_constant_folding.py | 1 + onnxscript/optimizer/_constant_folding_test.py | 15 +++++++++++++++ 2 files changed, 16 insertions(+) diff --git a/onnxscript/optimizer/_constant_folding.py b/onnxscript/optimizer/_constant_folding.py index 66179e8efa..f78acfe070 100644 --- a/onnxscript/optimizer/_constant_folding.py +++ b/onnxscript/optimizer/_constant_folding.py @@ -53,6 +53,7 @@ "RandomUniformLike", "RandomNormalLike", "Multinomial", + "Bernoulli", } ) diff --git a/onnxscript/optimizer/_constant_folding_test.py b/onnxscript/optimizer/_constant_folding_test.py index 4cac9debd1..5c76dad1a6 100644 --- a/onnxscript/optimizer/_constant_folding_test.py +++ b/onnxscript/optimizer/_constant_folding_test.py @@ -869,6 +869,21 @@ def test_dequantize_linear_is_not_folded(self): # DequantizeLinear should not be folded even when all inputs are constants self.assertEqual(ops, ["DequantizeLinear"]) + def test_bernoulli_is_not_folded(self): + model_text = """ + + agraph () => (float[4] z) + + { + z = Bernoulli (p) + } + """ + model = ir.from_onnx_text(model_text) + optimized = self._fold(model) + ops = [node.op_type for node in optimized.graph] + # Bernoulli is random, so folding it would freeze a single sample + self.assertEqual(ops, ["Bernoulli"]) + def test_multi_graph_identity_output_preserves_output_name(self): model = """