COURSE / SOURCE

triton_axpy.py

All lessons
Source filecode/day89-tool-map/triton_axpy.py

This is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.

#!/usr/bin/env python3
"""The same fused axpy-and-clamp as fused_axpy.cu, in Triton.

Same decomposition on purpose: one program per tile of 1024 elements, a
mask on the tail, the same alpha, the same 1,000,003 elements. What
differs is the column the decision map calls "what you still own". This
file never names a lane, a warp or a block size in threads. It names a
tile, and the compiler decides how many threads carry it. The README has
the grep that settles that claim against both files; run it rather than
trusting the sentence.

Versions this file is written against, read off PyPI on 2026-09-01
(https://pypi.org/pypi/triton/json and https://pypi.org/pypi/torch/json):

    python 3.12, torch==2.13.0 (which resolves triton==3.7.1 on Linux)

Triton publishes Linux wheels only; the 3.8.0 file list on PyPI is
manylinux aarch64 and x86_64 and nothing else, which is the machine-
readable form of the README's "Supported Platforms: Linux"
(https://github.com/triton-lang/triton/blob/main/README.md#compatibility
, checked 2026-09-01). The same section says "NVIDIA GPUs (Compute
Capability 8.0+)".

PARTIALLY VERIFIED on 2026-09-02. The Triton 3.7.1 interpreter path passed
both gates. The verification node is a Tesla T4 at compute capability 7.5,
one generation below the floor Triton documents, so the GPU path remains
blocked and the README says on what. The free correctness path is:

    TRITON_INTERPRET=1 python3 triton_axpy.py

The interpreter transcript and environment are retained in `evidence/`.
No Triton GPU result is claimed.

This file times nothing. Day 89 makes no tool-versus-tool performance
claim, because no comparison this project could verify exists after CUDA
13.3 (research/FACT-SHEET.md section 8).

Exit codes: 0 gates pass, 1 a gate failed, 2 the environment is missing.
"""

import os
import sys

N = 1000003
BLOCK = 1024
ALPHA = 2.5

# Depth 1, one multiply and one add per element, so this does not scale
# with N. Four float epsilons plus an absolute floor, matching
# fused_axpy.cu line for line: the floor is there because the clamp
# writes exact zeros and a relative tolerance says nothing at zero.
RTOL = 4.0 * 2.0**-23
ATOL = 1e-6

try:
    import torch
    import triton
    import triton.language as tl
except ImportError as exc:  # pragma: no cover - environment report
    print(f"missing dependency: {exc}", file=sys.stderr)
    print("pip install torch==2.13.0", file=sys.stderr)
    sys.exit(2)


# snippet: triton-kernel
@triton.jit
def fused_axpy_clamp(x_ptr, y_ptr, alpha, n_elements,
                     BLOCK_SIZE: tl.constexpr):
    pid = tl.program_id(axis=0)
    offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    mask = offsets < n_elements
    x = tl.load(x_ptr + offsets, mask=mask)
    y = tl.load(y_ptr + offsets, mask=mask)
    v = alpha * x + y
    tl.store(y_ptr + offsets, tl.maximum(v, 0.0), mask=mask)
# end snippet: triton-kernel


def build_inputs(device):
    """Same two sequences as fused_axpy.cu's inputX and inputY."""
    i = torch.arange(N, dtype=torch.float32, device=device)
    x = 1.0 / (torch.remainder(i, 13.0) + 1.0)
    # Keep the test away from the host-separate versus device-FMA boundary.
    y = -0.51 + torch.remainder(i, 17.0) / 17.0
    return x, y


def main():
    interpret = os.environ.get("TRITON_INTERPRET") == "1"
    device = "cpu" if interpret else "cuda"
    if not interpret and not torch.cuda.is_available():
        print("no CUDA device, and TRITON_INTERPRET is not 1", file=sys.stderr)
        return 2

    print(f"triton {triton.__version__}, torch {torch.__version__}")
    if interpret:
        print("TRITON_INTERPRET=1: the kernel runs on the CPU, no GPU used")
    else:
        cap = torch.cuda.get_device_capability()
        print(f"device: {torch.cuda.get_device_name()} "
              f"(compute capability {cap[0]}.{cap[1]})")
        print("Triton documents compute capability 8.0 and above")

    x, y = build_inputs(device)
    want = torch.clamp(ALPHA * x + y, min=0.0)
    clamped_in_reference = int((ALPHA * x + y <= 0.0).sum())

    got = y.clone()
    grid = (triton.cdiv(N, BLOCK),)
    fused_axpy_clamp[grid](x, got, ALPHA, N, BLOCK_SIZE=BLOCK)

    close = torch.isclose(got, want, rtol=RTOL, atol=ATOL)
    mismatches = int((~close).sum())
    clamped = int((got == 0.0).sum())

    failures = 0
    # Gate 1: agrees with the reference elementwise.
    if mismatches:
        first = int((~close).nonzero()[0])
        print(f"gate 1 FAIL: {mismatches} mismatches, first at {first}: "
              f"triton {got[first]!r} vs reference {want[first]!r}",
              file=sys.stderr)
        failures += 1
    # Gate 2: the clamp was exercised, and by the same elements. A relu
    # test on inputs that are never negative passes without testing the
    # branch.
    if clamped_in_reference == 0 or clamped != clamped_in_reference:
        print(f"gate 2 FAIL: clamped {clamped}, reference "
              f"{clamped_in_reference} (must be equal and nonzero)",
              file=sys.stderr)
        failures += 1

    if failures:
        return 1
    print(f"all gates pass: {N} elements, 0 mismatches, {clamped} clamped")
    return 0


if __name__ == "__main__":
    sys.exit(main())