Repository navigation
EnzymeHLOOpt: the live box of a masked scatter through broadcasts and constant masks - #3441
Merged
Merged
Conversation
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #3441 +/- ##
==========================================
+ Coverage 29.63% 29.96% +0.33%
==========================================
Files 240 240
Lines 48496 48546 +50
==========================================
+ Hits 14371 14548 +177
+ Misses 34125 33998 -127 ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
wsmoses
force-pushed
the
pb/scatter-masked-slice-constant
branch
from
October 8, 2026 21:52
983e55e to
099edfc
Compare
… constant masks ScatterMaskedIndexSlice read the mask as compares of iotas against constants anded together; once a kernel is specialized, those compares fold to constant prefixes along each lane dimension, broadcast onto the grid, and the pattern saw none of it. The terms are now read through the broadcasts that put them on the grid, and a constant along one dimension gives its interval. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_016zErYp7upmqr4NHfhod9UD
wsmoses
force-pushed
the
pb/scatter-masked-slice-constant
branch
from
October 8, 2026 22:03
099edfc to
c4358b7
Compare
This was referenced Oct 8, 2026
Member
Author
|
MFEM GPU unit suite (74 tests) with the runtime built from Enzyme-JAX main + #3437 #3439 #3440 #3441 #3442 and Reactant.jl #3423 + #3425, over objects from main + the open affine-cfg PRs (#3436 #3438 among them): 74/74 (sweep final55, 2026-10-08 17:38). ex1 (Poisson, order 3, PA, PCG to 1e-12), timed solves in one process (JIT excluded), native CUDA for reference:
|
This was referenced Oct 8, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
ScatterMaskedIndexSlice(#3191) finds the live box of a masked store from the mask'scompare(iota, c)terms. After exec-time specialization the extent those compares test against is a constant, and the runtime's const-prop folds each compare into a constant prefix along its lane dimension (dense<[true, true, true, true, false, ...]>), broadcast onto the lane grid andand-ed — a form the pattern did not read, so the scatter stayed at the padded size.Now each term is read through the
broadcast_in_dims that put it on the grid (a dimension map from the term's dims to the mask's), and a 1-D constant term gives its interval of true values (a splat false empties the box).The case: MFEM's
PADiffusionSetupkernels padded toMAX_Q1D(32×32 lanes per element) run with Q1D=4 — 16 live lanes of 1024. In ex1's exec-time module for it the scatters now slice to20480x4x4, and (withslice_elementwise, #3440'sslice_gather, and #3439's fold of the fully dead scatters) the arithmetic feeding them shrinks with them.A dimension of extent one that a broadcast expands holds one value for the whole grid dimension (a flag the kernel was specialized on,
broadcast_in_dim(dense<[true]> : tensor<1xi1>)): it gives no interval — a 1-element constant or a 1-extent iota compare is read as holding everywhere or nowhere. A term false everywhere is carried as an explicit "nowhere" flag, since a scatter of one point has a zero-dimensional lane grid with no box to carry it in. (The first version of this PR wrote the empty box into the grid extents, so a single masked-off store — the batched form of anscf.ifon a specialized-false flag — was read as live and its mask dropped:PA Convection,PA DG DiffusionandL2 Assembly Levelsfailed in the GPU suite; found by bisecting the runtime and an oracle comparison of the test's exec modules.)Tests:
@constant_maskinscatter_masked_slice.mlir(the folded form of@box: rows below 4, columns below 3 → one 4×3 dynamic_update_slice)@expanded_term(a 1-extent flag term broadcast over the columns: the box is the rows' alone), and@point_off(one point under a false mask: nothing written).🤖 Generated with Claude Code
https://claude.ai/code/session_016zErYp7upmqr4NHfhod9UD