Skip to content

[Metal] Faster logsumexp for short rows via a simdgroup-per-row kernel#3913

Open
qx555 wants to merge 1 commit into
ml-explore:mainfrom
qx555:logsumexp-simd-row
Open

[Metal] Faster logsumexp for short rows via a simdgroup-per-row kernel#3913
qx555 wants to merge 1 commit into
ml-explore:mainfrom
qx555:logsumexp-simd-row

Conversation

@qx555

@qx555 qx555 commented Jul 24, 2026

Copy link
Copy Markdown

One simdgroup per row (8 rows per 256-thread threadgroup) with the online max/normalizer rescaling already used by logsumexp_looped: reductions never leave the simdgroup, so the kernel needs no threadgroup memory and no barriers (the block kernel uses 5). Dispatched for axis_size <= 2048; the block kernel is unchanged and still serves 2048 < axis_size <= 4096.

On M4 Max: fp16 up to 2.8x, bf16 up to 2.6x, fp32 up to 1.5x on affected sizes; parity at the 2048 boundary and no change on block/looped paths.

One simdgroup per row (8 rows per 256-thread threadgroup) with the online
max/normalizer rescaling already used by logsumexp_looped: reductions never
leave the simdgroup, so the kernel needs no threadgroup memory and no
barriers (the block kernel uses 5). Dispatched for axis_size <= 2048; the
block kernel is unchanged and still serves 2048 < axis_size <= 4096.

On M4 Max: fp16 up to 2.8x, bf16 up to 2.6x, fp32 up to 1.5x on affected
sizes; parity at the 2048 boundary and no change on block/looped paths.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant