Skip to content
KernelIndex
Search⌘K

submission 661294

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v023a.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-661294?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
31.2µs
#18 of 766
2026-03-29

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

num-warps = 2num_warps=2,num_stages=1,waves_per_eu=4
split-kfor split_kv_id in range(0, num_valid_kv_splits):
stages = 1num_warps=2,num_stages=1,waves_per_eu=4

Kernel source

v023a.py183 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""V023a: v022o plus lighter custom 32x8k Triton stage2 launch."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
os.environ.setdefault("HIPBLASLT_ALLOW_FLUSH_DENORM", "1")
os.environ.setdefault("GPU_MAX_HW_QUEUES", "2")
import torch, triton, triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes, mla as aiter_mla
try:
    from aiter.jit.module_quant import static_per_tensor_quant as _quant
except Exception:
    from aiter.ops.quant import static_per_tensor_quant as _quant
try:
    from aiter.jit.module_mla_metadata import get_mla_metadata_v1 as _mv1
except Exception:
    _mv1 = aiter.get_mla_metadata_v1
try:
    from aiter.jit.module_mla_asm import mla_decode_stage1_asm_fwd as _s1
except Exception:
    _s1 = aiter.mla_decode_stage1_asm_fwd
try:
    from aiter.jit.module_mla_reduce import mla_reduce_v1 as _rd
except Exception:
    _rd = aiter.mla_reduce_v1
_mi = aiter.get_mla_metadata_info_v1
@triton.jit
def _safe_reduce_branchless(
    Mid_O, Mid_lse, O, qo_indptr, kv_indptr, num_kv_splits_indptr,
    stride_mid_ob: tl.int64, stride_mid_oh: tl.int64, stride_mid_os: tl.int64,
    stride_obs: tl.int64, stride_oh: tl.int64,
    MAYBE_FINAL_OUT: tl.constexpr, BATCH_NUM: tl.constexpr,
    BLOCK_DV: tl.constexpr, Lv: tl.constexpr, mgc: tl.constexpr,
):
    cur_batch = tl.program_id(0); cur_head = tl.program_id(1)
    cur_qo_start = tl.load(qo_indptr + cur_batch); cur_qo_end = tl.load(qo_indptr + cur_batch + 1)
    cur_split_start = tl.load(num_kv_splits_indptr + cur_batch); cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)
    num_max_kv_splits = tl.load(num_kv_splits_indptr + BATCH_NUM)
    cur_kv_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)
    offs_d = tl.arange(0, BLOCK_DV); mask_d = offs_d < Lv
    offs_logic = cur_qo_start * stride_mid_ob + cur_head * stride_mid_oh
    offs_v = offs_logic * Lv + offs_d
    num_valid_kv_splits = tl.minimum(cur_split_end - cur_split_start, tl.cdiv(cur_kv_seq_len, mgc))
    final_out = MAYBE_FINAL_OUT and num_max_kv_splits == BATCH_NUM
    for cur_qo in range(cur_qo_start, cur_qo_end):
        if final_out:
            input_ptr = Mid_O.to(tl.pointer_type(O.type.element_ty))
            out = tl.load(input_ptr + Lv * (cur_qo * stride_mid_os + cur_head * stride_mid_oh) + offs_d, mask=mask_d, other=0.0)
            tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, out, mask=mask_d)
        else:
            e_sum = 0.0; e_max = -float("inf"); acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)
            for split_kv_id in range(0, num_valid_kv_splits):
                tv = tl.load(Mid_O + offs_v + split_kv_id * stride_mid_os * Lv, mask=mask_d, other=0.0)
                tlogic = tl.load(Mid_lse + offs_logic + split_kv_id * stride_mid_os)
                tlogic = tl.where(tlogic == tlogic, tlogic, -1e30); tlogic = tl.minimum(tlogic, 1e30)
                n_e_max = tl.maximum(tlogic, e_max); old_scale = tl.exp(e_max - n_e_max)
                acc *= old_scale; exp_logic = tl.exp(tlogic - n_e_max)
                acc += exp_logic * tv; e_sum = e_sum * old_scale + exp_logic; e_max = n_e_max
            offs_logic += stride_mid_ob; offs_v += stride_mid_ob * Lv
            tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, acc / e_sum, mask=mask_d)
aiter.mla._fwd_kernel_stage2_asm = _safe_reduce_branchless
_fwd = aiter.mla.mla_decode_fwd
NH=16;QKD=576;VD=512;SM=float(1.0/(QKD**0.5));FP8=aiter_dtypes.fp8;BF16=torch.bfloat16
QSV=float(2.0/torch.finfo(FP8).max);MGC=64;_st={};_pc={};_nc={};_s2c={}
def _gs(d):
    t=_st.get(0)
    if t is None: t=torch.tensor([QSV],dtype=torch.float32,device=d);_st[0]=t
    return t
def _gnc(dev,bs,kvl,ps):
    key=(bs,kvl,ps);c=_nc.get(key)
    if c: return c
    ppb=kvl//ps;np_=bs*kvl//ps
    ki=torch.arange(np_,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb;out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev);c=(ki,kl,kip,out,qf);_nc[key]=c;return c
def _gpc(dev,qo,bs,kvl,ps,intra=True,split_override=None,kg_override=None):
    key=(bs,kvl,ps,intra,split_override,kg_override);c=_pc.get(key)
    if c: return c
    tot=bs*kvl;ppb=kvl//ps;np_=tot//ps
    ns = split_override if split_override is not None else aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)[0]
    ki=torch.arange(np_,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb;out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    kg=kg_override if kg_override is not None else max(ps,16)
    info=_mi(bs,1,NH,FP8,FP8,is_sparse=False,fast_mode=True,num_kv_splits=ns,intra_batch_mode=intra)
    w=[torch.empty(s,dtype=t,device=dev) for s,t in info]
    _mv1(qo,kip,kl,16,1,False,w[0],w[2],w[1],w[3],w[4],w[5],page_size=ps,kv_granularity=kg,max_seqlen_qo=1,uni_seqlen_qo=1,fast_mode=True,max_split_per_batch=ns,intra_batch_mode=intra,dtype_q=FP8,dtype_kv=FP8)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    pt=int(w[5].numel());po=torch.empty((pt,1,NH,VD),dtype=torch.float32,device=dev);pl=torch.empty((pt,1,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,po,pl);_pc[key]=c;return c
def _uniform_nsi(dev,bs,ns):
    return torch.arange(bs+1,dtype=torch.int32,device=dev) * ns
def _g32s2(dev,bs,kvl,ps):
    key=(bs,kvl,ps,"32split3");c=_s2c.get(key)
    if c: return c
    ppb=kvl//ps;np_=bs*kvl//ps;ns=3
    ki=torch.arange(np_,dtype=torch.int32,device=dev)
    kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
    out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    nsi=_uniform_nsi(dev,bs,ns)
    logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
    lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
def _g64s2(dev,bs,kvl,ps):
    key=(bs,kvl,ps,"64split2w2");c=_s2c.get(key)
    if c: return c
    ppb=kvl//ps;np_=bs*kvl//ps;ns=2
    ki=torch.arange(np_,dtype=torch.int32,device=dev)
    kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
    out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    nsi=_uniform_nsi(dev,bs,ns)
    logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
    lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
def _g256s2(dev,bs,kvl,ps):
    key=(bs,kvl,ps,"256split1");c=_s2c.get(key)
    if c: return c
    ppb=kvl//ps;np_=bs*kvl//ps;ns=1
    ki=torch.arange(np_,dtype=torch.int32,device=dev)
    kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
    kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
    out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
    qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
    nsi=_uniform_nsi(dev,bs,ns)
    logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
    lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
    c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
_CFG = {
    (4,8192):  (8, False, None, None),
    (64,8192): (8, False, 3, None),
    (256,8192):(8, False, None, None),
    (64,1024): (2, False, 3, None),
    (256,1024):(2, False, 1, 8),
}
def custom_kernel(data: input_t) -> output_t:
    q,kv_data,qo_indptr,kv_indptr,config=data
    bs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.device
    kv_fp8,kv_scale=kv_data["fp8"];qs=_gs(dev)
    if kvl == 1024 and bs <= 32:
        ps=2;kb=kv_fp8.view(-1,ps,1,QKD);c=_gnc(dev,bs,kvl,ps);_quant(c[4],q,qs)
        _fwd(c[4],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)
        return c[3]
    if bs == 32 and kvl == 8192:
        kb=kv_fp8.view(-1,8,1,QKD);c=_g32s2(dev,bs,kvl,8);_quant(c[4],q,qs)
        _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
        _safe_reduce_branchless[(bs,NH)](
            c[7],c[8],c[3],qo_indptr,c[2],c[6],
            c[8].stride(0),c[8].stride(2),c[8].stride(1),
            c[3].stride(0),c[3].stride(1),
            MAYBE_FINAL_OUT=False,BATCH_NUM=bs,BLOCK_DV=VD,Lv=VD,mgc=MGC,
            num_warps=2,num_stages=1,waves_per_eu=4
        )
        return c[3]
    if bs == 64 and kvl == 8192:
        kb=kv_fp8.view(-1,8,1,QKD);c=_g64s2(dev,bs,kvl,8);_quant(c[4],q,qs)
        _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
        _safe_reduce_branchless[(bs,NH)](
            c[7],c[8],c[3],qo_indptr,c[2],c[6],
            c[8].stride(0),c[8].stride(2),c[8].stride(1),
            c[3].stride(0),c[3].stride(1),
            MAYBE_FINAL_OUT=False,BATCH_NUM=bs,BLOCK_DV=VD,Lv=VD,mgc=MGC,
            num_warps=2,num_stages=1,waves_per_eu=4
        )
        return c[3]
    if bs == 256 and kvl == 8192:
        kb=kv_fp8.view(-1,8,1,QKD);c=_g256s2(dev,bs,kvl,8);_quant(c[4],q,qs)
        c[7].zero_()
        c[8].fill_(-float("inf"))
        c[3].zero_()
        _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
        return c[3]
    ps,intra,split_ov,kg_ov = _CFG.get((bs,kvl), (2, True, None, None))
    kb=kv_fp8.view(-1,ps,1,QKD);c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=intra,split_override=split_ov,kg_override=kg_ov);_quant(c[11],q,qs)
    _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
    if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
    return c[3]
scrolls · 183 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 635670.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """V1261: Best of v012g safety + v1260 ps=8 speed.
- NP only for 4×1K, 32×1K (safest, small batch, safe reduce patched).
- Persistent wrapper for 64×1K, 256×1K (from v012g — avoids NP for large batch 1K).
- Direct _s1+_rd ps=8 intra=False for ALL 8K (from v1260 — ps=8 is 10µs faster than ps=4).
- Target ~31.5µs."""
+ """V023a: v022o plus lighter custom 32x8k Triton stage2 launch."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
- import torch
- import triton
- import triton.language as tl
+ os.environ.setdefault("HIPBLASLT_ALLOW_FLUSH_DENORM", "1")
+ os.environ.setdefault("GPU_MAX_HW_QUEUES", "2")
+ import torch, triton, triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes, mla as aiter_mla
⋯ 14 unchanged lines
except Exception:
_rd = aiter.mla_reduce_v1
_mi = aiter.get_mla_metadata_info_v1
- _orig_fwd = aiter.mla.mla_decode_fwd
-
@triton.jit
def _safe_reduce_branchless(
Mid_O, Mid_lse, O, qo_indptr, kv_indptr, num_kv_splits_indptr,
⋯ 2 unchanged lines
MAYBE_FINAL_OUT: tl.constexpr, BATCH_NUM: tl.constexpr,
BLOCK_DV: tl.constexpr, Lv: tl.constexpr, mgc: tl.constexpr,
):
- cur_batch = tl.program_id(0)
- cur_head = tl.program_id(1)
- cur_qo_start = tl.load(qo_indptr + cur_batch)
- cur_qo_end = tl.load(qo_indptr + cur_batch + 1)
- cur_split_start = tl.load(num_kv_splits_indptr + cur_batch)
- cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)
+ cur_batch = tl.program_id(0); cur_head = tl.program_id(1)
+ cur_qo_start = tl.load(qo_indptr + cur_batch); cur_qo_end = tl.load(qo_indptr + cur_batch + 1)
+ cur_split_start = tl.load(num_kv_splits_indptr + cur_batch); cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)
num_max_kv_splits = tl.load(num_kv_splits_indptr + BATCH_NUM)
cur_kv_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)
- offs_d = tl.arange(0, BLOCK_DV)
- mask_d = offs_d < Lv
+ offs_d = tl.arange(0, BLOCK_DV); mask_d = offs_d < Lv
offs_logic = cur_qo_start * stride_mid_ob + cur_head * stride_mid_oh
offs_v = offs_logic * Lv + offs_d
num_valid_kv_splits = tl.minimum(cur_split_end - cur_split_start, tl.cdiv(cur_kv_seq_len, mgc))
⋯ 4 unchanged lines
out = tl.load(input_ptr + Lv * (cur_qo * stride_mid_os + cur_head * stride_mid_oh) + offs_d, mask=mask_d, other=0.0)
tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, out, mask=mask_d)
else:
- e_sum = 0.0
- e_max = -float("inf")
- acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)
+ e_sum = 0.0; e_max = -float("inf"); acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)
for split_kv_id in range(0, num_valid_kv_splits):
tv = tl.load(Mid_O + offs_v + split_kv_id * stride_mid_os * Lv, mask=mask_d, other=0.0)
tlogic = tl.load(Mid_lse + offs_logic + split_kv_id * stride_mid_os)
- tlogic = tl.where(tlogic == tlogic, tlogic, -1e30)
- tlogic = tl.minimum(tlogic, 1e30)
- n_e_max = tl.maximum(tlogic, e_max)
- old_scale = tl.exp(e_max - n_e_max)
- acc *= old_scale
- exp_logic = tl.exp(tlogic - n_e_max)
- acc += exp_logic * tv
- e_sum = e_sum * old_scale + exp_logic
- e_max = n_e_max
- offs_logic += stride_mid_ob
- offs_v += stride_mid_ob * Lv
+ tlogic = tl.where(tlogic == tlogic, tlogic, -1e30); tlogic = tl.minimum(tlogic, 1e30)
+ n_e_max = tl.maximum(tlogic, e_max); old_scale = tl.exp(e_max - n_e_max)
+ acc *= old_scale; exp_logic = tl.exp(tlogic - n_e_max)
+ acc += exp_logic * tv; e_sum = e_sum * old_scale + exp_logic; e_max = n_e_max
+ offs_logic += stride_mid_ob; offs_v += stride_mid_ob * Lv
tl.store(O + cur_qo * stride_obs + cur_head * stride_oh + offs_d, acc / e_sum, mask=mask_d)
-
aiter.mla._fwd_kernel_stage2_asm = _safe_reduce_branchless
_fwd = aiter.mla.mla_decode_fwd
-
NH=16;QKD=576;VD=512;SM=float(1.0/(QKD**0.5));FP8=aiter_dtypes.fp8;BF16=torch.bfloat16
- QSV=float(2.0/torch.finfo(FP8).max);_st={};_pc={};_nc={}
+ QSV=float(2.0/torch.finfo(FP8).max);MGC=64;_st={};_pc={};_nc={};_s2c={}
def _gs(d):
t=_st.get(0)
if t is None: t=torch.tensor([QSV],dtype=torch.float32,device=d);_st[0]=t
⋯ 5 unchanged lines
ki=torch.arange(np_,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb;out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev);c=(ki,kl,kip,out,qf);_nc[key]=c;return c
- def _gpc(dev,qo,bs,kvl,ps,intra=True):
- key=(bs,kvl,ps,intra);c=_pc.get(key)
+ def _gpc(dev,qo,bs,kvl,ps,intra=True,split_override=None,kg_override=None):
+ key=(bs,kvl,ps,intra,split_override,kg_override);c=_pc.get(key)
if c: return c
tot=bs*kvl;ppb=kvl//ps;np_=tot//ps
- ns,_=aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)
+ ns = split_override if split_override is not None else aiter_mla.get_meta_param(None,bs,tot,NH,1,FP8)[0]
ki=torch.arange(np_,dtype=torch.int32,device=dev);kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb;out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
- kg=max(ps,16)
+ kg=kg_override if kg_override is not None else max(ps,16)
info=_mi(bs,1,NH,FP8,FP8,is_sparse=False,fast_mode=True,num_kv_splits=ns,intra_batch_mode=intra)
w=[torch.empty(s,dtype=t,device=dev) for s,t in info]
_mv1(qo,kip,kl,16,1,False,w[0],w[2],w[1],w[3],w[4],w[5],page_size=ps,kv_granularity=kg,max_seqlen_qo=1,uni_seqlen_qo=1,fast_mode=True,max_split_per_batch=ns,intra_batch_mode=intra,dtype_q=FP8,dtype_kv=FP8)
qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
- pt=int(w[5].numel())
- po=torch.empty((pt,1,NH,VD),dtype=torch.float32,device=dev)
- pl=torch.empty((pt,1,NH,1),dtype=torch.float32,device=dev)
- c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,intra,po,pl);_pc[key]=c;return c
+ pt=int(w[5].numel());po=torch.empty((pt,1,NH,VD),dtype=torch.float32,device=dev);pl=torch.empty((pt,1,NH,1),dtype=torch.float32,device=dev)
+ c=(ki,kl,kip,out,w[0],w[1],w[2],w[3],w[4],w[5],ns,qf,po,pl);_pc[key]=c;return c
+ def _uniform_nsi(dev,bs,ns):
+ return torch.arange(bs+1,dtype=torch.int32,device=dev) * ns
+ def _g32s2(dev,bs,kvl,ps):
+ key=(bs,kvl,ps,"32split3");c=_s2c.get(key)
+ if c: return c
+ ppb=kvl//ps;np_=bs*kvl//ps;ns=3
+ ki=torch.arange(np_,dtype=torch.int32,device=dev)
+ kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
+ kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
+ out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
+ qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
+ nsi=_uniform_nsi(dev,bs,ns)
+ logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
+ lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
+ c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
+ def _g64s2(dev,bs,kvl,ps):
+ key=(bs,kvl,ps,"64split2w2");c=_s2c.get(key)
+ if c: return c
+ ppb=kvl//ps;np_=bs*kvl//ps;ns=2
+ ki=torch.arange(np_,dtype=torch.int32,device=dev)
+ kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
+ kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
+ out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
+ qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
+ nsi=_uniform_nsi(dev,bs,ns)
+ logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
+ lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
+ c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
+ def _g256s2(dev,bs,kvl,ps):
+ key=(bs,kvl,ps,"256split1");c=_s2c.get(key)
+ if c: return c
+ ppb=kvl//ps;np_=bs*kvl//ps;ns=1
+ ki=torch.arange(np_,dtype=torch.int32,device=dev)
+ kl=torch.full((bs,),ps,dtype=torch.int32,device=dev)
+ kip=torch.arange(bs+1,dtype=torch.int32,device=dev)*ppb
+ out=torch.empty((bs,NH,VD),dtype=BF16,device=dev)
+ qf=torch.empty((bs,NH,QKD),dtype=FP8,device=dev)
+ nsi=_uniform_nsi(dev,bs,ns)
+ logits=torch.empty((bs,ns,NH,VD),dtype=torch.float32,device=dev)
+ lse=torch.empty((bs,ns,NH,1),dtype=torch.float32,device=dev)
+ c=(ki,kl,kip,out,qf,ns,nsi,logits,lse);_s2c[key]=c;return c
+ _CFG = {
+ (4,8192): (8, False, None, None),
+ (64,8192): (8, False, 3, None),
+ (256,8192):(8, False, None, None),
+ (64,1024): (2, False, 3, None),
+ (256,1024):(2, False, 1, 8),
+ }
def custom_kernel(data: input_t) -> output_t:
q,kv_data,qo_indptr,kv_indptr,config=data
bs=int(config["batch_size"]);kvl=int(config["kv_seq_len"]);dev=q.device
- kv_fp8,kv_scale=kv_data["fp8"]
- qs=_gs(dev)
- # NP: only 4×1K, 32×1K (small batch, safe reduce patched)
+ kv_fp8,kv_scale=kv_data["fp8"];qs=_gs(dev)
if kvl == 1024 and bs <= 32:
- ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
- c=_gnc(dev,bs,kvl,ps);_quant(c[4],q,qs)
+ ps=2;kb=kv_fp8.view(-1,ps,1,QKD);c=_gnc(dev,bs,kvl,ps);_quant(c[4],q,qs)
_fwd(c[4],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,q_scale=qs,kv_scale=kv_scale,intra_batch_mode=True)
return c[3]
- # ALL 8K: direct _s1+_rd ps=8 intra=False
- if kvl == 8192:
- ps=8;kb=kv_fp8.view(-1,ps,1,QKD)
- c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=False);_quant(c[11],q,qs)
- _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,
- c[4],c[5],c[6],1,ps,1,SM,c[13],c[14],c[3],qs,kv_scale)
- if c[10]>1: _rd(c[13],c[14],c[7],c[8],c[9],1,c[3],None)
+ if bs == 32 and kvl == 8192:
+ kb=kv_fp8.view(-1,8,1,QKD);c=_g32s2(dev,bs,kvl,8);_quant(c[4],q,qs)
+ _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
+ _safe_reduce_branchless[(bs,NH)](
+ c[7],c[8],c[3],qo_indptr,c[2],c[6],
+ c[8].stride(0),c[8].stride(2),c[8].stride(1),
+ c[3].stride(0),c[3].stride(1),
+ MAYBE_FINAL_OUT=False,BATCH_NUM=bs,BLOCK_DV=VD,Lv=VD,mgc=MGC,
+ num_warps=2,num_stages=1,waves_per_eu=4
+ )
return c[3]
- # 64×1K, 256×1K: persistent wrapper (safe, ASM reduce internally)
- ps=2;kb=kv_fp8.view(-1,ps,1,QKD)
- c=_gpc(dev,qo_indptr,bs,kvl,ps);_quant(c[11],q,qs)
- _orig_fwd(c[11],kb,c[3],qo_indptr,c[2],c[0],c[1],1,ps,1,SM,
- num_kv_splits=c[10],work_meta_data=c[4],work_indptr=c[5],work_info_set=c[6],
- reduce_indptr=c[7],reduce_final_map=c[8],reduce_partial_map=c[9],
- q_scale=qs,kv_scale=kv_scale,intra_batch_mode=c[12])
+ if bs == 64 and kvl == 8192:
+ kb=kv_fp8.view(-1,8,1,QKD);c=_g64s2(dev,bs,kvl,8);_quant(c[4],q,qs)
+ _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
+ _safe_reduce_branchless[(bs,NH)](
+ c[7],c[8],c[3],qo_indptr,c[2],c[6],
+ c[8].stride(0),c[8].stride(2),c[8].stride(1),
+ c[3].stride(0),c[3].stride(1),
+ MAYBE_FINAL_OUT=False,BATCH_NUM=bs,BLOCK_DV=VD,Lv=VD,mgc=MGC,
+ num_warps=2,num_stages=1,waves_per_eu=4
+ )
+ return c[3]
+ if bs == 256 and kvl == 8192:
+ kb=kv_fp8.view(-1,8,1,QKD);c=_g256s2(dev,bs,kvl,8);_quant(c[4],q,qs)
+ c[7].zero_()
+ c[8].fill_(-float("inf"))
+ c[3].zero_()
+ _s1(c[4],kb,qo_indptr,c[2],c[0],c[1],c[6],None,None,None,1,8,1,SM,c[7],c[8],c[3],qs,kv_scale)
+ return c[3]
+ ps,intra,split_ov,kg_ov = _CFG.get((bs,kvl), (2, True, None, None))
+ kb=kv_fp8.view(-1,ps,1,QKD);c=_gpc(dev,qo_indptr,bs,kvl,ps,intra=intra,split_override=split_ov,kg_override=kg_ov);_quant(c[11],q,qs)
+ _s1(c[11],kb,qo_indptr,c[2],c[0],c[1],None,c[4],c[5],c[6],1,ps,1,SM,c[12],c[13],c[3],qs,kv_scale)
+ if c[10]>1: _rd(c[12],c[13],c[7],c[8],c[9],1,c[3],None)
return c[3]
scrolls · 223 diff lines total

Best evidence level for this revision: reported

JSON