Arm backend: Add static floating-point TopK lowering - #23377
Conversation
Delegate static rank-2 FP16/FP32 TopK with constant K=1..4 along
the last dimension. Support largest=True and sorted=True, using FP
for K=1 and FP+INT for K>1.
Previously, TopK remained outside the Arm delegate. Existing range
analysis narrowed only index paths proven safe for int32:
scores -> TopK
|-- values (FP)
`-- indices (int64)
|-- model output / gather / int64 arithmetic
`-- cast(int32) -> proven-safe consumers
Route all supported TopK index paths through int32 before
partitioning. Restore int64 separately for consumers that need it,
allowing delegated gathers to coexist with portable consumers.
Keep TopK, tuple extraction and narrowing together in one delegate.
Decompose TopK into repeated ARGMAX with cumulative masking,
unrolling K selections at compile time. Mask each selected position
with -inf while preserving previous masks. Gather values from the
original scores using the selected indices.
After decomposition (simplified):
scores -> decomposed TopK
|-- values (FP)
`-- indices (int32)
|-- proven-safe int32 consumers
`-- cast(int64) per remaining consumer
`-- model output / portable gather /
int64 arithmetic
Require finite scores without checking finiteness at runtime.
Select equal scores in increasing index order, which may differ
from PyTorch's tie ordering.
Change-Id: I2b9a0b9e4e6a11d13a84c473161b0ad330e20b14
Signed-off-by: Yufeng Shi <yufeng.shi@arm.com>
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/23377
Note: Links to docs will display an error until the docs builds have been completed. ❌ 1 New Failure, 1 Cancelled Job, 1 Unrelated FailureAs of commit c430ffd with merge base 5e21c13 ( NEW FAILURE - The following job has failed:
CANCELLED JOB - The following job was cancelled. Please retry:
BROKEN TRUNK - The following job failed but were 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. |
|
@rascani or @digantdesai do any of you have access to that job logs?! |
Sebastian-Larsson
left a comment
There was a problem hiding this comment.
Didn't see the failing job. Let's get that sorted first.
|
Hmm I think that job name is auto generated from the branch name of this PR. I see the same job failing on other PR but renamed to something kind of matching the PR title. |
|
I think the test might be related to NXP backend as the URL link is: https://bamboo3.sw.nxp.com/browse/MLTECE-GHPOC17-2 |
Sebastian-Larsson
left a comment
There was a problem hiding this comment.
unrelated failures, including unknown NXP failure which we cannot see the log of. Consider that failure is on all PR's it's likely unrelated.
Fix CI failure for VGF supported ops - the reason is that PR that adds a new op was submitted seconds before mine. #23377 The failing PR is #23407 cc @digantdesai @freddan80 @per @zingo @oscarandersson8218 @mansnils @Sebastian-Larsson @robell @rascani Signed-off-by: Elena Zhelezina <elena.zhelezina@arm.com>
Delegate static rank-2 FP16/FP32 TopK with constant K=1..4 along the last dimension. Support largest=True and sorted=True, using FP for K=1 and FP+INT for K>1.
Previously, TopK remained outside the Arm delegate. Existing range analysis narrowed only index paths proven safe for int32:
Route all supported TopK index paths through int32 before partitioning. Restore int64 separately for consumers that need it, allowing delegated gathers to coexist with portable consumers. Keep TopK, tuple extraction and narrowing together in one delegate.
Decompose TopK into repeated ARGMAX with cumulative masking, unrolling K selections at compile time. Mask each selected position with -inf while preserving previous masks. Gather values from the original scores using the selected indices.
After decomposition (simplified):
Require finite scores without checking finiteness at runtime. Select equal scores in increasing index order, which may differ from PyTorch's tie ordering.
Change-Id: I2b9a0b9e4e6a11d13a84c473161b0ad330e20b14
cc @digantdesai @freddan80 @per @zingo @oscarandersson8218 @mansnils @Sebastian-Larsson @robell @rascani