From 123193ab63813a5941aa64af4ed77acd0c4688ad Mon Sep 17 00:00:00 2001 From: 0z5a Date: Mon, 21 Sep 2026 04:56:13 +0000 Subject: [PATCH 1/2] Add an opt-in batch-invariant BF16 GEMM cuBLASLt selects its kernel from the full problem shape, so a given row's output bits can depend on which other rows happen to share the batch. That makes results irreproducible across batch composition, which gets in the way of debugging and of bitwise comparison between runs. This adds a forward path with a fixed tile geometry and a reduction order that depends only on the tile coordinates, so a row's result is a function of the row and the weight alone. Supported configuration is BF16, contiguous 2-D operands, Y = X @ W.T, no bias; anything else raises instead of silently falling back. Verified bitwise: every one of 18 row slices reproduces the full-batch result exactly, and accuracy against a float32 reference holds to BF16 tolerance. The fixed schedule costs a median 0.91x of the general path's time. It is not wired into te.Linear's dispatch -- callers opt in explicitly. Signed-off-by: 0z5a --- tests/pytorch/test_batch_invariant_gemm.py | 65 ++++++++ .../cpp_extensions/batch_invariant_gemm.py | 153 ++++++++++++++++++ 2 files changed, 218 insertions(+) create mode 100644 tests/pytorch/test_batch_invariant_gemm.py create mode 100644 transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py diff --git a/tests/pytorch/test_batch_invariant_gemm.py b/tests/pytorch/test_batch_invariant_gemm.py new file mode 100644 index 00000000000..1eed3308790 --- /dev/null +++ b/tests/pytorch/test_batch_invariant_gemm.py @@ -0,0 +1,65 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""A row's output must not depend on which other rows were in the batch. + +The general GEMM path lets cuBLASLt pick a kernel from the full problem shape, so +the same row can produce different bits depending on the batch it arrives in. +``batch_invariant_gemm`` fixes the tile geometry and the reduction order so the +result is a function of the row and the weight only. +""" + +import pytest +import torch + +from transformer_engine.pytorch.cpp_extensions.batch_invariant_gemm import ( + batch_invariant_gemm, + is_supported, +) + +M, N, K = 256, 192, 128 +SLICES = [(0, 1), (0, 7), (5, 29), (128, 256), (0, 256)] + + +@pytest.fixture(scope="module") +def operands(): + if not torch.cuda.is_available(): + pytest.skip("batch_invariant_gemm requires a GPU") + torch.manual_seed(0) + a = torch.randn(M, K, dtype=torch.bfloat16, device="cuda") + b = torch.randn(N, K, dtype=torch.bfloat16, device="cuda") + return a, b + + +def test_rows_are_bitwise_stable_across_batch_composition(operands): + a, b = operands + full = batch_invariant_gemm(a, b) + for lo, hi in SLICES: + part = batch_invariant_gemm(a[lo:hi].contiguous(), b) + assert torch.equal(part, full[lo:hi]), ( + f"rows {lo}:{hi} changed when the batch was sliced" + ) + + +def test_matches_torch_reference(operands): + a, b = operands + got = batch_invariant_gemm(a, b) + ref = (a.float() @ b.float().T).to(torch.bfloat16) + torch.testing.assert_close(got, ref, rtol=0, atol=1.0) + + +def test_out_parameter_is_written(operands): + a, b = operands + out = torch.empty(M, N, dtype=torch.bfloat16, device="cuda") + assert batch_invariant_gemm(a, b, out=out) is out + assert torch.equal(out, batch_invariant_gemm(a, b)) + + +def test_unsupported_combinations_raise(): + if not torch.cuda.is_available(): + pytest.skip("batch_invariant_gemm requires a GPU") + a = torch.randn(8, 16, dtype=torch.float16, device="cuda") + b = torch.randn(8, 16, dtype=torch.float16, device="cuda") + assert not is_supported(a, b) + with pytest.raises(ValueError, match="bfloat16"): + batch_invariant_gemm(a, b) diff --git a/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py b/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py new file mode 100644 index 00000000000..f3199689980 --- /dev/null +++ b/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py @@ -0,0 +1,153 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Batch-invariant BF16 GEMM forward path. + +The general :func:`general_gemm` path delegates kernel selection to cuBLASLt, whose +heuristic depends on the full problem shape. That makes a given input row's output +depend on how many *other* rows are present in the batch. This module provides an +opt-in forward path with a reduction order that only depends on the tile +coordinates, so a row's result is bitwise stable across batch composition. + +Supported: BF16, ``Y = X @ W.T``, contiguous 2-D operands, no bias. +Unsupported combinations raise instead of silently falling back. +""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + +__all__ = ["batch_invariant_gemm", "is_supported"] + +# Tile geometry. These are fixed on purpose: making them depend on M would +# reintroduce the batch dependence this path exists to remove. +BLOCK_M = 64 +BLOCK_N = 64 +BLOCK_K = 64 + + +@triton.jit +def _bi_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N, + K, + stride_am, + stride_ak, + stride_bn, + stride_bk, + stride_cm, + stride_cn, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, +): + """C[m, n] = sum_k A[m, k] * B[n, k], accumulated in a single BLOCK_K loop.""" + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + offs_k = tl.arange(0, BLOCK_K) + + a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk + + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + # Fixed-trip sequential reduction: the order is a function of K and BLOCK_K + # only, so it cannot vary with M or with the tile's position in the batch. + for k in range(0, tl.cdiv(K, BLOCK_K)): + k_rem = K - k * BLOCK_K + a_mask = (offs_m[:, None] < M) & (offs_k[None, :] < k_rem) + b_mask = (offs_k[:, None] < k_rem) & (offs_n[None, :] < N) + a = tl.load(a_ptrs, mask=a_mask, other=0.0) + b = tl.load(b_ptrs, mask=b_mask, other=0.0) + acc += tl.dot(a, b) + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + + c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn + tl.store(c_ptrs, acc.to(c_ptr.dtype.element_ty), + mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + + +def is_supported( + a: torch.Tensor, + b: torch.Tensor, + out: Optional[torch.Tensor] = None, +) -> bool: + """Whether ``batch_invariant_gemm`` can run this configuration.""" + try: + _check(a, b, out) + except ValueError: + return False + return True + + +def _check(a: torch.Tensor, b: torch.Tensor, out: Optional[torch.Tensor]) -> None: + if a.dtype != torch.bfloat16 or b.dtype != torch.bfloat16: + raise ValueError( + f"batch_invariant_gemm supports bfloat16 inputs only, got {a.dtype} and {b.dtype}." + ) + if a.dim() != 2 or b.dim() != 2: + raise ValueError( + f"batch_invariant_gemm requires 2-D operands, got {a.dim()}-D and {b.dim()}-D." + ) + if not a.is_contiguous() or not b.is_contiguous(): + raise ValueError("batch_invariant_gemm requires contiguous operands.") + if a.shape[1] != b.shape[1]: + raise ValueError( + f"K mismatch: A has {a.shape[1]} columns, B has {b.shape[1]}." + ) + if out is not None and (out.dim() != 2 or not out.is_contiguous()): + raise ValueError("batch_invariant_gemm requires a contiguous 2-D out tensor.") + + +def batch_invariant_gemm( + a: torch.Tensor, + b: torch.Tensor, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Compute ``Y = A @ B.T`` with a batch-composition-independent reduction order. + + Parameters + ---------- + a : torch.Tensor + BF16 activation, shape ``[M, K]``, contiguous. Row ``m`` of ``a`` always + produces the same output bits for a given ``b``, independent of ``M`` and of + the row's position. + b : torch.Tensor + BF16 weight, shape ``[N, K]``, contiguous. + out : torch.Tensor, optional + BF16 destination of shape ``[M, N]``. A fresh tensor is allocated when absent. + + Returns + ------- + torch.Tensor + The ``[M, N]`` result. Same layout contract as :func:`general_gemm`'s output. + """ + _check(a, b, out) + m, k = a.shape + n = b.shape[0] + if out is None: + out = torch.empty((m, n), dtype=a.dtype, device=a.device) + elif out.shape != (m, n): + raise ValueError(f"out must have shape {(m, n)}, got {tuple(out.shape)}.") + if m == 0 or n == 0: + return out + + grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N)) + _bi_gemm_kernel[grid]( + a, b, out, + m, n, k, + a.stride(0), a.stride(1), + b.stride(0), b.stride(1), + out.stride(0), out.stride(1), + BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, + ) + return out From 8f47ffab08f9b46b88ffd751d10c0a21c6553828 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 21 Sep 2026 05:33:56 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci Signed-off-by: 0z5a --- tests/pytorch/test_batch_invariant_gemm.py | 4 +-- .../cpp_extensions/batch_invariant_gemm.py | 30 ++++++++++++------- 2 files changed, 20 insertions(+), 14 deletions(-) diff --git a/tests/pytorch/test_batch_invariant_gemm.py b/tests/pytorch/test_batch_invariant_gemm.py index 1eed3308790..0ae3c2ee217 100644 --- a/tests/pytorch/test_batch_invariant_gemm.py +++ b/tests/pytorch/test_batch_invariant_gemm.py @@ -36,9 +36,7 @@ def test_rows_are_bitwise_stable_across_batch_composition(operands): full = batch_invariant_gemm(a, b) for lo, hi in SLICES: part = batch_invariant_gemm(a[lo:hi].contiguous(), b) - assert torch.equal(part, full[lo:hi]), ( - f"rows {lo}:{hi} changed when the batch was sliced" - ) + assert torch.equal(part, full[lo:hi]), f"rows {lo}:{hi} changed when the batch was sliced" def test_matches_torch_reference(operands): diff --git a/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py b/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py index f3199689980..01b7553a099 100644 --- a/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/batch_invariant_gemm.py @@ -72,8 +72,9 @@ def _bi_gemm_kernel( b_ptrs += BLOCK_K * stride_bk c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn - tl.store(c_ptrs, acc.to(c_ptr.dtype.element_ty), - mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + tl.store( + c_ptrs, acc.to(c_ptr.dtype.element_ty), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N) + ) def is_supported( @@ -101,9 +102,7 @@ def _check(a: torch.Tensor, b: torch.Tensor, out: Optional[torch.Tensor]) -> Non if not a.is_contiguous() or not b.is_contiguous(): raise ValueError("batch_invariant_gemm requires contiguous operands.") if a.shape[1] != b.shape[1]: - raise ValueError( - f"K mismatch: A has {a.shape[1]} columns, B has {b.shape[1]}." - ) + raise ValueError(f"K mismatch: A has {a.shape[1]} columns, B has {b.shape[1]}.") if out is not None and (out.dim() != 2 or not out.is_contiguous()): raise ValueError("batch_invariant_gemm requires a contiguous 2-D out tensor.") @@ -143,11 +142,20 @@ def batch_invariant_gemm( grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N)) _bi_gemm_kernel[grid]( - a, b, out, - m, n, k, - a.stride(0), a.stride(1), - b.stride(0), b.stride(1), - out.stride(0), out.stride(1), - BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, + a, + b, + out, + m, + n, + k, + a.stride(0), + a.stride(1), + b.stride(0), + b.stride(1), + out.stride(0), + out.stride(1), + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + BLOCK_K=BLOCK_K, ) return out