Repository navigation
add hadamard option to low precision attention api - #4194
Conversation
[ghstack-poisoned]
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4194
Note: Links to docs will display an error until the docs builds have been completed. ✅ You can merge normally! (8 Unrelated Failures)As of commit 3b80af2 with merge base f11eff8 ( FLAKY - The following job failed but was likely due to flakiness present on trunk:
BROKEN TRUNK - The following jobs failed but was present on the merge base:👉 Rebase onto the `viable/strict` branch to avoid these failures
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
|
@claude /pr-review detailed Please do a deep review focused on subtle correctness bugs and regressions that are easy for human reviewers to miss. Prioritize findings over summary. Specifically, look for: incorrect edge-case behavior |
|
@claude Specifically, look for: incorrect edge-case behavior |
|
Claude finished @drisspg's task in 4m 27s —— View job Deep Correctness Review — PR #4194Reviewed all 16 changed files focusing on subtle correctness bugs, edge cases, and regressions. Bug:
|
| Finding | Severity | Type |
|---|---|---|
hadamard=True (bool) in benchmark |
Medium | Type bug |
tl.debug_barrier() for production sync |
High | Correctness risk |
| Zero test coverage for hadamard paths | High | Testing gap |
kv_temp buffer reuse assumption |
Medium | Fragile correctness |
| Global memory for intra-block shuffle | Medium | Performance/correctness |
| Last-chunk loop overhead | Low | Performance |
_compute_num_chunks interface change |
Low | Refactoring |
str(HadamardMode.QKV) fragility |
Low | API design |
The core Hadamard butterfly math is correct. The main concerns are the reliance on tl.debug_barrier() for production synchronization, the complete absence of tests, and the type mismatch in the benchmark.
[ghstack-poisoned]
|
@drisspg Can you take another look? Just addressed some of the claude review comments:
A lot of the other comments I think aren't really applicable in my opinion or are wrong, but happy to look into it more if you think I should. |
| sqnr = compute_error(out_ref, out_fp8) | ||
| self.assertGreater( | ||
| sqnr.item(), | ||
| 25.0, |
There was a problem hiding this comment.
should we expect/check sqnr to be higher w/ hadamard?
There was a problem hiding this comment.
The tests here were more just making sure Hadamard was calculated correctly (with inverse and stuff). Because it's uniformly sampled, Hadamard doesn't actually improve SQNR. It only improves when there's outliers. I just added a separate set of tests to check that Hadamard improves accuracy when the inputs have outliers.
| try: | ||
| with torch.no_grad(): | ||
| out_fp8 = fp8_fa3_rope_sdpa(q, k, v, cos, sin, is_causal=False) | ||
| out_fp8 = fp8_fa3_rope_sdpa( |
## Summary - Added Hadamard on QKV tensor support for the low precision attention API, passed through from apply_low_precision_attention to select the hadamard fused kernel - Added new kernels (triton_hadamard_qkv_quantization.py and triton_hadamard_rope_qkv_quantization.py) in triton for fused hadamard and QKV quantization (with rope fusion as well) - Because of the way per-head quantization works with the sequence chunks, we need to store the hadamard outputs in a temp buffer, which eliminates some of the benefit of the fusion. However, it still saves one global read of the QKV tensors, which experiments show still benefits runtime quite a bit, so the fusion is still much better than just running hadamard separately. - Added hadamard option to the benchmarks - Replaced some duplicate compute_num_chunks code ## Results #### Single Attention Layer <img width="692" height="314" alt="image" src="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/user-attachments/assets/e31952d4-5cca-4677-8f3f-65b892ea6b23" /> For a single attention layer, it got slower (from 1.36x speedup to 1.29x speedup on the highest sequence length). The SQNR does not really improve because we're testing with random tensors, which are pretty uniform already (Hadamard is intended to spread out intensity for better quantization accuracy) #### LLaMA3 Prefill <img width="695" height="240" alt="image" src="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/user-attachments/assets/4c46259b-2eee-4700-bd82-a653abaa2cde" /> Perplexity used to go from 7.54 -> 7.62. Now it is noticeably better, going from 7.54 -> 7.57. The speedup dropped from 1.23x to 1.15x at the highest sequence length.
Stack from ghstack (oldest at bottom):
Summary
Results
Single Attention Layer
LLaMA3 Prefill