Skip to content
KernelIndex
Search⌘K

submission 754549

flower2123 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v3_flower_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754549?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
7.99µs
#15 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:996bca8c64910916ff537966c9145d2ef91e7405727e1da629b5b1625b8bb5e4
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15

Techniques

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

split-kand (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)

Kernel source

submission_v3_flower_mm.py323 lines
# Author: flower

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

_y_store = {}
_c_store = {}
_l_store = {}

TUNE = {
    (4, 2880, 512):    {"BLOCK_SIZE_M": 4,  "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
    (16, 2112, 7168):  {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 14},
    (32, 4096, 512):   {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32,  "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
    (32, 2880, 512):   {"BLOCK_SIZE_M": 8,  "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None,  "NUM_KSPLIT": 1},
    (64, 7168, 2048):  {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
    (256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None,  "NUM_KSPLIT": 1},
}
TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}

SEP_TUNE = {
    (32, 4096, 512):  {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
    (32, 2880, 512):  {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64,  "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
    (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},
}
SEP_TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}


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


@triton.jit
def _quantize_mxfp4(
    v, DM: tl.constexpr, DK: tl.constexpr,
):
    QG: tl.constexpr = 32
    NQ: tl.constexpr = DK // QG
    w = v.to(tl.float32).reshape(DM, NQ, QG)
    pk = tl.max(tl.abs(w), axis=-1, keep_dims=True)
    pk = (pk.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    lg = ((pk >> 23) & 0xFF).to(tl.int32) - 127
    ub = tl.minimum(tl.maximum(lg - 2, -127), 127)
    se = ub.to(tl.uint8) + 127
    hb = (ub.to(tl.int32) + 127).to(tl.uint32) << 23
    hf = hb.to(tl.float32, bitcast=True)
    hx = tl.broadcast_to(hf, (DM, NQ, QG)).reshape(DM, DK)
    pv = hx.reshape(DM, DK // 2, 2)
    ev, _ = tl.split(pv)
    ev = ev.reshape(DM, DK // 2)
    u16 = v.to(tl.uint16, bitcast=True).reshape(DM, DK // 2, 2)
    lo, hi = tl.split(u16)
    u32 = lo.to(tl.uint32) | (hi.to(tl.uint32) << 16)
    u32 = u32.reshape(DM, DK // 2)
    r = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2", "=v, v, v",
        [u32, ev], dtype=tl.uint32, is_pure=True, pack=1)
    return (r & 0xFF).to(tl.uint8).reshape(DM, DK // 2), se.reshape(DM, NQ)


@triton.jit
def _quant_block_kernel(
    src_p, fp4_p, sc_p, nrow, ncol,
    s0, s1, q0, q1, c0, c1,
    BM: tl.constexpr, BK: tl.constexpr,
):
    ri = tl.program_id(0) * BM + tl.arange(0, BM)
    ci = tl.program_id(1) * BK + tl.arange(0, BK)
    d = tl.load(src_p + ri[:, None] * s0 + ci[None, :] * s1,
                mask=(ri[:, None] < nrow) & (ci[None, :] < ncol), other=0.0)
    f4, sc = _quantize_mxfp4(d, BM, BK)
    HK: tl.constexpr = BK // 2
    hc = tl.program_id(1) * HK + tl.arange(0, HK)
    tl.store(fp4_p + ri[:, None] * q0 + hc[None, :] * q1,
             f4, mask=(ri[:, None] < nrow) & (hc[None, :] < ncol // 2))
    SK: tl.constexpr = BK // 32
    si = tl.program_id(1) * SK + tl.arange(0, SK)
    tl.store(sc_p + ri[:, None] * c0 + si[None, :] * c1,
             sc, mask=(ri[:, None] < nrow) & (si[None, :] < ncol // 32))


@triton.heuristics({"EVEN_K": 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 _gemm_inline_quant(
    ap, bp, cp, bsp, M, N, K,
    sa0, sa1, sb0, sb1, sck, 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,
    EVEN_K: 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)
    SG: 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; pid = uid // NUM_KSPLIT
    if NUM_KSPLIT == 1:
        gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
        gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gs
    else:
        im = pid // nn; jn = pid % nn
    tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
    if (sk * SPLITK_BLOCK_SIZE // 2) < K:
        ni = 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 = ap + rm[:, None] * sa0 + ck[None, :] * sa1
        sha = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
        bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
        pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1
        bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
        bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, BLOCK_SIZE_K // SG * 32)
        pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for ki in range(sk * ni, (sk + 1) * ni):
            if EVEN_K:
                va = tl.load(pa); vbs = tl.load(pbs, cache_modifier=cache_modifier)
                vb = tl.load(pb, cache_modifier=cache_modifier)
            else:
                lo = (ki - sk * ni) * 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=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * (BLOCK_SIZE_K // 2))) * 16), other=0, cache_modifier=cache_modifier)
            aq, asc = _quantize_mxfp4(va, BLOCK_SIZE_M, BLOCK_SIZE_K)
            ws = vbs.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SG // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SG)
            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)
            acc = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", acc)
            pa += BLOCK_SIZE_K * sa1; pb += (BLOCK_SIZE_K // 2) * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
        res = acc.to(cp.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(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))


@triton.jit
def _accum(inp, out, M, N, isk, ism, isn, osm, osn,
           TM: tl.constexpr, TN: tl.constexpr, RK: tl.constexpr, PK: tl.constexpr):
    rm = (tl.program_id(0) * TM + tl.arange(0, TM)) % M
    rn = (tl.program_id(1) * TN + tl.arange(0, TN)) % N
    b = inp + rm[:, None] * ism + rn[None, :] * isn
    s = tl.load(b).to(tl.float32)
    for j in tl.static_range(1, PK):
        if j < RK:
            s += tl.load(b + j * isk).to(tl.float32)
    tl.store(out + rm[:, None] * osm + rn[None, :] * osn, s.to(out.type.element_ty))


@triton.heuristics({"EVEN_K": 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 _gemm_preq(
    q4p, scp, bp, cp, bsp, M, N, K,
    sq0, sq1, ss0, ss1, sb0, sb1, sck, 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,
    EVEN_K: 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)
    SG: tl.constexpr = 32; HK: tl.constexpr = BLOCK_SIZE_K // 2; SCK: tl.constexpr = BLOCK_SIZE_K // SG
    nm = tl.cdiv(M, BLOCK_SIZE_M); nn = tl.cdiv(N, BLOCK_SIZE_N)
    uid = tl.program_id(0); sk = uid % NUM_KSPLIT; pid = uid // NUM_KSPLIT
    if NUM_KSPLIT == 1:
        gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
        gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gs
    else:
        im = pid // nn; jn = pid % nn
    tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
    if (sk * SPLITK_BLOCK_SIZE // 2) < K:
        ni = tl.cdiv(SPLITK_BLOCK_SIZE // 2, HK)
        rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        cq = sk * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)
        pq = q4p + rm[:, None] * sq0 + cq[None, :] * sq1
        cs = sk * (SPLITK_BLOCK_SIZE // SG) + tl.arange(0, SCK)
        ps = scp + rm[:, None] * ss0 + cs[None, :] * ss1
        sha = tl.arange(0, HK * 16); sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
        bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
        pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1
        bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
        bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, SCK * 32)
        pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for ki in range(sk * ni, (sk + 1) * ni):
            if EVEN_K:
                va = tl.load(pq, cache_modifier=cache_modifier); vas = tl.load(ps, cache_modifier=cache_modifier)
            else:
                lo = (ki - sk * ni) * HK; rem = K - (sk * (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) // SG - (sk * (SPLITK_BLOCK_SIZE // SG) + (ki - sk * ni) * SCK)
                vas = tl.load(ps, mask=tl.arange(0, SCK)[None, :] < sr, other=0, cache_modifier=cache_modifier)
            ws = tl.load(pbs, cache_modifier=cache_modifier).reshape(BLOCK_SIZE_N // 32, SCK // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, SCK)
            if EVEN_K:
                vb = tl.load(pb, cache_modifier=cache_modifier)
            else:
                vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * 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)
            acc = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", acc)
            pq += HK * sq1; ps += SCK * ss1; pb += HK * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
        res = acc.to(cp.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(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))


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


def _ybuf(m, n, ns, d):
    k = (m, n, ns)
    if k not in _y_store:
        _y_store[k] = (torch.empty((m, n), dtype=torch.bfloat16, device=d),
                       torch.empty((ns, m, n), dtype=torch.float32, device=d) if ns > 1 else None)
    return _y_store[k]


def _resolve(m, n, k):
    t = (m, n, k)
    if t not in _c_store:
        c = TUNE.get(t, TUNE_DEF).copy()
        kp = k // 2
        if c["NUM_KSPLIT"] > 1:
            sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
            c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
        else:
            c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
        if c["BLOCK_SIZE_K"] >= 2 * kp:
            c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp)
            c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
        c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
        _c_store[t] = c
    return _c_store[t]


def _pre(m, n, k, d):
    t = (m, n, k)
    if t not in _l_store:
        c = _resolve(m, n, k); kp = k // 2; ns = c["NUM_KSPLIT"]
        y, pp = _ybuf(m, n, ns, d)
        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))
        r = dict(c=c, kp=kp, g=g, ns=ns, ck=ck, cm=cm, cn=cn)
        if ns > 1:
            r["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
            r["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
            r["pk"] = triton.next_power_of_2(ns)
        _l_store[t] = r
    return _l_store[t]


def _go_fused(a, ws, wsc, m, n, k):
    p = _pre(m, n, k, a.device)
    y, pp = _ybuf(m, n, p["ns"], a.device)
    bw, bs = _w_view(ws, wsc, n, p["kp"])
    _gemm_inline_quant[p["g"]](
        a, bw, y if p["ns"] == 1 else pp, bs, m, n, p["kp"],
        a.stride(0), a.stride(1), bw.stride(0), bw.stride(1),
        p["ck"], p["cm"], p["cn"], bs.stride(0), bs.stride(1), **p["c"])
    if p["ns"] > 1:
        _accum[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 _go_sep(a, ws, wsc, m, n, k):
    kp, QM, QK = k // 2, 16, 256
    f4 = torch.empty((m, kp), dtype=torch.uint8, device=a.device)
    sc = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
    _quant_block_kernel[(triton.cdiv(m, QM), triton.cdiv(k, QK))](
        a, f4, sc, m, k, a.stride(0), a.stride(1), f4.stride(0), f4.stride(1), sc.stride(0), sc.stride(1), QM, QK)
    c = SEP_TUNE.get((m, n, k), SEP_TUNE_DEF).copy()
    if c["NUM_KSPLIT"] > 1:
        sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
        c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
    else:
        c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
    if c["BLOCK_SIZE_K"] >= 2 * kp:
        c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp); c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
    c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
    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, bs = _w_view(ws, wsc, n, kp)
    gf = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
    _gemm_preq[gf](f4, sc, bw, y if ns == 1 else pp, bs, m, n, kp,
        f4.stride(0), f4.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), bs.stride(0), bs.stride(1), **c)
    if ns > 1:
        rk = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
        _accum[(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, rk, triton.next_power_of_2(ns))
    return y


def custom_kernel(data: input_t) -> output_t:
    x = data[0]
    return _go_fused(x, data[3], data[4], x.shape[0], data[1].shape[0], x.shape[1])
scrolls · 323 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 754367.

+ # Author: flower
+
import torch
import triton
import triton.language as tl
from task import input_t, output_t
- _out_pool = {}
- _params = {}
- _grid_memo = {}
+ _y_store = {}
+ _c_store = {}
+ _l_store = {}
-
- 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),
+ TUNE = {
+ (4, 2880, 512): {"BLOCK_SIZE_M": 4, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
+ (16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 14},
+ (32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
+ (32, 2880, 512): {"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},
+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
+ (256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},
}
- FALLBACK = dict(tm=16, tn=32, tk=256, gm=1, warp=2, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)
+ TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 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_TUNE = {
+ (32, 4096, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
+ (32, 2880, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},
}
- SEP_FALLBACK = dict(tm=16, tn=64, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)
+ SEP_TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 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:
+ def _fix_ksplit(half, bk, nk):
+ sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk
+ while nk > 1 and bk > 16:
+ if half % (sb // 2) == 0 and sb % bk == 0 and half % (bk // 2) == 0:
break
- if kp % (sb // 2) != 0 and ns > 1:
- ns //= 2
+ if half % (sb // 2) != 0 and nk > 1:
+ nk //= 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:
+ nk = nk // 2 if nk > 1 else nk
+ if nk <= 1 and bk > 16: bk //= 2
+ elif half % (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)
+ sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk
+ return sb, bk, triton.cdiv(half, 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,
+ def _quantize_mxfp4(
+ v, DM: tl.constexpr, DK: 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
+ QG: tl.constexpr = 32
+ NQ: tl.constexpr = DK // QG
+ w = v.to(tl.float32).reshape(DM, NQ, QG)
+ pk = tl.max(tl.abs(w), axis=-1, keep_dims=True)
+ pk = (pk.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
+ lg = ((pk >> 23) & 0xFF).to(tl.int32) - 127
+ ub = tl.minimum(tl.maximum(lg - 2, -127), 127)
+ se = 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)
+ hx = tl.broadcast_to(hf, (DM, NQ, QG)).reshape(DM, DK)
+ pv = hx.reshape(DM, DK // 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(
+ ev = ev.reshape(DM, DK // 2)
+ u16 = v.to(tl.uint16, bitcast=True).reshape(DM, DK // 2, 2)
+ lo, hi = tl.split(u16)
+ u32 = lo.to(tl.uint32) | (hi.to(tl.uint32) << 16)
+ u32 = u32.reshape(DM, DK // 2)
+ r = 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)
+ [u32, ev], dtype=tl.uint32, is_pure=True, pack=1)
+ return (r & 0xFF).to(tl.uint8).reshape(DM, DK // 2), se.reshape(DM, NQ)
@triton.jit
- def _encode_block(
- inp, o_fp4, o_sc,
- nrow, ncol,
- si0, si1, sq0, sq1, ss0, ss1,
- TR: tl.constexpr, TC: tl.constexpr,
+ def _quant_block_kernel(
+ src_p, fp4_p, sc_p, nrow, ncol,
+ s0, s1, q0, q1, c0, c1,
+ BM: tl.constexpr, BK: 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))
+ ri = tl.program_id(0) * BM + tl.arange(0, BM)
+ ci = tl.program_id(1) * BK + tl.arange(0, BK)
+ d = tl.load(src_p + ri[:, None] * s0 + ci[None, :] * s1,
+ mask=(ri[:, None] < nrow) & (ci[None, :] < ncol), other=0.0)
+ f4, sc = _quantize_mxfp4(d, BM, BK)
+ HK: tl.constexpr = BK // 2
+ hc = tl.program_id(1) * HK + tl.arange(0, HK)
+ tl.store(fp4_p + ri[:, None] * q0 + hc[None, :] * q1,
+ f4, mask=(ri[:, None] < nrow) & (hc[None, :] < ncol // 2))
+ SK: tl.constexpr = BK // 32
+ si = tl.program_id(1) * SK + tl.arange(0, SK)
+ tl.store(sc_p + ri[:, None] * c0 + si[None, :] * c1,
+ sc, mask=(ri[:, None] < nrow) & (si[None, :] < ncol // 32))
- @triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
+ @triton.heuristics({"EVEN_K": 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,
+ def _gemm_inline_quant(
+ ap, bp, cp, bsp, M, N, K,
+ sa0, sa1, sb0, sb1, sck, 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,
+ GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
+ EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
- waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
- cache_modifier: 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
+ 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)
+ SG: 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; pid = 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
+ gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
+ gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gs
else:
- im = flat // nn
- jn = flat % nn
+ im = pid // nn; jn = pid % 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)
+ ni = 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
+ pa = ap + rm[:, None] * sa0 + ck[None, :] * sa1
+ sha = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
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
-
+ pb = bp + bn[:, None] * sb0 + sho[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)
+ bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, BLOCK_SIZE_K // SG * 32)
+ pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for ki in range(sk * ni, (sk + 1) * ni):
+ if EVEN_K:
+ 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
+ lo = (ki - sk * ni) * 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)
+ vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * (BLOCK_SIZE_K // 2))) * 16), other=0, cache_modifier=cache_modifier)
+ aq, asc = _quantize_mxfp4(va, BLOCK_SIZE_M, BLOCK_SIZE_K)
+ ws = vbs.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SG // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SG)
+ 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)
+ acc = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", acc)
+ pa += BLOCK_SIZE_K * sa1; pb += (BLOCK_SIZE_K // 2) * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
+ res = acc.to(cp.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))
+ tl.store(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, 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):
+ def _accum(inp, out, M, N, isk, ism, isn, osm, osn,
+ TM: tl.constexpr, TN: tl.constexpr, RK: tl.constexpr, PK: 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))
+ b = inp + rm[:, None] * ism + rn[None, :] * isn
+ s = tl.load(b).to(tl.float32)
+ for j in tl.static_range(1, PK):
+ if j < RK:
+ s += tl.load(b + j * isk).to(tl.float32)
+ tl.store(out + rm[:, None] * osm + rn[None, :] * osn, s.to(out.type.element_ty))
- @triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
+ @triton.heuristics({"EVEN_K": 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,
+ def _gemm_preq(
+ q4p, scp, bp, cp, bsp, M, N, K,
+ sq0, sq1, ss0, ss1, sb0, sb1, sck, 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,
+ GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
+ EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
- waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
- cache_modifier: 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(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
+ SG: tl.constexpr = 32; HK: tl.constexpr = BLOCK_SIZE_K // 2; SCK: tl.constexpr = BLOCK_SIZE_K // SG
+ nm = tl.cdiv(M, BLOCK_SIZE_M); nn = tl.cdiv(N, BLOCK_SIZE_N)
+ uid = tl.program_id(0); sk = uid % NUM_KSPLIT; pid = 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
+ gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
+ gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % 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)
+ im = pid // nn; jn = pid % nn
+ tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
+ if (sk * SPLITK_BLOCK_SIZE // 2) < K:
+ ni = 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
+ cq = sk * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)
+ pq = q4p + rm[:, None] * sq0 + cq[None, :] * sq1
+ cs = sk * (SPLITK_BLOCK_SIZE // SG) + tl.arange(0, SCK)
+ ps = scp + rm[:, None] * ss0 + cs[None, :] * ss1
+ sha = tl.arange(0, HK * 16); sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
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
-
+ pb = bp + bn[:, None] * sb0 + sho[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)
+ bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, SCK * 32)
+ pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for ki in range(sk * ni, (sk + 1) * ni):
+ if EVEN_K:
+ 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)
+ lo = (ki - sk * ni) * HK; rem = K - (sk * (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:
+ sr = (2 * K) // SG - (sk * (SPLITK_BLOCK_SIZE // SG) + (ki - sk * ni) * SCK)
+ vas = tl.load(ps, mask=tl.arange(0, SCK)[None, :] < sr, other=0, cache_modifier=cache_modifier)
+ ws = tl.load(pbs, cache_modifier=cache_modifier).reshape(BLOCK_SIZE_N // 32, SCK // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, SCK)
+ if EVEN_K:
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)
+ vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * 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)
+ acc = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", acc)
+ pq += HK * sq1; ps += SCK * ss1; pb += HK * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
+ res = acc.to(cp.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))
+ tl.store(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))
- def _weight_views(ws, wsc, n, kp):
+ def _w_view(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 _ybuf(m, n, ns, d):
+ k = (m, n, ns)
+ if k not in _y_store:
+ _y_store[k] = (torch.empty((m, n), dtype=torch.bfloat16, device=d),
+ torch.empty((ns, m, n), dtype=torch.float32, device=d) if ns > 1 else None)
+ return _y_store[k]
- def _setup(m, n, k, dev):
+ def _resolve(m, n, k):
t = (m, n, k)
- if t not in _grid_memo:
+ if t not in _c_store:
+ c = TUNE.get(t, TUNE_DEF).copy()
kp = k // 2
- c = _expand(PER_SHAPE.get(t, FALLBACK).copy(), kp)
- ns = c["NUM_KSPLIT"]
- y, pp = _bufs(m, n, ns, dev)
+ if c["NUM_KSPLIT"] > 1:
+ sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
+ c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
+ else:
+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
+ if c["BLOCK_SIZE_K"] >= 2 * kp:
+ c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp)
+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
+ c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
+ _c_store[t] = c
+ return _c_store[t]
+
+
+ def _pre(m, n, k, d):
+ t = (m, n, k)
+ if t not in _l_store:
+ c = _resolve(m, n, k); kp = k // 2; ns = c["NUM_KSPLIT"]
+ y, pp = _ybuf(m, n, ns, d)
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)
+ r = 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]
+ r["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
+ r["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
+ r["pk"] = triton.next_power_of_2(ns)
+ _l_store[t] = r
+ return _l_store[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"],
+ def _go_fused(a, ws, wsc, m, n, k):
+ p = _pre(m, n, k, a.device)
+ y, pp = _ybuf(m, n, p["ns"], a.device)
+ bw, bs = _w_view(ws, wsc, n, p["kp"])
+ _gemm_inline_quant[p["g"]](
+ a, bw, y if p["ns"] == 1 else pp, bs, 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"],
- )
+ p["ck"], p["cm"], p["cn"], bs.stride(0), bs.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"])
+ _accum[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)
+ def _go_sep(a, ws, wsc, m, n, k):
+ kp, QM, QK = k // 2, 16, 256
+ f4 = 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)
+ _quant_block_kernel[(triton.cdiv(m, QM), triton.cdiv(k, QK))](
+ a, f4, sc, m, k, a.stride(0), a.stride(1), f4.stride(0), f4.stride(1), sc.stride(0), sc.stride(1), QM, QK)
+ c = SEP_TUNE.get((m, n, k), SEP_TUNE_DEF).copy()
+ if c["NUM_KSPLIT"] > 1:
+ sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
+ c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
+ else:
+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
+ if c["BLOCK_SIZE_K"] >= 2 * kp:
+ c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp); c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
+ c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
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)
-
+ bw, bs = _w_view(ws, wsc, n, kp)
+ gf = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
+ _gemm_preq[gf](f4, sc, bw, y if ns == 1 else pp, bs, m, n, kp,
+ f4.stride(0), f4.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), bs.stride(0), bs.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))
+ rk = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
+ _accum[(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, rk, 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])
+ x = data[0]
+ return _go_fused(x, data[3], data[4], x.shape[0], data[1].shape[0], x.shape[1])
scrolls · 632 diff lines total

Best evidence level for this revision: reported

JSON