huggingface/transformers · #48191

[ONNX] Skip affected models on torch 2.13 (two dynamo regressions)

ydshieh · merged Aug 27, 20261 files · 102 + / 0
src/transformers/exporters/exporter_onnx.py102 + / 0
@@ -727,6 +727,108 @@ def _fix_sort_stable(gm: torch.fx.GraphModule, node: torch.fx.Node) -> bool:     return True  +@functools.cache+def _integral_scalar_promotion_ops() -> frozenset:+    """Ops where a Python float meeting an integral tensor needs the tensor promoted first.++    Named for torch 2.13, where the mishandling appeared, but applied on every version: both rewrites are+    semantics-preserving — a cast to the dtype the op already produces, and an overload swap with the same+    meaning — so gating them on a version would add a branch that changes nothing except which torch the+    path is exercised on.++    `sub`/`rsub` and `mul` are what the affected models spell (`1.0 - attention_mask`, `mask * 2.0`); both+    overloads appear, since decomposition rewrites `.Tensor` to `.Scalar` when the operand is a constant.++    Resolved on first use rather than at import: this module is importable without torch, and naming an+    `OpOverload` at module scope breaks that.+    """+    return frozenset(+        {+            torch.ops.aten.rsub.Scalar,+            torch.ops.aten.sub.Scalar,+            torch.ops.aten.sub.Tensor,+            torch.ops.aten.mul.Scalar,+            torch.ops.aten.mul.Tensor,+        }+    )+++@register_fx_node_fix("onnx")+def _fix_integral_tensor_float_scalar(gm: torch.fx.GraphModule, node: torch.fx.Node) -> bool:+    """Promote an integral tensor before it meets a Python float, which torch 2.13 mishandles.++    Two torch 2.13 regressions have the same shape — a float scalar against an *integral* tensor, whose+    promotion the export pipeline no longer gets right:++    - `1.0 - int_mask` (`aten.rsub.Scalar`) crashes the decomposition pass (pytorch/pytorch#194381),+    - `int_mask * 2.0` (`aten.mul.Tensor`, `aten.mul.Scalar` after decomposition) reaches translation with+      no ONNX decomposition registered for it (pytorch/pytorch#194382).++    Both go away once the tensor is already the dtype the op produces: `1.0 - float_tensor` and+    `float_tensor * 2.0` export fine on the same torch. So rather than rewriting the op — which would mean+    building the constant as a tensor and picking the right overload — cast its tensor operand up front and+    leave the op alone. The cast's value is the op's own output for these elementwise cases (same shape,+    the promoted dtype), so it carries `node.meta` unchanged.++    Self-limiting: once the operand is floating point the predicate no longer matches, so the walk cannot+    revisit it.+    """+    if node.target not in _integral_scalar_promotion_ops():+        return False+    if len(node.args) < 2:+        return False+    tensor_arg, scalar_arg = node.args[0], node.args[1]+    # A tensor on the left, a Python float on the right — a `Node` there is already a real tensor operand.+    if not isinstance(tensor_arg, torch.fx.Node) or not isinstance(scalar_arg, float):+        return False+    operand, result = tensor_arg.meta.get("val"), node.meta.get("val")+    if operand is None or result is None:+        return False+    if operand.dtype.is_floating_point or not result.dtype.is_floating_point:+        return False+    with gm.graph.inserting_before(node):+        promoted = gm.graph.call_function(+            torch.ops.aten._to_copy.default, args=(tensor_arg,), kwargs={"dtype": result.dtype}+        )+    promoted.meta.update(node.meta)+    node.replace_input_with(tensor_arg, promoted)+    return True+++@register_fx_node_fix("onnx")+def _fix_mul_scalar_symbolic(gm: torch.fx.GraphModule, node: torch.fx.Node) -> bool:+    """Rewrite `mul.Scalar` to `mul.Tensor` when its 'scalar' is a graph node.++    The other half of pytorch/pytorch#194382: torchlib registers no real-valued `aten.mul.Scalar`+    translation at all, and decomposition also produces that overload with a *symbolic* second operand — a+    division result rather than a literal (`mul.Scalar(x, %truediv_1)`) — which the promotion fix above+    cannot address, since there is no Python constant to promote against. `mul.Tensor` has the two-operand+    translation, which is the same rewrite `_fix_remainder_scalar` makes for the same reason.++    That operand is a `SymFloat`, not a tensor: across the affected families every one of the 33 sites is+    an `operator.truediv` result. `mul.Tensor`'s translation takes a symbolic scalar there — it becomes a+    graph value like any other — so the rewrite is sound, but the guard below names both accepted forms+    rather than trusting that a `Node` implies a tensor. Anything else (a nested list, an unbacked value+    with no `val`) is left as `mul.Scalar` to fail visibly in translation instead of silently here.++    Reached because the FX fixes run a second time right after `run_decompositions`, where this overload+    appears.+    """+    if node.target is not torch.ops.aten.mul.Scalar:+        return False+    if len(node.args) < 2 or not isinstance(node.args[1], torch.fx.Node):+        return False+    other = node.args[1].meta.get("val")+    if not isinstance(other, (torch.Tensor, torch.SymFloat, torch.SymInt, torch.SymBool)):+        return False+    with gm.graph.inserting_before(node):+        new = gm.graph.call_function(torch.ops.aten.mul.Tensor, args=node.args)+    new.meta.update(node.meta)+    node.replace_all_uses_with(new)+    gm.graph.erase_node(node)+    return True++ @register_fx_node_fix("onnx") def _fix_remainder_scalar(gm: torch.fx.GraphModule, node: torch.fx.Node) -> bool:     """Rewrite remainder.Scalar to remainder.Tensor when the 'scalar' arg is actually a tensor.