submission 745087
HorizonLiang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1304 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-745087?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:3bcfcb2de30cc10ca2b9b49ef7cd55971cee7336194599ad396882c60285b419
license declaredunknown
license concludedunknown
authorsHorizonLiang
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
os.environ.setdefault("OPTIMIZE_EPILOGUE", "1")num-warps = 4
TRITON_REDUCE_NUM_WARPS = 4stages = 1
TRITON_REDUCE_NUM_STAGES = 1Kernel source
submission.py1304 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""Submission version 033."""
from __future__ import annotations
import os
import sys
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("OPTIMIZE_EPILOGUE", "1")
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import (
dtypes as aiter_dtypes,
dynamic_per_tensor_quant,
get_mla_metadata_info_v1,
get_mla_metadata_v1,
mla_decode_stage1_asm_fwd,
)
from aiter.mla import mla_decode_fwd
from aiter.ops.triton.utils._triton.pid_preprocessing import remap_xcd
PAGE_SIZE = 1
SMALLKV_PAGE_SIZE = 2
LONGKV_PAGE_SIZE = 8
DEFAULT_KV_GRANULARITY = 16
DEFAULT_SPLIT_PER_BATCH = 32
FAST_MODE = True
INTRA_BATCH_MODE = False
FP8_DTYPE = aiter_dtypes.fp8
ZERO_REDUCE_PARTIAL_SLOTS = 1
NUM_HEADS = 16
V_HEAD_DIM = 512
TRITON_REDUCE_MAX_SPLITS = 6
TRITON_REDUCE_BLOCK_H = 8
TRITON_REDUCE_BLOCK_D = 256
TRITON_REDUCE_NUM_WARPS = 4
TRITON_REDUCE_NUM_STAGES = 1
TRITON_NUM_XCDS = 8
_META_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
_PAGED_META_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
_KV_INDEX_CACHE: dict[tuple, torch.Tensor] = {}
_KV_MISC_CACHE: dict[tuple, torch.Tensor] = {}
_OUTPUT_CACHE: dict[tuple, torch.Tensor] = {}
_PARTIAL_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
_Q_QUANT_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
_SEQ_INDEX_CACHE: dict[tuple, torch.Tensor] = {}
_DENSE_FASTPATH_PROBE_SEEN: set[tuple[int, int, str]] = set()
@triton.jit
def _mla_reduce_q1_longkv_kernel(
partial_output_ptr,
partial_lse_ptr,
reduce_indptr_ptr,
reduce_q_indices_ptr,
reduce_partial_map_ptr,
output_ptr,
stride_po0,
stride_po1,
stride_po2,
stride_pl0,
stride_pl1,
stride_o0,
stride_o1,
stride_o2,
grid_tiles,
tiles_per_group,
BLOCK_H: tl.constexpr,
BLOCK_D: tl.constexpr,
MAX_SPLITS: tl.constexpr,
NUM_XCDS: tl.constexpr,
):
pid = tl.program_id(0)
pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)
pid_group = pid // tiles_per_group
pid_tile = pid % tiles_per_group
num_pid_d = tl.cdiv(512, BLOCK_D)
pid_h = pid_tile // num_pid_d
pid_d = pid_tile % num_pid_d
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
mask_h = offs_h < 16
mask_d = offs_d < 512
reduce_start = tl.load(reduce_indptr_ptr + pid_group)
reduce_end = tl.load(reduce_indptr_ptr + pid_group + 1)
split_count = reduce_end - reduce_start
q_row = tl.load(reduce_q_indices_ptr + pid_group)
max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)
for split_idx in tl.static_range(MAX_SPLITS):
active = split_idx < split_count
partial_idx = tl.load(
reduce_partial_map_ptr + reduce_start + split_idx,
mask=active,
other=0,
)
lse = tl.load(
partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
mask=active & mask_h,
other=-float("inf"),
)
new_max_lse = tl.maximum(max_lse, lse)
old_scale = tl.exp(max_lse - new_max_lse)
new_scale = tl.exp(lse - new_max_lse)
partial = tl.load(
partial_output_ptr
+ partial_idx * stride_po0
+ offs_h[:, None] * stride_po1
+ offs_d[None, :] * stride_po2,
mask=(active & mask_h)[:, None] & mask_d[None, :],
other=0.0,
)
acc = acc * old_scale[:, None] + partial * new_scale[:, None]
sum_e_lse = sum_e_lse * old_scale + new_scale
max_lse = new_max_lse
denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
out = acc / denom[:, None]
tl.store(
output_ptr
+ q_row * stride_o0
+ offs_h[:, None] * stride_o1
+ offs_d[None, :] * stride_o2,
out.to(output_ptr.type.element_ty),
mask=mask_h[:, None] & mask_d[None, :],
)
@triton.jit
def _mla_reduce_q1_dense_exact_kernel(
partial_output_ptr,
partial_lse_ptr,
output_ptr,
stride_po0,
stride_po1,
stride_po2,
stride_pl0,
stride_pl1,
stride_o0,
stride_o1,
stride_o2,
grid_tiles,
tiles_per_group,
BLOCK_H: tl.constexpr,
BLOCK_D: tl.constexpr,
EXACT_SPLITS: tl.constexpr,
NUM_XCDS: tl.constexpr,
):
pid = tl.program_id(0)
pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)
pid_group = pid // tiles_per_group
pid_tile = pid % tiles_per_group
num_pid_d = tl.cdiv(512, BLOCK_D)
pid_h = pid_tile // num_pid_d
pid_d = pid_tile % num_pid_d
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
mask_h = offs_h < 16
mask_d = offs_d < 512
max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)
base_partial_idx = pid_group * EXACT_SPLITS
q_row = pid_group
for split_idx in tl.static_range(EXACT_SPLITS):
partial_idx = base_partial_idx + split_idx
lse = tl.load(
partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
mask=mask_h,
other=-float("inf"),
)
new_max_lse = tl.maximum(max_lse, lse)
old_scale = tl.exp(max_lse - new_max_lse)
new_scale = tl.exp(lse - new_max_lse)
partial = tl.load(
partial_output_ptr
+ partial_idx * stride_po0
+ offs_h[:, None] * stride_po1
+ offs_d[None, :] * stride_po2,
mask=mask_h[:, None] & mask_d[None, :],
other=0.0,
)
acc = acc * old_scale[:, None] + partial * new_scale[:, None]
sum_e_lse = sum_e_lse * old_scale + new_scale
max_lse = new_max_lse
denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
out = acc / denom[:, None]
tl.store(
output_ptr
+ q_row * stride_o0
+ offs_h[:, None] * stride_o1
+ offs_d[None, :] * stride_o2,
out.to(output_ptr.type.element_ty),
mask=mask_h[:, None] & mask_d[None, :],
)
@triton.jit
def _mla_reduce_q1_dense_upto_kernel(
partial_output_ptr,
partial_lse_ptr,
reduce_indptr_ptr,
output_ptr,
stride_po0,
stride_po1,
stride_po2,
stride_pl0,
stride_pl1,
stride_o0,
stride_o1,
stride_o2,
grid_tiles,
tiles_per_group,
BLOCK_H: tl.constexpr,
BLOCK_D: tl.constexpr,
MAX_SPLITS: tl.constexpr,
NUM_XCDS: tl.constexpr,
):
pid = tl.program_id(0)
pid = remap_xcd(pid, grid_tiles, NUM_XCDS=NUM_XCDS)
pid_group = pid // tiles_per_group
pid_tile = pid % tiles_per_group
num_pid_d = tl.cdiv(512, BLOCK_D)
pid_h = pid_tile // num_pid_d
pid_d = pid_tile % num_pid_d
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D)
mask_h = offs_h < 16
mask_d = offs_d < 512
reduce_start = tl.load(reduce_indptr_ptr + pid_group)
reduce_end = tl.load(reduce_indptr_ptr + pid_group + 1)
split_count = reduce_end - reduce_start
q_row = pid_group
max_lse = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32)
sum_e_lse = tl.zeros((BLOCK_H,), dtype=tl.float32)
acc = tl.zeros((BLOCK_H, BLOCK_D), dtype=tl.float32)
for split_idx in tl.static_range(MAX_SPLITS):
active = split_idx < split_count
partial_idx = reduce_start + split_idx
lse = tl.load(
partial_lse_ptr + partial_idx * stride_pl0 + offs_h * stride_pl1,
mask=active & mask_h,
other=-float("inf"),
)
new_max_lse = tl.maximum(max_lse, lse)
old_scale = tl.exp(max_lse - new_max_lse)
new_scale = tl.exp(lse - new_max_lse)
partial = tl.load(
partial_output_ptr
+ partial_idx * stride_po0
+ offs_h[:, None] * stride_po1
+ offs_d[None, :] * stride_po2,
mask=(active & mask_h)[:, None] & mask_d[None, :],
other=0.0,
)
acc = acc * old_scale[:, None] + partial * new_scale[:, None]
sum_e_lse = sum_e_lse * old_scale + new_scale
max_lse = new_max_lse
denom = tl.where(sum_e_lse > 0.0, sum_e_lse, 1.0)
out = acc / denom[:, None]
tl.store(
output_ptr
+ q_row * stride_o0
+ offs_h[:, None] * stride_o1
+ offs_d[None, :] * stride_o2,
out.to(output_ptr.type.element_ty),
mask=mask_h[:, None] & mask_d[None, :],
)
def _device_key(device: torch.device) -> tuple[str, int | None]:
return device.type, device.index
def _get_cached_seq_indices(length: int, device: torch.device) -> torch.Tensor:
key = (_device_key(device), length)
cached = _SEQ_INDEX_CACHE.get(key)
if cached is None:
cached = torch.arange(length, dtype=torch.int32, device=device)
_SEQ_INDEX_CACHE[key] = cached
return cached
def _probe_dense_fastpath_once(batch_size: int, kv_seq_len: int, message: str) -> None:
key = (batch_size, kv_seq_len, message)
if key in _DENSE_FASTPATH_PROBE_SEEN:
return
print(f"[probe] dense_tail_reduce batch={batch_size} kv={kv_seq_len} {message}", file=sys.stderr, flush=True)
_DENSE_FASTPATH_PROBE_SEEN.add(key)
def _use_bf16_q(kv_seq_len: int) -> bool:
return kv_seq_len == 1024
def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
split_table = {
(4, 8192): 24,
(32, 8192): 6,
(64, 1024): 16,
(64, 8192): 6,
(256, 1024): 1,
(256, 8192): 6,
}
return split_table.get((batch_size, kv_seq_len), DEFAULT_SPLIT_PER_BATCH)
def _choose_kv_granularity(batch_size: int, kv_seq_len: int) -> int:
if batch_size == 4 and kv_seq_len == 8192:
return 32
if batch_size == 64 and kv_seq_len == 1024:
return 8
if batch_size == 256 and kv_seq_len == 8192:
return 8
if kv_seq_len == 8192 and batch_size >= 32:
return 16
return DEFAULT_KV_GRANULARITY
def _use_paged_fast_mode(batch_size: int, kv_seq_len: int, page_size: int) -> bool:
return FAST_MODE
def _get_alt_page_size(batch_size: int, kv_seq_len: int) -> int | None:
if kv_seq_len == 1024 and batch_size in (4, 32, 64, 256):
return SMALLKV_PAGE_SIZE
if kv_seq_len == 8192 and batch_size >= 4:
return LONGKV_PAGE_SIZE
return None
def _get_smallkv_dense_exact_splits(
batch_size: int,
kv_seq_len: int,
page_size: int | None,
num_kv_splits: int,
) -> int | None:
if kv_seq_len != 1024 or page_size != SMALLKV_PAGE_SIZE:
return None
exact_split_table = {
(4, 32): 16,
(32, 32): 7,
(64, 16): 4,
}
return exact_split_table.get((batch_size, num_kv_splits))
def _use_stage1_direct(batch_size: int, kv_seq_len: int, num_kv_splits: int) -> bool:
return batch_size == 256 and kv_seq_len == 1024 and num_kv_splits == 1
def _use_paged_stage1_only(
batch_size: int,
kv_seq_len: int,
page_size: int | None,
num_kv_splits: int,
) -> bool:
if batch_size != 256:
return False
if kv_seq_len == 1024:
return page_size == SMALLKV_PAGE_SIZE and num_kv_splits == 1
return kv_seq_len == 8192 and page_size == LONGKV_PAGE_SIZE
def _use_triton_longkv_reduce(
batch_size: int,
kv_seq_len: int,
page_size: int | None,
num_kv_splits: int,
) -> bool:
return (
batch_size in (32, 64)
and kv_seq_len == 8192
and page_size == LONGKV_PAGE_SIZE
and num_kv_splits == TRITON_REDUCE_MAX_SPLITS
)
def _choose_stage1_partial_slots(
batch_size: int,
kv_seq_len: int,
page_size: int,
partial_tiles: int,
) -> int:
if batch_size == 256 and kv_seq_len == 8192 and page_size == LONGKV_PAGE_SIZE:
return ZERO_REDUCE_PARTIAL_SLOTS
return max(partial_tiles, 1)
def _get_used_group_count(indptr: torch.Tensor) -> int:
if indptr.numel() <= 1:
return 0
diffs = indptr[1:] - indptr[:-1]
nonzero = torch.nonzero(diffs, as_tuple=False)
if nonzero.numel() == 0:
return 0
return int(nonzero[-1, 0].item()) + 1
def _get_reduce_compact_metadata(metadata: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
cached = metadata.get("_reduce_compact_metadata")
if cached is not None:
return cached
reduce_indptr = metadata["reduce_indptr"]
used_reduce_groups = _get_used_group_count(reduce_indptr)
compact_reduce_indptr = reduce_indptr[: used_reduce_groups + 1]
num_partial = (
int(compact_reduce_indptr[-1].item())
if compact_reduce_indptr.numel() != 0
else 0
)
reduce_final_map = metadata["reduce_final_map"][:used_reduce_groups]
reduce_q_indices = reduce_final_map[:, 0].contiguous()
compact = dict(metadata)
compact.update(
{
"reduce_indptr": compact_reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_q_indices": reduce_q_indices,
"reduce_partial_map": metadata["reduce_partial_map"][: max(num_partial, 1)],
}
)
metadata["_reduce_compact_metadata"] = compact
return compact
def _get_dense_reduce_layout(metadata: dict[str, torch.Tensor]) -> dict[str, int | bool]:
cached = metadata.get("_dense_reduce_layout")
if cached is not None:
return cached
reduce_indptr = metadata["reduce_indptr"]
reduce_q_indices = metadata["reduce_q_indices"]
reduce_partial_map = metadata["reduce_partial_map"]
num_groups = int(reduce_q_indices.numel())
num_partial = int(reduce_partial_map.numel())
q_identity = torch.equal(
reduce_q_indices, _get_cached_seq_indices(num_groups, reduce_q_indices.device)
)
partial_identity = torch.equal(
reduce_partial_map,
_get_cached_seq_indices(num_partial, reduce_partial_map.device),
)
if num_groups > 0:
split_counts = reduce_indptr[1:] - reduce_indptr[:-1]
split_min = int(split_counts.min().item())
split_max = int(split_counts.max().item())
else:
split_min = 0
split_max = 0
cached = {
"num_groups": num_groups,
"num_partial": num_partial,
"q_identity": q_identity,
"partial_identity": partial_identity,
"split_min": split_min,
"split_max": split_max,
}
metadata["_dense_reduce_layout"] = cached
return cached
def _run_triton_longkv_reduce_q1(
partial_output: torch.Tensor,
partial_lse: torch.Tensor,
reduce_indptr: torch.Tensor,
reduce_q_indices: torch.Tensor,
reduce_partial_map: torch.Tensor,
output: torch.Tensor,
) -> None:
num_groups = int(reduce_q_indices.numel())
if num_groups == 0:
return
partial_output_view = partial_output[:, 0]
partial_lse_view = partial_lse[:, 0, :, 0]
tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
)
grid_tiles = num_groups * tiles_per_group
_mla_reduce_q1_longkv_kernel[(grid_tiles,)](
partial_output_view,
partial_lse_view,
reduce_indptr,
reduce_q_indices,
reduce_partial_map,
output,
partial_output_view.stride(0),
partial_output_view.stride(1),
partial_output_view.stride(2),
partial_lse_view.stride(0),
partial_lse_view.stride(1),
output.stride(0),
output.stride(1),
output.stride(2),
grid_tiles,
tiles_per_group,
BLOCK_H=TRITON_REDUCE_BLOCK_H,
BLOCK_D=TRITON_REDUCE_BLOCK_D,
MAX_SPLITS=TRITON_REDUCE_MAX_SPLITS,
NUM_XCDS=TRITON_NUM_XCDS,
num_warps=TRITON_REDUCE_NUM_WARPS,
num_stages=TRITON_REDUCE_NUM_STAGES,
)
def _run_dense_longkv_reduce_q1_exact(
partial_output: torch.Tensor,
partial_lse: torch.Tensor,
output: torch.Tensor,
exact_splits: int,
) -> None:
num_groups = output.shape[0]
if num_groups == 0:
return
partial_output_view = partial_output[:, 0]
partial_lse_view = partial_lse[:, 0, :, 0]
tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
)
grid_tiles = num_groups * tiles_per_group
_mla_reduce_q1_dense_exact_kernel[(grid_tiles,)](
partial_output_view,
partial_lse_view,
output,
partial_output_view.stride(0),
partial_output_view.stride(1),
partial_output_view.stride(2),
partial_lse_view.stride(0),
partial_lse_view.stride(1),
output.stride(0),
output.stride(1),
output.stride(2),
grid_tiles,
tiles_per_group,
BLOCK_H=TRITON_REDUCE_BLOCK_H,
BLOCK_D=TRITON_REDUCE_BLOCK_D,
EXACT_SPLITS=exact_splits,
NUM_XCDS=TRITON_NUM_XCDS,
num_warps=TRITON_REDUCE_NUM_WARPS,
num_stages=TRITON_REDUCE_NUM_STAGES,
)
def _run_dense_longkv_reduce_q1_upto(
partial_output: torch.Tensor,
partial_lse: torch.Tensor,
reduce_indptr: torch.Tensor,
output: torch.Tensor,
max_splits: int,
) -> None:
num_groups = output.shape[0]
if num_groups == 0:
return
partial_output_view = partial_output[:, 0]
partial_lse_view = partial_lse[:, 0, :, 0]
tiles_per_group = triton.cdiv(NUM_HEADS, TRITON_REDUCE_BLOCK_H) * triton.cdiv(
V_HEAD_DIM, TRITON_REDUCE_BLOCK_D
)
grid_tiles = num_groups * tiles_per_group
_mla_reduce_q1_dense_upto_kernel[(grid_tiles,)](
partial_output_view,
partial_lse_view,
reduce_indptr,
output,
partial_output_view.stride(0),
partial_output_view.stride(1),
partial_output_view.stride(2),
partial_lse_view.stride(0),
partial_lse_view.stride(1),
output.stride(0),
output.stride(1),
output.stride(2),
grid_tiles,
tiles_per_group,
BLOCK_H=TRITON_REDUCE_BLOCK_H,
BLOCK_D=TRITON_REDUCE_BLOCK_D,
MAX_SPLITS=max_splits,
NUM_XCDS=TRITON_NUM_XCDS,
num_warps=TRITON_REDUCE_NUM_WARPS,
num_stages=TRITON_REDUCE_NUM_STAGES,
)
def _get_cached_kv_indices(total_kv: int, device: torch.device) -> torch.Tensor:
key = (_device_key(device), total_kv)
cached = _KV_INDEX_CACHE.get(key)
if cached is None:
cached = torch.arange(total_kv, dtype=torch.int32, device=device)
_KV_INDEX_CACHE[key] = cached
return cached
def _get_kv_last_page_len(
batch_size: int, kv_seq_len: int, device: torch.device
) -> torch.Tensor:
key = (_device_key(device), batch_size, kv_seq_len)
cached = _KV_MISC_CACHE.get(key)
if cached is None:
cached = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
_KV_MISC_CACHE[key] = cached
return cached
def _get_cached_output(
total_q: int,
num_heads: int,
v_head_dim: int,
device: torch.device,
batch_size: int,
kv_seq_len: int,
) -> torch.Tensor:
key = (_device_key(device), total_q, num_heads, v_head_dim, batch_size, kv_seq_len)
cached = _OUTPUT_CACHE.get(key)
if cached is None:
cached = torch.empty(
(total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=device
)
_OUTPUT_CACHE[key] = cached
return cached
def _get_cached_partials(
partial_tiles: int,
q_seq_len: int,
num_heads: int,
v_head_dim: int,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
slots = max(partial_tiles, 1)
key = (_device_key(device), slots, q_seq_len, num_heads, v_head_dim)
cached = _PARTIAL_CACHE.get(key)
if cached is None:
cached = (
torch.empty(
(slots * q_seq_len, 1, num_heads, v_head_dim),
dtype=torch.float32,
device=device,
),
torch.empty(
(slots * q_seq_len, 1, num_heads, 1),
dtype=torch.float32,
device=device,
),
)
_PARTIAL_CACHE[key] = cached
return cached
def _get_cached_q_quant(
total_q: int, num_heads: int, qk_head_dim: int, device: torch.device
) -> tuple[torch.Tensor, torch.Tensor]:
key = (_device_key(device), total_q, num_heads, qk_head_dim)
cached = _Q_QUANT_CACHE.get(key)
if cached is None:
cached = (
torch.empty(
(total_q, num_heads, qk_head_dim), dtype=FP8_DTYPE, device=device
),
torch.empty((1,), dtype=torch.float32, device=device),
)
_Q_QUANT_CACHE[key] = cached
return cached
def _quantize_q_fp8(
q: torch.Tensor, num_heads: int, qk_head_dim: int
) -> tuple[torch.Tensor, torch.Tensor]:
q_fp8, q_scale = _get_cached_q_quant(q.shape[0], num_heads, qk_head_dim, q.device)
dynamic_per_tensor_quant(
q_fp8.view(-1, qk_head_dim),
q.view(-1, qk_head_dim),
q_scale,
)
return q_fp8, q_scale
def _get_cached_metadata(
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
batch_size: int,
q_seq_len: int,
kv_seq_len: int,
num_heads: int,
num_kv_heads: int,
num_kv_splits: int,
kv_granularity: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
device = qo_indptr.device
key = (
_device_key(device),
batch_size,
q_seq_len,
kv_seq_len,
num_heads,
num_kv_heads,
num_kv_splits,
kv_granularity,
str(q_dtype),
str(kv_dtype),
)
cached = _META_CACHE.get(key)
if cached is not None:
return cached
kv_last_page_len = _get_kv_last_page_len(batch_size, kv_seq_len, device)
info = get_mla_metadata_info_v1(
batch_size,
q_seq_len,
num_heads,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=FAST_MODE,
num_kv_splits=num_kv_splits,
intra_batch_mode=INTRA_BATCH_MODE,
)
work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr,
kv_last_page_len,
num_heads // num_kv_heads,
num_kv_heads,
True,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=kv_granularity,
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=FAST_MODE,
max_split_per_batch=num_kv_splits,
intra_batch_mode=INTRA_BATCH_MODE,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
cached = {
"kv_last_page_len": kv_last_page_len,
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
_META_CACHE[key] = cached
return cached
def _get_cached_paged_kv_indices(
batch_size: int, kv_seq_len: int, page_size: int, device: torch.device
) -> torch.Tensor:
pages_per_batch = kv_seq_len // page_size
total_pages = batch_size * pages_per_batch
return _get_cached_kv_indices(total_pages, device)
def _get_cached_paged_kv_indptr(
batch_size: int, kv_seq_len: int, page_size: int, device: torch.device
) -> torch.Tensor:
pages_per_batch = kv_seq_len // page_size
key = (_device_key(device), "paged_kv_indptr", batch_size, kv_seq_len, page_size)
cached = _KV_MISC_CACHE.get(key)
if cached is None:
cached = torch.arange(
0,
(batch_size + 1) * pages_per_batch,
pages_per_batch,
dtype=torch.int32,
device=device,
)
_KV_MISC_CACHE[key] = cached
return cached
def _get_cached_paged_last_page_len(
batch_size: int, page_size: int, device: torch.device
) -> torch.Tensor:
key = (_device_key(device), "paged_last_page_len", batch_size, page_size)
cached = _KV_MISC_CACHE.get(key)
if cached is None:
cached = torch.full((batch_size,), page_size, dtype=torch.int32, device=device)
_KV_MISC_CACHE[key] = cached
return cached
def _get_cached_paged_metadata(
qo_indptr: torch.Tensor,
batch_size: int,
q_seq_len: int,
kv_seq_len: int,
page_size: int,
num_heads: int,
num_kv_heads: int,
num_kv_splits: int,
kv_granularity: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
device = qo_indptr.device
paged_fast_mode = _use_paged_fast_mode(batch_size, kv_seq_len, page_size)
key = (
_device_key(device),
"paged",
batch_size,
q_seq_len,
kv_seq_len,
page_size,
num_heads,
num_kv_heads,
num_kv_splits,
kv_granularity,
paged_fast_mode,
str(q_dtype),
str(kv_dtype),
)
cached = _PAGED_META_CACHE.get(key)
if cached is not None:
return cached
paged_kv_indptr = _get_cached_paged_kv_indptr(
batch_size, kv_seq_len, page_size, device
)
paged_last_page_len = _get_cached_paged_last_page_len(batch_size, page_size, device)
info = get_mla_metadata_info_v1(
batch_size,
q_seq_len,
num_heads,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=paged_fast_mode,
num_kv_splits=num_kv_splits,
intra_batch_mode=INTRA_BATCH_MODE,
)
work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = work
get_mla_metadata_v1(
qo_indptr,
paged_kv_indptr,
paged_last_page_len,
num_heads // num_kv_heads,
num_kv_heads,
True,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
page_size=page_size,
kv_granularity=kv_granularity,
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=paged_fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=INTRA_BATCH_MODE,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
cached = {
"kv_indptr": paged_kv_indptr,
"kv_last_page_len": paged_last_page_len,
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
_PAGED_META_CACHE[key] = cached
return cached
def _view_kv_paged(
kv_buffer_fp8: torch.Tensor,
batch_size: int,
kv_seq_len: int,
num_kv_heads: int,
page_size: int,
) -> torch.Tensor:
return kv_buffer_fp8.view(
batch_size,
kv_seq_len // page_size,
page_size,
num_kv_heads,
kv_buffer_fp8.shape[-1],
).reshape(-1, page_size, num_kv_heads, kv_buffer_fp8.shape[-1])
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
num_heads = config["num_heads"]
num_kv_heads = config["num_kv_heads"]
qk_head_dim = config["qk_head_dim"]
v_head_dim = config["v_head_dim"]
sm_scale = config["sm_scale"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)
kv_granularity = _choose_kv_granularity(batch_size, kv_seq_len)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
if _use_bf16_q(kv_seq_len):
q_input = q
q_scale = None
else:
q_input, q_scale = _quantize_q_fp8(q, num_heads, qk_head_dim)
output = _get_cached_output(
q.shape[0],
num_heads,
v_head_dim,
q.device,
batch_size,
kv_seq_len,
)
alt_page_size = _get_alt_page_size(batch_size, kv_seq_len)
if alt_page_size is not None:
paged_kv = _view_kv_paged(
kv_buffer_fp8,
batch_size,
kv_seq_len,
num_kv_heads,
alt_page_size,
)
paged_metadata = _get_cached_paged_metadata(
qo_indptr,
batch_size,
q_seq_len,
kv_seq_len,
alt_page_size,
num_heads,
num_kv_heads,
num_kv_splits,
kv_granularity,
q_input.dtype,
kv_buffer_fp8.dtype,
)
paged_kv_indices = _get_cached_paged_kv_indices(
batch_size, kv_seq_len, alt_page_size, q.device
)
if _use_paged_stage1_only(batch_size, kv_seq_len, alt_page_size, num_kv_splits):
partial_output, partial_lse = _get_cached_partials(
_choose_stage1_partial_slots(
batch_size,
kv_seq_len,
alt_page_size,
int(paged_metadata["reduce_partial_map"].numel()),
),
q_seq_len,
num_heads,
v_head_dim,
q.device,
)
mla_decode_stage1_asm_fwd(
q_input.view(-1, num_heads, qk_head_dim),
paged_kv,
qo_indptr,
paged_metadata["kv_indptr"],
paged_kv_indices,
paged_metadata["kv_last_page_len"],
None,
paged_metadata["work_meta_data"],
paged_metadata["work_indptr"],
paged_metadata["work_info_set"],
q_seq_len,
alt_page_size,
num_kv_heads,
sm_scale,
partial_output,
partial_lse,
output,
q_scale,
kv_scale_fp8,
)
return output
if _use_triton_longkv_reduce(batch_size, kv_seq_len, alt_page_size, num_kv_splits):
compact_paged_metadata = _get_reduce_compact_metadata(paged_metadata)
partial_output, partial_lse = _get_cached_partials(
int(compact_paged_metadata["reduce_partial_map"].numel()),
q_seq_len,
num_heads,
v_head_dim,
q.device,
)
mla_decode_stage1_asm_fwd(
q_input.view(-1, num_heads, qk_head_dim),
paged_kv,
qo_indptr,
paged_metadata["kv_indptr"],
paged_kv_indices,
paged_metadata["kv_last_page_len"],
None,
paged_metadata["work_meta_data"],
paged_metadata["work_indptr"],
paged_metadata["work_info_set"],
q_seq_len,
alt_page_size,
num_kv_heads,
sm_scale,
partial_output,
partial_lse,
output,
q_scale,
kv_scale_fp8,
)
dense_layout = _get_dense_reduce_layout(compact_paged_metadata)
use_dense_exact6 = (
batch_size == 32
and dense_layout["q_identity"]
and dense_layout["partial_identity"]
and dense_layout["num_groups"] == batch_size
and dense_layout["num_partial"] == batch_size * 6
and dense_layout["split_min"] == 6
and dense_layout["split_max"] == 6
)
use_dense_upto5 = (
batch_size == 64
and dense_layout["q_identity"]
and dense_layout["partial_identity"]
and dense_layout["num_groups"] == batch_size
and dense_layout["num_partial"] == 304
and dense_layout["split_min"] == 4
and dense_layout["split_max"] == 5
)
if use_dense_exact6:
_probe_dense_fastpath_once(
batch_size,
kv_seq_len,
"dense_exact6 enabled",
)
_run_dense_longkv_reduce_q1_exact(
partial_output,
partial_lse,
output,
exact_splits=6,
)
elif use_dense_upto5:
_probe_dense_fastpath_once(
batch_size,
kv_seq_len,
"dense_upto5 enabled",
)
_run_dense_longkv_reduce_q1_upto(
partial_output,
partial_lse,
compact_paged_metadata["reduce_indptr"],
output,
max_splits=5,
)
else:
_probe_dense_fastpath_once(
batch_size,
kv_seq_len,
"fallback generic",
)
_run_triton_longkv_reduce_q1(
partial_output,
partial_lse,
compact_paged_metadata["reduce_indptr"],
compact_paged_metadata["reduce_q_indices"],
compact_paged_metadata["reduce_partial_map"],
output,
)
return output
smallkv_dense_exact_splits = _get_smallkv_dense_exact_splits(
batch_size,
kv_seq_len,
alt_page_size,
num_kv_splits,
)
if smallkv_dense_exact_splits is not None:
compact_paged_metadata = _get_reduce_compact_metadata(paged_metadata)
dense_layout = _get_dense_reduce_layout(compact_paged_metadata)
use_dense_exact = (
dense_layout["q_identity"]
and dense_layout["partial_identity"]
and dense_layout["num_groups"] == batch_size
and dense_layout["num_partial"] == batch_size * smallkv_dense_exact_splits
and dense_layout["split_min"] == smallkv_dense_exact_splits
and dense_layout["split_max"] == smallkv_dense_exact_splits
)
if use_dense_exact:
partial_output, partial_lse = _get_cached_partials(
int(compact_paged_metadata["reduce_partial_map"].numel()),
q_seq_len,
num_heads,
v_head_dim,
q.device,
)
mla_decode_stage1_asm_fwd(
q_input.view(-1, num_heads, qk_head_dim),
paged_kv,
qo_indptr,
paged_metadata["kv_indptr"],
paged_kv_indices,
paged_metadata["kv_last_page_len"],
None,
paged_metadata["work_meta_data"],
paged_metadata["work_indptr"],
paged_metadata["work_info_set"],
q_seq_len,
alt_page_size,
num_kv_heads,
sm_scale,
partial_output,
partial_lse,
output,
q_scale,
kv_scale_fp8,
)
_probe_dense_fastpath_once(
batch_size,
kv_seq_len,
f"dense_exact{smallkv_dense_exact_splits} enabled",
)
_run_dense_longkv_reduce_q1_exact(
partial_output,
partial_lse,
output,
exact_splits=smallkv_dense_exact_splits,
)
return output
mla_decode_fwd(
q_input.view(-1, num_heads, qk_head_dim),
paged_kv,
output,
qo_indptr,
paged_metadata["kv_indptr"],
paged_kv_indices,
paged_metadata["kv_last_page_len"],
q_seq_len,
page_size=alt_page_size,
nhead_kv=num_kv_heads,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale_fp8,
intra_batch_mode=INTRA_BATCH_MODE,
work_meta_data=paged_metadata["work_meta_data"],
work_indptr=paged_metadata["work_indptr"],
work_info_set=paged_metadata["work_info_set"],
reduce_indptr=paged_metadata["reduce_indptr"],
reduce_final_map=paged_metadata["reduce_final_map"],
reduce_partial_map=paged_metadata["reduce_partial_map"],
)
return output
metadata = _get_cached_metadata(
qo_indptr,
kv_indptr,
batch_size,
q_seq_len,
kv_seq_len,
num_heads,
num_kv_heads,
num_kv_splits,
kv_granularity,
q_input.dtype,
kv_buffer_fp8.dtype,
)
kv_indices = _get_cached_kv_indices(kv_buffer_fp8.shape[0], kv_buffer_fp8.device)
if _use_stage1_direct(batch_size, kv_seq_len, num_kv_splits):
partial_output, partial_lse = _get_cached_partials(
int(metadata["reduce_partial_map"].numel()),
q_seq_len,
num_heads,
v_head_dim,
q.device,
)
mla_decode_stage1_asm_fwd(
q_input.view(-1, num_heads, qk_head_dim),
kv_buffer_fp8.view(
kv_buffer_fp8.shape[0],
PAGE_SIZE,
num_kv_heads,
kv_buffer_fp8.shape[-1],
),
qo_indptr,
kv_indptr,
kv_indices,
metadata["kv_last_page_len"],
None,
metadata["work_meta_data"],
metadata["work_indptr"],
metadata["work_info_set"],
q_seq_len,
PAGE_SIZE,
num_kv_heads,
sm_scale,
partial_output,
partial_lse,
output,
q_scale,
kv_scale_fp8,
)
return output
mla_decode_fwd(
q_input.view(-1, num_heads, qk_head_dim),
kv_buffer_fp8.view(
kv_buffer_fp8.shape[0],
PAGE_SIZE,
num_kv_heads,
kv_buffer_fp8.shape[-1],
),
output,
qo_indptr,
kv_indptr,
kv_indices,
metadata["kv_last_page_len"],
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=num_kv_heads,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale_fp8,
intra_batch_mode=INTRA_BATCH_MODE,
work_meta_data=metadata["work_meta_data"],
work_indptr=metadata["work_indptr"],
work_info_set=metadata["work_info_set"],
reduce_indptr=metadata["reduce_indptr"],
reduce_final_map=metadata["reduce_final_map"],
reduce_partial_map=metadata["reduce_partial_map"],
)
return output
scrolls · 1304 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