Consecutive-B split-K NAX GEMM
Questi contenuti non sono ancora disponibili nella tua lingua.
iron_steel_gemm_splitk_nax_bt reduces repeated BF16 M5 GEMM GPU time by about 22–24% at M96 and 24–25% at M256, both with N512/K2048. Each primary case wins all 18 paired rounds against the fastest measured dev arm in that round. This explicit kernel consumes weights already stored as [N,K]; it does not change automatic dispatch.
Source commit: b3e3132008beba7620a8ce0e731c65cc751de11b; tree: 8ed7db5cb16fa9410f657ab9c6f70123ef22f245. Baseline: actual dev 74e03a762f96db37403ca4cd3114c8290b9b4c88. Issue: #341. The 119-line kernel body is byte-identical to #299 at 8831565c. All 12 pre-existing functions in its source file are byte-identical to dev. No compiler pass or fused-GEMM sibling is included.
Contract and numerical coverage
Section titled “Contract and numerical coverage”Inputs are row-major a [M,K] and b [N,K] in the selected dtype. Output is F32 partials [splits,M,N], reduced by an existing accumulation kernel. Use exactly 128 threads per group and grid [N/32,M/32,splits]. M, N, K and K-per-split must be positive multiples of 32. Every split starts below K and their total extent covers K; the final split is clamped. Inputs and partial output must not alias. Metal NAX requires Apple10 or newer and Metal 4. BF16 operands stage through F16, so callers must account for F16’s range; these tests use finite operands in that range.
24 new GPU cases cover eight shapes × F32/F16/BF16, including non-square matrices, a short final split, a single tile, production-sized matrices and long K. Every partial slab is checked against independent F64 dot products over the actual dtype-rounded inputs. The first input row is zero and the last has four times the ordinary scale. Outputs start at 7 and have 64 trailing guards initialized to 11. Input allocations have exact extents. All three initial partial-output tolerances remain 1e-4.
The new cases pass on M5 Metal and on GB10 under both memcheck and synccheck, with zero skips or sanitizer errors. A numerical control changes only the candidate’s partial store to zero: it compiles and all 24 cases fail. Restoring the original source makes all 24 pass. The geometry and barriers are unchanged in that control.
Comparison and independent confirmation
Section titled “Comparison and independent confirmation”The benchmark imports production kernels. Both [N,K] and [K,N] weight layouts are prepared once on the host and kept resident. Dev receives its preferred persistent layout at no per-call conversion cost. Thus the result concerns repeated GEMM with prepared weights; it does not establish a one-shot transpose benefit.
The initial screen compares every shape-legal option among existing dense GEMM, fused NAX, seven steel tile configurations, and NAX/two steel split-K configurations at 2/4/8/16 splits. The screen uses 13 dev arms at M96 and 21 at M256. Every split-K arm, including the candidate, uses the same existing elementwise accumulation kernel, and its complete two-pass GPU time is measured. Candidate splits 1/2/4/8/16 are screened separately. All 54 full-chain F64 oracle checks pass in the three-dtype M96 sweep; every timing process checks its complete output too.
Before independent sampling, the primary cases and split counts were fixed: BF16 M96/N512/K2048 at eight candidate splits and M256/N512/K2048 at two. Competitors are dev fused NAX and complete split-NAX chains at 2/4/8/16 splits. Compare to the fastest dev arm, including a fresh fastest-dev selection in every round. Acceptance required at least 5% reduction in both directions and a win in every primary round. The archived plan predates confirmation.
Confirmation contains 38,160 valid samples across 18 processes, with 212 complete-output F64 checks before/after timing. Each process has nine rounds, 20 warmups and 40 samples per arm per round. Arm order rotates for every sample; the second direction reverses arm and case order. Pack 1 measures one whole GEMM; pack 8 repeats eight independent whole GEMMs with the same immutable inputs and divides complete GPU time by eight. The main benchmark tolerances remain F32 1e-4, F16 2e-3, BF16 2e-2.
A local Cargo process overlapped one packed-operation control. Its 2,160 samples and 12 passing oracle checks are preserved but excluded. The entire process was repeated once quiet; all 18 accepted processes have no Cargo/rustc overlap observed at 250 ms intervals. No other process or source was interrupted or modified.
| Case | Splits | Dev → candidate µs, forward | Dev → candidate µs, reverse | Process reduction | Paired-round reduction | Round wins |
|---|---|---|---|---|---|---|
| BF16 96×512×2048, pack 1 | 8 | 15.000 → 11.625 | 67.000 → 50.708 | 22.50% / 24.32% | 22.18% / 23.74% | 18/18 |
| BF16 256×512×2048, pack 1 | 2 | 30.167 → 22.875 | 30.229 → 22.792 | 24.17% / 24.60% | 24.00% / 24.62% | 18/18 |
| BF16 96×512×2048, pack 8 | 8 | 13.625 → 10.714 | 13.646 → 10.568 | 21.37% / 22.56% | 21.63% / 22.79% | 18/18 |
| BF16 256×512×2048, pack 8 | 2 | 28.891 → 21.289 | 28.943 → 21.203 | 26.31% / 26.74% | 26.30% / 26.70% | 18/18 |
| F16 96×512×2048, pack 1 | 8 | 20.250 → 16.083 | 15.000 → 12.563 | 20.58% / 16.25% | 22.51% / 22.49% | 18/18 |
| F32 96×512×2048, pack 1 | 8 | 31.437 → 31.417 | 25.208 → 23.417 | 0.07% / 7.11% | 8.26% / 8.12% | 18/18 |
| BF16 32×512×2048, pack 1 | 16 | 41.333 → 37.625 | 41.292 → 37.500 | 8.97% / 9.18% | 8.79% / 8.87% | 17/18 |
| BF16 96×512×512, pack 1 | 4 | 31.375 → 28.833 | 31.500 → 29.000 | 8.10% / 7.94% | 7.60% / 7.58% | 18/18 |
| BF16 32×32×32, pack 1 | 1 | 3.875 → 4.583 | 14.250 → 16.917 | -18.28% / -18.71% | -18.92% / -19.19% | 0/18 |
The first two rows are the primary result: 36/36 round wins. “Process reduction” compares process medians against that process’s fastest dev arm. “Paired-round reduction” is the median of nine reductions against each round’s fastest dev arm. Absolute latencies vary substantially within and between processes; GPU clock state was not recorded. The relative paired comparisons support the primary claim, not a fixed absolute latency promise. F32’s process result is order-sensitive (0.07% / 7.11%), so no F32 performance promotion is claimed. The skinny control wins 17/18 rounds. Tiny matrices lose about 18–19%; existing fused paths remain preferable there.
The earlier non-split consecutive-B sibling loses every screened round against the best dev option. Its source, oracle results and raw pilot samples are retained under the nax-bt-pilot/nax-bt-fused-held filenames as a held alternative. Those samples are not part of confirmation. No H100 runtime, CUDA performance or whole-model throughput claim is made.
Repository checks and provenance
Section titled “Repository checks and provenance”- Full
cargo test --locked --workspace: 1,244 passed, zero failed, 59 ignored across 134 result groups, including the complete registered GPU harness. - Inventory: 5,497 cases / FNV
0x88ee6aff0fb7d1e0. Removing exactly the 24 additions restores all 5,473 old source/name/dtype/tolerance rows and old FNV0xcc275d4c17fe634f. - Workspace/all-target/all-feature Clippy with
-D warnings, formatting, configured typos, rendered inventory and buffer bindings (662 kernels, zero unbound) pass. - GB10 source audit: all 694 source/dependency files, including the explicit frozen Cargo.lock, match the local validated source, with no extra crate files. The locked build succeeds; CUDA runner
e08b6f7429cc6be08ceb56f05a93cb66e2785382da24c32aec70f913a7f93bd2is unchanged before/after both sanitizer runs. - M5 Max, 128 GiB, macOS 27.0 build 26A5353q; Rust 1.98.0 (
88d9e12ae), LLVM 22.1.8. Build environment: jobs 2, dev debug info 0, incremental off. Frozen lock SHA-256:2427939dd3932e3270c5d4b3119e23b3be2c67442cd6edd6a6b2251f672c8b0c. - Measured executable SHA-256:
d812b1af77f8b2436a5a41a0c0084d6c7cd628656c3b1e18efbc22cd2c8b7078; benchmark source SHA-256:9ecc99832308d16dad256c508e85a448ba8763738c1bd713bc016e8e43d1dd27.
The repository has no make bench target. Its native runner completed the six new partial-pass benchmarks with 20 warmups/40 runs and no observed compilation overlap. These rows are a separate native benchmark check; they exclude accumulation and are not the complete-operation proof above.
| Native partial-pass shape | Dtype | Minimum µs | Mean µs |
|---|---|---|---|
| partial pass m256 n512 k2048 splits2 | f32 | 99.125 | 104.165 |
| partial pass m256 n512 k2048 splits2 | f16 | 19.125 | 24.808 |
| partial pass m256 n512 k2048 splits2 | bf16 | 20.167 | 22.118 |
| partial pass m96 n512 k2048 splits8 | f32 | 20.625 | 21.184 |
| partial pass m96 n512 k2048 splits8 | f16 | 9.167 | 9.465 |
| partial pass m96 n512 k2048 splits8 | bf16 | 9.375 | 10.336 |
Reproduction
Section titled “Reproduction”Check out the source commit, restore Cargo.lock from nax-bt-split-source.tar.gz, and build cargo build --locked -p wh-iron-std --example nax_bt_ab --bin __iron_runner with the recorded environment. The archive also contains the exact 694 source/dependency files used for the remote audit. Run the benchmark from that checkout with GEMM_DTYPE=bf16, GEMM_M=96, GEMM_N=512, GEMM_K=2048, GEMM_ROUNDS=9, GEMM_WARMUP=20, GEMM_SAMPLES=40, GEMM_PACK=1, and GEMM_ARMS=candidate-bt-s8,dev-nax,dev-split-nax-s2,dev-split-nax-s4,dev-split-nax-s8,dev-split-nax-s16. Use GEMM_REVERSE=0 and 1 in separate runs while holding /tmp/iron-gpu.lock and avoiding concurrent compilation. The archived confirmation provenance contains every control’s exact command and environment.
The scripts pin original executable hashes deliberately. For a new machine/build, preserve these published receipts and record a new campaign manifest for the rebuilt executable instead of overwriting the original evidence. SHA256SUMS verifies all compressed receipts and the plain summary CSV; raw samples, invalid overlap, numerical control, frozen lock/source archive, exact commands and all audit results are retained. AI assistance was used for extraction, cleanup, validation and reporting.
