submission 480973
Cookie 🍪 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 479 lines, June 9 Researcher Reciprocity License v1.0.
trimul_H100_gpt-5_ka_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-480973?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:e6fff007a4e413a431a9c6ffbbdd8fd57f2a7b3044b1f4b41662a95e6b8f7f5c
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_proj += tl.dot(a_tile, wproj_tile) # [M_blk, U_blk]num-warps = 4
num_warps=4, num_stages=2stages = 2
num_warps=4, num_stages=2tile-k = 64
BLOCK_K = 64tile-m = 64
BLOCK_M = 64tile-n = 64
BLOCK_N = 64Kernel source
trimul_H100_gpt-5_ka_submission.py479 lines
# kernel.py
# Triangle Multiplicative Update (Outgoing) -- Forward-only Triton implementation.
# All compute is performed inside Triton kernels; Python only validates, allocates, and launches.
#
# Fused stages implemented across kernels (aggressively fused where feasible):
# 1) LayerNorm over D for x[b, i, j, :]
# 2) Fused linear(gate+proj) + sigmoid + masking for LEFT and RIGHT paths on rows [B, N, N]
# 3) Contraction over k: S[b, i, j, u] = sum_k LEFT[b, i, k, u] * RIGHT[b, j, k, u]
# 4) OUT gate per (b, i, j): sigmoid(x_norm[b,i,j,:] @ W_out_gate^T)
# 5) Final LayerNorm over U + affine + gate + matmul to D: out[b,i,j,:]
#
# Fusion boundaries:
# - The proj+gate+sigmoid(+mask) is fused per path to minimize memory traffic.
# - The N-reduction over k (stage 3) and the U-epilogue + projection (stage 5) use incompatible tilings:
# * The contraction reduces along N while keeping (i,j,u), whereas the epilogue normalizes along U and projects to D.
# * A single-kernel fusion would cause extreme register/shared-memory pressure and poor occupancy.
# Therefore, they are launched as separate kernels for performance and maintainability.
import triton
import triton.language as tl
import torch
@triton.jit
def layer_norm_lastdim_kernel(x_ptr, y_ptr, gamma_ptr, beta_ptr,
M, D, eps,
BLOCK_D: tl.constexpr):
# Each program normalizes one row of length D
pid = tl.program_id(axis=0)
if pid >= M:
return
row_x = x_ptr + pid * D
row_y = y_ptr + pid * D
offs = tl.arange(0, BLOCK_D)
# Pass 1: mean
sum_val = tl.zeros((), dtype=tl.float32)
for d0 in tl.range(0, D, BLOCK_D):
d_ids = d0 + offs
mask = d_ids < D
x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
sum_val += tl.sum(x, axis=0)
d_float = tl.cast(D, tl.float32)
mean = sum_val / d_float
# Pass 2: variance via squared deviations
sum_sqdiff = tl.zeros((), dtype=tl.float32)
for d0 in tl.range(0, D, BLOCK_D):
d_ids = d0 + offs
mask = d_ids < D
x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
diff = x - mean
sum_sqdiff += tl.sum(diff * diff, axis=0)
var = sum_sqdiff / d_float
inv_std = 1.0 / tl.sqrt(var + eps)
# Pass 3: normalize, scale, shift
for d0 in tl.range(0, D, BLOCK_D):
d_ids = d0 + offs
mask = d_ids < D
x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)
g = tl.load(gamma_ptr + d_ids, mask=mask, other=1.0).to(tl.float32)
b = tl.load(beta_ptr + d_ids, mask=mask, other=0.0).to(tl.float32)
y = (x - mean) * inv_std
y = y * g + b
tl.store(row_y + d_ids, y, mask=mask)
@triton.jit
def proj_gate_fused_rows_kernel(a_ptr, wproj_ptr, wgate_ptr, mask_ptr, out_ptr,
M, N, D, U,
apply_mask: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
"""
Fused projection and gate over rows of a 2D tensor [M, D], where M = B * N * N and each row maps to (b, i, j).
Produces: out[row, u] = (a_row @ wproj.T) * sigmoid(a_row @ wgate.T) * (optional mask[b,i,j]).
"""
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
acc_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Loop over K = D
for k0 in tl.range(0, D, BLOCK_K):
k_ids = k0 + tl.arange(0, BLOCK_K)
k_mask = k_ids < D
# A tile: [BLOCK_M, BLOCK_K]
a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
a_tile = tl.load(
a_ptrs,
mask=(offs_m[:, None] < M) & (k_mask[None, :]),
other=0.0
).to(tl.float32)
# Weight tiles: [BLOCK_K, BLOCK_N] (weights are [U, D])
wproj_ptrs = wproj_ptr + (offs_n[None, :] * D + k_ids[:, None])
wgate_ptrs = wgate_ptr + (offs_n[None, :] * D + k_ids[:, None])
wproj_tile = tl.load(
wproj_ptrs,
mask=(offs_n[None, :] < U) & (k_mask[:, None]),
other=0.0
).to(tl.float32)
wgate_tile = tl.load(
wgate_ptrs,
mask=(offs_n[None, :] < U) & (k_mask[:, None]),
other=0.0
).to(tl.float32)
acc_proj += tl.dot(a_tile, wproj_tile) # [M_blk, U_blk]
acc_gate += tl.dot(a_tile, wgate_tile)
gate = tl.sigmoid(acc_gate)
out = acc_proj * gate
if apply_mask:
NN = N * N
b_m = offs_m // NN
rem = offs_m - b_m * NN
i_m = rem // N
j_m = rem - i_m * N
m_ptrs = mask_ptr + b_m * NN + i_m * N + j_m
mvals = tl.load(m_ptrs, mask=offs_m < M, other=0.0).to(tl.float32) # [M_blk]
out = out * mvals[:, None]
out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
mask_out = (offs_m[:, None] < M) & (offs_n[None, :] < U)
tl.store(out_ptrs, out, mask=mask_out)
@triton.jit
def contract_k_kernel(L_ptr, R_ptr, S_ptr,
B, N, U,
stride_lb, stride_li, stride_lk, stride_lu,
stride_rb, stride_rj, stride_rk, stride_ru,
stride_sb, stride_si, stride_sj, stride_su,
BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr):
# 3D launch: I, J, (B*U)
pid_i = tl.program_id(axis=0)
pid_j = tl.program_id(axis=1)
pid_bu = tl.program_id(axis=2)
b = pid_bu // U
u = pid_bu % U
offs_i = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
offs_j = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
for k0 in tl.range(0, N, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
k_mask = offs_k < N
# L[b, i, k, u] -> [BLOCK_I, BLOCK_K]
l_ptrs = (L_ptr
+ b * stride_lb
+ offs_i[:, None] * stride_li
+ offs_k[None, :] * stride_lk
+ u * stride_lu)
L_tile = tl.load(l_ptrs, mask=(offs_i[:, None] < N) & (k_mask[None, :]), other=0.0).to(tl.float32)
# R[b, j, k, u] -> [BLOCK_K, BLOCK_J]
r_ptrs = (R_ptr
+ b * stride_rb
+ offs_j[None, :] * stride_rj
+ offs_k[:, None] * stride_rk
+ u * stride_ru)
R_tile_KJ = tl.load(r_ptrs, mask=(k_mask[:, None]) & (offs_j[None, :] < N), other=0.0).to(tl.float32)
acc += tl.dot(L_tile, R_tile_KJ)
s_ptrs = (S_ptr
+ b * stride_sb
+ offs_i[:, None] * stride_si
+ offs_j[None, :] * stride_sj
+ u * stride_su)
tl.store(s_ptrs, acc, mask=(offs_i[:, None] < N) & (offs_j[None, :] < N))
@triton.jit
def out_gate_kernel(a_ptr, w_ptr, out_ptr, M, D, U,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
# Compute: sigmoid(a @ w^T) for M rows, projecting D->U
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k0 in tl.range(0, D, BLOCK_K):
k_ids = k0 + tl.arange(0, BLOCK_K)
k_mask = k_ids < D
a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])
a_tile = tl.load(a_ptrs,
mask=(offs_m[:, None] < M) & (k_mask[None, :]),
other=0.0).to(tl.float32)
w_ptrs = w_ptr + (offs_n[None, :] * D + k_ids[:, None])
w_tile = tl.load(w_ptrs,
mask=(offs_n[None, :] < U) & (k_mask[:, None]),
other=0.0).to(tl.float32)
acc += tl.dot(a_tile, w_tile)
out = tl.sigmoid(acc)
out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])
tl.store(out_ptrs, out, mask=(offs_m[:, None] < M) & (offs_n[None, :] < U))
@triton.jit
def final_project_kernel_affine(S_ptr, G_ptr, W_ptr, gamma_ptr, beta_ptr, Out_ptr,
B, N, U, D,
stride_sb, stride_si, stride_sj, stride_su,
stride_gb, stride_gi, stride_gj, stride_gu,
BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_U: tl.constexpr, BLOCK_D: tl.constexpr):
"""
For each (b, i, j), perform:
Z = LayerNorm_U(S[b,i,j,:], eps=1e-5) with affine (gamma, beta)
Z = Z * G[b,i,j,:]
out[b,i,j,:] = Z @ W^T where W is [D, U]
We implement LN over U via a two-pass reduction (mean, var) to fp32 and project to D in tiles.
"""
pid_i = tl.program_id(axis=0)
pid_j = tl.program_id(axis=1)
pid_bd = tl.program_id(axis=2)
num_d_blocks = tl.cdiv(D, BLOCK_D)
b = pid_bd // num_d_blocks
d_block = pid_bd % num_d_blocks
d_offs = d_block * BLOCK_D + tl.arange(0, BLOCK_D)
i_offs = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)
j_offs = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)
# Pass 1: mean over U of S[b, i, j, :]
sum_vals = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
for u0 in tl.range(0, U, BLOCK_U):
u_ids = u0 + tl.arange(0, BLOCK_U)
u_mask = u_ids < U
s_ptrs = (S_ptr
+ b * stride_sb
+ i_offs[:, None, None] * stride_si
+ j_offs[None, :, None] * stride_sj
+ u_ids[None, None, :] * stride_su)
S_blk = tl.load(s_ptrs,
mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
other=0.0).to(tl.float32)
sum_vals += tl.sum(S_blk, axis=2)
U_f32 = tl.cast(U, tl.float32)
mean = sum_vals / U_f32
# Pass 2: variance via squared deviations
sum_sqdiff = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
for u0 in tl.range(0, U, BLOCK_U):
u_ids = u0 + tl.arange(0, BLOCK_U)
u_mask = u_ids < U
s_ptrs = (S_ptr
+ b * stride_sb
+ i_offs[:, None, None] * stride_si
+ j_offs[None, :, None] * stride_sj
+ u_ids[None, None, :] * stride_su)
S_blk = tl.load(s_ptrs,
mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
other=0.0).to(tl.float32)
diff = S_blk - mean[:, :, None]
sum_sqdiff += tl.sum(diff * diff, axis=2)
var = sum_sqdiff / U_f32
inv_std = 1.0 / tl.sqrt(var + 1.0e-5)
# Accumulator for final projection to D: [BLOCK_I, BLOCK_J, BLOCK_D]
acc_out = tl.zeros((BLOCK_I, BLOCK_J, BLOCK_D), dtype=tl.float32)
# Pass 3: normalize + affine + gate + matmul with W to D
for u0 in tl.range(0, U, BLOCK_U):
u_ids = u0 + tl.arange(0, BLOCK_U)
u_mask = u_ids < U
s_ptrs = (S_ptr
+ b * stride_sb
+ i_offs[:, None, None] * stride_si
+ j_offs[None, :, None] * stride_sj
+ u_ids[None, None, :] * stride_su)
g_ptrs = (G_ptr
+ b * stride_gb
+ i_offs[:, None, None] * stride_gi
+ j_offs[None, :, None] * stride_gj
+ u_ids[None, None, :] * stride_gu)
gamma_blk = tl.load(gamma_ptr + u_ids, mask=u_mask, other=1.0).to(tl.float32)
beta_blk = tl.load(beta_ptr + u_ids, mask=u_mask, other=0.0).to(tl.float32)
S_blk = tl.load(s_ptrs,
mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
other=0.0).to(tl.float32)
G_blk = tl.load(g_ptrs,
mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),
other=0.0).to(tl.float32)
# Normalize + affine + gate
S_norm = (S_blk - mean[:, :, None]) * inv_std[:, :, None]
S_affine = S_norm * gamma_blk[None, None, :] + beta_blk[None, None, :]
Z = S_affine * G_blk # [BLOCK_I, BLOCK_J, BLOCK_U]
# W_ptr is [D, U]; submatrix shaped [BLOCK_U, BLOCK_D] storing W^T
W_sub = tl.load(W_ptr + (d_offs[None, :] * U + u_ids[:, None]),
mask=(d_offs[None, :] < D) & (u_mask[:, None]),
other=0.0).to(tl.float32)
# Flatten Z to 2D: [(BLOCK_I*BLOCK_J), BLOCK_U] then matmul
Z_flat = tl.reshape(Z, (BLOCK_I * BLOCK_J, BLOCK_U))
tmp = tl.dot(Z_flat, W_sub) # [BLOCK_I*BLOCK_J, BLOCK_D]
tmp_3d = tl.reshape(tmp, (BLOCK_I, BLOCK_J, BLOCK_D))
acc_out += tmp_3d
out_ptrs = (Out_ptr
+ b * (N * N * D)
+ i_offs[:, None, None] * (N * D)
+ j_offs[None, :, None] * D
+ d_offs[None, None, :])
tl.store(out_ptrs, acc_out,
mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (d_offs[None, None, :] < D))
def _check_inputs(x, mask, weights, config):
assert isinstance(x, torch.Tensor) and x.is_cuda and x.dtype == torch.float32 and x.is_contiguous()
assert isinstance(mask, torch.Tensor) and mask.is_cuda and mask.dtype == torch.float32 and mask.is_contiguous()
assert x.ndim == 4
B, N, N2, D = x.shape
assert N == N2, "x must be [B, N, N, D]"
assert mask.shape == (B, N, N)
assert isinstance(weights, dict)
assert isinstance(config, dict)
U = int(config["hidden_dim"])
assert "dim" in config and int(config["dim"]) == D
def _ck(name, shape):
t = weights.get(name, None)
assert t is not None and isinstance(t, torch.Tensor) and t.is_cuda and t.dtype == torch.float32
assert t.is_contiguous()
assert tuple(t.shape) == tuple(shape)
_ck("norm.weight", (D,))
_ck("norm.bias", (D,))
_ck("left_proj.weight", (U, D))
_ck("right_proj.weight", (U, D))
_ck("left_gate.weight", (U, D))
_ck("right_gate.weight", (U, D))
_ck("out_gate.weight", (U, D))
_ck("to_out_norm.weight", (U,))
_ck("to_out_norm.bias", (U,))
_ck("to_out.weight", (D, U))
def kernel_function(*args):
# Accept either a single tuple or unpacked args
if len(args) == 1 and isinstance(args[0], tuple):
x, mask, weights, config = args[0]
else:
x, mask, weights, config = args
_check_inputs(x, mask, weights, config)
B, N, _, D = x.shape
U = int(config["hidden_dim"])
device = x.device
# Allocate intermediates
x_norm = torch.empty_like(x, dtype=torch.float32, device=device)
M = B * N * N
x_flat = x.view(M, D)
x_norm_flat = x_norm.view(M, D)
# 1) LayerNorm over last dim
eps = 1e-5
BLOCK_LN = 128
grid_ln = (M,)
layer_norm_lastdim_kernel[grid_ln](
x_flat, x_norm_flat, weights["norm.weight"], weights["norm.bias"],
M, D, eps,
BLOCK_D=BLOCK_LN,
num_warps=4, num_stages=2
)
# 2) left and right features via fused proj+gate (+ mask) over rows [B, N, N].
left = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
right = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
left_flat = left.view(M, U)
right_flat = right.view(M, U)
BLOCK_M = 64
BLOCK_N = 64
BLOCK_K = 64
grid_pg = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))
# LEFT on rows (b,i,j) with mask applied per row
proj_gate_fused_rows_kernel[grid_pg](
x_norm_flat, weights["left_proj.weight"], weights["left_gate.weight"], mask, left_flat,
M, N, D, U,
apply_mask=True,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
num_warps=4, num_stages=2
)
# RIGHT on rows (b,i,j) with its mask applied per row
proj_gate_fused_rows_kernel[grid_pg](
x_norm_flat, weights["right_proj.weight"], weights["right_gate.weight"], mask, right_flat,
M, N, D, U,
apply_mask=True,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
num_warps=4, num_stages=2
)
# 3) Contraction over k to produce S: [B, N, N, U]
S = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
# Strides for [B, N, N, U] contiguous
stride_b = N * N * U
stride_i = N * U
stride_j = U
stride_u = 1
BLOCK_I = 32
BLOCK_J = 32
BLOCK_KC = 64
grid_contract = (triton.cdiv(N, BLOCK_I), triton.cdiv(N, BLOCK_J), B * U)
contract_k_kernel[grid_contract](
left, right, S,
B, N, U,
stride_b, stride_i, stride_j, stride_u, # L strides: b,i,k(==j),u
stride_b, stride_i, stride_j, stride_u, # R strides: b,j(==i),k(==j),u
stride_b, stride_i, stride_j, stride_u, # S strides: b,i,j,u
BLOCK_I=BLOCK_I, BLOCK_J=BLOCK_J, BLOCK_K=BLOCK_KC,
num_warps=4, num_stages=2
)
# 4) Compute out_gate: G = sigmoid(x_norm @ W_out_gate^T)
G = torch.empty((B, N, N, U), device=device, dtype=torch.float32)
G_flat = G.view(M, U)
grid_gate = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))
out_gate_kernel[grid_gate](
x_norm_flat, weights["out_gate.weight"], G_flat, M, D, U,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
num_warps=4, num_stages=2
)
# 5) Final projection with LayerNorm(U) + affine + gate + matmul to D
out = torch.empty((B, N, N, D), device=device, dtype=torch.float32)
BLOCK_I2 = 16
BLOCK_J2 = 16
BLOCK_U2 = 32
BLOCK_D2 = 64
grid_final = (triton.cdiv(N, BLOCK_I2), triton.cdiv(N, BLOCK_J2), B * triton.cdiv(D, BLOCK_D2))
final_project_kernel_affine[grid_final](
S, G, weights["to_out.weight"], weights["to_out_norm.weight"], weights["to_out_norm.bias"], out,
B, N, U, D,
stride_b, stride_i, stride_j, stride_u, # S strides
stride_b, stride_i, stride_j, stride_u, # G strides
BLOCK_I=BLOCK_I2, BLOCK_J=BLOCK_J2, BLOCK_U=BLOCK_U2, BLOCK_D=BLOCK_D2,
num_warps=4, num_stages=2
)
return out
def custom_kernel(input):
return kernel_function(*input)
scrolls · 479 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 426455.
- import torch+ # kernel.py+ # Triangle Multiplicative Update (Outgoing) -- Forward-only Triton implementation.+ # All compute is performed inside Triton kernels; Python only validates, allocates, and launches.+ #+ # Fused stages implemented across kernels (aggressively fused where feasible):+ # 1) LayerNorm over D for x[b, i, j, :]+ # 2) Fused linear(gate+proj) + sigmoid + masking for LEFT and RIGHT paths on rows [B, N, N]+ # 3) Contraction over k: S[b, i, j, u] = sum_k LEFT[b, i, k, u] * RIGHT[b, j, k, u]+ # 4) OUT gate per (b, i, j): sigmoid(x_norm[b,i,j,:] @ W_out_gate^T)+ # 5) Final LayerNorm over U + affine + gate + matmul to D: out[b,i,j,:]+ #+ # Fusion boundaries:+ # - The proj+gate+sigmoid(+mask) is fused per path to minimize memory traffic.+ # - The N-reduction over k (stage 3) and the U-epilogue + projection (stage 5) use incompatible tilings:+ # * The contraction reduces along N while keeping (i,j,u), whereas the epilogue normalizes along U and projects to D.+ # * A single-kernel fusion would cause extreme register/shared-memory pressure and poor occupancy.+ # Therefore, they are launched as separate kernels for performance and maintainability.+import tritonimport triton.language as tl+ import torch@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+ def layer_norm_lastdim_kernel(x_ptr, y_ptr, gamma_ptr, beta_ptr,+ M, D, eps,+ BLOCK_D: tl.constexpr):+ # Each program normalizes one row of length D+ pid = tl.program_id(axis=0)+ if pid >= M:+ return- 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+ row_x = x_ptr + pid * D+ row_y = y_ptr + pid * D+ offs = tl.arange(0, BLOCK_D)- 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+ # Pass 1: mean+ sum_val = tl.zeros((), dtype=tl.float32)+ for d0 in tl.range(0, D, BLOCK_D):+ d_ids = d0 + offs+ mask = d_ids < D+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)+ sum_val += tl.sum(x, axis=0)+ d_float = tl.cast(D, tl.float32)+ mean = sum_val / d_float- out_ptr = Y_ptr + row_idx * stride_x- tl.store(out_ptr + col_offsets, y, mask=mask)+ # Pass 2: variance via squared deviations+ sum_sqdiff = tl.zeros((), dtype=tl.float32)+ for d0 in tl.range(0, D, BLOCK_D):+ d_ids = d0 + offs+ mask = d_ids < D+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)+ diff = x - mean+ sum_sqdiff += tl.sum(diff * diff, axis=0)+ var = sum_sqdiff / d_float+ inv_std = 1.0 / tl.sqrt(var + eps)+ # Pass 3: normalize, scale, shift+ for d0 in tl.range(0, D, BLOCK_D):+ d_ids = d0 + offs+ mask = d_ids < D+ x = tl.load(row_x + d_ids, mask=mask, other=0.0).to(tl.float32)+ g = tl.load(gamma_ptr + d_ids, mask=mask, other=1.0).to(tl.float32)+ b = tl.load(beta_ptr + d_ids, mask=mask, other=0.0).to(tl.float32)+ y = (x - mean) * inv_std+ y = y * g + b+ tl.store(row_y + d_ids, 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+ def proj_gate_fused_rows_kernel(a_ptr, wproj_ptr, wgate_ptr, mask_ptr, out_ptr,+ M, N, D, U,+ apply_mask: tl.constexpr,+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):+ """+ Fused projection and gate over rows of a 2D tensor [M, D], where M = B * N * N and each row maps to (b, i, j).+ Produces: out[row, u] = (a_row @ wproj.T) * sigmoid(a_row @ wgate.T) * (optional mask[b,i,j]).+ """+ pid_m = tl.program_id(axis=0)+ pid_n = tl.program_id(axis=1)- 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+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)- for h_start in range(0, H, BLOCK_H):- h_offs = h_start + tl.arange(0, BLOCK_H)- h_mask = h_offs < H+ acc_proj = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)- 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)+ # Loop over K = D+ for k0 in tl.range(0, D, BLOCK_K):+ k_ids = k0 + tl.arange(0, BLOCK_K)+ k_mask = k_ids < D- 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)+ # A tile: [BLOCK_M, BLOCK_K]+ a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])+ a_tile = tl.load(+ a_ptrs,+ mask=(offs_m[:, None] < M) & (k_mask[None, :]),+ other=0.0+ ).to(tl.float32)- weight_offsets = h_offs[:, None] * C + c_offs[None, :]- combined_mask = h_mask[:, None] & c_mask[None, :]+ # Weight tiles: [BLOCK_K, BLOCK_N] (weights are [U, D])+ wproj_ptrs = wproj_ptr + (offs_n[None, :] * D + k_ids[:, None])+ wgate_ptrs = wgate_ptr + (offs_n[None, :] * D + k_ids[:, 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)+ wproj_tile = tl.load(+ wproj_ptrs,+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),+ other=0.0+ ).to(tl.float32)+ wgate_tile = tl.load(+ wgate_ptrs,+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),+ 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)+ acc_proj += tl.dot(a_tile, wproj_tile) # [M_blk, U_blk]+ acc_gate += tl.dot(a_tile, wgate_tile)- left_gate_sig = tl.sigmoid(left_gate_acc)- right_gate_sig = tl.sigmoid(right_gate_acc)- out_gate_sig = tl.sigmoid(out_gate_acc)+ gate = tl.sigmoid(acc_gate)+ out = acc_proj * gate- left_result = left_proj_acc * mask_val * left_gate_sig- right_result = right_proj_acc * mask_val * right_gate_sig+ if apply_mask:+ NN = N * N+ b_m = offs_m // NN+ rem = offs_m - b_m * NN+ i_m = rem // N+ j_m = rem - i_m * N+ m_ptrs = mask_ptr + b_m * NN + i_m * N + j_m+ mvals = tl.load(m_ptrs, mask=offs_m < M, other=0.0).to(tl.float32) # [M_blk]+ out = out * mvals[:, None]- 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)+ out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])+ mask_out = (offs_m[:, None] < M) & (offs_n[None, :] < U)+ tl.store(out_ptrs, out, mask=mask_out)@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)+ def contract_k_kernel(L_ptr, R_ptr, S_ptr,+ B, N, U,+ stride_lb, stride_li, stride_lk, stride_lu,+ stride_rb, stride_rj, stride_rk, stride_ru,+ stride_sb, stride_si, stride_sj, stride_su,+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr):+ # 3D launch: I, J, (B*U)+ pid_i = tl.program_id(axis=0)+ pid_j = tl.program_id(axis=1)+ pid_bu = tl.program_id(axis=2)- num_tiles_j = tl.cdiv(N, BLOCK_J)- pid_i = pid_ij // num_tiles_j- pid_j = pid_ij % num_tiles_j+ b = pid_bu // U+ u = pid_bu % Uoffs_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 < Nacc = 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+ for k0 in tl.range(0, N, BLOCK_K):+ offs_k = k0 + tl.arange(0, BLOCK_K)+ k_mask = 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)+ # L[b, i, k, u] -> [BLOCK_I, BLOCK_K]+ l_ptrs = (L_ptr+ + b * stride_lb+ + offs_i[:, None] * stride_li+ + offs_k[None, :] * stride_lk+ + u * stride_lu)+ L_tile = tl.load(l_ptrs, mask=(offs_i[:, None] < N) & (k_mask[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)+ # R[b, j, k, u] -> [BLOCK_K, BLOCK_J]+ r_ptrs = (R_ptr+ + b * stride_rb+ + offs_j[None, :] * stride_rj+ + offs_k[:, None] * stride_rk+ + u * stride_ru)+ R_tile_KJ = tl.load(r_ptrs, mask=(k_mask[:, None]) & (offs_j[None, :] < N), other=0.0).to(tl.float32)- acc += tl.dot(x0, tl.trans(x1))+ acc += tl.dot(L_tile, R_tile_KJ)- 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, :])+ s_ptrs = (S_ptr+ + b * stride_sb+ + offs_i[:, None] * stride_si+ + offs_j[None, :] * stride_sj+ + u * stride_su)+ tl.store(s_ptrs, acc, mask=(offs_i[:, None] < N) & (offs_j[None, :] < N))@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+ def out_gate_kernel(a_ptr, w_ptr, out_ptr, M, D, U,+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):+ # Compute: sigmoid(a @ w^T) for M rows, projecting D->U+ pid_m = tl.program_id(axis=0)+ pid_n = tl.program_id(axis=1)- offs_h = tl.arange(0, BLOCK_H)- mask_h = offs_h < H+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)- 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+ acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)- 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+ for k0 in tl.range(0, D, BLOCK_K):+ k_ids = k0 + tl.arange(0, BLOCK_K)+ k_mask = k_ids < D- 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)+ a_ptrs = a_ptr + (offs_m[:, None] * D + k_ids[None, :])+ a_tile = tl.load(a_ptrs,+ mask=(offs_m[:, None] < M) & (k_mask[None, :]),+ other=0.0).to(tl.float32)- 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)+ w_ptrs = w_ptr + (offs_n[None, :] * D + k_ids[:, None])+ w_tile = tl.load(w_ptrs,+ mask=(offs_n[None, :] < U) & (k_mask[:, None]),+ other=0.0).to(tl.float32)+ acc += tl.dot(a_tile, w_tile)- def kernel_function(input_tensor, mask, weights, config):- B, N, _, C = input_tensor.shape- H = config["hidden_dim"]+ out = tl.sigmoid(acc)+ out_ptrs = out_ptr + (offs_m[:, None] * U + offs_n[None, :])+ tl.store(out_ptrs, out, mask=(offs_m[:, None] < M) & (offs_n[None, :] < U))- # 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')+ @triton.jit+ def final_project_kernel_affine(S_ptr, G_ptr, W_ptr, gamma_ptr, beta_ptr, Out_ptr,+ B, N, U, D,+ stride_sb, stride_si, stride_sj, stride_su,+ stride_gb, stride_gi, stride_gj, stride_gu,+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_U: tl.constexpr, BLOCK_D: tl.constexpr):+ """+ For each (b, i, j), perform:+ Z = LayerNorm_U(S[b,i,j,:], eps=1e-5) with affine (gamma, beta)+ Z = Z * G[b,i,j,:]+ out[b,i,j,:] = Z @ W^T where W is [D, U]+ We implement LN over U via a two-pass reduction (mean, var) to fp32 and project to D in tiles.+ """+ pid_i = tl.program_id(axis=0)+ pid_j = tl.program_id(axis=1)+ pid_bd = tl.program_id(axis=2)- 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- )+ num_d_blocks = tl.cdiv(D, BLOCK_D)+ b = pid_bd // num_d_blocks+ d_block = pid_bd % num_d_blocks+ d_offs = d_block * BLOCK_D + tl.arange(0, 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- )+ i_offs = pid_i * BLOCK_I + tl.arange(0, BLOCK_I)+ j_offs = pid_j * BLOCK_J + tl.arange(0, BLOCK_J)- # 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- )+ # Pass 1: mean over U of S[b, i, j, :]+ sum_vals = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)+ for u0 in tl.range(0, U, BLOCK_U):+ u_ids = u0 + tl.arange(0, BLOCK_U)+ u_mask = u_ids < U- return output+ s_ptrs = (S_ptr+ + b * stride_sb+ + i_offs[:, None, None] * stride_si+ + j_offs[None, :, None] * stride_sj+ + u_ids[None, None, :] * stride_su)+ S_blk = tl.load(s_ptrs,+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),+ other=0.0).to(tl.float32)+ sum_vals += tl.sum(S_blk, axis=2)+ U_f32 = tl.cast(U, tl.float32)+ mean = sum_vals / U_f32- def test_kernel():- torch.manual_seed(42)- B, N, C, H = 1, 32, 128, 128+ # Pass 2: variance via squared deviations+ sum_sqdiff = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)+ for u0 in tl.range(0, U, BLOCK_U):+ u_ids = u0 + tl.arange(0, BLOCK_U)+ u_mask = u_ids < U- input_tensor = torch.randn(B, N, N, C, device='cuda', dtype=torch.float32)- mask = torch.ones(B, N, N, device='cuda', dtype=torch.float32)+ s_ptrs = (S_ptr+ + b * stride_sb+ + i_offs[:, None, None] * stride_si+ + j_offs[None, :, None] * stride_sj+ + u_ids[None, None, :] * stride_su)+ S_blk = tl.load(s_ptrs,+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),+ other=0.0).to(tl.float32)+ diff = S_blk - mean[:, :, None]+ sum_sqdiff += tl.sum(diff * diff, axis=2)- 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}+ var = sum_sqdiff / U_f32+ inv_std = 1.0 / tl.sqrt(var + 1.0e-5)- out_triton = kernel_function(input_tensor, mask, weights, config)+ # Accumulator for final projection to D: [BLOCK_I, BLOCK_J, BLOCK_D]+ acc_out = tl.zeros((BLOCK_I, BLOCK_J, BLOCK_D), dtype=tl.float32)- # 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+ # Pass 3: normalize + affine + gate + matmul with W to D+ for u0 in tl.range(0, U, BLOCK_U):+ u_ids = u0 + tl.arange(0, BLOCK_U)+ u_mask = u_ids < U- if torch.allclose(out_triton, ref, rtol=2e-2, atol=2e-2):- print("PASS")+ s_ptrs = (S_ptr+ + b * stride_sb+ + i_offs[:, None, None] * stride_si+ + j_offs[None, :, None] * stride_sj+ + u_ids[None, None, :] * stride_su)+ g_ptrs = (G_ptr+ + b * stride_gb+ + i_offs[:, None, None] * stride_gi+ + j_offs[None, :, None] * stride_gj+ + u_ids[None, None, :] * stride_gu)++ gamma_blk = tl.load(gamma_ptr + u_ids, mask=u_mask, other=1.0).to(tl.float32)+ beta_blk = tl.load(beta_ptr + u_ids, mask=u_mask, other=0.0).to(tl.float32)++ S_blk = tl.load(s_ptrs,+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),+ other=0.0).to(tl.float32)+ G_blk = tl.load(g_ptrs,+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (u_mask[None, None, :]),+ other=0.0).to(tl.float32)++ # Normalize + affine + gate+ S_norm = (S_blk - mean[:, :, None]) * inv_std[:, :, None]+ S_affine = S_norm * gamma_blk[None, None, :] + beta_blk[None, None, :]+ Z = S_affine * G_blk # [BLOCK_I, BLOCK_J, BLOCK_U]++ # W_ptr is [D, U]; submatrix shaped [BLOCK_U, BLOCK_D] storing W^T+ W_sub = tl.load(W_ptr + (d_offs[None, :] * U + u_ids[:, None]),+ mask=(d_offs[None, :] < D) & (u_mask[:, None]),+ other=0.0).to(tl.float32)++ # Flatten Z to 2D: [(BLOCK_I*BLOCK_J), BLOCK_U] then matmul+ Z_flat = tl.reshape(Z, (BLOCK_I * BLOCK_J, BLOCK_U))+ tmp = tl.dot(Z_flat, W_sub) # [BLOCK_I*BLOCK_J, BLOCK_D]+ tmp_3d = tl.reshape(tmp, (BLOCK_I, BLOCK_J, BLOCK_D))+ acc_out += tmp_3d++ out_ptrs = (Out_ptr+ + b * (N * N * D)+ + i_offs[:, None, None] * (N * D)+ + j_offs[None, :, None] * D+ + d_offs[None, None, :])+ tl.store(out_ptrs, acc_out,+ mask=(i_offs[:, None, None] < N) & (j_offs[None, :, None] < N) & (d_offs[None, None, :] < D))+++ def _check_inputs(x, mask, weights, config):+ assert isinstance(x, torch.Tensor) and x.is_cuda and x.dtype == torch.float32 and x.is_contiguous()+ assert isinstance(mask, torch.Tensor) and mask.is_cuda and mask.dtype == torch.float32 and mask.is_contiguous()+ assert x.ndim == 4+ B, N, N2, D = x.shape+ assert N == N2, "x must be [B, N, N, D]"+ assert mask.shape == (B, N, N)+ assert isinstance(weights, dict)+ assert isinstance(config, dict)+ U = int(config["hidden_dim"])+ assert "dim" in config and int(config["dim"]) == D++ def _ck(name, shape):+ t = weights.get(name, None)+ assert t is not None and isinstance(t, torch.Tensor) and t.is_cuda and t.dtype == torch.float32+ assert t.is_contiguous()+ assert tuple(t.shape) == tuple(shape)++ _ck("norm.weight", (D,))+ _ck("norm.bias", (D,))+ _ck("left_proj.weight", (U, D))+ _ck("right_proj.weight", (U, D))+ _ck("left_gate.weight", (U, D))+ _ck("right_gate.weight", (U, D))+ _ck("out_gate.weight", (U, D))+ _ck("to_out_norm.weight", (U,))+ _ck("to_out_norm.bias", (U,))+ _ck("to_out.weight", (D, U))+++ def kernel_function(*args):+ # Accept either a single tuple or unpacked args+ if len(args) == 1 and isinstance(args[0], tuple):+ x, mask, weights, config = args[0]else:- print(f"FAIL: max diff = {(out_triton - ref).abs().max().item()}")+ x, mask, weights, config = args+ _check_inputs(x, mask, weights, config)+ B, N, _, D = x.shape+ U = int(config["hidden_dim"])+ device = x.device- if __name__ == "__main__":- test_kernel()+ # Allocate intermediates+ x_norm = torch.empty_like(x, dtype=torch.float32, device=device)+ M = B * N * N+ x_flat = x.view(M, D)+ x_norm_flat = x_norm.view(M, D)+ # 1) LayerNorm over last dim+ eps = 1e-5+ BLOCK_LN = 128+ grid_ln = (M,)+ layer_norm_lastdim_kernel[grid_ln](+ x_flat, x_norm_flat, weights["norm.weight"], weights["norm.bias"],+ M, D, eps,+ BLOCK_D=BLOCK_LN,+ num_warps=4, num_stages=2+ )+ # 2) left and right features via fused proj+gate (+ mask) over rows [B, N, N].+ left = torch.empty((B, N, N, U), device=device, dtype=torch.float32)+ right = torch.empty((B, N, N, U), device=device, dtype=torch.float32)+ left_flat = left.view(M, U)+ right_flat = right.view(M, U)++ BLOCK_M = 64+ BLOCK_N = 64+ BLOCK_K = 64+ grid_pg = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))++ # LEFT on rows (b,i,j) with mask applied per row+ proj_gate_fused_rows_kernel[grid_pg](+ x_norm_flat, weights["left_proj.weight"], weights["left_gate.weight"], mask, left_flat,+ M, N, D, U,+ apply_mask=True,+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,+ num_warps=4, num_stages=2+ )+ # RIGHT on rows (b,i,j) with its mask applied per row+ proj_gate_fused_rows_kernel[grid_pg](+ x_norm_flat, weights["right_proj.weight"], weights["right_gate.weight"], mask, right_flat,+ M, N, D, U,+ apply_mask=True,+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,+ num_warps=4, num_stages=2+ )++ # 3) Contraction over k to produce S: [B, N, N, U]+ S = torch.empty((B, N, N, U), device=device, dtype=torch.float32)++ # Strides for [B, N, N, U] contiguous+ stride_b = N * N * U+ stride_i = N * U+ stride_j = U+ stride_u = 1++ BLOCK_I = 32+ BLOCK_J = 32+ BLOCK_KC = 64+ grid_contract = (triton.cdiv(N, BLOCK_I), triton.cdiv(N, BLOCK_J), B * U)+ contract_k_kernel[grid_contract](+ left, right, S,+ B, N, U,+ stride_b, stride_i, stride_j, stride_u, # L strides: b,i,k(==j),u+ stride_b, stride_i, stride_j, stride_u, # R strides: b,j(==i),k(==j),u+ stride_b, stride_i, stride_j, stride_u, # S strides: b,i,j,u+ BLOCK_I=BLOCK_I, BLOCK_J=BLOCK_J, BLOCK_K=BLOCK_KC,+ num_warps=4, num_stages=2+ )++ # 4) Compute out_gate: G = sigmoid(x_norm @ W_out_gate^T)+ G = torch.empty((B, N, N, U), device=device, dtype=torch.float32)+ G_flat = G.view(M, U)+ grid_gate = (triton.cdiv(M, BLOCK_M), triton.cdiv(U, BLOCK_N))+ out_gate_kernel[grid_gate](+ x_norm_flat, weights["out_gate.weight"], G_flat, M, D, U,+ BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,+ num_warps=4, num_stages=2+ )++ # 5) Final projection with LayerNorm(U) + affine + gate + matmul to D+ out = torch.empty((B, N, N, D), device=device, dtype=torch.float32)++ BLOCK_I2 = 16+ BLOCK_J2 = 16+ BLOCK_U2 = 32+ BLOCK_D2 = 64+ grid_final = (triton.cdiv(N, BLOCK_I2), triton.cdiv(N, BLOCK_J2), B * triton.cdiv(D, BLOCK_D2))+ final_project_kernel_affine[grid_final](+ S, G, weights["to_out.weight"], weights["to_out_norm.weight"], weights["to_out_norm.bias"], out,+ B, N, U, D,+ stride_b, stride_i, stride_j, stride_u, # S strides+ stride_b, stride_i, stride_j, stride_u, # G strides+ BLOCK_I=BLOCK_I2, BLOCK_J=BLOCK_J2, BLOCK_U=BLOCK_U2, BLOCK_D=BLOCK_D2,+ num_warps=4, num_stages=2+ )++ return out+def custom_kernel(input):return kernel_function(*input)
scrolls · 684 diff lines total
Best evidence level for this revision: reported
JSON