Skip to content
KernelIndex
Search⌘K

submission 530427

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mxfp4_v87_best_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-530427?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
11.7µs
#349 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2c1d9519b5fe6e7795383818d22e49109080137c5b044a7811bb24e2179a2466
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Techniques

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

split-k- M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)
tile-m = 8- M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)

Kernel source

mxfp4_v87_best_hybrid.py161 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Version 87: Best hybrid combining ALL discoveries:
- M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)
- M=16,K=7168: Fused GEMM with KSPLIT=14 → 16.8µs (v79 finding)
- M=32,K=512: Fused GEMM with DEFAULT config → 9.47/9.50µs (v77 finding)
- M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)
- M=256,K=1536: ASM splitK=1 → 12.5µs (v62 finding)

Expected geomean: ~11.1µs (vs 12.7µs baseline)
"""
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout
from task import input_t, output_t


@triton.jit
def _mxfp4_quant_op_asm_exact(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    E8_BIAS: tl.constexpr = 127
    E2_BIAS: tl.constexpr = 1
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.log2(amax).floor() - 2
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
    quant_scale = tl.exp2(-scale_e8m0_unbiased)
    qx = x * quant_scale
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    e = (qx >> 23) & 0xFF
    m = qx & 0x7FFFFF
    adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
    m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
    e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
    e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
    e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
    e2m1_value = tl.reshape(
        e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact

from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0

_kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

# ASM splitK for large-M shapes
_ASM_SPLITK = {
    (64, 7168, 2048): 2,     # Proven best: 4 K-splits
    (256, 3072, 1536): 1,    # Proven best: 2 K-splits
}

# Fused configs - only override where needed
_FUSED_CONFIGS = {
    # M=4: BLOCK_SIZE_M=8 reduces padding waste (4→8 vs 4→32)
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
        "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    # M=16: Split-K=14 for better CU utilization
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 14,
    },
    # M=32: Use default config (BLOCK_SIZE_N=64 gives more N-tiles → better CU util)
    # No config override → uses default
}

_QUANT_BLOCK = 32
_QUANT_TILE = 128

_bufs = {}


def _get_asm_bufs(m, k, n, device):
    x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
    sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
    sN_pad = ((sN + 7) >> 3) << 3
    sM_pad = ((m + 255) >> 8) << 8
    scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
    padded_m = ((m + 31) >> 5) << 5
    out = torch.empty((padded_m, n), dtype=_bf16, device=device)
    return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]
    key = (m, n, k)

    if key in _ASM_SPLITK:
        # ASM 2-kernel path for M>=64 shapes
        if key not in _bufs:
            _bufs[key] = ('asm', _get_asm_bufs(m, k, n, A.device))
        _, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]

        grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
        _dynamic_mxfp4_quant_kernel_asm_layout[grid](
            A, x_fp4, scale,
            A.stride(0), A.stride(1),
            x_fp4.stride(0), x_fp4.stride(1),
            scale.stride(0), scale.stride(1),
            M=m, N=k, scaleN=sN,
            scaleM_pad=sM_pad, scaleN_pad=sN_pad,
            BLOCK_SIZE=_QUANT_TILE,
            MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
            SCALING_MODE=0, SHUFFLE=True,
        )

        splitK = _ASM_SPLITK[key]
        gemm_a4w4_asm(
            x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
            out, _kernel_32x128,
            bpreshuffle=True, log2_k_split=splitK,
        )
        return out[:m]
    else:
        # Fused preshuffle GEMM for M<=32 shapes
        if key not in _bufs:
            _bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
        _, out = _bufs[key]

        w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
        sm, sn = B_scale_sh.shape
        w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)

        config = _FUSED_CONFIGS.get(key)
        return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 161 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 529480.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- Autogenerated MXFP4 candidate: cached_quant_h8_2112_7168_h8_32x128_s3
- Strategy: cached_quant
- Generated by autosubmit_amd_mxfp4_mm.py.
- """
- from __future__ import annotations
+ Version 87: Best hybrid combining ALL discoveries:
+ - M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)
+ - M=16,K=7168: Fused GEMM with KSPLIT=14 → 16.8µs (v79 finding)
+ - M=32,K=512: Fused GEMM with DEFAULT config → 9.47/9.50µs (v77 finding)
+ - M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)
+ - M=256,K=1536: ASM splitK=1 → 12.5µs (v62 finding)
- import csv
- import os
- from pathlib import Path
-
+ Expected geomean: ~11.1µs (vs 12.7µs baseline)
+ """
import torch
+ import triton
+ import triton.language as tl
+ import aiter
+ from aiter import dtypes
+ from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
+ from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout
+ from task import input_t, output_t
- _AITER_BASE = Path("/home/runner/aiter")
- _CUSTOM_CONFIG_PATH = "/tmp/submission_20260311_auto_cached_quant_h8_2112_7168_h8_32x128_s3_6a1a37d4.csv"
- _CUSTOM_ROWS = (
- {'cu_num': '256', 'M': '4', 'N': '2880', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},
- {'cu_num': '256', 'M': '16', 'N': '2112', 'K': '7168', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '3'},
- {'cu_num': '256', 'M': '32', 'N': '2880', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},
- {'cu_num': '256', 'M': '32', 'N': '4096', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},
- {'cu_num': '256', 'M': '8', 'N': '2112', 'K': '7168', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '3'},
- )
- _CUSTOM_KEYS = {(row["M"], row["N"], row["K"]) for row in _CUSTOM_ROWS}
-
- def _looks_like_config(path: Path) -> bool:
- try:
- with path.open(newline="") as f:
- reader = csv.DictReader(f)
- return reader.fieldnames is not None and "kernelName" in reader.fieldnames and "splitK" in reader.fieldnames
- except Exception:
- return False
-
-
- def _find_default_config() -> Path | None:
- candidates = (
- _AITER_BASE / "aiter" / "configs" / "a4w4_blockscale_tuned_gemm.csv",
- _AITER_BASE / "hsa" / "configs" / "a4w4_tuned_gemm.csv",
- _AITER_BASE / "configs" / "a4w4_tuned_gemm.csv",
+ @triton.jit
+ def _mxfp4_quant_op_asm_exact(
+ x,
+ BLOCK_SIZE_N,
+ BLOCK_SIZE_M,
+ MXFP4_QUANT_BLOCK_SIZE,
+ ):
+ E8_BIAS: tl.constexpr = 127
+ E2_BIAS: tl.constexpr = 1
+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
+ x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
+ amax = amax.to(tl.int32, bitcast=True)
+ amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
+ amax = amax.to(tl.float32, bitcast=True)
+ scale_e8m0_unbiased = tl.log2(amax).floor() - 2
+ scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
+ bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
+ quant_scale = tl.exp2(-scale_e8m0_unbiased)
+ qx = x * quant_scale
+ qx = qx.to(tl.uint32, bitcast=True)
+ s = qx & 0x80000000
+ e = (qx >> 23) & 0xFF
+ m = qx & 0x7FFFFF
+ adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
+ m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
+ e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
+ e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
+ e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
+ e2m1_value = tl.reshape(
+ e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
- for path in candidates:
- if path.exists() and _looks_like_config(path):
- return path
- for path in sorted(_AITER_BASE.rglob("*a4w4*.csv")):
- if _looks_like_config(path):
- return path
- return None
+ evens, odds = tl.split(e2m1_value)
+ x_fp4 = evens | (odds << 4)
+ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
- def _write_merged_config():
- base_path = _find_default_config()
- fieldnames = ["cu_num", "M", "N", "K", "kernelName", "splitK"]
- rows = []
+ import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact
- if base_path is not None:
- with base_path.open(newline="") as f:
- reader = csv.DictReader(f)
- if reader.fieldnames:
- fieldnames = list(reader.fieldnames)
- for row in reader:
- if (row.get("M", ""), row.get("N", ""), row.get("K", "")) in _CUSTOM_KEYS:
- continue
- rows.append(row)
+ from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
- rows.extend(_CUSTOM_ROWS)
- with open(_CUSTOM_CONFIG_PATH, "w", newline="") as f:
- writer = csv.DictWriter(f, fieldnames=fieldnames)
- writer.writeheader()
- writer.writerows(rows)
+ _bf16 = dtypes.bf16
+ _fp4x2 = dtypes.fp4x2
+ _fp8_e8m0 = dtypes.fp8_e8m0
+ _kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- _write_merged_config()
- os.environ["AITER_CONFIG_GEMM_A4W4"] = _CUSTOM_CONFIG_PATH
+ # ASM splitK for large-M shapes
+ _ASM_SPLITK = {
+ (64, 7168, 2048): 2, # Proven best: 4 K-splits
+ (256, 3072, 1536): 1, # Proven best: 2 K-splits
+ }
- import aiter
- from aiter import QuantType, dtypes
- from aiter.utility import fp4_utils
+ # Fused configs - only override where needed
+ _FUSED_CONFIGS = {
+ # M=4: BLOCK_SIZE_M=8 reduces padding waste (4→8 vs 4→32)
+ (4, 2880, 512): {
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ },
+ # M=16: Split-K=14 for better CU utilization
+ (16, 2112, 7168): {
+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 14,
+ },
+ # M=32: Use default config (BLOCK_SIZE_N=64 gives more N-tiles → better CU util)
+ # No config override → uses default
+ }
- from task import input_t, output_t
+ _QUANT_BLOCK = 32
+ _QUANT_TILE = 128
- _DTYPES = dtypes
- _FP4_UTILS = fp4_utils
- _TRITON_QUANT = aiter.get_triton_quant(QuantType.per_1x32)
- _GEMM = aiter.gemm_a4w4
- _A_QUANT_BUFFERS = {}
- _OUT_CACHE = {}
+ _bufs = {}
- _QUANT_BLOCK_SIZE_BY_MK = {
- (4, 512): 32,
- (32, 512): 32,
- }
+ def _get_asm_bufs(m, k, n, device):
+ x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
+ sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
+ sN_pad = ((sN + 7) >> 3) << 3
+ sM_pad = ((m + 255) >> 8) << 8
+ scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
+ padded_m = ((m + 31) >> 5) << 5
+ out = torch.empty((padded_m, n), dtype=_bf16, device=device)
+ return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m
- def _ceil_div(x: int, y: int) -> int:
- return (x + y - 1) // y
+ @torch.inference_mode()
+ def custom_kernel(data: input_t) -> output_t:
+ A, B, B_q, B_shuffle, B_scale_sh = data
+ m, k = A.shape
+ n = B.shape[0]
+ key = (m, n, k)
- def _get_quant_buffers(m: int, k: int, device):
- scale_n = _ceil_div(k, 32)
- scale_n_pad = _ceil_div(scale_n, 8) * 8
- scale_m_pad = _ceil_div(m, 256) * 256
- key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)
- buffers = _A_QUANT_BUFFERS.get(key)
- if buffers is None:
- a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
- a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
- buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)
- _A_QUANT_BUFFERS[key] = buffers
- return buffers
+ if key in _ASM_SPLITK:
+ # ASM 2-kernel path for M>=64 shapes
+ if key not in _bufs:
+ _bufs[key] = ('asm', _get_asm_bufs(m, k, n, A.device))
+ _, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]
+ grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
+ _dynamic_mxfp4_quant_kernel_asm_layout[grid](
+ A, x_fp4, scale,
+ A.stride(0), A.stride(1),
+ x_fp4.stride(0), x_fp4.stride(1),
+ scale.stride(0), scale.stride(1),
+ M=m, N=k, scaleN=sN,
+ scaleM_pad=sM_pad, scaleN_pad=sN_pad,
+ BLOCK_SIZE=_QUANT_TILE,
+ MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
+ SCALING_MODE=0, SHUFFLE=True,
+ )
- def _quantize_a(a):
- m = int(a.shape[0])
- k = int(a.shape[1])
- try:
- a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)
- block_size = _QUANT_BLOCK_SIZE_BY_MK.get((m, k), 128)
- grid = (_ceil_div(m, block_size), scale_n_pad)
- _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](
- a,
- a_q,
- a_scale,
- *a.stride(),
- *a_q.stride(),
- *a_scale.stride(),
- M=m,
- N=k,
- scaleN=scale_n,
- scaleM_pad=scale_m_pad,
- scaleN_pad=scale_n_pad,
- BLOCK_SIZE=block_size,
- MXFP4_QUANT_BLOCK_SIZE=32,
- SCALING_MODE=0,
- SHUFFLE=True,
+ splitK = _ASM_SPLITK[key]
+ gemm_a4w4_asm(
+ x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
+ out, _kernel_32x128,
+ bpreshuffle=True, log2_k_split=splitK,
)
- return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)
- except Exception:
- return _TRITON_QUANT(a, shuffle=True)
+ return out[:m]
+ else:
+ # Fused preshuffle GEMM for M<=32 shapes
+ if key not in _bufs:
+ _bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
+ _, out = _bufs[key]
- def custom_kernel(data: input_t) -> output_t:
- a, _b, _b_q, b_shuffle, b_scale_sh = data
- a_q, a_scale_sh = _quantize_a(a)
- return _GEMM(
- a_q,
- b_shuffle,
- a_scale_sh,
- b_scale_sh,
- dtype=dtypes.bf16,
- bpreshuffle=True,
- )
+ w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
+ sm, sn = B_scale_sh.shape
+ w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
+
+ config = _FUSED_CONFIGS.get(key)
+ return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 286 diff lines total

Best evidence level for this revision: reported

JSON