Reorder channelwise gated delta rule chunked hot loops for autovectorization (#21021)#21021
Reorder channelwise gated delta rule chunked hot loops for autovectorization (#21021)#21021JakeStevens wants to merge 3 commits into
Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/21021
Note: Links to docs will display an error until the docs builds have been completed. ❌ 2 New Failures, 2 Unrelated FailuresAs of commit 3a6025b with merge base 2524691 ( NEW FAILURES - The following jobs have failed:
FLAKY - The following job failed but was likely due to flakiness present on trunk:
BROKEN TRUNK - The following job 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. |
|
@JakeStevens has exported this pull request. If you are a Meta employee, you can view the originating Diff in D112598714. |
This PR needs a
|
…ization (pytorch#21021) Summary: Reorder the chunked prefill inner loops (steps 1, 4, 5, 6) so the innermost loop runs contiguously over the head dimension (k or v) instead of striding down a column of the state / pv. This lets the compiler autovectorize the now-unit-stride AXPYs; hand-written at::vec was tried and was slower than the compiler output, so the loops stay scalar. Differential Revision: D112598714
b8d62b5 to
dac4505
Compare
…ization (pytorch#21021) Summary: Reorder the chunked prefill inner loops (steps 1, 4, 5, 6) so the innermost loop runs contiguously over the head dimension (k or v) instead of striding down a column of the state / pv. This lets the compiler autovectorize the now-unit-stride AXPYs; hand-written at::vec was tried and was slower than the compiler output, so the loops stay scalar. Reviewed By: billmguo Differential Revision: D112598714
dac4505 to
f5857be
Compare
…h#21020) Summary: Reduce the channelwise gated delta rule kernel from four state traversals to two per token. The first pass computes the predicted value while folding in decay; the second applies decay and the rank-one update while accumulating the output. This preserves operation grouping and numerical behavior while reducing state-memory traffic. Add a standalone microbenchmark covering decode and representative prefill lengths. Reviewed By: billmguo Differential Revision: D112596724
Summary: Route the channelwise gated delta rule by sequence length: T == 1 keeps the two-pass token recurrence for autoregressive decode, while T != 1 uses a chunkwise WY/UT formulation for prefill. The chunked path computes per-channel log-decay prefixes, causal query/key terms, the beta-folded triangular transform, WY pseudo-keys and pseudo-values, and inter-chunk state carry. It handles a ragged final chunk without a separate tail implementation. Parallelize independent (batch, head) work across the ExecuTorch threadpool. Each worker receives a disjoint slice of one temporary scratch arena, avoiding shared mutable buffers while amortizing allocation across chunks. Reviewed By: billmguo Differential Revision: D112597348
…ization (pytorch#21021) Summary: Reorder the chunked prefill inner loops (steps 1, 4, 5, 6) so the innermost loop runs contiguously over the head dimension (k or v) instead of striding down a column of the state / pv. This lets the compiler autovectorize the now-unit-stride AXPYs; hand-written at::vec was tried and was slower than the compiler output, so the loops stay scalar. Reviewed By: billmguo Differential Revision: D112598714
f5857be to
3a6025b
Compare
Summary:
Reorder the chunked prefill inner loops (steps 1, 4, 5, 6) so the innermost loop runs contiguously over the head dimension (k or v) instead of striding down a column of the state / pv. This lets the compiler autovectorize the now-unit-stride AXPYs; hand-written at::vec was tried and was slower than the compiler output, so the loops stay scalar.
Reviewed By: billmguo
Differential Revision: D112598714