Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 17 additions & 9 deletions coremltools/converters/mil/frontend/torch/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,10 @@
# searchsorted side
"left",
"right",

# gelu approximate
"none",
"tanh",
}


Expand Down Expand Up @@ -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)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
Expand Down