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
1 change: 1 addition & 0 deletions onnxscript/optimizer/_constant_folding.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
"RandomUniformLike",
"RandomNormalLike",
"Multinomial",
"Bernoulli",
}
)

Expand Down
15 changes: 15 additions & 0 deletions onnxscript/optimizer/_constant_folding_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = """
<ir_version: 10, opset_import: [ "" : 18]>
agraph () => (float[4] z)
<float[4] p = {0.5, 0.5, 0.5, 0.5}>
{
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 = """
<ir_version: 10, opset_import: ["" : 20]>
Expand Down