From b19bca4207351a11ef9e5f0f83d7d76b36d79a44 Mon Sep 17 00:00:00 2001 From: Yigtwxx Date: Thu, 17 Sep 2026 16:47:00 +0300 Subject: [PATCH] Fix torch gelu converter ignoring the approximate keyword argument --- .../converters/mil/frontend/torch/ops.py | 26 ++++++++++++------- .../mil/frontend/torch/test/test_torch_ops.py | 8 +++++- 2 files changed, 24 insertions(+), 10 deletions(-) diff --git a/coremltools/converters/mil/frontend/torch/ops.py b/coremltools/converters/mil/frontend/torch/ops.py index 6a32bebc5..153afcf14 100644 --- a/coremltools/converters/mil/frontend/torch/ops.py +++ b/coremltools/converters/mil/frontend/torch/ops.py @@ -80,6 +80,10 @@ # searchsorted side "left", "right", + + # gelu approximate + "none", + "tanh", } @@ -6429,16 +6433,20 @@ def hardsigmoid(context, node): @register_torch_op def gelu(context, node): - inputs = _get_inputs(context, node) - assert len(inputs) in (1, 2) + inputs = _get_inputs(context, node, expected=(1, 2)) + x = inputs[0] + # torch script serializes approximate positionally, torch.export keeps it as a keyword + approximate = inputs[1] if len(inputs) == 2 else None + approximate = _get_kwinputs(context, node, "approximate", default=[approximate])[0] + if isinstance(approximate, Var): + approximate = approximate.val + mode = None - if len(inputs) == 2: - approximate = inputs[1].val - if approximate == "tanh": - mode = "TANH_APPROXIMATION" - else: - assert approximate == "none" - res = mb.gelu(x=inputs[0], mode=mode, name=node.name) + if approximate == "tanh": + mode = "TANH_APPROXIMATION" + elif approximate not in (None, "none"): + raise ValueError(f"gelu: unsupported approximate mode {approximate}") + res = mb.gelu(x=x, mode=mode, name=node.name) context.add(res) diff --git a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py index 14c515be4..acfd14ea0 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -7076,7 +7076,7 @@ def test_hardswish(self, compute_unit, backend, frontend, shape, minimum_deploym def test_gelu(self, compute_unit, backend, frontend, shape, approximate): model = nn.GELU() if approximate is None else nn.GELU(approximate=approximate) model = model.eval() - self.run_compare_torch( + res = self.run_compare_torch( shape, model, atol=1e-3, @@ -7085,6 +7085,12 @@ def test_gelu(self, compute_unit, backend, frontend, shape, approximate): backend=backend, compute_unit=compute_unit, ) + # The tolerance above is loose enough to hide the difference between the exact + # and the tanh-approximated gelu, so check the mode of the emitted op explicitly + gelu_ops = res[1]._mil_program.find_ops(op_type="gelu") + assert len(gelu_ops) == 1 + expected_mode = "TANH_APPROXIMATION" if approximate == "tanh" else "EXACT" + assert gelu_ops[0].mode.val == expected_mode @pytest.mark.parametrize( "compute_unit, backend, frontend, inplace",