submission 755183
Will Fisher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 906 lines, June 9 Researcher Reciprocity License v1.0.
attempt_v103.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755183?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:e5f28330db81c8f0105f24a06d64e48f87d98dc42039091c576297f1c76524ba
license declaredunknown
license concludedunknown
authorsWill Fisher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_buffer, kv_scales = kv_data["mxfp4"]num-warps = 8
constexpr int NUM_WARPS = 8;online-softmax
float running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];shared-memory
extern __shared__ char smem[];vector-width = float2
float2 f = __half22float2(v_acc[i]);Kernel source
attempt_v103.py906 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import os
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/amd_detail/amd_hip_fp16.h>
#include <pybind11/pybind11.h>
constexpr int NUM_HEADS = 16;
constexpr int QK_HEAD_DIM = 576;
constexpr int V_HEAD_DIM = 512;
constexpr int KV_PACKED = 288;
constexpr int SCALE_COLS = 18;
constexpr int WAVEFRONT_SIZE = 64;
constexpr int NUM_WARPS = 8;
constexpr int BLOCK_THREADS = NUM_WARPS * WAVEFRONT_SIZE;
constexpr int HEADS_PER_LANE = 4;
constexpr float LOG2E = 1.4426950408889634f;
constexpr float SM_SCALE = 1.0f / 24.0f; // 1/sqrt(576), hardcoded for DeepSeek R1
constexpr float SM_LOG2E = SM_SCALE * LOG2E; // pre-multiplied, compile-time constant
constexpr int FP4_K = 128; // dims per K chunk
constexpr int FP4_K_CHUNKS = 5; // ceil(576/128)
constexpr int V_K_CHUNKS = 4; // 512/128 — chunks that contain V data
constexpr int TILE_TOKENS = 128; // tokens per tile
// LDS layout: KV double-buf | V_fp8/Q reuse | score | attn_fp8 | stats
// KV loaded via coalesced DMA → LDS, then each thread reads to registers
constexpr int BUF_KV_BYTES = TILE_TOKENS * KV_PACKED; // 36864
constexpr int BUF_SC_BYTES = TILE_TOKENS * SCALE_COLS; // 2304
constexpr int BUF_TOTAL = BUF_KV_BYTES + BUF_SC_BYTES; // 39168
constexpr int V_FP8_STRIDE = V_HEAD_DIM + 8; // 520: 8-byte aligned, breaks bank conflicts (520/4%64=2)
constexpr int V_FP8_BYTES = TILE_TOKENS * V_FP8_STRIDE; // 66560
constexpr int SCORE_STRIDE = NUM_HEADS + 2; // 18 bf16/row: pad for bank conflicts (gcd(9,64)=1)
constexpr int SCORE_BYTES = TILE_TOKENS * SCORE_STRIDE * 2; // 4608
constexpr int ATTN_FP8_BYTES = TILE_TOKENS * NUM_HEADS + 136; // 2184 (fp8 attn + non-linear bank shift)
constexpr int STATS_BYTES = NUM_HEADS * 2 * 4; // 128
// Q reuses the V_fp8 area (Q written in prologue, cached to regs, then V_fp8 overwrites)
constexpr int OFF_KV0 = 0;
constexpr int OFF_KV1 = BUF_TOTAL; // 39168
constexpr int OFF_VFP8 = BUF_TOTAL * 2; // 78336 (also OFF_Q during prologue)
constexpr int OFF_Q = OFF_VFP8; // reuses V_fp8 area
constexpr int OFF_QSC = OFF_Q + NUM_HEADS * KV_PACKED; // 82944
constexpr int OFF_SCORE = OFF_VFP8 + V_FP8_BYTES; // 143872
constexpr int OFF_ATTN = OFF_SCORE + SCORE_BYTES; // 147968
constexpr int OFF_STATS = OFF_ATTN + ATTN_FP8_BYTES; // 150016
constexpr int TOTAL_LDS = OFF_STATS + STATS_BYTES; // 150144 (~147KB)
typedef __bf16 bf16x2 __attribute__((__vector_size__(2 * sizeof(__bf16))));
typedef float f32x4 __attribute__((__vector_size__(4 * sizeof(float))));
typedef int v4i32 __attribute__((__vector_size__(4 * sizeof(int))));
typedef int v2i32 __attribute__((__vector_size__(2 * sizeof(int))));
typedef int v8i32 __attribute__((__vector_size__(8 * sizeof(int))));
typedef uint32_t u32x4_vec __attribute__((__vector_size__(4 * sizeof(uint32_t))));
typedef __bf16 bf16x4 __attribute__((__vector_size__(4 * sizeof(__bf16))));
__device__ __forceinline__ v8i32 to_v8(v4i32 x) {
v8i32 r = {x[0], x[1], x[2], x[3], 0, 0, 0, 0};
return r;
}
constexpr int GROUP_SZ = 32;
__device__ __forceinline__
uint8_t compute_e8m0_scale(const __hip_bfloat16* vals) {
uint16_t mx_bits = 0;
#pragma unroll
for (int i = 0; i < GROUP_SZ; i++) {
uint16_t b = *reinterpret_cast<const uint16_t*>(&vals[i]) & 0x7FFF;
mx_bits = max(mx_bits, b);
}
if (mx_bits == 0) return 0;
uint32_t bits = (uint32_t)mx_bits << 16;
bits = (bits + 0x200000u) & 0xFF800000u;
int e8m0 = (int)(bits >> 23) - 2;
return (uint8_t)max(0, min(254, e8m0));
}
__device__ __forceinline__
void quantize_group(const __hip_bfloat16* vals, v4i32& out, uint32_t& scale_out) {
uint8_t scale = compute_e8m0_scale(vals);
scale_out = scale;
float scale_f = __uint_as_float((uint32_t)scale << 23);
#define PACK_WORD(w) do { \
unsigned int packed = 0; \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+0], vals[(w)*8+1]}, scale_f, 0); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+2], vals[(w)*8+3]}, scale_f, 1); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+4], vals[(w)*8+5]}, scale_f, 2); \
packed = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
packed, __hip_bfloat162{vals[(w)*8+6], vals[(w)*8+7]}, scale_f, 3); \
out[w] = (int)packed; \
} while(0)
PACK_WORD(0);
PACK_WORD(1);
PACK_WORD(2);
PACK_WORD(3);
#undef PACK_WORD
}
// Fused DPP mirror butterfly: single-instruction v_max/v_add with DPP modifier
#define DPP_MAX(m, ctrl) asm volatile( \
"v_max_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(m))
#define DPP_ADD(s, ctrl) asm volatile( \
"v_add_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(s))
__device__ __forceinline__ float dpp_row_max(float m) {
DPP_MAX(m, "quad_perm:[1,0,3,2]");
DPP_MAX(m, "quad_perm:[2,3,0,1]");
DPP_MAX(m, "row_half_mirror");
DPP_MAX(m, "row_mirror");
return m;
}
__device__ __forceinline__ float dpp_row_sum(float s) {
DPP_ADD(s, "quad_perm:[1,0,3,2]");
DPP_ADD(s, "quad_perm:[2,3,0,1]");
DPP_ADD(s, "row_half_mirror");
DPP_ADD(s, "row_mirror");
return s;
}
using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using i32x2 = int32_t __attribute__((ext_vector_type(2)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;
using as3_bf16x4_ptr = bf16x4 __attribute__((address_space(3)))*;
using as3_i32x2_ptr = i32x2 __attribute__((address_space(3)))*;
using as3_v2i32_ptr = v2i32 __attribute__((address_space(3)))*;
typedef __bf16 bf16x8 __attribute__((__vector_size__(8 * sizeof(__bf16))));
typedef short ev_short2 __attribute__((ext_vector_type(2)));
extern "C" __device__ __attribute__((const)) __bf16 llvm_exp2_bf16(__bf16) __asm("llvm.exp2.bf16");
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset, int offset, int aux)
__asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const i32x4*>(&rsrc);
}
// ============================================================
// Grouped kernel: double-buffered KV LDS, per-warp V dequant
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
__global__ void __launch_bounds__(BLOCK_THREADS)
mla_grouped(
const __hip_bfloat16* __restrict__ q,
const uint8_t* __restrict__ kv_data,
const uint8_t* __restrict__ kv_scales,
const int32_t* __restrict__ qo_indptr,
const int32_t* __restrict__ kv_indptr,
__half* __restrict__ partial_out,
float* __restrict__ partial_max,
float* __restrict__ partial_sum,
__hip_bfloat16* __restrict__ final_output,
const float* __restrict__ kv_scale_ptr
) {
const int batch_idx = blockIdx.x / NUM_PARTIALS;
const int partial_idx = blockIdx.x - batch_idx * NUM_PARTIALS;
const int tile_base = partial_idx * TILES_PER_BLOCK;
const int warp_id = threadIdx.x >> 6;
const int lane_id = threadIdx.x & 63;
const int lane_mod16 = lane_id & 15;
// Read fp8 kv_scale once, round UP to next power of 2 to avoid NaN
// Hardware cvt instructions drop mantissa (E8M0), rounding scale DOWN.
// Smaller divisor → larger quotient → overflow. Fix: round UP.
float kv_scale_raw = *kv_scale_ptr;
uint32_t ks_bits = __float_as_uint(kv_scale_raw);
uint32_t ks_exp = (ks_bits >> 23) & 0xFF;
uint32_t ks_mant = ks_bits & 0x7FFFFF;
int safe_exp = (int)ks_exp + (ks_mant != 0 ? 1 : 0);
float safe_scale = __uint_as_float((uint32_t)safe_exp << 23);
int mfma_v_scale = safe_exp;
const int k_group = lane_id >> 4;
const int tok = warp_id * 16 + lane_mod16;
const int q_start = qo_indptr[batch_idx];
const int kv_start = kv_indptr[batch_idx];
const int kv_end = kv_indptr[batch_idx + 1];
extern __shared__ char smem[];
uint8_t* q_lds = reinterpret_cast<uint8_t*>(smem + OFF_Q);
uint8_t* qsc_lds = reinterpret_cast<uint8_t*>(smem + OFF_QSC);
uint8_t* v_fp8_lds = reinterpret_cast<uint8_t*>(smem + OFF_VFP8);
uint8_t* attn_fp8_lds = reinterpret_cast<uint8_t*>(smem + OFF_ATTN);
float* stats_lds = reinterpret_cast<float*>(smem + OFF_STATS);
// Buffer resource descriptors hoisted — use soffset for per-tile offset
const int kv_len = kv_end - kv_start;
i32x4 kv_srsrc = make_srsrc(kv_data + kv_start * KV_PACKED, kv_len * KV_PACKED);
i32x4 sc_srsrc = make_srsrc(kv_scales + kv_start * SCALE_COLS, kv_len * SCALE_COLS);
auto load_tile = [&](int tile_tok_offset, int buf) {
uint8_t* dst_kv = reinterpret_cast<uint8_t*>(smem + (buf == 0 ? OFF_KV0 : OFF_KV1));
uint8_t* dst_sc = dst_kv + BUF_KV_BYTES;
const int kv_soff = tile_tok_offset * KV_PACKED;
const int sc_soff = tile_tok_offset * SCALE_COLS;
constexpr int KV_LOADS = (BUF_KV_BYTES + 15) / 16;
for (int off = threadIdx.x; off < KV_LOADS; off += BLOCK_THREADS) {
int voff = off * 16;
llvm_amdgcn_raw_buffer_load_lds(kv_srsrc,
reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_kv) + voff),
16, voff, kv_soff, 0, 0);
}
constexpr int SC_LOADS = (BUF_SC_BYTES + 15) / 16;
for (int off = threadIdx.x; off < SC_LOADS; off += BLOCK_THREADS) {
int voff = off * 16;
llvm_amdgcn_raw_buffer_load_lds(sc_srsrc,
reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_sc) + voff),
16, voff, sc_soff, 0, 0);
}
};
// Issue tile 0 DMA (non-blocking, in flight during Q quantize)
load_tile(tile_base * TILE_TOKENS, 0);
// ---- Q quantize ----
{
const __hip_bfloat16* q_base = q + q_start * NUM_HEADS * QK_HEAD_DIM;
const int dc = warp_id;
if (dc < FP4_K_CHUNKS) {
const int qds = dc * FP4_K + (k_group << 5);
if (qds < QK_HEAD_DIM) {
// Vectorized 128-bit loads (4 × 16 bytes = 32 bf16)
__hip_bfloat16 qv[GROUP_SZ];
const __hip_bfloat16* qp = q_base + lane_mod16 * QK_HEAD_DIM + qds;
#pragma unroll
for (int i = 0; i < 4; i++)
*reinterpret_cast<u32x4_vec*>(&qv[i*8]) = *reinterpret_cast<const u32x4_vec*>(qp + i*8);
v4i32 ar;
uint32_t e8;
quantize_group(qv, ar, e8);
const int qb = qds >> 5;
*reinterpret_cast<u32x4_vec*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4)) = *reinterpret_cast<const u32x4_vec*>(&ar);
qsc_lds[lane_mod16 * SCALE_COLS + qb] = (uint8_t)e8;
}
}
}
// Running accumulators
f32x4 running_v[4];
float running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];
#pragma unroll
for (int i = 0; i < 4; i++) running_v[i] = {0,0,0,0};
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) { running_max[h] = -INFINITY; running_sum[h] = 0; }
v4i32 q_v4[FP4_K_CHUNKS];
int q_scale_cached[FP4_K_CHUNKS];
// SM_LOG2E = SM_LOG2E is constexpr (compile-time constant)
using f16x2 = _Float16 __attribute__((ext_vector_type(2)));
// ============ TILE LOOP ============
for (int ti = 0; ti < TILES_PER_BLOCK; ti++) {
const int cur = ti & 1;
const int tile_start = kv_start + (tile_base + ti) * TILE_TOKENS;
const int tile_len = max(0, min(TILE_TOKENS, kv_end - tile_start));
// Wait for KV DMA + ensure all threads' LDS writes visible
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
// Cache Q on first iteration (after sync ensures q_lds visible)
if (ti == 0) {
#pragma unroll
for (int d = 0; d < FP4_K_CHUNKS; d++) {
const int qb = d * (FP4_K / 32) + k_group;
v4i32 qr = *reinterpret_cast<const v4i32*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4));
q_scale_cached[d] = (int)qsc_lds[lane_mod16 * SCALE_COLS + qb];
q_v4[d] = qr;
}
if (k_group >= 2) { q_v4[4] = {0,0,0,0}; q_scale_cached[4] = 0; }
}
const uint8_t* kv_lds = reinterpret_cast<const uint8_t*>(smem + (cur == 0 ? OFF_KV0 : OFF_KV1));
const uint8_t* sc_lds = kv_lds + BUF_KV_BYTES;
const int tok_kv_base = tok * KV_PACKED + (k_group << 4);
const int tok_sc_base = tok * SCALE_COLS + k_group;
// Prefetch next tile (overlaps with entire tile compute)
if (ti + 1 < TILES_PER_BLOCK)
load_tile((tile_base + ti + 1) * TILE_TOKENS, 1 - cur);
// ---- Fused Score + V dequant, double-buffered LDS reads ----
f32x4 scores = {0,0,0,0};
const bool tok_valid = tok < tile_len;
// Prefetch chunk 0
v4i32 kv_next = {0,0,0,0};
int sc_next = 0;
if (tok_valid) {
kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base);
sc_next = (int)sc_lds[tok_sc_base];
}
#pragma unroll
for (int d = 0; d < FP4_K_CHUNKS; d++) {
v4i32 kv = kv_next;
int sc = sc_next;
// Score MFMA first (fires on matrix core immediately)
scores = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
to_v8(q_v4[d]), to_v8(kv),
scores, 4, 4, 0, q_scale_cached[d], 0, sc);
// Prefetch chunk d+1 (LDS read during MFMA + dequant)
if (d + 1 < FP4_K_CHUNKS) {
const int next_dim = (d + 1) * (FP4_K / 2);
if (tok_valid && next_dim + (k_group << 4) + 16 <= KV_PACKED) {
kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base + next_dim);
sc_next = (int)sc_lds[tok_sc_base + (d + 1) * (FP4_K / 32)];
} else {
kv_next = {0,0,0,0}; sc_next = 0;
}
}
// V dequant: fp4 → f16 → bf8 (VALU, overlaps with MFMA on matrix core)
if (d < V_K_CHUNKS) {
uint8_t* bf8_dst = v_fp8_lds + tok * V_FP8_STRIDE + d * FP4_K + (k_group << 5);
float block_sf = tok_valid ? __uint_as_float((uint32_t)sc << 23) : 0.0f;
ev_short2 ztmp = {0, 0};
#pragma unroll
for (int w = 0; w < 4; w++) {
uint32_t pk = (uint32_t)kv[w];
f16x2 h0 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 0);
f16x2 h1 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 1);
f16x2 h2 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 2);
f16x2 h3 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 3);
ev_short2 b01 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(ztmp, h0, safe_scale, false);
ev_short2 b0123 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(b01, h1, safe_scale, true);
ev_short2 b45 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(ztmp, h2, safe_scale, false);
ev_short2 b4567 = __builtin_amdgcn_cvt_scalef32_pk_fp8_f16(b45, h3, safe_scale, true);
*reinterpret_cast<ev_short2*>(bf8_dst + w * 8) = b0123;
*reinterpret_cast<ev_short2*>(bf8_dst + w * 8 + 4) = b4567;
}
}
}
// ---- Register-based softmax: no score_lds ----
if (!tok_valid) { scores[0] = scores[1] = scores[2] = scores[3] = -INFINITY; }
else { scores[0] *= SM_LOG2E; scores[1] *= SM_LOG2E; scores[2] *= SM_LOG2E; scores[3] *= SM_LOG2E; }
// Warp-local max over 16 tokens (DPP mirror butterfly within k_group row)
float warp_max[HEADS_PER_LANE];
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
warp_max[h] = dpp_row_max(scores[h]);
// Warp-local exp + sum (DPP mirror butterfly, no * LOG2E — already in log2 space)
float attn_w[HEADS_PER_LANE], warp_sum[HEADS_PER_LANE];
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
attn_w[h] = __builtin_amdgcn_exp2f(scores[h] - warp_max[h]);
warp_sum[h] = dpp_row_sum(attn_w[h]);
}
// Cross-warp softmax exchange
float* scratch = reinterpret_cast<float*>(smem + OFF_SCORE);
const int sg = k_group << 2;
if (lane_mod16 == 0) {
*reinterpret_cast<f32x4*>(&scratch[warp_id * NUM_HEADS + sg]) =
(f32x4){warp_max[0], warp_max[1], warp_max[2], warp_max[3]};
*reinterpret_cast<f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + warp_id * NUM_HEADS + sg]) =
(f32x4){warp_sum[0], warp_sum[1], warp_sum[2], warp_sum[3]};
}
__syncthreads();
// Global max: init from own warp (already in regs, skip -INFINITY)
float my_tm[HEADS_PER_LANE] = {warp_max[0], warp_max[1], warp_max[2], warp_max[3]};
#pragma unroll
for (int w = 0; w < NUM_WARPS; w++) {
f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + sg]);
my_tm[0] = fmaxf(my_tm[0], wm[0]); my_tm[1] = fmaxf(my_tm[1], wm[1]);
my_tm[2] = fmaxf(my_tm[2], wm[2]); my_tm[3] = fmaxf(my_tm[3], wm[3]);
}
// Corrected sum: peel warp 0
f32x4 wm0 = *reinterpret_cast<const f32x4*>(&scratch[sg]);
f32x4 ws0 = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + sg]);
float my_ts[HEADS_PER_LANE] = {
ws0[0] * __builtin_amdgcn_exp2f(wm0[0] - my_tm[0]),
ws0[1] * __builtin_amdgcn_exp2f(wm0[1] - my_tm[1]),
ws0[2] * __builtin_amdgcn_exp2f(wm0[2] - my_tm[2]),
ws0[3] * __builtin_amdgcn_exp2f(wm0[3] - my_tm[3])
};
#pragma unroll
for (int w = 1; w < NUM_WARPS; w++) {
f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + sg]);
f32x4 ws = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + w * NUM_HEADS + sg]);
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
my_ts[h] += ws[h] * __builtin_amdgcn_exp2f(wm[h] - my_tm[h]);
}
// Correct attn weights with global max + write fp8 to LDS
{
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
attn_w[h] *= __builtin_amdgcn_exp2f(warp_max[h] - my_tm[h]);
ev_short2 tmp = {0, 0};
ev_short2 fp8_lo = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(tmp, attn_w[0], attn_w[1], 1.0f, false);
ev_short2 fp8_all = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(fp8_lo, attn_w[2], attn_w[3], 1.0f, true);
int pad = ((tok & 16) >> 1) + ((tok & 32) << 2);
*reinterpret_cast<ev_short2*>(&attn_fp8_lds[tok * NUM_HEADS + pad + (k_group << 2)]) = fp8_all;
}
__syncthreads();
// ---- Pre-compute merge params (VALU, doesn't need v_out) ----
float co_arr[HEADS_PER_LANE], cn_arr[HEADS_PER_LANE];
if (ti > 0) {
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
float nm = fmaxf(running_max[h], my_tm[h]);
co_arr[h] = __builtin_amdgcn_exp2f(running_max[h] - nm);
cn_arr[h] = __builtin_amdgcn_exp2f(my_tm[h] - nm);
running_sum[h] = fmaf(running_sum[h], co_arr[h], my_ts[h] * cn_arr[h]);
running_max[h] = nm;
}
}
// ---- V multiply + interleaved merge pre-scale ----
f32x4 v_out[4];
#pragma unroll
for (int i = 0; i < 4; i++) v_out[i] = {0,0,0,0};
{
const int half = lane_mod16 & 1;
const int tok_in_grp = lane_mod16 >> 1;
const int half8 = half << 3;
const int kg16 = k_group << 4;
int tok_arr[4];
tok_arr[0] = kg16 + tok_in_grp;
tok_arr[1] = kg16 + 8 + tok_in_grp;
tok_arr[2] = 64 + kg16 + tok_in_grp;
tok_arr[3] = 64 + kg16 + 8 + tok_in_grp;
// Load A (attn fp8): 4 transpose reads, reused across all rounds
const uintptr_t attn_base = reinterpret_cast<uintptr_t>(attn_fp8_lds) + half8;
v8i32 a_reg;
{
v2i32* a_parts = reinterpret_cast<v2i32*>(&a_reg);
#pragma unroll
for (int c = 0; c < 4; c++) {
int t_c = tok_arr[c];
int pad_c = ((t_c & 16) >> 1) + ((t_c & 32) << 2);
a_parts[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(attn_base + t_c * NUM_HEADS + pad_c));
}
}
// V fp8: double-buffered B reads
const uintptr_t vfp8_base = reinterpret_cast<uintptr_t>(v_fp8_lds) + half8;
v8i32 b_reg[2];
// Pre-load round 0
{
const uintptr_t vb0 = vfp8_base + ((0 * 8 + warp_id) << 4);
v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[0]);
#pragma unroll
for (int c = 0; c < 4; c++)
bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(vb0 + tok_arr[c] * V_FP8_STRIDE));
}
#pragma unroll
for (int r = 0; r < 4; r++) {
int cur = r & 1, nxt = (r + 1) & 1;
// Prefetch next round
if (r + 1 < 4) {
const uintptr_t vb_next = vfp8_base + (((r + 1) * 8 + warp_id) << 4);
v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[nxt]);
#pragma unroll
for (int c = 0; c < 4; c++)
bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(vb_next + tok_arr[c] * V_FP8_STRIDE));
}
// Issue MFMA (matrix core)
v_out[r] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg[cur], v_out[r], 0, 0, 0, 127, 0, mfma_v_scale);
// Interleave merge pre-scale on VALU during MFMA execution
if (ti > 0) {
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
running_v[r][h] *= co_arr[h];
}
}
}
// ---- Merge: post-MFMA (just add v_out * cn, or copy for ti==0) ----
if (ti == 0) {
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
running_max[h] = my_tm[h];
running_sum[h] = my_ts[h];
#pragma unroll
for (int i = 0; i < 4; i++) running_v[i][h] = v_out[i][h];
}
} else {
#pragma unroll
for (int i = 0; i < 4; i++)
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
running_v[i][h] += v_out[i][h] * cn_arr[h];
}
}
// ---- Write output ----
if (NUM_PARTIALS == 1) {
// Direct output: divide by sum, write final result (skip reduce kernel)
__hip_bfloat16* ob = final_output + (q_start * NUM_HEADS) * V_HEAD_DIM;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
float inv = (running_sum[h] > 0.0f) ? (1.0f / running_sum[h]) : 0.0f;
#pragma unroll
for (int i = 0; i < 4; i++) {
const int vd = (i*8 + warp_id)*16 + lane_mod16;
ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2bfloat16(running_v[i][h] * inv);
}
}
} else {
// Write partial for reduce kernel
const int pidx = batch_idx * NUM_PARTIALS + partial_idx;
__half* ob = partial_out + pidx * (NUM_HEADS * V_HEAD_DIM);
#pragma unroll
for (int i = 0; i < 4; i++) {
const int vd = (i*8 + warp_id)*16 + lane_mod16;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2half(running_v[i][h]);
}
if (lane_mod16 == 0 && warp_id == 0) {
float* pm = partial_max + pidx * NUM_HEADS;
float* ps = partial_sum + pidx * NUM_HEADS;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
pm[(k_group<<2)+h] = running_max[h];
ps[(k_group<<2)+h] = running_sum[h];
}
}
}
}
// ============================================================
// Reduce kernel
// ============================================================
template<int NUM_PARTIALS>
__global__ void __launch_bounds__(WAVEFRONT_SIZE)
mla_reduce(
const __half* __restrict__ partial_out,
const float* __restrict__ partial_max,
const float* __restrict__ partial_sum,
const int32_t* __restrict__ qo_indptr,
__hip_bfloat16* __restrict__ output
) {
const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE;
constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;
const int bo = batch_idx * NUM_PARTIALS;
const float* pm = partial_max + bo * NUM_HEADS + head_idx;
const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
const __half* pv = partial_out + bo * NUM_HEADS * V_HEAD_DIM
+ head_idx * V_HEAD_DIM + lane_id * ELEMS;
// Peel first iteration
float local_max = *pm; pm += NUM_HEADS;
float local_sum = *ps; ps += NUM_HEADS;
union { u32x4_vec raw; __half2 h2[4]; } cur;
cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
__half2 v_acc[4];
#pragma unroll
for (int i = 0; i < 4; i++) v_acc[i] = cur.h2[i];
#pragma unroll
for (int p = 1; p < NUM_PARTIALS; p++) {
float cm = *pm; pm += NUM_HEADS;
float cs = *ps; ps += NUM_HEADS;
cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
float nm = fmaxf(local_max, cm);
float co = __builtin_amdgcn_exp2f(local_max - nm);
float cn = __builtin_amdgcn_exp2f(cm - nm);
local_sum = fmaf(local_sum, co, cs * cn);
local_max = nm;
__half2 co2 = __float2half2_rn(co);
__half2 cn2 = __float2half2_rn(cn);
#pragma unroll
for (int i = 0; i < 4; i++)
v_acc[i] = __hfma2(cur.h2[i], cn2, __hmul2(v_acc[i], co2));
}
// Normalize in f32, convert to bf16
float inv = (local_sum > 0.0f) ? __builtin_amdgcn_rcpf(local_sum) : 0.0f;
u32x4_vec out_pack;
#pragma unroll
for (int i = 0; i < 4; i++) {
float2 f = __half22float2(v_acc[i]);
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
}
*reinterpret_cast<u32x4_vec*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}
// ============================================================
// Multi-wavefront reduce for high-partial shapes (e.g. 64 partials)
// Each wavefront handles NUM_PARTIALS/NWARPS partials, then LDS merge
// ============================================================
template<int NUM_PARTIALS, int NWARPS>
__global__ void __launch_bounds__(NWARPS * WAVEFRONT_SIZE)
mla_reduce_multi(
const __half* __restrict__ partial_out,
const float* __restrict__ partial_max,
const float* __restrict__ partial_sum,
const int32_t* __restrict__ qo_indptr,
__hip_bfloat16* __restrict__ output
) {
const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
const int warp_id = threadIdx.x >> 6;
const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE;
constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;
constexpr int PARTIALS_PER_WARP = NUM_PARTIALS / NWARPS;
const int bo = batch_idx * NUM_PARTIALS;
const float* pm = partial_max + bo * NUM_HEADS + head_idx;
const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
const __half* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
+ head_idx * V_HEAD_DIM + lane_id * ELEMS;
// Peel first iteration per warp
const int p_start = warp_id * PARTIALS_PER_WARP;
const float* wpm = pm + p_start * NUM_HEADS;
const float* wps = ps + p_start * NUM_HEADS;
const __half* pv = pv_base + p_start * PV_STRIDE;
float local_max = *wpm; wpm += NUM_HEADS;
float local_sum = *wps; wps += NUM_HEADS;
union { u32x4_vec raw; __half2 h2[4]; } cur;
cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
__half2 v_acc[4];
#pragma unroll
for (int i = 0; i < 4; i++) v_acc[i] = cur.h2[i];
#pragma unroll
for (int p = 1; p < PARTIALS_PER_WARP; p++) {
float cm = *wpm; wpm += NUM_HEADS;
float cs = *wps; wps += NUM_HEADS;
cur.raw = *reinterpret_cast<const u32x4_vec*>(pv); pv += PV_STRIDE;
float nm = fmaxf(local_max, cm);
float co = __builtin_amdgcn_exp2f(local_max - nm);
float cn = __builtin_amdgcn_exp2f(cm - nm);
local_sum = fmaf(local_sum, co, cs * cn);
local_max = nm;
__half2 co2 = __float2half2_rn(co);
__half2 cn2 = __float2half2_rn(cn);
#pragma unroll
for (int i = 0; i < 4; i++)
v_acc[i] = __hfma2(cur.h2[i], cn2, __hmul2(v_acc[i], co2));
}
// F16 LDS merge: pre-scale in f16, write/read __half2, halved LDS bandwidth
extern __shared__ char reduce_smem[];
float* smem_max = reinterpret_cast<float*>(reduce_smem);
float* smem_sum = smem_max + NWARPS;
__half2* smem_vh = reinterpret_cast<__half2*>(smem_sum + NWARPS);
if (lane_id == 0) {
smem_max[warp_id] = local_max;
smem_sum[warp_id] = local_sum;
}
__syncthreads();
if constexpr (NWARPS <= 4) {
// === Flat merge: pre-scale in f16, warp 0 sums in f16 ===
float global_max = -INFINITY;
#pragma unroll
for (int w = 0; w < NWARPS; w++)
global_max = fmaxf(global_max, smem_max[w]);
__half2 scale2 = __float2half2_rn(__builtin_amdgcn_exp2f(local_max - global_max));
local_sum *= __builtin_amdgcn_exp2f(local_max - global_max);
#pragma unroll
for (int i = 0; i < 4; i++) v_acc[i] = __hmul2(v_acc[i], scale2);
__half2* wv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) wv_dst[i] = v_acc[i];
if (lane_id == 0) smem_sum[warp_id] = local_sum;
__syncthreads();
if (warp_id == 0) {
__half2 gv[4];
#pragma unroll
for (int i = 0; i < 4; i++) gv[i] = __float2half2_rn(0.0f);
float gs = 0;
#pragma unroll
for (int w = 0; w < NWARPS; w++) {
const __half2* wv = &smem_vh[(w * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) gv[i] = __hadd2(gv[i], wv[i]);
gs += smem_sum[w];
}
float inv = (gs > 0.0f) ? __builtin_amdgcn_rcpf(gs) : 0.0f;
u32x4_vec out_pack;
#pragma unroll
for (int i = 0; i < 4; i++) {
float2 f = __half22float2(gv[i]);
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
}
*reinterpret_cast<u32x4_vec*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}
} else {
// === Two-tier merge in f16 ===
constexpr int GROUP_SIZE = 4;
constexpr int NUM_GROUPS = NWARPS / GROUP_SIZE;
const int group_id = warp_id / GROUP_SIZE;
const int local_warp = warp_id & (GROUP_SIZE - 1);
float group_max = -INFINITY;
#pragma unroll
for (int w = group_id * GROUP_SIZE; w < (group_id + 1) * GROUP_SIZE; w++)
group_max = fmaxf(group_max, smem_max[w]);
__half2 scale2 = __float2half2_rn(__builtin_amdgcn_exp2f(local_max - group_max));
local_sum *= __builtin_amdgcn_exp2f(local_max - group_max);
#pragma unroll
for (int i = 0; i < 4; i++) v_acc[i] = __hmul2(v_acc[i], scale2);
__half2* wv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) wv_dst[i] = v_acc[i];
if (lane_id == 0) smem_sum[warp_id] = local_sum;
__syncthreads();
// Group leader sums in f16
if (local_warp == 0) {
__half2 gv[4];
#pragma unroll
for (int i = 0; i < 4; i++) gv[i] = __float2half2_rn(0.0f);
float gsum = 0;
#pragma unroll
for (int w = group_id * GROUP_SIZE; w < (group_id + 1) * GROUP_SIZE; w++) {
const __half2* wv = &smem_vh[(w * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) gv[i] = __hadd2(gv[i], wv[i]);
gsum += smem_sum[w];
}
if (lane_id == 0) { smem_max[warp_id] = group_max; smem_sum[warp_id] = gsum; }
__half2* gv_dst = &smem_vh[(warp_id * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) gv_dst[i] = gv[i];
}
__syncthreads();
// Warp 0 merges groups with rescaling in f16
if (warp_id == 0) {
float fmax = -INFINITY;
#pragma unroll
for (int g = 0; g < NUM_GROUPS; g++)
fmax = fmaxf(fmax, smem_max[g * GROUP_SIZE]);
__half2 fv[4];
#pragma unroll
for (int i = 0; i < 4; i++) fv[i] = __float2half2_rn(0.0f);
float fsum = 0;
#pragma unroll
for (int g = 0; g < NUM_GROUPS; g++) {
__half2 s2 = __float2half2_rn(__builtin_amdgcn_exp2f(smem_max[g * GROUP_SIZE] - fmax));
const __half2* gv = &smem_vh[(g * GROUP_SIZE * WAVEFRONT_SIZE + lane_id) * (ELEMS/2)];
#pragma unroll
for (int i = 0; i < 4; i++) fv[i] = __hfma2(gv[i], s2, fv[i]);
fsum += smem_sum[g * GROUP_SIZE] * __builtin_amdgcn_exp2f(smem_max[g * GROUP_SIZE] - fmax);
}
float inv = (fsum > 0.0f) ? __builtin_amdgcn_rcpf(fsum) : 0.0f;
u32x4_vec out_pack;
#pragma unroll
for (int i = 0; i < 4; i++) {
float2 f = __half22float2(fv[i]);
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[i]) : "v"(f.x * inv), "v"(f.y * inv));
}
*reinterpret_cast<u32x4_vec*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}
}
}
// ============================================================
// Dispatch
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
void run_impl(uintptr_t q_ptr, uintptr_t kv_data_ptr, uintptr_t kv_scales_ptr,
uintptr_t qo_indptr_ptr, uintptr_t kv_indptr_ptr, uintptr_t output_ptr,
uintptr_t pout_ptr, uintptr_t pmax_ptr, uintptr_t psum_ptr,
int batch_size, uintptr_t kv_scale_ptr) {
const auto* q_d = reinterpret_cast<const __hip_bfloat16*>(q_ptr);
const auto* kvd = reinterpret_cast<const uint8_t*>(kv_data_ptr);
const auto* kvs = reinterpret_cast<const uint8_t*>(kv_scales_ptr);
const auto* qoi = reinterpret_cast<const int32_t*>(qo_indptr_ptr);
const auto* kvi = reinterpret_cast<const int32_t*>(kv_indptr_ptr);
auto* out = reinterpret_cast<__hip_bfloat16*>(output_ptr);
auto* pout = reinterpret_cast<__half*>(pout_ptr);
auto* pmax = reinterpret_cast<float*>(pmax_ptr);
auto* psum = reinterpret_cast<float*>(psum_ptr);
const auto* kvsc = reinterpret_cast<const float*>(kv_scale_ptr);
const int tb = batch_size * NUM_PARTIALS;
hipLaunchKernelGGL((mla_grouped<NUM_PARTIALS, TILES_PER_BLOCK>),
dim3(tb), dim3(BLOCK_THREADS), TOTAL_LDS, 0,
q_d, kvd, kvs, qoi, kvi, pout, pmax, psum, out, kvsc);
if (NUM_PARTIALS > 1) {
if (NUM_PARTIALS >= 8) {
// Multi-wavefront reduce: 4 warps for 64 partials, 2 warps for 8
constexpr int REDUCE_WARPS = (NUM_PARTIALS >= 64) ? 16 : 2;
constexpr int REDUCE_THREADS = REDUCE_WARPS * WAVEFRONT_SIZE;
constexpr int REDUCE_LDS = REDUCE_WARPS * (V_HEAD_DIM * 2 + 2 * 4);
hipLaunchKernelGGL((mla_reduce_multi<NUM_PARTIALS, REDUCE_WARPS>),
dim3(batch_size * NUM_HEADS), dim3(REDUCE_THREADS), REDUCE_LDS, 0,
pout, pmax, psum, qoi, out);
} else {
hipLaunchKernelGGL((mla_reduce<NUM_PARTIALS>),
dim3(batch_size * NUM_HEADS), dim3(WAVEFRONT_SIZE), 0, 0,
pout, pmax, psum, qoi, out);
}
}
}
void run_s1(uintptr_t a, uintptr_t b, uintptr_t c, uintptr_t d, uintptr_t e, uintptr_t f,
uintptr_t po, uintptr_t pm, uintptr_t ps,
int i, int max_kvseqlen, uintptr_t kv_sc) {
const int total_tiles = max_kvseqlen / 128;
const int tpb = max(1, i * total_tiles / 256);
const int np = total_tiles / tpb;
if (np == 1 && tpb == 8) run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 1 && tpb == 64) run_impl<1, 32>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 4 && tpb == 2) run_impl<4, 2>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 4 && tpb == 16) run_impl<4, 8>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 8 && tpb == 1) run_impl<8, 1>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 8 && tpb == 8) run_impl<8, 4>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else if (np == 64 && tpb == 1) run_impl<64, 1>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
else run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, i, kv_sc);
}
PYBIND11_MODULE(mla_new_kernel_module, m) { m.def("run_s1", &run_s1); }
"""
hip_module = load_inline(
name="mla_new_kernel_module",
cpp_sources="",
cuda_sources=CUDA_SRC,
with_cuda=True,
verbose=True,
extra_cuda_cflags=["-std=c++20", "-O3"],
no_implicit_headers=True,
)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_buffer, kv_scales = kv_data["mxfp4"]
_, fp8_scale = kv_data["fp8"]
total_q_len = q.size(0)
max_kvseqlen = kv_buffer.size(0) // batch_size
total_tiles = max_kvseqlen // 128
tpb = max(1, batch_size * total_tiles // 256)
np_ = total_tiles // tpb
tb = batch_size * np_
pout = torch.empty(tb, NUM_HEADS, V_HEAD_DIM, dtype=torch.float16, device=q.device)
pmax = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
psum = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
output = torch.empty(total_q_len, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=q.device)
hip_module.run_s1(
q.data_ptr(), kv_buffer.data_ptr(), kv_scales.data_ptr(),
qo_indptr.data_ptr(), kv_indptr.data_ptr(), output.data_ptr(),
pout.data_ptr(), pmax.data_ptr(), psum.data_ptr(),
batch_size, max_kvseqlen, fp8_scale.data_ptr())
return output
scrolls · 906 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON