Skip to content

Arm backend: Add static floating-point TopK lowering - #23377

Merged
YufengShi-dudu merged 3 commits into
pytorch:mainfrom
YufengShi-dudu:support-static-topk-lowering
Oct 5, 2026
Merged

YufengShi-dudu merged 3 commits into
pytorch:mainfrom
YufengShi-dudu:support-static-topk-lowering

Conversation

@YufengShi-dudu

@YufengShi-dudu YufengShi-dudu commented Oct 2, 2026 •

Copy link
Copy Markdown
Collaborator

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

cc @digantdesai @freddan80 @per @zingo @oscarandersson8218 @mansnils @Sebastian-Larsson @robell @rascani

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>
@YufengShi-dudu YufengShi-dudu added partner: arm For backend delegation, kernels, demo, etc. from the 3rd-party partner, Arm ciflow/trunk release notes: arm Changes to the ARM backend delegate module: arm Issues related to arm backend labels Oct 2, 2026
@pytorch-bot

pytorch-bot Bot commented Oct 2, 2026 •

Copy link
Copy Markdown

🔗 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 Failure

As of commit c430ffd with merge base 5e21c13 (image):

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.

@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 Oct 2, 2026

@zingo zingo left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I see buck2 files is updated, Good!
There is a topk related fail MLTEC_Elasic that need to be checked before merging. Seems NXP base, I don't seem to have access so Icant see the logs but maybe NPX backend share some of our backend code?

@zingo

zingo commented Oct 3, 2026 •

Copy link
Copy Markdown
Collaborator

@rascani or @digantdesai do any of you have access to that job logs?!

@Sebastian-Larsson Sebastian-Larsson left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Didn't see the failing job. Let's get that sorted first.

@zingo

zingo commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

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.
But still no idea what the test does or if it important.

@YufengShi-dudu

Copy link
Copy Markdown
Collaborator Author

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 Sebastian-Larsson left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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.

@YufengShi-dudu
YufengShi-dudu merged commit 1b9ba6d into pytorch:main Oct 5, 2026
431 of 435 checks passed
wwwind added a commit that referenced this pull request Oct 5, 2026
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ciflow/trunk CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. module: arm Issues related to arm backend partner: arm For backend delegation, kernels, demo, etc. from the 3rd-party partner, Arm release notes: arm Changes to the ARM backend delegate

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants