submission 703731
tanibe5390 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 41 lines, June 9 Researcher Reciprocity License v1.0.
_bss_merged_s4_test_asm_s4_v78.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-703731?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:4c7fbc62baf66dce5c79ec9acc6c5157714a0c1022deeb50d31b6523bd64b9a2
license declaredunknown
license concludedunknown
authorstanibe5390
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
…NK\nSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)\n_cache = {}\n\n\n@triton.jit\ndef _flash_decode_fp8_full_s4_exact(\n Q_ptr, KV_ptr, Mid_O, Mid_lse,\n stride_kv: tl.int64, kv_scale…num-warps = 4
…mulator split, full-head QK MFMA path,\n exact-path `num_warps=4`, bf16 `Mid_O`, and `REDUCE_BLOCK_V=512`, but lower\n the exact-path launch from `num_stages=3` to `n…online-softmax
…_scale\n\n row_max = tl.max(scores, axis=1)\n m_new = tl.maximum(m_i, row_max)\n alpha = tl.exp(m_i - m_new)\n l_i = l_i * alpha\n exp_scores = t…persistent-kernel
…": o,\n }\n return _cache[key]\n\n\ndef _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):\n ke…split-k
…_idx * NH + h_offs, lse_vals)\n\n\n@triton.jit\ndef _reduce_splitk(Mid_O, Mid_lse, O_ptr, NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr):\n…stages = 3
…LOCK_V=512`, but lower\n the exact-path launch from `num_stages=3` to `num_stages=2`.\nRationale: `v77` ruled out higher warp residency on the kept `s4` V-tiling branch.\n …tile-k = 128
…ze * kv_seq_len\n NUM_SPLITS = 8\n BLOCK_N = 128\n BLOCK_K = 128\n EXACT_V_BLOCK = 256\n kv_fp8, kv_scale = kv_data["fp8"]\n kv_flat = kv_fp8.view(total_kv, QK_HE…tile-n = 128
…total_kv = batch_size * kv_seq_len\n NUM_SPLITS = 8\n BLOCK_N = 128\n BLOCK_K = 128\n EXACT_V_BLOCK = 256\n kv_fp8, kv_scale = kv_data["fp8"]\n kv_flat = kv_fp8.v…Kernel source
_bss_merged_s4_test_asm_s4_v78.py41 lines
# Auto-generated by submit-single-shape.py
# Target shape: s4 = {"batchsize": 32, "kvseqlen": 8192, "qseqlen": 1}
import importlib.util
import sys
import os
from pathlib import Path
import tempfile
_TARGET_SOURCE = '"""\ntest_asm_s4_v78: lower exact-path pipeline depth on the kept s4 V-tiling head\nBase: test_asm_s4_v76.py\nDirection: CONTINUING full-head exact-path V-tiling\nTarget: s4 (batch=32, kv_seq_len=8192)\nChange: Keep `v76`\'s exact `V_BLOCK=256` accumulator split, full-head QK MFMA path,\n exact-path `num_warps=4`, bf16 `Mid_O`, and `REDUCE_BLOCK_V=512`, but lower\n the exact-path launch from `num_stages=3` to `num_stages=2`.\nRationale: `v77` ruled out higher warp residency on the kept `s4` V-tiling branch.\n The next direct launch-side lever is whether the lighter exact accumulator\n shape also prefers less software pipelining.\nScale: MODERATE\n"""\nimport torch\nimport triton\nimport triton.language as tl\nfrom task import input_t, output_t\n\nNUM_HEADS: tl.constexpr = 16\nKV_LORA_RANK = 512\nQK_ROPE_HEAD_DIM = 64\nQK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM\nV_HEAD_DIM = KV_LORA_RANK\nSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)\n_cache = {}\n\n\n@triton.jit\ndef _flash_decode_fp8_full_s4_exact(\n Q_ptr, KV_ptr, Mid_O, Mid_lse,\n stride_kv: tl.int64, kv_scale,\n sm_scale: tl.constexpr,\n QK_DIM: tl.constexpr, V_DIM: tl.constexpr,\n BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, NH: tl.constexpr, V_BLOCK: tl.constexpr,\n):\n batch_id = tl.program_id(0)\n split_id = tl.program_id(1)\n v_block_id = tl.program_id(2)\n out_idx = batch_id * 8 + split_id\n kv_start = batch_id * 8192 + split_id * 1024\n h_offs = tl.arange(0, NH)\n v_start = v_block_id * V_BLOCK\n v_offs = v_start + tl.arange(0, V_BLOCK)\n v_mask = v_offs < V_DIM\n q_base = batch_id * NH * QK_DIM\n acc = tl.zeros([NH, V_BLOCK], dtype=tl.float32)\n m_i = tl.full([NH], float("-inf"), dtype=tl.float32)\n l_i = tl.zeros([NH], dtype=tl.float32)\n score_scale = sm_scale * kv_scale\n\n for kv_offset in range(0, 1024, BLOCK_N):\n kv_pos = kv_start + kv_offset\n n_offs = tl.arange(0, BLOCK_N)\n\n scores = tl.zeros([NH, BLOCK_N], dtype=tl.float32)\n for k_start in range(0, QK_DIM, BLOCK_K):\n k_offs = tl.arange(0, BLOCK_K)\n k_mask = (k_start + k_offs) < QK_DIM\n q_tile = tl.load(\n Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :],\n mask=k_mask[None, :], other=0.0,\n )\n k_tile = tl.load(\n KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :],\n mask=k_mask[None, :], other=0.0,\n )\n scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))\n\n scores *= score_scale\n\n row_max = tl.max(scores, axis=1)\n m_new = tl.maximum(m_i, row_max)\n alpha = tl.exp(m_i - m_new)\n l_i = l_i * alpha\n exp_scores = tl.exp(scores - m_new[:, None])\n l_i += tl.sum(exp_scores, axis=1)\n acc = acc * alpha[:, None]\n\n v_tile = tl.load(\n KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :],\n mask=v_mask[None, :],\n other=0.0,\n )\n acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)\n\n m_i = m_new\n\n acc = (acc * kv_scale) / l_i[:, None]\n lse_vals = m_i + tl.log(l_i)\n tl.store(\n Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :],\n acc,\n mask=v_mask[None, :],\n )\n tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)\n\n\n@triton.jit\ndef _flash_decode_fp8_full(\n Q_ptr, KV_ptr, kv_indptr, Mid_O, Mid_lse,\n stride_kv: tl.int64, kv_scale,\n sm_scale: tl.constexpr,\n QK_DIM: tl.constexpr, V_DIM: tl.constexpr, NUM_SPLITS: tl.constexpr,\n BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, NH: tl.constexpr,\n):\n batch_id = tl.program_id(0)\n split_id = tl.program_id(1)\n kv_start = tl.load(kv_indptr + batch_id)\n kv_end = tl.load(kv_indptr + batch_id + 1)\n kv_len = kv_end - kv_start\n split_len = tl.cdiv(kv_len, NUM_SPLITS)\n my_start = kv_start + split_id * split_len\n my_end = tl.minimum(kv_start + (split_id + 1) * split_len, kv_end)\n actual_len = my_end - my_start\n out_idx = batch_id * NUM_SPLITS + split_id\n h_offs = tl.arange(0, NH)\n v_offs = tl.arange(0, V_DIM)\n if actual_len <= 0:\n tl.store(Mid_lse + out_idx * NH + h_offs, tl.full([NH], float("-inf"), dtype=tl.float32))\n return\n q_base = batch_id * NH * QK_DIM\n acc = tl.zeros([NH, V_DIM], dtype=tl.float32)\n m_i = tl.full([NH], float("-inf"), dtype=tl.float32)\n l_i = tl.zeros([NH], dtype=tl.float32)\n score_scale = sm_scale * kv_scale\n\n for kv_offset in range(0, actual_len, BLOCK_N):\n n_valid = tl.minimum(BLOCK_N, actual_len - kv_offset)\n kv_pos = my_start + kv_offset\n n_offs = tl.arange(0, BLOCK_N)\n n_mask = n_offs < n_valid\n\n scores = tl.zeros([NH, BLOCK_N], dtype=tl.float32)\n for k_start in range(0, QK_DIM, BLOCK_K):\n k_offs = tl.arange(0, BLOCK_K)\n k_mask = (k_start + k_offs) < QK_DIM\n q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)\n k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=n_mask[:, None] & k_mask[None, :], other=0.0)\n scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))\n\n scores *= score_scale\n scores = tl.where(n_mask[None, :], scores, float("-inf"))\n\n row_max = tl.max(scores, axis=1)\n m_new = tl.maximum(m_i, row_max)\n alpha = tl.exp(m_i - m_new)\n l_i = l_i * alpha\n exp_scores = tl.exp(scores - m_new[:, None])\n l_i += tl.sum(exp_scores, axis=1)\n acc = acc * alpha[:, None]\n\n v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=n_mask[:, None], other=0.0)\n acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)\n\n m_i = m_new\n\n acc = (acc * kv_scale) / l_i[:, None]\n lse_vals = m_i + tl.log(l_i)\n tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc)\n tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)\n\n\n@triton.jit\ndef _reduce_splitk(Mid_O, Mid_lse, O_ptr, NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr):\n batch_id = tl.program_id(0)\n head_id = tl.program_id(1)\n v_block = tl.program_id(2)\n v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)\n v_mask = v_offs < V_DIM\n m_final = tl.full([], float("-inf"), dtype=tl.float32)\n l_final = tl.zeros([], dtype=tl.float32)\n acc = tl.zeros([BLOCK_V], dtype=tl.float32)\n for s in range(NUM_SPLITS):\n idx = batch_id * NUM_SPLITS + s\n lse = tl.load(Mid_lse + idx * NH + head_id)\n is_valid = lse > float("-inf")\n if is_valid:\n m_new = tl.maximum(m_final, lse)\n alpha = tl.exp(m_final - m_new)\n beta = tl.exp(lse - m_new)\n partial = tl.load(Mid_O + idx * NH * V_DIM + head_id * V_DIM + v_offs, mask=v_mask, other=0.0).to(tl.float32)\n acc = acc * alpha + beta * partial\n l_final = l_final * alpha + beta\n m_final = m_new\n acc = acc / l_final\n tl.store(O_ptr + batch_id * NH * V_DIM + head_id * V_DIM + v_offs, acc.to(tl.bfloat16), mask=v_mask)\n\n\ndef _ensure_cache(batch_size, kv_seq_len, total_q, num_splits, kv_scale_tensor):\n key = ("s4_v78", batch_size, kv_seq_len, num_splits)\n if key in _cache:\n return _cache[key]\n _cache[key] = {\n "mid_o": torch.empty((batch_size * num_splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),\n "mid_lse": torch.empty((batch_size * num_splits, NUM_HEADS), dtype=torch.float32, device="cuda"),\n "o": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),\n "kv_scale_val": kv_scale_tensor.item(),\n }\n return _cache[key]\n\n\ndef custom_kernel(data: input_t) -> output_t:\n q, kv_data, qo_indptr, kv_indptr, config = data\n batch_size = config["batch_size"]\n kv_seq_len = config["kv_seq_len"]\n total_q = q.shape[0]\n total_kv = batch_size * kv_seq_len\n NUM_SPLITS = 8\n BLOCK_N = 128\n BLOCK_K = 128\n EXACT_V_BLOCK = 256\n kv_fp8, kv_scale = kv_data["fp8"]\n kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)\n c = _ensure_cache(batch_size, kv_seq_len, total_q, NUM_SPLITS, kv_scale)\n if batch_size == 32 and kv_seq_len == 8192:\n _flash_decode_fp8_full_s4_exact[(batch_size, NUM_SPLITS, triton.cdiv(V_HEAD_DIM, EXACT_V_BLOCK))](\n q, kv_flat, c["mid_o"], c["mid_lse"],\n QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,\n QK_HEAD_DIM, V_HEAD_DIM, BLOCK_N, BLOCK_K, 16, EXACT_V_BLOCK,\n num_warps=4, num_stages=2,\n )\n else:\n _flash_decode_fp8_full[(batch_size, NUM_SPLITS)](\n q, kv_flat, kv_indptr, c["mid_o"], c["mid_lse"],\n QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,\n QK_HEAD_DIM, V_HEAD_DIM, NUM_SPLITS, BLOCK_N, BLOCK_K, 16,\n num_warps=4, num_stages=2,\n )\n REDUCE_BLOCK_V = 512\n n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)\n _reduce_splitk[(batch_size, 16, n_v_blocks)](\n c["mid_o"], c["mid_lse"], c["o"],\n NUM_SPLITS, V_HEAD_DIM, 16, REDUCE_BLOCK_V, num_warps=4,\n )\n return c["o"]\n'
_REF_SOURCE = '"""\ntest_v143_s6_splits4: s6 splits 8→4 (continue reduce optimization pattern)\nBase: test.py (v142)\nDirection: NEW — s6 splits tuning\nTarget: s6 (64,8192) — reduce overhead with kv=8192\nChange: s6 splits 8→4. batch=64 × splits=4 = 256 programs (100% CU fill).\n Follows s5 splits reduction pattern (v142 +8.5%).\nRationale: v140 profile s8 reduce=3.3us. s6 with splits=8 has more reduce overhead.\nScale: INCREMENTAL\n"""\nimport torch\nimport aiter\nimport triton\nfrom task import input_t, output_t\n\nfrom aiter import dtypes as aiter_dtypes\nfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1\nfrom aiter.mla import get_meta_param, _fwd_kernel_stage2_asm\n\nNUM_HEADS = 16\nNUM_KV_HEADS = 1\nKV_LORA_RANK = 512\nQK_ROPE_HEAD_DIM = 64\nQK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM\nV_HEAD_DIM = KV_LORA_RANK\nSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)\nPAGE_SIZE = 1\nFP8_DTYPE = aiter_dtypes.fp8\n\n_cache = {}\n\n\ndef _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q):\n key = ("npbf16", batch_size, kv_seq_len)\n if key in _cache:\n return _cache[key]\n\n nq = NUM_HEADS\n total_kv = batch_size * kv_seq_len\n\n kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")\n kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")\n num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, torch.bfloat16)\n o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")\n logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")\n attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")\n\n _cache[key] = {\n "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,\n "num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,\n "logits": logits, "attn_lse": attn_lse, "o": o,\n }\n return _cache[key]\n\n\ndef _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):\n key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)\n if key in _cache:\n return _cache[key]\n\n max_q_len = 1\n nq, nkv = NUM_HEADS, NUM_KV_HEADS\n total_kv = batch_size * kv_seq_len\n\n kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")\n kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")\n\n info = get_mla_metadata_info_v1(\n batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,\n is_sparse=False, fast_mode=fast_mode,\n num_kv_splits=persistent_splits, intra_batch_mode=True,\n )\n work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]\n (work_metadata, work_indptr, work_info_set,\n reduce_indptr, reduce_final_map, reduce_partial_map) = work\n\n get_mla_metadata_v1(\n qo_indptr, kv_indptr, kv_last_page_len,\n nq // nkv, nkv, True,\n work_metadata, work_info_set, work_indptr,\n reduce_indptr, reduce_final_map, reduce_partial_map,\n page_size=PAGE_SIZE,\n kv_granularity=max(PAGE_SIZE, kv_gran),\n max_seqlen_qo=max_q_len,\n uni_seqlen_qo=max_q_len,\n fast_mode=fast_mode,\n max_split_per_batch=persistent_splits,\n intra_batch_mode=True,\n dtype_q=FP8_DTYPE,\n dtype_kv=FP8_DTYPE,\n )\n\n num_partials = reduce_partial_map.size(0)\n logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")\n attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")\n o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")\n q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")\n q_scale = torch.ones(1, dtype=torch.float32, device="cuda")\n\n _cache[key] = {\n "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,\n "work_metadata": work_metadata, "work_indptr": work_indptr,\n "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,\n "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,\n "logits": logits, "attn_lse": attn_lse, "o": o,\n "q_fp8": q_fp8, "q_scale": q_scale,\n "num_partials": num_partials,\n }\n return _cache[key]\n\n\ndef custom_kernel(data: input_t) -> output_t:\n q, kv_data, qo_indptr, kv_indptr, config = data\n\n batch_size = config["batch_size"]\n kv_seq_len = config["kv_seq_len"]\n total_q = q.shape[0]\n total_kv = batch_size * kv_seq_len\n\n # ---- Tier 1: batch<=4 -> bf16/bf16 non-persistent (s1, s2) ----\n if batch_size <= 4:\n kv_bf16 = kv_data["bf16"]\n kv_4d = kv_bf16.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)\n c = _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q)\n\n aiter.mla_decode_stage1_asm_fwd(\n q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,\n qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],\n c["num_kv_splits_indptr"],\n None, None, None,\n 1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,\n c["logits"], c["attn_lse"], c["o"],\n None, None,\n )\n\n Lv = V_HEAD_DIM\n BLOCK_DV = triton.next_power_of_2(Lv)\n _fwd_kernel_stage2_asm[(batch_size, NUM_HEADS)](\n c["logits"], c["attn_lse"], c["o"],\n qo_indptr, kv_indptr, c["num_kv_splits_indptr"],\n c["attn_lse"].stride(0), c["attn_lse"].stride(2), c["attn_lse"].stride(1),\n c["o"].stride(0), c["o"].stride(1),\n MAYBE_FINAL_OUT=True,\n BATCH_NUM=batch_size,\n BLOCK_DV=BLOCK_DV,\n Lv=Lv,\n mgc=64,\n num_warps=4,\n num_stages=2,\n waves_per_eu=4,\n )\n return c["o"]\n\n # ---- Tier 2: ALL fp8 shapes -> persistent (s3-s8) ----\n # Non-persistent fp8 was faster but fails leaderboard correctness (v77, v78).\n # Persistent + mla_reduce_v1 is the only leaderboard-safe fp8 path.\n else:\n kv_buffer_fp8, kv_scale = kv_data["fp8"]\n kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)\n\n # Per-shape split tuning\n if total_kv >= 1000000:\n # s8 (256, 8192) -- splits=4 (from v77)\n splits, fast_mode = 4, False\n elif total_kv >= 300000:\n # s6 (64, 8192) -- splits=4 (from 8, 64*4=256 programs = 100% CU fill)\n splits, fast_mode = 4, False\n elif batch_size >= 256:\n # s7 (256, 1024) -- splits=4 with kv_gran=64 (v131 LB-safe config)\n # splits=1+kv_gran=64 FAILED LB in v136. splits=4 gives 1024 programs.\n splits, fast_mode = 4, False\n elif batch_size >= 64:\n # s5 (64, 1024) -- splits=2 (from 4, reduce=7us → ~3.5us, 128 programs = 50% CU)\n splits, fast_mode = 2, False\n else:\n # s3 (32, 1024) and s4 (32, 8192)\n if kv_seq_len <= 1024:\n splits, fast_mode = 4, True # s3: reduced from 8 to 4\n else:\n splits, fast_mode = 32, True # s4\n\n # Use kv_granularity=64 for ALL persistent shapes (matches v131 LB-safe config)\n # v131 (kv_gran=64 all) PASSED LB at 58.0us. v137/v138 (kv_gran=16 for s3/s5)\n # FAILED LB on s3. kv_gran=64 is required for LB correctness.\n kv_gran = 64\n c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)\n\n # Fast FP8 quant: copy_ cast (scale=1.0) -- from v63\n q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)\n c["q_fp8"].copy_(q_2d)\n\n aiter.mla_decode_stage1_asm_fwd(\n c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,\n qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],\n None, c["work_metadata"], c["work_indptr"], c["work_info_set"],\n 1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,\n c["logits"], c["attn_lse"], c["o"],\n c["q_scale"], kv_scale,\n )\n\n aiter.mla_reduce_v1(\n c["logits"], c["attn_lse"],\n c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],\n 1, c["o"], None,\n )\n return c["o"]\n'
def _load_module(name, source):
tmpdir = tempfile.mkdtemp()
path = os.path.join(tmpdir, name + '.py')
open(path, 'w').write(source)
spec = importlib.util.spec_from_file_location(name, path)
mod = importlib.util.module_from_spec(spec)
sys.modules[name] = mod
spec.loader.exec_module(mod)
return mod
# LAZY loading: modules are loaded on first use, not at import time.
# This prevents aiter (reference) from polluting Triton (target) state.
_target_mod = None
_ref_mod = None
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
global _target_mod, _ref_mod
_q, _cfg = data[0], data[4]
if (_q.shape[0] == 32 and _cfg["kv_seq_len"] == 8192):
if _target_mod is None:
_target_mod = _load_module('_bss_target', _TARGET_SOURCE)
return _target_mod.custom_kernel(data)
else:
if _ref_mod is None:
_ref_mod = _load_module('_bss_ref', _REF_SOURCE)
return _ref_mod.custom_kernel(data)
scrolls · 41 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