claude-opus-4-1 / triton8f159a
claude-opus-4-1_triton_8f159a · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 123 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-triton-8f159a?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
43 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Show all 43 measurements ›Showing all 43 measurements ⌄
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:b6a8b41d4edc7ecc23d239c0e96ae879d037f8af72a29c73dbcd374f818ac079
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
acc += tl.dot(a, tl.trans(b), allow_tf32=True)tile-k = 32
BLOCK_K = 32tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128Kernel source
main.py123 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gemm_kernel(
a_ptr, b_ptr, c_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
# Program ID
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# Block starting positions
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Accumulator with float32 for better precision
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Base pointers for A and B blocks
a_base = a_ptr + rm[:, None] * stride_am
b_base = b_ptr + rn[:, None] * stride_bn
# Iterate over K dimension
for k in range(0, K, BLOCK_K):
rk = k + tl.arange(0, BLOCK_K)
# Load A block [BLOCK_M, BLOCK_K]
a_mask = (rm[:, None] < M) & (rk[None, :] < K)
a_offs = a_base + rk[None, :] * stride_ak
a = tl.load(a_offs, mask=a_mask, other=0.0).to(tl.float32)
# Load B block [BLOCK_N, BLOCK_K]
b_mask = (rn[:, None] < N) & (rk[None, :] < K)
b_offs = b_base + rk[None, :] * stride_bk
b = tl.load(b_offs, mask=b_mask, other=0.0).to(tl.float32)
# Compute dot product A @ B.T
# A is [BLOCK_M, BLOCK_K], B is [BLOCK_N, BLOCK_K]
# We need B transposed to [BLOCK_K, BLOCK_N]
acc += tl.dot(a, tl.trans(b), allow_tf32=True)
# Store result
c_mask = (rm[:, None] < M) & (rn[None, :] < N)
c_offs = c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn
c = acc.to(tl.float16)
tl.store(c_offs, c, mask=c_mask)
def run(A, B):
# Handle device management
original_device_a = A.device
original_device_b = B.device
# Move to GPU if needed
if A.device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU tensors are required")
A = A.cuda()
if B.device.type == 'cpu':
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU tensors are required")
B = B.cuda()
# Ensure tensors are on the same device
if A.device != B.device:
B = B.to(A.device)
# Convert to float16 if needed
if A.dtype != torch.float16:
A = A.to(torch.float16)
if B.dtype != torch.float16:
B = B.to(torch.float16)
# Get dimensions
M = A.shape[0]
K = A.shape[1]
N = B.shape[0]
# Validate dimensions
assert A.shape[1] == B.shape[1], f"K dimension mismatch: {A.shape[1]} != {B.shape[1]}"
assert B.shape[0] == 4096, f"N dimension must be 4096, got {B.shape[0]}"
assert B.shape[1] == 14336, f"K dimension must be 14336, got {B.shape[1]}"
# Allocate output
C = torch.empty((M, N), dtype=torch.float16, device=A.device)
# Block sizes optimized for B200 with large K dimension
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 32
# Grid dimensions
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']),
triton.cdiv(N, META['BLOCK_N']),
)
# Launch kernel
gemm_kernel[grid](
A, B, C,
M, N, K,
A.stride(0), A.stride(1),
B.stride(0), B.stride(1),
C.stride(0), C.stride(1),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
)
# Move result back to original device
if original_device_a.type == 'cpu':
C = C.cpu()
return Cscrolls · 123 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON