Skip to content

Commit 1fbe60b

Browse files
authored
XNNPACK: Route tanh GELU to approxgelu (pytorch#22851)
### Summary The XNNPACK backend currently ignores GELU's `approximate` argument and serializes both modes as `XNNGelu`. As a result, `nn.GELU(approximate="tanh")` executes exact GELU after delegation. For FP32 input `[-2.7]`, the baseline differs from the tanh reference by approximately `4.73e-4`, failing comparison at `atol=rtol=1e-5`. This PR serializes tanh GELU as `XNNApproxGelu` and dispatches it to XNNPACK's existing `xnn_unary_approxgelu` operation. Default/exact GELU and FP16 fallback remain unchanged. The new node is appended to both FlatBuffer unions to preserve existing node IDs. Affected XNNPACK models need to be re-exported, and the new node requires an updated runtime. ### Test plan I tested the fix on macOS arm64 using rebuilt runtimes based on `500849ba5b`. Baseline validation reproduced the FP32 tanh GELU mismatch at input `[-2.7]`, failing comparison at `atol=rtol=1e-5`. The current parameterized regression tests retain this input and tolerance. All 7 GELU tests passed, covering default GELU, both approximation modes, static/dynamic shapes, and FP16 fallback. Local lintrunner passed. Authored with assistance from OpenAI Codex. cc @GregoryComer @digantdesai @cbilgin @JakeStevens
1 parent 67dc0ad commit 1fbe60b

7 files changed

Lines changed: 84 additions & 35 deletions

File tree

‎backends/xnnpack/operators/op_gelu.py‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
register_node_visitor,
1313
)
1414
from executorch.backends.xnnpack.serialization.xnnpack_graph_schema import (
15+
XNNApproxGelu,
1516
XNNGelu,
1617
XNNGraph,
1718
XNode,
@@ -41,8 +42,15 @@ def define_node(
4142
# output
4243
output_id = vals_to_ids[node]
4344

45+
approximate = node.kwargs.get("approximate", "none")
46+
if approximate == "none":
47+
gelu_node_type = XNNGelu
48+
elif approximate == "tanh":
49+
gelu_node_type = XNNApproxGelu
50+
else:
51+
raise ValueError(f"Unsupported GELU approximation: {approximate}")
4452
ser_node = XNode(
45-
xnode_union=XNNGelu(
53+
xnode_union=gelu_node_type(
4654
input_id=input_id,
4755
output_id=output_id,
4856
flags=0,

‎backends/xnnpack/runtime/XNNCompiler.cpp‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1980,6 +1980,7 @@ _DEFINE_UNARY_NODE_NO_PARAMS(
19801980
xnn_unary_reciprocal_square_root)
19811981
_DEFINE_UNARY_NODE_NO_PARAMS(Ceiling, xnn_unary_ceiling)
19821982
_DEFINE_UNARY_NODE_NO_PARAMS(Gelu, xnn_unary_gelu)
1983+
_DEFINE_UNARY_NODE_NO_PARAMS(ApproxGelu, xnn_unary_approxgelu)
19831984
_DEFINE_UNARY_NODE_NO_PARAMS(Hardswish, xnn_unary_hardswish)
19841985
_DEFINE_UNARY_NODE_NO_PARAMS(Log, xnn_unary_log)
19851986
_DEFINE_UNARY_NODE_NO_PARAMS(Negate, xnn_unary_negate)
@@ -2021,6 +2022,7 @@ DefineNodeFunc getDefineNodeFunc(fb_xnnpack::XNodeUnion nodeType) {
20212022
_DEFINE(ReciprocalSquareRoot)
20222023
_DEFINE(Ceiling)
20232024
_DEFINE(Gelu)
2025+
_DEFINE(ApproxGelu)
20242026
_DEFINE(Hardswish)
20252027
_DEFINE(Log)
20262028
_DEFINE(Tanh)

‎backends/xnnpack/serialization/runtime_schema.fbs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -159,6 +159,7 @@ union XNodeUnion {
159159
XNNSin: _XNNNode1x1,
160160
XNNCopy: _XNNNode1x1,
161161
XNNCos: _XNNNode1x1,
162+
XNNApproxGelu: _XNNNode1x1,
162163
}
163164

164165
union XValueUnion {

‎backends/xnnpack/serialization/schema.fbs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -155,6 +155,7 @@ union XNodeUnion {
155155
XNNSin: _XNNNode1x1,
156156
XNNCopy: _XNNNode1x1,
157157
XNNCos: _XNNNode1x1,
158+
XNNApproxGelu: _XNNNode1x1,
158159
}
159160

160161
union XValueUnion {

‎backends/xnnpack/serialization/xnnpack_graph_schema.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -301,6 +301,11 @@ class XNNGelu(XNNNode1x1):
301301
pass
302302

303303

304+
@dataclass
305+
class XNNApproxGelu(XNNNode1x1):
306+
pass
307+
308+
304309
@dataclass
305310
class XNNHardswish(XNNNode1x1):
306311
pass
@@ -421,6 +426,7 @@ class XNNScaledDotProductAttention:
421426
XNNSin,
422427
XNNCopy,
423428
XNNCos,
429+
XNNApproxGelu,
424430
]
425431

426432

‎backends/xnnpack/test/BUCK‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -94,7 +94,9 @@ fbcode_target(_kind = runtime.python_test,
9494
]) + [
9595
"test_xnnpack_utils.py",
9696
],
97+
supports_static_listing = False,
9798
deps = [
99+
"fbsource//third-party/pypi/parameterized:parameterized",
98100
"//executorch/backends/test/harness:tester",
99101
"//executorch/backends/xnnpack/partition:xnnpack_partitioner",
100102
"//executorch/backends/xnnpack/quantizer:xnnpack_quantizer",

‎backends/xnnpack/test/ops/test_gelu.py‎

Lines changed: 63 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import torch
1010
from executorch.backends.xnnpack.test.tester import Tester
11+
from parameterized import parameterized
1112

1213

1314
def calculate_fp16_gelu_tolerance(ref_output_tensor):
@@ -30,62 +31,90 @@ def setUp(self):
3031
torch._dynamo.reset()
3132

3233
class Gelu(torch.nn.Module):
33-
def __init__(self):
34+
def __init__(self, approximate="none"):
3435
super().__init__()
35-
self.gelu = torch.nn.GELU()
36+
self.gelu = torch.nn.GELU(approximate=approximate)
3637

3738
def forward(self, x):
3839
return self.gelu(x)
3940

40-
def run_gelu_test(self, inputs):
41+
def run_gelu_test(
42+
self, inputs, *, approximate="none", dynamic=False, atol=None, rtol=None
43+
):
4144
input_tensor = inputs[0]
45+
is_fp16 = input_tensor.dtype == torch.float16
4246

43-
if input_tensor.dtype == torch.float16:
47+
if is_fp16:
4448
with torch.no_grad():
4549
ref_output = torch.nn.functional.gelu(
46-
input_tensor.to(torch.float32)
50+
input_tensor.to(torch.float32), approximate=approximate
4751
).to(torch.float16)
48-
atol, rtol = calculate_fp16_gelu_tolerance(ref_output)
52+
default_atol, default_rtol = calculate_fp16_gelu_tolerance(ref_output)
4953
else:
50-
atol = 1e-03
51-
rtol = 1e-03
54+
default_atol, default_rtol = 1e-3, 1e-3
5255

53-
(
54-
Tester(self.Gelu(), inputs)
56+
if atol is None:
57+
atol = default_atol
58+
if rtol is None:
59+
rtol = default_rtol
60+
61+
dynamic_shapes = (
62+
({0: torch.export.Dim("length", min=2, max=32)},) if dynamic else None
63+
)
64+
tester = (
65+
Tester(
66+
self.Gelu(approximate=approximate),
67+
inputs,
68+
dynamic_shapes=dynamic_shapes,
69+
)
5570
.export()
5671
.check_count({"torch.ops.aten.gelu.default": 1})
5772
.to_edge_transform_and_lower()
58-
.check_count({"torch.ops.higher_order.executorch_call_delegate": 1})
59-
.check_not(["executorch_exir_dialects_edge__ops_aten_gelu_default"])
60-
.to_executorch()
61-
.serialize()
62-
.run_method_and_compare_outputs(atol=atol, rtol=rtol)
6373
)
6474

65-
def test_fp16_gelu(self):
66-
# Older versions of XNNPACK don't support fp16 GELU.
67-
# TODO (gjcomer) Remove this when we update XNNPACK. (#16679)
68-
inputs = (torch.randn(20).to(torch.float16),)
69-
70-
with torch.no_grad():
71-
ref_output = torch.nn.functional.gelu(inputs[0].to(torch.float32)).to(
72-
torch.float16
73-
)
74-
atol, rtol = calculate_fp16_gelu_tolerance(ref_output)
75+
if is_fp16:
76+
# Older versions of XNNPACK don't support fp16 GELU.
77+
# TODO (gjcomer) Remove this when we update XNNPACK. (#16679)
78+
tester.check(
79+
["executorch_exir_dialects_edge__ops_aten_gelu_default"]
80+
).check_not(["torch.ops.higher_order.executorch_call_delegate"])
81+
else:
82+
tester.check_count(
83+
{"torch.ops.higher_order.executorch_call_delegate": 1}
84+
).check_not(["executorch_exir_dialects_edge__ops_aten_gelu_default"])
7585

7686
(
77-
Tester(self.Gelu(), inputs)
78-
.export()
79-
.check_count({"torch.ops.aten.gelu.default": 1})
80-
.to_edge_transform_and_lower()
81-
# Expect no delegation
82-
.check(["executorch_exir_dialects_edge__ops_aten_gelu_default"])
83-
.check_not(["torch.ops.higher_order.executorch_call_delegate"])
84-
.to_executorch()
87+
tester.to_executorch()
8588
.serialize()
86-
.run_method_and_compare_outputs(atol=atol, rtol=rtol)
89+
.run_method_and_compare_outputs(inputs=inputs, atol=atol, rtol=rtol)
8790
)
91+
if dynamic:
92+
tester.run_method_and_compare_outputs(
93+
inputs=(torch.linspace(-6, 6, 19),), atol=atol, rtol=rtol
94+
)
95+
96+
@parameterized.expand([("none",), ("tanh",)])
97+
def test_fp16_gelu(self, approximate):
98+
inputs = (torch.randn(20).to(torch.float16),)
99+
self.run_gelu_test(inputs, approximate=approximate)
88100

89101
def test_fp32_gelu(self):
90102
inputs = (torch.randn(20),)
91103
self.run_gelu_test(inputs)
104+
105+
@parameterized.expand(
106+
[
107+
(approximate, dynamic)
108+
for approximate in ("none", "tanh")
109+
for dynamic in (False, True)
110+
]
111+
)
112+
def test_fp32_gelu_approximation(self, approximate, dynamic):
113+
inputs = (torch.tensor([-6.0, -2.7, -1.0, 0.0, 1.0, 2.7, 6.0]),)
114+
self.run_gelu_test(
115+
inputs,
116+
approximate=approximate,
117+
dynamic=dynamic,
118+
atol=1e-5,
119+
rtol=1e-5,
120+
)

0 commit comments

Comments
 (0)