Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 56 additions & 28 deletions aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,14 +153,14 @@ def _cache_load(


@gluon.jit
def _fp8_to_f32(x_u8):
return x_u8.to(gl.float8e4nv, bitcast=True).to(gl.float32)
def _fp8_to_f32(x_u8, FP8_TY: gl.constexpr):
return x_u8.to(FP8_TY, bitcast=True).to(gl.float32)


@gluon.jit
def _fp8_to_bf16(x_u8):
def _fp8_to_bf16(x_u8, FP8_TY: gl.constexpr):
# Exact: fp8's 3 mantissa bits fit bf16's 8.
return x_u8.to(gl.float8e4nv, bitcast=True).to(gl.bfloat16)
return x_u8.to(FP8_TY, bitcast=True).to(gl.bfloat16)


@gluon.jit
Expand Down Expand Up @@ -328,6 +328,14 @@ class Cfg:
IDX_BUFFER_LOAD: gl.constexpr
FP8_MFMA: gl.constexpr # "fp8_scalar" only: feed the matrix core the cache's
# own fp8 instead of dequantizing to bf16
# fp8 encoding of the cache, q and the dot operands: e4m3fnuz (gfx942's native,
# which its fp8 MFMA reads) or OCP e4m3 (gfx950's)
FP8_FNUZ: gl.constexpr
FP8_TY: gl.constexpr
FP8_MAX: gl.constexpr
# fp8 PV dots quantize p * P_SCALE, so small probabilities do not flush to
# zero; the epilogue divides it back out with the V-side scale.
P_SCALE: gl.constexpr
# Cache policy per load site
GATHER_CACHE: gl.constexpr
IDX_CACHE: gl.constexpr
Expand Down Expand Up @@ -386,6 +394,7 @@ def __init__(
DSV4_WALK=False,
SCL_DWORD=False,
STAGED_K32=False,
FP8_FNUZ=False,
):
self.BLOCK_M = gl.constexpr(BLOCK_M)
self.BLOCK_K = gl.constexpr(BLOCK_K)
Expand All @@ -407,6 +416,11 @@ def __init__(
self.HEAD_ALIGNED = gl.constexpr(HEAD_ALIGNED)
self.IDX_BUFFER_LOAD = gl.constexpr(IDX_BUFFER_LOAD)
self.FP8_MFMA = gl.constexpr(FP8_MFMA)
self.FP8_FNUZ = gl.constexpr(FP8_FNUZ)
self.FP8_TY = gl.constexpr(gl.float8e4b8 if FP8_FNUZ else gl.float8e4nv)
self.FP8_MAX = gl.constexpr(240.0 if FP8_FNUZ else 448.0)
# p <= 1, so 128 keeps p * P_SCALE under fnuz's 240 and is exact.
self.P_SCALE = gl.constexpr(128.0 if FP8_FNUZ else 1.0)
self.GATHER_CACHE = gl.constexpr(GATHER_CACHE)
self.IDX_CACHE = gl.constexpr(IDX_CACHE)
self.ASYNC_LDS = gl.constexpr(ASYNC_LDS)
Expand All @@ -418,21 +432,23 @@ def __init__(
# bf16 dots on 16x16x32 (gfx950's full rate) on the dsv4 and STAGED_K32 walks.
MFMA_K = 32 if FP8_MFMA or DSV4_WALK or STAGED_K32 else 16
self.MFMA_K = gl.constexpr(MFMA_K)
# fnuz fp8 operands only have CDNA3 (version 3) matrix-core intrinsics.
MFMA_VERSION = 3 if FP8_MFMA and FP8_FNUZ else 4

# Warps tile N; past 16 heads they also tile M, 16 heads per warp.
M_WARPS = max(1, min(BLOCK_M // 16, NUM_WARPS))
self.N_WARPS = gl.constexpr(NUM_WARPS // M_WARPS)
self.qk_layout = gl.constexpr(
gl.amd.AMDMFMALayout(
version=4,
version=MFMA_VERSION,
instr_shape=[16, 16, MFMA_K],
transposed=True,
warps_per_cta=[M_WARPS, NUM_WARPS // M_WARPS],
)
)
self.pv_layout = gl.constexpr(
gl.amd.AMDMFMALayout(
version=4,
version=MFMA_VERSION,
instr_shape=[16, 16, MFMA_K],
transposed=True,
warps_per_cta=[M_WARPS, NUM_WARPS // M_WARPS],
Expand Down Expand Up @@ -735,13 +751,13 @@ def _deq_store(x_u8, sc, kv_smem, off, cfg, fmt, AXIS: gl.constexpr):
x_u8.to(gl.float8e4nv, bitcast=True), sc, gl.bfloat16
)
elif fmt.KIND == "fp8_scalar":
val = _fp8_to_bf16(x_u8)
val = _fp8_to_bf16(x_u8, cfg.FP8_TY)
else:
gl.static_assert(
fmt.KIND == "fp8_g64" or fmt.KIND == "fp8_dsv32_mla",
"fp8_dsv4_mla dequantizes with DEQ upcast or asm",
)
val = (_fp8_to_f32(x_u8) * sc).to(gl.bfloat16)
val = (_fp8_to_f32(x_u8, cfg.FP8_TY) * sc).to(gl.bfloat16)
if AXIS == 1:
kv_smem.slice(off, x_u8.shape[1], dim=1).store(val)
else:
Expand Down Expand Up @@ -923,7 +939,7 @@ def _qk_scores(cfg, q_dot, q_rope_dot, kv_smem, rope_smem):
else:
k = kv_smem.permute([1, 0]).load(cfg.k_layout) # [KV_DIM, BLOCK_K]
if cfg.ASYNC_LDS:
k = k.to(gl.float8e4nv, bitcast=True) # raw cache bytes; layout-preserving
k = k.to(cfg.FP8_TY, bitcast=True) # raw cache bytes; layout-preserving
S = gl.amd.cdna4.mfma(q_dot, k, S)
if cfg.ROPE_SEPARATE:
if cfg.ASYNC_LDS and cfg.RELAXED_LOAD:
Expand All @@ -933,7 +949,7 @@ def _qk_scores(cfg, q_dot, q_rope_dot, kv_smem, rope_smem):
else:
k_rope = rope_smem.permute([1, 0]).load(cfg.k_layout)
if cfg.ASYNC_LDS:
k_rope = k_rope.to(gl.float8e4nv, bitcast=True)
k_rope = k_rope.to(cfg.FP8_TY, bitcast=True)
S = gl.amd.cdna4.mfma(q_rope_dot, k_rope, S)
return S

Expand Down Expand Up @@ -1197,17 +1213,17 @@ def _stage(cfg, seg, x_u8, sc, k_rope, kv_smem, rope_smem):
# No dequant: the scale is folded outside the loop (qk_scale on the K
# side, the accumulator on the V side), so what lands in LDS is exactly
# what the gather returned.
kv_smem.store(x_u8.to(gl.float8e4nv, bitcast=True))
kv_smem.store(x_u8.to(cfg.FP8_TY, bitcast=True))
if cfg.ROPE_SEPARATE:
rope_smem.store(k_rope.to(gl.float8e4nv, bitcast=True))
rope_smem.store(k_rope.to(cfg.FP8_TY, bitcast=True))
else:
_deq_store_tile(x_u8, sc, kv_smem, cfg, fmt)
if fmt.KIND == "fp8_dsv4_mla":
kv_smem.slice(fmt.NOPE_DIM, cfg.ROPE_DIM, dim=1).store(k_rope)
elif fmt.KIND == "fp8_dsv32_mla":
rope_smem.store(k_rope)
elif fmt.KIND == "fp8_scalar" and cfg.ROPE_SEPARATE:
rope_smem.store(_fp8_to_bf16(k_rope))
rope_smem.store(_fp8_to_bf16(k_rope, cfg.FP8_TY))
# "fp8_g64": the whole head is one fp8 tile; nothing else to store.


Expand Down Expand Up @@ -1332,13 +1348,13 @@ def _qkpv_lds(
else:
v = kv_smem.load(cfg.v_layout)
if cfg.ASYNC_LDS:
v = v.to(gl.float8e4nv, bitcast=True)
v = v.to(cfg.FP8_TY, bitcast=True)
# "fp8_scalar": V was staged as raw fp8 code points; apply the per-tensor scale
# on the small side (p) and leave l scale-free: out = sum(p*s*V)/l exactly.
if seg.fmt.KIND == "fp8_scalar" and not cfg.FP8_MFMA:
p = p * v_scale
if cfg.FP8_MFMA:
p_dot = gl.convert_layout(p.to(gl.float8e4nv), cfg.p_layout)
p_dot = gl.convert_layout((p * cfg.P_SCALE).to(cfg.FP8_TY), cfg.p_layout)
else:
p_dot = gl.convert_layout(p.to(gl.bfloat16), cfg.p_layout)
alpha_pv = gl.convert_layout(alpha, gl.SliceLayout(1, cfg.pv_layout))
Expand Down Expand Up @@ -2147,7 +2163,16 @@ def _xcd_work(GRID_ORDER: gl.constexpr, NUM_XCDS: gl.constexpr, SPLIT_K: gl.cons

_sparse_mla_repr = make_kernel_repr(
"_sparse_mla",
["BLOCK_M", "BLOCK_K", "HEAD_SIZE", "SPLIT_K", "MAIN_FMT", "ROPE_SEPARATE"],
[
"BLOCK_M",
"BLOCK_K",
"HEAD_SIZE",
"SPLIT_K",
"MAIN_FMT",
"ROPE_SEPARATE",
"FP8_MFMA",
"FP8_FNUZ",
],
)


Expand Down Expand Up @@ -2239,6 +2264,9 @@ def _sparse_mla(
IDX_BUFFER_LOAD: gl.constexpr,
HAS_INVALID: gl.constexpr,
FP8_MFMA: gl.constexpr = False,
# fp8 is e4m3fnuz (gfx942) rather than OCP e4m3: the cache, the q this kernel
# quantizes, and the fp8 dot operands.
FP8_FNUZ: gl.constexpr = False,
Comment thread
jin-amd marked this conversation as resolved.
# q already quantized to e4m3 by the caller, plus the scalar f32 scale it
# was quantized with. This is the calling convention aiter's asm
# mla_decode_fwd uses, where vLLM passes layer._q_scale.
Expand Down Expand Up @@ -2329,8 +2357,8 @@ def _sparse_mla(
"inconsistent",
)
# The fp8 path needs one positive scalar scale per cache, since that is what
# folds outside the loop, and OCP e4m3 code points, which is what the matrix
# core reads.
# folds outside the loop, and the arch's own fp8 code points (FP8_TY), which
# is what the matrix core reads.
gl.static_assert(
(not FP8_MFMA)
or (MAIN_FMT == "fp8_scalar" and (not HAS_EXTRA or EXTRA_FMT == "fp8_scalar")),
Expand Down Expand Up @@ -2397,6 +2425,7 @@ def _sparse_mla(
DSV4_WALK,
CS0_ALIGN >= 4,
STAGED_K32,
FP8_FNUZ,
)
main_fmt = Fmt(
cfg,
Expand Down Expand Up @@ -2466,7 +2495,6 @@ def _sparse_mla(
if FP8_MFMA and not Q_FP8:
# bf16 q: quantize here, one e4m3 scale for this program's whole Q tile
# (nope and rope), so the fold below is one extra factor on qk_scale.
E4M3_MAX: gl.constexpr = 448.0
q_amax = gl.max(gl.max(gl.abs(q).to(gl.float32), axis=1), axis=0)
if cfg.Q_LDS:
q_dot = gl.allocate_shared_memory(
Expand Down Expand Up @@ -2509,17 +2537,17 @@ def _sparse_mla(
q_amax, gl.max(gl.max(gl.abs(q_rope).to(gl.float32), axis=1), axis=0)
)
q_amax = gl.maximum(q_amax, 1e-30)
q_rcp = E4M3_MAX / q_amax
q_rcp = cfg.FP8_MAX / q_amax
q_dot = gl.convert_layout(
(q.to(gl.float32) * q_rcp).to(gl.float8e4nv), cfg.q_layout
(q.to(gl.float32) * q_rcp).to(cfg.FP8_TY), cfg.q_layout
)
if ROPE_SEPARATE:
q_rope_dot = gl.convert_layout(
(q_rope.to(gl.float32) * q_rcp).to(gl.float8e4nv), cfg.q_layout
(q_rope.to(gl.float32) * q_rcp).to(cfg.FP8_TY), cfg.q_layout
)
else:
q_rope_dot = q_dot
q_scale = q_amax / E4M3_MAX
q_scale = q_amax / cfg.FP8_MAX
main_qk_scale = main_qk_scale * q_scale
extra_qk_scale = extra_qk_scale * q_scale

Expand All @@ -2536,7 +2564,7 @@ def _sparse_mla(
acc = gl.zeros([BLOCK_M, HEAD_SIZE], gl.float32, layout=cfg.pv_layout)

# An fp8 buffer is half the bytes of the bf16 staging it replaces.
SMEM_DT: gl.constexpr = gl.float8e4nv if FP8_MFMA else gl.bfloat16
SMEM_DT: gl.constexpr = cfg.FP8_TY if FP8_MFMA else gl.bfloat16
# The LDS-DMA converts nothing, so an async buffer's element type has to be the
# cache's own u8; the dot operands bitcast on read (free, layout-preserving).
BUF_DT: gl.constexpr = gl.uint8 if ASYNC_LDS else SMEM_DT
Expand Down Expand Up @@ -2672,10 +2700,10 @@ def _sparse_mla(
)

if FP8_MFMA:
# The fp8 PV dot ran on raw code points, so the V-side scale comes off
# here, once per program instead of once per tile. l is untouched, so
# out = acc*s/l is what the bf16 path computes.
acc = acc * (extra_v_scale if HAS_EXTRA else main_v_scale)
# The fp8 PV dot ran on raw code points and p * P_SCALE, so the V-side
# scale and P_SCALE come off here, once per program instead of once per
# tile. l is untouched, so out = acc*s/l is what the bf16 path computes.
acc = acc * ((extra_v_scale if HAS_EXTRA else main_v_scale) / cfg.P_SCALE)

# Move the row reductions into pv-slice space for output/partials.
m_pv = gl.convert_layout(m_i, gl.SliceLayout(1, cfg.pv_layout))
Expand Down
Loading
Loading