code/day89-tool-map/triton_axpy.pyThis 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())