#!/usr/bin/env python3 """One operator expressed in PyTorch and Triton with a correctness receipt.""" from __future__ import annotations import json import torch import triton import triton.language as tl @triton.jit def add_kernel(x, y, output, n_elements: tl.constexpr, BLOCK: tl.constexpr): offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) mask = offsets < n_elements tl.store(output + offsets, tl.load(x + offsets, mask=mask) + tl.load(y + offsets, mask=mask), mask=mask) def triton_add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: output = torch.empty_like(x) grid = (triton.cdiv(x.numel(), 256),) add_kernel[grid](x, y, output, x.numel(), BLOCK=256) return output def main() -> None: if not torch.cuda.is_available(): raise SystemExit("CUDA GPU required") x = torch.randn(1 << 20, device="cuda") y = torch.randn_like(x) reference = x + y candidate = triton_add(x, y) print( json.dumps( { "torch": torch.__version__, "triton": triton.__version__, "max_abs_error": float((reference - candidate).abs().max()), "matches": bool(torch.allclose(reference, candidate)), "receipt_scope": "toy_operator_not_glm_5_2", }, indent=2, ) ) if __name__ == "__main__": main()