Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
117.6µs
#486 of 766
2026-03-30

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