COURSE / SOURCE

triton_softmax_matmul.py

All lessons
Source filecode/day86-triton/triton_softmax_matmul.py

This 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())