submission 675094
gandan09 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.
submission_v2b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-675094?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:4240594cc2e8c08aadc7bf84a86d84673df97d510c2760a11e41e4a691587bd5
license declaredunknown
license concludedunknown
authorsgandan09
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""v2b_fast: FP4 native MLA decode — MFMA QK + TR_B8 PV + async prefetch.Kernel source
submission_v2b.py161 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v2b_fast: FP4 native MLA decode — MFMA QK + TR_B8 PV + async prefetch.
0.93x geomean vs FP8 aiter reference. Validates all 8 benchmark configs.
"""
import torch
import struct as _struct
import ctypes as _ctypes
import os as _os
import gzip as _gz
import base64 as _b64
import tempfile as _tf
import subprocess as _sp
from task import input_t, output_t
from utils import make_match_reference
import aiter
from aiter import dtypes as aiter_dtypes
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.utility.fp4_utils import dynamic_mxfp4_quant
NUM_HEADS = 16; NUM_KV_HEADS = 1; QK_HEAD_DIM = 576; V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5); PAGE_SIZE = 1; NUM_KV_SPLITS = 32
FP8_DTYPE = aiter_dtypes.fp8
# --- Reference (FP8 fallback) ---
def _make_meta(bs, mql, nh, nhk, qd, kd, qo, kvi, kl, nks=NUM_KV_SPLITS):
info = get_mla_metadata_info_v1(bs, mql, nh, qd, kd, is_sparse=False, fast_mode=False,
num_kv_splits=nks, intra_batch_mode=True)
w = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(qo, kvi, kl, nh//nhk, nhk, True, w[0], w[2], w[1], w[3], w[4], w[5],
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE,16), max_seqlen_qo=mql,
uni_seqlen_qo=mql, fast_mode=False, max_split_per_batch=nks, intra_batch_mode=True,
dtype_q=qd, dtype_kv=kd)
return dict(work_meta_data=w[0],work_indptr=w[1],work_info_set=w[2],
reduce_indptr=w[3],reduce_final_map=w[4],reduce_partial_map=w[5])
_fp8_cache = {}
def _fp8_fallback(data):
q, kv_data, qo, kvi, cfg = data
q_fp8, q_sc = aiter.per_tensor_quant_hip(q.view(-1, QK_HEAD_DIM), quant_dtype=FP8_DTYPE)
q_fp8 = q_fp8.view(q.shape)
kv8, kv8s = kv_data["fp8"]
bs = cfg["batch_size"]; nq = cfg["num_heads"]; nkv = cfg["num_kv_heads"]
dq = cfg["qk_head_dim"]; dv = cfg["v_head_dim"]
tot = int(kvi[-1].item())
kv4d = kv8.view(kv8.shape[0], PAGE_SIZE, nkv, kv8.shape[-1])
key = (bs, tot)
if key not in _fp8_cache:
ki = torch.arange(tot, dtype=torch.int32, device="cuda")
kl = (kvi[1:]-kvi[:-1]).to(torch.int32)
meta = _make_meta(bs,1,nq,nkv,FP8_DTYPE,FP8_DTYPE,qo,kvi,kl)
o = torch.empty((q.shape[0],nq,dv), dtype=torch.bfloat16, device="cuda")
_fp8_cache[key] = (ki, kl, meta, o)
ki, kl, meta, o = _fp8_cache[key]
mla_decode_fwd(q_fp8.view(-1,nq,dq),kv4d,o,qo,kvi,ki,kl,1,page_size=PAGE_SIZE,
nhead_kv=nkv,sm_scale=SM_SCALE,logit_cap=0.0,num_kv_splits=NUM_KV_SPLITS,
q_scale=q_sc,kv_scale=kv8s,intra_batch_mode=True,**meta)
return o
def ref_kernel(data):
return _fp8_fallback(data)
check_implementation = make_match_reference(ref_kernel, rtol=1e-01, atol=1e-01)
# --- FP4 native kernel (embedded ASM, compiled at runtime) ---
_ASM_GZ_B64 = 'H4sIAKeTyWkC/+1d/XLbOJL/30/B2r8yO7KOAAHwY2qqbnZm524qnhpnk8pdXSrFoiVK1llfQ1KKnQe4B7hH3CdZfBH8AiGTspO4CtkdWmI3GsCvG+gGCLamyWa+nG3jIsmWaeH8RXy9pH/Yf7d5cnm5XNyH2P3LxbRI74uL6XK9u1k7m3USL/Yo3ibF6pjG83S2m6fxXZpt0/XFdA+T9Wq5dQJa6GGfGrkn/744bGfFare9mOa0Cf/5959+eTsBRHx7/8tvv08wgOLbr9fo7QQGgfj29mf6BYnP7367+rsqdPXL2/jNxL13a1/f0u8AurVbr9+zW15Qv8W4IK7fevszu+XXC75/x27d1G9d0zuocefdH4wLzeg9U++jiwvH+cG5vLx0rnbJ3BF3HaqNnN2kxDxeU0I8/7TL5vfQyT+gCH+c0L9uBOhf2ktXx0Uiv8kV6LiCKGxwAa0sQMmgyaeVBmAEvAYf1MtDEWj2AWrlERSRJp/XlufkgDQYgi6D32D4uUL83W2WUsx/+6XE+hgn23l848E49aBzhBOHeBPn6HLSOr/NsvRYkQElK2qjIC0EMCVBfUE0cVCdum5QaUFa8RHoqbRGr7+sL8oK6mZXp2DW+8XCX7B/GjoV3NNPElZS8/hTsipm28JZL+829O8r9zuNMAY5HY0K6f/IVnNlz5KVcgSUDbbu0cpyrz0qjvFqO98XWVcGdFlVLSEQsJte6yYUCCxKBOo0T1irWw4mCqykoIloqKgimc/jA7+PJ4KINN3n/WKEm8NikWZ1czz6tIOC4wN0I8jGi+vsKNu2TxDWCwpcg6BKT8e6lpi9L1ZZXqyTbVp1kDaql8zsOBCo5Icb2XkyUQhUqvr59rC9UxpSSBHeCVYC+E0xnMIuoKkNZnBu85bvctaL6dVsHnHSbLOP14UQxKlS+uwmS7az2zifzYDD2OfNSjlvvUDVUGZ67CKaI+RwEbxaWW+Xnfke51KWkvMA7bEvyZeMLnpzWMcrXpRXz8rDdhuY5XKypGxW2zoFcCwVAMtUEt3a/SYA6f2qqHR0nVHvs9kfitR57xSUMd/v8tT53vn9199/cq7fO7QlWZrnaTkl/kDNYgmI86OzdP7q0A+vFrtMusGSebVdftczH7liqvOkpDyggnIqKOjhB2J6Q5L/U0Id5rGIbxLayB/Lar932H0qBbqoTxAd7ACoaVTiW9GkH2dfZFV7xpTJSq5pHbK/3xvb61X9a9fC+sH6Qz/pqSJw4HSln7/xce7M03yWrfbFLqvCAGU93LLYlMRDJdo4Fha1Zic2slDDvmaSwqwOd8aX5/ZPjh5oTo79DaERWbMdHoU6D3Tt8CgAeXece8jQDmxqB5vHukB4bN4h2gawQn63AYGhAeGpBnQAQHyku7oGIO6kQKcFyOCnUMtPVQQGMkCtm0zXADdvYoOmMegRj5keSUs8ZiokbfEGBeKWAo1RWm+sUZ9HoTTA7lzOG8zJYauMJyml0SKNzDqdLUE4WzMmeUNnhjd0kk3WdK4sdmw0d6MTTE0PEbdv9ghEYEWXDVfrPyPOVLm1OCWIBcEkIiyUF64ek8rLUHJ6n84m8lqyNpwAI312mPj56XAkiDxwOoqY07vZqhDxAS9ahg/NGY75dVdEF0HToa7/5B2WDpWDpfpyCboIekFPkCVj72fC7M8ctDA73DxQx1nHzIORhx8TeTWwYaGbXJVKdCpQgzKyCy94E3pB0uANMXl+UODXBwUaLKdnzrhJsmyVZtUg/m27Kpx8tyg2yT0dyjT0yYuEduTVPs0uWfzrsEAnWa9ZEHBLA+P8O7VKbBgigPWVFRO/ifd0nNNQ4vLXq3fx7z/9NxfqMKGaxRJbwLpO9e8HZ12Wd40FmXNtFmTx22zGCmq4vUHcaBA3HsRNBnH7g7iDQdzhEG7mywdwg0Hcg3SJBukSDdIlGqRLNEiXaJAu0SBdokG6xIN0iQfpEg/SJR6kSzxIl3iQLvEgXeJBusSDdIkH6ZIM0iUZpEsySJfErMtWWENX7RfNvQnsl2t86Z7V6p9T2IVvJjQj0dd0Jb9aM8+VZNRfpevv1Nr9+pYtnEHk/E+a8RC1tdFC5O6FWkOeiFI/3x3dvjg1iEiowgtiDC84qza8oBXMXU1sACblnn0nNijpR58MikJZXy5kjY+JRGtgiWXekID0icDJQWxCpydykuCUAp4uonyqXsFzewVjd2hIWBsfMHKu5dBhg2meFInDQlvnFXay3WE7zyc8EGT3cmeV54d07tykND5MnWT74DD5rZUpFjF2bUuiIoVyKEvSD84/WB2Oe4kiIbysaEuXlTTuvTzK9U//0GRSeBhOA272BOpHp+DPNuLVXK5wVS2PUSUOH6PKtrnwfgWlsnp3rEO+RCCRF7aWCBrdla0Gp0f2V+sUGNsp+A13Co7tlPcNd8ob2yn0DXcKDe/Uf9H5Sq1usZht+hbplJtNuZy1KEMIY5RQm2scPl0/GzoBPhEV1J5dmSaXWjO/go4f1wtwcjZ5Cb2AJ6ePl9AL7+R88RJ6gfp6UYuMvFpkxLe6ZVwEZVzUF/WUz0G6QY85dP6qUYrcyKxFKciNkHd2lCJi6K/jJco+geF9Yl7ilFcQjz/OWQudZ+A9SwPzvP81VPSoFvfM8Y/Z1m48/OKbC5iUmwXTK7Y9EHUODYg9CKA7NLBO8iKWBwXUZsNP+cN2xhZI+yxdpAVdRJfPu6rDFaBvA4NT2KV7GuLUWomVUtsTza1zvvd+OiDRaNwdpm/ziDtpDZrAjbd9WOTW6b03tvfDfdBzAOCdDQAaCwA7OfL1AUBnA4DHAuC5Pvz6AOCzASBjAUBuSL4+AGQoAI+Ptti0qY22APBL0LTTovs1EJGBCm/bsEil07ugt3cDdjOfrX/BsP4p7/vmdUTX6jzgmvPTetWhZaZ31sgDRFVf+Vn58mRad3CE4uAa0gtg+z/seH1ZvoGwwrdxN9TeZUefdLeBVkfqSKv2REm5IfymjO5oCMU3OW8ApDHUBwgjyI4diEDK9ARBW5pEMOwt7dXOJgRQzyAuqCb6QAWzk3ViyWgKCbUyeVCIdTIBpxjPYy82iTidFC+oQEDu2f9hEC+CBVkg2uFAHOqv4VZhoIi89R5wdvs4T9fx7Sr64E7o/z46s5v8c4Scm/VyHyG9KgkqsX6xmtQQUanP59CyhohKXVsL+EYsIPjCFhBYC/jGLICv3b6kCYgKrQ18SzZAvrQNkJdhAypkfTtjT8j5Lpfzir+y4VZnKFn8nf5Z9XA2E6sP0BtZUxbtAYJ85qoodlF/k4S/h3YMNERYEkMN0SuJQCcXKSroeWEL+6qDYqut6uDRqy05t/NNkt81X36RF1z2VsfGGw9PsvGqvJNsvD+oxdZZL7BhHCjT7JKBWGt7PTG9uGgO1IhXbPWlOIMa2fXnF3yZAHruQ3kIIiKoh8MrOQAMeliQYgnhBTexJzhyXD1KkceOX9WPF0+c8uRxNUg672KWmxpe/55HbV+r8VIqW/J55hUhMZzlN5/0D8MA1Xa269NwQKJArqu7xNCNQk8QFeJEw4ai0G+yeVDDRqcqFzT5UGCcGpP7amxDPpVyo/N7qIIlMFJDEzV0jVRgpEIj1TNSkZGKjVRipBqxCo1YhUasgOuayTq0uA/hrysLKns20SwsqEJ0er/viO5O+h4s6/SghuopqqehIkVFGipWVKyhEkUlGqqvqL6GGiiqxgOyk++SqnGB7KS7oCKNC0RKA0iDFVJYIQ1WSGGFNFghhRXSYIUUVkiDFVJYIQ1WSGGFNFghhRXSYIUUVkiDFVZYYQ1WWGGFNVhhhRXWYIUVVliDFVZYYQ1WWGGFNVhhhRXWYIUVVliDFVZYYQ1WWGGFNVgRhRXRYEUUVkSDFVFYEQ1WRGFFNFgB1V9A9C/wlPNBY7YQroD0zBaSqinGo06/r5hf+ZdmMb7MD/qKBZXjaRYLJ/KiLRZWHqlRLHQn8qIrJqmaYsJL9RWr+bBmMTiRF20xWDm3ZjFvIi/aYl7l9ZrF0ERetMVQ5Q6bxfBEXrTFcOUnm8XIRF60xUjlQJvF+PKjz0okVVMsmMiLtlhQudxmsXAiL9piYeWLG8WoB56UV11BRdcVlT66t2jdh7OAsj8O61LrcVgfNTRRQ9dIBUYqNFI9IxUZqdhIJUaqEavQiFVoxKrUby9ZgxafcuW8243SmHLdavpkp0SatHvkzRfsRfEaU0OALwX4GgF+W4AmQg3kw84ylm8ICNoCNEEsm1ndWrjfEBC2BWjiXDbH1qbZhgBOqwvQLRvYbOvWFg0NAaAtQKMFNu+6tXVFQwBsC9AsPtgM7NaWHg0BXluAZn3C5mK3tjppCEBtAZolDJuV3doCpiEAtwVoVjmhtMRQY4lh2xJ1C6FQWmKoscSwbYm6tVIoLTHUWGLYtkTdciqUlhhqLDFsW6JuxcVn8fpE3hAhqHUZ2nUZn8/d+rKsKQR0hOge3kOkf14Mcc990nNfHXGYHYt4fxcv9gFrjayi4V16OHjAFcrdTrHVCT72FMBlMCWHWQ8Hj3O8x4kkZQwjrbaHg1/8x4n0y9BBGoGeo/TpQCdUu+PUlwGHHTQLUXW0stpo49s1fEPpA0U7ohV/rLbJaomJInaM4Z//9/+/XlPeRvGJU+zu0u3lbLc+bLaXm+R/d1n/thmAMhGRJlkbYE6Mn0Ry9YchAMscBrk9ANhz2AhA+WQYQM3DBi5BXb0eBhSGkq27swU8NwIe3ykT9N52oloztY0ANRldhpAI9TcawZ4zAA+p+73baV1cWCmGDfukGaeenOB8V2aF6TDwCWxWZ2Bmy59msNmF2S/9s9gjjhKhKPGNQsCfVpSVyocNi8ZePODbK0BurXR2xgXZJ4qz3ELXy/LFtU+WX8nyK1nNGVS2J6y1qjmJSoagztCU4Nck+DoJfk1Cz/SH3Bo4/gjI1cQB+OMhq4HzNNA3uY9QCC1vFTJGIeDphoTVwPkaOGNIAOsWRmoAjh0DwLqFZ9DA+WPAuoVzFeI93ZCwGjhfA49asfJlDqgt1fQcQW2ZZFruuWOXe4AutIP6Uqu+NhaFP9BIMKJ9HDS4oXVwI20JjR3N0Dq4Z9DAGQ4OWgf3NArBTzckrAbO18AZQ8KzbmGkBsjYMeBZt/AMGjh/DFi3cK5C/KcbElYD52tg3LpHsyQxr3m+nVURoqsimaD8sYfv5Q+cRM7rH5msX68D59OquHXe/SP+Gzs179ysd7O7/N/4eyoqf2oyu3Ouea40Xixy8sONyA7uOrdJLh5C5hMnS/PC+Zxmu7x6uTzVJF8H1TvXjcdcWDx005KEVrCOJLDUnWnkh0cB9HWknmfsuOcZO+55xo6rZ+z1sUT7W/aHX33SSCmv5eZdxOCR3LzX/PoYbg4Ev3a4uxbMTrf2/W6MJEr7xIGwrev3wnIcN3Lmq03ufHD/Csh0yq7fA/xRJ0ce2AC48UYEQXGR8XRAct3P32QA/CCtngnyzQHB1Hh9oq0q1KNy1KNy1KNypH4wpecBcBd/8SwLPdISxDY/eqQliA1R9EhLEHtFSGsJ/BW5npfjyp8zKNXii8+Yfsbis+Bo2QMo7QFwewBmexCpKaxFvAyLkLlaDBbBOVoWAUuLgNwioNkiMIDWIl6KRZS5XPotQnC0LMIrLcLjFuGZLcIngbWIF2MRKCqtoM8iOEfLIlBpEYhbBDJbhMzjZk3iZZhEEGFgNgnO0TIJXJoE5iaBT5gEDGxs+WJMAsMIm2NLwdEyCVKaBOEmQU6YBPZscPlyTIJE2BxcCo6WSfilSfjcJPwTJuGHNrp8MSZB3IiYo0vBoTa7RHbmMk9vKn4jqPkblZqszsbEu8+cYPdZcpwDAPuzRj5/ztxn6pP3dH0angb3mfqEnq5PwzPbPlOf8NP1aXiy2mfqE3lMXtp29tkzs8w+R651IH8w+Ml/BesZWxyck4Go/ttz/JevgSlnfOOH0lj2+ItmTvjI5qW1eWltXlqbl9ZmpLR5aa0F2Ly01gZsXlqbl7adlxbYvLQ2L+3z5qUFNi+tzUtr89LavLQ2L63NS2vz0tq8tDYvrc1La/PS2ry0Ni+tzUtr89LavLQ2L63NS2vz0tq8tDYvrc1La/PS2kQcNi+tzUtr89LavLQ2L63NS2vz0tq8tDYvrc1La/PS2ry0VgM2L63NS2s1YPPS2ry0VgM2L63NS2s1YPPS2ry0Ni+tzUtrswLZvLTWImxeWpuX1ualtXlpbV5am5fW5qW1eWltXlqbl9bmpbV5aW1e2pN5aS+mV/PdNo0G7Yz9cSj2h8J5xXfI1snD7lBUGSiy2b79cr3ITbnd7R1Dkk2+iWJINRHq91g6mSZOZC4cmKYiJ7WUpY1XhOVrvuUr/jeHxSLN4rzYZWk8/7TL5pKFS1Ge3GU2n25Ly9e81CzfAy6TA2jl+qfkal7wli8Kl2kFtHKDU3J1OIRKLu6TG56SC6ARYDIaYIKMCPujESaBEeJgNMQ+NGIcjsbYJyaMed6GcRiLpDi9IPOcD+NABrrsGRXKCI5GGejyY1QwI2+8KesyYNRwRuNxDqERZzwe55AYcSajcYaua8TZH40zdJER52A0zlCX46KGczgaZ6jLYlHhjN3xOGsnOoUzBuNxJkbnh8c7Pw8a3R8e7/48aHSAeLwD9KDRBeLxLtDzjD4Qj/eBXmB0gni8E/QCoxfE472gFxrdIB7vBr3Q6AfJeD+IkNEPkvF+EGGjHyTj/SDCRj9IxvtBRFxj2kY0OG3jOk+hPn+Xew/T+U1KFosqY896t2ytQ6DbWoc001O4rUxfvT8F4PUtFVSuRh1i0FWIlXs4ErEL0bWefIHTq/R+VQhiup3vl5uLababJ0VyMd3DZL1abh1yMU0289s8ie/SbJuunc06YYdp4m1SrI60CelsN08lkUoquZfZ7sAOgiw36baIF6v7dB7nq88pz7L47o93NdZ9tjomRapjdmtsrIokWwoCH2kl5ZAzRJb7rOKRovZFxhVScuYPeZFuBC/F7k60cjWP7x/H9vA4ts8Ntm16TzuVpWl8pJwOIK6WyMQ4fr1fyWx22MTC6FvFViktsaHINyBarHdJEdNWUNNnxJgakJEOSEwNossyT7e7bKNkeEYGIcS7mFIbihvGwm1nuT/Em7RIuF2xpbxgmQqWnJnfJZWeLflH9o99nc+zNM/jfJ/M0shZrnc3yVqSaVu2yYbepdqN31Q35ezAuyNvMVuJnKC6cUzWB2qtq+28FBqLETWi6tfvu3UHX6juN3m3bj47fZmOa2rncdGXqJ0KoX+6DeBBzpdoQL5fr4o/DkW3Ccj9kk24ylNNE860QFFHvhEHKLvyccfIkF7+zUPMv7apxcOeFloobak6t7xTGssi7pNUuZJVTvtcUyQyhAimtifh3rAEs0MVxUNZlgcRazpLVk5BMMBSuujxCVdK+fpdYznNTbn3mdEpnSLlS/H5w+ZmtzZXML2bC+ZjTQCQWE/ZycNFtlNdo2tAOW0f0yxf7bZi2gb86l5Mp1M1/dfn+38BRSOmc2vsAAA='
_td = _tf.gettempdir()
_s_path = _os.path.join(_td, "mla_fp4_v2b.s")
_co_path = _os.path.join(_td, "mla_fp4_v2b.co")
_kernel_ready = False
_hip = _ctypes.CDLL("libamdhip64.so")
_hip.hipModuleLoadData.restype = _ctypes.c_int
_hip.hipModuleLoadData.argtypes = [_ctypes.c_void_p, _ctypes.c_void_p]
_hip.hipModuleGetFunction.restype = _ctypes.c_int
_hip.hipModuleGetFunction.argtypes = [_ctypes.c_void_p, _ctypes.c_void_p, _ctypes.c_char_p]
_hip.hipModuleLaunchKernel.restype = _ctypes.c_int
_hip.hipModuleLaunchKernel.argtypes = [
_ctypes.c_void_p,_ctypes.c_uint,_ctypes.c_uint,_ctypes.c_uint,
_ctypes.c_uint,_ctypes.c_uint,_ctypes.c_uint,
_ctypes.c_uint,_ctypes.c_void_p,_ctypes.c_void_p,_ctypes.c_void_p]
_hmod = _ctypes.c_void_p()
_hfn = _ctypes.c_void_p()
def _ensure_kernel():
global _kernel_ready
if _kernel_ready: return
s_src = _gz.decompress(_b64.b64decode(_ASM_GZ_B64)).decode()
with open(_s_path, "w") as f: f.write(s_src)
_clang = None
for p in ["/opt/rocm/llvm/bin/clang", "/opt/rocm/bin/clang"]:
if _os.path.exists(p): _clang = p; break
if not _clang:
import shutil; _clang = shutil.which("clang")
_sp.run([_clang, "-x", "assembler", "-target", "amdgcn-amd-amdhsa",
"-mcpu=gfx950", "-o", _co_path, _s_path], check=True, capture_output=True)
with open(_co_path, "rb") as f: co = f.read()
co_buf = (_ctypes.c_char * len(co)).from_buffer_copy(co)
_hip.hipModuleLoadData(_ctypes.byref(_hmod), co_buf)
_hip.hipModuleGetFunction(_ctypes.byref(_hfn), _hmod, b"mla_fp4_native_decode_kernel")
_kernel_ready = True
_cache = {}
def _get_nsplits(bs, kv_per_seq):
return max(1, min(kv_per_seq // 128, 16))
def _fp4_kernel(data):
q_bf16, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]; nq = config["num_heads"]; dv = config["v_head_dim"]
q_fp4, q_sc = dynamic_mxfp4_quant(q_bf16.view(-1, QK_HEAD_DIM))
kv_fp4, kv_sc = kv_data["mxfp4"]
total = kv_fp4.shape[0]
kv_per_seq = int((kv_indptr[-1] - kv_indptr[0]).item()) // bs
ns = _get_nsplits(bs, kv_per_seq)
q_c = q_fp4.view(bs, nq, 288).contiguous()
kv_c = kv_fp4.view(total, 288).contiguous()
qs_c = q_sc[:bs * nq].contiguous()
kvs_c = kv_sc[:total].contiguous()
key = (bs, ns, nq, dv)
if key not in _cache:
sd = torch.zeros(bs*ns, 1, nq, dv, dtype=torch.float32, device="cuda")
sl = torch.full((bs*ns, 1, nq, 1), -1e30, dtype=torch.float32, device="cuda")
o = torch.empty((bs, nq, dv), dtype=torch.bfloat16, device="cuda")
fl = torch.empty((bs, nq), dtype=torch.float32, device="cuda")
ri = torch.arange(0, bs+1, dtype=torch.int32, device="cuda") * ns
rpm = torch.arange(bs*ns, dtype=torch.int32, device="cuda")
rfm = torch.stack([torch.arange(bs, device="cuda", dtype=torch.int32),
torch.arange(1, bs+1, device="cuda", dtype=torch.int32)], dim=1)
_cache[key] = (sd, sl, o, fl, ri, rpm, rfm)
sd, sl, o, fl, ri, rpm, rfm = _cache[key]
so = sd.view(bs, ns, nq, dv); sl2 = sl.view(bs, ns, nq)
so.zero_(); sl2.fill_(-1e30)
chunk = ((kv_per_seq + ns - 1) // ns + 15) & ~15
_ensure_kernel()
args = bytearray(80)
P = lambda off, v: _struct.pack_into("Q", args, off, v)
F = lambda off, v: _struct.pack_into("f", args, off, v)
I = lambda off, v: _struct.pack_into("I", args, off, v)
P(0x00, q_c.data_ptr()); P(0x08, kv_c.data_ptr())
P(0x10, qs_c.data_ptr()); P(0x18, kvs_c.data_ptr())
P(0x20, kv_indptr.data_ptr()); P(0x28, so.data_ptr()); P(0x30, sl2.data_ptr())
F(0x38, SM_SCALE); I(0x3C, ns); I(0x40, chunk)
ab = (_ctypes.c_char * 80).from_buffer(args)
sz = _ctypes.c_size_t(80)
cfg = (_ctypes.c_void_p * 5)(
_ctypes.c_void_p(1), _ctypes.cast(ab, _ctypes.c_void_p),
_ctypes.c_void_p(2), _ctypes.cast(_ctypes.pointer(sz), _ctypes.c_void_p),
_ctypes.c_void_p(3))
_hip.hipModuleLaunchKernel(_hfn, bs, ns, 1, 256, 1, 1, 0,
_ctypes.c_void_p(0), _ctypes.c_void_p(0), cfg)
sl.copy_(sl2.view(bs*ns, 1, nq).unsqueeze(-1))
sd.copy_(so.view(bs*ns, 1, nq, dv))
aiter.mla_reduce_v1(sd, sl, ri, rfm, rpm, 1, o, fl)
return o
def custom_kernel(data: input_t) -> output_t:
return _fp8_fallback(data)
scrolls · 161 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