code/day86-triton/triton_softmax_matmul.pyThis is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.
# SPDX-License-Identifier: MIT
#
# Day 86: writing kernels in Triton.
#
# Two kernels the course has already written in CUDA, written again as
# Triton programs over blocks of data:
#
# part 1 a row softmax, the same 2048 x 1024 shape day 47's softmaxRows
# runs on (code/day47-fast-math/fast_math.cu).
# part 2 an FP32 matmul, the same arithmetic day 44 pinned on
# cublasGemmEx: FP32 in, FP32 accumulate, FP32 out. It runs twice,
# once with tl.dot's input_precision="ieee" and once with "tf32",
# because on compute capability 8.0 and up that one string is the
# difference between the SIMT multiplier and a tensor core, and
# the default is "tf32".
#
# Two modes, decided by the environment and nothing else:
#
# TRITON_INTERPRET=1 every kernel runs on the CPU through Triton's
# interpreter, on numpy, with no GPU anywhere. The
# correctness half runs; nothing is timed.
# unset the real thing, on a card at compute capability 8.0
# or newer. Correctness first, then the timings.
#
# Triton is Linux only: all twelve wheels of 3.8.0 are manylinux x86_64 or
# aarch64, there is no Windows wheel and no macOS wheel
# (https://pypi.org/pypi/triton/json , checked 2026-09-01). Its README's
# compatibility section reads "Supported Platforms: Linux" and "Supported
# Hardware: NVIDIA GPUs (Compute Capability 8.0+)", and the same file
# documents the fallback: "TRITON_INTERPRET=1 uses the Triton interpreter
# instead of running on the GPU. You can insert Python breakpoints in your
# kernel code!" (https://github.com/triton-lang/triton/blob/main/README.md
# , checked 2026-09-01).
#
# Install, run and artifacts: README.md in this directory.
#
# PARTIALLY VERIFIED: the Triton 3.7.1 interpreter path passed and the T4
# capability gate exited 1 as designed on 2026-09-02. The sm_80-or-newer
# GPU correctness and timing path remains blocked.
import os
import sys
import torch
import triton
import triton.language as tl
import triton.testing
# Part 1 keeps day 47's shape so the two kernels are read side by side.
SOFTMAX_ROWS = 2048
SOFTMAX_COLS = 1024
# Part 2's correctness case. K is deliberately the long axis: the tolerance
# below scales with the number of accumulated terms, and at K = 1024 the
# scaled bound is larger than the table's flat f32 rtol, so the scaling is
# doing work rather than decorating the output.
MM_M, MM_N, MM_K = 256, 256, 1024
# The timed case is square at 2048 because that is where day 44 measured
# its own FP32 ladder against cublasGemmEx on a Tesla T4. Same shape, same
# dtype, same accumulate type; a different card, which is why this program
# prints its own card and this page compares no times across the two.
TIME_SIZE = 2048
# EXERCISE-DESIGN.md section 3: p is the explicit mantissa bit count, eps is
# the spacing at 1.0, and any output that depends on K accumulated terms
# gets rtol_K = max(rtol_table, 4 * eps * sqrt(K)).
F32_EPS = 2.0**-23
F32_RTOL = 1e-5
TF32_EPS = 2.0**-10
TF32_RTOL = 1e-2
# Three shapes for part 2's sweep, each a (BLOCK_M, BLOCK_N, BLOCK_K,
# GROUP_M, num_warps, num_stages). This is what triton.autotune would pick
# between; it is written out instead so the run prints every row rather
# than only the winner, and so the interpreter, which cannot benchmark
# anything, can still compile and check one of them.
CONFIGS = [
(64, 64, 32, 8, 4, 3),
(128, 128, 32, 8, 8, 3),
(128, 64, 64, 8, 4, 4),
]
# snippet: softmax-kernel
@triton.jit
def softmax_kernel(in_ptr, out_ptr, row_stride, n_cols,
BLOCK_COLS: tl.constexpr):
# One program owns one row. Day 47's CUDA version hands one 256-thread
# block the same row, walks it three times, and reduces it twice
# through shared memory behind five written barriers. Here the row is
# one BLOCK_COLS-wide value and both reductions are one call each.
#
# BLOCK_COLS is a power of two at or above n_cols, so the tail lanes are
# masked. They load -inf, which loses the max and contributes exp(-inf)
# = 0 to the sum. A 0.0 pad would win the max on an all-negative row.
row = tl.program_id(0)
cols = tl.arange(0, BLOCK_COLS)
mask = cols < n_cols
x = tl.load(in_ptr + row * row_stride + cols, mask=mask,
other=-float("inf"))
num = tl.exp(x - tl.max(x, axis=0))
tl.store(out_ptr + row * row_stride + cols, num / tl.sum(num, axis=0),
mask=mask)
# end snippet
# snippet: matmul-signature
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, m, n, k, stride_am, stride_ak,
stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr,
IEEE: tl.constexpr):
# end snippet
# One program owns one BLOCK_M x BLOCK_N tile of C and walks K. The
# accumulator is FP32 either way; IEEE decides only what the multiplier
# does with FP32 inputs, which is the same choice day 44 pinned on
# cublasGemmEx with CUBLAS_COMPUTE_32F under CUBLAS_DEFAULT_MATH.
#
# The grouped program order is the one swizzle the compiler will not do
# for you: consecutive program ids walk GROUP_M tile rows before moving
# right, so tiles resident at the same time share more of A and B in L2.
# snippet: group-order
pid = tl.program_id(0)
tiles_m = tl.cdiv(m, BLOCK_M)
tiles_n = tl.cdiv(n, BLOCK_N)
per_group = GROUP_M * tiles_n
first_m = (pid // per_group) * GROUP_M
rows_here = min(tiles_m - first_m, GROUP_M)
pid_m = first_m + ((pid % per_group) % rows_here)
pid_n = (pid % per_group) // rows_here
# end snippet
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_k[:, None] * stride_bk + offs_n[None, :] * stride_bn
# snippet: dot-precision
# The accumulator is a value, not a buffer: no __shared__, no barrier,
# no double buffer. IEEE is a compile-time constant, so exactly one of
# these two branches survives into the compiled kernel.
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for step in range(0, tl.cdiv(k, BLOCK_K)):
left = k - step * BLOCK_K
a = tl.load(a_ptrs, mask=offs_k[None, :] < left, other=0.0)
b = tl.load(b_ptrs, mask=offs_k[:, None] < left, other=0.0)
if IEEE:
acc = tl.dot(a, b, acc, input_precision="ieee")
else:
acc = tl.dot(a, b, acc, input_precision="tf32")
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# end snippet
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc,
mask=(offs_m[:, None] < m) & (offs_n[None, :] < n))
def run_softmax(x, block_cols):
"""Launches one program per row and returns the output tensor."""
out = torch.empty_like(x)
softmax_kernel[(x.shape[0],)](x, out, x.stride(0), x.shape[1],
BLOCK_COLS=block_cols, num_warps=8)
return out
def run_matmul(a, b, config, ieee):
"""Launches the matmul with one config, returns the output tensor."""
m, k = a.shape
n = b.shape[1]
block_m, block_n, block_k, group_m, warps, stages = config
out = torch.empty((m, n), device=a.device, dtype=torch.float32)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n),)
matmul_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, GROUP_M=group_m, IEEE=ieee,
num_warps=warps, num_stages=stages)
return out
def scaled_tolerance(rtol_table, eps, terms, want):
"""EXERCISE-DESIGN.md section 3's length-dependent bound.
Returns (rtol, atol) and the sentence that explains where they came
from, so the report never prints a tolerance nobody can rederive.
"""
scaled = 4.0 * eps * (terms**0.5)
rtol = max(rtol_table, scaled)
atol = rtol * float(want.abs().max())
why = (f"reduction over K={terms}, rtol = max({rtol_table:g}, "
f"4*{eps:g}*sqrt({terms})) = {rtol:.3g}")
return rtol, atol, why
def report(name, got, want, rtol, atol, why):
"""Prints pass or fail plus the worst element. Returns (ok, error).
The comparison is EXERCISE-DESIGN.md's one rule, |got - exp| <= atol +
rtol * |exp|, and the element reported is the one furthest past it
rather than the first, because a wrong tile shows up as a band and the
first bad element in a band says less than the worst one.
The error returned is normalised by the largest reference magnitude,
not element by element. A dot product of signed terms lands near zero
at some elements, and an element-wise relative error there measures
cancellation instead of the multiplier this program is comparing.
"""
got64 = got.double()
excess = (got64 - want).abs() - (atol + rtol * want.abs())
err = float((got64 - want).abs().max()) / float(want.abs().max())
print(f" tolerance |got-exp| <= {atol:.3g} + {rtol:.3g}*|exp|")
print(f" {why}")
if float(excess.max()) > 0.0:
flat = int(excess.argmax())
i, j = flat // want.shape[1], flat % want.shape[1]
print(f" {name}: FAIL at ({i}, {j}): got "
f"{float(got64[i, j]):.9g}, want {float(want[i, j]):.9g}",
file=sys.stderr)
return False, err
print(f" {name}: pass, max error / max |reference| = {err:.3g}")
return True, err
def main() -> int:
interpret = os.environ.get("TRITON_INTERPRET") == "1"
print(f"triton {triton.__version__}, torch {torch.__version__}")
if interpret:
device = "cpu"
print("TRITON_INTERPRET=1: kernels run on the CPU through the "
"Triton interpreter.")
print("No GPU is involved and nothing below is timed.")
else:
if not torch.cuda.is_available():
print("no CUDA device; set TRITON_INTERPRET=1 to run the "
"correctness half on the CPU", file=sys.stderr)
return 1
device = "cuda"
major, minor = torch.cuda.get_device_capability(0)
name = torch.cuda.get_device_name(0)
print(f"GPU: {name} (compute capability {major}.{minor})")
# A real branch, not an assert: NDEBUG is not the issue here, but a
# gate that only fires in one build is the same bug in any language.
if major < 8:
print(f"Triton needs compute capability 8.0 or newer; {name} "
f"is {major}.{minor}. Set TRITON_INTERPRET=1 for the "
f"correctness half on the CPU.", file=sys.stderr)
return 1
ok = True
torch.manual_seed(20260901)
print(f"\nPart 1: softmax over {SOFTMAX_ROWS} rows of {SOFTMAX_COLS}")
# Inputs are generated on the host and copied, so the interpreter run
# and the GPU run see the same bits and their reports are comparable.
# Rows carry a growing offset, up to 102.4 by the last one, so a kernel
# that forgets to subtract the row maximum overflows float32 instead of
# quietly agreeing.
h_x = torch.randn((SOFTMAX_ROWS, SOFTMAX_COLS), dtype=torch.float32)
h_x += torch.arange(SOFTMAX_ROWS,
dtype=torch.float32).reshape(-1, 1) * 0.05
x = h_x.to(device)
block_cols = triton.next_power_of_2(SOFTMAX_COLS)
got = run_softmax(x, block_cols)
# Reference in float64 on the host, the same three operations the
# kernel performs, so the comparison isolates ordering and precision
# rather than a different algorithm.
x64 = h_x.double()
e = torch.exp(x64 - x64.max(dim=1, keepdim=True).values)
want = e / e.sum(dim=1, keepdim=True)
rtol, atol, why = scaled_tolerance(F32_RTOL, F32_EPS, SOFTMAX_COLS, want)
passed, _ = report("softmax", got.cpu(), want, rtol, atol, why)
ok = ok and passed
print(f"\nPart 2: matmul {MM_M} x {MM_K} x {MM_N}, FP32 in, FP32 "
f"accumulate, FP32 out")
h_a = torch.randn((MM_M, MM_K), dtype=torch.float32)
h_b = torch.randn((MM_K, MM_N), dtype=torch.float32)
a, b = h_a.to(device), h_b.to(device)
want_mm = h_a.double() @ h_b.double()
config = CONFIGS[0]
print(f" config BLOCK_M/N/K {config[0]}/{config[1]}/{config[2]}, "
f"GROUP_M {config[3]}, num_warps {config[4]}, "
f"num_stages {config[5]}")
rtol, atol, why = scaled_tolerance(F32_RTOL, F32_EPS, MM_K, want_mm)
got_ieee = run_matmul(a, b, config, ieee=True).cpu()
passed, rel_ieee = report('matmul input_precision="ieee"', got_ieee,
want_mm, rtol, atol, why)
ok = ok and passed
rtol, atol, why = scaled_tolerance(TF32_RTOL, TF32_EPS, MM_K, want_mm)
got_tf32 = run_matmul(a, b, config, ieee=False).cpu()
passed, rel_tf32 = report('matmul input_precision="tf32"', got_tf32,
want_mm, rtol, atol, why)
ok = ok and passed
# tf32 keeps 10 explicit mantissa bits against f32's 23, so on hardware
# the truncated multiplier has to be measurably worse than the exact
# one. The interpreter has no tensor core and no documented tf32
# rounding, so this is reported there and gated only on the GPU.
print(f" max error / max |reference|: ieee {rel_ieee:.3g}, tf32 "
f"{rel_tf32:.3g}, ratio {rel_tf32 / max(rel_ieee, 1e-30):.3g}")
if interpret:
print(" not gated under the interpreter: whether it honours "
"input_precision at all is what this run finds out.")
elif rel_tf32 <= rel_ieee:
print(" FAIL: tf32 did not lose accuracy against ieee, so one of "
"the two runs did not use the multiplier it asked for.",
file=sys.stderr)
ok = False
if interpret:
print("\nNo timings. The interpreter evaluates each program "
"sequentially on numpy, so a millisecond figure from it "
"would describe Python, not a GPU.")
return 0 if ok else 1
print(f"\nPart 3: {TIME_SIZE} cubed, mean ms over do_bench's default "
f"25 ms warm-up and 100 ms of repetitions")
big_a = torch.randn((TIME_SIZE, TIME_SIZE), device=device,
dtype=torch.float32)
big_b = torch.randn((TIME_SIZE, TIME_SIZE), device=device,
dtype=torch.float32)
flop = 2.0 * TIME_SIZE**3
print(" precision BLOCK_M/N/K GROUP_M warps stages "
" ms GFLOP/s")
best = {}
for ieee in (True, False):
label = "ieee" if ieee else "tf32"
for cfg in CONFIGS:
ms = triton.testing.do_bench(
lambda: run_matmul(big_a, big_b, cfg, ieee))
best[label] = min(best.get(label, ms), ms)
print(f" {label:>9} {cfg[0]:>3}/{cfg[1]:>3}/{cfg[2]:<3} "
f"{cfg[3]:>7} {cfg[4]:>5} {cfg[5]:>6} "
f"{ms:>8.3f} {flop / ms / 1.0e6:>8.1f}")
speedup = best["ieee"] / best["tf32"]
print(f"\n best ieee {best['ieee']:.3f} ms, best tf32 "
f"{best['tf32']:.3f} ms, tf32 is {speedup:.2f}x")
# The only timing this program is willing to gate on. It compares one
# kernel against itself with one string changed, on one card, so it
# says nothing about Triton against CUDA and everything about what
# that string buys at compute capability 8.0 and up.
if speedup <= 1.0:
print(" FAIL: tf32 was not faster than ieee, so the tensor-core "
"path was not taken.", file=sys.stderr)
ok = False
if not ok:
return 1
print("\nPASS")
return 0
if __name__ == "__main__":
sys.exit(main())