submission 35764
msuiche · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 376 lines, June 9 Researcher Reciprocity License v1.0.
triton_optimized_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-35764?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:31ee82a2e8db09b7e8d173c394498cf3adf8e8711f120e3d161cc53589c45ae4
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
def epilogue_ln_gate_kernel(mma
acc_l += tl.dot(X_blk, tl.trans(LW_blk), allow_tf32=True)num-warps = 4
num_warps=4, num_stages=2,stages = 2
num_warps=4, num_stages=2,tile-k = 32
BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,tile-m = 64
BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,tile-n = 64
BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,Kernel source
triton_optimized_v3.py376 lines
#!POPCORN leaderboard trimul
#!POPCORN gpu H100
from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
import math
# Keep harness globals unchanged
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = False
# ============================================================
# 1) Fused 5× projections + gates + mask
# ============================================================
@triton.jit
def proj5_gated_mask_kernel(
X_ptr, # float32 [M, D]
LW_ptr, RW_ptr, LGW_ptr, RGW_ptr, OGW_ptr, # float32 [H, D]
MASK_ptr, # float32 [M] (0/1)
LEFT_ptr, RIGHT_ptr, OG_ptr, # float32 [M, H]
M, D, H,
stride_x_m, stride_x_d,
stride_w_h, stride_w_d,
stride_o_m, stride_o_h,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_h = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
m_mask = offs_m < M
h_mask = offs_h < H
acc_l = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_r = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_lg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_rg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_og = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
tl.multiple_of(offs_k, 16)
tl.multiple_of(offs_h, 16)
num_k = tl.cdiv(D, BLOCK_K)
for kb in range(num_k):
k = kb * BLOCK_K + offs_k
k_mask = k < D
# X tile [M, K]
x_ptrs = X_ptr + offs_m[:, None] * stride_x_m + k[None, :] * stride_x_d
X_blk = tl.load(x_ptrs, mask=(m_mask[:, None] & k_mask[None, :]), other=0.0)
# Five weight tiles [H, K]
lw_ptrs = LW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
rw_ptrs = RW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
lgw_ptrs = LGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
rgw_ptrs = RGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
ogw_ptrs = OGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d
LW_blk = tl.load(lw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
RW_blk = tl.load(rw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
LGW_blk = tl.load(lgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
RGW_blk = tl.load(rgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
OGW_blk = tl.load(ogw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)
# FP32 matmul (no TF32)
acc_l += tl.dot(X_blk, tl.trans(LW_blk), allow_tf32=True)
acc_r += tl.dot(X_blk, tl.trans(RW_blk), allow_tf32=True)
acc_lg += tl.dot(X_blk, tl.trans(LGW_blk), allow_tf32=True)
acc_rg += tl.dot(X_blk, tl.trans(RGW_blk), allow_tf32=True)
acc_og += tl.dot(X_blk, tl.trans(OGW_blk), allow_tf32=True)
# Gates + mask
lgate = tl.sigmoid(acc_lg)
rgate = tl.sigmoid(acc_rg)
ogate = tl.sigmoid(acc_og)
mval = tl.load(MASK_ptr + offs_m, mask=m_mask, other=0.0) # [M]
mval = mval[:, None] # [M,1]
left = acc_l * lgate * mval
right = acc_r * rgate * mval
# Stores
left_ptrs = LEFT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
right_ptrs = RIGHT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
og_ptrs = OG_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h
tl.store(left_ptrs, left, mask=(m_mask[:, None] & h_mask[None, :]))
tl.store(right_ptrs, right, mask=(m_mask[:, None] & h_mask[None, :]))
tl.store(og_ptrs, ogate, mask=(m_mask[:, None] & h_mask[None, :]))
# ============================================================
# 2) Contraction: EIN[b,i,j,h] = sum_k LEFT[b,i,k,h] * RIGHT[b,j,k,h]
# Vectorized: broadcast over I/J, reduce over K (no per-h indexing)
# ============================================================
@triton.jit
def contraction_kernel(
LEFT_ptr, RIGHT_ptr, OUT_ptr, # float32
B, N, H,
stride_l_b, stride_l_i, stride_l_k, stride_l_h,
stride_r_b, stride_r_j, stride_r_k, stride_r_h,
stride_o_b, stride_o_i, stride_o_j, stride_o_h,
BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr,
BLOCK_K: tl.constexpr, BLOCK_H: tl.constexpr,
):
# Grid mapping: x-dim covers (b, h-tile); y -> i-tiles; z -> j-tiles
pid_bh = tl.program_id(0)
pid_i = tl.program_id(1)
pid_j = tl.program_id(2)
# Decode i/j tiles
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
# Decode (b, h_start) from pid_bh
tiles_h = (H + BLOCK_H - 1) // BLOCK_H # runtime integer ok
b = pid_bh // tiles_h
h_tile = pid_bh % tiles_h
h_start = h_tile * BLOCK_H
# Iterate over the H micro-tile with compile-time unrolling
for h_rel in tl.static_range(0, BLOCK_H):
h = h_start + h_rel
h_valid = h < H
# Accumulator for this single h
acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)
# Stream over K dimension
for k0 in range(0, N, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
mask_k = offs_k < N
# LEFT[b, i, k, h] -> [I, K]
l_ptrs = (LEFT_ptr
+ b * stride_l_b
+ offs_i[:, None] * stride_l_i
+ offs_k[None, :] * stride_l_k
+ h * stride_l_h)
L = tl.load(l_ptrs,
mask=(mask_i[:, None] & mask_k[None, :] & h_valid),
other=0.0)
# RIGHT[b, j, k, h] -> [J, K]
r_ptrs = (RIGHT_ptr
+ b * stride_r_b
+ offs_j[:, None] * stride_r_j
+ offs_k[None, :] * stride_r_k
+ h * stride_r_h)
R = tl.load(r_ptrs,
mask=(mask_j[:, None] & mask_k[None, :] & h_valid),
other=0.0)
# Use TF32 on tensor cores where available
acc += tl.dot(L, tl.trans(R), allow_tf32=True)
# Store EIN[b, i, j, h] for this h
o_ptrs = (OUT_ptr
+ b * stride_o_b
+ offs_i[:, None] * stride_o_i
+ offs_j[None, :] * stride_o_j
+ h * stride_o_h)
tl.store(o_ptrs, acc, mask=(mask_i[:, None] & mask_j[None, :] & h_valid))
# ============================================================
# 3) Epilogue: LN over H (no clamp; eps=1e-5) -> * out_gate_sigmoid -> final W[D,H]
# ============================================================
@triton.jit
def epilogue_ln_gate_kernel(
EIN_ptr, OG_ptr, # float32 [B, N, N, H]
LN_w_ptr, LN_b_ptr, # float32 [H]
G_ptr, # float32 [B, N, N, H] (output: ln(ein)*og)
B, N, H,
stride_e_b, stride_e_i, stride_e_j, stride_e_h,
stride_g_b, stride_g_i, stride_g_j, stride_g_h,
BLOCK_H: tl.constexpr,
):
pid_pos = tl.program_id(0) # (b,i,j)
total_pos = B * N * N
if pid_pos >= total_pos:
return
b = pid_pos // (N * N)
rem = pid_pos % (N * N)
i = rem // N
j = rem % N
# Stats over H (no clamping)
sum_x = 0.0
sum_x2 = 0.0
for h0 in range(0, H, BLOCK_H):
offs_h = h0 + tl.arange(0, BLOCK_H)
mask_h = offs_h < H
e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
vals = tl.load(e_ptrs, mask=mask_h, other=0.0) # [BLOCK_H]
sum_x += tl.sum(vals)
sum_x2 += tl.sum(vals * vals)
Hf = tl.full((1,), H, tl.float32)
mean = sum_x / Hf
var = sum_x2 / Hf - mean * mean
inv_std = tl.rsqrt(var + 1e-5)
# Write normalized-and-gated vector to G
for h0 in range(0, H, BLOCK_H):
offs_h = h0 + tl.arange(0, BLOCK_H)
mask_h = offs_h < H
e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
og_ptrs = OG_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h
g_ptrs = G_ptr + b * stride_g_b + i * stride_g_i + j * stride_g_j + offs_h * stride_g_h
lnw = tl.load(LN_w_ptr + offs_h, mask=mask_h, other=1.0)
lnb = tl.load(LN_b_ptr + offs_h, mask=mask_h, other=0.0)
ein = tl.load(e_ptrs, mask=mask_h, other=0.0)
og = tl.load(og_ptrs, mask=mask_h, other=0.0) # already sigmoid'd
normed = ((ein - mean) * inv_std) * lnw + lnb
gated = normed * og # [BLOCK_H]
tl.store(g_ptrs, gated, mask=mask_h)
# ============================================================
# Python wrapper
# ============================================================
def custom_kernel(data: input_t) -> output_t:
with DisableCuDNNTF32():
input_tensor, mask, weights, config = data
B, N, _, D = input_tensor.shape
H = config["hidden_dim"]
# Prefer Tensor Cores / TF32 for speed on Ampere+/Hopper
prev_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = True
prev_prec = torch.get_float32_matmul_precision() if hasattr(torch, "get_float32_matmul_precision") else None
if hasattr(torch, "set_float32_matmul_precision"):
torch.set_float32_matmul_precision("high")
try:
# 0) Input LayerNorm (fused), FP32, eps=1e-5 (no clamp)
x = F.layer_norm(
input_tensor, (D,),
weight=weights["norm.weight"],
bias=weights["norm.bias"],
eps=1e-5,
).contiguous()
# Flatten to [M, D]
M = B * N * N
x2d = x.view(M, D)
mask_f = mask.to(dtype=torch.float32).reshape(M).contiguous()
# Contiguous weights
LW = weights['left_proj.weight' ].contiguous() # [H,D]
RW = weights['right_proj.weight'].contiguous()
LGW = weights['left_gate.weight' ].contiguous()
RGW = weights['right_gate.weight'].contiguous()
OGW = weights['out_gate.weight' ].contiguous()
# Outputs of projection kernel
LEFT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
RIGHT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
OG2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)
# Launch fused projections (small tiles to fit SMEM)
grid_proj = (triton.cdiv(M, 64), triton.cdiv(H, 64))
proj5_gated_mask_kernel[grid_proj](
x2d, LW, RW, LGW, RGW, OGW, mask_f,
LEFT2D, RIGHT2D, OG2D,
M, D, H,
x2d.stride(0), x2d.stride(1),
LW.stride(0), LW.stride(1),
LEFT2D.stride(0), LEFT2D.stride(1),
BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,
num_warps=4, num_stages=2,
)
LEFT = LEFT2D.view(B, N, N, H)
RIGHT = RIGHT2D.view(B, N, N, H)
OG = OG2D.view(B, N, N, H)
# Contraction via batched GEMM over (b,h): for each h, L[i,k] @ R[j,k]^T -> [i,j]
Left_h = LEFT.permute(0, 3, 1, 2).contiguous().view(B * H, N, N)
Right_h = RIGHT.permute(0, 3, 2, 1).contiguous().view(B * H, N, N)
EIN_h = torch.bmm(Left_h, Right_h) # [B*H, N, N]
EIN = EIN_h.view(B, H, N, N).permute(0, 2, 3, 1).contiguous()
# Epilogue split: (1) Triton LN+gate per (b,i,j,h) -> G; (2) cuBLAS GEMM G @ W^T
W = weights['to_out.weight' ].contiguous() # [D,H]
LNw = weights['to_out_norm.weight' ].contiguous() # [H]
LNb = weights['to_out_norm.bias' ].contiguous() # [H]
G = torch.empty_like(OG) # [B,N,N,H]
grid_epi = (B * N * N,)
epilogue_ln_gate_kernel[grid_epi](
EIN, OG, LNw, LNb, G,
B, N, H,
EIN.stride(0), EIN.stride(1), EIN.stride(2), EIN.stride(3),
G.stride(0), G.stride(1), G.stride(2), G.stride(3),
BLOCK_H=64,
num_warps=4, num_stages=2,
)
M = B * N * N
OUT2D = torch.matmul(G.view(M, H), W.t()) # [M,D]
OUT = OUT2D.view(B, N, N, D)
return OUT
finally:
torch.backends.cuda.matmul.allow_tf32 = prev_tf32
if hasattr(torch, "set_float32_matmul_precision") and prev_prec is not None:
torch.set_float32_matmul_precision(prev_prec)
# ============================================================
# Input generation (unchanged)
def generate_input(seqlen: int, bs: int, dim: int, hiddendim: int,
seed: int, nomask: bool, distribution: str) -> input_t:
batch_size = bs
seq_len = seqlen
hidden_dim = hiddendim
no_mask = nomask
config = {"hidden_dim": hidden_dim, "dim": dim}
gen = torch.Generator(device='cuda')
gen.manual_seed(seed)
weights = {}
if distribution == "cauchy":
input_tensor = torch.distributions.Cauchy(0, 2).sample(
(batch_size, seq_len, seq_len, dim)
).to(device='cuda', dtype=torch.float32)
else:
input_tensor = torch.randn(
(batch_size, seq_len, seq_len, dim),
device='cuda', dtype=torch.float32, generator=gen
).contiguous()
if no_mask:
mask = torch.ones(batch_size, seq_len, seq_len, device=input_tensor.device)
else:
mask = torch.randint(0, 2, (batch_size, seq_len, seq_len),
device=input_tensor.device, generator=gen)
weights["norm.weight"] = torch.randn(dim, device="cuda", dtype=torch.float32)
weights["norm.bias"] = torch.randn(dim, device="cuda", dtype=torch.float32)
weights["left_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["right_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["left_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["right_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["out_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)
weights["to_out_norm.weight"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
weights["to_out.weight"] = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)
weights["to_out_norm.bias"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)
return (input_tensor, mask, weights, config)
# Correctness check
check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)
scrolls · 376 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 35747.
⋯ 3 unchanged linesfrom task import input_t, output_timport torch- import torch.nn as nnimport torch.nn.functional as Fimport tritonimport triton.language as tlimport math- # Enable TF32 for H100 tensor cores while maintaining FP32 precision+ # Keep harness globals unchangedtorch.backends.cuda.matmul.allow_tf32 = True- torch.backends.cudnn.allow_tf32 = False # Keep cuDNN FP32 for accuracy+ torch.backends.cudnn.allow_tf32 = False- # Ultra-optimized LayerNorm kernel with fused operations- @triton.jit- def optimized_layernorm_kernel(- x_ptr, ln_w_ptr, ln_b_ptr, y_ptr,- mean_ptr, inv_std_ptr,- N, D,- stride_x_n, stride_x_d,- stride_y_n, stride_y_d,- BLOCK_D: tl.constexpr- ):- pid = tl.program_id(0)- offs_n = pid- offs_d = tl.arange(0, BLOCK_D)-- mask_d = offs_d < D-- # Load input block- x_ptrs = x_ptr + offs_n * stride_x_n + offs_d * stride_x_d- x_block = tl.load(x_ptrs, mask=mask_d, other=0.0)-- # Compute mean and variance in one pass- sum_x = tl.sum(x_block)- sum_x2 = tl.sum(x_block * x_block)-- mean = sum_x / D- var = tl.maximum(sum_x2 / D - mean * mean, 0.0)- inv_std = tl.rsqrt(var + 1e-5)-- # Store statistics for potential reuse- if mean_ptr is not None:- tl.store(mean_ptr + offs_n, mean)- if inv_std_ptr is not None:- tl.store(inv_std_ptr + offs_n, inv_std)-- # Load weights and apply normalization- ln_w = tl.load(ln_w_ptr + offs_d, mask=mask_d, other=1.0)- ln_b = tl.load(ln_b_ptr + offs_d, mask=mask_d, other=0.0)-- # Fused normalization- scale = inv_std * ln_w- y_block = (x_block - mean) * scale + ln_b-- # Store output- y_ptrs = y_ptr + offs_n * stride_y_n + offs_d * stride_y_d- tl.store(y_ptrs, y_block, mask=mask_d)- # Fully fused projection kernel with optimized memory access+ # ============================================================+ # 1) Fused 5× projections + gates + mask+ # ============================================================@triton.jit- def fused_projections_kernel(- x_ptr, weights_ptr, out_ptr,- N, D, H,- stride_x_n, stride_x_d,+ def proj5_gated_mask_kernel(+ X_ptr, # float32 [M, D]+ LW_ptr, RW_ptr, LGW_ptr, RGW_ptr, OGW_ptr, # float32 [H, D]+ MASK_ptr, # float32 [M] (0/1)+ LEFT_ptr, RIGHT_ptr, OG_ptr, # float32 [M, H]+ M, D, H,+ stride_x_m, stride_x_d,stride_w_h, stride_w_d,- stride_out_n, stride_out_h,- BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr+ stride_o_m, stride_o_h,+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,):pid_m = tl.program_id(0)pid_n = tl.program_id(1)-+offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)- offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)+ offs_h = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)offs_k = tl.arange(0, BLOCK_K)-- mask_m = offs_m < N- mask_n = offs_n < (5 * H) # 5 projections: left, right, left_gate, right_gate, out_gate-- acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)-- num_k_blocks = tl.cdiv(D, BLOCK_K)- for kb in range(num_k_blocks):- k_idx = kb * BLOCK_K + offs_k- valid_k = k_idx < D-- # Load input block with vectorized access- x_ptrs = x_ptr + offs_m[:, None] * stride_x_n + k_idx[None, :] * stride_x_d- x_block = tl.load(x_ptrs, mask=(mask_m[:, None] & valid_k[None, :]), other=0.0)-- # Load weight block with coalesced access- w_ptrs = weights_ptr + offs_n[:, None] * stride_w_h + k_idx[None, :] * stride_w_d- w_block = tl.load(w_ptrs, mask=(mask_n[:, None] & valid_k[None, :]), other=0.0)-- # Accumulate using tensor cores when available- acc += tl.dot(x_block, tl.trans(w_block), allow_tf32=False)-- # Store output with vectorized writes- out_ptrs = out_ptr + offs_m[:, None] * stride_out_n + offs_n[None, :] * stride_out_h- store_mask = mask_m[:, None] & mask_n[None, :]- tl.store(out_ptrs, acc, mask=store_mask)- # Highly optimized einsum kernel with tiling and vectorization+ m_mask = offs_m < M+ h_mask = offs_h < H++ acc_l = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ acc_r = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ acc_lg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ acc_rg = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ acc_og = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)++ tl.multiple_of(offs_k, 16)+ tl.multiple_of(offs_h, 16)++ num_k = tl.cdiv(D, BLOCK_K)+ for kb in range(num_k):+ k = kb * BLOCK_K + offs_k+ k_mask = k < D++ # X tile [M, K]+ x_ptrs = X_ptr + offs_m[:, None] * stride_x_m + k[None, :] * stride_x_d+ X_blk = tl.load(x_ptrs, mask=(m_mask[:, None] & k_mask[None, :]), other=0.0)++ # Five weight tiles [H, K]+ lw_ptrs = LW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d+ rw_ptrs = RW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d+ lgw_ptrs = LGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d+ rgw_ptrs = RGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d+ ogw_ptrs = OGW_ptr + offs_h[:, None] * stride_w_h + k[None, :] * stride_w_d++ LW_blk = tl.load(lw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)+ RW_blk = tl.load(rw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)+ LGW_blk = tl.load(lgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)+ RGW_blk = tl.load(rgw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)+ OGW_blk = tl.load(ogw_ptrs, mask=(h_mask[:, None] & k_mask[None, :]), other=0.0)++ # FP32 matmul (no TF32)+ acc_l += tl.dot(X_blk, tl.trans(LW_blk), allow_tf32=True)+ acc_r += tl.dot(X_blk, tl.trans(RW_blk), allow_tf32=True)+ acc_lg += tl.dot(X_blk, tl.trans(LGW_blk), allow_tf32=True)+ acc_rg += tl.dot(X_blk, tl.trans(RGW_blk), allow_tf32=True)+ acc_og += tl.dot(X_blk, tl.trans(OGW_blk), allow_tf32=True)++ # Gates + mask+ lgate = tl.sigmoid(acc_lg)+ rgate = tl.sigmoid(acc_rg)+ ogate = tl.sigmoid(acc_og)++ mval = tl.load(MASK_ptr + offs_m, mask=m_mask, other=0.0) # [M]+ mval = mval[:, None] # [M,1]++ left = acc_l * lgate * mval+ right = acc_r * rgate * mval++ # Stores+ left_ptrs = LEFT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h+ right_ptrs = RIGHT_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h+ og_ptrs = OG_ptr + offs_m[:, None] * stride_o_m + offs_h[None, :] * stride_o_h++ tl.store(left_ptrs, left, mask=(m_mask[:, None] & h_mask[None, :]))+ tl.store(right_ptrs, right, mask=(m_mask[:, None] & h_mask[None, :]))+ tl.store(og_ptrs, ogate, mask=(m_mask[:, None] & h_mask[None, :]))+++ # ============================================================+ # 2) Contraction: EIN[b,i,j,h] = sum_k LEFT[b,i,k,h] * RIGHT[b,j,k,h]+ # Vectorized: broadcast over I/J, reduce over K (no per-h indexing)+ # ============================================================@triton.jit- def optimized_einsum_kernel(- left_ptr, right_ptr, out_ptr,+ def contraction_kernel(+ LEFT_ptr, RIGHT_ptr, OUT_ptr, # float32B, N, H,- stride_left_b, stride_left_i, stride_left_k, stride_left_h,- stride_right_b, stride_right_j, stride_right_k, stride_right_h,- stride_out_b, stride_out_i, stride_out_j, stride_out_h,- BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_H: tl.constexpr+ stride_l_b, stride_l_i, stride_l_k, stride_l_h,+ stride_r_b, stride_r_j, stride_r_k, stride_r_h,+ stride_o_b, stride_o_i, stride_o_j, stride_o_h,+ BLOCK_I: tl.constexpr, BLOCK_J: tl.constexpr,+ BLOCK_K: tl.constexpr, BLOCK_H: tl.constexpr,):- pid_b = tl.program_id(0)- pid_i = tl.program_id(1)- pid_j = tl.program_id(2)-+ # Grid mapping: x-dim covers (b, h-tile); y -> i-tiles; z -> j-tiles+ pid_bh = tl.program_id(0)+ pid_i = tl.program_id(1)+ pid_j = tl.program_id(2)++ # Decode i/j tilesoffs_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 < Nmask_j = offs_j < N-- # Process H dimension in blocks for better vectorization- for h_start in range(0, H, BLOCK_H):- offs_h = h_start + tl.arange(0, BLOCK_H)- mask_h = offs_h < H-- # Initialize accumulator for this H block- acc = tl.zeros((BLOCK_I, BLOCK_J, BLOCK_H), dtype=tl.float32)-- # Process K dimension in blocks for cache efficiency- for k_start in range(0, N, BLOCK_K):- offs_k = k_start + tl.arange(0, BLOCK_K)++ # Decode (b, h_start) from pid_bh+ tiles_h = (H + BLOCK_H - 1) // BLOCK_H # runtime integer ok+ b = pid_bh // tiles_h+ h_tile = pid_bh % tiles_h+ h_start = h_tile * BLOCK_H++ # Iterate over the H micro-tile with compile-time unrolling+ for h_rel in tl.static_range(0, BLOCK_H):+ h = h_start + h_rel+ h_valid = h < H++ # Accumulator for this single h+ acc = tl.zeros((BLOCK_I, BLOCK_J), dtype=tl.float32)++ # Stream over K dimension+ for k0 in range(0, N, BLOCK_K):+ offs_k = k0 + tl.arange(0, BLOCK_K)mask_k = offs_k < N-- # Load left block: [BLOCK_I, BLOCK_K, BLOCK_H]- left_ptrs = left_ptr + pid_b * stride_left_b + \- offs_i[:, None, None] * stride_left_i + \- offs_k[None, :, None] * stride_left_k + \- offs_h[None, None, :] * stride_left_h- left_block = tl.load(left_ptrs,- mask=mask_i[:, None, None] & mask_k[None, :, None] & mask_h[None, None, :],- other=0.0)-- # Load right block: [BLOCK_J, BLOCK_K, BLOCK_H]- right_ptrs = right_ptr + pid_b * stride_right_b + \- offs_j[:, None, None] * stride_right_j + \- offs_k[None, :, None] * stride_right_k + \- offs_h[None, None, :] * stride_right_h- right_block = tl.load(right_ptrs,- mask=mask_j[:, None, None] & mask_k[None, :, None] & mask_h[None, None, :],- other=0.0)-- # Accumulate using vectorized operations- # Sum over K dimension: left[i,k,h] * right[j,k,h] -> out[i,j,h]- acc += tl.sum(left_block[:, :, None, :] * right_block[None, :, :, :], axis=2)-- # Store output block with vectorized writes- out_ptrs = out_ptr + pid_b * stride_out_b + \- offs_i[:, None, None] * stride_out_i + \- offs_j[None, :, None] * stride_out_j + \- offs_h[None, None, :] * stride_out_h- tl.store(out_ptrs, acc,- mask=mask_i[:, None, None] & mask_j[None, :, None] & mask_h[None, None, :])- # Fused output processing kernel+ # LEFT[b, i, k, h] -> [I, K]+ l_ptrs = (LEFT_ptr+ + b * stride_l_b+ + offs_i[:, None] * stride_l_i+ + offs_k[None, :] * stride_l_k+ + h * stride_l_h)+ L = tl.load(l_ptrs,+ mask=(mask_i[:, None] & mask_k[None, :] & h_valid),+ other=0.0)++ # RIGHT[b, j, k, h] -> [J, K]+ r_ptrs = (RIGHT_ptr+ + b * stride_r_b+ + offs_j[:, None] * stride_r_j+ + offs_k[None, :] * stride_r_k+ + h * stride_r_h)+ R = tl.load(r_ptrs,+ mask=(mask_j[:, None] & mask_k[None, :] & h_valid),+ other=0.0)++ # Use TF32 on tensor cores where available+ acc += tl.dot(L, tl.trans(R), allow_tf32=True)++ # Store EIN[b, i, j, h] for this h+ o_ptrs = (OUT_ptr+ + b * stride_o_b+ + offs_i[:, None] * stride_o_i+ + offs_j[None, :] * stride_o_j+ + h * stride_o_h)+ tl.store(o_ptrs, acc, mask=(mask_i[:, None] & mask_j[None, :] & h_valid))+++ # ============================================================+ # 3) Epilogue: LN over H (no clamp; eps=1e-5) -> * out_gate_sigmoid -> final W[D,H]+ # ============================================================@triton.jit- def fused_output_kernel(- einsum_ptr, out_gate_ptr,- norm_w_ptr, norm_b_ptr, to_out_w_ptr,- out_ptr,- B, N, D, H,- stride_ein_b, stride_ein_i, stride_ein_j, stride_ein_h,- stride_out_b, stride_out_i, stride_out_j, stride_out_d,- BLOCK_SIZE: tl.constexpr+ def epilogue_ln_gate_kernel(+ EIN_ptr, OG_ptr, # float32 [B, N, N, H]+ LN_w_ptr, LN_b_ptr, # float32 [H]+ G_ptr, # float32 [B, N, N, H] (output: ln(ein)*og)+ B, N, H,+ stride_e_b, stride_e_i, stride_e_j, stride_e_h,+ stride_g_b, stride_g_i, stride_g_j, stride_g_h,+ BLOCK_H: tl.constexpr,):- pid = tl.program_id(0)-- # Calculate which (b,i,j) position we're processing- total_positions = B * N * N- if pid >= total_positions:+ pid_pos = tl.program_id(0) # (b,i,j)++ total_pos = B * N * N+ if pid_pos >= total_pos:return-- b = pid // (N * N)- rem = pid % (N * N)++ b = pid_pos // (N * N)+ rem = pid_pos % (N * N)i = rem // Nj = rem % N-- # Process D dimension in blocks- for d_start in range(0, D, BLOCK_SIZE):- offs_d = d_start + tl.arange(0, BLOCK_SIZE)- mask_d = offs_d < D-- # Load einsum output and apply LayerNorm- h_offs = tl.arange(0, H)-- einsum_ptrs = einsum_ptr + b * stride_ein_b + i * stride_ein_i + j * stride_ein_j + h_offs * stride_ein_h- einsum_vals = tl.load(einsum_ptrs, mask=h_offs < H, other=0.0)-- # Compute LayerNorm statistics- mean = tl.sum(einsum_vals) / H- var = tl.sum((einsum_vals - mean) * (einsum_vals - mean)) / H- inv_std = tl.rsqrt(var + 1e-5)-- # Load normalization weights- ln_w = tl.load(norm_w_ptr + h_offs, mask=h_offs < H, other=1.0)- ln_b = tl.load(norm_b_ptr + h_offs, mask=h_offs < H, other=0.0)-- # Apply normalization- normed = (einsum_vals - mean) * inv_std * ln_w + ln_b-- # Load and apply out_gate- gate_vals = tl.load(out_gate_ptr + b * N * N * H + i * N * H + j * H + h_offs, mask=h_offs < H, other=0.0)- gated = normed * tl.sigmoid(gate_vals)-- # Final projection to output dimension- to_out_ptrs = to_out_w_ptr + offs_d[:, None] * H + h_offs[None, :]- to_out_vals = tl.load(to_out_ptrs, mask=mask_d[:, None] & (h_offs < H)[None, :], other=0.0)-- # Matrix multiplication: gated @ to_out_w.T- output_vals = tl.sum(to_out_vals * gated[None, :], axis=1)-- # Store final output- out_ptrs = out_ptr + b * stride_out_b + i * stride_out_i + j * stride_out_j + offs_d * stride_out_d- tl.store(out_ptrs, output_vals, mask=mask_d)- class H100OptimizedTriMul(nn.Module):- """- H100-optimized TriMul with tensor core acceleration and aggressive fusion- """-- def __init__(self, dim: int, hidden_dim: int):- super().__init__()- self.dim = dim- self.hidden_dim = hidden_dim-- # Single fused weight matrix for maximum tensor core utilization- self.fused_weights = nn.Parameter(torch.empty(5 * hidden_dim, dim))-- # LayerNorm parameters (weights applied separately for flexibility)- self.norm_weight = nn.Parameter(torch.ones(dim))- self.norm_bias = nn.Parameter(torch.zeros(dim))- self.out_norm_weight = nn.Parameter(torch.ones(hidden_dim))- self.out_norm_bias = nn.Parameter(torch.zeros(hidden_dim))-- # Final projection weight- self.to_out_weight = nn.Parameter(torch.empty(dim, hidden_dim))-- # Initialize weights for better numerical stability- nn.init.kaiming_normal_(self.fused_weights, mode='fan_out', nonlinearity='linear')- nn.init.kaiming_normal_(self.to_out_weight, mode='fan_out', nonlinearity='linear')-- # Note: torch.compile disabled due to compilation overhead outweighing benefits- # The model already performs well with the other optimizations applied-- def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:- B, N, _, D = x.shape- H = self.hidden_dim-- # Step 1: Apply LayerNorm using PyTorch's highly optimized implementation- # Reshape once and keep in flat form for subsequent operations- x_flat = x.reshape(B * N * N, D)- x_norm_flat = F.layer_norm(x_flat, normalized_shape=(D,),- weight=self.norm_weight, bias=self.norm_bias, eps=1e-5)-- # Step 2: Ultra-fused projections using single large matmul- # Use the already flattened tensor - avoid extra reshape- all_projections = torch.mm(x_norm_flat, self.fused_weights.T)- all_projections = all_projections.view(B, N, N, 5, H)-- # Extract components efficiently- left_proj = all_projections[..., 0, :]- right_proj = all_projections[..., 1, :]- left_gate = all_projections[..., 2, :]- right_gate = all_projections[..., 3, :]- out_gate = all_projections[..., 4, :]-- # Apply sigmoid gates with efficient computation- left_gate = torch.sigmoid(left_gate)- right_gate = torch.sigmoid(right_gate)- out_gate = torch.sigmoid(out_gate)-- # Apply mask and gates efficiently- mask_expanded = mask.unsqueeze(-1)- left = left_proj * mask_expanded * left_gate- right = right_proj * mask_expanded * right_gate-- # Ensure contiguous layout for optimal einsum performance- left = left.contiguous()- right = right.contiguous()-- # Step 3: H100-optimized einsum - keeping einsum as it's already well optimized- # The einsum 'bikd,bjkd->bijd' computes: for each batch, sum over k dimension- # torch.einsum is highly optimized on H100 with TF32, so we keep it- einsum_out = torch.einsum('bikd,bjkd->bijd', left, right)-- # Step 4: Fused output processing with minimal reshapes- # Reshape once for LayerNorm and keep flat for final operations- einsum_flat = einsum_out.reshape(B * N * N, H)- normed_out_flat = F.layer_norm(einsum_flat, normalized_shape=(H,),- weight=self.out_norm_weight, bias=self.out_norm_bias, eps=1e-5)-- # Apply out_gate (already in correct shape from earlier extraction)- gated_flat = normed_out_flat * out_gate.reshape(B * N * N, H)-- # Final projection using tensor cores - output already in correct shape- output_flat = torch.mm(gated_flat, self.to_out_weight.T)- output = output_flat.view(B, N, N, D)-- return output+ # Stats over H (no clamping)+ sum_x = 0.0+ sum_x2 = 0.0+ for h0 in range(0, H, BLOCK_H):+ offs_h = h0 + tl.arange(0, BLOCK_H)+ mask_h = offs_h < H+ e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h+ vals = tl.load(e_ptrs, mask=mask_h, other=0.0) # [BLOCK_H]+ sum_x += tl.sum(vals)+ sum_x2 += tl.sum(vals * vals)+ Hf = tl.full((1,), H, tl.float32)+ mean = sum_x / Hf+ var = sum_x2 / Hf - mean * mean+ inv_std = tl.rsqrt(var + 1e-5)+ # Write normalized-and-gated vector to G+ for h0 in range(0, H, BLOCK_H):+ offs_h = h0 + tl.arange(0, BLOCK_H)+ mask_h = offs_h < H++ e_ptrs = EIN_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h+ og_ptrs = OG_ptr + b * stride_e_b + i * stride_e_i + j * stride_e_j + offs_h * stride_e_h+ g_ptrs = G_ptr + b * stride_g_b + i * stride_g_i + j * stride_g_j + offs_h * stride_g_h++ lnw = tl.load(LN_w_ptr + offs_h, mask=mask_h, other=1.0)+ lnb = tl.load(LN_b_ptr + offs_h, mask=mask_h, other=0.0)+ ein = tl.load(e_ptrs, mask=mask_h, other=0.0)+ og = tl.load(og_ptrs, mask=mask_h, other=0.0) # already sigmoid'd++ normed = ((ein - mean) * inv_std) * lnw + lnb+ gated = normed * og # [BLOCK_H]++ tl.store(g_ptrs, gated, mask=mask_h)+++ # ============================================================+ # Python wrapper+ # ============================================================def custom_kernel(data: input_t) -> output_t:- """- H100-optimized custom kernel with tensor core acceleration- """with DisableCuDNNTF32():input_tensor, mask, weights, config = data-- dim = config["dim"]- hidden_dim = config["hidden_dim"]-- # Ensure contiguous tensors for H100 memory efficiency- input_tensor = input_tensor.contiguous()- mask = mask.contiguous()-- # Create H100-optimized model- model = H100OptimizedTriMul(dim, hidden_dim).to(input_tensor.device)-- # Optimized weight loading - pre-concatenate and use direct assignment- with torch.no_grad():- # Pre-concatenate all projection/gate weights in one operation- fused_weights_data = torch.cat([- weights['left_proj.weight'],- weights['right_proj.weight'],- weights['left_gate.weight'],- weights['right_gate.weight'],- weights['out_gate.weight']- ], dim=0)-- # Single copy operation for fused weights- model.fused_weights.copy_(fused_weights_data)-- # Direct assignment for remaining weights (minimal overhead)- model.norm_weight[:] = weights['norm.weight']- model.norm_bias[:] = weights['norm.bias']- model.out_norm_weight[:] = weights['to_out_norm.weight']- model.out_norm_bias[:] = weights['to_out_norm.bias']- model.to_out_weight[:] = weights['to_out.weight']-- # Run with H100 optimizations- with torch.no_grad():- output = model(input_tensor, mask)-- return output+ B, N, _, D = input_tensor.shape+ H = config["hidden_dim"]+ # Prefer Tensor Cores / TF32 for speed on Ampere+/Hopper+ prev_tf32 = torch.backends.cuda.matmul.allow_tf32+ torch.backends.cuda.matmul.allow_tf32 = True+ prev_prec = torch.get_float32_matmul_precision() if hasattr(torch, "get_float32_matmul_precision") else None+ if hasattr(torch, "set_float32_matmul_precision"):+ torch.set_float32_matmul_precision("high")+ try:+ # 0) Input LayerNorm (fused), FP32, eps=1e-5 (no clamp)+ x = F.layer_norm(+ input_tensor, (D,),+ weight=weights["norm.weight"],+ bias=weights["norm.bias"],+ eps=1e-5,+ ).contiguous()- # Input generation function (same as reference)+ # Flatten to [M, D]+ M = B * N * N+ x2d = x.view(M, D)+ mask_f = mask.to(dtype=torch.float32).reshape(M).contiguous()++ # Contiguous weights+ LW = weights['left_proj.weight' ].contiguous() # [H,D]+ RW = weights['right_proj.weight'].contiguous()+ LGW = weights['left_gate.weight' ].contiguous()+ RGW = weights['right_gate.weight'].contiguous()+ OGW = weights['out_gate.weight' ].contiguous()++ # Outputs of projection kernel+ LEFT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)+ RIGHT2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)+ OG2D = torch.empty((M, H), device=x2d.device, dtype=torch.float32)++ # Launch fused projections (small tiles to fit SMEM)+ grid_proj = (triton.cdiv(M, 64), triton.cdiv(H, 64))+ proj5_gated_mask_kernel[grid_proj](+ x2d, LW, RW, LGW, RGW, OGW, mask_f,+ LEFT2D, RIGHT2D, OG2D,+ M, D, H,+ x2d.stride(0), x2d.stride(1),+ LW.stride(0), LW.stride(1),+ LEFT2D.stride(0), LEFT2D.stride(1),+ BLOCK_M=64, BLOCK_N=64, BLOCK_K=32,+ num_warps=4, num_stages=2,+ )++ LEFT = LEFT2D.view(B, N, N, H)+ RIGHT = RIGHT2D.view(B, N, N, H)+ OG = OG2D.view(B, N, N, H)++ # Contraction via batched GEMM over (b,h): for each h, L[i,k] @ R[j,k]^T -> [i,j]+ Left_h = LEFT.permute(0, 3, 1, 2).contiguous().view(B * H, N, N)+ Right_h = RIGHT.permute(0, 3, 2, 1).contiguous().view(B * H, N, N)+ EIN_h = torch.bmm(Left_h, Right_h) # [B*H, N, N]+ EIN = EIN_h.view(B, H, N, N).permute(0, 2, 3, 1).contiguous()+++ # Epilogue split: (1) Triton LN+gate per (b,i,j,h) -> G; (2) cuBLAS GEMM G @ W^T+ W = weights['to_out.weight' ].contiguous() # [D,H]+ LNw = weights['to_out_norm.weight' ].contiguous() # [H]+ LNb = weights['to_out_norm.bias' ].contiguous() # [H]++ G = torch.empty_like(OG) # [B,N,N,H]+ grid_epi = (B * N * N,)+ epilogue_ln_gate_kernel[grid_epi](+ EIN, OG, LNw, LNb, G,+ B, N, H,+ EIN.stride(0), EIN.stride(1), EIN.stride(2), EIN.stride(3),+ G.stride(0), G.stride(1), G.stride(2), G.stride(3),+ BLOCK_H=64,+ num_warps=4, num_stages=2,+ )++ M = B * N * N+ OUT2D = torch.matmul(G.view(M, H), W.t()) # [M,D]+ OUT = OUT2D.view(B, N, N, D)+ return OUT+ finally:+ torch.backends.cuda.matmul.allow_tf32 = prev_tf32+ if hasattr(torch, "set_float32_matmul_precision") and prev_prec is not None:+ torch.set_float32_matmul_precision(prev_prec)+++ # ============================================================+ # Input generation (unchanged)def generate_input(seqlen: int, bs: int, dim: int, hiddendim: int,seed: int, nomask: bool, distribution: str) -> input_t:batch_size = bsseq_len = seqlenhidden_dim = hiddendimno_mask = nomask-+config = {"hidden_dim": hidden_dim, "dim": dim}-+gen = torch.Generator(device='cuda')gen.manual_seed(seed)-+weights = {}-+if distribution == "cauchy":input_tensor = torch.distributions.Cauchy(0, 2).sample((batch_size, seq_len, seq_len, dim)⋯ 3 unchanged lines(batch_size, seq_len, seq_len, dim),device='cuda', dtype=torch.float32, generator=gen).contiguous()-+if no_mask:mask = torch.ones(batch_size, seq_len, seq_len, device=input_tensor.device)else:mask = torch.randint(0, 2, (batch_size, seq_len, seq_len),- device=input_tensor.device, generator=gen)-+ device=input_tensor.device, generator=gen)+weights["norm.weight"] = torch.randn(dim, device="cuda", dtype=torch.float32)weights["norm.bias"] = torch.randn(dim, device="cuda", dtype=torch.float32)- weights["left_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)+ weights["left_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)weights["right_proj.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)- weights["left_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)+ weights["left_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)weights["right_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)- weights["out_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)+ weights["out_gate.weight"] = torch.randn(hidden_dim, dim, device="cuda", dtype=torch.float32) / math.sqrt(hidden_dim)weights["to_out_norm.weight"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)- weights["to_out.weight"] = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)- weights["to_out_norm.bias"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)-+ weights["to_out.weight"] = torch.randn(dim, hidden_dim, device="cuda", dtype=torch.float32) / math.sqrt(dim)+ weights["to_out_norm.bias"] = torch.randn(hidden_dim, device="cuda", dtype=torch.float32)+return (input_tensor, mask, weights, config)- # Check implementation correctness- check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)No newline at end of file+ # Correctness check+ check_implementation = make_match_reference(custom_kernel, rtol=2e-2, atol=2e-2)
scrolls · 704 diff lines total
Best evidence level for this revision: reported
JSON