Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI300X
5.40ms
#5 of 19
2025-09-19

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