Skip to content
KernelIndex
Search⌘K

submission 723086

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 589 lines, June 9 Researcher Reciprocity License v1.0.

Submission.v208.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-723086?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
33.0µs
#39 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:061d3e148ef19bf616b5fb60d635de3e59cb4f79d5c0dd571ad2c5fd53951cb3
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernel_NON_PERSISTENT_SHAPES: set[tuple[int, int]] = set()

Kernel source

Submission.v208.py589 lines
# gpumode leaderboard reference
"""
MLA-SESSION true b4 hybrid.

Keep the live `v158` runner unchanged for every non-`b4` shape.
Only the two `b4` shapes get the old exact-shape-bank runner family:
- `(4, 1024)` uses the `v167` exact-shape runner depth
- `(4, 8192)` uses the `v173` deeper-Q-ring exact-shape runner depth

This targets the only remaining artifact-backed ranked gap after the `b256`
hybrid lane died in `v206` and `v207`.
"""

import os

os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")

import torch
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1


NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
Q_DTYPE = os.getenv("AMD2_MLA_Q_DTYPE", "fp8").lower()
KV_DTYPE = os.getenv("AMD2_MLA_KV_DTYPE", "fp8").lower()

if Q_DTYPE not in {"fp8", "bf16"}:
    raise ValueError(f"Unsupported AMD2_MLA_Q_DTYPE={Q_DTYPE!r}")
if KV_DTYPE not in {"fp8", "bf16"}:
    raise ValueError(f"Unsupported AMD2_MLA_KV_DTYPE={KV_DTYPE!r}")

_NON_PERSISTENT_SHAPES: set[tuple[int, int]] = set()
_SHAPE_CONFIG: dict[tuple[int, int], tuple[int, int, int, bool]] = {
    (4, 1024): (1, 8, 128, False),
    (4, 8192): (8, 32, 128, False),
    (32, 1024): (1, 4, 128, False),
    (32, 8192): (8, 32, 32, True),
    (64, 1024): (2, 8, 128, False),
    (64, 8192): (8, 32, 32, False),
    (256, 1024): (2, 8, 32, False),
    (256, 8192): (8, 32, 32, False),
}
_Q_CACHE_SLOTS = 16
_OUTPUT_CACHE_SLOTS = 4
_FAST_PATH_SHAPES = {
    (4, 1024),
    (4, 8192),
}
_EXACT_Q_CACHE_SLOTS_BY_SHAPE = {
    (4, 1024): 16,
    (4, 8192): 32,
}
_EXACT_SAFE_Q_CACHE_SLOTS = 16

torch.set_grad_enabled(False)

_KV_INDICES_CACHE = {}
_PERSISTENT_SETUP_CACHE = {}
_NON_PERSISTENT_SETUP_CACHE = {}
_PERSISTENT_RUNNER_CACHE = {}
_EXACT_SHAPE_RUNNER_BANK = {}
_UNIT_Q_SCALE_CACHE = {}
_EXACT_SHAPE_CONFIGS = {
    shape: {
        "batch_size": shape[0],
        "q_seq_len": 1,
        "kv_seq_len": shape[1],
        "num_heads": NUM_HEADS,
        "num_kv_heads": NUM_KV_HEADS,
        "qk_head_dim": QK_HEAD_DIM,
        "v_head_dim": V_HEAD_DIM,
        "sm_scale": SM_SCALE,
    }
    for shape in _FAST_PATH_SHAPES
}


def _get_unit_q_scale(device: torch.device) -> torch.Tensor:
    key = device.index or 0
    cached = _UNIT_Q_SCALE_CACHE.get(key)
    if cached is None:
        cached = torch.ones((1,), dtype=torch.float32, device=device)
        _UNIT_Q_SCALE_CACHE[key] = cached
    return cached


def _quantize_fp8_copy(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    out = torch.empty(tensor.shape, dtype=FP8_DTYPE, device=tensor.device)
    out.copy_(tensor)
    return out, _get_unit_q_scale(tensor.device)


def _get_shape_config(config: dict) -> tuple[int, int, int, bool]:
    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    shape = (batch_size, kv_seq_len)
    exact = _SHAPE_CONFIG.get(shape)
    if exact is not None:
        return exact
    if kv_seq_len >= 8192:
        if batch_size <= 4:
            return 8, 32, 128, False
        if batch_size <= 32:
            return 8, 32, 32, True
        return 8, 32, 32, False
    if batch_size <= 32:
        return 1, 8, 128, False
    if batch_size <= 64:
        return 2, 8, 128, False
    return 2, 8, 32, False


def _is_non_persistent_case(config: dict) -> bool:
    return (int(config["batch_size"]), int(config["kv_seq_len"])) in _NON_PERSISTENT_SHAPES


def _get_kv_indices(total_pages: int, device: torch.device) -> torch.Tensor:
    key = (device.index or 0, total_pages)
    cached = _KV_INDICES_CACHE.get(key)
    if cached is None:
        cached = torch.arange(total_pages, dtype=torch.int32, device=device)
        _KV_INDICES_CACHE[key] = cached
    return cached


def _make_metadata(
    qo_indptr,
    kv_indptr,
    kv_last_page_len,
    nhead,
    nhead_kv,
    q_dtype,
    kv_dtype,
    page_size,
    num_kv_splits,
    kv_granularity,
    intra_batch_mode,
):
    batch_size = qo_indptr.numel() - 1
    max_q_len = int((qo_indptr[1] - qo_indptr[0]).item())
    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nhead, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=False, num_kv_splits=num_kv_splits, intra_batch_mode=intra_batch_mode,
    )
    work = [torch.empty(shape, dtype=dtype, device=qo_indptr.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,
        nhead // nhead_kv, nhead_kv, 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=max_q_len, uni_seqlen_qo=max_q_len,
        fast_mode=False, max_split_per_batch=num_kv_splits, intra_batch_mode=intra_batch_mode,
        dtype_q=q_dtype, dtype_kv=kv_dtype,
    )
    return {
        "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,
    }


def _get_non_persistent_setup(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device):
    key = (device.index or 0, batch_size, q_seq_len, kv_seq_len)
    cached = _NON_PERSISTENT_SETUP_CACHE.get(key)
    if cached is None:
        qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
        kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
        cached = {
            "qo_indptr": qo_indptr,
            "kv_indptr": kv_indptr,
            "kv_indices": _get_kv_indices(batch_size * kv_seq_len, device),
            "kv_last_page_len": torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
        }
        _NON_PERSISTENT_SETUP_CACHE[key] = cached
    return cached


def _get_persistent_setup(
    config: dict,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    page_size: int,
    num_kv_splits: int,
    kv_granularity: int,
    intra_batch_mode: bool,
    device: torch.device,
):
    batch_size = int(config["batch_size"])
    q_seq_len = int(config["q_seq_len"])
    kv_seq_len = int(config["kv_seq_len"])
    nhead = int(config["num_heads"])
    nhead_kv = int(config["num_kv_heads"])
    pages_per_seq = (kv_seq_len + page_size - 1) // page_size
    key = (
        device.index or 0,
        batch_size,
        q_seq_len,
        kv_seq_len,
        nhead,
        nhead_kv,
        q_dtype,
        kv_dtype,
        page_size,
        num_kv_splits,
        kv_granularity,
        intra_batch_mode,
    )
    cached = _PERSISTENT_SETUP_CACHE.get(key)
    if cached is None:
        qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
        kv_page_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * pages_per_seq
        kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
        cached = {
            "qo_indptr": qo_indptr,
            "kv_page_indptr": kv_page_indptr,
            "kv_page_indices": _get_kv_indices(batch_size * pages_per_seq, device),
            "kv_last_page_len": kv_last_page_len,
            "meta": _make_metadata(
                qo_indptr,
                kv_page_indptr,
                kv_last_page_len,
                nhead,
                nhead_kv,
                q_dtype,
                kv_dtype,
                page_size,
                num_kv_splits,
                kv_granularity,
                intra_batch_mode,
            ),
        }
        _PERSISTENT_SETUP_CACHE[key] = cached
    return cached


def _reshape_paged_kv(kv_buffer: torch.Tensor, batch_size: int, kv_seq_len: int, page_size: int, nhead_kv: int, qk_head_dim: int) -> torch.Tensor:
    total_kv = batch_size * kv_seq_len
    if kv_buffer.shape[0] != total_kv:
        raise ValueError(f"Expected uniform total_kv={total_kv}, got {kv_buffer.shape[0]}")
    if kv_seq_len % page_size != 0:
        pages_per_seq = (kv_seq_len + page_size - 1) // page_size
        kv_paged = torch.zeros((batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim), dtype=kv_buffer.dtype, device=kv_buffer.device)
        kv_seq = kv_buffer.reshape(batch_size, kv_seq_len, nhead_kv, qk_head_dim)
        for batch_idx in range(batch_size):
            flat = kv_paged[batch_idx * pages_per_seq:(batch_idx + 1) * pages_per_seq].view(pages_per_seq * page_size, nhead_kv, qk_head_dim)
            flat[:kv_seq_len].copy_(kv_seq[batch_idx])
        return kv_paged
    pages_per_seq = kv_seq_len // page_size
    return kv_buffer.reshape(batch_size, kv_seq_len, nhead_kv, qk_head_dim).reshape(batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)


def _run_non_persistent(q, kv_buffer, config, q_scale, kv_scale):
    batch_size = int(config["batch_size"])
    q_seq_len = int(config["q_seq_len"])
    kv_seq_len = int(config["kv_seq_len"])
    nhead = int(config["num_heads"])
    nhead_kv = int(config["num_kv_heads"])
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])
    setup = _get_non_persistent_setup(batch_size, q_seq_len, kv_seq_len, q.device)
    out = _get_output_workspace(q.shape[0], nhead, v_head_dim, q.device)
    mla_decode_fwd(
        q.reshape(-1, nhead, qk_head_dim),
        kv_buffer.reshape(kv_buffer.shape[0], 1, nhead_kv, qk_head_dim),
        out,
        setup["qo_indptr"], setup["kv_indptr"], setup["kv_indices"], setup["kv_last_page_len"], q_seq_len,
        page_size=1, nhead_kv=nhead_kv, sm_scale=float(config.get("sm_scale", SM_SCALE)), logit_cap=0.0,
        q_scale=q_scale, kv_scale=kv_scale,
    )
    return out


def _run_persistent(q, kv_buffer, config, q_scale, kv_scale):
    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    nhead = int(config["num_heads"])
    nhead_kv = int(config["num_kv_heads"])
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])
    q_seq_len = int(config["q_seq_len"])
    page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
    setup = _get_persistent_setup(
        config,
        q.dtype,
        kv_buffer.dtype,
        page_size,
        num_kv_splits,
        kv_granularity,
        intra_batch_mode,
        q.device,
    )
    out = _get_output_workspace(q.shape[0], nhead, v_head_dim, q.device)
    mla_decode_fwd(
        q.reshape(-1, nhead, qk_head_dim),
        _reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim),
        out,
        setup["qo_indptr"], setup["kv_page_indptr"], setup["kv_page_indices"], setup["kv_last_page_len"], q_seq_len,
        page_size=page_size, nhead_kv=nhead_kv, sm_scale=float(config.get("sm_scale", SM_SCALE)), logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=intra_batch_mode,
        **setup["meta"],
    )
    return out


def _make_persistent_runner(config: dict, kv_dtype: torch.dtype, device: torch.device):
    batch_size = int(config["batch_size"])
    q_seq_len = int(config["q_seq_len"])
    kv_seq_len = int(config["kv_seq_len"])
    nhead = int(config["num_heads"])
    nhead_kv = int(config["num_kv_heads"])
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])
    num_tokens = batch_size * q_seq_len
    page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
    q_storage_dtype = FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16
    setup = _get_persistent_setup(
        config,
        q_storage_dtype,
        kv_dtype,
        page_size,
        num_kv_splits,
        kv_granularity,
        intra_batch_mode,
        device,
    )
    qo_indptr = setup["qo_indptr"]
    kv_page_indptr = setup["kv_page_indptr"]
    kv_page_indices = setup["kv_page_indices"]
    kv_last_page_len = setup["kv_last_page_len"]
    meta = setup["meta"]
    work_meta_data = meta["work_meta_data"]
    work_indptr = meta["work_indptr"]
    work_info_set = meta["work_info_set"]
    reduce_indptr = meta["reduce_indptr"]
    reduce_final_map = meta["reduce_final_map"]
    reduce_partial_map = meta["reduce_partial_map"]
    sm_scale = float(config.get("sm_scale", SM_SCALE))
    pages_per_seq = (kv_seq_len + page_size - 1) // page_size
    output_buffers = [
        torch.empty((num_tokens, nhead, v_head_dim), dtype=torch.bfloat16, device=device)
        for _ in range(_OUTPUT_CACHE_SLOTS)
    ]
    output_index = 0

    if Q_DTYPE == "fp8":
        q_buffers = [
            torch.empty((num_tokens, nhead, qk_head_dim), dtype=FP8_DTYPE, device=device)
            for _ in range(_Q_CACHE_SLOTS)
        ]
        q_index = 0
        unit_q_scale = _get_unit_q_scale(device)
    else:
        q_buffers = []
        q_index = 0
        unit_q_scale = None

    if kv_seq_len % page_size == 0:
        def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
            return kv_buffer.reshape(batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)
    else:
        def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
            return _reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim)

    def runner(q: torch.Tensor, kv_buffer: torch.Tensor, kv_scale: torch.Tensor | None) -> torch.Tensor:
        nonlocal output_index, q_index

        if Q_DTYPE == "fp8":
            q_input = q_buffers[q_index]
            q_input.copy_(q)
            q_index = (q_index + 1) % len(q_buffers)
            q_scale = unit_q_scale
        else:
            q_input = q
            q_scale = None

        out = output_buffers[output_index]
        output_index = (output_index + 1) % len(output_buffers)

        mla_decode_fwd(
            q_input.reshape(-1, nhead, qk_head_dim),
            reshape_kv(kv_buffer),
            out,
            qo_indptr,
            kv_page_indptr,
            kv_page_indices,
            kv_last_page_len,
            q_seq_len,
            page_size=page_size,
            nhead_kv=nhead_kv,
            sm_scale=sm_scale,
            logit_cap=0.0,
            num_kv_splits=num_kv_splits,
            q_scale=q_scale,
            kv_scale=kv_scale,
            intra_batch_mode=intra_batch_mode,
            work_meta_data=work_meta_data,
            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,
        )
        return out

    return runner


def _get_persistent_runner(config: dict, kv_dtype: torch.dtype, device: torch.device):
    key = (
        device.index or 0,
        int(config["batch_size"]),
        int(config["q_seq_len"]),
        int(config["kv_seq_len"]),
        int(config["num_heads"]),
        int(config["num_kv_heads"]),
        int(config["qk_head_dim"]),
        int(config["v_head_dim"]),
        FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16,
        kv_dtype,
    )
    cached = _PERSISTENT_RUNNER_CACHE.get(key)
    if cached is None:
        cached = _make_persistent_runner(config, kv_dtype, device)
        _PERSISTENT_RUNNER_CACHE[key] = cached
    return cached


def _make_exact_shape_runner(shape: tuple[int, int], kv_dtype: torch.dtype, device: torch.device):
    config = _EXACT_SHAPE_CONFIGS[shape]
    batch_size = int(config["batch_size"])
    q_seq_len = int(config["q_seq_len"])
    kv_seq_len = int(config["kv_seq_len"])
    nhead = int(config["num_heads"])
    nhead_kv = int(config["num_kv_heads"])
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])
    num_tokens = batch_size * q_seq_len
    page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
    q_storage_dtype = FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16
    setup = _get_persistent_setup(
        config,
        q_storage_dtype,
        kv_dtype,
        page_size,
        num_kv_splits,
        kv_granularity,
        intra_batch_mode,
        device,
    )
    qo_indptr = setup["qo_indptr"]
    kv_page_indptr = setup["kv_page_indptr"]
    kv_page_indices = setup["kv_page_indices"]
    kv_last_page_len = setup["kv_last_page_len"]
    meta = setup["meta"]
    work_meta_data = meta["work_meta_data"]
    work_indptr = meta["work_indptr"]
    work_info_set = meta["work_info_set"]
    reduce_indptr = meta["reduce_indptr"]
    reduce_final_map = meta["reduce_final_map"]
    reduce_partial_map = meta["reduce_partial_map"]
    sm_scale = float(config.get("sm_scale", SM_SCALE))
    pages_per_seq = (kv_seq_len + page_size - 1) // page_size

    output_buffers = [
        torch.empty((num_tokens, nhead, v_head_dim), dtype=torch.bfloat16, device=device)
        for _ in range(_OUTPUT_CACHE_SLOTS)
    ]
    output_index = 0
    num_output_buffers = len(output_buffers)

    if Q_DTYPE == "fp8":
        q_cache_slots = _EXACT_Q_CACHE_SLOTS_BY_SHAPE.get(shape, _EXACT_SAFE_Q_CACHE_SLOTS)
        q_buffers = [
            torch.empty((num_tokens, nhead, qk_head_dim), dtype=FP8_DTYPE, device=device)
            for _ in range(q_cache_slots)
        ]
        q_index = 0
        num_q_buffers = len(q_buffers)
        unit_q_scale = _get_unit_q_scale(device)
    else:
        q_buffers = []
        q_index = 0
        num_q_buffers = 0
        unit_q_scale = None

    if kv_seq_len % page_size == 0:
        kv_view_shape = (batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)

        def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
            return kv_buffer.view(*kv_view_shape)
    else:
        def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
            return _reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim)

    def runner(q: torch.Tensor, kv_buffer: torch.Tensor, kv_scale: torch.Tensor | None) -> torch.Tensor:
        nonlocal output_index, q_index

        if Q_DTYPE == "fp8":
            q_input = q_buffers[q_index]
            q_input.copy_(q)
            q_index += 1
            if q_index == num_q_buffers:
                q_index = 0
            q_scale = unit_q_scale
        else:
            q_input = q
            q_scale = None

        out = output_buffers[output_index]
        output_index += 1
        if output_index == num_output_buffers:
            output_index = 0

        mla_decode_fwd(
            q_input.reshape(-1, nhead, qk_head_dim),
            reshape_kv(kv_buffer),
            out,
            qo_indptr,
            kv_page_indptr,
            kv_page_indices,
            kv_last_page_len,
            q_seq_len,
            page_size=page_size,
            nhead_kv=nhead_kv,
            sm_scale=sm_scale,
            logit_cap=0.0,
            num_kv_splits=num_kv_splits,
            q_scale=q_scale,
            kv_scale=kv_scale,
            intra_batch_mode=intra_batch_mode,
            work_meta_data=work_meta_data,
            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,
        )
        return out

    return runner


def _get_exact_shape_runner(shape: tuple[int, int], kv_dtype: torch.dtype, device: torch.device):
    bank_key = (device.index or 0, kv_dtype)
    bank = _EXACT_SHAPE_RUNNER_BANK.get(bank_key)
    if bank is None:
        bank = {}
        _EXACT_SHAPE_RUNNER_BANK[bank_key] = bank
    runner = bank.get(shape)
    if runner is None:
        config = _EXACT_SHAPE_CONFIGS.get(shape)
        if config is None:
            return None
        runner = _make_exact_shape_runner(shape, kv_dtype, device)
        bank[shape] = runner
    return runner


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, _qo_indptr, _kv_indptr, config = data
    if int(config["q_seq_len"]) != 1:
        raise ValueError(f"Expected decode q_seq_len=1, got {config['q_seq_len']}")
    shape = (int(config["batch_size"]), int(config["kv_seq_len"]))
    if KV_DTYPE == "fp8":
        kv_input, kv_scale = kv_data["fp8"]
    else:
        kv_input, kv_scale = kv_data["bf16"], None
    if shape in _FAST_PATH_SHAPES:
        runner = _get_exact_shape_runner(shape, kv_input.dtype, q.device)
        if runner is not None:
            return runner(q, kv_input, kv_scale)
    return _get_persistent_runner(config, kv_input.dtype, q.device)(q, kv_input, kv_scale)
scrolls · 589 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