Skip to content

Add WAN style rope detection pattern for low precision attention API - #4173

Merged
howardzhang-cv merged 3 commits into
mainfrom
gh/howardzhang-cv/37/head
Mar 26, 2026
Merged

howardzhang-cv merged 3 commits into
mainfrom
gh/howardzhang-cv/37/head

Conversation

@howardzhang-cv

@howardzhang-cv howardzhang-cv commented Mar 25, 2026 •

Copy link
Copy Markdown
Contributor

Stack from ghstack (oldest at bottom):

Quick Summary

  • Added rope fusion detection for WAN style models that use RoPE. This is for the low precision attention API, to enable rope fusion.
  • Previously, the fallback path would work, allowing for torch.compile to work with the monkey patched F.SDPA. However, the RoPE operation was not fused into a single kernel with the quantization to fp8.

Results

Previous results on WAN model:

Config Median Time (s) Speedup
bf16 baseline 100.85 1.00x
fp8_attn 67.15 1.50x
bf16 baseline + torch.compile 81.32 1.00x
fp8_attn + compile 46.88 1.73x

New results with rope fusion:

Config Median Time (s) Speedup
fp8_attn + rope fusion + compile 44.24 1.84x

[ghstack-poisoned]
[ghstack-poisoned]
@pytorch-bot

pytorch-bot Bot commented Mar 25, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4173

Note: Links to docs will display an error until the docs builds have been completed.

✅ No Failures

As of commit 0b9efbb with merge base ac0b820 (image):
💚 Looks good so far! There are no failures yet. 💚

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Mar 25, 2026
howardzhang-cv added a commit that referenced this pull request Mar 25, 2026
@howardzhang-cv howardzhang-cv added the module: not user facing Use this tag if you don't want this PR to show up in release notes label Mar 25, 2026
Comment on lines +695 to +715
def _get_setitem_stride2_slice(node: Node) -> Optional[Tuple[int, Node]]:
"""Check if node is ``operator.setitem(tensor, (..., slice(start, None, 2)), value)``.

Returns ``(start, value_node)`` or None.
"""
if not _is_op(node, operator.setitem):
return None
if len(node.args) < 3:
return None
idx = node.args[1]
value = node.args[2]
if not isinstance(value, Node):
return None
# idx should be (Ellipsis, slice(start, None, 2))
if not isinstance(idx, tuple) or len(idx) < 1:
return None
last = idx[-1]
if not isinstance(last, slice) or last.step != 2:
return None
# Check preceding entries are Ellipsis or slice(None)
for entry in idx[:-1]:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

these should share ode right?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I debated that but kept it separate just for cleaner separation of operations when we're trying to pattern match rope. Like for example, when we pattern match, it'll say ("find setitem operations with stride 2 slices, do stuff etc., find getitem operations with stride 2 slices"). I thought it was slightly cleaner to keep them separate.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Actually, I just thought about it a bit more and I think it's better to combine them. Probably better to avoid the case where someone changes one but forgets to change the other.

@drisspg drisspg left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think that if we are goin to add more patterns we should have some tests for the patterns we want to match against besides just integration tests

[ghstack-poisoned]
howardzhang-cv added a commit that referenced this pull request Mar 26, 2026
@howardzhang-cv

Copy link
Copy Markdown
Contributor Author

I think that if we are goin to add more patterns we should have some tests for the patterns we want to match against besides just integration tests

Good point. I'll work on setting those up. I'll do it in a separate PR to keep it clean.

@howardzhang-cv
howardzhang-cv changed the base branch from gh/howardzhang-cv/37/base to main March 26, 2026 02:28
@howardzhang-cv
howardzhang-cv merged commit 6f56403 into main Mar 26, 2026
35 of 36 checks passed
Priyjain-amd pushed a commit to Priyjain-amd/ao that referenced this pull request May 26, 2026
…ytorch#4173)

## Quick Summary
- Added rope fusion detection for WAN style models that use RoPE. This is for the low precision attention API, to enable rope fusion.
- Previously, the fallback path would work, allowing for torch.compile to work with the monkey patched F.SDPA. However, the RoPE operation was not fused into a single kernel with the quantization to fp8.

## Results
Previous results on WAN model:
| Config | Median Time (s) | Speedup |
| --- | --- | --- |
| bf16 baseline | 100.85 | 1.00x |
| fp8_attn | 67.15 | **1.50x** |
| bf16 baseline + torch.compile | 81.32 | 1.00x |
| fp8_attn + compile | 46.88 | **1.73x** |

New results with rope fusion:
| Config | Median Time (s) | Speedup |
| --- | --- | --- |
| fp8_attn + rope fusion + compile | 44.24 | **1.84x** |
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. module: not user facing Use this tag if you don't want this PR to show up in release notes

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants