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
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.
fp8
Q_nope = tl.load(q_base + offs_m[:, None] * KV_STRIDE + offs_nope[None, :]).to(tl.float8e4nv)mma
S = qk_s * (tl.dot(Q_nope, K_nope) + tl.dot(Q_rope, K_rope))num-warps = 4
num_warps=4, num_stages=1,split-k
FP8_SPLITK_CONFIGS = {stages = 1
num_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 torchimport tritonimport triton.language as tl⋯ 14 unchanged linesNUM_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_idxsegm_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 lineskv_s = tl.load(kv_scale_ptr)qk_s = kv_s * SM_SCALE+ v_s = kv_stiles_per_segment = tl.cdiv(kv_len, NUM_SEGMENTS * TILE_SIZE)seg_tile_start = segm_idx * tiles_per_segmentseg_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 chunkq_base = q_bf16_ptr + batch_idx * BLOCK_M * KV_STRIDEQ_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 linesK_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"))+ # softmaxm_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_jM_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_Mtl.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_segmentseg_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_STRIDEQ_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_ttile_mask = seq_offset < kv_lenkv_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 linesacc = acc * alpha[:, None]L = L * alpha + l_jM_val = m_j-V = tl.trans(K_nope)acc += tl.dot(P.to(tl.float8e4nv), V)-- acc = acc * v_s+ acc = acc * kv_so_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 linestl.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 linesreturn o- # ═══════════════════════════════════════════════════════════════════════- # Champ: hip_v16_mk4j split-K fp8pv (bs256/8192)- # In-kernel Q quantization + split-K=2 for 2 waves/SIMD latency hiding- # ═══════════════════════════════════════════════════════════════════════-@triton.jitdef _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_segmentseg_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_STRIDEQ_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_MAXinv_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_ttile_mask = seq_offset < kv_lenkv_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 linesacc = acc * alpha[:, None]L = L * alpha + l_jM_val = m_j-V = tl.trans(K_nope)acc += tl.dot(P.to(tl.float8e4nv), V)-acc = acc * kv_so_base = (segm_output_ptr + batch_idx * NUM_SEGMENTS * BLOCK_M * V_DIM+ segm_idx * BLOCK_M * V_DIM)⋯ 29 unchanged linesreturn 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 linesV_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 linesintra_batch_mode=True, **meta)return o+def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = databs = 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 splitkif 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 fusedif 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