Skip to content

Fix torch gelu converter ignoring the approximate keyword argument - #2863

Open
Yigtwxx wants to merge 1 commit into
apple:mainfrom
Yigtwxx:gelu-approximate-kwarg
Open

Yigtwxx wants to merge 1 commit into
apple:mainfrom
Yigtwxx:gelu-approximate-kwarg

Conversation

@Yigtwxx

@Yigtwxx Yigtwxx commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

Summary

The torch gelu converter only reads approximate when it arrives as a positional input, which is how TorchScript serializes it. torch.export and the ExecuTorch edge dialect keep it as a keyword argument (aten.gelu.default(x, approximate='tanh')), so nn.GELU(approximate="tanh") silently converts to the exact gelu on those frontends.

This affects every model that uses the tanh variant, most visibly the Hugging Face gelu_pytorch_tanh activation (ACT2FN["gelu_pytorch_tanh"] wraps nn.functional.gelu(approximate="tanh")), which is the default activation of the Gemma / Gemma 2 / Gemma 3 configs and of SigLIP and RecurrentGemma among others. The two gelu variants differ by up to ~5e-4 per activation, and the existing test_gelu runs with atol=1e-3, which is why the mismatch never showed up.

Minimal reproduction on main:

import torch, coremltools as ct

model = torch.nn.GELU(approximate="tanh").eval()
x = torch.randn(2, 8)
ep = torch.export.export(model, (x,)).run_decompositions({})
prog = ct.convert(ep, inputs=[ct.TensorType(shape=x.shape)], convert_to="milinternal")
print(prog.functions["main"].find_ops(op_type="gelu")[0].mode.val)
# main: EXACT        (torch.jit.trace of the same module gives TANH_APPROXIMATION)
# this PR: TANH_APPROXIMATION

Implementation

  • gelu now falls back to _get_kwinputs(context, node, "approximate") when the argument is not positional, unwraps a Var, and maps "tanh" to TANH_APPROXIMATION. Any value other than "none" / "tanh" raises a ValueError instead of tripping an assertion.
  • "none" and "tanh" are added to TORCH_STRING_ARGS so the export frontend binds them without the "neither a name of existing var nor a torch string argument" warning.

TorchScript behaviour is unchanged.

Tests

test_gelu keeps its tolerance but now also asserts the mode of the emitted gelu op (EXACT vs TANH_APPROXIMATION) for every frontend. On main the new assertion fails for the TORCHEXPORT and EXECUTORCH approximate="tanh" cases and passes for TorchScript; with this change all 45 test_gelu cases pass, and TestActivation shows no new failures compared to main.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant