Skip to content

AutoBatching: batch through a nested loop whose trip count the data gives - #3437

Merged
wsmoses merged 1 commit into
mainfrom
pb/batch-through-invariant-trip
Oct 8, 2026
Merged

wsmoses merged 1 commit into
mainfrom
pb/batch-through-invariant-trip

Conversation

@wsmoses

@wsmoses wsmoses commented Oct 8, 2026

Copy link
Copy Markdown
Member

ParallelWhileBatcher (#3261) runs a loop nested in an enzymexla.parallel while over the batched values — interchanging it with the parallel loop — only when its trip count is a constant. What the interchange needs is that the count be the same for every iteration of the parallel loop, and analyzeWhile already refuses a condition reading anything that varies with it; so the gate is now a constant start and step, with any invariant limit.

The case: a CSR transpose padded to its longest row (the next affine-cfg PR), for i: for k < M: j = off[i] + k; if j < off[i+1]: acc += x[idx[j]], where M = max_i (off[i+1] - off[i]) is a reduce over the data. With this the row loop batches into while k < M over gathers of every row at once, where before it stayed a host-driven loop over the rows (MFEM's ElementRestriction::MultTranspose, the dominant cost of ex1's CG iteration on XLA).

Test parallel_while_invariant_trip.mlir (that kernel, full-line golden); checked against NumPy on random CSR structures with empty rows (exact before and after). The batcher's other tests unchanged.

🤖 Generated with Claude Code

https://claude.ai/code/session_016zErYp7upmqr4NHfhod9UD

…ives

The parallel-while batcher ran a nested loop over the batched values only
where its trip count was a constant, though all it needs is that the
count be the same for every iteration of the parallel loop, which its
condition check already demands: a CSR transpose padded to its longest
row runs its inner loop to a max read from the offsets.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_016zErYp7upmqr4NHfhod9UD
@codecov

codecov Bot commented Oct 8, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 29.96%. Comparing base (2d3ff23) to head (2b5f29c).
⚠️ Report is 5 commits behind head on main.

Additional details and impacted files
@@            Coverage Diff             @@
##             main    #3437      +/-   ##
==========================================
+ 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.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@wsmoses

wsmoses commented Oct 8, 2026

Copy link
Copy Markdown
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:

  • star (185k dofs, 400 iterations): 0.64 s (native 0.047 s) — the morning's state ran ~70 iterations per 600 s
  • fichera (802k dofs, 270 iterations): 1.08 s (native 0.14 s)

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