submission 748568
coderwhisper · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 638 lines, June 9 Researcher Reciprocity License v1.0.
submission_v104_fixptr2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-748568?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:521e61c3f675a087d0ca3e68a12de2d1b96d84dbdf1c1e177bbea46e650f1f7e
license declaredunknown
license concludedunknown
authorscoderwhisper
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)stages = 1
num_warps=q_NW, waves_per_eu=0, num_stages=1,Kernel source
submission_v104_fixptr2.py638 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
V104: Fixed direct launcher.launch bypass (attempt 2).
V103's fix failed because captured args contain tensor OBJECTS, not int pointers.
isinstance(val, int) matched nothing -> ptr_subs=[0] for all shapes.
Fix: match both torch.Tensor (by data_ptr()) and int args. For tensor args,
substitute with a tensor sharing the new data's storage. For int args, substitute
with new data_ptr() values.
"""
import torch
import triton
import triton.language as tl
import sys
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_preshuffle_kernel,
)
from task import input_t, output_t
SCALE_GROUP = 32
_cache = {}
_bf16 = torch.bfloat16
def p(*args):
print(*args, file=sys.stderr)
# ===================== A16WFP4 preshuffle kernel with inline KSPLIT reduction =====================
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
}
)
@triton.jit
def _a16wfp4_inline_reduce_kernel(
a_ptr, b_ptr, c_ptr, c_final_ptr, b_scales_ptr, counter_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
stride_cf_m, stride_cf_n,
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,
INLINE_REDUCE: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
tl.assume(pid_k >= 0)
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in tl.range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter, num_stages=num_stages):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
4, 16, 2, 2, 1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.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)
)
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
if INLINE_REDUCE:
c_ptrs = (
c_ptr
+ pid_k * stride_ck
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, accumulator, mask=c_mask, cache_modifier=".wt")
tile_id = pid_m * num_pid_n + pid_n
old_count = tl.atomic_add(counter_ptr + tile_id, 1)
if old_count == NUM_KSPLIT - 1:
total = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ks in range(NUM_KSPLIT):
partial_ptrs = (
c_ptr
+ ks * stride_ck
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
)
partial = tl.load(partial_ptrs, mask=c_mask)
total += partial
result = total.to(c_final_ptr.type.element_ty)
cf_ptrs = (
c_final_ptr
+ stride_cf_m * offs_cm[:, None]
+ stride_cf_n * offs_cn[None, :]
)
tl.store(cf_ptrs, result, mask=c_mask, cache_modifier=".wt")
tl.atomic_xchg(counter_ptr + tile_id, 0)
else:
c = accumulator.to(c_final_ptr.type.element_ty)
cf_ptrs = (
c_final_ptr
+ stride_cf_m * offs_cm[:, None]
+ stride_cf_n * offs_cn[None, :]
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(cf_ptrs, c, mask=c_mask, cache_modifier=".wt")
# ===================== Fused quant+shuffle kernel =====================
@triton.jit
def _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m_in, stride_x_n_in,
stride_x_fp4_m_in, stride_x_fp4_n_in,
M, N, padded_sn,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SCALING_MODE: tl.constexpr, EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
i = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
j = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
shuffled_off = (
(i // 32)[:, None] * (32 * padded_sn)
+ (j // 8)[None, :] * 256 + (j % 4)[None, :] * 64
+ (i % 16)[:, None] * 4 + ((j % 8) // 4)[None, :] * 2
+ ((i % 32) // 16)[:, None]
)
if EVEN_M_N:
tl.store(bs_ptr + shuffled_off, bs_e8m0)
else:
scale_n = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_mask = (i < M)[:, None] & (j < scale_n)[None, :]
tl.store(bs_ptr + shuffled_off, bs_e8m0, mask=bs_mask)
# ===================== Split-K helper =====================
def _get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
if (K % (SPLITK_BLOCK_SIZE // 2) == 0
and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
and K % (BLOCK_SIZE_K // 2) == 0):
break
elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
else:
break
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
# ===================== Per-shape configs =====================
_PS_CONFIGS = {
"k_small_m4": {
"BLOCK_SIZE_M": 4, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
"k_small_m32": {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
"k7168": {
"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": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 7,
},
}
_AFP4_CONFIGS = {
"k2048_m64": {
"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
"k1536_m256": {
"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": None, "NUM_KSPLIT": 1,
},
}
# ===================== A16WFP4 Preshuffle path =====================
def _prepare_preshuffle(M, K, N, dev):
if K <= 1024:
cfg = dict(_PS_CONFIGS["k_small_m4" if M <= 4 else "k_small_m32"])
else:
cfg = dict(_PS_CONFIGS["k7168"])
K_packed = K // 2
if cfg["NUM_KSPLIT"] > 1:
sbs, bsk, nks = _get_splitk(K_packed, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])
cfg["SPLITK_BLOCK_SIZE"] = sbs
cfg["BLOCK_SIZE_K"] = bsk
cfg["NUM_KSPLIT"] = nks
if cfg["BLOCK_SIZE_K"] >= 2 * K_packed:
cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
cfg["NUM_KSPLIT"] = 1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
use_splitk = cfg["NUM_KSPLIT"] > 1
y = torch.empty((M, N), dtype=_bf16, device=dev)
if use_splitk:
y_pp = torch.empty((cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=dev)
num_pid_m = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
num_pid_n = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
counter = torch.zeros(num_pid_m * num_pid_n, dtype=torch.int32, device=dev)
else:
cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
y_pp = None
counter = None
grid = (cfg["NUM_KSPLIT"] * triton.cdiv(M, cfg["BLOCK_SIZE_M"]) * triton.cdiv(N, cfg["BLOCK_SIZE_N"]),)
return ("ps", cfg, grid, y, y_pp, use_splitk, K_packed, counter)
def _run_preshuffle(data, M, K, N, cached):
_, cfg, grid, y, y_pp, use_splitk, K_packed, counter = cached
w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
bs_uint8 = data[4].view(torch.uint8)
sm, sn = bs_uint8.shape
w_scales = bs_uint8.reshape(sm // 32, sn * 32)
if use_splitk:
_a16wfp4_inline_reduce_kernel[grid](
data[0], w_ps, y_pp, y, w_scales, counter,
M, N, K_packed,
data[0].stride(0), data[0].stride(1),
w_ps.stride(0), w_ps.stride(1),
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
w_scales.stride(0), w_scales.stride(1),
y.stride(0), y.stride(1),
INLINE_REDUCE=True,
**cfg,
)
else:
_a16wfp4_inline_reduce_kernel[grid](
data[0], w_ps, None, y, w_scales, None,
M, N, K_packed,
data[0].stride(0), data[0].stride(1),
w_ps.stride(0), w_ps.stride(1),
0, y.stride(0), y.stride(1),
w_scales.stride(0), w_scales.stride(1),
y.stride(0), y.stride(1),
INLINE_REDUCE=False,
**cfg,
)
return y
# ===================== AFP4WFP4 preshuffle path (medium K) =====================
def _prepare_afp4(M, K, N, dev):
K_packed = K // 2
x_fp4 = torch.empty((M, K_packed), dtype=torch.uint8, device=dev)
sn = triton.cdiv(K, SCALE_GROUP)
padded_sm = triton.cdiv(M, 256) * 256
padded_sn = triton.cdiv(sn, 8) * 8
bs_shuffled = torch.zeros((padded_sm, padded_sn), dtype=torch.uint8, device=dev)
if M <= 32:
q_NI, q_BSM, q_BSN, q_NW, q_NS = 1, triton.next_power_of_2(M), 32, 1, 1
else:
q_NI, q_BSM, q_BSN, q_NW, q_NS = 4, 32, 128, 4, 2
q_EVEN = (M % q_BSM == 0) and (K % q_BSN == 0)
grid_q = (triton.cdiv(M, q_BSM), triton.cdiv(K, q_BSN * q_NI))
fp4_s = x_fp4.stride()
if K == 2048 and M <= 64:
cfg = dict(_AFP4_CONFIGS["k2048_m64"])
else:
cfg = dict(_AFP4_CONFIGS["k1536_m256"])
if cfg["NUM_KSPLIT"] > 1:
sbs, bsk, nks = _get_splitk(K_packed, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])
cfg["SPLITK_BLOCK_SIZE"] = sbs
cfg["BLOCK_SIZE_K"] = bsk
cfg["NUM_KSPLIT"] = nks
else:
cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
if cfg["BLOCK_SIZE_K"] >= 2 * K_packed:
cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
grid_gemm = (
cfg["NUM_KSPLIT"]
* triton.cdiv(M, cfg["BLOCK_SIZE_M"])
* triton.cdiv(N, cfg["BLOCK_SIZE_N"]),
)
y = torch.empty((M, N), dtype=_bf16, device=dev)
return ("afp4", x_fp4, bs_shuffled, padded_sm, padded_sn,
cfg, y, grid_q, q_BSM, q_BSN, q_NI, q_NS, q_NW, q_EVEN,
fp4_s[0], fp4_s[1], grid_gemm, K_packed)
def _run_afp4(data, M, K, N, cached):
(_, x_fp4, bs_sh, padded_sm, padded_sn,
cfg, y, grid_q, q_BSM, q_BSN, q_NI, q_NS, q_NW, q_EVEN,
fp4s0, fp4s1, grid_gemm, K_packed) = cached
A = data[0]
_fused_quant_shuffle_kernel[grid_q](
A, x_fp4, bs_sh,
A.stride(0), A.stride(1), fp4s0, fp4s1,
M, K, padded_sn,
BLOCK_SIZE_M=q_BSM, BLOCK_SIZE_N=q_BSN,
NUM_ITER=q_NI, NUM_STAGES=q_NS,
MXFP4_QUANT_BLOCK_SIZE=SCALE_GROUP,
SCALING_MODE=0, EVEN_M_N=q_EVEN,
num_warps=q_NW, waves_per_eu=0, num_stages=1,
)
a_scales = bs_sh.reshape(padded_sm // 32, padded_sn * 32)
w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
bs_uint8 = data[4].view(torch.uint8)
bsm, bsn = bs_uint8.shape
w_scales = bs_uint8.reshape(bsm // 32, bsn * 32)
_gemm_afp4wfp4_preshuffle_kernel[grid_gemm](
x_fp4, w_ps, y, a_scales, w_scales,
M, N, K_packed,
x_fp4.stride(0), x_fp4.stride(1),
w_ps.stride(0), w_ps.stride(1),
0, y.stride(0), y.stride(1),
a_scales.stride(0), a_scales.stride(1),
w_scales.stride(0), w_scales.stride(1),
**cfg,
)
return y
# ===================== Direct launch bypass infrastructure =====================
# key -> list of (launch_fn, args_list, ptr_subs)
# ptr_subs: list of (position, data_index, is_tensor)
_fast_cache = {}
_outputs = {} # key -> output tensor
def _get_jit_fn(kfn):
while not hasattr(kfn, 'device_caches') and hasattr(kfn, 'fn'):
kfn = kfn.fn
return kfn
def _patch_and_capture(kernel_fns, run_fn):
captured = []
patches = []
for kfn in kernel_fns:
jit_fn = _get_jit_fn(kfn)
if not hasattr(jit_fn, 'device_caches') or 0 not in jit_fn.device_caches:
continue
cache_dict = jit_fn.device_caches[0][0]
for k, ck in cache_dict.items():
try:
orig = ck.run.launch
def make_hook(o):
def hook(*a):
captured.append((o, a))
return o(*a)
return hook
new_fn = make_hook(orig)
ck.run.launch = new_fn
if ck.run.launch is new_fn:
patches.append((ck.run, orig))
else:
ck.run.launch = orig
except (AttributeError, TypeError):
pass
result = run_fn()
for launcher_obj, orig in patches:
try:
launcher_obj.launch = orig
except (AttributeError, TypeError):
pass
return captured, result
# ===================== Derive input tensors for fast path =====================
def _derive_tensors_ps(data, K_packed, N):
"""Derive kernel arg tensors for preshuffle path."""
w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
bs_uint8 = data[4].view(torch.uint8)
sm, sn = bs_uint8.shape
w_scales = bs_uint8.reshape(sm // 32, sn * 32)
return data[0], w_ps, w_scales
def _derive_tensors_afp4(data, K_packed, N):
"""Derive kernel arg tensors for AFP4WFP4 path."""
w_ps = data[3].view(torch.uint8).reshape(N // 16, K_packed * 16)
bs_uint8 = data[4].view(torch.uint8)
bsm, bsn = bs_uint8.shape
w_scales = bs_uint8.reshape(bsm // 32, bsn * 32)
return data[0], w_ps, w_scales
# ===================== Router =====================
def custom_kernel(data: input_t) -> output_t:
A = data[0]
M, K = A.shape
N = data[3].shape[0]
key = (M, K, N)
# Fast path: replay with updated input tensor pointers
if key in _fast_cache:
fc = _fast_cache[key]
if fc is not None:
entries, derive_fn, K_packed = fc
a_new, w_ps_new, w_scales_new = derive_fn(data, K_packed, N)
new_tensors = {0: a_new, 3: w_ps_new, 4: w_scales_new}
for launch_fn, args_list, ptr_subs in entries:
for pos, data_idx, is_tensor in ptr_subs:
if is_tensor:
args_list[pos] = new_tensors[data_idx]
else:
args_list[pos] = new_tensors[data_idx].data_ptr()
launch_fn(*args_list)
return _outputs[key]
else:
cached = _cache[key]
if cached[0] == "ps":
return _run_preshuffle(data, M, K, N, cached)
else:
return _run_afp4(data, M, K, N, cached)
# First call: prepare, compile, run, and capture
if key not in _cache:
if K <= 1024 or K >= 4096:
_cache[key] = _prepare_preshuffle(M, K, N, A.device)
else:
_cache[key] = _prepare_afp4(M, K, N, A.device)
cached = _cache[key]
# Run once to compile
if cached[0] == "ps":
result = _run_preshuffle(data, M, K, N, cached)
kfns = [_a16wfp4_inline_reduce_kernel]
run_fn = lambda: _run_preshuffle(data, M, K, N, cached)
derive_fn = _derive_tensors_ps
K_packed = K // 2
else:
result = _run_afp4(data, M, K, N, cached)
kfns = [_fused_quant_shuffle_kernel, _gemm_afp4wfp4_preshuffle_kernel]
run_fn = lambda: _run_afp4(data, M, K, N, cached)
derive_fn = _derive_tensors_afp4
K_packed = K // 2
# Record input tensor pointers for substitution mapping
# data[0]=A, data[3]=B_shuffle (views share same data_ptr), data[4]=B_scale_sh
input_ptrs = {}
for idx in (0, 3, 4):
ptr = data[idx].data_ptr()
if ptr not in input_ptrs:
input_ptrs[ptr] = idx
# Capture launcher.launch args by running again with patches
captures, _ = _patch_and_capture(kfns, run_fn)
if captures:
fast_entries = []
for fn, args in captures:
args_list = list(args)
ptr_subs = []
for pos, val in enumerate(args_list):
# Check int pointers
if isinstance(val, int) and val in input_ptrs:
ptr_subs.append((pos, input_ptrs[val], False))
# Check tensor objects (Triton may pass tensors, not ints)
elif isinstance(val, torch.Tensor):
try:
dptr = val.data_ptr()
if dptr in input_ptrs:
ptr_subs.append((pos, input_ptrs[dptr], True))
except Exception:
pass
fast_entries.append((fn, args_list, ptr_subs))
_fast_cache[key] = (fast_entries, derive_fn, K_packed)
_outputs[key] = result
p(f"[V104] {key}: captured {len(captures)} launches, ptr_subs={[len(e[2]) for e in fast_entries]}")
else:
_fast_cache[key] = None
p(f"[V104] {key}: capture FAILED, using normal path")
return result
scrolls · 638 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