onnxruntime
7529033f - Fix ReshapeFusion dropping allowzero on inferred 0-sized intermediate dims (#28349)

Commit
154 days ago
Fix ReshapeFusion dropping allowzero on inferred 0-sized intermediate dims (#28349) ### Description `ReshapeFusion::FuseContiguousReshapes` collapses a chain of `Reshape` / `Squeeze` / `Unsqueeze` nodes into a single `Reshape` whose shape data is taken verbatim from the fully-inferred output shape of the last node in the chain. The new node is created without an `allowzero` attribute, so it defaults to `allowzero = 0`. When that inferred shape contains a literal `0` dim (legitimate when the original chain used `allowzero=1`, or when intermediate tensors had zero-sized dimensions), the fused `Reshape` misinterprets the `0` as "copy the corresponding dim from the input tensor" — but the input here is the original input of the *first* reshape in the chain, with unrelated dims. The result is a silently wrong output shape (and a benign-looking `MergeShapeInfo` warning at graph load). ### Repro (before the fix) ```python import numpy as np, onnx, onnxruntime as ort, onnx.reference from onnx import helper, TensorProto X = helper.make_tensor_value_info("X", TensorProto.FLOAT, [0, 6, 2]) Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [None, None, None]) s1 = helper.make_tensor("s1", TensorProto.INT64, [3], [3, 2, -1]) s2 = helper.make_tensor("s2", TensorProto.INT64, [3], [0, 0, 3]) n1 = helper.make_node("Reshape", ["X", "s1"], ["mid"]) n2 = helper.make_node("Reshape", ["mid", "s2"], ["Y"], allowzero=1) m = helper.make_model(helper.make_graph([n1, n2], "g", [X], [Y], initializer=[s1, s2]), opset_imports=[helper.make_opsetid("", 18)]) inp = np.random.default_rng(7).random((0, 6, 2), dtype=np.float32) print("REF:", onnx.reference.ReferenceEvaluator(m).run(None, {"X": inp})[0].shape) print("ORT:", ort.InferenceSession(m.SerializeToString(), providers=["CPUExecutionProvider"]).run(None, {"X": inp})[0].shape) ``` Output on `main` (`40c9f85f69`): ``` REF: (0, 0, 3) [W ... graph.cc:122 MergeShapeInfo] Error merging shape info for output. 'Y' source:{0,6,3} target:{0,0,3}. Falling back to lenient merge. ORT: (0, 6, 3) ❌ ``` ### Fix Setting `allowzero=1` on the fused node would also work but requires opset >= 14, which this transformer cannot assume (it accepts `Reshape` opset 5+). Bail out of fusion conservatively when `shape_value` contains any literal `0` dim. ### Test Adds `ReshapeFusionContiguousReshapesWithZeroDim` that builds the bug repro programmatically and asserts: - the two reshapes are NOT collapsed - the inferred output shape stays `(0, 0, 3)` The existing happy-path test `ReshapeFusion_Contiguous_Reshape` (added in #22494) is unaffected — its inferred output shape `(2, 1, 64, 32)` contains no zero dims, so the new guard does not trigger. ### Provenance `FuseContiguousReshapes` was introduced in #22494 (Feb 2025). The bug has been latent in `main` since then. ### Motivation and Context Found while reviewing https://github.com/microsoft/onnxscript/pull/2907 — the rewriter rule under test there is semantically correct, but its numerical-equivalence check using ORT as the oracle fails because of this fusion bug. Fixes #28348. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Author
Parents
Loading