huggingface/transformers · #48191
[ONNX] Skip affected models on torch 2.13 (two dynamo regressions)
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.