Ir al contenido

Wider consecutive-weight NAX GEMM

Esta página aún no está disponible en tu idioma.

The explicit iron_steel_gemm_splitk_nax_bt_direct_32x64 kernel reduces repeated BF16 M5 whole-operation GPU time by 16.76% / 15.64% at M256/N512/K2048 and 22.21% / 11.40% at M256/N1024/K2048, in forward/reverse process order. All 36 primary paired rounds beat the fastest dev arm measured in that round. The wider tile shares A staging across 64 output columns and writes F32 partials directly from cooperative-tile elements.

Source commit bb4ba6d8c84a869c11c9f4efd54208595ccea911, tree 1c14f5a8601734dfd9e72042315343763bd2a75b. Baseline is actual dev ce34672a933a6f88cb56ebbdb83edb77f24d3a02, including the consecutive-weight NAX kernel from #342. Issue #344. The 124-line production body matches #299 at 2628da110d6cb6878ba5e06fb801be91beb29038 exactly. All 25 pre-existing functions in its file, the compiler, runtime and old benchmark example are unchanged. This is an explicit option; automatic selection is unchanged.

Inputs are row-major a [M,K] and b [N,K] in F32/F16/BF16. Output is F32 partials [splits,M,N]; an existing accumulation kernel produces final [M,N] output. Dispatch exactly 128 threads with grid [N/64,M/32,splits]. M, K and K-per-split are positive multiples of 32; N is a positive multiple of 64. Every split begins below K, their total extent covers K, and the final split clamps to K. Allocate exact input extents and separate output storage without aliasing. Metal NAX requires Apple10 or newer and Metal 4. BF16 stages through F16, so callers must account for F16 range; these operands are finite and within that range.

Thirty new partial-output tests cover ten shapes × three dtypes: a single tile, non-square matrices, a short final split, long K, skinny and larger matrices, and both primary benchmark shapes. Independent F64 dot products check every partial slab from the actual dtype-rounded inputs. The first A row is zero; the last has four times the ordinary scale. Partial buffers begin at 7, with 64 trailing guards at 11. All three initial tolerances remain 1e-4.

All 30 pass on M5 Metal and GB10 CUDA under both memcheck and synccheck, with zero skips or sanitizer errors. A numerical negative changes only the candidate’s final partial store to zero: the same legal geometry and barriers compile, and all 30 tests fail. Restoring the source makes all 30 pass. The source, commands, hashes and full failures are archived.

The benchmark imports production kernels and prepares both [N,K] and [K,N] weight layouts once, outside timing. Both remain resident, giving each dev arm its preferred layout without conversion cost. Every split-K arm includes the same existing elementwise accumulation kernel, with both passes measured in one command chain. The results concern repeated GEMM with prepared weights, not one-shot transpose savings.

Exploration compares all geometrically legal dev dense, fused NAX, seven Steel tile configurations, ordinary NAX/two Steel split families at 2/4/8/16 splits, and the #342 consecutive-weight family at 1/2/4/8/16 splits. Candidate splits 1/2/4/8/16 are screened. Fourteen valid exploratory processes contain 12,168 samples and 676 pre/post full-output checks; the separate three-dtype sweep adds 69 full-output checks. One extended-screen process overlapped compilation; its 1,116 samples and 62 passing oracles are retained but excluded and the same case/order was repeated. Skinny M32/N512/K2048 and short-K M96/N512/K512 lose roughly 3–4% and 2–3%, respectively. All slower configurations remain in the raw evidence.

The archived plan fixed both primary BF16 cases and split counts before independent confirmation: M256/N512/K2048 at splits 4, and M256/N1024/K2048 at splits 2. It required at least 5% gain in both process and paired-round medians in both directions, plus every primary round faster. The confirmation dev set is fused NAX, all four ordinary NAX split counts, and all five #342 split counts. The exploratory screen found no dense or Steel arm faster than this set on the primary cases. The tiny control includes dense and all three legal Steel configurations too.

Confirmation has 68,400 valid samples across 18 processes, with 380 full-output F64 checks before/after timing. Nine rounds, 20 warmups and 40 samples per arm per round; arm order rotates each repetition. A separate direction reverses arm and case order. Pack 8 repeats eight whole GEMMs with immutable resident inputs, measuring the complete chain and dividing GPU time by eight. Final-output tolerances stay F32 1e-4, F16 2e-3, BF16 2e-2.

No confirmation process had observed Cargo/rustc overlap; no confirmation process was excluded. Every accepted process has no compilation overlap observed at 250 ms intervals. Other agents were not interrupted. Source and measured executable hashes remain frozen throughout confirmation.

Case M×N×K Splits / pack Dev → candidate µs, forward Dev → candidate µs, reverse Process reduction Paired-round reduction Round wins
BF16 256×512×2048 4 / 1 22.375 → 18.625 22.375 → 18.875 16.76% / 15.64% 16.51% / 15.64% 18/18
BF16 256×1024×2048 2 / 1 179.375 → 139.542 39.125 → 34.667 22.21% / 11.40% 22.23% / 12.19% 18/18
BF16 256×512×2048 4 / 8 20.724 → 16.974 20.755 → 17.344 18.09% / 16.44% 18.13% / 16.55% 18/18
BF16 256×1024×2048 2 / 8 36.719 → 28.797 37.193 → 30.284 21.57% / 18.58% 21.60% / 18.52% 18/18
F16 256×512×2048 4 / 1 21.708 → 17.875 21.458 → 18.000 17.66% / 16.12% 17.73% / 16.44% 18/18
F32 256×512×2048 4 / 1 84.729 → 82.188 97.500 → 101.000 3.00% / -3.59% -3.49% / -3.77% 0/18
BF16 128×512×2048 8 / 1 13.917 → 12.708 14.000 → 12.875 8.68% / 8.04% 8.83% / 8.78% 17/18
BF16 96×512×2048 8 / 1 11.667 → 11.333 50.875 → 47.937 2.86% / 5.77% 3.41% / 5.15% 18/18
BF16 32×64×32 1 / 1 13.208 → 14.937 11.750 → 14.812 -13.09% / -26.06% -16.42% / -18.42% 0/18

The first two rows are the promotion claim. Process reduction compares process medians against the fastest dev process median. Paired-round reduction is the median of nine within-round reductions, selecting the fastest dev separately each round. Absolute latencies shift within and between processes; GPU clock state was not measured. No timing-level shift is used to discard an observation. F32 loses every paired round in both directions despite a misleading +3.00% forward process comparison; no F32 improvement is claimed. The minimum 32×64×32 shape loses 13–26% in process medians and 16–18% in paired-round medians. M128 wins 17/18 rounds. These controls are not a routing rule. The separate 32×32 direct-store candidate in #343 remains held after failing its own predeclared gate. No CUDA performance, H100 runtime or whole-model throughput claim follows from these results.

  • Full cargo test --locked --workspace: 1,244 passed, zero failed, 59 ignored across 134 result groups, including the complete registered GPU harness.
  • Inventory 5,527 / FNV 0x583050e7e9987729 preserves all 5,497 old source/name/dtype/tolerance rows and FNV 0x88ee6aff0fb7d1e0, adding exactly 30 cases.
  • Workspace/all-target/all-feature Clippy with -D warnings, formatting, rendered inventory, configured typos and buffer-binding checks pass. No dependencies were added.
  • All 695 remote source/dependency files match the validated local source, including the explicit frozen Cargo.lock, with no extra crate files. The locked GB10 build succeeds. CUDA runner SHA-256 066a7266c5302091e60ff2c75108d6e02d2ea65ba3fab996d4f4653e493756d9 is unchanged before/after both sanitizers.
  • M5 Max, 128 GiB, macOS 27.0 build 26A5353q; Rust 1.98.0, LLVM 22.1.8. Build jobs 2, dev debug info 0, incremental off. Cargo.lock SHA-256 2427939dd3932e3270c5d4b3119e23b3be2c67442cd6edd6a6b2251f672c8b0c.
  • Confirmation executable SHA-256 62c11d2a40c30c711afc3ff0af468979dafa94bfbe8bdd40c7eb67178d60ae4b; example SHA-256 2a24c98e17494275c1ea514630b243c22127e64a083d75a574d892d8d06dfb19.

The repository has no make bench target. The native runner passes six new benchmarks with 20 warmups/40 runs and no observed compilation overlap. The first native process also passed all six rows but overlapped another local build; its output is preserved as excluded-attempt1 and the entire native process was repeated once quiet. These partial-pass timings exclude accumulation and are separate from the complete-operation proof above.

Partial-pass shape Dtype Minimum µs Mean µs
partial pass m256 n1024 k2048 splits2 f32 189.000 198.132
partial pass m256 n1024 k2048 splits2 f16 121.125 154.178
partial pass m256 n1024 k2048 splits2 bf16 128.083 159.547
partial pass m256 n512 k2048 splits4 f32 48.292 50.114
partial pass m256 n512 k2048 splits4 f16 20.500 24.482
partial pass m256 n512 k2048 splits4 bf16 73.125 101.179

Check out the source commit, restore Cargo.lock from nax-wide-source.tar.gz, and build cargo build --locked -p wh-iron-std --example nax_wide_ab --bin __iron_runner with the recorded environment. The archive contains the exact 695 audited source/dependency files. Run the example with GEMM_DTYPE=bf16, GEMM_M=256, GEMM_N=512, GEMM_K=2048, GEMM_ROUNDS=9, GEMM_WARMUP=20, GEMM_SAMPLES=40, GEMM_PACK=1, and GEMM_ARMS=candidate-wide-s4,dev-nax,dev-split-nax-s2,dev-split-nax-s4,dev-split-nax-s8,dev-split-nax-s16,dev-bt-s1,dev-bt-s2,dev-bt-s4,dev-bt-s8,dev-bt-s16. Run GEMM_REVERSE=0 and 1 separately while holding /tmp/iron-gpu.lock, avoiding concurrent compilation. Exact commands and environment for every control are in confirmation provenance.

Driver scripts deliberately pin the original executable hashes. For a rebuilt binary on another machine, retain the published evidence and record a new campaign manifest. SHA256SUMS verifies compressed receipts and the plain summary CSV. Full raw samples, excluded processes, oracles, numerical control, source archive, plan, commands and validation logs are retained. AI assistance was used for extraction, cleanup, validation and reporting.