submission 426455
Cookie 🍪 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 272 lines, June 9 Researcher Reciprocity License v1.0.
trimul_opus_real.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-426455?include=source"interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:a0710768fa9aaeea2b29faf25b8df6e013f299a2b1ae9f57881ce3674a244a61
license declaredunknown
license concludedunknown
authorsCookie 🍪
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
acc += tl.dot(x0, tl.trans(x1))tile-k = 16
BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16Kernel source
trimul_opus_real.py272 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _layernorm_kernel(
X_ptr, Y_ptr, W_ptr, B_ptr,
stride_x, C, eps,
BLOCK_SIZE: tl.constexpr,
):
row_idx = tl.program_id(0)
row_start_ptr = X_ptr + row_idx * stride_x
col_offsets = tl.arange(0, BLOCK_SIZE)
mask = col_offsets < C
x = tl.load(row_start_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
x_sum = tl.sum(x, axis=0)
mean = x_sum / C
x_centered = tl.where(mask, x - mean, 0.0)
var_sum = tl.sum(x_centered * x_centered, axis=0)
var = var_sum / C
rstd = 1.0 / tl.sqrt(var + eps)
x_norm = x_centered * rstd
w = tl.load(W_ptr + col_offsets, mask=mask, other=1.0).to(tl.float32)
b = tl.load(B_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)
y = x_norm * w + b
out_ptr = Y_ptr + row_idx * stride_x
tl.store(out_ptr + col_offsets, y, mask=mask)
@triton.jit
def _projection_gating_kernel(
x_ptr, mask_ptr,
left_proj_weight_ptr, right_proj_weight_ptr,
left_gate_weight_ptr, right_gate_weight_ptr, out_gate_weight_ptr,
left_out_ptr, right_out_ptr, out_gate_ptr,
B, N, C, H,
BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
pid = tl.program_id(0)
b = pid // (N * N)
remainder = pid % (N * N)
i = remainder // N
j = remainder % N
x_base = b * (N * N * C) + i * (N * C) + j * C
mask_idx = b * (N * N) + i * N + j
mask_val = tl.load(mask_ptr + mask_idx).to(tl.float32)
out_base = b * (N * N * H) + i * (N * H) + j * H
for h_start in range(0, H, BLOCK_H):
h_offs = h_start + tl.arange(0, BLOCK_H)
h_mask = h_offs < H
left_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
right_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
left_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
right_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
out_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)
for c_start in range(0, C, BLOCK_C):
c_offs = c_start + tl.arange(0, BLOCK_C)
c_mask = c_offs < C
x_vals = tl.load(x_ptr + x_base + c_offs, mask=c_mask, other=0.0).to(tl.float32)
weight_offsets = h_offs[:, None] * C + c_offs[None, :]
combined_mask = h_mask[:, None] & c_mask[None, :]
left_proj_w = tl.load(left_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
right_proj_w = tl.load(right_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
left_gate_w = tl.load(left_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
right_gate_w = tl.load(right_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
out_gate_w = tl.load(out_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)
left_proj_acc += tl.sum(left_proj_w * x_vals[None, :], axis=1)
right_proj_acc += tl.sum(right_proj_w * x_vals[None, :], axis=1)
left_gate_acc += tl.sum(left_gate_w * x_vals[None, :], axis=1)
right_gate_acc += tl.sum(right_gate_w * x_vals[None, :], axis=1)
out_gate_acc += tl.sum(out_gate_w * x_vals[None, :], axis=1)
left_gate_sig = tl.sigmoid(left_gate_acc)
right_gate_sig = tl.sigmoid(right_gate_acc)
out_gate_sig = tl.sigmoid(out_gate_acc)
left_result = left_proj_acc * mask_val * left_gate_sig
right_result = right_proj_acc * mask_val * right_gate_sig
out_offs = out_base + h_offs
tl.store(left_out_ptr + out_offs, left_result, mask=h_mask)
tl.store(right_out_ptr + out_offs, right_result, mask=h_mask)
tl.store(out_gate_ptr + out_offs, out_gate_sig, mask=h_mask)
@triton.jit
def _triangular_mul_kernel(
x0_ptr, x1_ptr, out_ptr,
B, N, H,
BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr,
):
pid_b = tl.program_id(0)
pid_ij = tl.program_id(1)
pid_d = tl.program_id(2)
num_tiles_j = tl.cdiv(N, BLOCK_J)
pid_i = pid_ij // num_tiles_j
pid_j = pid_ij % num_tiles_j
offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
mask_i = offs_i < N
mask_j = offs_j < N
acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
stride_b = N * N * H
stride_i = N * H
stride_k = H
base_b = pid_b * stride_b
base_d = pid_d
for k_start in range(0, N, BLOCK_K):
offs_k = k_start + tl.arange(0, BLOCK_K)
mask_k = offs_k < N
x0_offsets = base_b + offs_i[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
x0 = tl.load(x0_ptr + x0_offsets, mask=mask_i[:, None] & mask_k[None, :], other=0.0).to(tl.float32)
x1_offsets = base_b + offs_j[:, None] * stride_i + offs_k[None, :] * stride_k + base_d
x1 = tl.load(x1_ptr + x1_offsets, mask=mask_j[:, None] & mask_k[None, :], other=0.0).to(tl.float32)
acc += tl.dot(x0, tl.trans(x1))
out_offsets = base_b + offs_i[:, None] * stride_i + offs_j[None, :] * stride_k + base_d
tl.store(out_ptr + out_offsets, acc, mask=mask_i[:, None] & mask_j[None, :])
@triton.jit
def _output_norm_gate_proj_kernel(
x0_ptr, x1_ptr, ln_weight_ptr, ln_bias_ptr, linear_weight_ptr, out_ptr,
num_rows, H: tl.constexpr, C, eps,
BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,
):
row_idx = tl.program_id(0)
if row_idx >= num_rows:
return
offs_h = tl.arange(0, BLOCK_H)
mask_h = offs_h < H
x0 = tl.load(x0_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
sum_x = tl.sum(x0, axis=0)
mean = sum_x / H
x0_centered = x0 - mean
sum_sq = tl.sum(x0_centered * x0_centered, axis=0)
var = sum_sq / H
rstd = tl.rsqrt(var + eps)
x_norm = x0_centered * rstd
ln_w = tl.load(ln_weight_ptr + offs_h, mask=mask_h, other=1.0).to(tl.float32)
ln_b = tl.load(ln_bias_ptr + offs_h, mask=mask_h, other=0.0).to(tl.float32)
x_ln = x_norm * ln_w + ln_b
x1 = tl.load(x1_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)
x_gated = x_ln * x1
x_gated = tl.where(mask_h, x_gated, 0.0)
for c_start in range(0, C, BLOCK_C):
offs_c = c_start + tl.arange(0, BLOCK_C)
mask_c = offs_c < C
weight_ptrs = linear_weight_ptr + offs_c[:, None] * H + offs_h[None, :]
weights = tl.load(weight_ptrs, mask=mask_c[:, None] & mask_h[None, :], other=0.0).to(tl.float32)
acc = tl.sum(weights * x_gated[None, :], axis=1)
tl.store(out_ptr + row_idx * C + offs_c, acc, mask=mask_c)
def kernel_function(input_tensor, mask, weights, config):
B, N, _, C = input_tensor.shape
H = config["hidden_dim"]
# Stage 1: LayerNorm
x_norm = torch.empty_like(input_tensor)
n_rows = B * N * N
BLOCK_SIZE = triton.next_power_of_2(C)
BLOCK_SIZE = min(BLOCK_SIZE, 1024)
_layernorm_kernel[(n_rows,)](
input_tensor, x_norm, weights['norm.weight'], weights['norm.bias'],
C, C, 1e-5, BLOCK_SIZE=BLOCK_SIZE
)
# Stage 2: Projection + Gating
left = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
right = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
out_gate = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
BLOCK_H = min(32, H)
BLOCK_C = min(32, C)
_projection_gating_kernel[(B * N * N,)](
x_norm, mask,
weights['left_proj.weight'], weights['right_proj.weight'],
weights['left_gate.weight'], weights['right_gate.weight'], weights['out_gate.weight'],
left, right, out_gate,
B, N, C, H, BLOCK_H, BLOCK_C
)
# Stage 3: Triangular multiplication
tri_out = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')
BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16
num_tiles_i = triton.cdiv(N, BLOCK_I)
num_tiles_j = triton.cdiv(N, BLOCK_J)
_triangular_mul_kernel[(B, num_tiles_i * num_tiles_j, H)](
left, right, tri_out, B, N, H, BLOCK_I, BLOCK_J, BLOCK_K
)
# Stage 4: Output LayerNorm + Gate + Linear
output = torch.empty((B, N, N, C), dtype=torch.float32, device='cuda')
BLOCK_H_out = triton.next_power_of_2(H)
BLOCK_C_out = min(128, triton.next_power_of_2(C))
_output_norm_gate_proj_kernel[(n_rows,)](
tri_out.reshape(-1, H), out_gate.reshape(-1, H),
weights['to_out_norm.weight'], weights['to_out_norm.bias'], weights['to_out.weight'],
output.reshape(-1, C), n_rows, H, C, 1e-5, BLOCK_H_out, BLOCK_C_out
)
return output
def test_kernel():
torch.manual_seed(42)
B, N, C, H = 1, 32, 128, 128
input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)
mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)
weights = {
'norm.weight': torch.randn(C, device='cuda'), 'norm.bias': torch.randn(C, device='cuda'),
'left_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
'right_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),
'left_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
'right_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
'out_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),
'to_out_norm.weight': torch.randn(H, device='cuda'), 'to_out_norm.bias': torch.randn(H, device='cuda'),
'to_out.weight': torch.randn(C, H, device='cuda') / (C**0.5),
}
config = {"dim": C, "hidden_dim": H}
out_triton = kernel_function(input_tensor, mask, weights, config)
# Reference
from torch import nn, einsum
x = torch.nn.functional.layer_norm(input_tensor, [C], weights['norm.weight'], weights['norm.bias'])
left = x @ weights['left_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['left_gate.weight'].T)
right = x @ weights['right_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['right_gate.weight'].T)
out_gate = torch.sigmoid(x @ weights['out_gate.weight'].T)
tri = einsum('bikd,bjkd->bijd', left, right)
tri_norm = torch.nn.functional.layer_norm(tri, [H], weights['to_out_norm.weight'], weights['to_out_norm.bias'])
ref = (tri_norm * out_gate) @ weights['to_out.weight'].T
if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):
print("PASS")
else:
print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")
if __name__ == "__main__":
test_kernel()
def custom_kernel(input):
return kernel_function(*input)
scrolls · 272 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 416249.
- from typing import Dict, Tuple, TypeVar-import torchimport tritonimport triton.language as tl⋯ 1 unchanged lines@triton.jitdef _layernorm_kernel(- x_ptr,- out_ptr,- gamma_ptr,- beta_ptr,- N,- D,- stride_n,- stride_d,- eps: tl.constexpr,- BLOCK_D: tl.constexpr,+ X_ptr, Y_ptr, W_ptr, B_ptr,+ stride_x, C, eps,+ BLOCK_SIZE: tl.constexpr,):- pid = tl.program_id(0)- offs_d = tl.arange(0, BLOCK_D)- mask = offs_d < D+ row_idx = tl.program_id(0)+ row_start_ptr = X_ptr + row_idx * stride_x+ col_offsets = tl.arange(0, BLOCK_SIZE)+ mask = col_offsets < C- x_ptrs = x_ptr + pid * stride_n + offs_d * stride_d- x = tl.load(x_ptrs, mask=mask, other=0.0)-- mean = tl.sum(x, axis=0) / D- x_centered = x - mean- var = tl.sum(x_centered * x_centered, axis=0) / D+ x = tl.load(row_start_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)+ x_sum = tl.sum(x, axis=0)+ mean = x_sum / C+ x_centered = tl.where(mask, x - mean, 0.0)+ var_sum = tl.sum(x_centered * x_centered, axis=0)+ var = var_sum / Crstd = 1.0 / tl.sqrt(var + eps)+ x_norm = x_centered * rstd- gamma = tl.load(gamma_ptr + offs_d, mask=mask, other=1.0)- beta = tl.load(beta_ptr + offs_d, mask=mask, other=0.0)+ w = tl.load(W_ptr + col_offsets, mask=mask, other=1.0).to(tl.float32)+ b = tl.load(B_ptr + col_offsets, mask=mask, other=0.0).to(tl.float32)+ y = x_norm * w + b- out = x_centered * rstd * gamma + beta- out_ptrs = out_ptr + pid * stride_n + offs_d * stride_d- tl.store(out_ptrs, out, mask=mask)+ out_ptr = Y_ptr + row_idx * stride_x+ tl.store(out_ptr + col_offsets, y, mask=mask)@triton.jit- def _linear_sigmoid_kernel(- x_ptr,- w_ptr,- out_ptr,- M,- K,- N,- stride_xm,- stride_xk,- stride_wn,- stride_wk,- stride_om,- stride_on,- apply_sigmoid: tl.constexpr,- BLOCK_M: tl.constexpr,- BLOCK_N: tl.constexpr,- BLOCK_K: tl.constexpr,+ def _projection_gating_kernel(+ x_ptr, mask_ptr,+ left_proj_weight_ptr, right_proj_weight_ptr,+ left_gate_weight_ptr, right_gate_weight_ptr, out_gate_weight_ptr,+ left_out_ptr, right_out_ptr, out_gate_ptr,+ B, N, C, H,+ BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,):- pid_m = tl.program_id(0)- pid_n = tl.program_id(1)+ pid = tl.program_id(0)+ b = pid // (N * N)+ remainder = pid % (N * N)+ i = remainder // N+ j = remainder % N- 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)+ x_base = b * (N * N * C) + i * (N * C) + j * C+ mask_idx = b * (N * N) + i * N + j+ mask_val = tl.load(mask_ptr + mask_idx).to(tl.float32)+ out_base = b * (N * N * H) + i * (N * H) + j * H- acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ for h_start in range(0, H, BLOCK_H):+ h_offs = h_start + tl.arange(0, BLOCK_H)+ h_mask = h_offs < H- for k_start in range(0, K, BLOCK_K):- k_offs = k_start + offs_k- k_mask = k_offs < K+ left_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)+ right_proj_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)+ left_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)+ right_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)+ out_gate_acc = tl.zeros((BLOCK_H,), dtype=tl.float32)- x_ptrs = x_ptr + offs_m[:, None] * stride_xm + k_offs[None, :] * stride_xk- w_ptrs = w_ptr + offs_n[:, None] * stride_wn + k_offs[None, :] * stride_wk+ for c_start in range(0, C, BLOCK_C):+ c_offs = c_start + tl.arange(0, BLOCK_C)+ c_mask = c_offs < C+ x_vals = tl.load(x_ptr + x_base + c_offs, mask=c_mask, other=0.0).to(tl.float32)- x_mask = (offs_m[:, None] < M) & k_mask[None, :]- w_mask = (offs_n[:, None] < N) & k_mask[None, :]+ weight_offsets = h_offs[:, None] * C + c_offs[None, :]+ combined_mask = h_mask[:, None] & c_mask[None, :]- x_vals = tl.load(x_ptrs, mask=x_mask, other=0.0)- w_vals = tl.load(w_ptrs, mask=w_mask, other=0.0)+ left_proj_w = tl.load(left_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)+ right_proj_w = tl.load(right_proj_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)+ left_gate_w = tl.load(left_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)+ right_gate_w = tl.load(right_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)+ out_gate_w = tl.load(out_gate_weight_ptr + weight_offsets, mask=combined_mask, other=0.0).to(tl.float32)- acc += tl.dot(x_vals, tl.trans(w_vals))+ left_proj_acc += tl.sum(left_proj_w * x_vals[None, :], axis=1)+ right_proj_acc += tl.sum(right_proj_w * x_vals[None, :], axis=1)+ left_gate_acc += tl.sum(left_gate_w * x_vals[None, :], axis=1)+ right_gate_acc += tl.sum(right_gate_w * x_vals[None, :], axis=1)+ out_gate_acc += tl.sum(out_gate_w * x_vals[None, :], axis=1)- if apply_sigmoid:- acc = tl.sigmoid(acc)+ left_gate_sig = tl.sigmoid(left_gate_acc)+ right_gate_sig = tl.sigmoid(right_gate_acc)+ out_gate_sig = tl.sigmoid(out_gate_acc)- out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on- out_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)- tl.store(out_ptrs, acc, mask=out_mask)+ left_result = left_proj_acc * mask_val * left_gate_sig+ right_result = right_proj_acc * mask_val * right_gate_sig+ out_offs = out_base + h_offs+ tl.store(left_out_ptr + out_offs, left_result, mask=h_mask)+ tl.store(right_out_ptr + out_offs, right_result, mask=h_mask)+ tl.store(out_gate_ptr + out_offs, out_gate_sig, mask=h_mask)+@triton.jit- def _masked_gating_einsum_kernel(- left_ptr,- right_ptr,- left_gate_ptr,- right_gate_ptr,- mask_ptr,- out_ptr,- B,- seq_len,- hidden_dim,- stride_lb,- stride_li,- stride_lk,- stride_ld,- stride_rb,- stride_ri,- stride_rk,- stride_rd,- stride_lgb,- stride_lgi,- stride_lgk,- stride_lgd,- stride_rgb,- stride_rgi,- stride_rgk,- stride_rgd,- stride_mb,- stride_mi,- stride_mk,- stride_ob,- stride_oi,- stride_oj,- stride_od,- BLOCK_D: tl.constexpr,+ def _triangular_mul_kernel(+ x0_ptr, x1_ptr, out_ptr,+ B, N, H,+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr,):- pid_bi = tl.program_id(0)- pid_j = tl.program_id(1)+ pid_b = tl.program_id(0)+ pid_ij = tl.program_id(1)pid_d = tl.program_id(2)- pid_b = pid_bi // seq_len- pid_i = pid_bi % seq_len+ num_tiles_j = tl.cdiv(N, BLOCK_J)+ pid_i = pid_ij // num_tiles_j+ pid_j = pid_ij % num_tiles_j- d_offs = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)- d_mask = d_offs < hidden_dim+ offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)+ offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)+ mask_i = offs_i < N+ mask_j = offs_j < N- acc = tl.zeros((BLOCK_D,), dtype=tl.float32)+ acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)+ stride_b = N * N * H+ stride_i = N * H+ stride_k = H+ base_b = pid_b * stride_b+ base_d = pid_d- for k in range(seq_len):- mask_l = tl.load(- mask_ptr + pid_b * stride_mb + pid_i * stride_mi + k * stride_mk- )- mask_r = tl.load(- mask_ptr + pid_b * stride_mb + pid_j * stride_mi + k * stride_mk- )+ for k_start in range(0, N, BLOCK_K):+ offs_k = k_start + tl.arange(0, BLOCK_K)+ mask_k = offs_k < N- left_off = (- pid_b * stride_lb + pid_i * stride_li + k * stride_lk + d_offs * stride_ld- )- left_val = tl.load(left_ptr + left_off, mask=d_mask, other=0.0)- lg_off = (- pid_b * stride_lgb- + pid_i * stride_lgi- + k * stride_lgk- + d_offs * stride_lgd- )- lg_val = tl.load(left_gate_ptr + lg_off, mask=d_mask, other=0.0)- left_gated = left_val * mask_l * lg_val+ x0_offsets = base_b + offs_i[:, None] * stride_i + offs_k[None, :] * stride_k + base_d+ x0 = tl.load(x0_ptr + x0_offsets, mask=mask_i[:, None] & mask_k[None, :], other=0.0).to(tl.float32)- right_off = (- pid_b * stride_rb + pid_j * stride_ri + k * stride_rk + d_offs * stride_rd- )- right_val = tl.load(right_ptr + right_off, mask=d_mask, other=0.0)- rg_off = (- pid_b * stride_rgb- + pid_j * stride_rgi- + k * stride_rgk- + d_offs * stride_rgd- )- rg_val = tl.load(right_gate_ptr + rg_off, mask=d_mask, other=0.0)- right_gated = right_val * mask_r * rg_val+ x1_offsets = base_b + offs_j[:, None] * stride_i + offs_k[None, :] * stride_k + base_d+ x1 = tl.load(x1_ptr + x1_offsets, mask=mask_j[:, None] & mask_k[None, :], other=0.0).to(tl.float32)- acc += left_gated * right_gated+ acc += tl.dot(x0, tl.trans(x1))- out_off = (- pid_b * stride_ob + pid_i * stride_oi + pid_j * stride_oj + d_offs * stride_od- )- tl.store(out_ptr + out_off, acc, mask=d_mask)+ out_offsets = base_b + offs_i[:, None] * stride_i + offs_j[None, :] * stride_k + base_d+ tl.store(out_ptr + out_offsets, acc, mask=mask_i[:, None] & mask_j[None, :])@triton.jit- def _layernorm_gate_linear_kernel(- x_ptr,- gate_ptr,- gamma_ptr,- beta_ptr,- w_ptr,- out_ptr,- N,- hidden_dim,- out_dim,- stride_xn,- stride_xd,- stride_gn,- stride_gd,- stride_wout,- stride_whid,- stride_on,- stride_od,- eps: tl.constexpr,- BLOCK_HD: tl.constexpr,- BLOCK_OUT: tl.constexpr,+ def _output_norm_gate_proj_kernel(+ x0_ptr, x1_ptr, ln_weight_ptr, ln_bias_ptr, linear_weight_ptr, out_ptr,+ num_rows, H: tl.constexpr, C, eps,+ BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr,):- pid_n = tl.program_id(0)- pid_out = tl.program_id(1)+ row_idx = tl.program_id(0)+ if row_idx >= num_rows:+ return- offs_hd = tl.arange(0, BLOCK_HD)- hd_mask = offs_hd < hidden_dim+ offs_h = tl.arange(0, BLOCK_H)+ mask_h = offs_h < H- x_ptrs = x_ptr + pid_n * stride_xn + offs_hd * stride_xd- x = tl.load(x_ptrs, mask=hd_mask, other=0.0)+ x0 = tl.load(x0_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)+ sum_x = tl.sum(x0, axis=0)+ mean = sum_x / H+ x0_centered = x0 - mean+ sum_sq = tl.sum(x0_centered * x0_centered, axis=0)+ var = sum_sq / H+ rstd = tl.rsqrt(var + eps)+ x_norm = x0_centered * rstd- mean = tl.sum(x, axis=0) / hidden_dim- x_centered = x - mean- var = tl.sum(x_centered * x_centered, axis=0) / hidden_dim- rstd = 1.0 / tl.sqrt(var + eps)+ ln_w = tl.load(ln_weight_ptr + offs_h, mask=mask_h, other=1.0).to(tl.float32)+ ln_b = tl.load(ln_bias_ptr + offs_h, mask=mask_h, other=0.0).to(tl.float32)+ x_ln = x_norm * ln_w + ln_b- gamma = tl.load(gamma_ptr + offs_hd, mask=hd_mask, other=1.0)- beta = tl.load(beta_ptr + offs_hd, mask=hd_mask, other=0.0)- x_normed = x_centered * rstd * gamma + beta+ x1 = tl.load(x1_ptr + row_idx * H + offs_h, mask=mask_h, other=0.0).to(tl.float32)+ x_gated = x_ln * x1+ x_gated = tl.where(mask_h, x_gated, 0.0)- gate_ptrs = gate_ptr + pid_n * stride_gn + offs_hd * stride_gd- gate = tl.load(gate_ptrs, mask=hd_mask, other=0.0)- x_gated = x_normed * gate+ for c_start in range(0, C, BLOCK_C):+ offs_c = c_start + tl.arange(0, BLOCK_C)+ mask_c = offs_c < C+ weight_ptrs = linear_weight_ptr + offs_c[:, None] * H + offs_h[None, :]+ weights = tl.load(weight_ptrs, mask=mask_c[:, None] & mask_h[None, :], other=0.0).to(tl.float32)+ acc = tl.sum(weights * x_gated[None, :], axis=1)+ tl.store(out_ptr + row_idx * C + offs_c, acc, mask=mask_c)- offs_out = pid_out * BLOCK_OUT + tl.arange(0, BLOCK_OUT)- out_mask = offs_out < out_dim- acc = tl.zeros((BLOCK_OUT,), dtype=tl.float32)- for hd_start in range(0, hidden_dim, BLOCK_HD):- hd_offs = hd_start + tl.arange(0, BLOCK_HD)- hd_m = hd_offs < hidden_dim+ def kernel_function(input_tensor, mask, weights, config):+ B, N, _, C = input_tensor.shape+ H = config["hidden_dim"]- if hd_start == 0:- x_chunk = x_gated- else:- x_ptrs2 = x_ptr + pid_n * stride_xn + hd_offs * stride_xd- x2 = tl.load(x_ptrs2, mask=hd_m, other=0.0)- mean2 = tl.sum(x2, axis=0) / hidden_dim- var2 = tl.sum((x2 - mean2) * (x2 - mean2), axis=0) / hidden_dim- rstd2 = 1.0 / tl.sqrt(var2 + eps)- gamma2 = tl.load(gamma_ptr + hd_offs, mask=hd_m, other=1.0)- beta2 = tl.load(beta_ptr + hd_offs, mask=hd_m, other=0.0)- x_normed2 = (x2 - mean2) * rstd2 * gamma2 + beta2- gate2 = tl.load(- gate_ptr + pid_n * stride_gn + hd_offs * stride_gd, mask=hd_m, other=0.0- )- x_chunk = x_normed2 * gate2+ # Stage 1: LayerNorm+ x_norm = torch.empty_like(input_tensor)+ n_rows = B * N * N+ BLOCK_SIZE = triton.next_power_of_2(C)+ BLOCK_SIZE = min(BLOCK_SIZE, 1024)+ _layernorm_kernel[(n_rows,)](+ input_tensor, x_norm, weights['norm.weight'], weights['norm.bias'],+ C, C, 1e-5, BLOCK_SIZE=BLOCK_SIZE+ )- for out_idx in range(BLOCK_OUT):- if pid_out * BLOCK_OUT + out_idx < out_dim:- w_ptrs = (- w_ptr- + (pid_out * BLOCK_OUT + out_idx) * stride_wout- + hd_offs * stride_whid- )- w_vals = tl.load(w_ptrs, mask=hd_m, other=0.0)- acc = tl.where(- tl.arange(0, BLOCK_OUT) == out_idx,- acc + tl.sum(x_chunk * w_vals),- acc,- )- break+ # Stage 2: Projection + Gating+ left = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')+ right = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')+ out_gate = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')- out_ptrs = out_ptr + pid_n * stride_on + offs_out * stride_od- tl.store(out_ptrs, acc, mask=out_mask)--- def kernel_function(x, mask, weights, config):- B, seq_len, _, dim = x.shape- hidden_dim = config["hidden_dim"]- device = x.device- dtype = x.dtype-- # Flatten for processing- N = B * seq_len * seq_len- x_flat = x.reshape(N, dim).contiguous()-- # LayerNorm- x_normed = torch.empty_like(x_flat)- BLOCK_D = triton.next_power_of_2(dim)- _layernorm_kernel[(N,)](- x_flat,- x_normed,- weights["norm.weight"],- weights["norm.bias"],- N,- dim,- dim,- 1,- eps=1e-5,- BLOCK_D=BLOCK_D,+ BLOCK_H = min(32, H)+ BLOCK_C = min(32, C)+ _projection_gating_kernel[(B * N * N,)](+ x_norm, mask,+ weights['left_proj.weight'], weights['right_proj.weight'],+ weights['left_gate.weight'], weights['right_gate.weight'], weights['out_gate.weight'],+ left, right, out_gate,+ B, N, C, H, BLOCK_H, BLOCK_C)- # Linear projections using torch.mm (for simplicity, we use matmul here)- x_normed_2d = x_normed.view(N, dim)- left = x_normed_2d @ weights["left_proj.weight"].T- right = x_normed_2d @ weights["right_proj.weight"].T- left_gate = torch.sigmoid(x_normed_2d @ weights["left_gate.weight"].T)- right_gate = torch.sigmoid(x_normed_2d @ weights["right_gate.weight"].T)- out_gate = torch.sigmoid(x_normed_2d @ weights["out_gate.weight"].T)-- # Reshape back- left = left.view(B, seq_len, seq_len, hidden_dim)- right = right.view(B, seq_len, seq_len, hidden_dim)- left_gate = left_gate.view(B, seq_len, seq_len, hidden_dim)- right_gate = right_gate.view(B, seq_len, seq_len, hidden_dim)- out_gate = out_gate.view(B, seq_len, seq_len, hidden_dim)-- # Masked gating einsum- out = torch.empty((B, seq_len, seq_len, hidden_dim), device=device, dtype=dtype)- BLOCK_D = min(64, triton.next_power_of_2(hidden_dim))- grid = (B * seq_len, seq_len, triton.cdiv(hidden_dim, BLOCK_D))- _masked_gating_einsum_kernel[grid](- left,- right,- left_gate,- right_gate,- mask.float(),- out,- B,- seq_len,- hidden_dim,- *left.stride(),- *right.stride(),- *left_gate.stride(),- *right_gate.stride(),- *mask.stride(),- *out.stride(),- BLOCK_D=BLOCK_D,+ # Stage 3: Triangular multiplication+ tri_out = torch.empty((B, N, N, H), dtype=torch.float32, device='cuda')+ BLOCK_I, BLOCK_J, BLOCK_K = 16, 16, 16+ num_tiles_i = triton.cdiv(N, BLOCK_I)+ num_tiles_j = triton.cdiv(N, BLOCK_J)+ _triangular_mul_kernel[(B, num_tiles_i * num_tiles_j, H)](+ left, right, tri_out, B, N, H, BLOCK_I, BLOCK_J, BLOCK_K)- # Output: LayerNorm, gate, linear- out_flat = out.reshape(N, hidden_dim)- out_normed = torch.empty_like(out_flat)- BLOCK_HD = triton.next_power_of_2(hidden_dim)- _layernorm_kernel[(N,)](- out_flat,- out_normed,- weights["to_out_norm.weight"],- weights["to_out_norm.bias"],- N,- hidden_dim,- hidden_dim,- 1,- eps=1e-5,- BLOCK_D=BLOCK_HD,+ # Stage 4: Output LayerNorm + Gate + Linear+ output = torch.empty((B, N, N, C), dtype=torch.float32, device='cuda')+ BLOCK_H_out = triton.next_power_of_2(H)+ BLOCK_C_out = min(128, triton.next_power_of_2(C))+ _output_norm_gate_proj_kernel[(n_rows,)](+ tri_out.reshape(-1, H), out_gate.reshape(-1, H),+ weights['to_out_norm.weight'], weights['to_out_norm.bias'], weights['to_out.weight'],+ output.reshape(-1, C), n_rows, H, C, 1e-5, BLOCK_H_out, BLOCK_C_out)- out_gated = out_normed * out_gate.reshape(N, hidden_dim)- result = out_gated @ weights["to_out.weight"].T- return result.view(B, seq_len, seq_len, dim)+ return output- def test_kernel():- from math import sqrt+ def test_kernel():torch.manual_seed(42)- B, seq_len, dim, hidden_dim = 2, 8, 32, 64- x = torch.randn(B, seq_len, seq_len, dim, device="cuda")- mask = torch.randint(0, 2, (B, seq_len, seq_len), device="cuda").float()+ B, N, C, H = 1, 32, 128, 128++ input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)+ mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)+weights = {- "norm.weight": torch.randn(dim, device="cuda"),- "norm.bias": torch.randn(dim, device="cuda"),- "left_proj.weight": torch.randn(hidden_dim, dim, device="cuda")- / sqrt(hidden_dim),- "right_proj.weight": torch.randn(hidden_dim, dim, device="cuda")- / sqrt(hidden_dim),- "left_gate.weight": torch.randn(hidden_dim, dim, device="cuda")- / sqrt(hidden_dim),- "right_gate.weight": torch.randn(hidden_dim, dim, device="cuda")- / sqrt(hidden_dim),- "out_gate.weight": torch.randn(hidden_dim, dim, device="cuda")- / sqrt(hidden_dim),- "to_out_norm.weight": torch.randn(hidden_dim, device="cuda"),- "to_out_norm.bias": torch.randn(hidden_dim, device="cuda"),- "to_out.weight": torch.randn(dim, hidden_dim, device="cuda") / sqrt(dim),+ 'norm.weight': torch.randn(C, device='cuda'), 'norm.bias': torch.randn(C, device='cuda'),+ 'left_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),+ 'right_proj.weight': torch.randn(H, C, device='cuda') / (H**0.5),+ 'left_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),+ 'right_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),+ 'out_gate.weight': torch.randn(H, C, device='cuda') / (H**0.5),+ 'to_out_norm.weight': torch.randn(H, device='cuda'), 'to_out_norm.bias': torch.randn(H, device='cuda'),+ 'to_out.weight': torch.randn(C, H, device='cuda') / (C**0.5),}- config = {"dim": dim, "hidden_dim": hidden_dim}- out_triton = kernel_function(x, mask, weights, config)+ config = {"dim": C, "hidden_dim": H}- from torch import nn+ out_triton = kernel_function(input_tensor, mask, weights, config)- model = nn.Module()- x_n = nn.functional.layer_norm(- x, [dim], weights["norm.weight"], weights["norm.bias"]- )- left = x_n @ weights["left_proj.weight"].T- right = x_n @ weights["right_proj.weight"].T- lg = torch.sigmoid(x_n @ weights["left_gate.weight"].T)- rg = torch.sigmoid(x_n @ weights["right_gate.weight"].T)- og = torch.sigmoid(x_n @ weights["out_gate.weight"].T)- m = mask.unsqueeze(-1)- left = left * m * lg- right = right * m * rg- out = torch.einsum("bikd,bjkd->bijd", left, right)- out = nn.functional.layer_norm(- out, [hidden_dim], weights["to_out_norm.weight"], weights["to_out_norm.bias"]- )- out = out * og- out_ref = out @ weights["to_out.weight"].T+ # Reference+ from torch import nn, einsum+ x = torch.nn.functional.layer_norm(input_tensor, [C], weights['norm.weight'], weights['norm.bias'])+ left = x @ weights['left_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['left_gate.weight'].T)+ right = x @ weights['right_proj.weight'].T * mask.unsqueeze(-1) * torch.sigmoid(x @ weights['right_gate.weight'].T)+ out_gate = torch.sigmoid(x @ weights['out_gate.weight'].T)+ tri = einsum('bikd,bjkd->bijd', left, right)+ tri_norm = torch.nn.functional.layer_norm(tri, [H], weights['to_out_norm.weight'], weights['to_out_norm.bias'])+ ref = (tri_norm * out_gate) @ weights['to_out.weight'].T- if torch.allclose(out_triton, out_ref, rtol=2e-2, atol=2e-2):+ if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):print("PASS")else:- print(f"FAIL: max diff = {(out_triton - out_ref).abs().max()}")+ print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")- # if __name__ == "__main__":- # test_kernel()+ if __name__ == "__main__":+ test_kernel()- input_t = TypeVar(- "input_t", bound=Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict]- )--- def custom_kernel(input: input_t) -> torch.Tensor:- x, mask, weights, config = input- return kernel_function(x, mask, weights, config)+ def custom_kernel(input):+ return kernel_function(*input)
scrolls · 619 diff lines total
Best evidence level for this revision: reported
JSON