Repository navigation
Add WAN style rope detection pattern for low precision attention API - #4173
Conversation
[ghstack-poisoned]
🔗 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 FailuresAs of commit 0b9efbb with merge base ac0b820 ( This comment was automatically generated by Dr. CI and updates every 15 minutes. |
| 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]: |
There was a problem hiding this comment.
these should share ode right?
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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
left a comment
There was a problem hiding this comment.
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. |
…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** |
Stack from ghstack (oldest at bottom):
Quick Summary
Results
Previous results on WAN model:
New results with rope fusion: