submission 40869
msuiche · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 108 lines, June 9 Researcher Reciprocity License v1.0.
submission_baseline_tuned.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-40869?include=source"interfacepython
Compatibility
measured onAMD Instinct MI300X
declared hardwareAMD Instinct MI300X
architecturesgfx942
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:2073c900dbd960b7d4eeea626d32fa8064c71ebd563cb3afe9b3a155bff3725c
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-15
Kernel source
submission_baseline_tuned.py108 lines
"""
Baseline copy with tuned einsum paths (batched GEMM) and small memory/layout tweaks.
Single-file, entry: custom_kernel(data). Keeps DisableCuDNNTF32 semantics.
"""
import torch
import torch.nn.functional as F
from task import input_t, output_t
from utils import DisableCuDNNTF32
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = False
def _einsum_opt(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
"""Compute einsum('bikh,bjkh->bijh') via batched GEMM on (B,H).
left/right: [B,N,N,H]
returns: [B,N,N,H]
"""
B, N, _, H = left.shape
L = left.permute(0, 3, 1, 2).contiguous() # [B,H,N,N]
R = right.permute(0, 3, 1, 2).contiguous() # [B,H,N,N]
out_bh = torch.matmul(L.bfloat16(), R.transpose(-2, -1).bfloat16()).float() # [B,H,N,N]
return out_bh.permute(0, 2, 3, 1).contiguous()
def _custom_kernel_core(data: input_t) -> output_t:
input_tensor, mask, weights, config = data
B, N, _, D = input_tensor.shape
H = config["hidden_dim"]
device = input_tensor.device
M = B * N * N
# Heuristic low-rank path as in baseline
use_lr = (N >= 512 and H >= 384)
x = F.layer_norm(
input_tensor, (D,),
weight=weights["norm.weight"],
bias=weights["norm.bias"],
eps=1e-5,
)
W_key = "__W_concat__"
if W_key not in weights or weights[W_key].shape != (5 * H, D):
weights[W_key] = torch.cat([
weights['left_proj.weight'],
weights['right_proj.weight'],
weights['left_gate.weight'],
weights['right_gate.weight'],
weights['out_gate.weight'],
], dim=0).contiguous().half()
W = weights[W_key]
x_T = x.view(M, D).t().half()
P = torch.matmul(W, x_T).view(5, H, M)
LEFT_T = torch.sigmoid(P[2]) * P[0]
if mask.min() < 1.0:
LEFT_T *= mask.view(1, M).half()
RIGHT_T = torch.sigmoid(P[3]) * P[1]
OG_T = torch.sigmoid(P[4])
LEFT = LEFT_T.view(H, B, N, N).permute(1, 2, 3, 0)
RIGHT = RIGHT_T.view(H, B, N, N).permute(1, 2, 3, 0)
OG = OG_T.view(H, B, N, N).permute(1, 2, 3, 0)
if use_lr:
RANK = min(64, H // 4)
LEFT_lr = LEFT[..., :RANK].contiguous()
RIGHT_lr = RIGHT[..., :RANK].contiguous()
EIN_lr = _einsum_opt(LEFT_lr, RIGHT_lr)
proj_key = "__proj_lr__"
if proj_key not in weights or weights[proj_key].shape != (H, RANK):
weights[proj_key] = torch.eye(H, device=device)[:, :RANK].contiguous()
EIN = torch.matmul(EIN_lr, weights[proj_key].t())
if H > RANK:
LEFT_res = LEFT[..., RANK:min(RANK*2, H)]
RIGHT_res = RIGHT[..., RANK:min(RANK*2, H)]
EIN_res = _einsum_opt(LEFT_res, RIGHT_res)
EIN[..., RANK:min(RANK*2, H)] += EIN_res
else:
EIN = _einsum_opt(LEFT, RIGHT)
G = F.layer_norm(
EIN, (H,),
weight=weights['to_out_norm.weight'],
bias=weights['to_out_norm.bias'],
eps=1e-5
) * OG.float()
Wt_out_key = "__Wt_out__"
if Wt_out_key not in weights or weights[Wt_out_key].shape != (H, D):
weights[Wt_out_key] = weights['to_out.weight'].t().half()
OUT = torch.matmul(G.view(M, H).half(), weights[Wt_out_key]).float()
return OUT.view(B, N, N, D)
def custom_kernel(data: input_t) -> output_t:
with DisableCuDNNTF32():
# Keep matmul precision limited but safe
torch.set_float32_matmul_precision('medium')
return _custom_kernel_core(data)
scrolls · 108 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 40862.
Best evidence level for this revision: reported
JSON