Skip to content

Commit 3b065a6

Browse files
authored
Add fused rope pass (pytorch#22918)
Summary: FuseRoPEPass folds the HuggingFace rotate_half rotary-embedding apply into a single native::rope op. This diff adds the op as well as the pass. custom_ops.py registers `native::rope(Tensor input, Tensor cos, Tensor sin, bool interleaved=False) -> Tensor`, with a fake (meta) impl used by AOT export/serialization and a CompositeExplicitAutograd reference body used for eager execution. Registration is an import side effect, so the module is imported from the backend entry points to guarantee the op exists before lowering. HF apply_rotary_pos_emb lowers each rotated tensor to `x*cos + rotate_half(x)*sin`, where `rotate_half(x) = cat([-x2, x1], -1)` splits the last dim in half. The pass matches that ~7-op chain (an add over two muls, over cat / neg / slice_copy) and rewrites it to one `native::rope(x, cos, sin, interleaved=False)`. Partial rotary is folded too: when only the leading rotary_dim channels are rotated and the apply result is concatenated with the pass-through `x[..., rotary_dim:]`, the whole slice/apply/cat is replaced by a single native::rope over the full tensor. The op rotates the first `cos.shape[-1]` channels and passes the remainder through, so the rope input stays contiguous with no leftover strided slice. cos/sin are captured as the already-scaled tensors feeding the multiplies, so any rope variant (default, YaRN, long-rope) that ends in this apply is handled without the pass modelling frequency construction. Only the split-half layout is matched; the interleaved (GPT-NeoX "traditional") layout is left decomposed, though the op implements it for completeness. Every intermediate node must be single-use so the decomposition DCEs away after the rewrite. The pass is added to get_default_passes(), so it runs pre-partition on the edge-dialect graph, and the partitioner claims native::rope via _SUPPORTED_NON_CORE_OPS. Reviewed By: digantdesai Differential Revision: D111753295 Pull Request resolved: pytorch#22918
1 parent 89a8e59 commit 3b065a6

7 files changed

Lines changed: 1000 additions & 1 deletion

File tree

‎backends/native/BUCK‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,11 @@ fbcode_target(
1515
name = "lib",
1616
srcs = [
1717
"__init__.py",
18+
"custom_ops.py",
1819
"partitioner.py",
1920
"passes/__init__.py",
21+
"passes/_utils.py",
22+
"passes/fuse_rope.py",
2023
"passes/reinplace.py",
2124
"passes/replace_copy_with_alias.py",
2225
"preprocess.py",

‎backends/native/custom_ops.py‎

Lines changed: 107 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,107 @@
1+
# Copyright (c) Meta Platforms, Inc. and affiliates.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the BSD-style license found in the
5+
# LICENSE file in the root directory of this source tree.
6+
7+
# pyre-strict
8+
9+
"""
10+
Custom ops for the native backend.
11+
12+
Registers ``native::rope``, the fused rotary position embedding op produced by
13+
``FuseRoPEPass``. Registration happens as an import side effect, so this module is
14+
imported from the backend's entry points (``passes`` / ``partitioner.py``) to
15+
guarantee the op exists before lowering. The reference ``CompositeExplicitAutograd``
16+
body is only used for eager execution; AOT export/serialization uses the fake
17+
(meta) impl.
18+
"""
19+
20+
import torch
21+
from torch import Tensor
22+
23+
lib: torch.library.Library = torch.library.Library("native", "DEF")
24+
25+
26+
lib.define(
27+
"rope(Tensor input, Tensor position_ids, Tensor inv_freq, "
28+
"bool interleaved=False, float attention_scale=1.0) -> Tensor"
29+
)
30+
31+
32+
def _check_rope_shapes(input: Tensor, position_ids: Tensor, inv_freq: Tensor) -> None:
33+
torch._check(input.is_floating_point(), lambda: "rope input must be floating point")
34+
torch._check(input.dim() == 4, lambda: "rope input must have shape [B, H, T, D]")
35+
torch._check(
36+
position_ids.dim() == 2, lambda: "rope position_ids must have shape [Bpos, T]"
37+
)
38+
torch._check(inv_freq.dim() == 1, lambda: "rope inv_freq must have shape [R / 2]")
39+
torch._check(
40+
(position_ids.shape[0] == 1) | (position_ids.shape[0] == input.shape[0]),
41+
lambda: "rope position batch must be 1 or match the input batch",
42+
)
43+
torch._check(
44+
position_ids.shape[1] == input.shape[2],
45+
lambda: "rope position sequence length must match the input",
46+
)
47+
torch._check(
48+
2 * inv_freq.shape[0] <= input.shape[3],
49+
lambda: "rope rotary width must not exceed the input width",
50+
)
51+
52+
53+
@torch.library.register_fake("native::rope", lib=lib)
54+
def _rope_fake(
55+
input: Tensor,
56+
position_ids: Tensor,
57+
inv_freq: Tensor,
58+
interleaved: bool = False,
59+
attention_scale: float = 1.0,
60+
) -> Tensor:
61+
_check_rope_shapes(input, position_ids, inv_freq)
62+
return input.new_empty(input.shape, dtype=input.dtype)
63+
64+
65+
@torch.library.impl("native::rope", "CompositeExplicitAutograd", lib=lib)
66+
def _rope_impl(
67+
input: Tensor,
68+
position_ids: Tensor,
69+
inv_freq: Tensor,
70+
interleaved: bool = False,
71+
attention_scale: float = 1.0,
72+
) -> Tensor:
73+
"""Rotate the leading R channels of [B, H, T, D], preserving the tail.
74+
75+
Positions [Bpos, T] broadcast over heads and, when Bpos is 1, batches.
76+
inv_freq [R / 2] determines the rotary width. Phase, trig and attention
77+
scaling use FP32; the tables are then cast to the input dtype.
78+
"""
79+
_check_rope_shapes(input, position_ids, inv_freq)
80+
# Elementwise multiplication keeps phase construction in FP32 under autocast,
81+
# unlike the equivalent batched matmul used by HF's source pattern.
82+
phase = position_ids.float().unsqueeze(-1) * inv_freq.float()
83+
if interleaved:
84+
phase = phase.repeat_interleave(2, dim=-1)
85+
else:
86+
phase = torch.cat((phase, phase), dim=-1)
87+
cos = (phase.cos() * attention_scale).to(input.dtype).unsqueeze(1)
88+
sin = (phase.sin() * attention_scale).to(input.dtype).unsqueeze(1)
89+
rotary = 2 * inv_freq.shape[0]
90+
x_rot = input[..., :rotary]
91+
x_pass = input[..., rotary:]
92+
if interleaved:
93+
x1 = x_rot[..., 0::2]
94+
x2 = x_rot[..., 1::2]
95+
rotated = torch.stack((-x2, x1), dim=-1).flatten(-2)
96+
else:
97+
half = rotary // 2
98+
x1 = x_rot[..., :half]
99+
x2 = x_rot[..., half:]
100+
rotated = torch.cat((-x2, x1), dim=-1)
101+
out = x_rot * cos + rotated * sin
102+
if x_pass.shape[-1] == 0:
103+
return out
104+
return torch.cat((out, x_pass), dim=-1)
105+
106+
107+
rope_op: torch._ops.OpOverload = torch.ops.native.rope.default

‎backends/native/partitioner.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
# Registers torch.ops.torchao.dequantize_gguf, referenced in _SUPPORTED_NON_CORE_OPS.
1919
import executorch.extension.llm.export.gguf # noqa: F401
2020
import torch
21+
from executorch.backends.native.custom_ops import rope_op
2122
from executorch.backends.native.passes import backend_inplace_aten_variants
2223

2324
from executorch.exir.backend.compile_spec_schema import CompileSpec
@@ -43,6 +44,7 @@
4344
# GGUF weight dequantize stays in the delegate; the serializer folds it into a
4445
# PackedQuant weight on the consuming op.
4546
torch.ops.torchao.dequantize_gguf.default,
47+
rope_op,
4648
torch.ops.aten.rms_norm.default,
4749
]
4850

‎backends/native/passes/__init__.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929

3030
from typing import List, Union
3131

32+
from executorch.backends.native.passes.fuse_rope import FuseRoPEPass
3233
from executorch.backends.native.passes.reinplace import (
3334
backend_inplace_aten_variants,
3435
BACKEND_INPLACE_OPS,
@@ -54,6 +55,7 @@
5455
"CollapseViewCopyPass",
5556
"FuseGQAWithSDPAPass",
5657
"FuseRMSNormPass",
58+
"FuseRoPEPass",
5759
"get_default_passes",
5860
"NativeReinplacePass",
5961
"NormalizeSDPAInputRankPass",
@@ -69,6 +71,7 @@ def get_default_passes() -> List[Union[ExportPass, ExportedProgramPassBase]]:
6971
their in-place edge forms.
7072
"""
7173
return [
74+
FuseRoPEPass(),
7275
FuseGQAWithSDPAPass(),
7376
NormalizeSDPAInputRankPass(),
7477
FuseRMSNormPass(fold_dtype_casts=True, allow_lossy_weight_casts=True),

‎backends/native/passes/_utils.py‎

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
# Copyright (c) Meta Platforms, Inc. and affiliates.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the BSD-style license found in the
5+
# LICENSE file in the root directory of this source tree.
6+
7+
"""Shared helpers for native backend graph passes."""
8+
9+
import torch
10+
from torch.fx import Node
11+
12+
13+
def _resolve_aten(target: object) -> "torch._ops.OpOverload | None":
14+
"""Return the underlying aten OpOverload for a node target, or None.
15+
16+
Handles both edge-dialect ops (EdgeOpOverload, whose ``_op`` is the aten op)
17+
and plain aten OpOverloads (whose ``_op`` is the C++ builtin, so it must not
18+
be unwrapped). See graph_serialize._resolve_op_overload for the same logic.
19+
"""
20+
inner = getattr(target, "_op", None)
21+
if isinstance(inner, torch._ops.OpOverload):
22+
return inner
23+
if isinstance(target, torch._ops.OpOverload):
24+
return target
25+
return None
26+
27+
28+
def _single_user(node: object) -> bool:
29+
return isinstance(node, Node) and len(node.users) == 1
30+
31+
32+
def _fake_tensor(node: object) -> "torch.Tensor | None":
33+
val = node.meta.get("val") if isinstance(node, Node) else None
34+
return val if isinstance(val, torch.Tensor) else None

0 commit comments

Comments
 (0)