Taming Bitwise Behavior in GPU Kernels with Tensor Core
arXiv:2609.11356v1 [cs.DC] 10 Sep 2026
Black-Box Reconstruction, Compiler Enforcement, and Static Verification Ziteng Yang
Nicholas J. Riasanovsky
Georgia Institute of Technology Atlanta, Georgia, USA [email protected]
Meta Menlo Park, California, USA [email protected]
Warren Deng
Vivek Sarkar
Meta Menlo Park, California, USA [email protected]
Georgia Institute of Technology Atlanta, Georgia, USA [email protected]
Abstract Determinism and numerical reproducibility are increasingly required of the GPU kernels under machine learning systems, yet two deterministic implementations of one kernel differ bit for bit. The result is decided by the floating-point (FP) reduction order above all, and by partial-sum precision, multiply-add fusion and where the rounding falls: written by hand, left to the compiler by a block-level language like Triton, or owned by a closed library like cuBLAS or rocBLAS. The tile shape chosen for speed therefore chooses the arithmetic, and a request stops being batch invariant. Holding the order still costs performance, reported at up to 20%. Also, an autotuner searching hundreds of configurations cannot tell which of them agree bit for bit. In this work we characterize what decides the bitwise behavior of GPU kernels, for reductions and General Matrix Multiply (GEMM). i) We give a descriptor that records the parameters fixing a GEMM’s reduction order, such as where the K axis is cut in a split-K GEMM. On that basis we give the first black-box reconstruction of a closed source library’s arithmetic towards bit-level correctness, on NVIDIA cuBLAS, and rebuild it as a bitwise equivalent GEMM family in Triton, a block-level GPU language. It matches cuBLAS bit for bit at a 100% rate on Blackwell and Hopper, and with the epilogue fused from realistic LLM shape, the performance matches or even exceed torch.compile. ii) Lowering to backend, we enforce balanced tree reduction in Triton’s compiler, with a data-layout optimization that brings 19 of 27 kernels on GB300 and H100 within 10% of the free-order mode and 5 even exceed. iii) We build the first sound static equivalence checkers to decide bitwise equivalence between compiled GPU kernels, and the first to partition two vendors’ instruction sets, NVIDIA PTX and AMD GCN with integration into Triton’s autotuner as a static pruning predicate, so a search runs inside a single bitwise-equivalence class.
CCS Concepts: • Software and its engineering;
Keywords: bitwise reproducibility, numerical determinism, floating-point reduction order, compilers, compiler backends, code generation, program equivalence, static equivalence checking, black-box reconstruction, GPU kernels, GEMM, tensor compilers, Triton, autotuning, machine learning systems
1
Introduction
Numerical determinism is now an important need for the systems that train and serve large models: the same computation, the same bits, whenever and wherever it runs. Without determinism during numerical computation, a model may answer differently when asked the same question twice with the same weights and the same seed. A kernel is batch invariant when one request’s output is the same whatever else shares its batch, and the kernels serving these models are not. Serving a request alongside others changes how many rows a kernel reduces at once, which changes the algorithm its partial sums are combined by, giving a different floatingpoint reduction order, while floating-point addition is not associative, so the bits move [24]. Measured across batch size, GPU count and GPU generation on one 7B model in bfloat16, the difference is worth up to nine percentage points of accuracy and generations that differ in length by thousands of tokens [63]. The same instability reaches training. Two runs of one configuration, differing only in effects this small, end at meaningfully different models, and a retrained model changes its mind about individual examples it used to get right [10, 51]. Reinforcement learning is worse than either, because it runs two engines at once: a rollout engine generates, a training engine scores, and when the two disagree about a token’s log-probability the objective being optimized is no longer the one that was written down [48]. Floating-point accumulation order dominates [50]. The precision the partial sums are held at, whether a multiply and an add were contracted into one instruction (FMA fusion),
Yang et al.
and where the rounding to the output type falls all move the result as well. At different infrastructure stack levels, such problems are addressed with heavy effort: i) Compute in higher precision. A fixed-point accumulator of 2098 bits for a sum, 4288 for a dot product, holds every fp64 magnitude, so every addition is exact and one rounding closes it [2, 15, 19]. The Ozaki scheme reaches matrix multiplication the other way, splitting each input into int8 pieces a GPU’s integer matrix unit multiplies exactly and accumulates in int32 [56]. Both cost several times the work. ii) Set library mode. To give callers the property without asking them to change an algorithm, production libraries attach conditions to it. Intel’s oneMKL returns bit-identical results run to run only while the executable, the instruction-set code path and the thread count all stay fixed, and warns that pinning the code path can more than halve its speed; PyTorch’s deterministic mode substitutes deterministic operators, raises an error where none exists, and scopes its guarantee to one release on one platform [26, 47]. The price is the conditions themselves, which the caller must hold. iii) Rewrite the GPU kernels. Nondeterministic LLM inference has been traced to kernels whose reduction order moves with batch size, and rewriting them to fix the order costs about 20% of matrix-multiply throughput against cuBLAS [24]; RepDL enforces correct rounding and order invariance to keep training and inference bit-identical across machines [60]; LayerCast holds weights in sixteen bits while performing every computation in fp32, paying in bandwidth rather than in kernel structure [63]. DeepSeek-V4 [18] makes a frontier training stack bitwise batch-invariant and deterministic end to end. a) cuBLAS picks its kernel from the problem dimensions behind a closed cost model, so the arithmetic moves with the batch; a matrixmultiply library of their own replaces it, and split-K, taken only at small batch sizes and so batch-dependent in the same way, is dropped wherever a kernel can do without it. b) Atomic adds land in thread arrival order, so the backward pass gives each streaming multiprocessor its own buffer and sums those in a fixed order. c) Its compiler turns fast-math off by default, aligns its lowering with the reference CUDA toolchain so no transformation moves a bit, and calls an SMT solver for the integer facts its layout inference rests on. Their matrix multiplication is reported to match or surpass standard split-K in most major scenarios. But none of them reproduces the discipline a closed library follows. iv) Align two engines. A reinforcement-learning rollout engine and training engine that disagree on a token’s log-probability shift the objective being optimized; one answer moves the whole pipeline from bf16 to fp16, trading dynamic range for the rounding headroom that keeps the two consistent [48]. Holding the bits still is a challenge at different layers: In the training algorithm, a policy update weighs each token by the ratio of two probabilities, so when the engine that generated it and the engine that scores it disagree the
ratio is taken between two different models rather than between two policies of one [48]. In the ML system infrastructure, a stack is assembled from independently built parts, and one that accumulates differently breaks the agreement for all of them. Below the kernel, the vendor’s assembler may reassociate an accumulation it judges safe to move, and where no machine code was shipped the driver assembles at load time, so a driver update can change bit-level semantics the program did not. In the GPU kernel (scope of this work). Various obstacles exist to reaching numerical consistency: i) A vendor library will not fully disclose what arithmetic it performed, so a kernel written elsewhere has difficulty matching it. ii) Restricting the computation order reduces the flexibility of optimization towards parallelism. PyTorch’s deterministic mode stops tuning a reduction kernel’s configurations, and when its full search is enabled it takes the vendor’s matrix multiply instead of a generated one [47]. Where agreement cannot be established, speed is what gets given up. iii) One source compiled twice by one compiler under different parameters may or may not agree, and an autotuner [13, 64] searching hundreds of configurations for speed has no way to do equivalence partition. This paper makes four contributions. A mechanistic account of bitwise behavior. We show that FP reduction order is the dominant factor that decides the numeric bits. We give a theoretical characterization of the factors deciding the FP reduction order of reduction kernels and General Matrix Multiply (GEMM) kernels. We give a formal descriptor for GEMM kernels, GEMMDesc, that describes the full bitwise behavior of the most commonly used GEMM variants. It records the parameters that decide the reduction order, for example where to cut the K axis in a split-K GEMM, and holds nothing that is only a speed knob. Two machines handed the same descriptor owe each other the same bits. The first Triton GEMM family that is bit-identical to a closedsource vendor library. NVIDIA’s cuBLAS [42] is a black-box GPU kernel library that only provides incomplete computation information that decides the FP computation order. We applied the mechanism above, conducted a set of offline experiments (runtime profiling, numerical inference, and others), and give the first black-box reconstruction of such a library towards bit-level correctness: a Triton GEMM family over fp16, bf16 and fp8 e4m3 that reaches 100% bit-matches with cuBLAS GEMM on GB300, GB200 and H100. As an extra bonus, we discovered a logical bug in cuBLAS GEMM during that offline work (Appendix D).1 First Balanced Tree reduction in and layout optimization in Triton compiler We implement compiler enforcement for balanced tree reduction (“inner tree” mode) in the compiler backend of Triton [54] (none before), the block-level GPU 1 The 100% rate excludes this bug.
Taming Bitwise Behavior in GPU Kernels with Tensor Core
kernel language. We analyze and implement the data-layout optimization opportunity that the enforcement opens, bringing 19 of 27 kernels across GB300 and H100 within 10% of the free-order mode, and even past it on five of the ten on H100. The first Static & sound equivalence checkers for compiled Triton kernels and integration in autotuner. We implemented the first static checker that decides the bitwise equivalence of two NVIDIA PTX kernel assembly, sound by construction, and the first symmetric one for a second vendor’s instruction set, AMD GCN. It acts both as a static equivalence partitioner during the autotuning stage of a GPU kernel and as a verification of the mechanism we claim. We incorparate it into Triton’s autotuner and achieved static pruning before the actual pruning happens for the first time. The checker is validated to be robustly sound across a large kernel suite, including fused kernels generated by PyTorch Inductor [8] and Flash Attention kernels [17]. It recovers the exact partition on the GEMM family and on Inductor’s fused kernels, and comes within a factor 2.5 on reductions and normalizations.
2
Background
2.1
GPU kernels and the programming model
A GPU kernel is a function the host launches as a grid of thread blocks [6, 39, 43]; every thread runs the same body and reads an index to find its slice of the data. i) A block is at most 1024 threads and goes whole onto one streaming multiprocessor (SM) on NVIDIA, one compute unit (CU) on AMD [33], staying there until it finishes; an SM holds the tensor cores (AMD’s matrix cores) that execute matrix-multiply instructions, and up to 32 blocks share one when registers and scratchpad allow [44]. ii) Within a block, threads are cut into fixed groups, a warp of 32 on NVIDIA and a wavefront of 64 on AMD’s CDNA [5], which is the unit of instruction issue; lanes on a different branch sit an instruction out, and since Volta [40] each lane carries its own program counter, so a value passed between lanes through memory needs an explicit __syncwarp(). iii) Above the grid, ordering across launches comes from a kernel ending, and combining results across GPUs is a collective library’s job, NCCL or RCCL [52]. Registers are private to a thread, at most 255 of an SM’s 64K. Shared memory (AMD’s local data share) is a scratchpad, fast on-chip memory the program fills and empties itself, carved per resident block from the SM’s 228 KB [44]. Global memory is off-chip DRAM at roughly twenty times a sharedmemory read. Only block scope costs a barrier: __sync threads() publishes a block’s shared-memory writes, a shuffle between lanes is free, and grid-wide ordering comes from ending the kernel. Thread-level programming and the vendor libraries. CUDA C++ asks for the body of one thread and runs it across the grid, leaving cooperation to be written by hand:
which thread reads which element, when a tile is staged into shared memory, where the barriers go. A caller can reach for a closed-source vendor library instead, cuBLAS [42], cuDNN [14] on NVIDIA or rocBLAS [7] on AMD, and then chooses the result while the library chooses the arithmetic that produces it. Block-level programming and the compilation pipeline. Block-level programming asks instead for a block of threads over a tile of data and leaves the thread detail to the compiler. Triton [54], TileLang [58] and NVIDIA’s cuTile [41] implement it on GPUs, as do Pallas [27] on Google’s TPUs [28] and the Neuron Kernel Interface [4] on AWS’s training and inference chips [3]. Tile operations carry no hardware in them; the tile-level IR fixes a layout per tensor, saying algebraically which lane and register hold which element [65]; an LLVM backend lowers that per thread; and GPU assembly, PTX or AMDGCN, is the last textual form before machine code. The tile-level IR is where a reduction’s order first becomes visible, and the assembly is the lowest level we can read. Autotuning. Tile sizes, the pipeline depth over which a loop’s loads run ahead of its arithmetic, and warp count can be left as template parameters (knobs). A kernel written this way is a kernel template, one assignment to its knobs a configuration, compiling under one a kernel instance, and the legal assignments its configuration space, thousands for a GEMM; they are the configurations fanning out from one template. An autotuner picks among them by measurement at run time, because the winner turns on occupancy, on whether the tile shape matches the tensor core’s, and on register allocation inside ptxas, none of it readable off the source [13, 64]. Autotuners have become routine recently: PyTorch’s compiler generates the Triton kernels with autotuning decorators and autotunes them [8, 54]. Triton spells this as an annotation on the kernel: @triton.autotune(configs=[ triton.Config({'BLOCK_M': 128, 'BLOCK_K': 64}, num_warps=8, ...), triton.Config({'BLOCK_M': 64, 'BLOCK_K': 32}, num_warps=4, ...)], ..., key=['M', 'N', 'K']) @triton.jit def matmul(a, b, c, M, N, K, BLOCK_M, BLOCK_K): ...
The decorator compiles and times every listed configuration and caches the fastest; new shapes restart the measurement. Two instances of one template compute the same mathematics, but not necessarily the same bits.
2.2
Related work on compiler correctness
Correct by construction. A verified compiler carries a machine-checked proof that every program it compiles keeps its source semantics. CompCert is the first realistic instance for C [32], and its floating-point semantics are verified in Coq
Yang et al.
bits X
+ + a1
+
+
+ a0
layout A — 4 threads × 4 elements, 2 warps + bits X
a2
+
+ a3
a4
a5
a6
a7
+ +
+
+
+ + +
+
+
+ +
+
+
+
+ +
+
+
+
+
+
+ +
+
+
+
+ + +
+
+ + + + bits Y ≠ X + the same bits X layout B — 2 threads × 4 elements, 1 warp layout B — 8 threads × 2 elements, 1 warp within a thread (registers) within a warp (lane shuffle) across warps (shared memory) (a) Order free: two layouts, two trees
before
(b) Order pinned: two layouts, equivalent tree
0 0 2 2 0 0 2 2
0 0 2 2 0 0 2 2
after 1 1 3 3 1 1 3 3
1 1 3 3 1 1 3 3
warp 0 reduce
+
reduce
layout A — 4 threads × 2 elements, 2 warps
warp 1
0 1 2 3 0 1 2 3
0 1 2 3 0 1 2 3
0 1 2 3 0 1 2 3
0 1 2 3 0 1 2 3
warp 0 warp 1
(c) The layout optimization
Figure 1. Pinning the reduction order, and what it buys. (now Rocq [53]), so IEEE-754 arithmetic survives the transformations a compiler is otherwise tempted to make [12]. The same has since been done for instruction scheduling [62], tensor-language lowering [34], a GPU memory consistency model [36] and the semantics of GPU assembly [21]. The price is the proof effort, and a toolchain of its own to run under. Translation validation. It drops the claim about the compiler and checks the one compilation in front of it [45]: encode both IRs and ask a solver whether the output refines the input. Alive2 does this for LLVM IR [35] and MLIR-TV for the multi-level IR deep-learning compilers are built on [9]. Alive-FP brings floating point in, verifying LLVM’s floatingpoint and fast-math peepholes under one SMT encoding per reading of the under-specification around signed zeros, NaNs and infinities [37]. The price is what the encoding admits: MLIR-TV over-approximates floating-point arithmetic and reductions so that a solver finish in bounded time. Fuzzing. Fuzzing proves nothing and locates real incorrectness. Csmith generates random C free of undefined behavior and differentially tests compilers against each other [61], and MLIRSmith carries the idea to the multilevel IR [57]. Floating point moves the target from one wrong compilation to two that disagree [31, 49]. The price is every test-based method’s: finding nothing is not showing nothing is there. None of them decides whether two compiled GPU kernels return identical bits, let alone autotuner integration.
3
Understand the bit behavior under GPU parallelism
3.1
Dependence Tree of reduction
An addition of data through a fp32 accumulator it runs is written ⊕fp32 . Addition commutes, but does not associate: (𝑎 ⊕fp32 𝑏) ⊕fp32 𝑐 ≠ 𝑎 ⊕fp32 (𝑏 ⊕fp32 𝑐). A reduction of 𝑛 values on a GPU is a directed acyclic graph (DAG). Its nodes are of two kinds: a data node holds one floating-point value,
and an operation node is one floating-point addition. An edge runs from a node to the operation that consumes its value, so every edge is a data dependence and points the way the value travels. The graph is a tree in most cases for one reduction result (or can be expanded into a tree), and we call it the dependence tree of the reduction. The reduction order therefore is represented by a dependence tree: Exchanging an operation’s two incoming edges result a different tree holding the same relation, and we call them equivalent trees. A kernel may induces more than one result, and we write T (𝐾) for the family of their dependence tree. Two kernels over the same inputs return the same bits when the trees at every output coordinate are equivalent. A dependence tree also says how much of the reduction can run at once. Its work, the count of operation nodes, is 𝑛 − 1 whatever shape it takes, so a tree’s parallelism, its work divided by its longest dependence chain, is settled by that chain alone. Across the trees a compiler may build over the same values the chain runs from 𝑛 − 1, a left fold whose parallelism is 1, down to ⌈log2 𝑛⌉ for a balanced tree. The order Section 5 pins is the balanced one, so enforcing it gives up no parallelism chance in theory [11, 16, 30] and still preserves data locality. A GPU builds that tree in three levels. A thread folds the nodes held in its own registers, in index order. The threads of a warp fold their partials through lane exchanges. The warps holding a partial fold theirs through shared memory. The tree is settled only once all three are fixed, T = Tsmem ◦Twarp ◦Treg . Thread count, the elements one thread holds and the warp count each decide which nodes share an accumulator at some level, so each of them moves the tree: Figure 1(a) folds the same eight values under two layouts and the two roots differ. Figure 1(b) folds sixteen values under two layouts, the second over twice the threads of the first, and the two roots agree. Section 5 is how it is pinned by compiler backend. Pipeline depth leaves the tree alone. A softwarepipelined loop issues the loads for later iterations while an
Taming Bitwise Behavior in GPU Kernels with Tensor Core
one accumulate step
one accumulator, closed by one rounding
BLOCK_K instruction_k
k
span
an accumulator starts here
k_cuts is a descriptor field
k_cuts[0].span 0 1 2 3 0 1 2 3 0 1 2 3 0 1 2 3
0
one accumulator, never reopened round to output_dtype summed in index order
(a) Plain GEMM
(b) Split-K one thread block
the inner cut divides one outer part
1
2
3
layout: STRIDED, count = 4 lane j takes every fourth step
(c) Chained chunks
(d) GEMV
a second kernel launch
Figure 2. Where each GEMM algorithm family cuts the contracted dimension. earlier one computes, and the pipeline depth, num_stages in Triton, is how many iterations are in flight at once. What it sets is arrival time. Iteration 𝑖 still accumulates iteration 𝑖’s values into the same accumulator and in the same order, so every node reaches the accumulator it reached before and the tree is what it was. The premise is that pipelining moves the loads and leaves the loop’s accumulation chain where it stands. Warp specialization leaves it alone. A warp-specialized kernel splits the block’s warps into a producer set 𝑃, which issues asynchronous copies into shared memory, and a consumer set 𝐶, which does the arithmetic. Three facts settle the question. A node of the tree is a floating-point addition and a producer warp executes none, so 𝑃 contributes no node. Treg is fixed by the map from a node to the lane holding it, and that map is computed from the tile shape and the layout, which the split leaves untouched. Specialization adds the producers alongside the consumers rather than taking consumers away, and Tsmem spans the same partials. All three levels are unchanged, so T is, and so are the bits. The third fact is a premise about the compiler rather than about the hardware: a compiler that re-partitioned a tile over fewer compute warps when it specialized them would break the claim. 3.2
GEMM reduction
3.2.1 Tensor core semantics. A matrix instruction takes a tile of 𝐴 and a tile of 𝐵, multiplies them, and adds the products into an accumulator that already holds a value. Where the roundings fall inside that is a black box. For fp16, bf16 and tf32 inputs this section assumes the semantics below, and every measurement in Section 7 agrees with it: one instruction folds instruction_k products and the incoming accumula Í⊕fp32 tor into a single rounding: acc ← acc 𝑗 <instruction_k 𝑎 𝑗𝑏 𝑗 , where the sum is exact, so the instruction’s one rounding is the ⊕fp32 itself. The products inside one instruction therefore have no order of their own, and what survives as a reduction order is the chain of ⊕fp32 steps across instructions, one step per instruction. The scalar path rounds once per element,
acc ← fma(𝑎, 𝑏, acc), and twice when the multiply is left uncontracted. Three instructions carry the matrix path on NVIDIA: mma.sync in a warp, wgmma in a Hopper warpgroup, and tcgen05.mma from a single Blackwell thread with its accumulator in tensor memory. On GB300 mma.sync and tcgen05.mma are bitwise equivalent at fp16 and at bf16,2 so the descriptor names the instructions one by one rather than carrying a single tensor-core value. 3.2.2 The algorithm families. What separates one GEMM from another is where the contracted axis is cut and what happens at each cut. Four GEMM algorithm cover every library kernel that is used, and Figure 2 draws them. Plain GEMM hands one output tile to one threadblock, gives every element of that tile an accumulator, and walks the whole contracted axis inside that block, the uncut axis of Figure 2(a). Three things about that walk decide the order. A matrix instruction does not add one product at a time: it folds a fixed number of them into a single rounding, and that number sets how long the chain of roundings is, so we record it as instruction_k. The instruction’s own result then has to reach the accumulator, and it can do so inside that same rounding or in a second one after it, which is one rounding per step against two; we record which as use_fast_accum. And 𝐾 is rarely a multiple of the mainloop’s step, so one turn of the loop is short, and whether that short turn runs first or last changes which products share the first rounding and therefore every partial sum after it. We record it as k_loop_step. What tiles the output stays out of all three. BLOCK_M and BLOCK_N decide which thread owns which output element, and every element has an accumulator of its own either way, so re-tiling moves work between threads and leaves each accumulator’s chain where it was. BLOCK_K sets how many instructions one turn of the mainloop issues, and the accumulator is carried across turns, so by the semantics of 2 At fp8 the two are not bitwise equivalent, because fp8 accumulation runs
at reduced precision with a promotion cadence that need not match across two lowerings [29, 59].
Yang et al.
Section 3.2.1 the same instruction results reach it in the same order however the turns are cut. The three block sizes are speed knobs, which is why instruction_k and not BLOCK_K is the number the order depends on. Split-K [1] cuts the axis into contiguous parts and gives each part to a different threadblock, which starts it at zero in an accumulator of its own, the four parts and their merge in Figure 2(b). We record the length of a part as span. Some kernels cut again inside a part, closing an accumulator every so many elements before the part is finished, so what records the cutting is a nest, outermost first, which we write as k_cuts. A second kernel sums the parts. A finished part is written at one type and the running sum kept at another, and the two are set independently: parts at fp32 summed at fp32, parts at the output type summed at fp32, or both at the output type. We record them as partial_dtype and merge_dtype. Chained chunks cuts twice and stays inside one threadblock, the two levels of Figure 2(c). Several threads share one output element, and the outer cut gives each of them a contiguous chunk of the axis; the inner cut divides a chunk into the sub-blocks one accumulator closes over. Both lengths are the span of their level of the nest. This is what a kernel does on the scalar path, where no matrix instruction folds a group of products for it and the accumulation is an explicit chain of fused multiply-adds. GEMV [20] multiplies a matrix by a vector. One output element per row leaves no output tile to spread over threads, so the fold itself is spread instead, and that makes its order the most exposed of the four. The lanes of one warp each take every countth tile of the axis rather than a contiguous slice of it, the deal Figure 2(d) draws, and which of the two a cut does is recorded as layout. The lanes then fold their totals into one through shuffles, and that fold is where the order stops being a chain: a butterfly pairs lanes a power of two apart, which is a balanced tree, and a butterfly that counts its offsets down rather than up is a balanced tree over the lanes in bit-reversed order. Those are different sums, and only the second is what these kernels do, so we record the fold as one two-part cut per round rather than as a single tree over the lanes. We formalise that whole set of factors as the GEMM computation descriptor, GEMMDesc for short, and Section 4 is where it is put to work.
3.3
Attention
Readers may refer to Appendix F for the reduction order of a complex kernel such as flash attention, where this work offers a theoretical description and leaves the implementation open.
4
Black-box reconstruction
Although cuBLAS3 exposes a cost model API that names the GEMM algorithm it would run for a shape, that answer stops short of the parameters of Section 3.2.2 that fix the bitwise semantics. A GEMMDesc is the frozen record holding every one of those factors and nothing else, so two equal records owe each other the same bits. What we recover is not one of them. For a fixed library generation, cuBLAS 12 or cuBLAS 13, we recover the function that generation computes, (SM version, tensor shape) ↦−→ GEMMDesc, and a single shape’s descriptor is one point of it. Building that function is this section’s subject: how one point is recovered, and how the points become a function. Two instruments. Runtime profiling says what the library actually did: how many kernels it launched, over what grid, and whether it took a workspace. That narrows a shape to a family cheaply. The arithmetic is settled by numerical experiment: two extreme values of opposite sign and one extreme small value, +𝐿, −𝐿 and 𝑟 , are placed on an otherwise empty axis. For fp16, 𝐿 = 1024 and 𝑟 = 2−15 , which is under half an ulp of fp32 at 𝐿. The pair cancels exactly, so 𝑟 reaches the output from an accumulator that has already cancelled it and is swallowed by one still holding half of it. A tensor core folds its own 𝑘 before the accumulator sees anything, so 𝑟 goes one instruction group away from its 𝐿 rather than beside it. Two placements read the two shapes a grouping takes. i) Contiguous grouping. Eight terms summed as (𝑎 0 + 𝑎 1 + 𝑎 2 + 𝑎 3 ) + (𝑎 4 + 𝑎 5 + 𝑎 6 + 𝑎 7 ) answer differently from the same eight in one chain, and what has to be found is where a group ends. Put 𝑟 at index 0 and walk +𝐿 and −𝐿 as an adjacent pair, at (1, 2), then (2, 3), and on. While the pair sits inside 𝑟 ’s own group the +𝐿 arrives after 𝑟 and absorbs it, and the output is zero. At (3, 4) the +𝐿 is still in that group and absorbs 𝑟 there, so the output is zero once more. At (4, 5) the pair lies wholly in the second group, which cancels to zero on its own, and 𝑟 reaches the output intact. The first placement that returns 𝑟 is 4, so the first group is 𝑎 0 through 𝑎 3 ; pinning 𝑟 at 4 and walking again gives the next boundary. ii) Strided grouping. The same eight summed as (𝑎 0 + 𝑎 2 + 𝑎 4 + 𝑎 6 ) + (𝑎 1 + 𝑎 3 + 𝑎 5 + 𝑎 7 ) hold every other term in one group, and what has to be found is which terms a group holds. Here the pair is pinned and 𝑟 walks. Under the layout being tested, +𝐿 goes at a group’s first term and −𝐿 at its last, 0 and 6 here, so that group carries an uncancelled 𝐿 throughout, and 𝑟 takes each remaining index in turn. It is absorbed at 2 and 4 and survives at 1, 3, 5 and 7, so the group holds {0, 2, 4, 6}: stride two, two groups. A layout guessed wrong puts the −𝐿 outside the group its +𝐿 is in, and the pattern that comes back is no longer periodic.
3 CUDA 13.3 is the latest release at the completion of this work.
Taming Bitwise Behavior in GPU Kernels with Tensor Core
Taking split-K as the example. Three things that fix split-K’s bits are missing from that answer: k_cuts, how the axis is cut, span, how long a part is at each level of that nest, and the pair partial_dtype and merge_dtype. Reading span takes one row of an fp16 GEMM on GB300. An otherwise empty row of a 𝐾 = 576 GEMM carries 𝑟 at 𝑘 = 0 and the pair +𝐿, −𝐿 at 𝑘 = 𝑚 and 𝑚 + 1, walked along the axis. While the pair shares 𝑟 ’s part the +𝐿 arrives after 𝑟 and absorbs it, so the output is zero from 𝑚 = 1 to 191. From 192 the pair lies wholly in a later part, which cancels to zero on its own, and 𝑟 reaches the output. It falls back to zero at 𝑚 = 383 alone, where the pair straddles the next boundary, so its halves reach the merge apart and the +𝐿 absorbs 𝑟 there. One walk gives both boundaries, and the axis is in three parts of 192. Recovering one point runs the two instruments in order. A profile of the shape narrows it to a family and drops the candidates that disagree with what was launched. Each grouping still open is then read off the axis with +𝐿, −𝐿 and 𝑟 . What is left over is separated by running each candidate’s own walk against the library on fresh draws until one survives. From points to a function. A point at a time would never finish: the shapes are unbounded and the cost model is free to answer differently at each of them. What makes the function finite is that the descriptor depends on the shape only through the cost model’s answer for it, and those answers are finite in number. Recovering the function is therefore recovering one descriptor per reachable answer, which is a table, and we call one such table an arch profile. Following this principle we reconstructed the tables for fp16, bf16 and fp8 e4m3 on GB300, GB200 and H100 against cuBLAS 12 and 13, one per SM version and library generation. The answers a table has to key on are enumerated rather than sampled: a scan of 8.25 million cost-model queries over vector lengths to 106 found every answer the vector family is reached with, seven of them only at very long vectors or very deep 𝐾. On GB300, GB200, H100, excluding a located cuBLAS bug, our reconstruction using Triton (block-level programming) returns cuBLAS’s own bytes on every shape whose kernel computes the full sum, and holds at 100%. A non-cuBLAS GEMM taken alone falls short of highly optimized cuBLAS whether or not we require bitwise equivalent: above 5 GFLOP of arithmetic the Triton GEMM torch. compile generates (FP order free) runs at 61–88% of it over the layer groups, while our bit-exact Triton kernel runs at 56–93%, and a large-tensor acceleration (Appendix E) at 69– 91%. Below 5 GFLOP the free one runs at 59–88% while ours runs at 50–100%, which is where holding the order costs. Kernel fusion turns that around on realistic shapes: over 96 of them the fused pair torch.compile generates runs at 85–125% of the speed of a cuBLAS GEMM followed by its epilogue as a second kernel over the epilogue groups, while our bit-exact fused pair runs at 95–168%. Fixing the reduction order is usually expected to cost about 20% of matrix-multiply
throughput [24], but these results open the opposite prospect, that bitwise consistency and performance can be pursued together. Section 7 carries the complete evaluation result.
5
Compiler enforcement and layout optimization
Triton has no balanced tree reduction. Its lowering folds the values a thread holds one at a time, left to right, and the layout decides which values those are, so the tree a kernel computes moves with every layout the autotuner tries. We extend Triton’s compiler with one, and with a data layout optimization pass over it, at an acceptable performance cost against the unordered mode under autotuning. 5.1
Backend lowering
The lane exchange is where the tree is built out of shuffles, and Figure 3 is the function that emits them. Both modes walk the same sequence of shuffle-xor steps and differ in direction alone: the pinned mode counts the offset up from one, so neighbouring lanes pair first and the tree has the same shape whatever the warp count, while the default counts down from half the lane count. A switch of this shape sits at every level of Section 3.1’s hierarchy, so the within-thread fold and the cross-warp stage are pinned the same way. void warpReduce(ConversionPatternRewriter &rw, Location loc, SmallVector<Value> &acc, ReduceOp op, unsigned nLane, ...) { if (targetInfo.warpReduce(rw, loc, acc, op, nLane, ...)) return; if (isInnerTree(op)) { // count up: neighbours pair first for (unsigned N = 1; N <= nLane / 2; N <<= 1) { SmallVector<Value> shfl(acc.size()); for (unsigned i = 0; i < acc.size(); ++i) shfl[i] = targetInfo.shuffleXor( rw, loc, acc[i], N * interleave); accumulate(loc, rw, op.getCombineOp(), acc, shfl, pred); ...
Figure 3. The two orders, in the lane exchange. The dispatch above it reaches the pinned path by two routes. The attribute is one. The other is a layout whose reduced register-lane extent on the axis is still greater than one, which the default lowering cannot express, so it is sent down the same path and gets the same tree. AMD reaches the same tree on different instructions. A wave64 reduction folds within a row with row_shr steps, and the two orders are the two directions through the same step sequence: 8, 4, 2, 1 counting down for the default, and 1, 2, 4, 8 counting up for the pinned mode. The cross-row broadcast steps that follow preserve the tree and are identical either way. Where the assumption behind that sequence fails, on a wave32 target or a partial warp, the backend declines the instruction-level path and the shared count-up shuffle tree performs the reduction instead, which is the same tree reached more slowly.
Yang et al. A: 4 threads × 4 elements, 2 warps
B: 8 threads × 2 elements, 1 warp
mov.u32 %r1, %tid.x mul.wide.u32 %rd1, %r1, 16 add.s64 %rd2, %rd0, %rd1 // %rd2 = %rd0 + 16t, t = 0..3 ld.global.f32 %f1, [%rd2] ld.global.f32 %f2, [%rd2+4] add.f32 %f3, %f1, %f2 ... 4 elements, 2 levels shfl.sync.bfly.b32 %f4, %f3, 1, 31, -1 add.f32 %f5, %f3, %f4 st.shared.f32 [%rd6], %f5 bar.sync 0 ld.shared.f32 %f7, [%rd7] add.f32 %f8, %f5, %f7 setp.eq.s32 %p1, %r1, 0 @%p1 st.global.f32 [%rd4], %f8
mov.u32 %r1, %tid.x and.b32 %r2, %r1, 7 shl.b32 %r3, %r2, 3 cvt.u64.u32 %rd1, %r3 add.s64 %rd2, %rd0, %rd1 // %rd2 = %rd0 + 8t, t = 0..7 ld.global.f32 %f1, [%rd2] ld.global.f32 %f2, [%rd2+4] add.f32 %f3, %f2, %f1 shfl.sync.bfly.b32 %f4, %f3, 1, 31, -1 add.f32 %f5, %f3, %f4 ... offsets 2, 4 setp.eq.s32 %p1, %r1, 0 @%p1 st.global.f32 [%rd4], %f5
one signature: add.f32, height 4, first offset 1
Figure 4. Two layouts of one reduction, one signature. 5.2
The data-layout optimization
A reduction whose axis is spread across warps pays a crosswarp stage: two barriers, a shared-memory round trip and a second shuffle sequence. Moving the operand to a layout that puts the axis inside one warp escapes that stage, and it normally moves the tree along with it, so the escape costs bitwise equivalence. A pinned tree makes the move safe. The lowering above fixes the association independently of where the elements live, so for a reduction whose tree is pinned the entire space of valid layouts is a free performance knob. Algorithm 1, in Appendix B, is the pass that spends it. It builds the layout that would reduce the axis inside one warp, puts the warps on the kept dimensions so the axis stays warp-synchronous, carries whatever axis extent is left in registers as a within-thread fold, and then asks whether that is worth the convert_layout it has to insert. The op, its axis, its extent and its ordering are never touched, so the result is bit-identical by construction. Figure 1(c) is that rewrite on one operand. Before it a lane holds a two-by-two block, so the eight-element axis runs past the warp boundary; after it a lane holds a row of the kept dimension and each warp owns whole reductions. Loading in the reduce-friendly layout would reach the same place with no conversion at all, and the pass rejects it: a strided load is paid on every launch and the conversion once per tile. On GB300 the optimization recovers most of what the constraint costs, and on H100 it carries six of ten kernels to or past the unordered mode (see Section 7).
6 Static equivalence checking and partition The question the checker answers is narrow: given two compiled kernels, are they bitwise equivalent. It answers from the assembly statically. What it reads is PTX for NVIDIA and AMDGCN for AMD. 6.1
The algorithm
The walk takes one entry function in program order over a symbolic thread, one thread whose index stays a symbol,
carrying a map from register to the node standing for that register’s contents. Each instruction is a transfer function on that map: it looks its operands up rather than tracing them backwards, and writes a node into its destination. A global load becomes a data node carrying the address the load read. A floating-point combine becomes an arithmetic node over the nodes its operand registers already hold. A lane shuffle whose result is combined with the partial that was shuffled becomes one within-warp exchange node, and a store to shared memory, followed by a barrier, followed by a load, becomes one cross-warp exchange node. Figure 4 colours two assemblies by scope, thread green, warp yellow and shared memory pink, as Figure 1 colours its nodes. A loop that accumulates into a carried register becomes a fold over what one iteration contributes. A predicate is dropped when its truth is a function of the thread and block coordinates, since which threads run is a launch fact; a predicate carrying a loaded value is data, and the entry it guards is compared configuration by configuration instead. The nodes reaching a global store are the roots, one per output element, and each carries the dependence tree of that element. Addresses. Address registers are evaluated symbolically. Over a basis S of the integer symbols a kernel can read (the thread and block coordinates, the block extents, the parameters and theÍbase pointers), every integer register lands in A = { 𝑐 0 + 𝑠 ∈ S 𝑐𝑠 𝑠 } ∪ { ⊤𝑒 }: an affine form, or an opaque token ⊤𝑒 carrying the expression that produced it. Moves and address-space casts are the identity, addition adds termwise, and a multiply or a left shift by a literal scales every coefficient. A bitwise operation is admitted only where a rule proves it exact. Writing tz(𝑎) for the trailing zeros common to 𝑐 0 and every 𝑐𝑠 , a right shift by 𝑘 is exact when every termÍis a multiple of 2𝑘 , so tz(𝑎) ≥ 𝑘 gives 𝑎 ≫ 𝑘 = (𝑐 0 ≫ 𝑘) + 𝑠 (𝑐𝑠 ≫ 𝑘) 𝑠; a mask is the identity when it covers every bit 𝑎 can set, and an or over two values that share no set bit is an addition. What is left is a mask or a shift on a thread index, and the launch bound makes that Í exact too: %ntid.𝑑 = 𝑁𝑑 is a power of two, so %tid.𝑑 = 𝑖<log2 𝑁𝑑 2𝑖 %tid.𝑑.bit𝑖 covers exactly [0, 𝑁𝑑 ), and in that basis a mask keeps the bits it selects while a shift renumbers them. A leaf’s identity is the affine form of the address it reads, so two kernels reaching one element by different index arithmetic give one leaf. An opaque token equals only an identical token, which splits a class rather than merging two. Canonicalization sorts the children of every commutative node, so a tree is compared up to exactly the commutativity Section 3.1 allows. A bottom-up hash then reduces the tree to one signature, and two kernels are bitwise equivalent when their signatures agree. A balanced tree reduction collapses one step further. Its shape is fixed by the logical element order (Section 3.1), so the collapsed node keeps four facts and drops the physical structure that produced them: it keeps the combine with its
Taming Bitwise Behavior in GPU Kernels with Tensor Core
rounding, the leaf computation with its coordinate blanked, the reduction height, which is the base-two logarithm of the elements folded in, and the offset of the butterfly step nearest the leaves; it drops how many exchanges ran and which thread held which element. A shared-memory exchange relocates one value and combines nothing, so it adds no height, and the height then counts elements and holds when the warp count changes. Figure 4 reads the two layouts of Figure 1(b) out of assembly, sixteen elements over four threads and two warps and then over eight threads and one: the cut between the scopes moves, the height of four holds, and the two reach one signature. Balance is the guard, an equalheight pair at every combine, so a left fold keeps its physical structure and its own class. A construct the walk does not model becomes opaque syntax, matched on the instruction’s own text, which keeps the answer sound (see Appendix).
trip constants go into the signature and two kernels differing only in a split count separate. Across 51,152 compiled configurations on GB300, 4,500 on gfx942, and minor-scale benchmarks on GB200 and H100, equivalence the checker certified is equivalence the hardware confirmed. On most kernels, the TorchInductor-generated ones included, it hands back 1.0 to 2.5 times the classes that really exist (Section 7).
6.2
7.1
Implementation
The NVIDIA checker is about 2,500 lines of Python over a PTX parser: the walk, the symbolic address evaluator, the tree representation and its canonical form, and the loop summariser. The AMDGCN checker is its twin over the same tree representation, with the address evaluator and the exchange recognisers replaced for that instruction set, and it is graded against a corpus of its own. Triton’s autotuner already takes a pruning predicate, so the PTX checker becomes one with no change to the autotuner: a search keeps the configurations bitwise equivalent to a reference, either its own first configuration or one the caller compiles and hands in, and benchmarks those. A loop is handled where the walk meets it. A loop-carried accumulation is recognised by its shape in the assembly, an MMA or a combine whose destination is also one of its sources, as in fma.rn.f32 %f5, %f1, %f2, %f5. The walk then replaces the one-iteration add(seed, chunk) it would otherwise produce with a single fold node over the chunk, dropping the pre-loop seed, a value every configuration of one kernel shares. The node is keyed on the loop’s own increment, the BLOCK_N of a chunked reduction, since re-chunking a scalar fold regroups the sum. The key comes off for a 𝐾 loop whose tensor cores accumulate exactly (Section 3.2.1) and whose accumulator nothing outside the MMAs touches before the loop ends: every BLOCK_K then issues the same products into the same accumulator in the same order, so the chunk size leaves the signature and those configurations merge. An entry whose reduction the walk could not reconstruct keeps its launch geometry as the signature, so two kernels that both reconstruct to nothing stay apart. A reduction whose accumulation crosses a loop back edge is summarised from one iteration, which is exact for a single fold; a nest of folds would make that summary a guess, so the loop’s own
7
Evaluation
GB300 (sm_103) carries every measurement below. The cuBLAS reconstruction and the checker’s soundness gate run on GB200 (sm_100) and H100 (sm_90) too, H100 carries the layout optimization as well, and gfx942 (CDNA3) carries the AMDGCN checker over a corpus of its own. The evaluation splits in two: bit-level correctness first, then what holding it costs. Benchmarks
cuBLAS 13, Triton 3.8, PyTorch 2.12 and CUDA 13 throughout; bit-identical means the outputs agree byte for byte on every input draw. Shapes. A measurement that feeds a kernel draws its shapes from one of two sets. The realistic static set is 390 fp16 (𝑀, 𝑁 , 𝐾) triples read from the layer dimensions of open-weight models, each carrying its layer and the operation that follows it, with the ranges in Table 1. The fuzzing set is drawn at random over two dtypes and reaches the corners a layer list never visits, 𝑀 = 1 and a 𝐾 in the hundreds of thousands: 110,813 shapes for the reconstruction and 1,400 for the performance work. Kernels. A measurement that needs the kernel itself takes it from three sources: microkernels written for this work, each isolating one ordering or layout question; the Triton kernels TorchInductor [8] emits, copied verbatim; and a zoo of 97 kernels adapted from open-source complex benchmarks [22, 23, 38, 46, 55]. The ordering work runs 24 (kernel, dtype) pairs over the first two sources at fp16, fp8 and fp32, plus a LayerNorm weightgradient reduction from the zoo; the checkers are graded on one PTX file per autotuner configuration, drawn from all three. 7.2
Evaluation on bit-level correctness & static checker soundness
Reconstructing cuBLAS. Table 2 puts a Triton GEMM against cuBLAS on 110,813 random fp16 shapes on GB300, at ten input draws each. Every shape whose cuBLAS kernel computes the full sum comes back byte for byte identical, across all eight algorithm families the library dispatches to. The 60 shapes in the last column are one defect in cuBLAS itself: at very deep contraction lengths the library sums a whole number of blocks and drops the tail, which we reproduce with no Triton in the picture. The same holds at fp8 and
16
Lo
RA
B
(1
) oE
up
oE
M
w do
n
4)
(3
d
LM
M
a he
8)
3)
(1
9 (1
all
ge
ns
lar
tio
ion
c oje pr
4)
(4
ge
ad
LM
lar
ge
wn
he
oE
do
M
lar
3)
(2
LP
ge
up
lar
(2
do
ge
wn
lar
1)
(2
ge
oE
up
lar
ge
RA
B
lar
1)
(2
e
all
g lar
0.79
0.71
0.69
0.88
0.60
0.65
0.70
1)
(2
Lo
M
M
0.56
0.61
0.75
0.72
0.65
1)
LP
M
0.74
0.73
0.66
0.69
0.62
1)
(4
0.64
0.93
0.91
0.88
0.77
0.69
0.60
0.65
0.73
0.66
1.00
0.91
0.88
0.58
0.64
4)
(3
0.53
0.60
0.50
0.58
0.79
1.25 1.00 0.75 0.50 0.25
0.59
torch-generated GEMM, max-autotune (TRITON) bit-exact Triton GEMM, autotuned bit-exact Triton GEMM, large tensor accelerated 0.74
speed relative to cuBLAS (higher is faster)
Yang et al.
)
92
(1
(a) a GEMM on its own, against cuBLAS
0.88
0.91
1.00
0.67
0.86
0.59
0.80
1.00
launch overhead kernel time 1.00
1.05
0.97
0.78
1.00
bit-exact, autotuned 1.04
1.00
1.04
1.18
1.00
1.0
cuBLAS, then the epilogue torch.compile, max-autotune (ATEN + TRITON) 0.93
1.5
0.88
2.0
1.00
launch + kernel time, relative to the baseline (lower is faster)
nt
te at
0.5
8)
LU iG Sw
(3
al
u sid
re
2)
(2
S
RM
w
h eig
)
0)
)
14
t(
U
Sw
L iG
cw
(1
A
R Lo
)
0 (1
LU
d
re
ua
sq
(2
all
Re
6)
(9
(b) a GEMM with its epilogue folded in
Figure 5. What the bitwise constraint costs, on a GEMM and on a fused epilogue. Table 1. The static GEMM shapes, by layer. layer
shapes
𝑀
𝑁
Table 2. Reconstructing cuBLAS bit for bit, on three GPU generations. 𝐾
attention projections MLP up MLP down MoE up MoE down LM head LoRA B
45 256–16,384 2,048–18,432 2,048–16,384 21 256–16,384 6,144–33,792 2,048–7,168 21 256–16,384 2,048–7,168 6,144–33,792 55 16–3,072 512–3,072 2,048–8,192 57 16–6,144 2,048–8,192 512–3,072 54 1–256 100,352–248,320 2,048–8,192 137 80–16,384 768–32,768 8–64
all
390
1–16,384
512–248,320
8–33,792
Read from the layer dimensions of 16 open-weight models: Qwen3.8-2.4TA95B, Qwen3.8-27B, Qwen3.6-35B-A3B, Qwen3-30B-A3B, DeepSeek-V4 Pro and Flash, Kimi-K2.6 and K3, GLM-5.2, GLM-4.7-Flash, gpt-oss-120b and 20b, Ling-3.0-flash, MiniMax-M2, Nemotron-3.5-L, and granite-4.1-8b. Every dimension is a named key in that model’s own config.json; Appendix C gives them per model.
GEMM algorithm (cuBLAS kernel)
tested bit identical cuBLAS bug
Single-pass accumulation (nvjet) Split-K (nvjet) Per-MMA accumulation (cutlass) Split-K, per-MMA (cutlass) Three-level chain (gemmSN_NN) Lane-tree GEMV (gemv2T) Contiguous-slice GEMV (gemv2T) Workspace GEMV (reduce_1Block)
19,449 18,636 13,473 33,677 4,731 14,897 1,722 4,228
19,449 18,576 13,473 33,677 4,731 14,897 1,722 4,228
0 60 0 0 0 0 0 0
total
110,813
110,753
60
GB300 (sm_103) against cuBLAS 13, fp16, 110,813 random shapes at ten input draws each, from six regimes spanning 𝑀 and 𝑁 from 1 to 120,000 and 𝐾 from 8 to 300,000. Appendix D is the defect behind the last column, with its standalone reproduction. architecture
on the two older generations: 100% bit-identical on every shape outside that defect. Soundness. Over 51,152 configurations on GB300, from 47 Triton kernels in five floating-point formats, every pair the checker certified as bitwise equivalent returned the same bytes on every input draw. The same gate holds over the AMDGCN checker’s own 4,500 configurations, and over minor-scale benchmarks on GB200 and H100. Table 3 reads that corpus one dtype per kernel: on the reductions and
tested bit identical cuBLAS bug
GB300, fp8 42,793 GB200, fp16 + fp8 626,522 H100, fp16 + fp8 648,720
99.93% 99.81% 99.82%
0.07% 0.19% 0.18%
The fp8 half of the same GB300 campaign, then the same evaluation on GB200 (sm_100) and H100 (sm_90), whose records hold the two dtypes together. Counts are shapes; rates are over byte comparisons.
GEMMs this paper targets the checker over-splits the true
Table 3. Static equivalence checking on NVIDIA PTX and on AMD GCN.
d_
_2 m su
0 (0) 0 (0) 0 (0) 0 (0) 0 (0) 0 (0) 0 (0) 0 (–) 0 (–) 0 (0) 0 (–) 0 (–) 0 (–)
4 torch.compile emits a Triton GEMM of its own rather than reproducing
cuBLAS’s algorithm, so it is not necessarily the faster kernel.
ol
c d_
2
2
f3
ig
f3
b ol_
d_ _2 d_c m _3 m _2 u m s su
su
141.9% 100.0%
118.9%
2
16
2 2 2 2 f3 f3 f3 f3 0 0 op db m m i i o w l d _ _d _d l_b exp d_ um um ra wd co l_ l_s _g _b ols s co c co m _ r bia ue no r_ er og to pil lay duc e r_ in to uc ind
er
t ou
2
f3
6 f1
bf
m
f3
u _s
94.5%
100.2%
98.6%
90.5%
57.3%
91.4%
85.4%
90.8%
88.7%
72.6%
90.0%
80.5%
96.4%
68.7%
75%
72.5%
100%
50% 25%
6 8 6 8 6 8 6 8 6 6 8 6 8 2 2 2 2 f1 fp f1 fp f1 fp f1 fp bf1 f1 fp f1 fp f3 f3 f3 f3 0 0 ol ol ig ig er er m m p p b d 0 0 xis axis d_c d_c l_b l_b out out bf16 _su _su _loo _loo dwd ope dim dim a _ d_ _2 _2 co co d_ d_ l_ xp xp m m d_ _lo d_ m_ d _ _ o _2 _2 m m d d _3 _3 c l_e l_e l_su l_su _bw um _gra lsu m m su su _2 _2 um um co co co co rm n_s ias _co m m s s su su i su su no la _b ue er _p tor log lay ctor duc epi u in or_ t ind uc ind
Evaluation on performance cost
A GEMM on its own. Panel (a) of Figure 5 splits the 390 shapes at 5 GFLOP of arithmetic (2𝑀𝑁 𝐾), where a GEMM gains enough work to saturate the machine. Above it, where the Triton GEMM torch.compile generates with the order left free reaches 61–88% of cuBLAS over the layer groups,4 the bit-exact GEMM reaches 56–93% and the large-tensor accelerated one 69–91%. Below the line the free one reaches 59–88% and ours 50–100%. A GEMM with its epilogue. Panel (b) folds the operation that follows each GEMM in its model’s own code into the kernel, over 96 shapes in six epilogue groups. Each bar is a call’s launch overhead plus its kernel time over the same two parts of a cuBLAS GEMM and a separate epilogue kernel, where launch is 82% of the baseline. Where torch.compile reaches 85–125% of that baseline, the bit-exact fused kernel reaches 95–168%: it runs 1.8 times as long as the baseline’s kernel half and still wins the pair by removing one launch. The enforced reduction order. Figure 6 reads the fixbits layout optimization against the free-order mode, which lets the compiler pick any reduction order. Both modes are tuned, so each bar is that mode’s own best; the white part of a bar is where the ordering constraint alone leaves the kernel. 19 of the 27 bars finish within 10% of the free-order mode, and on H100 the optimization carries six of the ten kernels to or past it, by as much as 42% on the column sum a matrix-multiply epilogue emits. On GB300 it passes that mode once and otherwise reaches 57% to 99% of it.
f3
(a) H100
Black is the PTX checker on GB300; blue, in brackets, is the AMDGCN checker on gfx942 over a corpus of its own, where a dash is a kernel that corpus does not carry.
7.3
is ax
99.5%
reductions on different dims/axes (16), fp16 2,304 (1,920) 204 (144) 81 (87) 2.5 (1.7) softmax, fp16 144 (120) 8 (10) 4 (3) 2.0 (3.3) layernorm, fp16 144 (120) 19 (20) 9 (3) 2.1 (6.7) rmsnorm, fp16 144 (120) 17 (20) 8 (3) 2.1 (6.7) gemm, fp16 6,304 (72) 1 (4) 1 (4) 1.0 (1.0) gemm_bias_relu_fp_fusion, bf16 192 (24) 1 (2) 1 (2) 1.0 (1.0) gemm_kgroup, bf16 1,152 (132) 138 (1) 6 (1) 23.0 (1.0) gemm_reduce_sum, fp32 270 (–) 60 (–) 9 (–) 6.7 (–) gemm_softmax, fp32 270 (–) 60 (–) 9 (–) 6.7 (–) gemm_tma_store, fp16 18 (12) 1 (1) 1 (1) 1.0 (1.0) inductor_sum_loop, fp16 20 (–) 20 (–) 20 (–) 1.0 (–) inductor_splitk_gemm, fp16 4 (–) 4 (–) 4 (–) 1.0 (–) flash_attention, fp16 2,720 (–) 2,655 (–) 266 (–) 10.0 (–)
102.6%
25%
0
split merges
97.7%
50%
72.8%
configs checker bytes
111.0%
75%
over-
speed, as a share of the free-order mode (both tuned)
kernel
over-
98.3%
100%
2
classes
99.5%
99.1%
speed, as a share of the free-order mode (both tuned)
partition by 1.0 to 2.5, reaching it exactly on the GEMM family, and four kernels outside of our major focus run from 6.7 to 23.0.
101.0%
Taming Bitwise Behavior in GPU Kernels with Tensor Core
(b) GB300 raw inner-tree mode
speedup of layout optimization
Figure 6. What the reduction-ordering constraint costs, and what the layout optimization takes back.
8
Discussion and future work
Several limitations remain, and each is a future direction this work leaves open. i) Section 3 says enough about bit-level semantics for a descriptor to hold across machines and across library versions: none of its fields names a machine, and the split is written down rather than read off the hardware, so it survives the change in SM count at which cuBLAS’s own guarantee lapses [42]. Measuring that cross-machine robustness is what the resources behind this work did not reach, due to resource limitations. ii) Taming a GEMM’s arithmetic could go up a level, from one kernel to a system that composes many of them; and how much contribution this work can benefits the ML/RL training job is not measured either due to resource limitation. iii) It could go down a level equally, from the PTX and AMDGCN the checkers read to the machine code ptxas emits below them: every soundness claim here is relative to the assembly, and a checker reading SASS would close that gap at the cost of a format the vendor does not document. iv) Attention and the GEMM fusions whose epilogue impact the reduction order does not deeply covered but only a theoretical analysis in Appendix F; v) The
Yang et al.
principle could be built into automatic kernel generation: a generator such as TorchInductor [8] could pin the order it emits and tune inside one bitwise-equivalence class.
9
Conclusion
A GPU kernel’s bit-level semantics are decided by the structure of its floating-point accumulation, and that structure is chosen by a compiler and a hardware mapping rather than by whoever wrote the kernel. This paper makes the structure a thing that can be handled. It can be written down, as a descriptor that holds what decides the bit-level semantics. It can be recovered from a closed library by profiling and numerical experiment, our first black-box reconstruction to reach bit-level correctness, which reproduces cuBLAS’s own bytes across three GPU generations. It can be requirement of a compiler, and we show that such requirement costs can be won back at acceptable rate. And it can be decided statically from compiled code, soundly, by our first checkers to settle bitwise equivalence between compiled GPU kernels on two vendors’ instruction sets, over tens of thousands of autotuner configurations, and that decision goes back to the autotuner, which then searches inside one class. Our work was also heavily evaluated at the view of performance cost, which is much more acceptable than what the community usually worried on enforcing numerical correcteness.
A
High level graph
Figure 7 places the four mechanisms on the path a kernel takes from a template to the bytes it returns, beside the closed library whose arithmetic Section 4 reconstructs.
B
The layout optimization
The pass takes three numbers besides the module: the warp size 𝑤, a cap 𝑐 on the elements one thread may hold once the operand has moved, and a spread factor 𝑢. Every test is a skip, so a reduction it cannot place is left exactly as it was. A reduction crossing a thread block is out of scope, and one already inside a warp has nothing to win. The last two guards ask whether the conversion pays: the candidate has to put more lanes on the axis than the current layout does, and the axis has to be under-spread by a factor 𝑢 before the conversion earns its cost.
Algorithm 1 The layout optimization for a pinned reduction. Require: module 𝑀, warp size 𝑤, per-thread element cap 𝑐, spread factor 𝑢 1: for all reductions 𝑟 ∈ 𝑀 with reduction_ordering = inner_tree do 2: 𝐿 ← operand layout of 𝑟 ; skip unless 𝐿 is blocked 3: skip if 𝑟 crosses a thread block, or is already warpsynchronous 4: skip if |tile|/(𝑤 · 𝑛 warps ) > 𝑐 ⊲ the relayout would spill 5: 𝐶 ← Ideal(𝑟 ): lanes onto the axis first and the rest onto the kept dimensions; warps onto the kept dimensions only; the remaining axis extent into registers 6: skip if 𝐶 = 𝐿, or if 𝐶 puts no more lanes on the axis than 𝐿 does 7: skip if extent(axis) < 𝑢 · 𝐿.lanes(axis) ⊲ already well spread convert each operand to 𝐶, clone 𝑟 onto it, convert the 8: results back 9: end for
C
The models the shapes are read from
A row of Table 4 becomes a set of GEMM shapes by reading the weight each layer multiplies by. In 𝑀 × 𝑁 × 𝐾, 𝑁 is that weight’s output width, 𝐾 its input width, and 𝑀 a token count. Attention projections take their widths from hidden and from the head counts in the same file; the MLP pair takes 𝑁 from FFN going up and 𝐾 from FFN coming down; the MoE pair does the same with expert in place of FFN ; and the LM head takes 𝑁 from vocab, which is where the widest shapes in Table 1 come from. A LoRA adapter contributes the second of its two matrix multiplications, whose 𝐾 is the adapter rank rather than a model width. That is why the LoRA rows of Table 1 reach down to 𝐾 = 8 while their 𝑁 stays at hidden or FFN, and it is the one layer group whose contracted dimension a model’s configuration does not fix.
D
The cuBLAS bug on split-K tail
The defect sits in the nvjet split-K path cuBLASLt reaches at ALGO_ID 66. It reproduces on GB200, GB300 and H100, and under cuBLAS 12.8.5 as well as 13.1.1. Which shape loses its tail moves with the architecture and the library together, so Table 5 gives three that separate the pairings. Fill 𝐴 and 𝐵 with ones. Every element of 𝐶 must then be exactly 𝐾, being a sum of 𝐾 products of one with one. A product of two ones is exact in fp32, the accumulator is fp32, and an fp32 output holds every whole number below 224 , so the number that comes back is the count of 𝑘 values that were summed. On some shapes it comes back short by a whole number, which is that many missing terms of 1 × 1. Where the piece is lost. Figure 8 draws 𝐾 = 8648. The 𝑘 step of one threadblock in this fp16 kernel family is 64, and 8648 holds 135 whole steps with 8 values over. cuBLASLt spreads the work over nine threadblocks, and 135/9 = 15
Taming Bitwise Behavior in GPU Kernels with Tensor Core
ML Framework
Triton kernel template autotuner config. 1
config. i
…
…
config. n
GEMM descriptor fixes the reduction order offline derived
§5 Compiler enforcement & Layout Optimization
Triton compiler (MLIR) / LLVM Backend
tile IR
tile IR
tile IR
GPU assembly
PTX/GCN
PTX/GCN
PTX/GCN
SASS
SASS
SASS
execution code
execution code
execution code
ptxas machine program
cuBLAS black box api with limited algorithm info.
§4 Black-box reconstruction
closed source CUDA compiler
§6 Static equivalence checking & partition FP computation reconstruction
bit equivalent?
SASS execution code
Figure 7. High level overview of this work. Table 4. The open-weight models the static shapes are read from. model
hidden
Qwen3.8-2.4T-A95B DeepSeek-V4-Pro DeepSeek-V4-Flash GLM-5.2 GLM-4.7-Flash Kimi-K2.6 Kimi-K3 gpt-oss-120b
8192 7168 4096 6144 2048 7168 7168 2880
FFN expert experts active
vocab
model
– – – 12288 10240 18432 33792 –
248,320 129,280 129,280 154,880 154,880 163,840 163,840 201,088
gpt-oss-20b Nemotron-3.5-L MiniMax-M2 Qwen3-30B-A3B Ling-3.0-flash Qwen3.6-35B-A3B Qwen3.8-27B granite-4.1-8b
2048 3072 2048 2048 1536 2048 3072 2880
512 384 256 256 64 384 896 128
10 6 6 8 4 8 16 4
hidden
FFN expert experts active
2880 – 2688 1856 3072 – 2048 6144 2560 6144 2048 – 5120 17408 4096 12800
2880 1856 1536 768 768 512 – –
32 128 256 128 512 256 – –
4 6 8 8 8 8 – –
vocab 201,088 131,072 200,064 151,936 157,184 248,320 248,320 100,352
Every number is the value of a named key in that model’s own config.json, read from the Hugging Face hub [25] on 2026-08-16. FFN is the dense feed-forward width and expert the routed one; a dash means the key is absent, so a model with no FFN entry is mixture-of-experts in every layer and one with no expert entry is dense in every layer. active is how many experts a token is routed to. The 16 are drawn from the highest trending and most downloaded models in that hub’s text-generation listing on the day the set was fixed. DeepSeek-V4 [18] appears in two sizes.
k=0
8648
CTA 0
CTA 1
CTA 2
CTA 3
CTA 4
CTA 5
CTA 6
CTA 7
CTA 8
nine threadblocks, 960 values of k apiece
8
8640
CTA 8 7680 64
64
…
64
fifteen whole blocks of 64, the last ending at 8640 p0
p1
p2
p3
p4
p5
p6
p7
p8
splitKreduce C[m][n] = 8640 all nine partials are added, and 9 × 960 = 8640
tail out of scale
Figure 8. Where the last eight values of 𝑘 go, at 𝐾 = 8648.
exactly, so each takes 15 whole steps, or 960 values of 𝑘. Every threadblock is full, and the last step of the last one is a whole 64 ending at 8640. The reduction then adds all nine partial results, so everything computed reaches the output
and nothing is added twice. The nine ranges cover [0, 8640), and 𝐾 is 8648. The condition. Write 𝑏 for the 𝑘 step, 𝑞 = ⌊𝐾/𝑏⌋, 𝑡 = 𝐾 mod 𝑏, and 𝑠 for the split count cuBLASLt picks. The tail is lost exactly when ALGO_ID = 66, 𝑡 ≠ 0, 𝑞 mod 𝑠 = 0 and 𝑠 > 𝑡. At 𝐾 = 8648 that reads 𝑞 = 135, 𝑡 = 8, 𝑠 = 9. The last clause is what makes 𝑞 divide evenly into ranges of whole steps with the tail left over: a split that took an uneven share would give one threadblock the short step and sum it. The ALGO_ID clause carries its own weight. On H100 the shape 1 × 2 × 1032 satisfies all three arithmetic clauses and loses nothing, because it lands on ALGO_ID 23, a CUTLASS kernel rather than an nvjet one. Over 1,542,555 real GEMMs per library on H100 the four-clause condition predicted every loss with no false positives and no false negatives, and dropping the ALGO_ID clause misclassified 3,557 shapes under each library. Which library, which machine. Table 5 gives three shapes and where each one loses. No single shape covers
Yang et al.
Table 5. Which shape loses its tail on which pairing.
𝐾 8,648 57,608 11,528
GB200 13.1.1
GB300 13.1.x
−8
−8 ok
GB300 H100 12.8.5 13.1.1 ok −8
ok −8
Table 6. What the acceleration is worth, by plan mode.
H100 12.8.5
plan mode single-pass accumulation split-K split-K, grouped per-MMA accumulation three-level chain lane-tree GEMV contiguous-slice GEMV workspace GEMV
ok −8 −8
A dash of −8 is eight values of 𝑘 left out of the sum, ok is the full sum, and a blank cell was not measured. Architectures are sm_100, sm_103 and sm_90.
every pairing, because the loss needs cuBLASLt to split 𝐾 at all and that decision moves with both the architecture and the library. On GB300 the two libraries were nearly disjoint: over 12,282 shapes, 13.1.1 met the condition on 193 and 12.8.5 on 309, with no overlap, and at 𝐾 = 8648 the older library returns a split count of one and stays whole. On H100 they nearly coincide instead, and 11,528 is the smallest losing 𝐾 for both. Every pairing runs the same tile family 64x8_64x16_1x1 in the nvjet split-K path, followed by one splitKreduce launch. The reproducer. It calls cuBLASLt directly, since torch.matmul never reaches this algorithm on its own path. Two things in it are load bearing. The workspace is where split-K writes its partial results, so a zero-size workspace keeps cuBLASLt whole and both shapes come back right. The output is fp32 because at part two’s depth an fp16 output cannot be trusted: above 16,384 the fp16 step is 16, so a 𝐾 of the form 64𝑞 + 8 sits halfway between two fp16 values and rounds down on its own. The control is 𝐾 = 57,616, a shape whose sum is complete, which an fp16 output reads as 57,600 and an fp32 output reads correctly. Asking for fp32 leaves the heuristic’s choice alone, checked over 600 random shapes that returned a bit-identical configuration either way. M, N, K = 1, 8, 8648 # q = 135, t = 8, splits 9 a = torch.ones(M, K, dtype=torch.float16, device="cuda") b = torch.ones(K, N, dtype=torch.float16, device="cuda") c = torch.empty(M, N, dtype=torch.float32, device="cuda") # algo is NULL, so cuBLASLt picks the split; the # workspace is what lets it split at all cublasLtMatmul(handle, desc, one, a, a_layout, b, b_layout, zero, c, c_layout, c, c_layout, None, workspace, WORKSPACE_BYTES, stream) assert c[0, 0].item() == K
# reads back 8640
How far it reaches. The three shapes are the visible end of a sweep. Over 496,906 random shapes on a GB300 under 13.1.1, 1,062 came back other than bit-identical to a full-𝐾 reference and 1,059 of those meet the condition; under 12.8.5 on the same GPU, 179 shapes were flagged and all 179 lose exactly 𝑡 values of 𝑘, while 9,525,600 GEMMs at smaller 𝐾 stay whole. On H100 the counts are 2,714 under 13.1.1 and 2,995 under 12.8.5, each out of 1,542,555 real GEMMs. Every
reference
accelerated
0.303 0.426 0.624 0.654 0.184 0.210 0.596 0.853
0.802 0.857 1.639 2.086 0.998 0.843 0.962 1.057
Device geometric mean of cuBLAS’s kernel time over the arm’s, by CUDAgraph replay with L2 flushed between replays, on an idle GB300. Above 1.0 is faster than cuBLAS, and every such entry is in the unaligned regime, where the win comes from repacking the operands to restore alignment.
one of the 5,709 observed losses was exactly 𝐾 mod 64. It reaches ordinary shapes as well as skinny ones: 363 of 885 (𝑀, 𝑁 ) pairs under 13.1.1 and 416 of 885 under 12.8.5 lose 𝑘 somewhere, and 32 × 32, 64 × 64 and 128 × 128 all lose 8 at 𝐾 = 11,528.
E
The large-tensor acceleration on GB300
A descriptor fixes the arithmetic and says nothing about the schedule, so a second kernel that realises one descriptor may choose any schedule it likes and still return the same bytes. The accelerated arm is that freedom taken up on GB300: each launcher gained one branch at the top, taken on sm_103 and falling through to the original body everywhere else, so the kernels the other two generations run are the ones they ran before. The branch carries a GROUP_M tile swizzle, tensor-memory descriptors, a persistent grid with warp specialisation, deeper pipelining, 32-bit indexing wherever an operand cannot reach 231 elements, a copy-free operand path, a single launch across the split-K slices, and a tile rule for each of the eight plan modes a GB300 reaches. One change needed an argument that no value moves. The per-MMA plan accumulates in groups of KPD real 𝑘 elements, and the reference kernel spends one tl.dot on each group. The accelerated kernel spends one on 𝐺 groups at once. A tl.dot whose 𝑘 extent is 16𝐺 lowers to 𝐺 chained 𝑘 = 16 MMAs on one fp32 accumulator in increasing 𝑘, and each rounds the accumulator exactly once, so lanes [16𝑔, 16𝑔 +16) of the wide dot are the accumulator update the 𝑔-th narrow dot performed. Where a group is narrower than an MMA, at KPD 8, the tail of each group is masked to zero, which is the stand-in the narrow kernel already used, repeated 𝐺 times in one tile. The residue tile keeps one group per dot, since its last group can be short: # residue tile: one group per dot, the last one may be short for g in tl.range(0, nfirst, num_stages=1): off = g * KPD real = tl.arange(0, 16) < tl.minimum(KPD, rbk - off) kk = k0 + off + tl.arange(0, 16)
Taming Bitwise Behavior in GPU Kernels with Tensor Core a = tl.load(A + om[:, None]*am + kk[None, :]*ak, mask=real[None, :], other=0.0) b = tl.load(B + kk[:, None]*bk + on[None, :]*bn, mask=real[:, None], other=0.0) acc = tl.dot(a, b, acc) # the rest: G groups per dot, every group full BKW: tl.constexpr = 16 * G STEP: tl.constexpr = G * KPD ix = tl.arange(0, BKW) # lane ix holds k off+koff[ix]; KPD 8 zeroes each tail koff = (ix // 16) * KPD + ix % 16 wr = ix % 16 < KPD for _ in tl.range(0, (klen - rbk) // STEP, num_stages=NSTAGE): a = tl.load(ap, mask=wr[None, :], other=0.0) b = tl.load(bp, mask=wr[:, None], other=0.0) acc = tl.dot(a, b, acc) ap += STEP * ak bp += STEP * bk
Across the whole arm, 154,904 byte comparisons returned zero differences. The load-bearing half of that is a sweep of the plan-parameter space rather than of shapes: 7,868 parameter combinations over the nine modes at ten input draws each. cuBLAS’s own heuristic hands any one shape a small corner of a plan’s parameter space, so the sweep walks that space directly, and it was checked for power against a deliberately broken kernel whose backwards split-K partials it caught. Every input set was drawn twice over, once ordinarily and once with the exponents spread across the dtype’s usable range, since narrow exponents hide a regrouping. The descriptor path costs about 100 𝜇s of host time per call against about 20 𝜇s for the pointer launch, which is what builds the tensor descriptors. Device time does not see it and a caller does, so it is visible on a GEMM under roughly 100 𝜇s.
F
The reduction order of attention
One attention kernel holds three reductions. The score matrix contracts 𝑄 against 𝐾 over the head dimension, and the output contracts the probabilities against 𝑉 over the key axis; both are matrix multiplications, and their order is the object of Section 3.2. Between them sits an accumulation over the key axis that carries a scale factor, and that factor is what separates attention from every reduction in Section 3.1. The recurrence. Cut the keys of one query row into blocks B1, . . . , B𝐵 of width 𝑤, and write 𝑠𝑡 for the scaled score of key 𝑡 and 𝑣𝑡 for its value row. The kernel carries three running quantities in fp32: a maximum 𝑚 𝑗 , a denominator ℓ 𝑗 and an output row 𝑂 𝑗 . At block 𝑗 it folds the block’s own maximum into the running one and rescales what it already holds by 𝛼 𝑗 before it adds. Writing 𝑚 𝑗 = max(𝑚 𝑗 −1, max𝑡 ∈ B 𝑗 𝑠𝑡 ), 𝛼 𝑗 = 2 𝑚 𝑗 −1 −𝑚 𝑗 and 𝑝𝑡 = 2 𝑠𝑡 −𝑚 𝑗 , Éfp32 the two running sums are ℓ 𝑗 = 𝛼 𝑗 ℓ 𝑗 −1 ⊕fp32 𝑡 ∈ B 𝑗 𝑝𝑡 and
Éfp32 𝑂 𝑗 = 𝛼 𝑗 𝑂 𝑗 −1 ⊕fp32 𝑡 ∈ B 𝑗 𝑝𝑡 𝑣 𝑡 . The base is two because the kernel folds log2 𝑒 into the score scale and calls the hardware’s base-two exponential. The answer is 𝑂 𝐵 /ℓ𝐵 , taken once at the end. Every block boundary is a rounding. In exact arithmetic 𝛼 𝑗 undoes the previous normalisation, and the recurrence returns one value at every 𝑤. In floating point 𝛼 𝑗 𝑂 𝑗 −1 is a multiply over the whole running output row, so each boundary costs one rounding on every element of 𝑂 and one on ℓ, and a row of 𝑁 keys pays ⌈𝑁 /𝑤⌉ of them. The key block width, BLOCK_N in Triton, therefore sets how many roundings the answer carries and where they fall. Section 3.2.2 put BLOCK_K on the other side of that line: an accumulator carried across the turns of a tensor-core mainloop receives the same instruction results in the same order however the turns are cut, so the turn boundary is a speed knob there. Here it is a term of the arithmetic. Masking puts the query block into the answer as well. Left unmasked, the key walk runs from the first key to the last in steps of 𝑤, so the boundaries are multiples of 𝑤 alone and the query block width, BLOCK_M, chooses only which rows share a kernel instance. Causal masking splits the walk in two: a run of blocks lying entirely below the diagonal, and one block straddling it, whose entries above the diagonal are offset by −106 before the maximum is taken, so their probabilities underflow to zero. Both bounds of that split are multiples of BLOCK_M, so the query block width decides where the straddling block begins, how many keys inside it are live, and how many rescales precede it. Two causal configurations differing only in BLOCK_M accumulate over different boundaries, and two unmasked ones accumulate over the same boundaries at any BLOCK_M. What the layout still moves. Two reductions run inside a block. The row maximum selects one of its inputs and rounds nothing, so its dependence tree is free. The row sum Éfp32 𝑡 ∈ B 𝑗 𝑝𝑡 is a plain sum of 𝑤 values, which makes it the object of Section 3.1 exactly, and its tree moves with the thread count and the lane layout in the way that section describes. The probabilities are narrowed before the second matmul. A tensor core takes its operands at the input type, so 𝑝𝑡 is rounded from the fp32 it was computed in down to fp16, bf16 or fp8 before it meets 𝑉 . That rounding lands on every probability of every block, and it sits between the two reductions the kernel is composing, so the operand type is as much a part of the order as the accumulator type is. The descriptor. An attention descriptor carries the key block width, the masking convention together with the query block width when masking is on, the tree of the row sum inside a block, the type the probabilities are narrowed to, the descriptors of the two matrix multiplications, and the place of the final division.
Yang et al.
References [1] R. C. Agarwal, S. M. Balle, F. G. Gustavson, M. Joshi, and P. Palkar. 1995. A Three-Dimensional Approach to Parallel Matrix Multiplication. IBM Journal of Research and Development 39, 5 (Sept. 1995), 575–582. doi:10.1147/rd.395.0575 [2] Willow Ahrens, James Demmel, and Hong Diep Nguyen. 2020. Algorithms for Efficient Reproducible Floating Point Summation. ACM Trans. Math. Software 46, 3 (2020), 1–49. doi:10.1145/3389360 [3] Amazon Web Services. 2026. AWS Trainium and Inferentia. https:// aws.amazon.com/ai/machine-learning/trainium/. Amazon’s training and inference accelerators. [4] Amazon Web Services. 2026. Neuron Kernel Interface (NKI). https: //awsdocs-neuron.readthedocs-hosted.com/en/latest/nki/. Tile-level Python kernel programming for AWS Trainium and Inferentia. [5] AMD. 2023. AMD CDNA 3 Architecture. https://www.amd. com/content/dam/amd/en/documents/instinct-tech-docs/whitepapers/amd-cdna-3-white-paper.pdf. Compute unit: four SIMD units, 64 KB local data share, 64-work-item wavefront. [6] AMD. 2026. HIP Programming Model. https://rocm.docs.amd.com/ projects/HIP/en/latest/understand/programming_model.html. Workgroup, wavefront and compute unit; wavefront 64 on CDNA, 32 or 64 on RDNA. [7] AMD. 2026. rocBLAS Documentation. https://rocm.docs.amd.com/ projects/rocBLAS/en/latest/. The ROCm BLAS library, implemented in HIP and tuned for AMD GPUs. [8] Jason Ansel, Edward Yang, Horace He, Natalia Gimelshein, Animesh Jain, Michael Voznesensky, Bin Bao, Peter Bell, David Berard, Evgeni Burovski, Geeta Chauhan, Anjali Chourdia, Will Constable, Alban Desmaison, Zachary DeVito, Elias Ellison, Will Feng, Jiong Gong, Michael Gschwind, Brian Hirsh, Sherlock Huang, Kshiteej Kalambarkar, Laurent Kirsch, Michael Lazos, Mario Lezcano, Yanbo Liang, Jason Liang, Yinghai Lu, C. K. Luk, Bert Maher, Yunjie Pan, Christian Puhrsch, Matthias Reso, Mark Saroufim, Marcos Yukio Siraichi, Helen Susnea, Shunting Zhang, Michael Zhang, Matei Zaharia, and Soumith Chintala. 2024. PyTorch 2: Faster Machine Learning Through Dynamic Python Bytecode Transformation and Graph Compilation. In Proceedings of the 29th ACM International Conference on Architectural Support for Programming Languages and Operating Systems (ASPLOS), Volume 2. doi:10.1145/3620665.3640366 [9] Seongwon Bang, Seunghyeon Nam, Inwhan Chun, Ho Young Jhoo, and Juneyoung Lee. 2022. SMT-Based Translation Validation for Machine Learning Compiler. In Computer Aided Verification (CAV), Part II. 386–407. doi:10.1007/978-3-031-13188-2_19 [10] Srinadh Bhojanapalli, Kimberly Wilber, Andreas Veit, Ankit Singh Rawat, Seungyeon Kim, Aditya Krishna Menon, and Sanjiv Kumar. 2021. On the Reproducibility of Neural Network Predictions. arXiv:2102.03349 [cs.LG] [11] Guy E. Blelloch and Bruce M. Maggs. 2010. Parallel Algorithms. In Algorithms and Theory of Computation Handbook (2 ed.), Mikhail J. Atallah (Ed.). Chapman and Hall/CRC, 25–1–25–43. doi:10.1201/ 9781584888215-c25 [12] Sylvie Boldo, Jacques-Henri Jourdan, Xavier Leroy, and Guillaume Melquiond. 2015. Verified Compilation of Floating-Point Computations. Journal of Automated Reasoning 54, 2 (2015), 135–163. doi:10.1007/s10817-014-9317-x [13] Tianqi Chen, Thierry Moreau, Ziheng Jiang, Lianmin Zheng, Eddie Yan, Meghan Cowan, Haichen Shen, Leyuan Wang, Yuwei Hu, Luis Ceze, Carlos Guestrin, and Arvind Krishnamurthy. 2018. TVM: An Automated End-to-End Optimizing Compiler for Deep Learning. In 13th USENIX Symposium on Operating Systems Design and Implementation (OSDI). 578–594. [14] Sharan Chetlur, Cliff Woolley, Philippe Vandermersch, Jonathan Cohen, John Tran, Bryan Catanzaro, and Evan Shelhamer. 2014. cuDNN: Efficient Primitives for Deep Learning. arXiv:1410.0759 [cs.NE]
[15] Caroline Collange, David Defour, Stef Graillat, and Roman Iakymchuk. 2015. Numerical Reproducibility for the Parallel Reduction on Multiand Many-Core Architectures. Parallel Comput. 49 (2015), 83–97. doi:10.1016/j.parco.2015.09.001 [16] Thomas H. Cormen, Charles E. Leiserson, Ronald L. Rivest, and Clifford Stein. 2022. Introduction to Algorithms (4 ed.). MIT Press, Cambridge, MA. [17] Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, and Christopher Ré. 2022. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. In Advances in Neural Information Processing Systems 35 (NeurIPS). [18] DeepSeek-AI. 2026. DeepSeek-V4: Towards Highly Efficient MillionToken Context Intelligence. arXiv:2606.19348 [cs.CL] [19] James Demmel and Hong Diep Nguyen. 2013. Fast Reproducible Floating-Point Summation. In 2013 IEEE 21st Symposium on Computer Arithmetic (ARITH). 163–172. doi:10.1109/ARITH.2013.9 [20] Jack J. Dongarra, Jeremy Du Croz, Sven Hammarling, and Richard J. Hanson. 1988. An Extended Set of FORTRAN Basic Linear Algebra Subprograms. ACM Trans. Math. Software 14, 1 (March 1988), 1–17. doi:10.1145/42288.42291 [21] Benjamin Ferrell, Jun Duan, and Kevin W. Hamlen. 2019. CUDA au Coq: A Framework for Machine-validating GPU Assembly Programs. In Design, Automation and Test in Europe Conference (DATE). 474–479. doi:10.23919/DATE.2019.8715160 [22] FLA Organization. 2026. Flash Linear Attention: Efficient Triton Implementations for Emerging Model Architectures. https://github.com/flaorg/flash-linear-attention. Accessed 2026-09-06. [23] FlagOpen. 2026. FlagGems: A Triton-Powered Operator Library. https: //github.com/FlagOpen/FlagGems. Accessed 2026-09-06. [24] Horace He and Thinking Machines Lab. 2025. Defeating Nondeterminism in LLM Inference. https://thinkingmachines.ai/blog/defeatingnondeterminism-in-llm-inference/. Thinking Machines Lab blog; accessed 2026-09-05. [25] Hugging Face. 2026. The Model Hub. https://huggingface.co/docs/ hub/models. Each model’s config.json is the source of the layer dimensions used here; the text-generation listing is the source of the model selection. Read 2026-08-16. [26] Intel. 2023. Obtaining Numerically Reproducible Results (Conditional Numerical Reproducibility). https://www.intel.com/content/ www/us/en/docs/onemkl/developer-guide-linux/2023-0/obtainingnumerically-reproducible-results.html. oneMKL Developer Guide for Linux; accessed 2026-09-05. [27] JAX Developers. 2026. Pallas: a JAX Kernel Language. https://docs. jax.dev/en/latest/pallas/index.html. Lowers to Mosaic on TPU and to Triton on GPU. [28] Norman P. Jouppi, Cliff Young, Nishant Patil, David Patterson, Gaurav Agrawal, Raminder Bajwa, Sarah Bates, Suresh Bhatia, Nan Boden, Al Borchers, et al. 2017. In-Datacenter Performance Analysis of a Tensor Processing Unit. In Proceedings of the 44th Annual International Symposium on Computer Architecture (ISCA). 1–12. doi:10.1145/3079856.3080246 [29] Faizan A. Khattak and Mantas Mikaitis. 2025. Accurate Models of NVIDIA Tensor Cores. arXiv:2512.07004 [cs.MS] Bit-accurate innerproduct models for V100, A100, H100 and B200 at 8-, 16- and 19-bit inputs, verified against the hardware. [30] David J. Kuck and Yoichi Muraoka. 1974. Bounds on the Parallel Evaluation of Arithmetic Expressions Using Associativity and Commutativity. Acta Informatica 3, 3 (1974), 203–216. doi:10.1007/BF00288634 [31] Ignacio Laguna. 2020. Varity: Quantifying Floating-Point Variations in HPC Systems Through Randomized Testing. In 2020 IEEE International Parallel and Distributed Processing Symposium (IPDPS). IEEE, 622–633. doi:10.1109/IPDPS47924.2020.00070 [32] Xavier Leroy. 2009. Formal Verification of a Realistic Compiler. Commun. ACM 52, 7 (2009), 107–115.
Taming Bitwise Behavior in GPU Kernels with Tensor Core
[33] Erik Lindholm, John Nickolls, Stuart Oberman, and John Montrym. 2008. NVIDIA Tesla: A Unified Graphics and Computing Architecture. IEEE Micro 28, 2 (2008), 39–55. doi:10.1109/MM.2008.31 [34] Amanda Liu, Gilbert Bernstein, Adam Chlipala, and Jonathan RaganKelley. 2024. A Verified Compiler for a Functional Tensor Language. Proceedings of the ACM on Programming Languages 8, PLDI (2024), 320–342. doi:10.1145/3656390 [35] Nuno P. Lopes, Juneyoung Lee, Chung-Kil Hur, Zhengyang Liu, and John Regehr. 2021. Alive2: Bounded Translation Validation for LLVM. In Proceedings of the 42nd ACM SIGPLAN Conference on Programming Language Design and Implementation (PLDI). [36] Daniel Lustig, Sameer Sahasrabuddhe, and Olivier Giroux. 2019. A Formal Analysis of the NVIDIA PTX Memory Consistency Model. In Proceedings of the 24th ACM International Conference on Architectural Support for Programming Languages and Operating Systems (ASPLOS). 257–270. doi:10.1145/3297858.3304043 [37] David Menendez, Santosh Nagarakatte, and Aarti Gupta. 2016. AliveFP: Automated Verification of Floating Point Based Peephole Optimizations in LLVM. In Static Analysis (SAS). 317–337. doi:10.1007/9783-662-53413-7_16 [38] Meta PyTorch. 2026. Tritonbench: A Collection of PyTorch Custom Operators with Example Inputs. https://github.com/pytorch-labs/ tritonbench. Accessed 2026-09-06. [39] John Nickolls, Ian Buck, Michael Garland, and Kevin Skadron. 2008. Scalable Parallel Programming with CUDA. ACM Queue 6, 2 (2008), 40–53. doi:10.1145/1365490.1365500 NVIDIA Tesla V100 GPU Architecture. [40] NVIDIA. 2017. https://images.nvidia.com/content/volta-architecture/pdf/voltaarchitecture-whitepaper.pdf. WP-08608-001_v1.1; independent thread scheduling, per-thread program counters. [41] NVIDIA. 2025. cuTile Python and the CUDA Tile IR. https://docs. nvidia.com/cuda/tile-ir/latest/. A tile-level Python kernel language for CUDA, lowering through an MLIR-based tile intermediate representation. [42] NVIDIA. 2026. cuBLAS Library Documentation. https://docs.nvidia. com/cuda/cublas/. Results reproducibility: bitwise repeatability only on GPUs of the same architecture with the same number of SMs, and not across toolkit versions. [43] NVIDIA. 2026. CUDA C++ Programming Guide. https://docs.nvidia. com/cuda/cuda-c-programming-guide/. Thread hierarchy: grid, thread block, warp; 1024 threads per block; thread block clusters at compute capability 9.0. [44] NVIDIA. 2026. NVIDIA Blackwell Tuning Guide. https://docs.nvidia. com/cuda/blackwell-tuning-guide/index.html. Compute capability 10.0: 228 KB shared memory per SM, 227 KB the most one thread block may claim, 32 thread blocks per SM. [45] Amir Pnueli, Michael Siegel, and Eli Singerman. 1998. Translation Validation. In Tools and Algorithms for the Construction and Analysis of Systems (TACAS). [46] PyTorch. 2026. torchao: PyTorch Architecture Optimization. https: //github.com/pytorch/ao. Accessed 2026-09-06. [47] PyTorch Team. 2026. Reproducibility and Deterministic Algorithms. https://pytorch.org/docs/stable/notes/randomness.html. PyTorch documentation; accessed 2026-09-05. [48] Penghui Qi, Zichen Liu, Xiangxin Zhou, Tianyu Pang, Chao Du, Wee Sun Lee, and Min Lin. 2025. Defeating the Training-Inference Mismatch via FP16. arXiv:2510.26788 [cs.LG] [49] Geoffrey Sawaya, Michael Bentley, Ian Briggs, Ganesh Gopalakrishnan, and Dong H. Ahn. 2017. FLiT: Cross-Platform Floating-Point Result-Consistency Tester and Workload. In 2017 IEEE International Symposium on Workload Characterization (IISWC). IEEE, 229–238. doi:10.1109/IISWC.2017.8167780
[50] Sanjif Shanmugavelu et al. 2024. Impacts of Floating-Point NonAssociativity on Reproducibility for HPC and Deep Learning Applications. arXiv:2408.05148 [cs.DC] [51] Cecilia Summers and Michael J. Dinneen. 2021. Nondeterminism and Instability in Neural Network Optimization. In Proceedings of the 38th International Conference on Machine Learning (ICML). [52] Rajeev Thakur, Rolf Rabenseifner, and William Gropp. 2005. Optimization of Collective Communication Operations in MPICH. International Journal of High Performance Computing Applications 19, 1 (2005), 49–66. doi:10.1177/1094342005051521 [53] The Rocq Prover Team. 2026. About The Rocq Prover. https://rocqprover.org/about. States that the Rocq Prover was formerly known as the Coq Proof Assistant; accessed 2026-09-05. [54] Philippe Tillet, H. T. Kung, and David Cox. 2019. Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations. In Proceedings of the 3rd ACM SIGPLAN International Workshop on Machine Learning and Programming Languages (MAPL). [55] Triton Developers. 2026. Triton Tutorials. https://github.com/tritonlang/triton/tree/main/python/tutorials. The fused-attention tutorial is the source of the flash-attention kernel graded here; accessed 202609-06. [56] Yuki Uchino, Katsuhisa Ozaki, and Toshiyuki Imamura. 2025. Performance Enhancement of the Ozaki Scheme on Integer Matrix Multiplication Unit. The International Journal of High Performance Computing Applications 39, 3 (2025), 462–476. doi:10.1177/10943420241313064 [57] Haoyu Wang, Junjie Chen, Chuyue Xie, Shuang Liu, Zan Wang, Qingchao Shen, and Yingquan Zhao. 2023. MLIRSmith: Random Program Generation for Fuzzing MLIR Compiler Infrastructure. In Proceedings of the 38th IEEE/ACM International Conference on Automated Software Engineering (ASE). 1555–1566. doi:10.1109/ASE56229.2023. 00120 [58] Lei Wang, Yu Cheng, Yining Shi, Zhiwen Mo, Zhengju Tang, Wenhao Xie, Tong Wu, Lingxiao Ma, Yuqing Xia, Jilong Xue, Fan Yang, and Zhi Yang. 2026. TileLang: Bridge Programmability and Performance in Modern Neural Kernels. In International Conference on Learning Representations (ICLR). Oral. [59] Peichen Xie, Shuotao Xu, Yang Wang, Fan Yang, and Mao Yang. 2025. Bit-Accurate Modeling of GPU Matrix MultiplyAccumulate Units: Demystifying Numerical Discrepancy and Accuracy. arXiv:2511.10909 [cs.AR] Bit-accurate models of every MMA instruction on ten GPU architectures, NVIDIA Volta through Blackwell and AMD CDNA1 through CDNA3; open source as MMA-Sim. [60] Peichen Xie, Xian Zhang, and Shuo Chen. 2025. RepDL: Bit-level Reproducible Deep Learning Training and Inference. arXiv:2510.09180 [cs.LG] [61] Xuejun Yang, Yang Chen, Eric Eide, and John Regehr. 2011. Finding and Understanding Bugs in C Compilers. In Proceedings of the 32nd ACM SIGPLAN Conference on Programming Language Design and Implementation (PLDI). 283–294. doi:10.1145/1993498.1993532 [62] Ziteng Yang, Jun Shirako, and Vivek Sarkar. 2024. Fully Verified Instruction Scheduling. Proceedings of the ACM on Programming Languages 8, OOPSLA2 (2024), 791–816. doi:10.1145/3689739 [63] Jiayi Yuan, Hao Li, Xinheng Ding, Wenya Xie, Yu-Jhe Li, Wentian Zhao, Kun Wan, Jing Shi, Xia Hu, and Zirui Liu. 2025. Understanding and Mitigating Numerical Sources of Nondeterminism in LLM Inference. arXiv:2506.09501 [cs.LG] [64] Lianmin Zheng, Chengfan Jia, Minmin Sun, Zhao Wu, Cody Hao Yu, Ameer Haj-Ali, Yida Wang, Jun Yang, Danyang Zhuo, Koushik Sen, Joseph E. Gonzalez, and Ion Stoica. 2020. Ansor: Generating HighPerformance Tensor Programs for Deep Learning. In 14th USENIX Symposium on Operating Systems Design and Implementation (OSDI). 863–879. [65] Keren Zhou, Mario Lezcano, Adam Goucher, Akhmed Rakhmati, Jeff Niu, Justin Lebar, Pawel Szczerbuk, Peter Bell, Phil Tillet, Thomas
Yang et al.
Raoux, and Zahi Moudallal. 2026. Linear Layouts: Robust Code Generation of Efficient Tensor Computation Using F2 . In Proceedings of the 31st ACM International Conference on Architectural Support for
Programming Languages and Operating Systems (ASPLOS), Volume 1. doi:10.1145/3760250.3762221