Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
561.4µs
#41 of 337
2026-07-29

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.

mmav21 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 this
vector-width = float4float4 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