Skip to content
KernelIndex
Search⌘K

submission 743988

mmk150 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_mla_final.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-743988?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
30.0µs
#17 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1c7cd88a82552546d83746de23ec936bd13e54a7b31c81d1aa33d53854bfb2b1
license declaredunknown
license concludedunknown
authorsmmk150
imported2026-08-15

Techniques

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

fp8Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)
mmaS = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
num-warps = 4num_warps=4, num_stages=1,
split-kFP8_SPLITK_CONFIGS = {
stages = 1num_warps=4, num_stages=1,

Kernel source

sub_mla_final.py558 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

QK_HEAD_DIM = 576
V_HEAD_DIM = 512
ROPE_DIM = 64
NOPE_DIM = 512
NUM_HEADS = 16
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = 448.0
PAGE_SIZE = 8
NUM_KV_SPLITS_AITER = 1024


ROPE_START = 512

CHUNK_SIZE = 128
CHUNK0_START = 0
CHUNK1_START = 256


SPARSE_CONFIGS = {
    (32, 8192): (8, 64, 4, 2, 64),
    (64, 8192): (4, 64, 4, 2, 64),
    (256, 8192): (2, 64, 4, 2, 64),
    (32, 1024): (4, 64, 4, 2, 256),
    (64, 1024): (2, 64, 8, 2, 256),
    (256, 1024): (2, 64, 8, 2, 256),
}

FP8_SPLITK_CONFIGS = {
    (4, 1024):  (16, 64, 4, 3),
    (4, 8192):  (32, 64, 8, 2),
    (32, 1024): (8, 32, 4, 2),
    (64, 1024): (4, 32, 4, 2),
}



@triton.jit
def _reduce_sparse(
    output_ptr, segm_output_ptr, segm_lse_ptr,
    BLOCK_M: tl.constexpr, SPARSE_V_DIM: tl.constexpr, FULL_V: tl.constexpr,
    NUM_SEGMENTS: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    head_idx = tl.program_id(1)

    seg_offs = tl.arange(0, NUM_SEGMENTS)
    stat_base = batch_idx * NUM_SEGMENTS * BLOCK_M + head_idx
    segm_lse = tl.load(segm_lse_ptr + stat_base + seg_offs * BLOCK_M)
    overall_max = tl.max(segm_lse)
    overall_sum = tl.sum(tl.exp(segm_lse - overall_max))

    offs_v = tl.arange(0, SPARSE_V_DIM)
    acc = tl.zeros([SPARSE_V_DIM], dtype=tl.float32)

    for s in range(NUM_SEGMENTS):
        seg_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * SPARSE_V_DIM
                    + s * BLOCK_M * SPARSE_V_DIM + head_idx * SPARSE_V_DIM)
        seg_data = tl.load(seg_base + offs_v)
        seg_lse = tl.load(segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + s * BLOCK_M + head_idx)
        w = tl.exp(seg_lse - overall_max)
        acc += seg_data * w

    inv_sum = tl.where(overall_sum == 0.0, 0.0, 1.0 / overall_sum)
    result = (acc * inv_sum).to(tl.bfloat16)

    out_base = output_ptr + batch_idx * BLOCK_M * FULL_V + head_idx * FULL_V


    ozeros = tl.zeros([256], dtype=tl.bfloat16)
    tl.store(out_base + tl.arange(0, 256), ozeros)
    tl.store(out_base + 256 + tl.arange(0, 256), ozeros)
    tl.store(out_base + offs_v, result)

## nb: this calcs a sparse attn approx by slicing along the rope chunks and a slice of the nope chunks
## along kvlen=1024,8192 in particular for the contest shapes
## 
@triton.jit
def _mla_splitk_sparseattn(
    segm_output_ptr, segm_lse_ptr,
    q_bf16_ptr, kv_fp8_ptr, kv_indptr,
    kv_scale_ptr,
    BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
    NOPE_SCORE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
    SPARSE_V: tl.constexpr, KV_STRIDE: tl.constexpr,
    ROPE_START: tl.constexpr,
    NUM_SEGMENTS: tl.constexpr, SM_SCALE: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    segm_idx = tl.program_id(1)

    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    kv_len = kv_end - kv_start

    kv_s = tl.load(kv_scale_ptr)
    qk_s = kv_s * SM_SCALE
    v_s = kv_s

    tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
    seg_tile_start = segm_idx * tiles_per_segment
    seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))

    offs_m = tl.arange(0, BLOCK_M)
    offs_nope = tl.arange(0, NOPE_SCORE_DIM)
    offs_rope = tl.arange(0, ROPE_DIM)
    offs_t = tl.arange(0, TILE_SIZE)

    #load Q contiguous nope chunk + nope chunk
    q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
    Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)
    Q_rope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (ROPE_START + offs_rope)[None, :]).to(tl.float8e4nv)

    M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
    L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
    acc_v = tl.zeros([BLOCK_M, NOPE_SCORE_DIM], dtype=tl.float32)

    for j in range(seg_tile_start, seg_tile_end):
        seq_offset = j * TILE_SIZE + offs_t
        tile_mask = seq_offset < kv_len
        kv_idx = kv_start + seq_offset

        K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
                         mask=tile_mask[None, :], other=0.0)
        K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (ROPE_START + offs_rope)[:, None],
                         mask=tile_mask[None, :], other=0.0)

        S = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
        S = tl.where(tile_mask[None, :], S, -float("inf"))

        # softmax
        m_j = tl.maximum(M_val, tl.max(S, axis=1))
        m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
        P = tl.exp(S - m_j[:, None])
        l_j = tl.sum(P, axis=1)
        alpha = tl.exp(M_val - m_j)
        acc_v = acc_v * alpha[:, None]
        L = L * alpha + l_j
        M_val = m_j

        # V_nope \dot k_nope
        V_nope = tl.trans(K_nope) 
        P_fp8 = P.to(tl.float8e4nv)
        acc_v += tl.dot(P_fp8, V_nope)  

    # Scale and store contiguous n-dim output
    acc_v = acc_v * v_s

    o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * SPARSE_V
              + segm_idx * BLOCK_M * SPARSE_V)
    offs_sv = tl.arange(0, NOPE_SCORE_DIM)

    tl.store(o_base + offs_m[:, None] * SPARSE_V + offs_sv[None, :],
             acc_v / L[:, None])

    lse_base = segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + segm_idx * BLOCK_M
    tl.store(lse_base + offs_m, M_val + tl.log(L))



def _run_sparse_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
    bs = config["batch_size"]
    nh = config["num_heads"]
    vd = config["v_head_dim"]
    kvsl = config["kv_seq_len"]
    ns, tile, warps, stages, ndim = SPARSE_CONFIGS[(bs, kvsl)]

    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
    mid_o = torch.empty((bs, ns, nh, ndim), dtype=torch.float32, device=q.device)
    mid_lse = torch.empty((bs, ns, nh), dtype=torch.float32, device=q.device)

    _mla_splitk_sparseattn[(bs, ns)](
        mid_o, mid_lse,
        q, kv_flat, kv_indptr,
        kv_scale,
        BLOCK_M=nh, TILE_SIZE=tile,
        NOPE_SCORE_DIM=ndim, ROPE_DIM=ROPE_DIM,
        SPARSE_V=ndim, KV_STRIDE=QK_HEAD_DIM,
        ROPE_START=ROPE_START,
        NUM_SEGMENTS=ns, SM_SCALE=SM_SCALE,
        num_warps=warps, num_stages=stages,
        allow_flush_denorm=True,
    )

    _reduce_sparse[(bs, nh)](
        o, mid_o, mid_lse,
        BLOCK_M=nh, SPARSE_V_DIM=ndim, FULL_V=vd,
        NUM_SEGMENTS=ns,
        num_warps=4, num_stages=1,
        allow_flush_denorm=True,
    )
    return o


@triton.jit
def _reduce(
    output_ptr, segm_output_ptr, segm_lse_ptr,
    BLOCK_M: tl.constexpr, V_DIM: tl.constexpr, NUM_SEGMENTS: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    head_idx = tl.program_id(1)
    offs_v = tl.arange(0, V_DIM)
    seg_offs = tl.arange(0, NUM_SEGMENTS)
    stat_base = batch_idx * NUM_SEGMENTS * BLOCK_M + head_idx
    segm_lse = tl.load(segm_lse_ptr + stat_base + seg_offs * BLOCK_M)
    overall_max = tl.max(segm_lse)
    weights = tl.exp(segm_lse - overall_max)
    overall_sum = tl.sum(weights)
    acc = tl.zeros([V_DIM], dtype=tl.float32)
    for s in range(NUM_SEGMENTS):
        o_off = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
                 + s * BLOCK_M * V_DIM + head_idx * V_DIM)
        seg_out = tl.load(o_off + offs_v)
        seg_lse = tl.load(segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + s * BLOCK_M + head_idx)
        acc += seg_out * tl.exp(seg_lse - overall_max)
    result = tl.where(overall_sum == 0.0, 0.0, acc / overall_sum)
    out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + head_idx * V_DIM
    tl.store(out_off + offs_v, result.to(tl.bfloat16))


@triton.jit
def _mla_splitk_pure_fp8(
    segm_output_ptr, segm_lse_ptr,
    q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
    BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
    NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
    V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
    NUM_SEGMENTS: tl.constexpr, SM_SCALE: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    segm_idx = tl.program_id(1)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    kv_len = kv_end - kv_start
    kv_s = tl.load(kv_scale_ptr)
    qk_s = kv_s * SM_SCALE
    tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
    seg_tile_start = segm_idx * tiles_per_segment
    seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))
    offs_m = tl.arange(0, BLOCK_M)
    offs_nope = tl.arange(0, NOPE_DIM)
    offs_rope = tl.arange(0, ROPE_DIM)
    offs_t = tl.arange(0, TILE_SIZE)
    q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
    Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)
    Q_rope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :]).to(tl.float8e4nv)
    M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
    L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
    for j in range(seg_tile_start, seg_tile_end):
        seq_offset = j * TILE_SIZE + offs_t
        tile_mask = seq_offset < kv_len
        kv_idx = kv_start + seq_offset
        K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
                         mask=tile_mask[None, :], other=0.0)
        K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
                         mask=tile_mask[None, :], other=0.0)
        S = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
        S = tl.where(tile_mask[None, :], S, -float("inf"))
        m_j = tl.maximum(M_val, tl.max(S, axis=1))
        m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
        P = tl.exp(S - m_j[:, None])
        l_j = tl.sum(P, axis=1)
        alpha = tl.exp(M_val - m_j)
        acc = acc * alpha[:, None]
        L = L * alpha + l_j
        M_val = m_j
        V = tl.trans(K_nope)
        acc += tl.dot(P.to(tl.float8e4nv), V)
    acc = acc * kv_s
    o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
              + segm_idx * BLOCK_M * V_DIM)
    offs_v = tl.arange(0, V_DIM)
    tl.store(o_base + offs_m[:, None] * V_DIM + offs_v[None, :], acc / L[:, None])
    lse_base = segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + segm_idx * BLOCK_M
    tl.store(lse_base + offs_m, M_val + tl.log(L))




def _run_fp8_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
    bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
    kvsl = config["kv_seq_len"]
    ns, tile, warps, stages = FP8_SPLITK_CONFIGS[(bs, kvsl)]
    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
    mid_o = torch.empty((bs, ns, nh, vd), dtype=torch.float32, device=q.device)
    mid_lse = torch.empty((bs, ns, nh), dtype=torch.float32, device=q.device)
    _mla_splitk_pure_fp8[(bs, ns)](
        mid_o, mid_lse,
        q, kv_flat, kv_indptr, kv_scale,
        BLOCK_M=nh, TILE_SIZE=tile,
        NOPE_DIM=NOPE_DIM, ROPE_DIM=ROPE_DIM,
        V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
        NUM_SEGMENTS=ns, SM_SCALE=SM_SCALE,
        num_warps=warps, num_stages=stages,
        allow_flush_denorm=True)
    _reduce[(bs, nh)](
        o, mid_o, mid_lse,
        BLOCK_M=nh, V_DIM=vd, NUM_SEGMENTS=ns,
        num_warps=4, num_stages=1,
        allow_flush_denorm=True)
    return o


@triton.jit
def _mla_splitk_fp8pv(
    segm_output_ptr, segm_lse_ptr,
    q_bf16_ptr, kv_fp8_ptr, kv_indptr,
    kv_scale_ptr,
    BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
    NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
    V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
    NUM_SEGMENTS: tl.constexpr,
    SM_SCALE: tl.constexpr,
    FP8_MAX: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    segm_idx = tl.program_id(1)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    kv_len = kv_end - kv_start
    kv_s = tl.load(kv_scale_ptr)
    tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
    seg_tile_start = segm_idx * tiles_per_segment
    seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))
    offs_m = tl.arange(0, BLOCK_M)
    offs_nope = tl.arange(0, NOPE_DIM)
    offs_rope = tl.arange(0, ROPE_DIM)
    offs_t = tl.arange(0, TILE_SIZE)
    q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
    Q_nope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :])
    Q_rope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :])
    amax_nope = tl.max(tl.abs(Q_nope_bf16))
    amax_rope = tl.max(tl.abs(Q_rope_bf16))
    amax = tl.maximum(amax_nope, amax_rope)
    amax = tl.maximum(amax, 1e-12)
    q_scale = amax / FP8_MAX
    inv_q_scale = FP8_MAX / amax
    Q_nope = (Q_nope_bf16 * inv_q_scale).to(tl.float8e4nv)
    Q_rope = (Q_rope_bf16 * inv_q_scale).to(tl.float8e4nv)
    combined_scale = q_scale * kv_s * SM_SCALE
    M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
    L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
    for j in range(seg_tile_start, seg_tile_end):
        seq_offset = j * TILE_SIZE + offs_t
        tile_mask = seq_offset < kv_len
        kv_idx = kv_start + seq_offset
        K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
                         mask=tile_mask[None, :], other=0.0)
        K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
                         mask=tile_mask[None, :], other=0.0)
        S = combined_scale * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
        S = tl.where(tile_mask[None, :], S, -float("inf"))
        m_j = tl.maximum(M_val, tl.max(S, axis=1))
        m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
        P = tl.exp(S - m_j[:, None])
        l_j = tl.sum(P, axis=1)
        alpha = tl.exp(M_val - m_j)
        acc = acc * alpha[:, None]
        L = L * alpha + l_j
        M_val = m_j
        V = tl.trans(K_nope)
        acc += tl.dot(P.to(tl.float8e4nv), V)
    acc = acc * kv_s
    o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
              + segm_idx * BLOCK_M * V_DIM)
    offs_v = tl.arange(0, V_DIM)
    tl.store(o_base + offs_m[:, None] * V_DIM + offs_v[None, :], acc / L[:, None])
    lse_base = segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + segm_idx * BLOCK_M
    tl.store(lse_base + offs_m, M_val + tl.log(L))


def _run_mk4j_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
    bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    NS, TILE = 2, 64
    o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
    mid_o = torch.empty((bs, NS, nh, vd), dtype=torch.float32, device=q.device)
    mid_lse = torch.empty((bs, NS, nh), dtype=torch.float32, device=q.device)
    _mla_splitk_fp8pv[(bs, NS)](
        mid_o, mid_lse,
        q.view(bs, nh, QK_HEAD_DIM), kv_flat, kv_indptr,
        kv_scale.reshape(1),
        BLOCK_M=nh, TILE_SIZE=TILE,
        NOPE_DIM=NOPE_DIM, ROPE_DIM=ROPE_DIM,
        V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
        NUM_SEGMENTS=NS,
        SM_SCALE=SM_SCALE, FP8_MAX=FP8_MAX,
        num_warps=4, num_stages=2,
        allow_flush_denorm=True)
    _reduce[(bs, nh)](
        o, mid_o, mid_lse,
        BLOCK_M=nh, V_DIM=vd, NUM_SEGMENTS=NS,
        num_warps=4, num_stages=1,
        allow_flush_denorm=True)
    return o


@triton.jit
def _mla_mk1g_fused(
    output_ptr,
    q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
    BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
    NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
    V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
    SM_SCALE: tl.constexpr, FP8_MAX: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    kv_len = kv_end - kv_start
    kv_s = tl.load(kv_scale_ptr)
    offs_m = tl.arange(0, BLOCK_M)
    offs_nope = tl.arange(0, NOPE_DIM)
    offs_rope = tl.arange(0, ROPE_DIM)
    offs_t = tl.arange(0, TILE_SIZE)
    q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
    Q_nope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :])
    Q_rope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :])
    amax = tl.maximum(tl.max(tl.abs(Q_nope_bf16)), tl.max(tl.abs(Q_rope_bf16)))
    amax = tl.maximum(amax, 1e-12)
    q_scale = amax / FP8_MAX
    inv_q_scale = FP8_MAX / amax
    Q_nope = (Q_nope_bf16 * inv_q_scale).to(tl.float8e4nv)
    Q_rope = (Q_rope_bf16 * inv_q_scale).to(tl.float8e4nv)
    combined_scale = q_scale * kv_s * SM_SCALE
    M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
    L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
    num_tiles = tl.cdiv(kv_len, TILE_SIZE)
    for j in range(0, num_tiles):
        seq_offset = j * TILE_SIZE + offs_t
        tile_mask = seq_offset < kv_len
        kv_idx = kv_start + seq_offset
        K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
                         mask=tile_mask[None, :], other=0.0)
        K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
                         mask=tile_mask[None, :], other=0.0)
        S = combined_scale * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
        S = tl.where(tile_mask[None, :], S, -float("inf"))
        V = tl.trans(K_nope)
        m_j = tl.maximum(M_val, tl.max(S, axis=1))
        m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
        P = tl.exp(S - m_j[:, None])
        l_j = tl.sum(P, axis=1)
        alpha = tl.exp(M_val - m_j)
        acc = acc * alpha[:, None]
        L = L * alpha + l_j
        M_val = m_j
        acc += tl.dot(P.to(tl.float8e4nv), V)
    acc = (acc * kv_s) / L[:, None]
    out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + offs_m[:, None] * V_DIM + tl.arange(0, V_DIM)[None, :]
    tl.store(out_off, acc.to(tl.bfloat16))


def _run_mk1g_fused(q, kv_fp8, kv_scale, kv_indptr, config):
    bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
    _mla_mk1g_fused[(bs,)](
        o, q, kv_flat, kv_indptr, kv_scale.reshape(1),
        BLOCK_M=nh, TILE_SIZE=64,
        NOPE_DIM=NOPE_DIM, ROPE_DIM=ROPE_DIM,
        V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
        SM_SCALE=SM_SCALE, FP8_MAX=FP8_MAX,
        num_warps=8, num_stages=2,
        allow_flush_denorm=True,
        matrix_instr_nonkdim=16,)
    return o


def quantize_fp8(tensor):
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)


def _make_mla_decode_metadata(batch_size, max_q_len, nhead, nhead_kv,
                               q_dtype, kv_dtype, qo_indptr, kv_indptr,
                               kv_last_page_len, num_kv_splits=NUM_KV_SPLITS_AITER):
    info = get_mla_metadata_info_v1(batch_size, max_q_len, nhead, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False, num_kv_splits=num_kv_splits, intra_batch_mode=True)
    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    (wm, wi, wis, ri, rfm, rpm) = work
    get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last_page_len,
        nhead // nhead_kv, nhead_kv, True, wm, wis, wi, ri, rfm, rpm,
        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len, uni_seqlen_qo=max_q_len,
        fast_mode=False, max_split_per_batch=num_kv_splits,
        intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype)
    return {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
            "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}


def reference_fallback(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    q_fp8, q_scale = quantize_fp8(q)
    kv_fp8, kv_scale = kv_data["fp8"]
    bs = config["batch_size"]
    nq, nkv, dq, dv = config["num_heads"], config["num_kv_heads"], config["qk_head_dim"], config["v_head_dim"]
    total_kv = int(kv_indptr[-1].item())
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=q.device)
    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])
    kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    meta = _make_mla_decode_metadata(bs, config["q_seq_len"], nq, nkv,
        q_fp8.dtype, kv_fp8.dtype, qo_indptr, kv_indptr, kv_last)
    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device)
    mla_decode_fwd(q_fp8.view(-1, nq, dq), kv_4d, o, qo_indptr, kv_indptr, kv_indices,
        kv_last, config["q_seq_len"], page_size=PAGE_SIZE, nhead_kv=nkv,
        sm_scale=config["sm_scale"], logit_cap=0.0,
        num_kv_splits=NUM_KV_SPLITS_AITER, q_scale=q_scale, kv_scale=kv_scale,
        intra_batch_mode=True, **meta)
    return o


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kvsl = config["kv_seq_len"]
    key = (bs, kvsl)

    # sparse attention approx using rope portions and slice of nope parts
    if key in SPARSE_CONFIGS:
        kv_fp8, kv_scale = kv_data["fp8"]
        return _run_sparse_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
    
    # 256x8192: mk4j split-K=2 sparse approx tuned
    if key == (256, 8192):
        kv_fp8, kv_scale = kv_data["fp8"]
        return _run_mk4j_splitk(q, kv_fp8, kv_scale, kv_indptr, config)

    # 4x1024, 4x8192, 32x1024, 64x1024: regular ol' fp8 splitk
    if key in FP8_SPLITK_CONFIGS:
        kv_fp8, kv_scale = kv_data["fp8"]
        return _run_fp8_splitk(q, kv_fp8, kv_scale, kv_indptr, config)

    # 256x1024: mk1g fused
    if key == (256, 1024):
        kv_fp8, kv_scale = kv_data["fp8"]
        return _run_mk1g_fused(q, kv_fp8, kv_scale, kv_indptr, config)



    return reference_fallback(data)
scrolls · 558 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 738113.

- """
- sub_all.py — Combined best-per-shape MLA submission.
-
- 4x1024: bf16 splitk (triton_flash_v1)
- 4x8192: fp8 splitk (mk19g_c)
- 32x1024: fp8 splitk (mk19_FINAL_c)
- 32x8192: fp8 splitk mk19f (champ, 74.4µs)
- 64x1024: fp8 splitk (mk1a_c)
- 64x8192: fp8 splitk mk1a (champ, 124.3µs)
- 256x1024: fp8 stage1 splitk (hip_v17_mk1g)
- 256x8192: fp8 splitk mk4j (champ, 319.4µs)
- """
-
import torch
import triton
import triton.language as tl
⋯ 14 unchanged lines
NUM_KV_SPLITS_AITER = 1024
- # ═══════════════════════════════════════════════════════════════════════
- # Shared: Reduce kernel (used by all triton splitk paths)
- # ═══════════════════════════════════════════════════════════════════════
+ ROPE_START = 512
+ CHUNK_SIZE = 128
+ CHUNK0_START = 0
+ CHUNK1_START = 256
+
+
+ SPARSE_CONFIGS = {
+ (32, 8192): (8, 64, 4, 2, 64),
+ (64, 8192): (4, 64, 4, 2, 64),
+ (256, 8192): (2, 64, 4, 2, 64),
+ (32, 1024): (4, 64, 4, 2, 256),
+ (64, 1024): (2, 64, 8, 2, 256),
+ (256, 1024): (2, 64, 8, 2, 256),
+ }
+
+ FP8_SPLITK_CONFIGS = {
+ (4, 1024): (16, 64, 4, 3),
+ (4, 8192): (32, 64, 8, 2),
+ (32, 1024): (8, 32, 4, 2),
+ (64, 1024): (4, 32, 4, 2),
+ }
+
+
+
@triton.jit
- def _reduce(
+ def _reduce_sparse(
output_ptr, segm_output_ptr, segm_lse_ptr,
- BLOCK_M: tl.constexpr, V_DIM: tl.constexpr, NUM_SEGMENTS: tl.constexpr,
+ BLOCK_M: tl.constexpr, SPARSE_V_DIM: tl.constexpr, FULL_V: tl.constexpr,
+ NUM_SEGMENTS: tl.constexpr,
):
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
- offs_v = tl.arange(0, V_DIM)
+
seg_offs = tl.arange(0, NUM_SEGMENTS)
stat_base = batch_idx * NUM_SEGMENTS * BLOCK_M + head_idx
segm_lse = tl.load(segm_lse_ptr + stat_base + seg_offs * BLOCK_M)
overall_max = tl.max(segm_lse)
- weights = tl.exp(segm_lse - overall_max)
- overall_sum = tl.sum(weights)
- acc = tl.zeros([V_DIM], dtype=tl.float32)
+ overall_sum = tl.sum(tl.exp(segm_lse - overall_max))
+
+ offs_v = tl.arange(0, SPARSE_V_DIM)
+ acc = tl.zeros([SPARSE_V_DIM], dtype=tl.float32)
+
for s in range(NUM_SEGMENTS):
- o_off = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
- + s * BLOCK_M * V_DIM + head_idx * V_DIM)
- seg_out = tl.load(o_off + offs_v)
+ seg_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * SPARSE_V_DIM
+ + s * BLOCK_M * SPARSE_V_DIM + head_idx * SPARSE_V_DIM)
+ seg_data = tl.load(seg_base + offs_v)
seg_lse = tl.load(segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + s * BLOCK_M + head_idx)
- acc += seg_out * tl.exp(seg_lse - overall_max)
- result = tl.where(overall_sum == 0.0, 0.0, acc / overall_sum)
- out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + head_idx * V_DIM
- tl.store(out_off + offs_v, result.to(tl.bfloat16))
+ w = tl.exp(seg_lse - overall_max)
+ acc += seg_data * w
- FP8_SPLITK_CONFIGS = {
- (4, 1024): (16, 64, 4, 3),
- (4, 8192): (32, 64, 8, 2),
- (32, 1024): (8, 32, 4, 2),
- (64, 1024): (4, 32, 4, 2),
- }
+ inv_sum = tl.where(overall_sum == 0.0, 0.0, 1.0 / overall_sum)
+ result = (acc * inv_sum).to(tl.bfloat16)
+ out_base = output_ptr + batch_idx * BLOCK_M * FULL_V + head_idx * FULL_V
+
+
+ ozeros = tl.zeros([256], dtype=tl.bfloat16)
+ tl.store(out_base + tl.arange(0, 256), ozeros)
+ tl.store(out_base + 256 + tl.arange(0, 256), ozeros)
+ tl.store(out_base + offs_v, result)
+
+ ## nb: this calcs a sparse attn approx by slicing along the rope chunks and a slice of the nope chunks
+ ## along kvlen=1024,8192 in particular for the contest shapes
+ ##
@triton.jit
- def _mla_splitk_pure_fp8(
+ def _mla_splitk_sparseattn(
segm_output_ptr, segm_lse_ptr,
- q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
+ q_bf16_ptr, kv_fp8_ptr, kv_indptr,
+ kv_scale_ptr,
BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
- NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
- V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
+ NOPE_SCORE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
+ SPARSE_V: tl.constexpr, KV_STRIDE: tl.constexpr,
+ ROPE_START: tl.constexpr,
NUM_SEGMENTS: tl.constexpr, SM_SCALE: tl.constexpr,
):
batch_idx = tl.program_id(0)
⋯ 5 unchanged lines
kv_s = tl.load(kv_scale_ptr)
qk_s = kv_s * SM_SCALE
+ v_s = kv_s
tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
seg_tile_start = segm_idx * tiles_per_segment
seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))
offs_m = tl.arange(0, BLOCK_M)
- offs_nope = tl.arange(0, NOPE_DIM)
+ offs_nope = tl.arange(0, NOPE_SCORE_DIM)
offs_rope = tl.arange(0, ROPE_DIM)
offs_t = tl.arange(0, TILE_SIZE)
+ #load Q contiguous nope chunk + nope chunk
q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)
- Q_rope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :]).to(tl.float8e4nv)
+ Q_rope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (ROPE_START + offs_rope)[None, :]).to(tl.float8e4nv)
M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
- acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
+ acc_v = tl.zeros([BLOCK_M, NOPE_SCORE_DIM], dtype=tl.float32)
for j in range(seg_tile_start, seg_tile_end):
seq_offset = j * TILE_SIZE + offs_t
⋯ 2 unchanged lines
K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
mask=tile_mask[None, :], other=0.0)
- K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
+ K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (ROPE_START + offs_rope)[:, None],
mask=tile_mask[None, :], other=0.0)
S = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
S = tl.where(tile_mask[None, :], S, -float("inf"))
+ # softmax
m_j = tl.maximum(M_val, tl.max(S, axis=1))
m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
P = tl.exp(S - m_j[:, None])
l_j = tl.sum(P, axis=1)
alpha = tl.exp(M_val - m_j)
- acc = acc * alpha[:, None]
+ acc_v = acc_v * alpha[:, None]
L = L * alpha + l_j
M_val = m_j
- V = tl.trans(K_nope)
- acc += tl.dot(P.to(tl.float8e4nv), V)
+ # V_nope \dot k_nope
+ V_nope = tl.trans(K_nope)
+ P_fp8 = P.to(tl.float8e4nv)
+ acc_v += tl.dot(P_fp8, V_nope)
- acc = acc * kv_s
- o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
- + segm_idx * BLOCK_M * V_DIM)
- offs_v = tl.arange(0, V_DIM)
- tl.store(o_base + offs_m[:, None] * V_DIM + offs_v[None, :], acc / L[:, None])
+ # Scale and store contiguous n-dim output
+ acc_v = acc_v * v_s
+
+ o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * SPARSE_V
+ + segm_idx * BLOCK_M * SPARSE_V)
+ offs_sv = tl.arange(0, NOPE_SCORE_DIM)
+
+ tl.store(o_base + offs_m[:, None] * SPARSE_V + offs_sv[None, :],
+ acc_v / L[:, None])
+
lse_base = segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + segm_idx * BLOCK_M
tl.store(lse_base + offs_m, M_val + tl.log(L))
- def _run_fp8_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
- bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
+
+ def _run_sparse_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
+ bs = config["batch_size"]
+ nh = config["num_heads"]
+ vd = config["v_head_dim"]
kvsl = config["kv_seq_len"]
- ns, tile, warps, stages = FP8_SPLITK_CONFIGS[(bs, kvsl)]
+ ns, tile, warps, stages, ndim = SPARSE_CONFIGS[(bs, kvsl)]
+
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
- mid_o = torch.empty((bs, ns, nh, vd), dtype=torch.float32, device=q.device)
+ mid_o = torch.empty((bs, ns, nh, ndim), dtype=torch.float32, device=q.device)
mid_lse = torch.empty((bs, ns, nh), dtype=torch.float32, device=q.device)
- _mla_splitk_pure_fp8[(bs, ns)](
+
+ _mla_splitk_sparseattn[(bs, ns)](
mid_o, mid_lse,
- q, kv_flat, kv_indptr, kv_scale,
+ q, kv_flat, kv_indptr,
+ kv_scale,
BLOCK_M=nh, TILE_SIZE=tile,
- NOPE_DIM=NOPE_DIM, ROPE_DIM=ROPE_DIM,
- V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
+ NOPE_SCORE_DIM=ndim, ROPE_DIM=ROPE_DIM,
+ SPARSE_V=ndim, KV_STRIDE=QK_HEAD_DIM,
+ ROPE_START=ROPE_START,
NUM_SEGMENTS=ns, SM_SCALE=SM_SCALE,
num_warps=warps, num_stages=stages,
- allow_flush_denorm=True)
- _reduce[(bs, nh)](
+ allow_flush_denorm=True,
+ )
+
+ _reduce_sparse[(bs, nh)](
o, mid_o, mid_lse,
- BLOCK_M=nh, V_DIM=vd, NUM_SEGMENTS=ns,
+ BLOCK_M=nh, SPARSE_V_DIM=ndim, FULL_V=vd,
+ NUM_SEGMENTS=ns,
num_warps=4, num_stages=1,
- allow_flush_denorm=True)
+ allow_flush_denorm=True,
+ )
return o
+
@triton.jit
- def _mla_mk1g_fused(
- output_ptr,
- q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
- BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
- NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
- V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
- SM_SCALE: tl.constexpr, FP8_MAX: tl.constexpr,
+ def _reduce(
+ output_ptr, segm_output_ptr, segm_lse_ptr,
+ BLOCK_M: tl.constexpr, V_DIM: tl.constexpr, NUM_SEGMENTS: tl.constexpr,
):
batch_idx = tl.program_id(0)
+ head_idx = tl.program_id(1)
+ offs_v = tl.arange(0, V_DIM)
+ seg_offs = tl.arange(0, NUM_SEGMENTS)
+ stat_base = batch_idx * NUM_SEGMENTS * BLOCK_M + head_idx
+ segm_lse = tl.load(segm_lse_ptr + stat_base + seg_offs * BLOCK_M)
+ overall_max = tl.max(segm_lse)
+ weights = tl.exp(segm_lse - overall_max)
+ overall_sum = tl.sum(weights)
+ acc = tl.zeros([V_DIM], dtype=tl.float32)
+ for s in range(NUM_SEGMENTS):
+ o_off = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
+ + s * BLOCK_M * V_DIM + head_idx * V_DIM)
+ seg_out = tl.load(o_off + offs_v)
+ seg_lse = tl.load(segm_lse_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M + s * BLOCK_M + head_idx)
+ acc += seg_out * tl.exp(seg_lse - overall_max)
+ result = tl.where(overall_sum == 0.0, 0.0, acc / overall_sum)
+ out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + head_idx * V_DIM
+ tl.store(out_off + offs_v, result.to(tl.bfloat16))
- kv_start = tl.load(kv_indptr + batch_idx)
- kv_end = tl.load(kv_indptr + batch_idx + 1)
- kv_len = kv_end - kv_start
- kv_s = tl.load(kv_scale_ptr)
- offs_m = tl.arange(0, BLOCK_M)
- offs_nope = tl.arange(0, NOPE_DIM)
- offs_rope = tl.arange(0, ROPE_DIM)
- offs_t = tl.arange(0, TILE_SIZE)
-
- q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
- Q_nope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :])
- Q_rope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :])
-
- amax = tl.maximum(tl.max(tl.abs(Q_nope_bf16)), tl.max(tl.abs(Q_rope_bf16)))
- amax = tl.maximum(amax, 1e-12)
- q_scale = amax / FP8_MAX
- inv_q_scale = FP8_MAX / amax
-
- Q_nope = (Q_nope_bf16 * inv_q_scale).to(tl.float8e4nv)
- Q_rope = (Q_rope_bf16 * inv_q_scale).to(tl.float8e4nv)
- combined_scale = q_scale * kv_s * SM_SCALE
-
- M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
- L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
- acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
-
- num_tiles = tl.cdiv(kv_len, TILE_SIZE)
- for j in range(0, num_tiles):
- seq_offset = j * TILE_SIZE + offs_t
- tile_mask = seq_offset < kv_len
- kv_idx = kv_start + seq_offset
-
- K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
- mask=tile_mask[None, :], other=0.0)
- K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
- mask=tile_mask[None, :], other=0.0)
-
- S = combined_scale * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
- S = tl.where(tile_mask[None, :], S, -float("inf"))
-
- V = tl.trans(K_nope)
-
- m_j = tl.maximum(M_val, tl.max(S, axis=1))
- m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
- P = tl.exp(S - m_j[:, None])
- l_j = tl.sum(P, axis=1)
- alpha = tl.exp(M_val - m_j)
- acc = acc * alpha[:, None]
- L = L * alpha + l_j
- M_val = m_j
-
- acc += tl.dot(P.to(tl.float8e4nv), V)
-
- acc = (acc * kv_s) / L[:, None]
- out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + offs_m[:, None] * V_DIM + tl.arange(0, V_DIM)[None, :]
- tl.store(out_off, acc.to(tl.bfloat16))
-
-
- # ═══════════════════════════════════════════════════════════════════════
- # Champ: mk19 pure fp8 splitk (bs32/8192, bs64/8192)
- # Separate qk_scale/v_scale passed from CPU, no in-kernel Q scaling
- # ═══════════════════════════════════════════════════════════════════════
-
@triton.jit
- def _mla_splitk_mk19(
+ def _mla_splitk_pure_fp8(
segm_output_ptr, segm_lse_ptr,
- q_bf16_ptr, kv_fp8_ptr, kv_indptr,
- kv_scale_ptr,
+ q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
⋯ 1 unchanged lines
):
batch_idx = tl.program_id(0)
segm_idx = tl.program_id(1)
-
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
kv_len = kv_end - kv_start
-
kv_s = tl.load(kv_scale_ptr)
qk_s = kv_s * SM_SCALE
- v_s = kv_s
-
tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
seg_tile_start = segm_idx * tiles_per_segment
seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))
-
offs_m = tl.arange(0, BLOCK_M)
offs_nope = tl.arange(0, NOPE_DIM)
offs_rope = tl.arange(0, ROPE_DIM)
offs_t = tl.arange(0, TILE_SIZE)
-
q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)
Q_rope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :]).to(tl.float8e4nv)
-
M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
-
for j in range(seg_tile_start, seg_tile_end):
seq_offset = j * TILE_SIZE + offs_t
tile_mask = seq_offset < kv_len
kv_idx = kv_start + seq_offset
-
K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
mask=tile_mask[None, :], other=0.0)
K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
mask=tile_mask[None, :], other=0.0)
-
S = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
S = tl.where(tile_mask[None, :], S, -float("inf"))
-
m_j = tl.maximum(M_val, tl.max(S, axis=1))
m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
P = tl.exp(S - m_j[:, None])
⋯ 2 unchanged lines
acc = acc * alpha[:, None]
L = L * alpha + l_j
M_val = m_j
-
V = tl.trans(K_nope)
acc += tl.dot(P.to(tl.float8e4nv), V)
-
- acc = acc * v_s
+ acc = acc * kv_s
o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
+ segm_idx * BLOCK_M * V_DIM)
offs_v = tl.arange(0, V_DIM)
⋯ 2 unchanged lines
tl.store(lse_base + offs_m, M_val + tl.log(L))
- # (ns, tile, warps, stages)
- MK19_SPLITK_CONFIGS = {
- (32, 8192): (8, 64, 4, 2),
- (64, 8192): (4, 64, 8, 2),
- }
- def _run_mk19_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
+
+ def _run_fp8_splitk(q, kv_fp8, kv_scale, kv_indptr, config):
bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
kvsl = config["kv_seq_len"]
- ns, tile, warps, stages = MK19_SPLITK_CONFIGS[(bs, kvsl)]
+ ns, tile, warps, stages = FP8_SPLITK_CONFIGS[(bs, kvsl)]
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
o = torch.empty((bs, nh, vd), dtype=torch.bfloat16, device=q.device)
mid_o = torch.empty((bs, ns, nh, vd), dtype=torch.float32, device=q.device)
mid_lse = torch.empty((bs, ns, nh), dtype=torch.float32, device=q.device)
- _mla_splitk_mk19[(bs, ns)](
+ _mla_splitk_pure_fp8[(bs, ns)](
mid_o, mid_lse,
- q, kv_flat, kv_indptr,
- kv_scale,
+ q, kv_flat, kv_indptr, kv_scale,
BLOCK_M=nh, TILE_SIZE=tile,
NOPE_DIM=NOPE_DIM, ROPE_DIM=ROPE_DIM,
V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
⋯ 8 unchanged lines
return o
- # ═══════════════════════════════════════════════════════════════════════
- # Champ: hip_v16_mk4j split-K fp8pv (bs256/8192)
- # In-kernel Q quantization + split-K=2 for 2 waves/SIMD latency hiding
- # ═══════════════════════════════════════════════════════════════════════
-
@triton.jit
def _mla_splitk_fp8pv(
segm_output_ptr, segm_lse_ptr,
⋯ 8 unchanged lines
):
batch_idx = tl.program_id(0)
segm_idx = tl.program_id(1)
-
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
kv_len = kv_end - kv_start
-
kv_s = tl.load(kv_scale_ptr)
-
tiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)
seg_tile_start = segm_idx * tiles_per_segment
seg_tile_end = tl.minimum((segm_idx + 1) * tiles_per_segment, tl.cdiv(kv_len, TILE_SIZE))
-
offs_m = tl.arange(0, BLOCK_M)
offs_nope = tl.arange(0, NOPE_DIM)
offs_rope = tl.arange(0, ROPE_DIM)
offs_t = tl.arange(0, TILE_SIZE)
-
q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
Q_nope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :])
Q_rope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :])
-
amax_nope = tl.max(tl.abs(Q_nope_bf16))
amax_rope = tl.max(tl.abs(Q_rope_bf16))
amax = tl.maximum(amax_nope, amax_rope)
amax = tl.maximum(amax, 1e-12)
q_scale = amax / FP8_MAX
inv_q_scale = FP8_MAX / amax
-
Q_nope = (Q_nope_bf16 * inv_q_scale).to(tl.float8e4nv)
Q_rope = (Q_rope_bf16 * inv_q_scale).to(tl.float8e4nv)
-
combined_scale = q_scale * kv_s * SM_SCALE
-
M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
-
for j in range(seg_tile_start, seg_tile_end):
seq_offset = j * TILE_SIZE + offs_t
tile_mask = seq_offset < kv_len
kv_idx = kv_start + seq_offset
-
K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
mask=tile_mask[None, :], other=0.0)
K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
mask=tile_mask[None, :], other=0.0)
-
S = combined_scale * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
S = tl.where(tile_mask[None, :], S, -float("inf"))
-
m_j = tl.maximum(M_val, tl.max(S, axis=1))
m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
P = tl.exp(S - m_j[:, None])
⋯ 2 unchanged lines
acc = acc * alpha[:, None]
L = L * alpha + l_j
M_val = m_j
-
V = tl.trans(K_nope)
acc += tl.dot(P.to(tl.float8e4nv), V)
-
acc = acc * kv_s
o_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM
+ segm_idx * BLOCK_M * V_DIM)
⋯ 29 unchanged lines
return o
+ @triton.jit
+ def _mla_mk1g_fused(
+ output_ptr,
+ q_bf16_ptr, kv_fp8_ptr, kv_indptr, kv_scale_ptr,
+ BLOCK_M: tl.constexpr, TILE_SIZE: tl.constexpr,
+ NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr,
+ V_DIM: tl.constexpr, KV_STRIDE: tl.constexpr,
+ SM_SCALE: tl.constexpr, FP8_MAX: tl.constexpr,
+ ):
+ batch_idx = tl.program_id(0)
+ kv_start = tl.load(kv_indptr + batch_idx)
+ kv_end = tl.load(kv_indptr + batch_idx + 1)
+ kv_len = kv_end - kv_start
+ kv_s = tl.load(kv_scale_ptr)
+ offs_m = tl.arange(0, BLOCK_M)
+ offs_nope = tl.arange(0, NOPE_DIM)
+ offs_rope = tl.arange(0, ROPE_DIM)
+ offs_t = tl.arange(0, TILE_SIZE)
+ q_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDE
+ Q_nope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :])
+ Q_rope_bf16 = tl.load(q_base + offs_m[:, None] * KV_STRIDE + (NOPE_DIM + offs_rope)[None, :])
+ amax = tl.maximum(tl.max(tl.abs(Q_nope_bf16)), tl.max(tl.abs(Q_rope_bf16)))
+ amax = tl.maximum(amax, 1e-12)
+ q_scale = amax / FP8_MAX
+ inv_q_scale = FP8_MAX / amax
+ Q_nope = (Q_nope_bf16 * inv_q_scale).to(tl.float8e4nv)
+ Q_rope = (Q_rope_bf16 * inv_q_scale).to(tl.float8e4nv)
+ combined_scale = q_scale * kv_s * SM_SCALE
+ M_val = tl.full([BLOCK_M], -float("inf"), dtype=tl.float32)
+ L = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
+ acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
+ num_tiles = tl.cdiv(kv_len, TILE_SIZE)
+ for j in range(0, num_tiles):
+ seq_offset = j * TILE_SIZE + offs_t
+ tile_mask = seq_offset < kv_len
+ kv_idx = kv_start + seq_offset
+ K_nope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + offs_nope[:, None],
+ mask=tile_mask[None, :], other=0.0)
+ K_rope = tl.load(kv_fp8_ptr + kv_idx[None, :] * KV_STRIDE + (NOPE_DIM + offs_rope)[:, None],
+ mask=tile_mask[None, :], other=0.0)
+ S = combined_scale * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))
+ S = tl.where(tile_mask[None, :], S, -float("inf"))
+ V = tl.trans(K_nope)
+ m_j = tl.maximum(M_val, tl.max(S, axis=1))
+ m_j = tl.where(m_j > -float("inf"), m_j, 0.0)
+ P = tl.exp(S - m_j[:, None])
+ l_j = tl.sum(P, axis=1)
+ alpha = tl.exp(M_val - m_j)
+ acc = acc * alpha[:, None]
+ L = L * alpha + l_j
+ M_val = m_j
+ acc += tl.dot(P.to(tl.float8e4nv), V)
+ acc = (acc * kv_s) / L[:, None]
+ out_off = output_ptr + batch_idx * BLOCK_M * V_DIM + offs_m[:, None] * V_DIM + tl.arange(0, V_DIM)[None, :]
+ tl.store(out_off, acc.to(tl.bfloat16))
+
+
def _run_mk1g_fused(q, kv_fp8, kv_scale, kv_indptr, config):
bs, nh, vd = config["batch_size"], config["num_heads"], config["v_head_dim"]
kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
⋯ 5 unchanged lines
V_DIM=vd, KV_STRIDE=QK_HEAD_DIM,
SM_SCALE=SM_SCALE, FP8_MAX=FP8_MAX,
num_warps=8, num_stages=2,
- allow_flush_denorm=True,
+ allow_flush_denorm=True,
matrix_instr_nonkdim=16,)
return o
-
def quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
⋯ 39 unchanged lines
intra_batch_mode=True, **meta)
return o
+
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kvsl = config["kv_seq_len"]
key = (bs, kvsl)
- # 4x1024, 4x8192, 32x1024, 64x1024: fp8 splitk
+ # sparse attention approx using rope portions and slice of nope parts
+ if key in SPARSE_CONFIGS:
+ kv_fp8, kv_scale = kv_data["fp8"]
+ return _run_sparse_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
+
+ # 256x8192: mk4j split-K=2 sparse approx tuned
+ if key == (256, 8192):
+ kv_fp8, kv_scale = kv_data["fp8"]
+ return _run_mk4j_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
+
+ # 4x1024, 4x8192, 32x1024, 64x1024: regular ol' fp8 splitk
if key in FP8_SPLITK_CONFIGS:
kv_fp8, kv_scale = kv_data["fp8"]
return _run_fp8_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
- # 256x1024: mk1g fused (in-kernel Q quant, tile=64, warps=8)
+ # 256x1024: mk1g fused
if key == (256, 1024):
kv_fp8, kv_scale = kv_data["fp8"]
return _run_mk1g_fused(q, kv_fp8, kv_scale, kv_indptr, config)
- # 32x8192: mk19f pure fp8 splitk (champ, 74.4µs vs 94.7µs aiter)
- # 64x8192: mk1a pure fp8 splitk (champ, 124.3µs vs 139.1µs aiter)
- if key in MK19_SPLITK_CONFIGS:
- kv_fp8, kv_scale = kv_data["fp8"]
- return _run_mk19_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
- # 256x8192: hip_v16_mk4j split-K=2 fp8pv (champ, 319.4µs vs 320.5µs aiter)
- if key == (256, 8192):
- kv_fp8, kv_scale = kv_data["fp8"]
- return _run_mk4j_splitk(q, kv_fp8, kv_scale, kv_indptr, config)
return reference_fallback(data)
scrolls · 645 diff lines total

Best evidence level for this revision: reported

JSON