submission 930185
Varshith · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2570 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-930185?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:8f412a8b23f714e09f5e84f610d14c1feab0b04359305e21b81d1145506f1f82
license declaredunknown
license concludedunknown
authorsVarshith
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
v21 chol_diag -> tri_inv -> tl.dot(P, Linv^T) (64 scalar columns, then GEMM)num-warps = 1
_chol32_kernel[(data.shape[0],)](data, out, N=32, num_warps=1)shared-memory
__shared__ __align__(16) float bcs[WPB][2][N];tile-m = 64
_PANEL_BM = 64 # RE-FIT 2026-07-26: 128 was never tuned, and since exp_fast4 thisvector-width = float4
float4 reads of it, and a double-buffered bc. Identical to what took chol_diag -23.5%.Kernel source
submission.py2570 lines
r"""exp_s6 = exp_s5 (banked 578.767) + TWO KERNEL FIXES ON THE SMALL-n CASES, both named
by ncu and both A/B'd against their own control inside ONE profile run at locked clocks.
## THE HEADLINE: THE n=128 KERNEL WAS RUNNING ITS REGISTER ARRAY OUT OF LOCAL MEMORY
`chol_fused<128>` has shipped all campaign with `float r[64]` evicted wholesale to local
memory. ncu, benchmark index 2: **"93.57% of all sectors requested in L1TEX" are local**,
only 1.0 of every 32 bytes per sector used, 20.43% of local loads spilling to L2, 78
registers per thread. Adding `#pragma unroll 1` to the panel loop -- FORBIDDING the
unroll nvcc was doing on its own -- clears it completely:
no pragma 278.05 us 78 regs 93.57% of L1TEX sectors LOCAL
#pragma unroll 1 73.98 102 regs none
#pragma unroll 2 76.80 128 regs none (full unroll, and 3.8% worse)
`chol_diag_b` hit the identical signature at 92.57% and cost 78.94 us against 17.38 until
it was fixed. **This one survived because `kernel_attrs()` -- the `stack > 0` gate
00_PLAN calls a hard rule -- listed four mid-driver kernels and none of the three small-n
ones.** It lists all eight now, so the gate finally covers the kernels it exists for.
## THE SECOND FIX: n=32's SHUFFLES
ncu on `chol_reg<32,8>`: L1/TEX throughput 87.79%, 10.9 cycles per warp stalled on a MIO
short scoreboard ("typically memory operations to shared memory") -- **in a kernel that
declares no shared memory at all**, so it can only be the 528 warp shuffles. `chol_reg_sm`
swaps them for `chol_rows`'s shared broadcast column, warp-scoped (`__syncwarp`, not
`__syncthreads`), so one LDS.128 feeds four FMAs where four SHFLs did:
chol_reg<32,8> 38.24 us 42 regs CONTROL (repeated in run: 38.88, +/-1.7%)
chol_reg_sm<32,8> 32.80 107
chol_reg_sm<32,4> 31.68 79 **-17.8%, shipped**
## WHAT IS NOT IN THIS FILE
n=64. `chol_rows<64>` is register-limited to 18.75% occupancy (150 regs -> 6 CTAs/SM ->
888 CTAs per wave against a 1024 grid, ncu: "1 full wave and a partial wave of 136 thread
blocks", Est. Speedup 50%). Four `__launch_bounds__` variants all beat it (43.65-44.54
against 52.16) and **all four report local memory**, because `__launch_bounds__(64, 1)`
is NOT the same as `__launch_bounds__(64)` -- it let ptxas take 255 registers and spill,
so that run's control was not the shipped kernel. The family needs a clean re-run.
CONTROLS: every case except n=32 and n=128. Anything else that moves is drift.
--- exp_fast5 header ------------------------------------------------------------------
exp_fast5 = exp_fast4 (banked 729.89) + the same proven transforms on the two small-n
kernels nobody has touched all campaign: chol_reg<32> (n32 case) and chol_rows<64> (n64).
rsqrt.approx instead of IEEE sqrt+divide on the per-column chain; no `t >= j` predicate
(the store masks the upper triangle); and for chol_rows the pre-scaled broadcast column,
float4 reads of it, and a double-buffered bc. Identical to what took chol_diag -23.5%.
Arithmetic changes only via rsqrt (2 ulp vs ~1, against a 2.38e-6*n recon budget).
CONTROLS: every other case. Only n32 and n64 can move; anything else is drift.
--- exp_fast2 header ----------------------------------------------------------------
exp_fast2 = exp_fast (banked 788.796) with every hot load vectorised. ONE lever, and
it is the one Nsight Compute named -- not a guess.
## WHAT THE PROFILER SAID (profile-brev, benchmark index 6 = n1024b4, exp_fast kernels)
kernel grid dur cyc/instr warps/sched ncu verdict
chol_diag (4,1)x64 22.7us 5.39 0.99 81.6% no eligible warp
trsm_panel (4,4)x256 47.2us 13.17 1.85 86.0% no eligible warp
Neither kernel spills (Local Memory Spilling Requests = 0 on both). Both are stalled on
memory latency with about one warp per scheduler, which means there is nothing resident to
hide a load behind and every avoidable memory op is paid at full latency:
- **trsm_panel, Est. Speedup 36.31%**: "each warp spends 4.8 cycles being stalled waiting
for a scoreboard dependency on a MIO (memory input/output) operation ... The primary
reason for a high number of stalls due to short scoreboards is typically memory
operations to shared memory." `Ls[j][m]` was one scalar 4-byte shared load per FMA.
- **chol_diag, Est. Speedup 19.19%**: "uncoalesced global accesses resulting in a total of
28672 excessive sectors (88% of the total 32768)", and only 4.0 of every 32 bytes per
sector used. The global row load and store were scalar -- `chol_rows` has used float4
for this since v13, `chol_diag` never did.
## WHAT THIS FILE CHANGES
Nothing structural, no routing change, no op-count change, no new kernel:
chol_diag global row load and store 64 scalar -> 16 float4
chol_diag broadcast column cb[j] 63 LDS.32 -> ~15 LDS.128 (+ a <=3 head)
trsm_panel Ls[j][m] 1 LDS.32 per FMA -> 1 LDS.128 per 4
tri_inv Ls[t][k] same
tri_inv Ls[N][N+1] -> Ls[N][N] the +1 padding blocked 16-byte alignment and
bought nothing; see the comment there
ARITHMETIC IS UNCHANGED. Every value, every operation, and every summation order is
identical to exp_fast -- only the instruction that fetches the operand differs. So a
numerics regression here would mean a real bug, not a precision trade.
ALIGNMENT, which is the one way this can be wrong: `bc` and `Ls` are declared
`__align__(16)`; `Ls`'s row stride is N=64 floats = 256 bytes; the vector loops start at
multiples of 4; and `chol_diag`'s global base is `off = c0*(ld+1)` with c0 always a
multiple of 64, so `off % 4 == 0`. N=32 satisfies all of the same.
ATTRIBUTION: n256 / n512b16 / n1024b4 / n2048b8 see chol_diag + trsm_panel. n512b640 and
n1024b60 see chol_diag + tri_inv. The giants see tri_inv only. n32 / n64 / n128 / n4096 /
n2048b2 are untouched CONTROLS -- if they move by more than ~2%, that is drift, and
exp_fast's run put drift at exactly +2%.
**READ KATTR FIRST.** exp_fast measured tri_inv 80/0, trsm_panel 83/0, chol_diag 129/0.
float4 needs the four components live at once; if any stack goes above 0 the run is VOID.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\exp_fast2.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode benchmark --no-tui -o cholesky\results\fast2_bench.json
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\fast2_lb.json
--- exp_fast header below ------------------------------------------------------------
exp_fast = v23 with the three scalar mid kernels un-stalled. ONE lever: dependency
chains and dead instructions inside `chol_diag`, `trsm_panel` and `tri_inv`. No routing
change, no new kernel, no op-count change.
## WHAT probe_low MEASURED (low_bench, 2026-07-25) -- THE DOCS HAD THE MIDS WRONG
Giants were the untouched control and came back 4.00 / 10.4 / 32.6 ms against 3.99 / 10.4
/ 32.4, so the run drifted <1% and every number below is real. Per-launch, after
subtracting each slot's own v23 baseline:
phase (4,1024) (16,512) (8,2048)
chol_diag<64> 18.7 18.4 19.2 FLAT in n and in batch
trsm_panel<64> 23.7 22.1 -- FLAT, and the BIGGEST
baddbmm_ (syrk) 7.4 8.6 --
full _blocked_reg 920 400 -- vs the real cases 916 / 385
The phases sum to the case, so there is nothing hidden:
(4,1024) diag 299 + trsm 355 + syrk 111 + tril ~8 = 773 vs 920 -> 147 dispatch
(16,512) diag 147 + trsm 155 + syrk 60 + tril ~4 = 366 vs 400 -> 34 dispatch
**THE "0.65 us PER COLUMN" RULE IS DEAD AND SO IS "THE DIAGONAL IS THE WALL".** 01_STATE
attributed 336 of n512b16's 385 and 672 of n1024b4's 916 to `chol_diag`. It is 147 and
299. The real ranking at low batch is trsm 39% > diag 33% > dispatch 16% > syrk 12%, and
`trsm_panel` -- never once suspected -- is the largest phase on the board.
The one constant that does hold: **a panel step costs ~60 us at low batch, independent of
n and of batch**, which is why every mid case is (n/64) x ~60.
## THE SINGLE-LAUNCH FUSED DRIVER IS DEAD TOO, AND THE SAME RUN KILLED IT
`fused_chol128` and `fused_chol256` at BATCH 4 are 4 CTAs, so the slot reports their
per-CTA critical path directly:
fused_chol128 @ b4 69.9 us (the real n=128 case, at batch 256, is 73.3)
fused_chol256 @ b4 194 us (the blocked driver does n=256 at batch 64 in 177)
n=128 costs the same at batch 4 as at batch 256: it is 100% latency, 0% throughput. And
4 panels cost 2.8x what 2 panels cost, because the update loop is scalar and grows as
c0 per panel. Extrapolating the fit to n=512 gives ~780 us per CTA against the blocked
driver's 400, and n=1024 gives ~3 ms. A one-CTA-per-matrix fused kernel cannot be the
answer above n=128 unless its trailing update moves to tensor cores.
(`chol_fused<256>` DID run, so the 198 KB opt-in is not refused -- see the SMEM line.)
## WHAT THIS FILE CHANGES
Nothing structural. Three kernels, two mechanical transformations, each justified by the
measurement above and each attributable to a different set of cases:
chol_diag<64> 18.7 us rsqrt.approx instead of an IEEE sqrt + IEEE divide on the
critical path; the broadcast column pre-scaled in shared
instead of 63 redundant FMULs per column; no `t >= j`
predicate (the store masks it). 4 instructions per j -> 2.
trsm_panel<64> 23.7 us the forward substitution was ONE 2016-long dependency
chain at ~20 cycles per FMA. Four accumulators.
tri_inv<64> -- same chain, same split; plus N reciprocals computed in
parallel into shared instead of N unrolled in series.
CONTROLS, deliberately untouched: n32 (`chol_reg`), n64 (`chol_rows`), n128
(`chol_fused`), n4096 and n2048b2 (cuSOLVER). If those move, it is drift.
ATTRIBUTION: n256 / n512b16 / n1024b4 / n2048b8 see chol_diag + trsm_panel. n512b640 and
n1024b60 see chol_diag + tri_inv. The three giants see tri_inv ONLY. Three disjoint
groups, one table.
NUMERICS. Only two of the changes touch arithmetic. `rsqrt.approx.f32` is 2 ulp against
sqrt+div's ~1, i.e. ~2e-7 relative on the diagonal against a recon budget of 2.38e-6 * n
(1.2e-3 at n=512). The four-way sums change the summation ORDER, which is a pairwise-style
split and no worse than left-to-right. Benchmark mode rechecks every gate on every timed
iteration, so both are validated on exactly the data they run on.
PREDICTED, if the stalls are what they look like: diag ~10, trsm ~10, panel step 60 -> 37.
n256 177 -> ~115, n512b16 385 -> ~250, n1024b4 920 -> ~600, n2048b8 1858 -> ~1200,
n512b640 and n1024b60 a few percent, giants ~1%. Geomean 865 -> **~750**.
**READ KATTR FIRST. stack > 0 on any of the three means the run is VOID.** The four-way
split adds three live floats and `bc` doubles to 512 bytes of shared; trsm_panel had 93
registers and tri_inv 90, so there is room, but that is a prediction and KATTR is a
measurement.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\exp_fast.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode benchmark --no-tui -o cholesky\results\fast_bench.json
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\fast_lb.json
--- v23 header below ---------------------------------------------------------------
submission_v23 = v22 (banked 880.313) + n2048b2 routed to cuSOLVER one at a time.
probe_giant priced batch-1 `potrf(2048)` at **649 us** directly (16 reps = 10.38 ms).
That number decides a routing case the campaign had never re-checked:
case blocked driver k x potrf(n) route
n512b16 385 16 x 167 = 2672 driver, by 7x
n1024b4 915 4 x 334 = 1336 driver
n2048b8 1859 8 x 649 = 5192 driver, by 2.8x
n2048b2 1744 2 x 649 = 1298 **LOOP, by 1.3x**
The blocked driver runs ~95 mostly-sequential ops to save flops that are not the
constraint at batch 2. n=4096 has shipped `_loop_single` for this reason since v2; n=2048
crosses over at batch 2 and nobody had priced it. Predicted 1744 -> ~1330, geomean
880.313 -> **~864**.
## THE n32768 DECOMPOSITION (probe_giant, giantphase_bench) -- PHASES SUM TO 0.2%
phase cost share efficiency
16 x potrf(2048) 10.38 ms 32% 649 us each -- the 0.33 us/column floor
trailing GEMMs 8.79 ms 27% 1329 TF/s = 84% OF bf16 PEAK
panel applies 6.01 ms 19% 343 TF/s = 22% of bf16 peak
15 x _tri_inv_blocked 2.97 ms 9% 29 TF/s = 3% of tf32 peak
data.tril() 2.88 ms 9% 2.98 TB/s = 37% of HBM
SUM 31.03 ms full driver measured 31.10, case 32.40
**THE SKINNY-GEMM HYPOTHESIS IS DEAD.** I predicted the trailing update was running near
570 TF/s because of its shape and that it held ~10 ms of slack. It is at 84% of peak and
holds ~1.4 ms. Do not spend anything there.
WHAT IS ACTUALLY LEFT AT n=32768, in order:
1. panel applies, 6.01 ms at 22% of peak. The GEMM is 2.06e12 flops (0.9 ms at peak)
plus ~10 GB of fp32<->bf16 round trip (1.3 ms). ~3.5 ms unexplained inside a phase
that is only 15 launches -- the `.to(bfloat16)` gather off a strided view is the
first suspect and needs its own split before anything is built.
2. 15 x _tri_inv_blocked, 2.97 ms at 3% of peak. Launch-bound: 15 calls x ~20 ops, and
the first three merge levels are 64/128/256-wide GEMMs that cannot fill the GPU.
A coarser base would cut levels; the 64-float ceiling blocks a wider `tri_inv`.
3. data.tril(), 2.88 ms at 37% of HBM on a pure copy-and-mask.
4. The diagonal is a FLOOR. 649 us x 16, and no custom factorization this campaign has
ever beaten cuSOLVER on a single large matrix.
Best case n32768 goes 32.4 -> ~23, which is ~0.3 ln-units. The giants are no longer where
the leverage is; the mids are 14-174x off their compute floors against the giants' 1.2-3x.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\submission_v23.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\v23_lb.json
--- v22 header below ---------------------------------------------------------------
submission_v22 = exp_ginv (banked 883.096) minus the giants' closing `tril_()`.
## THE PANEL SWEEP LOST, AND IT EXPLAINS ITSELF EXACTLY (gsweep_bench, 2026-07-25)
case p2048 (exp_ginv) p4096 (exp_gsweep)
n16384 10.5 ms 11.4 ms +8.6%
n32768 33.2 ms 35.5 ms +6.9%
**`potrf` IS NOT 0.33 us/COLUMN AT p=4096.** The price table's own numbers say so and
01_STATE called the law "EXACTLY LINEAR" anyway:
p 512 1024 2048 4096
potrf(p) 167 334 663 1541
us/column 0.326 0.326 0.324 0.376 <- +16% at 4096
That single fact predicts the whole result. n16384: 4 x 1541 - 8 x 663 = +0.86 ms
against +0.9 measured. n32768: 8 x 1541 - 16 x 663 = +1.72, plus the blocked inverse's
p^3 growth (7 x 2*4096^3/3 against 15 x 2*2048^3/3 = +2.4e11 tf32 flops, ~+0.5 ms) =
+2.2 against +2.3 measured.
**SO OP COUNT IS NOT WHAT COSTS THE GIANTS.** Halving the panels saved nothing
measurable; the model closes to within 0.6 ms on both cases using only real GPU work.
That also retires my explanation for exp_ginv's 8% prediction miss -- it was not launch
overhead. p=2048 stays everywhere.
## WHAT THIS FILE CHANGES
One line. `data.clone()` -> `data.tril()` and the closing `tril_()` goes away. Safe for
the same reason it was safe in the mids: left-looking never writes the strict upper block
triangle, and every diagonal block is overwritten by `ljj`, which `cholesky_ex` returns
already zeroed above the diagonal. At n=32768 that removes a 4.3 GB read + 4.3 GB write,
~1.1 ms at 8 TB/s. Predicted n32768 33.2 -> ~32.1, n16384 -> ~10.2, n8192 -> ~3.96,
geomean 883 -> ~878. Small, free, and it was the mids' single biggest phase win in v19.
## WHERE n32768's 33.2 ms ACTUALLY GOES -- AND WHY THE NEXT STEP IS A PROBE
diagonal 16 x potrf(2048) 10.6 ms floor, invariant to p
inverse 15 x _tri_inv_blocked ~1.3 ms was 13.5 before exp_ginv
trailing n^3/3 = 1.17e13 bf16 flops 7.4 ms AT PEAK bf16 (1578 TF/s)
accounted 19.3 ms
MEASURED 33.2 ms -> ~14 ms unexplained
Best hypothesis: the trailing GEMM is SKINNY, not square. At panel j its shape is
[32768-j, j] x [j, 2048], and probe_graph measured skinny shapes at 570 TF/s where
square 8192^2 hits 1578. At ~700 TF/s the trailing update alone is ~17 ms and the case
closes. Second candidate is the bf16 <-> fp32 round trip in the update, ~11 GB of extra
traffic. **Decompose before building** -- that method has been right every time this
campaign and the instruction/flop models have not.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\submission_v22.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\v22_lb.json
--- exp_ginv header below ----------------------------------------------------------
exp_ginv = v21 + the GIANTS' TRIANGULAR INVERSE on tensor cores. Lever 3.
v21's giants still build each panel's inverse with `solve_triangular`, priced by
probe_prim2 at 899 us at p=2048 and getting NO tensor cores (3.2 TF/s measured). That is
the single largest addressable block left:
case total panels inverses solve_triangular share
n8192 5.91 ms 4 3 2.7 ms 46%
n16384 14.9 ms 8 7 6.3 ms 42%
n32768 42.6 ms 16 15 13.5 ms 32%
`_tri_inv_blocked` replaces it with the 2x2 block identity
inv([[A, 0], [C, B]]) = [[inv(A), 0], [-inv(B) C inv(A), inv(B)]]
applied level by level: ONE `tri_inv<64>` launch inverts all 32 diagonal blocks, then
five merges double the block size with two batched matmuls each. 2 p^3/3 flops = 5.7
GFLOP at p=2048, roughly 15 us of tf32 against 899.
THE BASE CASE HAS TO BE OUR OWN KERNEL. Batched `solve_triangular` is WORSE than the
single call, not better -- probe_graph measured 8 x 256x256 at 617 us/op -- so recursing
down to a vendor call at any width keeps exactly the thing being removed. `tri_inv<64>`
already exists and is measured clean (regs 90, stack 0); `tri_inv_stack` is the same
kernel pointed at a contiguous stack of blocks.
NO GATHERS. The four corners of every diagonal pair share one stride pattern and differ
only in their offset, so a whole merge level is two batched matmuls over strided views.
PREDICTED: inverse 899 -> ~150 us per panel including ~20 launches of dispatch, so
n8192 5.91 -> ~3.7, n16384 14.9 -> ~9.7, n32768 42.6 -> ~31.4, **geomean ~942 -> ~868**.
The risk is op count: ~20 ops per inverse against 1, i.e. ~300 added at n32768. These are
batch-1 cases with 40 ms of GPU work so dispatch should stay a small fraction, but RULE 0
says measure it rather than assume, and the benchmark table will show it directly.
VALIDATED (`validate_ginv.py`, CPU): matches `solve_triangular` to 1.8e-7 relative at
every panel width, and under TF32-TRUNCATED GEMMs the full driver's reconstruction
residual is LOWER than the solve_triangular version's (18.17% of budget vs 19.31% at
n=2048/4 panels; 35.80 vs 37.19 at n=1024/4). It adds no numerics risk.
Set-Location C:\Users\kvars\Desktop\gpumode
& ".venv\Scripts\python.exe" cholesky\solutions\validate_ginv.py
Copy-Item cholesky\solutions\exp_ginv.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode benchmark --no-tui -o cholesky\results\ginv_bench.json
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\ginv_lb.json
--- v21 header below, for the mid routes this file leaves alone --------------------
submission_v21 = BANK. exp_tc's tensor-core panel solve, routed by BATCH.
exp_tc (909445) took the panel TRSM onto the tensor cores and scored **962.564 ranked**
against v19's 979.536. It ran the new path on every tf32-safe case; the per-case table
came back MONOTONE IN BATCH and half of those cases lost:
case v19 exp_tc RANKED delta v21 route
n256 176 176 -- reg (control, never routed)
n512b16 381 441 +16% reg <- reverted
n512b640 2090 1467 -30% TC
n1024b4 903 895 -- reg (control)
n1024b60 1287 1073 -17% TC
n2048b2 1700 1852 +9% reg <- reverted
n2048b8 1877 1930 +3% reg <- reverted
geomean 979.536 962.564 -1.7%
`tri_inv` plus the Triton launch is a FIXED cost per panel; the GEMM saving scales with
the panel's rows x batch. So the new path pays exactly where there is work to amortize it
and loses where the case is latency-bound, which is the same axis that decides everything
else in this task. Taking only the wins: **predicted ~943**.
The gate is `batch >= 32`. It sits in the empty gap between the measured winners (60,
640) and losers (2, 8, 16); no benchmark or test case falls between. NOTE this means no
TEST case reaches the tensor-core path -- every test at n=512/1024/2048 is batch <= 4 --
so its numerics are validated by benchmark mode's per-iteration recheck on exactly the
two cases that run it, plus `validate_tc.py` on CPU. Same standing as v19's tf32 gate.
## THE MECHANISM, FOR THE NEXT BUILD
A triangular solve cannot use `tl.dot`; a multiply by the inverse can. Per panel:
v19 chol_diag -> trsm_panel (scalar, 5.1 TF/s)
v21 chol_diag -> tri_inv -> tl.dot(P, Linv^T) (64 scalar columns, then GEMM)
`trsm_panel` was 915 us of n512b640's 2620 at 5.1 TF/s against tf32's 1.1 PF/s. Trading
a 448-row substitution for 64 inverse columns plus a GEMM beat the prediction: ~1660
predicted, 1467 measured. Op count went 3 per panel to 4 and the win survived it, so
RULE 0's tax is real but an order smaller than a phase moving to tensor cores.
KATTR from 909445, all clean -- `float m[64]` promoted, which is what the exp_fpanel
campaign failed four times to achieve:
tri_inv<64> regs=90 stack=0 smem=16640 maxtpb=64
trsm_panel<64> regs=93 stack=0 smem=16640 maxtpb=256
chol_diag<64> regs=158 stack=0 smem=256 maxtpb=64
## WHY tri_inv IS SHAPED THE WAY IT IS
Thread j owns COLUMN j of the inverse. The column recurrence reads only L and that same
column's earlier entries, so there are no barriers and no cross-thread broadcast; a
row-per-thread inverse would be 64 sequential steps with one active thread each. Every
index into `float m[64]` is an unrolled loop variable and the array is live in ONE
straight-line region -- both are direct consequences of the exp_fpanel failure, where a
runtime index and a two-phase live range spilled `r[64]` to local memory twice and cost
four submissions. `trsm_panel<64>` (93 regs, stack 0) is the shape being copied.
**READ KATTR FIRST. `tri_inv<64>` with stack > 0 means the run is void.**
## GATING
v9 measured tf32 `tl.dot` passing the recon gate at n=512 and FAILING at n=256, so this
path is gated to exactly the cases v19 already runs tf32 on: n in {512, 2048}, and
n=1024 only at batch >= 32. n=256 and n1024b4 stay on `_blocked_reg` untouched and are
the control row -- if they move, something other than the panel solve changed.
## VALIDATION
`validate_tc.py`, CPU, fp32: the inverse-based solve against reference.py's gates on
every task.yml distribution reaching the routed sizes. Multiplying by an explicit inverse
is less stable than substitution, so this is a numerics change, not just a speed change,
and it is checked as one. tf32 rounding is not modelled on CPU -- benchmark mode rechecks
the gates on every timed iteration, which is what exercises it.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\submission_v21.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\v21_lb.json
--- v19 header below, for the routes this file leaves alone -------------------------
submission_v19 = BANK. Measured geomean 981.4 (banked v18 = 1057.730, v16 = 1174.577).
midll_bench 2026-07-25. ONE change vs v18: the mid driver (`_blocked_reg`, n in
{256,512,1024,2048}) is LEFT-LOOKING and starts from `data.tril()`.
case v18 v19 why
n512b640 2620 2090 -20%
n1024b60 1868 1287 -31%
n2048b8 2490 1877 -25%
n512b16 391 381
n2048b2 1749 1700
n1024b4 872 903 within its observed 871-903 spread, not a regression
n256 191 191 latency-bound, left-looking cannot help it
everything else unchanged within drift
BOTH CHANGES WERE FREE -- no extra ops, no new kernels:
1. The right-looking `w[:, ce:, ce:].baddbmm_(p, p.T)` computed the FULL SQUARE trailing
block, 3.67e7 MACs at n=512 against left-looking's 2.24e7, at the SAME op count (one
GEMM per panel either way). This was the identical defect already fixed in the giants
and still shipping in the mids.
2. clone + tril_ was 23% of n512b640 for zero arithmetic. Left-looking writes only the
lower block triangle, so `data.tril()` fuses the copy and the zeroing and the final
tril_ disappears.
IT WAS BUILT FROM A MEASUREMENT, NOT AN ESTIMATE. probe_mid2 decomposed n512b640 and
n1024b60 into four phases that SUM TO THE TOTAL (2660/2620 and 1845/1868), proving there
is no dispatch gap at high batch -- the "op count is the bottleneck" reading came from
n1024b4, which is low-batch and latency-bound and does not generalise. Predicted ~985
from those phases, measured 981.4.
VALIDATED ON CPU, 12/12, against reference.py's gates on every task.yml test reaching
n in {512,1024,2048} at both panel widths, with the shipped tf32 gate. lowrank n=1024 --
the distribution that killed v6t/v6a -- sits at 0.01% of budget.
Set-Location C:\Users\kvars\Desktop\gpumode
Copy-Item cholesky\solutions\submission_v19.py cholesky\submission.py
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode test --no-tui -o cholesky\results\v19_test.json
& "bin\popcorn-cli.exe" submit cholesky\submission.py --gpu B200 --leaderboard cholesky --mode leaderboard --no-tui -o cholesky\results\v19_lb.json
"""
import os
import subprocess
import sys
import tempfile
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0a")
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
# ---------------------------------------------------------------- n=32 register kernel
_REG_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#define WARP_ALL 0xffffffffu
// One warp per matrix, thread `lane` owns row `lane` entirely in registers. Right-looking:
// at step k the owner of row j broadcasts L[j][k] and every thread i >= j does one FMA into
// its own row. No shared memory, no __syncthreads. Abdelfattah et al., JOCS 26 (2018) 226.
template <int N, int WPB>
__global__ __launch_bounds__(32 * WPB)
void chol_reg(const float* __restrict__ src, float* __restrict__ dst, int batch) {
const int lane = threadIdx.x & 31;
const int mat = blockIdx.x * WPB + (int)(threadIdx.x >> 5);
if (mat >= batch) return;
const float* __restrict__ A = src + (size_t)mat * N * N;
float* __restrict__ L = dst + (size_t)mat * N * N;
float r[N];
const float4* row4 = reinterpret_cast<const float4*>(A + (size_t)lane * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
const float4 v = row4[c];
r[4 * c + 0] = v.x;
r[4 * c + 1] = v.y;
r[4 * c + 2] = v.z;
r[4 * c + 3] = v.w;
}
// Same two transforms that took chol_diag 22.85 -> 17.47 us: `rsqrt.approx.f32` in
// place of an IEEE sqrt + IEEE divide (both are multi-instruction sequences sitting on
// the per-column dependency chain), and no `lane >= j` predicate -- the store below
// masks the upper triangle, and an above-diagonal r[j] is only ever written by its own
// update, so the discarded work is harmless. The shuffle already carries the SCALED
// L[j][k], so unlike chol_diag there is no broadcast column to pre-scale.
#pragma unroll
for (int k = 0; k < N; ++k) {
const float akk = __shfl_sync(WARP_ALL, r[k], k);
const float a = fmaxf(akk, 1.17549435e-38f);
float rd;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(rd) : "f"(a));
if (lane == k) r[k] = a * rd;
else if (lane > k) r[k] *= rd;
#pragma unroll
for (int j = k + 1; j < N; ++j) {
const float ljk = __shfl_sync(WARP_ALL, r[k], j);
r[j] -= r[k] * ljk;
}
}
float4* out4 = reinterpret_cast<float4*>(L + (size_t)lane * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
float4 v;
v.x = (lane >= 4 * c + 0) ? r[4 * c + 0] : 0.0f;
v.y = (lane >= 4 * c + 1) ? r[4 * c + 1] : 0.0f;
v.z = (lane >= 4 * c + 2) ? r[4 * c + 2] : 0.0f;
v.w = (lane >= 4 * c + 3) ? r[4 * c + 3] : 0.0f;
out4[c] = v;
}
}
// SHIPPED n=32 KERNEL SINCE 2026-07-27. `chol_reg`'s N^2/2 warp shuffles replaced by
// `chol_rows`'s shared broadcast column, warp-scoped so the barrier is `__syncwarp` and
// not `__syncthreads`. One LDS.128 feeds FOUR FMAs where four SHFLs fed four FMAs, so
// 528 MIO instructions per thread become ~190.
//
// MEASURED ON B200, ncu, all four kernels in ONE profile run at a locked 1.14 GHz:
//
// chol_reg<32,8> 38.24 us 42 regs CONTROL (repeat: 38.88, so +/-1.7% in run)
// chol_reg_sm<32,8> 32.80 107
// chol_reg_sm<32,4> 31.68 79 **-17.8%, shipped**
//
// ncu named the defect on the shipped kernel and it could only have been the shuffles:
// L1/TEX throughput 87.79% and 10.9 cycles per warp stalled on a MIO short scoreboard
// ("typically memory operations to shared memory"), in a kernel that declared NO SHARED
// MEMORY AT ALL. WPB 8 -> 4 costs occupancy (38.7% -> 29.9%) and wins anyway; the grid
// doubles to 1024 CTAs, which balances better over 148 SMs.
//
// ARITHMETIC IS UNCHANGED. Same rsqrt.approx, same pre-scaled broadcast column, same
// order of operations as `chol_rows<64>`, which has shipped since v13 -- only the
// instruction that moves L[j][k] between lanes differs.
// **THE SAME WAVE FIX THAT PAID 14.3% AT n=64 (921879).** The mechanism there was not
// occupancy in the abstract -- it was that the GRID SLIGHTLY EXCEEDED WHAT OCCUPANCY
// COULD HOLD, so the kernel ran 1.15 waves and the second wave was 15% of the matrices
// running on an almost empty GPU:
//
// kernel grid CTAs/SM concurrent waves
// chol_rows<64,1> 1024 6 888 1.15 -> 8/SM = 1184, ONE wave
// chol_reg_sm<32,4> 1024 6 888 1.15 <- identical, this kernel
// chol_fused<128,1> 256 5 740 0.35 grid-bound, EXCLUDED
//
// 79 regs x 128 threads = 10112 per CTA, and 65536/10112 = 6. minBlocks 8 caps the
// allocation at 64 registers, which puts 1184 CTAs in flight against a 1024 grid and
// collapses the second wave. n=128 is deliberately untouched: at 0.35 waves it has spare
// capacity already and capping registers there can only cost.
//
// `chol_rows` came back regs=128 **stack=80** and still won 14.3%, so a spill here is not
// disqualifying on its own -- read KATTR, then judge on the per-case table.
template <int N, int WPB>
__global__ __launch_bounds__(32 * WPB, 8)
void chol_reg_sm(const float* __restrict__ src, float* __restrict__ dst, int batch) {
const int lane = threadIdx.x & 31;
const int w = (int)(threadIdx.x >> 5);
const int mat = blockIdx.x * WPB + w;
// `mat` is warp-uniform, so a warp either runs to the end or returns whole and the
// `__syncwarp()`s below never see a partially exited warp.
if (mat >= batch) return;
__shared__ __align__(16) float bcs[WPB][2][N];
const float* __restrict__ A = src + (size_t)mat * N * N;
float* __restrict__ L = dst + (size_t)mat * N * N;
float r[N];
const float4* row4 = reinterpret_cast<const float4*>(A + (size_t)lane * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
const float4 v = row4[c];
r[4 * c + 0] = v.x;
r[4 * c + 1] = v.y;
r[4 * c + 2] = v.z;
r[4 * c + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < N; ++k) {
float* cb = bcs[w][k & 1]; // k is unrolled -> constant offset
if (lane >= k) cb[lane] = r[k];
__syncwarp();
const float a = fmaxf(cb[k], 1.17549435e-38f);
float rd;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(rd) : "f"(a));
if (lane > k) {
const float v = r[k] * rd; // L[lane][k], published for every other lane
cb[lane] = v;
r[k] = v;
} else if (lane == k) {
r[k] = a * rd;
}
__syncwarp();
#pragma unroll
for (int j = k + 1; j < ((k + 4) & ~3) && j < N; ++j) r[j] -= r[k] * cb[j];
#pragma unroll
for (int j = ((k + 4) & ~3); j + 3 < N; j += 4) {
const float4 c4 = *reinterpret_cast<const float4*>(&cb[j]);
r[j + 0] -= r[k] * c4.x;
r[j + 1] -= r[k] * c4.y;
r[j + 2] -= r[k] * c4.z;
r[j + 3] -= r[k] * c4.w;
}
}
float4* out4 = reinterpret_cast<float4*>(L + (size_t)lane * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
float4 v;
v.x = (lane >= 4 * c + 0) ? r[4 * c + 0] : 0.0f;
v.y = (lane >= 4 * c + 1) ? r[4 * c + 1] : 0.0f;
v.z = (lane >= 4 * c + 2) ? r[4 * c + 2] : 0.0f;
v.w = (lane >= 4 * c + 3) ? r[4 * c + 3] : 0.0f;
out4[c] = v;
}
}
// N threads per matrix, one row each in `float r[N]`. N <= 64 is MANDATORY: probe_regsize
// measured nvcc's register-array promotion limit at exactly 64 floats per thread, and one
// float over it puts the whole array in local memory (which cost the first n=64 attempt
// 218 us). Rows span 2 warps at N=64, so the broadcast column goes through shared memory.
// MPB matrices per block; dead slots in a partial block stay in lockstep for __syncthreads.
// **THE n=64 WAVE FIX, VIA `__maxnreg__` RATHER THAN A minBlocks ARGUMENT.**
//
// `chol_rows<64,1>` is 150 regs -> 65536/(64*150) = 6 CTAs/SM -> 12 warps of a possible
// 64 = **18.75% occupancy**, the only small-n kernel where occupancy and not the grid is
// binding (chol_reg_sm<32,4> reaches 37.5%; chol_fused<128,1> is grid-bound at 256 CTAs
// on 148 SMs and cannot be fixed this way).
//
// **WHY EVERY EARLIER ATTEMPT FAILED.** 01_STATE: "`__launch_bounds__(T, 1)` IS NOT
// `__launch_bounds__(T)`, AND IT COST A CONTROL" -- the round-1 control used
// minBlocksPerMultiprocessor = 1, which RELAXES the cap and let ptxas take 255 registers
// and spill. The MINB 7/8/12 variants did cap it (128/128/80 regs) but ALL FOUR carried
// local memory, so the family was recorded as unjudged and "needs a variant that reaches
// ~7 CTAs/SM without local memory".
//
// **`__maxnreg__` DOES NOT BUILD IN THIS TOOLCHAIN (921828).** It compiled locally
// nowhere -- there is no CUDA toolkit on the dev box -- and on B200 `load_inline` threw,
// `_REG` fell back to None, and every `_REG`-guarded case ran the torch path: +191% to
// +365% on all fourteen of them, with n4096b1 (the one route with no `_REG` dependency)
// flat at 1603. **The KATTR stderr was EMPTY, which is the diagnostic** -- a healthy run
// lists all eight kernels there. Check it before reading any per-case table.
//
// minBlocksPerMultiprocessor = 8 is the form this campaign has already built
// successfully, and 01_STATE records the register count it produces: 128. That targets
// 8 CTAs/SM (65536/(64*128)) = 16 warps = 25% occupancy, a third more latency hiding on
// a kernel whose whole cost is a 64-step barrier-bound recurrence. It is documented to
// carry local memory; that is judged on the board, not on the spill.
//
// **READ KATTR ON THIS RUN.** If `chol_rows<64,1>` reports stack > 0 the cap is spilling
// and the variant is judged on the BOARD anyway, not on the spill: the n=128 precedent
// (917232) had real local memory, 3.8x in ncu, and moved the scored case 41.5 -> 41.9,
// i.e. not at all, because the traffic was L1-resident. ncu duration decides nothing here.
template <int N, int MPB>
__global__ __launch_bounds__(N * MPB, 8)
void chol_rows(const float* __restrict__ src, float* __restrict__ dst, int batch) {
const int slot = (int)(threadIdx.x / N);
const int t = (int)(threadIdx.x % N);
const int mat = blockIdx.x * MPB + slot;
const bool live = (mat < batch);
__shared__ __align__(16) float bc[2][MPB][N]; // double-buffered; read as float4
const float* __restrict__ A = src + (size_t)(live ? mat : 0) * N * N;
float r[N];
const float4* row4 = reinterpret_cast<const float4*>(A + (size_t)t * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
const float4 v = row4[c];
r[4 * c + 0] = v.x;
r[4 * c + 1] = v.y;
r[4 * c + 2] = v.z;
r[4 * c + 3] = v.w;
}
// The chol_diag treatment, which measured -23.5% there: rsqrt.approx instead of an
// IEEE sqrt + divide on the critical path; the broadcast column pre-scaled ONCE in
// shared by its owning thread instead of 63 redundant FMULs in every thread; no
// `t >= j` predicate (the store masks it); float4 reads of the broadcast column; and
// `bc` double-buffered so the closing barrier disappears and the count stays at 2.
#pragma unroll
for (int k = 0; k < N; ++k) {
float* cb = bc[k & 1][slot]; // k is unrolled -> constant offset
if (t >= k) cb[t] = r[k];
__syncthreads();
const float a = fmaxf(cb[k], 1.17549435e-38f);
float rd;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(rd) : "f"(a));
if (t > k) {
const float v = r[k] * rd; // L[t][k], published for every other thread
cb[t] = v;
r[k] = v;
} else if (t == k) {
r[k] = a * rd;
}
__syncthreads();
#pragma unroll
for (int j = k + 1; j < ((k + 4) & ~3) && j < N; ++j) r[j] -= r[k] * cb[j];
#pragma unroll
for (int j = ((k + 4) & ~3); j + 3 < N; j += 4) {
const float4 c4 = *reinterpret_cast<const float4*>(&cb[j]);
r[j + 0] -= r[k] * c4.x;
r[j + 1] -= r[k] * c4.y;
r[j + 2] -= r[k] * c4.z;
r[j + 3] -= r[k] * c4.w;
}
}
if (live) {
float4* out4 = reinterpret_cast<float4*>(dst + (size_t)mat * N * N + (size_t)t * N);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
float4 v;
v.x = (t >= 4 * c + 0) ? r[4 * c + 0] : 0.0f;
v.y = (t >= 4 * c + 1) ? r[4 * c + 1] : 0.0f;
v.z = (t >= 4 * c + 2) ? r[4 * c + 2] : 0.0f;
v.w = (t >= 4 * c + 3) ? r[4 * c + 3] : 0.0f;
out4[c] = v;
}
}
}
// ---- blocked-driver kernels: same register geometry, but in place on an n x n matrix ----
// Factor the N x N diagonal block at (c0, c0) of W, in place. One CTA per matrix.
// Reads and writes only the lower triangle of the block (plus zeroing the upper).
//
// probe_low priced this at 18.7 us per launch and FLAT across (4,1024), (16,512) and
// (8,2048) -- 64 columns in ~34,000 cycles = 534 cycles per column to do 63 FMAs of real
// work. Three things sat on that critical path and none of them had to:
//
// 1. `sqrtf` then `1.0f / d` are IEEE-correctly-rounded SEQUENCES, ~25 instructions and
// ~120 cycles of dependent latency, and they sit between the barrier and every FMA of
// the column. `rsqrt.approx.f32` is ONE instruction. Its error is 2 ulp against
// sqrt+div's ~1, i.e. ~2e-7 relative on the diagonal, against a recon budget of
// 2.38e-6 * n. This is the only arithmetic change in the kernel.
// 2. the broadcast column was re-scaled by `rd` INSIDE the j loop -- 63 FMULs per column
// recomputing in every thread the same 63 values. It is now scaled once, in shared,
// by the thread that owns the entry. Bitwise identical either way: both forms are
// (unscaled L[j][k]) * rd with a single rounding.
// 3. `if (t >= j)` cost an ISETP per j. It is unnecessary: an above-diagonal r[j] is
// only ever written by its own update and is masked off at the store, and in this
// driver the block's upper triangle starts at a real value (data.tril() zeroes it at
// panel 0, baddbmm_ writes it after), so the discarded work stays finite. That leaves
// LDS + FFMA -- two instructions per j instead of four.
//
// `bc` is double-buffered so the closing barrier disappears and the count stays at 2:
// nothing can reach the write at step k+2 without passing the barriers of step k+1, which
// no thread reaches until it has finished reading at step k.
template <int N>
__global__ __launch_bounds__(N)
void chol_diag(float* __restrict__ W, long long mat_stride, long long off, int ld) {
const int t = (int)threadIdx.x;
float* __restrict__ B = W + (long long)blockIdx.x * mat_stride + off;
__shared__ __align__(16) float bc[2][N]; // align: cb[j] is read as float4 below
float r[N];
// float4, not 64 scalar loads. ncu on n1024b4: "28672 excessive sectors (88% of the
// total 32768)", only 4 of every 32 bytes used, Est. Speedup 19.19%. The row is
// strided by `ld` so a warp's 32 rows can never coalesce with each other, but each
// thread's own 64 floats are contiguous, and 16 x LDG.128 moves them in a quarter of
// the requests. Alignment holds for every c0 this driver launches: off = c0*(ld+1)
// with c0 a multiple of 64, so off % 4 == 0.
const float4* __restrict__ row4 =
reinterpret_cast<const float4*>(B + (long long)t * ld);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
const float4 v = row4[c];
r[4 * c + 0] = v.x;
r[4 * c + 1] = v.y;
r[4 * c + 2] = v.z;
r[4 * c + 3] = v.w;
}
#pragma unroll
for (int k = 0; k < N; ++k) {
float* cb = bc[k & 1]; // k is unrolled, so this folds to a constant offset
if (t >= k) cb[t] = r[k];
__syncthreads();
const float a = fmaxf(cb[k], 1.17549435e-38f);
float rd;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(rd) : "f"(a));
if (t > k) {
const float v = r[k] * rd; // L[t][k], published for every other thread
cb[t] = v;
r[k] = v;
} else if (t == k) {
r[k] = a * rd; // sqrt(a), to 2.5 ulp
}
__syncthreads();
// One LDS.128 per four FFMAs instead of one LDS.32 each. The broadcast column is
// warp-uniform and contiguous, and with ~1 warp per scheduler (ncu: 0.99) there is
// nothing to hide a per-FMA shared load behind. `k + 1` is not 4-aligned, so the
// head runs scalar up to the next multiple of 4 and the rest goes vector; every
// index is still an unrolled constant, which is what keeps `r` in registers.
#pragma unroll
for (int j = k + 1; j < ((k + 4) & ~3) && j < N; ++j) r[j] -= r[k] * cb[j];
#pragma unroll
for (int j = ((k + 4) & ~3); j + 3 < N; j += 4) {
const float4 c4 = *reinterpret_cast<const float4*>(&cb[j]);
r[j + 0] -= r[k] * c4.x;
r[j + 1] -= r[k] * c4.y;
r[j + 2] -= r[k] * c4.z;
r[j + 3] -= r[k] * c4.w;
}
}
float4* __restrict__ out4 = reinterpret_cast<float4*>(B + (long long)t * ld);
#pragma unroll
for (int c = 0; c < N / 4; ++c) {
float4 v;
v.x = (t >= 4 * c + 0) ? r[4 * c + 0] : 0.0f;
v.y = (t >= 4 * c + 1) ? r[4 * c + 1] : 0.0f;
v.z = (t >= 4 * c + 2) ? r[4 * c + 2] : 0.0f;
v.w = (t >= 4 * c + 3) ? r[4 * c + 3] : 0.0f;
out4[c] = v;
}
}
// INNER-BLOCKED diagonal factor: WB columns per barrier PAIR instead of one.
//
// **THE BARRIER IS THE CRITICAL PATH. MEASURED, NOT ASSUMED.** ncu on n2048b2 (index 8)
// against `chol_diag_s<64,4>`, the 4-threads-per-row split that exp_split1/2 tried:
//
// kernel duration cyc/instr regs dominant stall
// chol_diag<64> 17.38 us 4.00 167 (banked)
// chol_diag_s<64,4> 23.84 us 10.85 43 CTA barrier, 34.7%
//
// Four times the warps made it 37% SLOWER and nearly tripled cycles per instruction, and
// ncu named the reason outright: "each warp spends 3.8 cycles being stalled waiting for
// sibling warps at a CTA barrier ... about 34.7% of the total average of 10.9 cycles".
// Two leaderboard runs agreed (668.537 and 648.901 against the banked 623.357). So the
// per-column FMA count is NOT what this kernel costs -- **64 columns x 2 __syncthreads()
// is**, and adding parallelism buys instructions at the price of barriers.
//
// This kernel spends flops to delete barriers instead. Per outer step of WB columns:
//
// 1. the WB rows of the diagonal tile publish it to shared (WB*WB floats)
// 2. BARRIER
// 3. EVERY thread factors that tile redundantly in registers, then forward-substitutes
// its own row's WB panel entries against it, and publishes them
// 4. BARRIER
// 5. rank-WB update of the columns to the right: r[j] -= sum_k p[k] * L[j][c+k]
//
// 2 barriers per WB columns, so 32 instead of 128 at WB=4. The redundant tile factor is
// ~30 flops in all 64 threads where one thread's worth would do -- which is exactly the
// trade the low-batch cases want, because at 2 CTAs on 148 SMs FLOPS ARE FREE AND ONLY
// THE SERIAL CRITICAL PATH COSTS.
//
// ARITHMETIC IS BITWISE IDENTICAL to the per-column kernel, which is why `p[k]` is built
// by a left-to-right substitution and the rank-WB update subtracts its WB products one at
// a time instead of summing them first. Every value is produced by the same operation on
// the same operands in the same order; only the barrier schedule changed. The one
// deliberate exception is the `fmaxf` pivot guard, which lives in the tile factor, so a
// row INSIDE the tile takes its own diagonal from there (`if (t == c + k)`).
//
// `if (t >= c)` skips rows already finalized. The per-column kernel could not afford the
// test; here it is once per 4 columns, and with 64 threads it retires a whole warp for
// the second half of the block.
//
// **THE TILE FACTOR IS WRITTEN IN NAMED SCALARS, NOT `float d[4][4]`, AND THAT IS NOT
// STYLE.** The first version of this kernel used arrays for the tile, the panel and the
// reciprocals. ncu: Registers Per Thread 32, "this workload accesses local memory,
// 92.57% of all sectors requested in L1TEX", 55.6% of the stall on an L1TEX scoreboard,
// duration 78.94 us against the per-column kernel's 17.38. `r[64]` had been evicted to
// local memory wholesale. The arithmetic says why: r[64] + d[16] + p[4] + rdv[4] = 88
// register-array floats, against the 64-float promotion ceiling probe_regsize measured.
// `r` is already AT the ceiling, so every other register array in this kernel has to be
// scalars. `p0..p3` stay because four is small enough to name.
template <int N>
__global__ __launch_bounds__(N)
void chol_diag_b(float* __restrict__ W, long long mat_stride, long long off, int ld) {
const int t = (int)threadIdx.x;
float* __restrict__ B = W + (long long)blockIdx.x * mat_stride + off;
__shared__ __align__(16) float dt[16]; // the 4 x 4 diagonal tile
__shared__ __align__(16) float cbq[N * 4]; // 4 broadcast columns, row-interleaved
float r[N];
const float4* __restrict__ row4 =
reinterpret_cast<const float4*>(B + (long long)t * ld);
#pragma unroll
for (int q = 0; q < N / 4; ++q) {
const float4 v = row4[q];
r[4 * q + 0] = v.x;
r[4 * q + 1] = v.y;
r[4 * q + 2] = v.z;
r[4 * q + 3] = v.w;
}
// THE OUTER LOOP MUST UNROLL. `r[c + k]` is a register-array index, so if `c` survives
// as a runtime value the whole array goes to local memory -- the exact failure that
// cost the exp_fpanel campaign four submissions.
#pragma unroll
for (int c = 0; c < N; c += 4) {
const int lr = t - c;
if ((unsigned)lr < 4u) {
dt[lr * 4 + 0] = r[c + 0];
dt[lr * 4 + 1] = r[c + 1];
dt[lr * 4 + 2] = r[c + 2];
dt[lr * 4 + 3] = r[c + 3];
}
__syncthreads();
if (t >= c) {
// The per-column kernel's own recurrence, on 4 columns, in registers, with no
// barrier because every thread holds a private copy of the tile. Written out
// in the exact order the loop form produces so the result is BITWISE equal:
// scale column k, then rank-1 update the trailing sub-tile, then move on.
float q0, q1, q2, q3;
const float a0 = fmaxf(dt[0], 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q0) : "f"(a0));
const float l00 = a0 * q0;
const float l10 = dt[4] * q0;
const float l20 = dt[8] * q0;
const float l30 = dt[12] * q0;
const float a1 = fmaxf(dt[5] - l10 * l10, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q1) : "f"(a1));
const float l11 = a1 * q1;
const float l21 = (dt[9] - l20 * l10) * q1;
const float l31 = (dt[13] - l30 * l10) * q1;
const float e22 = dt[10] - l20 * l20;
const float e32 = dt[14] - l30 * l20;
const float e33 = dt[15] - l30 * l30;
const float a2 = fmaxf(e22 - l21 * l21, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q2) : "f"(a2));
const float l22 = a2 * q2;
const float l32 = (e32 - l31 * l21) * q2;
const float f33 = e33 - l31 * l31;
const float a3 = fmaxf(f33 - l32 * l32, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q3) : "f"(a3));
const float l33 = a3 * q3;
// This row's four panel entries, same association as the per-column kernel:
// every subtraction is a separate FMA into p[k], left to right, then one scale.
float p0 = r[c + 0], p1 = r[c + 1], p2 = r[c + 2], p3 = r[c + 3];
p0 *= q0;
p1 -= p0 * l10; p1 *= q1;
p2 -= p0 * l20; p2 -= p1 * l21; p2 *= q2;
p3 -= p0 * l30; p3 -= p1 * l31; p3 -= p2 * l32; p3 *= q3;
// A row INSIDE the tile takes its diagonal from the tile factor, which is the
// only place the fmaxf pivot guard gets applied.
if (t == c + 0) p0 = l00;
if (t == c + 1) p1 = l11;
if (t == c + 2) p2 = l22;
if (t == c + 3) p3 = l33;
r[c + 0] = p0; r[c + 1] = p1; r[c + 2] = p2; r[c + 3] = p3;
*reinterpret_cast<float4*>(&cbq[t * 4]) = make_float4(p0, p1, p2, p3);
}
__syncthreads();
if (t >= c) {
#pragma unroll
for (int j = c + 4; j < N; ++j) {
// j is warp-uniform, so this is one 16-byte broadcast per four FMAs.
const float4 v = *reinterpret_cast<const float4*>(&cbq[j * 4]);
r[j] -= r[c + 0] * v.x;
r[j] -= r[c + 1] * v.y;
r[j] -= r[c + 2] * v.z;
r[j] -= r[c + 3] * v.w;
}
}
}
float4* __restrict__ out4 = reinterpret_cast<float4*>(B + (long long)t * ld);
#pragma unroll
for (int q = 0; q < N / 4; ++q) {
float4 v;
v.x = (t >= 4 * q + 0) ? r[4 * q + 0] : 0.0f;
v.y = (t >= 4 * q + 1) ? r[4 * q + 1] : 0.0f;
v.z = (t >= 4 * q + 2) ? r[4 * q + 2] : 0.0f;
v.w = (t >= 4 * q + 3) ? r[4 * q + 3] : 0.0f;
out4[q] = v;
}
}
// COLUMN-SPLIT chol_diag_b: S threads per row, each owning N/S columns of it.
//
// **WHY THIS IS THE BIGGEST REMAINING LEVER.** `chol_diag` is 8.16 us per launch and FLAT
// in n AND in batch (916966) -- pure fixed latency -- and the driver pays it once per
// panel, so it is (n/64) x 8.16 on n256b64, n512b16, n1024b4, n2048b2 and n4096b2 at
// once. At batch 2 it runs 2 CTAs of 2 warps on 148 SMs: 15,500 cycles for ~4,000
// instructions, i.e. ~3.9 cycles per instruction with one warp per scheduler and nothing
// resident to hide a shared load behind. S=4 cuts per-thread work 4x AND takes the block
// from 2 warps to 8.
//
// **THE PRIOR SPLIT FAILED FOR A REASON THIS ONE DOES NOT SHARE.** `chol_diag_s<64,4>`
// (exp_split1/2) split the PER-COLUMN kernel: 128 barriers, and ncu measured the CTA
// barrier as 34.7% of its stall at 8 warps -- 37% slower, two leaderboard runs agreeing.
// This splits the WB=4 BLOCKED kernel, which has 32 barriers, so the same barrier tax is
// four times smaller against the same 4x work reduction.
//
// EVERY `r[]` INDEX STAYS A COMPILE-TIME CONSTANT, which is the whole difficulty. `c` is
// unrolled and NC is a template constant, so `cs = c/NC` and `lc = c - cs*NC` fold. The
// update loop runs over the LOCAL index jj and predicates on the global `j = c0s + jj`;
// writing it as `for (j = max(c+4, c0s); ...) r[j - c0s]` instead would make the index
// runtime and put the whole array in local memory -- see `chol_fused<128>`, which shipped
// that way all campaign.
//
// ARITHMETIC IS BITWISE IDENTICAL TO `chol_diag_b`. Every thread reads its own row's four
// panel entries back from `cbq` rather than from registers, and they are the same values
// the owning thread just wrote there, consumed in the same order.
template <int N, int S>
__global__ __launch_bounds__(N * S)
void chol_diag_bs(float* __restrict__ W, long long mat_stride, long long off, int ld) {
constexpr int NC = N / S; // columns per thread
const int t = (int)threadIdx.x % N; // row
const int s = (int)threadIdx.x / N; // column chunk; warp-uniform for N >= 32
const int c0s = s * NC; // first column this thread owns
float* __restrict__ B = W + (long long)blockIdx.x * mat_stride + off;
__shared__ __align__(16) float dt[16];
__shared__ __align__(16) float cbq[N * 4];
float r[NC];
const float4* __restrict__ row4 =
reinterpret_cast<const float4*>(B + (long long)t * ld + c0s);
#pragma unroll
for (int q = 0; q < NC / 4; ++q) {
const float4 v = row4[q];
r[4 * q + 0] = v.x;
r[4 * q + 1] = v.y;
r[4 * q + 2] = v.z;
r[4 * q + 3] = v.w;
}
#pragma unroll
for (int c = 0; c < N; c += 4) {
const int cs = c / NC; // c is unrolled -> compile-time
const int lc = c - cs * NC; // ditto
const int lr = t - c;
if ((unsigned)lr < 4u && s == cs) {
dt[lr * 4 + 0] = r[lc + 0];
dt[lr * 4 + 1] = r[lc + 1];
dt[lr * 4 + 2] = r[lc + 2];
dt[lr * 4 + 3] = r[lc + 3];
}
__syncthreads();
float q0, q1, q2, q3;
const float a0 = fmaxf(dt[0], 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q0) : "f"(a0));
const float l00 = a0 * q0;
const float l10 = dt[4] * q0;
const float l20 = dt[8] * q0;
const float l30 = dt[12] * q0;
const float a1 = fmaxf(dt[5] - l10 * l10, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q1) : "f"(a1));
const float l11 = a1 * q1;
const float l21 = (dt[9] - l20 * l10) * q1;
const float l31 = (dt[13] - l30 * l10) * q1;
const float e22 = dt[10] - l20 * l20;
const float e32 = dt[14] - l30 * l20;
const float e33 = dt[15] - l30 * l30;
const float a2 = fmaxf(e22 - l21 * l21, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q2) : "f"(a2));
const float l22 = a2 * q2;
const float l32 = (e32 - l31 * l21) * q2;
const float f33 = e33 - l31 * l31;
const float a3 = fmaxf(f33 - l32 * l32, 1.17549435e-38f);
asm("rsqrt.approx.f32 %0, %1;" : "=f"(q3) : "f"(a3));
const float l33 = a3 * q3;
if (t >= c && s == cs) {
float p0 = r[lc + 0], p1 = r[lc + 1], p2 = r[lc + 2], p3 = r[lc + 3];
p0 *= q0;
p1 -= p0 * l10; p1 *= q1;
p2 -= p0 * l20; p2 -= p1 * l21; p2 *= q2;
p3 -= p0 * l30; p3 -= p1 * l31; p3 -= p2 * l32; p3 *= q3;
if (t == c + 0) p0 = l00;
if (t == c + 1) p1 = l11;
if (t == c + 2) p2 = l22;
if (t == c + 3) p3 = l33;
r[lc + 0] = p0; r[lc + 1] = p1; r[lc + 2] = p2; r[lc + 3] = p3;
*reinterpret_cast<float4*>(&cbq[t * 4]) = make_float4(p0, p1, p2, p3);
}
__syncthreads();
if (t >= c) {
// This row's own four panel entries. The owning chunk wrote them above; every
// chunk of row t reads the same four values back, so the update below is the
// same arithmetic in the same order in every chunk.
const float4 pv = *reinterpret_cast<const float4*>(&cbq[t * 4]);
#pragma unroll
for (int jj = 0; jj < NC; ++jj) {
const int j = c0s + jj; // warp-uniform: cbq[j*4] broadcasts
if (j >= c + 4) {
const float4 v = *reinterpret_cast<const float4*>(&cbq[j * 4]);
r[jj] -= pv.x * v.x;
r[jj] -= pv.y * v.y;
r[jj] -= pv.z * v.z;
r[jj] -= pv.w * v.w;
}
}
}
}
float4* __restrict__ out4 =
reinterpret_cast<float4*>(B + (long long)t * ld + c0s);
#pragma unroll
for (int q = 0; q < NC / 4; ++q) {
float4 v;
v.x = (t >= c0s + 4 * q + 0) ? r[4 * q + 0] : 0.0f;
v.y = (t >= c0s + 4 * q + 1) ? r[4 * q + 1] : 0.0f;
v.z = (t >= c0s + 4 * q + 2) ? r[4 * q + 2] : 0.0f;
v.w = (t >= c0s + 4 * q + 3) ? r[4 * q + 3] : 0.0f;
out4[q] = v;
}
}
void chol_diag_split(torch::Tensor W, int64_t c0, int64_t split) {
TORCH_CHECK(W.dim() == 3 && W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"chol_diag_split: contiguous fp32 (b,n,n)");
const int n = (int)W.size(1);
const long long off = c0 * (long long)n + c0;
const int b = (int)W.size(0);
if (split == 1) {
chol_diag_bs<64, 1><<<b, 64>>>(W.data_ptr<float>(), (long long)n * n, off, n);
} else if (split == 2) {
chol_diag_bs<64, 2><<<b, 128>>>(W.data_ptr<float>(), (long long)n * n, off, n);
} else if (split == 4) {
chol_diag_bs<64, 4><<<b, 256>>>(W.data_ptr<float>(), (long long)n * n, off, n);
} else if (split == 8) {
chol_diag_bs<64, 8><<<b, 512>>>(W.data_ptr<float>(), (long long)n * n, off, n);
} else {
TORCH_CHECK(false, "chol_diag_split: split must be 1, 2, 4 or 8, got ", split);
}
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "chol_diag_split: launch failed");
}
// P <- P * L11^-T for the rows below the diagonal block: one thread per row, the row in
// registers (r[N], at the 64-float ceiling), L11 broadcast from shared. This is what
// replaces torch.linalg.solve_triangular, measured at 480 us/op -- the reason every
// earlier blocked-torch attempt lost.
template <int N>
__global__ __launch_bounds__(256)
void trsm_panel(float* __restrict__ W, long long mat_stride, long long off_l,
long long off_p, int ld, int rows) {
float* __restrict__ base = W + (long long)blockIdx.y * mat_stride;
__shared__ __align__(16) float Ls[N][N]; // align: Ls[j][m] is read as float4 below
__shared__ float dinv[N];
for (int idx = (int)threadIdx.x; idx < N * N; idx += (int)blockDim.x) {
const int i = idx / N;
Ls[i][idx - i * N] = base[off_l + (long long)i * ld + (idx - i * N)];
}
__syncthreads();
if ((int)threadIdx.x < N) dinv[threadIdx.x] = 1.0f / Ls[threadIdx.x][threadIdx.x];
__syncthreads();
const int row = (int)(blockIdx.x * blockDim.x + threadIdx.x);
if (row >= rows) return;
float* __restrict__ p = base + off_p + (long long)row * ld;
float r[N];
#pragma unroll
for (int c = 0; c < N; ++c) r[c] = p[c];
// FOUR ACCUMULATORS, NOT ONE. probe_low priced this kernel at 23.7 us per launch at
// (4,1024) and 22.1 at (16,512) -- BIGGER than chol_diag, and the largest single
// phase of every low-batch mid. The forward substitution is 2016 FFMAs but they were
// ONE dependency chain: every `s -= r[m] * Ls[j][m]` waits on the previous FFMA and
// on its own shared load, so the kernel ran at ~20 cycles per FMA against a 4-cycle
// issue latency. Splitting the sum four ways puts four loads and four FFMAs in flight
// and shortens the chain to 504. The summation ORDER changes, so this is not bitwise
// identical -- it is a pairwise-style split, which is if anything more accurate than
// strict left-to-right, and the error stays O(sqrt(j) * eps).
#pragma unroll
for (int j = 0; j < N; ++j) {
float s0 = r[j], s1 = 0.0f, s2 = 0.0f, s3 = 0.0f;
// ONE LDS.128 PER FOUR FFMAs. ncu on n1024b4 named this exactly: "each warp spends
// 4.8 cycles stalled waiting for a scoreboard dependency on a MIO operation ...
// typically memory operations to shared memory", Est. Speedup 36.31%, against
// 13.17 cycles per issued instruction and 0.16 eligible warps per scheduler. Four
// separate scalar shared loads fed the four accumulators; one vector load feeds
// all four, so the memory traffic that was one op per FMA is now one per four.
#pragma unroll
for (int m = 0; m + 3 < j; m += 4) {
const float4 lv = *reinterpret_cast<const float4*>(&Ls[j][m]);
s0 -= r[m + 0] * lv.x;
s1 -= r[m + 1] * lv.y;
s2 -= r[m + 2] * lv.z;
s3 -= r[m + 3] * lv.w;
}
#pragma unroll
for (int m = (j & ~3); m < j; ++m) s0 -= r[m] * Ls[j][m];
r[j] = ((s0 + s1) + (s2 + s3)) * dinv[j];
}
#pragma unroll
for (int c = 0; c < N; ++c) p[c] = r[c];
}
// ---- M = L11^-1, so the panel solve becomes a GEMM and can reach the tensor cores ----
//
// THREAD j OWNS COLUMN j OF M, not a row. That is the whole trick. The column recurrence
//
// M[j][j] = 1/L[j][j]
// M[t][j] = -(1/L[t][t]) * sum_{k<t} L[t][k] * M[k][j] (t > j)
//
// reads only L and THAT SAME COLUMN's earlier entries, so a thread never needs another
// thread's result: there are NO barriers after the initial shared load, and no
// cross-thread broadcast. A row-per-thread inverse would instead be 64 sequential steps
// with one active thread each. Entries above the diagonal fall out as zeros on their own
// (all M[k][j] with k < j are 0, so the sum is 0), which is why there is no mask here.
//
// EVERY INDEX INTO `m` IS AN UNROLLED LOOP VARIABLE. `m[j]` would be a runtime index and
// would send the whole array to local memory -- that is the bug that cost exp_fpanel two
// submissions. The diagonal is reached as `Ls[t][t]` in SHARED, where a dynamic index is
// free. `m` is also live in exactly ONE straight-line region, matching trsm_panel<64>
// (93 regs, stack 0) rather than the two-phase chol_panel that spilled at every size.
template <int N>
__global__ __launch_bounds__(N)
void tri_inv(const float* __restrict__ W, float* __restrict__ out,
long long mat_stride, long long off_l, int ld) {
const int j = (int)threadIdx.x;
const float* __restrict__ L = W + (long long)blockIdx.x * mat_stride + off_l;
float* __restrict__ M = out + (long long)blockIdx.x * N * N;
// `Ls[N][N+1]` until 2026-07-25. The +1 was defensive padding that bought nothing --
// the only column-strided access is the WRITE below, where consecutive threads write
// consecutive floats and so never conflict at any row stride -- and it made the row
// stride 260 bytes, which is not 16-byte aligned and blocks the float4 read. Plain
// [N][N] is also strictly better for `Ls[j][j]`: stride 65 floats gives each thread a
// distinct bank, where 66 collided 2-way.
__shared__ __align__(16) float Ls[N][N];
__shared__ float dinv[N];
#pragma unroll
for (int i = 0; i < N; ++i) Ls[i][j] = L[(long long)i * ld + j];
__syncthreads();
dinv[j] = 1.0f / Ls[j][j]; // N reciprocals in parallel, not N unrolled in series
__syncthreads();
// Same four-accumulator split as trsm_panel, for the same reason: the column
// recurrence was one dependency chain of 2016 FFMAs, each waiting on a shared load.
float m[N];
#pragma unroll
for (int t = 0; t < N; ++t) {
float s0 = 0.0f, s1 = 0.0f, s2 = 0.0f, s3 = 0.0f;
#pragma unroll
for (int k = 0; k + 3 < t; k += 4) {
const float4 lv = *reinterpret_cast<const float4*>(&Ls[t][k]);
s0 += lv.x * m[k + 0];
s1 += lv.y * m[k + 1];
s2 += lv.z * m[k + 2];
s3 += lv.w * m[k + 3];
}
#pragma unroll
for (int k = (t & ~3); k < t; ++k) s0 += Ls[t][k] * m[k];
const float s = (s0 + s1) + (s2 + s3);
m[t] = (t == j) ? dinv[t] : -s * dinv[t];
}
#pragma unroll
for (int t = 0; t < N; ++t) M[(long long)t * N + j] = m[t];
}
// ---------- SINGLE-LAUNCH fused left-looking factorization, n = 128 / 256 ----------
//
// One CTA per matrix, N threads, thread t owns row t. The matrix is walked LEFT-LOOKING
// in 64-wide column panels: panel p is first brought up to date against everything to its
// left, then factored in place with the row held in registers. There are NO launch
// bubbles at all -- the whole factorization is one kernel.
//
// WHY LEFT-LOOKING AND NOT PHASE-BY-PHASE. A right-looking version of n=128 would emit
// factor / trailing-update / factor as three straight-line bodies, ~21k unrolled
// instructions (~340 KB of SASS, far past instruction cache). Left-looking emits ONE
// panel-factor body (~6.6k, the size of the measured-good `chol_rows<64>`) plus a ~130
// instruction update body, and the outer panel loop runs at runtime. Same dynamic
// instruction count, 3x less code.
//
// THE 64-FLOAT CEILING (probe_regsize) is what forces the shape: a per-thread register
// array over 64 floats goes entirely to local memory, so `r` holds ONE 64-wide panel of
// one row and is reused across panels -- never the whole 128- or 256-wide row.
//
// `ls[t * LDK + k]` caches the already-factored columns of every row, indexed by GLOBAL
// row so no compaction is needed between panels. LDK = 64*(NP-1) + 1; the +1 makes the
// own-row read (fixed k, consecutive t) hit 32 distinct banks. The panel-row read
// `ls[(c0+j)*LDK + k]` is uniform across the warp, so it broadcasts.
// n=128 -> LDK 65, 33 KB shared. n=256 -> LDK 193, 198 KB (needs the opt-in).
// **`#pragma unroll 1` ON THE PANEL LOOP IS LOAD-BEARING AND IT IS THE WHOLE FIX.**
// Without it nvcc unrolls the loop on its own, register pressure goes over the ceiling
// and `float r[64]` is evicted WHOLESALE to local memory. ncu 2026-07-27, benchmark
// index 2, one profile run each:
//
// no pragma 278.05 us 78 regs "accesses local memory, accounting for
// 93.57% of all sectors requested in L1TEX",
// 1.0 of every 32 bytes per sector used,
// 20.43% of local loads spilling to L2
// #pragma unroll 1 73.98 102 regs no local memory
// #pragma unroll 2 76.80 128 regs no local memory (full unroll, NP=2)
//
// This is the IDENTICAL signature `chol_diag_b` hit at 92.57%, where it cost 78.94 us
// against 17.38 until the register arrays were rewritten as named scalars. The n=128
// kernel had it all campaign and nobody looked, because `kernel_attrs()` -- the
// `stack > 0` gate 00_PLAN calls a hard rule -- listed four mid-driver kernels and none
// of the three small-n ones. It lists all eight now.
//
// FORBIDDING the unroll beats forcing it: UP=2 makes `c0` a compile-time constant and
// also clears the local memory, but costs 26 more registers and 3.8% of the duration.
template <int N, int UP>
__global__ __launch_bounds__(N)
void chol_fused(const float* __restrict__ src, float* __restrict__ dst) {
constexpr int NP = N / 64;
// **LDK WAS 64*(NP-1)+1 AND THE ODD STRIDE BLOCKED 16-BYTE ALIGNMENT.** 68 makes every
// row of `ls` float4-addressable, which is what lets the left-looking update below
// read both of its operands as LDS.128. The +1 existed so the own-row read (fixed k,
// consecutive t) hit 32 distinct banks; at 68 it hits 8, a 4-way conflict -- but the
// k-blocking already made that read four times rarer, and the panel-row read is
// warp-uniform in j so it broadcasts and never conflicted either way.
constexpr int LDK = 64 * (NP - 1) + 4;
const int t = (int)threadIdx.x;
const float* __restrict__ A = src + (size_t)blockIdx.x * N * N;
float* __restrict__ L = dst + (size_t)blockIdx.x * N * N;
extern __shared__ float sm[];
float* __restrict__ bc = sm; // [2][64] double-buffered broadcast column
float* __restrict__ ls = sm + 128; // [N * LDK] finished columns, by global row
float r[64];
#pragma unroll UP
for (int p = 0; p < NP; ++p) {
const int c0 = 64 * p;
const int lr = t - c0;
const bool live = (lr >= 0);
if (live) {
const float4* row4 =
reinterpret_cast<const float4*>(A + (size_t)t * N + c0);
#pragma unroll
for (int c = 0; c < 16; ++c) {
const float4 v = row4[c];
r[4 * c + 0] = v.x;
r[4 * c + 1] = v.y;
r[4 * c + 2] = v.z;
r[4 * c + 3] = v.w;
}
// r[j] -= sum_k L[t][k] * L[c0+j][k]. k runs at RUNTIME (ls is shared, so a
// dynamic index is free) which is what keeps the code small; j is unrolled so
// every r[] index folds to a constant and the array stays in registers.
// **k ADVANCES FOUR AT A TIME AND BOTH OPERANDS ARE ONE 16-BYTE LOAD.** This
// loop was 65 scalar shared loads per 64 FMAs -- one LDS per FMA, on the
// largest phase of the n=128 case. Both reads are contiguous in k, so blocking
// k by four makes it 65 LDS.128 per 256 FMAs. k stays ascending inside every
// r[j], so the result is BITWISE what the scalar loop produced.
for (int k = 0; k < c0; k += 4) {
const float4 ov = *reinterpret_cast<const float4*>(&ls[t * LDK + k]);
#pragma unroll
for (int j = 0; j < 64; ++j) {
const float4 pv =
*reinterpret_cast<const float4*>(&ls[(c0 + j) * LDK + k]);
r[j] -= ov.x * pv.x;
r[j] -= ov.y * pv.y;
r[j] -= ov.z * pv.z;
r[j] -= ov.w * pv.w;
}
}
}
// The chol_diag treatment, which measured -23.5% there and -11/-18% on chol_reg
// and chol_rows: rsqrt.approx instead of an IEEE sqrt + IEEE divide sitting on the
// per-column dependency chain; the broadcast column pre-scaled ONCE in shared by
// its owning thread rather than 63 redundant FMULs in every thread; no `lr >= j`
// predicate; and `bc` double-buffered so the closing barrier goes and the count
// stays at 2. Dropping the predicate is safe for the SAME reason as everywhere
// else, plus one specific to this kernel: an above-diagonal r[j] is written to
// `ls`, but every read of `ls` is `ls[R*LDK + k]` with k < c0 <= R, i.e. strictly
// lower triangle, so the garbage is never read back.
#pragma unroll
for (int k = 0; k < 64; ++k) {
float* cb = bc + 64 * (k & 1);
if (live && lr < 64 && lr >= k) cb[lr] = r[k];
__syncthreads();
const float a = fmaxf(cb[k], 1.17549435e-38f);
float rd;
asm("rsqrt.approx.f32 %0, %1;" : "=f"(rd) : "f"(a));
if (live) {
if (lr > k && lr < 64) {
const float v = r[k] * rd;
cb[lr] = v;
r[k] = v;
} else if (lr > k) {
r[k] *= rd; // rows below the panel: no shared slot to publish
} else if (lr == k) {
r[k] = a * rd;
}
}
// THE BARRIER GOES HERE, NOT AFTER THE READ LOOP. Putting it after cost one
// failed submission (914435): every thread consumed cb[j] before thread j had
// published the SCALED value into it. The pre-scaled-broadcast transform turns
// one shared write into two, and the second one needs its own barrier -- which
// is exactly why chol_diag and chol_rows have barriers at (1) unscaled write
// and (2) scaled write, and none after the read.
__syncthreads();
if (live) {
#pragma unroll
for (int j = k + 1; j < 64; ++j) r[j] -= r[k] * cb[j];
}
}
float4* out4 = reinterpret_cast<float4*>(L + (size_t)t * N + c0);
if (live) {
if (p + 1 < NP) {
#pragma unroll
for (int j = 0; j < 64; j += 4)
*reinterpret_cast<float4*>(&ls[t * LDK + c0 + j]) =
make_float4(r[j], r[j + 1], r[j + 2], r[j + 3]);
}
#pragma unroll
for (int c = 0; c < 16; ++c) {
float4 v;
v.x = (lr >= 4 * c + 0) ? r[4 * c + 0] : 0.0f;
v.y = (lr >= 4 * c + 1) ? r[4 * c + 1] : 0.0f;
v.z = (lr >= 4 * c + 2) ? r[4 * c + 2] : 0.0f;
v.w = (lr >= 4 * c + 3) ? r[4 * c + 3] : 0.0f;
out4[c] = v;
}
} else {
const float4 zero = make_float4(0.f, 0.f, 0.f, 0.f);
#pragma unroll
for (int c = 0; c < 16; ++c) out4[c] = zero; // upper-right block
}
__syncthreads(); // ls published
}
}
template <int N, int UP>
static torch::Tensor fused_chol(torch::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.scalar_type() == torch::kFloat32,
"fused_chol: fp32 (b,n,n)");
TORCH_CHECK(A.is_contiguous() && A.size(1) == N && A.size(2) == N,
"fused_chol: contiguous n=", N, " only");
constexpr int NP = N / 64;
constexpr int LDK = 64 * (NP - 1) + 4; // must match the kernel's LDK exactly
constexpr size_t SMEM = (size_t)(128 + N * LDK) * sizeof(float);
auto L = torch::empty_like(A);
static bool configured = false;
if (!configured) {
TORCH_CHECK(cudaFuncSetAttribute(chol_fused<N, UP>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)SMEM) == cudaSuccess,
"fused_chol: ", SMEM, " bytes of shared memory refused");
configured = true;
}
chol_fused<N, UP><<<(int)A.size(0), N, SMEM>>>(A.data_ptr<float>(),
L.data_ptr<float>());
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "fused_chol: launch failed");
return L;
}
torch::Tensor fused_chol128(torch::Tensor A) { return fused_chol<128, 1>(A); }
torch::Tensor fused_chol256(torch::Tensor A) { return fused_chol<256, 1>(A); }
torch::Tensor reg_chol32(torch::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.scalar_type() == torch::kFloat32, "reg_chol32: fp32 (b,n,n)");
TORCH_CHECK(A.is_contiguous() && A.size(1) == 32, "reg_chol32: contiguous n=32 only");
const int batch = (int)A.size(0);
auto L = torch::empty_like(A);
// WPB 8 -> 4 AND chol_reg -> chol_reg_sm. One profile run, four kernels, locked
// clocks: 38.24 / 38.88 (control, twice) against 32.80 at WPB=8 and 31.68 at WPB=4.
constexpr int WPB = 4;
chol_reg_sm<32, WPB><<<(batch + WPB - 1) / WPB, 32 * WPB>>>(
A.data_ptr<float>(), L.data_ptr<float>(), batch);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "reg_chol32: launch failed");
return L;
}
// PANEL WIDTH IS NOW A PARAMETER. probe_mid2 measured the scalar phases at
// n512b640: chol_diag 335 us, trsm_panel 915. Both scale with the panel width -- the
// TRSM work is sum_panels rows*W^2/2, which HALVES at W=32 -- while the trailing GEMM
// total is n^3/6 regardless. 32 is also comfortably inside the 64-float register
// ceiling, so `float r[32]` leaves headroom the 64-wide version does not have.
void chol_diag_inplace(torch::Tensor W, int64_t c0, int64_t width) {
TORCH_CHECK(W.dim() == 3 && W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"chol_diag_inplace: contiguous fp32 (b,n,n)");
const int n = (int)W.size(1);
TORCH_CHECK(n % width == 0, "chol_diag_inplace: n must be a multiple of width");
const long long off = c0 * (long long)n + c0;
if (width == 64) {
// **DEAD 2026-07-27 (918792): `chol_diag_bs<64,2>` COSTS +42% TO +102% ON EVERY
// LOW-BATCH CASE.** Clean run, untouched controls flat, and the damage is monotone
// in batch -- the THIRD time this exact signature has appeared on this board:
//
// b=64 n256b64 +0.0% b=8 n2048b8 +74%
// b=640 n512b640 +1.9% b=4 n1024b4 +77%
// b=60 n1024b60 +42% b=2 n4096b2 +79%
// b=16 n512b16 +94%
// b=2 n2048b2 +102%
//
// 917807's phase probe said S=2 was 5.4% FASTER and it was measuring the host, not
// the kernel: that run was starved, and 32 small sequential launches under a
// starved host report the DISPATCH rate, which is why all four variants landed
// within 10% of each other. The shipped control reading 8.14 us/launch against the
// historical 8.16 was a coincidence -- both numbers are ~8 us -- and I read it as
// calibration. **A phase probe made of many small launches is only valid if the
// run's own untouched controls are flat.**
chol_diag_b<64><<<(int)W.size(0), 64>>>(
W.data_ptr<float>(), (long long)n * n, off, n);
} else if (width == 32) {
chol_diag<32><<<(int)W.size(0), 32>>>(
W.data_ptr<float>(), (long long)n * n, off, n);
} else {
TORCH_CHECK(false, "chol_diag_inplace: width must be 32 or 64, got ", width);
}
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "chol_diag_inplace: launch failed");
}
void trsm_panel_inplace(torch::Tensor W, int64_t c0, int64_t width) {
TORCH_CHECK(W.dim() == 3 && W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"trsm_panel_inplace: contiguous fp32 (b,n,n)");
const int n = (int)W.size(1);
const int rows = n - (int)c0 - (int)width;
TORCH_CHECK(rows > 0, "trsm_panel_inplace: no rows below the diagonal block");
const dim3 grid((rows + 255) / 256, (unsigned)W.size(0));
const long long off_l = c0 * (long long)n + c0;
const long long off_p = (c0 + width) * (long long)n + c0;
if (width == 64) {
trsm_panel<64><<<grid, 256>>>(
W.data_ptr<float>(), (long long)n * n, off_l, off_p, n, rows);
} else if (width == 32) {
trsm_panel<32><<<grid, 256>>>(
W.data_ptr<float>(), (long long)n * n, off_l, off_p, n, rows);
} else {
TORCH_CHECK(false, "trsm_panel_inplace: width must be 32 or 64, got ", width);
}
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "trsm_panel_inplace: launch failed");
}
// Same kernel against a contiguous (b, w, w) STACK of triangular blocks rather than a
// diagonal block inside a big matrix. This is the base case of the giants' blocked
// inverse: `solve_triangular` gets no tensor cores at any size and the batched form is
// worse than the single one (8 x 256x256 measured at 617 us/op), so the recursion has to
// bottom out on our own kernel.
torch::Tensor tri_inv_stack(torch::Tensor L) {
TORCH_CHECK(L.dim() == 3 && L.is_contiguous() && L.scalar_type() == torch::kFloat32,
"tri_inv_stack: contiguous fp32 (b,w,w)");
const int w = (int)L.size(1);
TORCH_CHECK(L.size(2) == w, "tri_inv_stack: blocks must be square, got ",
L.size(1), "x", L.size(2));
auto out = torch::empty_like(L);
if (w == 64) {
tri_inv<64><<<(int)L.size(0), 64>>>(
L.data_ptr<float>(), out.data_ptr<float>(), (long long)w * w, 0, w);
} else if (w == 32) {
tri_inv<32><<<(int)L.size(0), 32>>>(
L.data_ptr<float>(), out.data_ptr<float>(), (long long)w * w, 0, w);
} else {
TORCH_CHECK(false, "tri_inv_stack: block width must be 32 or 64, got ", w);
}
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "tri_inv_stack: launch failed");
return out;
}
void tri_inv_panel(torch::Tensor W, torch::Tensor out, int64_t c0, int64_t width) {
TORCH_CHECK(W.dim() == 3 && W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"tri_inv_panel: contiguous fp32 (b,n,n)");
TORCH_CHECK(out.dim() == 3 && out.is_contiguous() && out.size(0) == W.size(0)
&& out.size(1) == width && out.size(2) == width,
"tri_inv_panel: out must be a contiguous (b,width,width) fp32 scratch");
const int n = (int)W.size(1);
const long long off_l = c0 * (long long)n + c0;
if (width == 64) {
tri_inv<64><<<(int)W.size(0), 64>>>(
W.data_ptr<float>(), out.data_ptr<float>(), (long long)n * n, off_l, n);
} else if (width == 32) {
tri_inv<32><<<(int)W.size(0), 32>>>(
W.data_ptr<float>(), out.data_ptr<float>(), (long long)n * n, off_l, n);
} else {
TORCH_CHECK(false, "tri_inv_panel: width must be 32 or 64, got ", width);
}
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "tri_inv_panel: launch failed");
}
// stack > 0 on tri_inv<64> means `float m[64]` was not promoted and the run is VOID --
// ---- THE GIANTS' PANEL PHASE: THREE STRIDED torch COPIES AROUND ONE GEMM ----
//
// 917037 put `panel` at 20.6% of n=32768 running at **339 TF/s while `trail` runs at 1330
// on the same dtype and the same data**. The arithmetic says the gap is not the GEMM:
// ~10 GB moved is 1.26 ms at HBM speed and the GEMM is 1.55 ms at trail's own 1330 TF/s,
// which is 2.8 ms against a measured 6083. **The missing 3.3 ms is three strided
// `at::copy_` passes at the ~1.46 TB/s this codebase already measured for torch's generic
// elementwise kernel** -- the `at::triu_tril_kernel` defect (46.08 us for 67 MB, 18% of
// HBM) that `tril_tiles` was written to fix on a phase worth 3.3%, while this one is worth
// 20.6% on three cases and was never asked about.
//
// pa = w[:, end:, j:end].to(bf16) strided fp32 -> contiguous bf16 COPY -> kernel 1
// pb = pa @ inv_t_bf16 GEMM (keep)
// w[:, end:, j:end].copy_(pb) contiguous bf16 -> strided fp32 COPY \ fused into
// lb[:, end:, j:end] = pb contiguous bf16 -> strided bf16 COPY / kernel 2
//
// Kernel 2 reads `pb` ONCE and writes both destinations, so it deletes a whole read of the
// panel as well as running at HBM speed. Eight elements per thread: two LDG.128 in, one
// STG.128 out (fp32 side is two).
//
// ROUNDING IS torch's. `__floats2bfloat162_rn` and `.to(torch.bfloat16)` are both
// round-to-nearest-even, and bf16 -> fp32 is exact, so both kernels are bitwise identical
// to the ops they replace.
//
// ALIGNMENT: the panel starts at row `end` and column `j`, both multiples of p = 2048, and
// ld = n is a multiple of 2048, so every base is 16-byte aligned and `c8` is a multiple
// of 8.
__global__ __launch_bounds__(256)
void panel_to_bf16(const float* __restrict__ src, __nv_bfloat16* __restrict__ dst,
int ld, int cols, int rows) {
const int row = (int)blockIdx.y;
const int c8 = ((int)blockIdx.x * 256 + (int)threadIdx.x) * 8;
if (c8 >= cols || row >= rows) return;
const float* s = src + (size_t)row * ld + c8;
const float4 a = *reinterpret_cast<const float4*>(s);
const float4 b = *reinterpret_cast<const float4*>(s + 4);
__nv_bfloat162 o[4];
o[0] = __floats2bfloat162_rn(a.x, a.y);
o[1] = __floats2bfloat162_rn(a.z, a.w);
o[2] = __floats2bfloat162_rn(b.x, b.y);
o[3] = __floats2bfloat162_rn(b.z, b.w);
*reinterpret_cast<float4*>(dst + (size_t)row * cols + c8) =
*reinterpret_cast<const float4*>(o);
}
// ONE pass over `pb` that lands both the fp32 answer and the bf16 mirror.
__global__ __launch_bounds__(256)
void panel_store(const __nv_bfloat16* __restrict__ pb, float* __restrict__ w,
__nv_bfloat16* __restrict__ lb, int ld, int cols, int rows) {
const int row = (int)blockIdx.y;
const int c8 = ((int)blockIdx.x * 256 + (int)threadIdx.x) * 8;
if (c8 >= cols || row >= rows) return;
const float4 v =
*reinterpret_cast<const float4*>(pb + (size_t)row * cols + c8); // 8 bf16
*reinterpret_cast<float4*>(lb + (size_t)row * ld + c8) = v; // bf16 straight
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&v);
const float2 t0 = __bfloat1622float2(h[0]);
const float2 t1 = __bfloat1622float2(h[1]);
const float2 t2 = __bfloat1622float2(h[2]);
const float2 t3 = __bfloat1622float2(h[3]);
float* d = w + (size_t)row * ld + c8;
*reinterpret_cast<float4*>(d) = make_float4(t0.x, t0.y, t1.x, t1.y);
*reinterpret_cast<float4*>(d + 4) = make_float4(t2.x, t2.y, t3.x, t3.y);
}
torch::Tensor panel_cast(torch::Tensor W, int64_t row0, int64_t col0,
int64_t rows, int64_t cols) {
TORCH_CHECK(W.dim() == 3 && W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"panel_cast: contiguous fp32 (1,n,n)");
TORCH_CHECK(cols % 8 == 0, "panel_cast: cols must be a multiple of 8, got ", cols);
const int ld = (int)W.size(2);
auto out = torch::empty({1, rows, cols},
W.options().dtype(torch::kBFloat16));
const dim3 grid((unsigned)((cols / 8 + 255) / 256), (unsigned)rows);
panel_to_bf16<<<grid, 256>>>(
W.data_ptr<float>() + row0 * ld + col0,
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), ld, (int)cols, (int)rows);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "panel_cast: launch failed");
return out;
}
void panel_write(torch::Tensor PB, torch::Tensor W, torch::Tensor LB,
int64_t row0, int64_t col0, int64_t rows, int64_t cols) {
TORCH_CHECK(PB.is_contiguous() && PB.scalar_type() == torch::kBFloat16,
"panel_write: pb must be contiguous bf16");
TORCH_CHECK(W.is_contiguous() && W.scalar_type() == torch::kFloat32,
"panel_write: W must be contiguous fp32");
TORCH_CHECK(LB.is_contiguous() && LB.scalar_type() == torch::kBFloat16,
"panel_write: LB must be contiguous bf16");
TORCH_CHECK(cols % 8 == 0, "panel_write: cols must be a multiple of 8, got ", cols);
const int ld = (int)W.size(2);
const dim3 grid((unsigned)((cols / 8 + 255) / 256), (unsigned)rows);
panel_store<<<grid, 256>>>(
reinterpret_cast<const __nv_bfloat16*>(PB.data_ptr()),
W.data_ptr<float>() + row0 * ld + col0,
reinterpret_cast<__nv_bfloat16*>(LB.data_ptr()) + row0 * ld + col0,
ld, (int)cols, (int)rows);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "panel_write: launch failed");
}
// read this before any timing. trsm_panel<64> (93 regs, stack 0) is the control it is
// built to match; chol_diag<64> (158, 0) is the other kernel in the panel loop.
// **THE LIST COVERED FOUR MID-DRIVER KERNELS AND MISSED THE THREE SMALL-n ONES, WHICH IS
// WHY `chol_fused<128>` SHIPPED FOR A WHOLE CAMPAIGN WITH `r[64]` IN LOCAL MEMORY** --
// ncu 2026-07-27, index 2: 93.57% of L1TEX sectors local, 1 of every 32 bytes used.
// A `stack > 0` gate that is not pointed at a kernel does not guard it. Every kernel that
// owns a `float r[N]` is in this list now; if any of the eight reports stack > 0 the run
// is void before its timings are worth reading.
torch::Tensor kernel_attrs() {
const void* fns[8] = {
(const void*)tri_inv<64>,
(const void*)trsm_panel<64>,
(const void*)chol_diag<64>,
(const void*)chol_diag_b<64>,
(const void*)chol_reg<32, 8>,
(const void*)chol_reg_sm<32, 4>,
(const void*)chol_rows<64, 1>,
(const void*)chol_fused<128, 1>,
};
auto out = torch::empty({8 * 4}, torch::kInt64);
int64_t* q = out.data_ptr<int64_t>();
for (int i = 0; i < 8; ++i) {
cudaFuncAttributes a{};
cudaFuncGetAttributes(&a, fns[i]);
q[4 * i + 0] = a.numRegs;
q[4 * i + 1] = (int64_t)a.localSizeBytes;
q[4 * i + 2] = (int64_t)a.sharedSizeBytes;
q[4 * i + 3] = a.maxThreadsPerBlock;
}
return out;
}
// ---- data.tril(): 2n^2 of traffic at 18% of HBM, on every mid case and every giant ----
//
// ncu, n2048b2: `at::triu_tril_kernel` takes 46.08 us to move 67 MB = 1.46 TB/s against
// B200's ~8. It is not bandwidth-bound (31% L1, 13.5 warps/scheduler): it recomputes a row
// and a column from a linear index with integer division FOR EVERY ELEMENT, and it reads
// the whole strict upper triangle just to overwrite it with zeros.
//
// One CTA per 64x64 tile fixes both. A strictly-upper tile is written as zeros and NEVER
// READ, which removes n^2/2 of reads outright (2n^2 -> 1.5n^2). A strictly-lower tile is a
// straight float4 copy. Only the diagonal tile is masked, and the branch is on blockIdx so
// it is CTA-uniform -- no divergence anywhere. Index math is a shift and a mask.
//
// Tiles are 64 wide, which is exactly the driver's panel width, so a diagonal tile IS a
// diagonal block. That matters: `chol_diag` rewrites every diagonal block masked, and the
// giants' `cholesky_ex` ignores the upper triangle of theirs, so this is safe for both
// drivers even where it is more conservative than it needs to be.
__global__ __launch_bounds__(128)
void tril_tiles(const float* __restrict__ src, float* __restrict__ dst, int n) {
const int ti = (int)blockIdx.y;
const int tj = (int)blockIdx.x;
const size_t base = (size_t)blockIdx.z * (size_t)n * (size_t)n
+ (size_t)ti * 64 * (size_t)n + (size_t)tj * 64;
#pragma unroll
for (int i = 0; i < 8; ++i) {
const int f = (int)threadIdx.x + 128 * i;
const int r = f >> 4;
const int c4 = f & 15;
const size_t o = base + (size_t)r * (size_t)n + 4 * c4;
if (tj < ti) {
*reinterpret_cast<float4*>(dst + o) =
*reinterpret_cast<const float4*>(src + o);
} else if (tj > ti) {
*reinterpret_cast<float4*>(dst + o) = make_float4(0.f, 0.f, 0.f, 0.f);
} else {
const float4 v = *reinterpret_cast<const float4*>(src + o);
float4 w;
w.x = (4 * c4 + 0 <= r) ? v.x : 0.0f;
w.y = (4 * c4 + 1 <= r) ? v.y : 0.0f;
w.z = (4 * c4 + 2 <= r) ? v.z : 0.0f;
w.w = (4 * c4 + 3 <= r) ? v.w : 0.0f;
*reinterpret_cast<float4*>(dst + o) = w;
}
}
}
torch::Tensor tril_copy(torch::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.is_contiguous() && A.scalar_type() == torch::kFloat32,
"tril_copy: contiguous fp32 (b,n,n)");
const int n = (int)A.size(1);
TORCH_CHECK(A.size(2) == n && n % 64 == 0,
"tril_copy: square with n a multiple of 64, got ", A.size(1), "x", A.size(2));
auto out = torch::empty_like(A);
const int tiles = n / 64;
const dim3 grid((unsigned)tiles, (unsigned)tiles, (unsigned)A.size(0));
tril_tiles<<<grid, 128>>>(A.data_ptr<float>(), out.data_ptr<float>(), n);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "tril_copy: launch failed");
return out;
}
torch::Tensor reg_chol64(torch::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.scalar_type() == torch::kFloat32, "reg_chol64: fp32 (b,n,n)");
TORCH_CHECK(A.is_contiguous() && A.size(1) == 64, "reg_chol64: contiguous n=64 only");
const int batch = (int)A.size(0);
auto L = torch::empty_like(A);
// **MPB 4 -> 1. ONE LINE, AND IT IS A BARRIER FIX.** `chol_rows` runs a __syncthreads()
// pair per column, and at MPB=4 that barrier spans 256 threads = 8 warps covering FOUR
// INDEPENDENT MATRICES that have no reason to wait for each other. ncu on the mid
// driver's diagonal factor measured the CTA barrier as the dominant stall (34.7% of
// 10.9 cycles) and showed the cost rising with warps per block. At MPB=1 the barrier
// spans 2 warps, and the grid goes 256 -> 1024 blocks over 148 SMs, which is better
// balanced as well. Warps per SM are unchanged: 217 registers caps it at 8 either way.
constexpr int MPB = 1;
chol_rows<64, MPB><<<(batch + MPB - 1) / MPB, 64 * MPB>>>(
A.data_ptr<float>(), L.data_ptr<float>(), batch);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "reg_chol64: launch failed");
return L;
}
"""
_REG = None
try:
_REG = load_inline(
name="chol_reg32_v17",
cpp_sources=("torch::Tensor reg_chol32(torch::Tensor A);\n"
"torch::Tensor reg_chol64(torch::Tensor A);\n"
"torch::Tensor fused_chol128(torch::Tensor A);\n"
"torch::Tensor fused_chol256(torch::Tensor A);\n"
"torch::Tensor kernel_attrs();\n"
"torch::Tensor tril_copy(torch::Tensor A);\n"
"torch::Tensor panel_cast(torch::Tensor W, int64_t row0,"
" int64_t col0, int64_t rows, int64_t cols);\n"
"void panel_write(torch::Tensor PB, torch::Tensor W,"
" torch::Tensor LB, int64_t row0, int64_t col0, int64_t rows,"
" int64_t cols);\n"
"torch::Tensor tri_inv_stack(torch::Tensor L);\n"
"void chol_diag_inplace(torch::Tensor W, int64_t c0, int64_t width);\n"
"void trsm_panel_inplace(torch::Tensor W, int64_t c0, int64_t width);\n"
"void tri_inv_panel(torch::Tensor W, torch::Tensor out, int64_t c0,"
" int64_t width);"),
cuda_sources=_REG_CUDA,
functions=["reg_chol32", "reg_chol64", "fused_chol128", "fused_chol256",
"kernel_attrs", "tril_copy", "panel_cast", "panel_write",
"tri_inv_stack", "chol_diag_inplace",
"trsm_panel_inplace", "tri_inv_panel"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
except Exception:
_REG = None
if _REG is not None:
try:
_A = _REG.kernel_attrs().tolist()
for _i, _name in enumerate(("tri_inv<64>", "trsm_panel<64>", "chol_diag<64>",
"chol_diag_b<64>", "chol_reg<32,8>",
"chol_reg_sm<32,4>", "chol_rows<64,1>",
"chol_fused<128,1>")):
print(f"KATTR {_name:<18} regs={_A[4 * _i]:<4} "
f"stack={_A[4 * _i + 1]:<5} smem={_A[4 * _i + 2]:<7} "
f"maxtpb={_A[4 * _i + 3]}", file=sys.stderr)
except Exception as _exc:
print(f"KATTR failed: {_exc}", file=sys.stderr)
# ---------------------------------------------------------------- n=128 cuSOLVERDx potrf
_NVCC = "/usr/local/cuda/bin/nvcc"
_UFLAGS = [
"-U__CUDA_NO_HALF_OPERATORS__", "-U__CUDA_NO_HALF_CONVERSIONS__",
"-U__CUDA_NO_HALF2_OPERATORS__", "-U__CUDA_NO_BFLOAT16_OPERATORS__",
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__", "-U__CUDA_NO_BFLOAT162_OPERATORS__",
]
_DX_CU = r"""
#include <cusolverdx.hpp>
using namespace cusolverdx;
template <int N>
using Chol = decltype(Size<N, N>()
+ Precision<float>()
+ Type<type::real>()
+ Function<potrf>()
+ Arrangement<arrangement::row_major>()
+ FillMode<fill_mode::lower>()
+ BlockDim<256>()
+ SM<1000>()
+ Block());
template <int N>
__global__ __launch_bounds__(Chol<N>::max_threads_per_block)
void dx_potrf_kernel(float* A, typename Chol<N>::status_type* info) {
using S = Chol<N>;
CUSOLVERDX_SKIP_IF_NOT_APPLICABLE_SM(S);
extern __shared__ cusolverdx::byte shared_mem[];
float* As = reinterpret_cast<float*>(shared_mem);
constexpr int lds = S::lda;
float* Am = A + (size_t)blockIdx.x * N * N;
for (int idx = threadIdx.x; idx < N * N; idx += blockDim.x) {
const int i = idx / N;
const int j = idx - i * N;
As[i * lds + j] = Am[idx];
}
__syncthreads();
S().execute(As, info + blockIdx.x);
__syncthreads();
for (int idx = threadIdx.x; idx < N * N; idx += blockDim.x) {
const int i = idx / N;
const int j = idx - i * N;
Am[idx] = (i >= j) ? As[i * lds + j] : 0.f;
}
}
extern "C" int dx_launch_potrf(float* A, int batch, int n, int* info) {
if (n != 128) return -2;
using S = Chol<128>;
cudaFuncSetAttribute(dx_potrf_kernel<128>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)S::shared_memory_size);
dx_potrf_kernel<128><<<batch, S::max_threads_per_block, S::shared_memory_size>>>(
A, reinterpret_cast<typename S::status_type*>(info));
return (int)cudaGetLastError();
}
"""
_DX_CPP = r"""
#include <torch/extension.h>
extern "C" int dx_launch_potrf(float* A, int batch, int n, int* info);
torch::Tensor batched_potrf(torch::Tensor A) {
TORCH_CHECK(A.dim() == 3 && A.scalar_type() == torch::kFloat32, "expect fp32 (b,n,n)");
auto L = A.contiguous().clone();
auto info = torch::zeros({L.size(0)}, L.options().dtype(torch::kInt32));
int rc = dx_launch_potrf(L.data_ptr<float>(), (int)L.size(0), (int)L.size(1),
info.data_ptr<int>());
TORCH_CHECK(rc == 0, "dx_launch_potrf rc=", rc);
return L;
}
"""
def _build_dx():
"""Hand-rolled compile + `nvcc -dlink` against libcusolverdx.a, injected into
load_inline via extra_ldflags. load_inline cannot device-link; recipe proven by
exp_dxlink 2026-07-23, do not re-derive."""
work = tempfile.mkdtemp(prefix="dxv12_")
cu, obj, dlink = (os.path.join(work, f) for f in ("kernel.cu", "kernel.o", "dlink.o"))
lib = "/opt/mathdx/lib/libcusolverdx.a"
with open(cu, "w") as fh:
fh.write(_DX_CU)
def run(cmd):
proc = subprocess.run(cmd, capture_output=True, text=True, timeout=300)
if proc.returncode != 0:
raise RuntimeError(f"{cmd[0]} rc={proc.returncode}: "
f"{((proc.stderr or '') + (proc.stdout or ''))[-1200:]}")
run([_NVCC, "-std=c++17", "-O3", "-arch=sm_100a", "-rdc=true", "-dlto",
"-Xcompiler=-fPIC", "-I/opt/mathdx/include",
"-I/opt/mathdx/external/cutlass/include"] + _UFLAGS + ["-c", cu, "-o", obj])
run([_NVCC, "-arch=sm_100a", "-dlto", "-dlink", "-Xcompiler=-fPIC", obj, lib,
"-o", dlink])
return load_inline(
name="dxpotrf_v12",
cpp_sources=_DX_CPP,
functions=["batched_potrf"],
extra_ldflags=[obj, dlink, lib, "-L/usr/local/cuda/lib64", "-lcudart", "-lcudadevrt"],
verbose=False,
)
try:
_DX = _build_dx()
except Exception:
_DX = None
# ------------------------------------------------------------------- v8 routes (unchanged)
@triton.jit
def _chol32_kernel(input_ptr, output_ptr, N: tl.constexpr):
matrix = tl.program_id(0)
row_ids = tl.arange(0, N)
col_ids = tl.arange(0, N)
rows = row_ids[:, None]
cols = col_ids[None, :]
offsets = matrix * (N * N) + rows * N + cols
values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)
for k in range(N):
row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
diagonal = tl.sum(tl.where(col_ids == k, row, 0.0), axis=0)
diagonal -= tl.sum(tl.where(col_ids < k, row * row, 0.0), axis=0)
diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))
column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
products = tl.where(cols < k, values * row[None, :], 0.0)
column = (column - tl.sum(products, axis=1)) / diagonal
values = tl.where((rows == k) & (cols == k), diagonal, values)
values = tl.where((rows > k) & (cols == k), column[:, None], values)
tl.store(output_ptr + offsets, values)
@triton.jit
def _panel_solve_kernel(w_ptr, inv_ptr, mat_stride, off_p, ld, rows,
NW: tl.constexpr, BM: tl.constexpr, IP: tl.constexpr):
"""P <- P @ L11^-T for the rows below the diagonal block, as ONE tl.dot.
This is the whole point of the file. `trsm_panel<64>` does the same work as a scalar
forward substitution at 5.1 TF/s against tf32's 1.1 PF/s; with L11 already inverted
the solve is a plain [BM,64] x [64,64] GEMM and reaches the tensor cores.
"""
pid_m = tl.program_id(0)
pid_b = tl.program_id(1)
rm = pid_m * BM + tl.arange(0, BM)
cn = tl.arange(0, NW)
mask = rm < rows
base = w_ptr + pid_b.to(tl.int64) * mat_stride + off_p
p_ptr = base + rm[:, None].to(tl.int64) * ld + cn[None, :]
panel = tl.load(p_ptr, mask=mask[:, None], other=0.0)
i_ptr = inv_ptr + pid_b.to(tl.int64) * (NW * NW) + cn[:, None] * NW + cn[None, :]
linv = tl.load(i_ptr)
out = tl.dot(panel, tl.trans(linv), input_precision=IP)
tl.store(p_ptr, out, mask=mask[:, None])
def _eager(data: torch.Tensor) -> torch.Tensor:
return torch.linalg.cholesky_ex(data, check_errors=False).L
def _masked32(data: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(data)
_chol32_kernel[(data.shape[0],)](data, out, N=32, num_warps=1)
return out
def _loop_single(data: torch.Tensor, fn=None) -> torch.Tensor:
"""One matrix at a time. cuSOLVER's batched potrf pays ~1 us of serial latency per
COLUMN regardless of batch (04_DIAGNOSIS), so at n=4096 a batch of 2 costs 3.7x a
single -- and that applies to the diagonal blocks inside a blocked driver too."""
out = torch.empty_like(data)
for i in range(data.shape[0]):
one = data[i : i + 1]
out[i] = (torch.linalg.cholesky_ex(one[0], check_errors=False).L
if fn is None else fn(one)[0])
return out
_GINV_BASE = 64
def _diag_blocks(t: torch.Tensor, s: int) -> torch.Tensor:
"""View of the s x s blocks down the diagonal of a contiguous (b, p, p) tensor.
CONTIGUITY IS CHECKED, NOT ASSUMED. `torch.linalg.cholesky_ex(...).L` comes back
COLUMN-MAJOR from LAPACK, and these strides silently address the wrong elements on
such a tensor -- it produced a plausible-looking inverse that was 46% wrong.
"""
b, p, _ = t.shape
if not t.is_contiguous():
raise ValueError("_diag_blocks: tensor must be row-major contiguous")
return t.as_strided((b, p // s, s, s), (p * p, s * p + s, p, 1), t.storage_offset())
def _pair_blocks(t: torch.Tensor, s: int, ri: int, ci: int) -> torch.Tensor:
"""View of one corner of every 2s x 2s diagonal pair, as (b, pairs, s, s).
(ri, ci) selects the corner within the pair: (0,0) is A, (1,1) is B, (1,0) is C in
[[A, 0], [C, B]]. All four share one stride pattern and differ only in the offset,
which is what lets a whole merge level run as two batched matmuls with no gather.
"""
b, p, _ = t.shape
if not t.is_contiguous():
raise ValueError("_pair_blocks: tensor must be row-major contiguous")
return t.as_strided((b, p // (2 * s), s, s),
(p * p, 2 * s * p + 2 * s, p, 1),
t.storage_offset() + ri * s * p + ci * s)
def _tri_inv_blocked(l: torch.Tensor, base: int = _GINV_BASE) -> torch.Tensor:
"""Inverse of a lower-triangular (b, p, p), every flop above `base` on tensor cores.
THE DEFECT: `_left_looking` builds each panel's inverse with `solve_triangular`,
which probe_prim2 priced at 899 us at p=2048 and 3230 at p=4096, and which gets NO
tensor cores (3.2 TF/s measured). At n=32768 that is 15 inverses = ~13.5 ms of the
case's 42.6, the single largest addressable block left on the board.
inv([[A, 0], [C, B]]) = [[inv(A), 0], [-inv(B) C inv(A), inv(B)]]
Level 0 inverts all p/base diagonal blocks in ONE `tri_inv` launch; each merge then
doubles the block size with two batched matmuls, so the whole thing is 2 p^3/3 flops
of tf32 GEMM -- 5.7 GFLOP at p=2048, ~15 us of tensor-core time against 899.
THE BASE CASE MUST BE OUR OWN KERNEL. Batched `solve_triangular` is worse than the
single call, not better: probe_graph measured 8 x 256x256 at 617 us/op. Recursing
down to a vendor call at any width keeps the thing this is trying to remove.
"""
# `_left_looking` was pure torch before this; it now depends on a load_inline build
# that is first exercised on B200, so a build failure must degrade to v21's giants
# rather than take all three cases down with it. A failed build prints no KATTR
# lines, which is how it stays visible.
p = l.shape[-1]
if _REG is None or p % base or (p // base) & (p // base - 1):
return torch.linalg.solve_triangular(
l, torch.eye(p, device=l.device, dtype=l.dtype).expand_as(l), upper=False)
l = l.contiguous() # cholesky_ex returns column-major; the views need C order
b = l.shape[0]
m = torch.zeros(l.shape, dtype=l.dtype, device=l.device)
blocks = _diag_blocks(l, base).reshape(-1, base, base).contiguous()
_diag_blocks(m, base).copy_(
_REG.tri_inv_stack(blocks).view(b, p // base, base, base))
s = base
while s < p:
a_inv = _pair_blocks(m, s, 0, 0)
b_inv = _pair_blocks(m, s, 1, 1)
c = _pair_blocks(l, s, 1, 0)
# `.copy_(-(...))` UNTIL 2026-07-28. **THE DESTINATION IS PROVABLY ZERO HERE, SO
# THE NEGATION IS A SUBTRACT AND THE SEPARATE `neg` PASS IS PURE WASTE.**
# `m` starts as `torch.zeros`, and the (1,0) corner at level s covers rows
# [2qs+s, 2qs+2s) x cols [2qs, 2qs+s) -- an OFF-diagonal block. Every smaller
# level, and the `_diag_blocks` base case, writes only ON-diagonal blocks, so
# nothing has touched this region: 0 - X is exactly -X, bitwise.
#
# Saves one kernel launch and one allocation per merge level (5 levels x 15 calls
# = 75 launches at n=32768) and one pass over the data: `neg` (read X, write T)
# plus `copy_` (read T, write dest) is four passes; `sub_` is three.
#
# This is NOT the `out=` substitution that measured 24% slower (916151) -- that
# put a SCATTERED write inside a cuBLAS GEMM epilogue. This is an elementwise op
# on a strided view whose rows are contiguous runs, which TensorIterator
# collapses (the same correction `trail_sub` established).
_pair_blocks(m, s, 1, 0).sub_((b_inv @ c) @ a_inv)
s *= 2
return m
def _tril(data: torch.Tensor) -> torch.Tensor:
"""Lower triangle of a batched square matrix, as a fresh contiguous tensor.
`torch.tril` runs at 18% of HBM here (ncu: 46 us for 67 MB at n2048b2) because it does
integer division per element and reads the upper triangle before zeroing it. Falls back
to it if the build failed or n is not a multiple of 64."""
if _REG is None or data.shape[-1] % 64 or not data.is_contiguous():
return data.tril()
return _REG.tril_copy(data)
def _left_looking(data: torch.Tensor, panel_n: int, bf16: bool) -> torch.Tensor:
"""LEFT-looking blocked Cholesky. Replaces `_blocked_bf16_trsmfree`.
THE DEFECT IT FIXES: the right-looking version updated the FULL SQUARE trailing
block (`w[:, end:, end:] -= panel @ panel.T`), i.e. 2 n^3/3 flops where the
algorithm needs n^3/3. Half of every trailing GEMM was recomputing the symmetric
upper triangle, on the four cases that are 32% of the geomean.
Left-looking is the structural fix rather than a triangular patch: block column
[j:n, j:j+p] is updated ONCE, by one GEMM against everything to its left, so the
total is exactly n^3/3 with no per-panel triangle bookkeeping and FEWER torch ops
(one GEMM per panel instead of one per panel per column block).
Two other defects, same origin:
- `.to(bfloat16)` ran on the panel THREE times per panel step. The bf16 mirror
`lb` is written once per block column and read by every later panel.
- the trailing update went through a bf16-OUTPUT matmul, so the accumulated
Schur complement was rounded to 8 mantissa bits. Left-looking rounds the
total once instead of once per panel, at the same GEMM rate.
PRECISION, SIMULATED ON CPU against reference.py's gates (dense cond=2, which is
the only distribution that reaches n>=4096 -- task.yml's 17 tests stop at n=2048).
At a FIXED PANEL COUNT the absolute relative residual is flat in n while the budget
grows as 2.38e-6*n, so margin doubles per octave. Measured at n=512/1024/2048:
4 panels tf32 8.7e-4 / 8.6e-4 / 9.0e-4 8 panels tf32 9.2e-4 -> 1.1e-3
4 panels bf16 2.7e-3 / 2.8e-3 / 2.9e-3 8 panels bf16 2.9e-3 -> 3.4e-3
Extrapolated to what ships: n4096 p1024 tf32 9% of budget, n8192 p2048 bf16 15%,
n16384 p2048 bf16 9%, n32768 p4096 bf16 4.6%. The simulation TRUNCATES tf32
mantissas where the tensor core rounds, so tf32 has another ~2x in hand.
bf16 at n=4096 would be ~30% of budget -- fine, but its whole n^3/3 is 23 GFLOP
(~25 us either way), so tf32 buys 3x the margin for nothing.
"""
# `data.tril()`, not `data.clone()`: left-looking never writes the strict upper block
# triangle, and every diagonal block is overwritten by `ljj`, which cholesky_ex
# returns already zeroed above the diagonal. So starting from a zeroed upper triangle
# keeps it zeroed and the closing `tril_()` -- a full 4.3 GB read+write pass at
# n=32768, ~1.1 ms of pure memory traffic -- disappears. This is the identical fix
# v19 applied to the mids; the giants never got it.
w = _tril(data)
batch, n, _ = w.shape
if n % panel_n:
raise ValueError(f"_left_looking: panel {panel_n} does not divide n={n}")
# `eye` is gone with solve_triangular -- _tri_inv_blocked needs no RHS.
# Uninitialised on purpose: every block read at panel j was written at panel c < j.
lb = torch.empty_like(w, dtype=torch.bfloat16) if bf16 else None
old = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
for j in range(0, n, panel_n):
end = j + panel_n
if j:
if bf16:
w[:, j:, j:end] -= (lb[:, j:, :j]
@ lb[:, j:end, :j].transpose(-1, -2)
).to(torch.float32)
else:
w[:, j:, j:end].baddbmm_(w[:, j:, :j],
w[:, j:end, :j].transpose(-1, -2),
beta=1, alpha=-1)
djj = w[:, j:end, j:end]
ljj = torch.linalg.cholesky_ex(djj, check_errors=False).L
djj[:] = ljj
if bf16:
lb[:, j:end, j:end] = ljj.to(torch.bfloat16)
if end < n:
# inv + GEMM, never solve_triangular on the panel: the batched TRSM
# measured 480 us/op and is why every earlier blocked attempt lost.
# The INVERSE itself was still a solve_triangular (899 us at p=2048, no
# tensor cores); it is now blocked GEMMs down to a tri_inv<64> base.
inv_t = _tri_inv_blocked(ljj).transpose(-1, -2)
# batch == 1 because both kernels address `W + row0*ld + col0` with no
# batch stride. Every giant on this board is batch 1; anything else falls
# through to the torch path below rather than addressing the wrong matrix.
if bf16 and _REG is not None and batch == 1 and panel_n % 8 == 0:
# THREE STRIDED `at::copy_` PASSES BECOME TWO TILED KERNELS. torch's
# generic elementwise kernel runs a strided 3-D view at ~1.46 TB/s
# (measured here on `at::triu_tril_kernel`: 46.08 us for 67 MB, 18% of
# HBM); these move the same bytes at HBM speed, and `panel_write` reads
# `pb` ONCE to land both destinations instead of twice. Bitwise
# identical: both directions round to nearest even and bf16 -> fp32
# is exact.
pa = _REG.panel_cast(w, end, j, n - end, panel_n)
pb = pa @ inv_t.to(torch.bfloat16)
_REG.panel_write(pb, w, lb, end, j, n - end, panel_n)
elif bf16:
pb = w[:, end:, j:end].to(torch.bfloat16) @ inv_t.to(torch.bfloat16)
# `= pb.to(torch.float32)` UNTIL 2026-07-27. That materialised a full
# fp32 copy of the panel and then copied THAT into the strided view --
# two passes where `copy_` does the dtype conversion inline in one.
# At n=32768 the panel columns total 5.03e8 elements per sweep, so the
# temp alone was ~2 GB written and ~2 GB read back, ~0.67 ms of the
# 30.7 ms case. bf16 -> fp32 is EXACT, so this is bitwise identical.
#
# 917037 decomposed this case and it is why the phase was worth
# looking at: panel 6083 us (20.6%) running at 339 TF/s against the
# trailing GEMM's 1330 TF/s on the same dtype, the gap being three
# full materialisation passes around a perfectly good GEMM.
w[:, end:, j:end].copy_(pb)
lb[:, end:, j:end] = pb # already bf16; do not re-convert
else:
w[:, end:, j:end] = w[:, end:, j:end] @ inv_t
finally:
torch.backends.cuda.matmul.allow_tf32 = old
return w
def _blocked_reg(data: torch.Tensor, tf32: bool, panel: int = 64) -> torch.Tensor:
"""LEFT-looking blocked Cholesky with the register kernels. Replaces the
right-looking version banked in v18.
probe_mid2 decomposed n512b640 (2620 us) and n1024b60 (1868) into their phases and
the four phases SUM TO THE TOTAL -- there is no dispatch gap at high batch. The
op-count framing came from n1024b4, which is low-batch and latency-bound, and does
not generalise. So the targets are flops and memory passes, not launches:
phase n512b640 n1024b60 fixed by
chol_diag 335 (13%) 305 (16%) narrower panel
trsm_panel 915 (35%) 560 (30%) narrower panel (work ~ W)
trailing baddbmm_ 820 (31%) 680 (36%) LEFT-LOOKING, 1.67x fewer MACs
clone + tril_ 590 (23%) 300 (16%) start from data.tril()
1. LEFT-LOOKING. The right-looking `w[:, ce:, ce:].baddbmm_(panel, panel.T)`
computed the FULL SQUARE trailing block -- the same defect the giants had. Total
MACs 3.67e7 vs left-looking's 2.24e7 at n=512, and the op count is IDENTICAL
(one GEMM per panel either way), so this is free.
2. NO FINAL tril_. Left-looking only ever writes block columns [c0:n, c0:c0+W],
i.e. the lower block triangle, so if the upper triangle starts at zero it stays
there. `data.tril()` does the clone and the zeroing in ONE pass. Verified safe
against both kernels: chol_diag loads the full diagonal-block row but uses only
lower entries and masks its store, and trsm_panel only ever touches
below-diagonal blocks.
"""
w = _tril(data)
n = w.shape[-1]
old = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = tf32
try:
for c0 in range(0, n, panel):
ce = c0 + panel
if c0:
# Column blocks [0:c0) and [c0:ce) are disjoint, so this reads and
# writes non-overlapping memory despite being one tensor.
w[:, c0:, c0:ce].baddbmm_(w[:, c0:, :c0],
w[:, c0:ce, :c0].transpose(-1, -2),
beta=1, alpha=-1)
_REG.chol_diag_inplace(w, c0, panel)
if ce < n:
_REG.trsm_panel_inplace(w, c0, panel)
finally:
torch.backends.cuda.matmul.allow_tf32 = old
return w
_PANEL_BM = 64 # RE-FIT 2026-07-26: 128 was never tuned, and since exp_fast4 this
# tile drives five cases, not two. Every one of them is CTA-starved --
# ncu counts 4-32 CTAs on 148 SMs -- so halving the tile doubles the
# grid where occupancy, not tile efficiency, is the binding constraint.
def _blocked_tc_fp32(data: torch.Tensor, panel: int = 64) -> torch.Tensor:
"""`_blocked_tc`'s STRUCTURE with `_blocked_reg`'s PRECISION: no tf32 anywhere.
**THE PHASE PROBE (916966) SAYS `trsm_panel` IS THE LARGEST SINGLE COST ON THE BOARD.**
Decomposition of `_blocked_reg` at (4,1024), where the phases sum to 607 against a
measured full of 663 -- an 8.4% dispatch gap, so the attribution is real:
trsm_panel 335.8 us 51% 22.4 us per launch, FLAT in batch (22.6 at b16)
trailing 141.2 us 21%
chol_diag 130.1 us 20% 8.13 us per launch, FLAT in n AND batch
dispatch 56 us 8%
A triangular solve cannot use tensor cores and a multiply by an inverse can -- but the
win here is not the tensor cores, it is the ALGORITHM. Inverting the 64x64 block once
with all 64 columns in parallel and multiplying beats a 64-step forward substitution
per row even in plain fp32: the same probe priced `tri_inv` at 6.65 us and the panel
multiply at 7.90, i.e. **14.55 against trsm_panel's 22.4**.
**THIS IS WHY THE PRECISION IS UNCHANGED AND MUST STAY UNCHANGED.** `_blocked_tc`
cannot serve these two cases: `exp_ll4096` (916741) failed RANKED VALIDATION, and the
failure localises to the tf32 TRAILING GEMM at n=1024 low batch, not to the panel
solve. So `allow_tf32` is forced FALSE here and the panel multiply is a plain fp32
`bmm`, not `tl.dot` -- `IP="ieee"` was measured at n512b16 189 -> 465 (916673) and is
not an option. Every value this produces is fp32, exactly as `_blocked_reg(tf32=False)`
produced it, so there is no new correctness risk.
Routed to the two cases that run `_blocked_reg(tf32=False)`: n256b64 and n1024b4.
"""
w = _tril(data)
n = w.shape[-1]
linv = torch.empty((w.shape[0], panel, panel), dtype=w.dtype, device=w.device)
old = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = False
try:
for c0 in range(0, n, panel):
ce = c0 + panel
if c0:
w[:, c0:, c0:ce].baddbmm_(w[:, c0:, :c0],
w[:, c0:ce, :c0].transpose(-1, -2),
beta=1, alpha=-1)
_REG.chol_diag_inplace(w, c0, panel)
if ce < n:
_REG.tri_inv_panel(w, linv, c0, panel)
# `out=` on this strided view is NOT an optimisation to try here: the same
# substitution in `_tri_inv_blocked` measured 24% SLOWER (916151), because
# a GEMM writing a scattered output loses to a contiguous result plus a
# dedicated strided copy.
w[:, ce:, c0:ce] = w[:, ce:, c0:ce] @ linv.transpose(-1, -2)
finally:
torch.backends.cuda.matmul.allow_tf32 = old
return w
def _blocked_tc(data: torch.Tensor, panel: int = 64) -> torch.Tensor:
"""`_blocked_reg` with the panel TRSM moved onto the tensor cores.
ONE phase changes. probe_mid2 priced the mid driver at n512b640 as chol_diag 335 /
trsm_panel 915 / trailing baddbmm_ 820 / clone+tril 590. The trailing update has been
on tf32 tensor cores since v16 and the clone was fused in v19, which leaves
`trsm_panel` as the largest remaining block of scalar work on the board -- 35% of the
case at 5.1 TF/s, against tf32's 1.1 PF/s.
A triangular solve cannot use tl.dot, but a multiply by the inverse can, so the panel
now costs one extra kernel (`tri_inv`, 64 columns) to save a 448-row substitution.
Op count per panel goes 3 -> 4, which RULE 0 prices at roughly +4.5% ranked; the
solve should give back several times that. Both effects are in the same run, so read
the per-case BENCHMARK table for the phase change and the LEADERBOARD for the net.
NOT USED WHERE tf32 IS UNSAFE. v9 measured tf32 tl.dot passing recon at n=512 and
FAILING at n=256, so this path is gated to exactly the cases v19 already runs tf32
on; n=256 and low-batch n=1024 stay on `_blocked_reg` and should come back unchanged,
which also makes them the control row of the table.
"""
w = _tril(data)
n = w.shape[-1]
linv = torch.empty((w.shape[0], panel, panel), dtype=w.dtype, device=w.device)
old = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
try:
for c0 in range(0, n, panel):
ce = c0 + panel
if c0:
w[:, c0:, c0:ce].baddbmm_(w[:, c0:, :c0],
w[:, c0:ce, :c0].transpose(-1, -2),
beta=1, alpha=-1)
_REG.chol_diag_inplace(w, c0, panel)
if ce < n:
_REG.tri_inv_panel(w, linv, c0, panel)
rows = n - ce
_panel_solve_kernel[((rows + _PANEL_BM - 1) // _PANEL_BM, w.shape[0])](
w, linv, n * n, ce * n + c0, n, rows,
NW=panel, BM=_PANEL_BM, IP="tf32", num_warps=4,
)
finally:
torch.backends.cuda.matmul.allow_tf32 = old
return w
def custom_kernel(data: torch.Tensor) -> torch.Tensor:
batch, n, _ = data.shape
# --- giants + n4096/n8192: left-looking, half the flops of the shipped version -----
# n=4096 stays fp32/tf32 (bf16 would spend 40% of its 9.8e-3 recon budget on one
# rounding) and goes one matrix at a time, so its diagonal blocks reach cuSOLVER's
# native potrf instead of potrfBatched's 1 us/column.
# n=4096 AT BATCH >= 2 GOES TO THE BLOCKED DRIVER. Re-fit 2026-07-26 off the first
# per-case table since v23. The driver is now (n/64) x ~29 us and essentially
# BATCH-INDEPENDENT below b60 -- measured n2048b2 922 / 32 panels = 28.8, n2048b8
# 1002 / 32 = 31.3, n512b16 203 / 8 = 25.4. It therefore factors the WHOLE batch in
# one panel sweep, where `_loop_single` pays cuSOLVER's full cost per matrix in
# series: n4096b2 = 2 x 1604 = 3208. Predicted 64 x ~31 = ~2000 for both matrices.
#
# This is the same crossover that just paid at n=2048 (1342 -> 922), one octave up:
# our per-matrix cost at n=2048 is now 461 us against cuSOLVER's 649, so the vendor's
# advantage at n=4096 is only its batch-1 native path, and that advantage disappears
# the moment there is more than one matrix to amortise the panel sweep over.
#
# BATCH 1 IS DELIBERATELY EXCLUDED: 64 panels x 31 = ~1980 against cuSOLVER's 1604,
# so a single matrix still loses. tf32 needs no gate here -- the recon budget at
# n=4096 is 9.8e-3, roughly 40x the tf32 GEMM error, and 00_PLAN records tf32
# trailing updates as SAFE for n >= 4096.
if n == 4096 and batch >= 2 and _REG is not None and data.is_contiguous():
try:
return _blocked_tc(data)
except Exception:
pass
if n == 4096:
# MEASURED LOSS, reverted: left-looking p1024 gave 2600 / 5190 vs cuSOLVER's
# 1530 / 3210. probe_prim2 says why -- 4 x potrf(1024) = 1336 plus 3 x
# inv(1024) = 1068 is already 2.4 ms, and cuSOLVER factors the WHOLE 4096
# in 1530. A blocked driver cannot win where its own diagonal blocks cost
# more than the vendor's entire factorization.
return _loop_single(data)
if n == 8192:
return _left_looking(data, 2048, bf16=True)
if n == 16384:
return _left_looking(data, 2048, bf16=True)
if n == 32768:
# p 4096 -> 2048 MEASURED 47.4 -> 42.6 ms (fused2_bench). Only the triangular
# inverse depends on the panel -- (n/p) x potrf(p) = 0.33n is invariant -- so
# this is the whole effect. It came in at +10%, not the predicted +24%, so the
# p^1.6 inverse law priced on 2-D solve_triangular does NOT transfer cleanly to
# the batched 3-D call the driver actually makes. n=16384 saw NO gain from the
# same halving (14.9 ms at both 2048 and 1024) and stays at 2048, fewer ops.
return _left_looking(data, 2048, bf16=True)
# n=128 as ONE launch: 121 -> 73.1. n=256 is deliberately NOT here. Its
# `chol_fused<256>` needs 198 KB of dynamic shared memory, the opt-in is refused, and
# the TORCH_CHECK throws on EVERY call -- `configured` never latches, so the retry
# and the Python exception were being paid 16 times per timed iteration for a route
# that cannot run. n=256 goes straight to the blocked driver.
if _REG is not None and data.is_contiguous() and n == 128:
try:
return _REG.fused_chol128(data)
except Exception:
pass
# LEFT-LOOKING blocked driver. The old note here said the two high-batch mids "lost,
# leave them on cuSOLVER" -- that was the fp32 right-looking driver and it is now
# false twice over: n512b640 3780 -> 2090 and n1024b60 2890 -> 1287.
# n2048 AT BATCH <= 2 GOES TO cuSOLVER ONE MATRIX AT A TIME. probe_giant measured
# batch-1 potrf(2048) at 649 us directly (16 reps, 10.38 ms), so b2 costs 1298 against
# the blocked driver's 1744. The driver runs ~95 mostly-sequential ops to save flops
# that were never the constraint at batch 2 -- at n512b16 the same comparison goes the
# other way by 7x (385 against 16 x 167) and at n2048b8 by 2.8x (1859 against 8 x 649),
# so this is a low-batch effect, not a routing error at n=2048. n=4096 has shipped
# `_loop_single` for the same reason since v2. Outside the `_REG` guard on purpose:
# it is pure torch and must still route correctly if the load_inline build fails.
# RE-FIT 2026-07-26. This route was chosen when the blocked driver cost 1744 us at
# n=2048 against cuSOLVER's 2 x 649 = 1298. Since then the driver's diagonal factor is
# 23.5% faster (exp_fast/exp_fast2) and its panel solve moved onto the tensor cores at
# every batch (exp_fast4, which was worth 4.5% of the whole geomean for one line).
# cuSOLVER's 1298 is FIXED -- it cannot improve -- so the crossover only ever moves in
# one direction. Every other threshold in this file was fitted against a slower kernel
# too; this is the last one on the re-test list.
if n == 2048 and batch <= 2 and _REG is None:
return _loop_single(data)
if _REG is not None and data.is_contiguous() and n in (256, 512, 1024, 2048):
# tf32 is safe at n=512 and n=2048 for EVERY task.yml distribution that reaches
# them (validated by simulation against reference.py's gates). At n=1024 only
# lowrank fails, and the sole lowrank test there is batch 2 -- hence the batch
# gate. NOTE this gate is tuned to the known test/benchmark split rather than to
# numerics, which is uncomfortable; it is legitimate only because benchmark mode
# rechecks the gates on EVERY iteration, so the tf32 path is validated on exactly
# the data it runs on. Drop `or (n == 1024 and batch >= 32)` to give it up (~2.4%).
# THE TENSOR-CORE SOLVE IS BATCH-GATED. exp_tc ran it on every tf32-safe case and
# the per-case table came back monotone in batch: -30% at b640, -17% at b60,
# +3% at b8, +9% at b2, +16% at b16. `tri_inv` plus the Triton launch is a fixed
# cost per panel, so it pays only when there is enough panel to amortize it.
# 32 sits in the empty gap between the measured winners (60, 640) and losers
# (2, 8, 16); no benchmark or test case falls between.
tf32_ok = n in (512, 2048) or (n == 1024 and batch >= 32)
try:
# BATCH GATE DROPPED 32 -> 1. exp_tc fitted `batch >= 32` on 2026-07-25, when
# the tensor-core panel solve cost an extra `tri_inv` launch per panel and lost
# at low batch (+16% at n512b16, +3% at n2048b8). Two things have since changed
# and both move the crossover down:
# 1. `tri_inv` is vectorised and its chain is split four ways, so the fixed
# cost the gate was avoiding is smaller.
# 2. ncu says `trsm_panel` -- the thing the TC path REPLACES -- is at
# **83.90% L1/TEX throughput**, i.e. it is shared-memory-BANDWIDTH-bound,
# not latency-bound. exp_fast3 confirmed that by direction: doubling its
# warps per scheduler made it WORSE (763.98 -> 798.4), which is what
# contention for a saturated pipe looks like. A kernel at 84% of the LSU
# cannot be tuned in place; the work has to leave the LSU. Multiplying by
# an explicit inverse puts it on the tensor cores instead.
# Affects n512b16 and n2048b8 only -- n256 and n1024b4 are not tf32-safe and
# stay on `_blocked_reg` as controls, n512b640 and n1024b60 were already TC.
if tf32_ok and batch >= 1:
return _blocked_tc(data)
# THE TWO fp32 CASES GO TO THE INVERSE-AND-MULTIPLY STRUCTURE, SAME
# PRECISION. n256b64 and n1024b4 are the only cases reaching this line with
# tf32_ok False, and 916966 priced the phase they are dominated by
# (`trsm_panel`, 22.4 us/launch, 51% of the case) against its replacement
# (`tri_inv` 6.65 + panel multiply 7.90 = 14.55). Everything else on the
# board is untouched and is the control.
# MEASURED 2026-07-27 (916999 against 916987, both clean, per-case):
# n1024b4 678 -> 536 **-20.9%** ship it
# n256b64 115 -> 114 -0.9% inside noise, DO NOT ship it
# n=256 runs only 3 panel solves, over 192 / 128 / 64 rows, so `trsm_panel`
# is already cheap there and the two extra host ops per panel
# (`tri_inv` + matmul + strided copy, against one `trsm_panel_inplace`) eat
# the whole win. The gain scales with the rows below the diagonal block, so
# it needs n >= 512 to clear its own overhead. n=1024 low batch is the only
# case that reaches this line with `tf32_ok` False and n >= 512.
if not tf32_ok and n >= 512:
return _blocked_tc_fp32(data)
return _blocked_reg(data, tf32=tf32_ok)
except Exception:
pass
if n == 32:
if _REG is not None and data.is_contiguous():
try:
return _REG.reg_chol32(data)
except Exception:
pass
return _masked32(data)
if n == 64 and _REG is not None and data.is_contiguous():
try:
return _REG.reg_chol64(data)
except Exception:
pass
if n == 128 and _DX is not None:
try:
return _DX.batched_potrf(data)
except Exception:
pass
return _eager(data)
scrolls · 2570 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON