Skip to content
KernelIndex
Search⌘K

submission 754367

flower2123 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 414 lines, June 9 Researcher Reciprocity License v1.0.

submission_v1_flower_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754367?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.10µs
#24 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2e655cfed189ae7fcb9826895c4d3ccc40e1f198360c6ff601aeba04a60a0990
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4fp4 = torch.empty((m, kp), dtype=torch.uint8, device=a.device)
split-kc.update(SPLITK_BLOCK_SIZE=sb, BLOCK_SIZE_K=bk, NUM_KSPLIT=ns)

Kernel source

submission_v1_flower_mm.py414 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

_out_pool = {}
_params = {}
_grid_memo = {}


PER_SHAPE = {
    (4, 2880, 512):    dict(tm=4,  tn=128, tk=256, gm=1, warp=4, pipe=2, occ=1, kdim=16, cmod=None,  ks=1),
    (16, 2112, 7168):  dict(tm=16, tn=128, tk=512, gm=1, warp=4, pipe=2, occ=3, kdim=16, cmod=".cg", ks=14),
    (32, 4096, 512):   dict(tm=16, tn=32,  tk=256, gm=1, warp=4, pipe=3, occ=3, kdim=16, cmod=".cg", ks=1),
    (32, 2880, 512):   dict(tm=8,  tn=128, tk=256, gm=1, warp=4, pipe=2, occ=2, kdim=16, cmod=None,  ks=1),
    (64, 7168, 2048):  dict(tm=16, tn=128, tk=256, gm=1, warp=4, pipe=2, occ=2, kdim=16, cmod=".cg", ks=1),
    (256, 3072, 1536): dict(tm=16, tn=256, tk=512, gm=1, warp=8, pipe=2, occ=2, kdim=16, cmod=None,  ks=1),
}
FALLBACK = dict(tm=16, tn=32, tk=256, gm=1, warp=2, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)

SEP_SHAPE = {
    (32, 4096, 512):  dict(tm=32, tn=128, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1),
    (32, 2880, 512):  dict(tm=32, tn=64,  tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1),
    (64, 7168, 2048): dict(tm=16, tn=128, tk=512, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=2),
}
SEP_FALLBACK = dict(tm=16, tn=64, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)


def _adj_split(kp, bk, ns):
    sb = triton.cdiv(2 * triton.cdiv(kp, ns), bk) * bk
    while ns > 1 and bk > 16:
        ok = kp % (sb // 2) == 0 and sb % bk == 0 and kp % (bk // 2) == 0
        if ok:
            break
        if kp % (sb // 2) != 0 and ns > 1:
            ns //= 2
        elif sb % bk != 0:
            ns = ns // 2 if ns > 1 else ns
            if ns <= 1 and bk > 16:
                bk //= 2
        elif kp % (bk // 2) != 0 and bk > 16:
            bk //= 2
        else:
            break
        sb = triton.cdiv(2 * triton.cdiv(kp, ns), bk) * bk
    return sb, bk, triton.cdiv(kp, sb // 2)


def _expand(raw, kp):
    c = {
        "BLOCK_SIZE_M": raw["tm"], "BLOCK_SIZE_N": max(raw["tn"], 32),
        "BLOCK_SIZE_K": raw["tk"], "GROUP_SIZE_M": raw["gm"],
        "num_warps": raw["warp"], "num_stages": raw["pipe"],
        "waves_per_eu": raw["occ"], "matrix_instr_nonkdim": raw["kdim"],
        "cache_modifier": raw["cmod"], "NUM_KSPLIT": raw["ks"],
    }
    if c["NUM_KSPLIT"] > 1:
        sb, bk, ns = _adj_split(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
        c.update(SPLITK_BLOCK_SIZE=sb, BLOCK_SIZE_K=bk, NUM_KSPLIT=ns)
    else:
        c.update(SPLITK_BLOCK_SIZE=2 * kp, NUM_KSPLIT=1)
    if c["BLOCK_SIZE_K"] >= 2 * kp:
        c.update(BLOCK_SIZE_K=triton.next_power_of_2(2 * kp), SPLITK_BLOCK_SIZE=2 * kp, NUM_KSPLIT=1)
    return c


@triton.jit
def _fp4_encode(
    src, R: tl.constexpr, C: tl.constexpr,
):
    Q: tl.constexpr = 32
    NB: tl.constexpr = C // Q
    f32 = src.to(tl.float32).reshape(R, NB, Q)
    mx = tl.max(tl.abs(f32), axis=-1, keep_dims=True)
    mx = (mx.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    lx = ((mx >> 23) & 0xFF).to(tl.int32) - 127
    ub = tl.minimum(tl.maximum(lx - 2, -127), 127)
    e8 = ub.to(tl.uint8) + 127
    hb = (ub.to(tl.int32) + 127).to(tl.uint32) << 23
    hf = hb.to(tl.float32, bitcast=True)
    hf_full = tl.broadcast_to(hf, (R, NB, Q)).reshape(R, C)
    pv = hf_full.reshape(R, C // 2, 2)
    ev, _ = tl.split(pv)
    ev = ev.reshape(R, C // 2)
    raw16 = src.to(tl.uint16, bitcast=True).reshape(R, C // 2, 2)
    lo16, hi16 = tl.split(raw16)
    p32 = lo16.to(tl.uint32) | (hi16.to(tl.uint32) << 16)
    p32 = p32.reshape(R, C // 2)
    enc = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2", "=v, v, v",
        [p32, ev], dtype=tl.uint32, is_pure=True, pack=1,
    )
    return (enc & 0xFF).to(tl.uint8).reshape(R, C // 2), e8.reshape(R, NB)


@triton.jit
def _encode_block(
    inp, o_fp4, o_sc,
    nrow, ncol,
    si0, si1, sq0, sq1, ss0, ss1,
    TR: tl.constexpr, TC: tl.constexpr,
):
    r = tl.program_id(0) * TR + tl.arange(0, TR)
    c = tl.program_id(1) * TC + tl.arange(0, TC)
    dat = tl.load(inp + r[:, None] * si0 + c[None, :] * si1,
                  mask=(r[:, None] < nrow) & (c[None, :] < ncol), other=0.0)
    f4, sc = _fp4_encode(dat, TR, TC)
    HC: tl.constexpr = TC // 2
    qc = tl.program_id(1) * HC + tl.arange(0, HC)
    tl.store(o_fp4 + r[:, None] * sq0 + qc[None, :] * sq1,
             f4, mask=(r[:, None] < nrow) & (qc[None, :] < ncol // 2))
    SC: tl.constexpr = TC // 32
    sc_c = tl.program_id(1) * SC + tl.arange(0, SC)
    tl.store(o_sc + r[:, None] * ss0 + sc_c[None, :] * ss1,
             sc, mask=(r[:, None] < nrow) & (sc_c[None, :] < ncol // 32))


@triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
                     and (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)
                     and (a["K"] % (a["SPLITK_BLOCK_SIZE"] // 2) == 0)})
@triton.jit
def _matmul_fused(
    a_p, b_p, c_p, bs_p,
    M, N, K,
    sa0, sa1, sb0, sb1, sc_k, sc0, sc1, sbs0, sbs1,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
    OK: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    tl.assume(sa0 > 0); tl.assume(sa1 > 0)
    tl.assume(sb0 > 0); tl.assume(sb1 > 0)
    tl.assume(sc0 > 0); tl.assume(sc1 > 0)
    tl.assume(sbs0 > 0); tl.assume(sbs1 > 0)

    G: tl.constexpr = 32
    nm = tl.cdiv(M, BLOCK_SIZE_M)
    nn = tl.cdiv(N, BLOCK_SIZE_N)
    uid = tl.program_id(0)
    sk = uid % NUM_KSPLIT
    flat = uid // NUM_KSPLIT
    if NUM_KSPLIT == 1:
        gn = GROUP_SIZE_M * nn
        gi = flat // gn
        fm = gi * GROUP_SIZE_M
        gs = min(nm - fm, GROUP_SIZE_M)
        im = fm + (flat % gn) % gs
        jn = (flat % gn) // gs
    else:
        im = flat // nn
        jn = flat % nn
    tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)

    if (sk * SPLITK_BLOCK_SIZE // 2) < K:
        niter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        ck = sk * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
        pa = a_p + rm[:, None] * sa0 + ck[None, :] * sa1

        sh_a = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        sh_o = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sh_a
        bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
        pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1

        bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
        bsk = sk * (SPLITK_BLOCK_SIZE // G) * 32 + tl.arange(0, BLOCK_SIZE_K // G * 32)
        pbs = bs_p + bsn[:, None] * sbs0 + bsk[None, :] * sbs1

        dot = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for ki in range(sk * niter, (sk + 1) * niter):
            if OK:
                va = tl.load(pa)
                vbs = tl.load(pbs, cache_modifier=cache_modifier)
                vb = tl.load(pb, cache_modifier=cache_modifier)
            else:
                lo = (ki - sk * niter) * BLOCK_SIZE_K
                va = tl.load(pa, mask=tl.arange(0, BLOCK_SIZE_K)[None, :] < (2 * K - sk * SPLITK_BLOCK_SIZE - lo), other=0.0)
                vbs = tl.load(pbs, cache_modifier=cache_modifier)
                vb = tl.load(pb,
                    mask=sh_a[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * niter) * (BLOCK_SIZE_K // 2))) * 16),
                    other=0, cache_modifier=cache_modifier)

            aq, asc = _fp4_encode(va, BLOCK_SIZE_M, BLOCK_SIZE_K)
            ws = (vbs
                .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // G // 8, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // G))
            bd = (vb.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0))
            dot = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", dot)

            pa += BLOCK_SIZE_K * sa1
            pb += (BLOCK_SIZE_K // 2) * 16 * sb1
            pbs += BLOCK_SIZE_K * sbs1

        out = dot.to(c_p.type.element_ty)
        om = im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        on = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        tl.store(c_p + sc0 * om[:, None] + sc1 * on[None, :] + sk * sc_k,
                 out, mask=(om[:, None] < M) & (on[None, :] < N))


@triton.jit
def _merge(src, dst, M, N, sk0, sm0, sn0, dm0, dn0,
           TM: tl.constexpr, TN: tl.constexpr,
           REAL: tl.constexpr, CAP: tl.constexpr):
    rm = (tl.program_id(0) * TM + tl.arange(0, TM)) % M
    rn = (tl.program_id(1) * TN + tl.arange(0, TN)) % N
    base = src + rm[:, None] * sm0 + rn[None, :] * sn0
    s = tl.load(base).to(tl.float32)
    for j in tl.static_range(1, CAP):
        if j < REAL:
            s += tl.load(base + j * sk0).to(tl.float32)
    tl.store(dst + rm[:, None] * dm0 + rn[None, :] * dn0, s.to(dst.type.element_ty))


@triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
                     and (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)
                     and (a["K"] % (a["SPLITK_BLOCK_SIZE"] // 2) == 0)})
@triton.jit
def _matmul_sep(
    q4, qsc, b_p, c_p, bs_p,
    M, N, K,
    sq0, sq1, ss0, ss1,
    sb0, sb1, sc_k, sc0, sc1, sbs0, sbs1,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
    OK: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    tl.assume(sq0 > 0); tl.assume(sq1 > 0)
    tl.assume(ss0 > 0); tl.assume(ss1 > 0)
    tl.assume(sb0 > 0); tl.assume(sb1 > 0)
    tl.assume(sc0 > 0); tl.assume(sc1 > 0)
    tl.assume(sbs0 > 0); tl.assume(sbs1 > 0)

    G: tl.constexpr = 32
    HK: tl.constexpr = BLOCK_SIZE_K // 2
    SK: tl.constexpr = BLOCK_SIZE_K // G
    nm = tl.cdiv(M, BLOCK_SIZE_M)
    nn = tl.cdiv(N, BLOCK_SIZE_N)
    uid = tl.program_id(0)
    ks = uid % NUM_KSPLIT
    flat = uid // NUM_KSPLIT
    if NUM_KSPLIT == 1:
        gn = GROUP_SIZE_M * nn
        gi = flat // gn
        fm = gi * GROUP_SIZE_M
        gs = min(nm - fm, GROUP_SIZE_M)
        im = fm + (flat % gn) % gs
        jn = (flat % gn) // gs
    else:
        im = flat // nn
        jn = flat % nn
    tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(ks >= 0)

    if (ks * SPLITK_BLOCK_SIZE // 2) < K:
        niter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, HK)
        rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        cq = ks * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)
        pq = q4 + rm[:, None] * sq0 + cq[None, :] * sq1
        cs = ks * (SPLITK_BLOCK_SIZE // G) + tl.arange(0, SK)
        ps = qsc + rm[:, None] * ss0 + cs[None, :] * ss1

        sh_a = tl.arange(0, HK * 16)
        sh_o = ks * (SPLITK_BLOCK_SIZE // 2) * 16 + sh_a
        bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
        pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1

        bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
        bsk = ks * (SPLITK_BLOCK_SIZE // G) * 32 + tl.arange(0, SK * 32)
        pbs = bs_p + bsn[:, None] * sbs0 + bsk[None, :] * sbs1

        dot = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for ki in range(ks * niter, (ks + 1) * niter):
            if OK:
                va = tl.load(pq, cache_modifier=cache_modifier)
                vas = tl.load(ps, cache_modifier=cache_modifier)
            else:
                lo = (ki - ks * niter) * HK
                rem = K - (ks * (SPLITK_BLOCK_SIZE // 2) + lo)
                va = tl.load(pq, mask=tl.arange(0, HK)[None, :] < rem, other=0, cache_modifier=cache_modifier)
                sr = (2 * K) // G - (ks * (SPLITK_BLOCK_SIZE // G) + (ki - ks * niter) * SK)
                vas = tl.load(ps, mask=tl.arange(0, SK)[None, :] < sr, other=0, cache_modifier=cache_modifier)

            ws = (tl.load(pbs, cache_modifier=cache_modifier)
                .reshape(BLOCK_SIZE_N // 32, SK // 8, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, SK))

            if OK:
                vb = tl.load(pb, cache_modifier=cache_modifier)
            else:
                vb = tl.load(pb,
                    mask=sh_a[None, :] < ((K - (ks * (SPLITK_BLOCK_SIZE // 2) + (ki - ks * niter) * HK)) * 16),
                    other=0, cache_modifier=cache_modifier)

            vb = (vb.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, HK)
                .trans(1, 0))
            dot = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", dot)

            pq += HK * sq1
            ps += SK * ss1
            pb += HK * 16 * sb1
            pbs += BLOCK_SIZE_K * sbs1

        out = dot.to(c_p.type.element_ty)
        om = im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        on = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        tl.store(c_p + sc0 * om[:, None] + sc1 * on[None, :] + ks * sc_k,
                 out, mask=(om[:, None] < M) & (on[None, :] < N))


def _weight_views(ws, wsc, n, kp):
    return ws.view(torch.uint8).reshape(n // 16, kp * 16), wsc.view(torch.uint8)


def _bufs(m, n, ns, dev):
    t = (m, n, ns)
    if t not in _out_pool:
        y = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
        pp = torch.empty((ns, m, n), dtype=torch.float32, device=dev) if ns > 1 else None
        _out_pool[t] = (y, pp)
    return _out_pool[t]


def _setup(m, n, k, dev):
    t = (m, n, k)
    if t not in _grid_memo:
        kp = k // 2
        c = _expand(PER_SHAPE.get(t, FALLBACK).copy(), kp)
        ns = c["NUM_KSPLIT"]
        y, pp = _bufs(m, n, ns, dev)
        g = (ns * triton.cdiv(m, c["BLOCK_SIZE_M"]) * triton.cdiv(n, c["BLOCK_SIZE_N"]),)
        ck, cm, cn = (0, y.stride(0), y.stride(1)) if ns == 1 else (pp.stride(0), pp.stride(1), pp.stride(2))
        info = dict(c=c, kp=kp, g=g, ns=ns, ck=ck, cm=cm, cn=cn)
        if ns > 1:
            info["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
            info["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
            info["pk"] = triton.next_power_of_2(ns)
        _grid_memo[t] = info
    return _grid_memo[t]


def _exec_fused(a, ws, wsc, m, n, k):
    p = _setup(m, n, k, a.device)
    y, pp = _bufs(m, n, p["ns"], a.device)
    bw, bsc = _weight_views(ws, wsc, n, p["kp"])
    _matmul_fused[p["g"]](
        a, bw, y if p["ns"] == 1 else pp, bsc,
        m, n, p["kp"],
        a.stride(0), a.stride(1), bw.stride(0), bw.stride(1),
        p["ck"], p["cm"], p["cn"], bsc.stride(0), bsc.stride(1),
        **p["c"],
    )
    if p["ns"] > 1:
        _merge[p["rg"]](pp, y, m, n,
            pp.stride(0), pp.stride(1), pp.stride(2), y.stride(0), y.stride(1),
            16, 64, p["rk"], p["pk"])
    return y


def _exec_sep(a, ws, wsc, m, n, k):
    kp = k // 2
    QR, QC = 16, 256
    fp4 = torch.empty((m, kp), dtype=torch.uint8, device=a.device)
    sc = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
    _encode_block[(triton.cdiv(m, QR), triton.cdiv(k, QC))](
        a, fp4, sc, m, k,
        a.stride(0), a.stride(1), fp4.stride(0), fp4.stride(1), sc.stride(0), sc.stride(1),
        QR, QC)

    c = _expand(SEP_SHAPE.get((m, n, k), SEP_FALLBACK).copy(), kp)
    ns = c["NUM_KSPLIT"]
    y = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
    pp = torch.empty((ns, m, n), dtype=torch.float32, device=a.device) if ns > 1 else None
    bw, bsc = _weight_views(ws, wsc, n, kp)

    grid_fn = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
    _matmul_sep[grid_fn](
        fp4, sc, bw, y if ns == 1 else pp, bsc,
        m, n, kp,
        fp4.stride(0), fp4.stride(1), sc.stride(0), sc.stride(1),
        bw.stride(0), bw.stride(1),
        0 if ns == 1 else pp.stride(0),
        y.stride(0) if ns == 1 else pp.stride(1),
        y.stride(1) if ns == 1 else pp.stride(2),
        bsc.stride(0), bsc.stride(1), **c)

    if ns > 1:
        real = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
        _merge[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
            pp, y, m, n,
            pp.stride(0), pp.stride(1), pp.stride(2), y.stride(0), y.stride(1),
            16, 64, real, triton.next_power_of_2(ns))
    return y


def custom_kernel(data: input_t) -> output_t:
    inp = data[0]
    return _exec_fused(inp, data[3], data[4], inp.shape[0], data[1].shape[0], inp.shape[1])
scrolls · 414 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON