Skip to content
KernelIndex
Search⌘K

submission 670551

南北绿豆糕 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_0330_2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-670551?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
52.3µs
#169 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dba16a52fb65b85f75da7e1ccbb08a16f13a34d6bb7221d4b744f74f4761c46e
license declaredunknown
license concludedunknown
authors南北绿豆糕
imported2026-08-15

Techniques

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

persistent-kernelmla_decode_fwd(*args, persistent=True, **kwargs)

Kernel source

submission_0330_2.py687 lines
import os
import sys

import torch
from task import input_t, output_t


NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
PAGE_SIZE = 1
PG2_PAGE_SIZE = 2
NUM_KV_SPLITS = 64
KV_GRANULARITY = max(PAGE_SIZE, 16)
PG2_KV_GRANULARITY = 16  # Toggle to 64 to test the alternate leaderboard hint.
SM_SCALE = QK_HEAD_DIM ** -0.5


_DEVNULL_HANDLE = None


def _mute_process_stderr_forever() -> None:
    global _DEVNULL_HANDLE

    try:
        sys.stderr.flush()
    except Exception:
        pass

    try:
        devnull_fd = os.open(os.devnull, os.O_WRONLY)
        os.dup2(devnull_fd, 2)
        _DEVNULL_HANDLE = os.fdopen(devnull_fd, "w")
    except Exception:
        return

    try:
        sys.stderr = _DEVNULL_HANDLE
    except Exception:
        pass


_mute_process_stderr_forever()


try:
    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

    FP8_DTYPE = aiter_dtypes.fp8
    _AITER_AVAILABLE = True
except Exception:  # pragma: no cover - runtime availability depends on remote env
    mla_decode_fwd = None
    aiter_dtypes = None
    get_mla_metadata_info_v1 = None
    get_mla_metadata_v1 = None
    FP8_DTYPE = getattr(torch, "float8_e4m3fnuz", torch.uint8)
    _AITER_AVAILABLE = False


class _Meta:
    __slots__ = (
        "work_meta_data",
        "work_indptr",
        "work_info_set",
        "reduce_indptr",
        "reduce_final_map",
        "reduce_partial_map",
    )


class _PagedCache:
    __slots__ = (
        "page_size",
        "kv_granularity",
        "kv_indptr_i32",
        "kv_last_page_len",
        "kv_indices",
        "num_pages",
        "meta_bf16_fp8",
        "meta_fp8_fp8",
    )


class _CaseState:
    __slots__ = (
        "batch_size",
        "q_seq_len",
        "kv_seq_len",
        "device",
        "device_key",
        "total_q",
        "total_kv",
        "prefer_a8w8",
        "dispatch_mode",
        "dispatch_page_size",
        "a16w8_use_unit_scale",
        "qo_indptr_i32",
        "page1",
        "page2",
        "active_page",
        "out",
        "unit_scale",
        "q_raw_ptr",
        "q_bf16",
        "q_fp8_src_ptr",
        "q_fp8",
        "q_scale",
        "kv_raw_ptr",
        "kv_fp8",
        "kv_buffer_4d",
        "kv_scale_raw_ptr",
        "kv_scale",
    )


_KV_INDICES_CACHE: dict[tuple[int, tuple[str, int | None]], torch.Tensor] = {}
_STATE_CACHE: dict[tuple[int, int, tuple[str, int | None]], _CaseState] = {}
_LAST_CASE_KEY: tuple[int, int, tuple[str, int | None]] | None = None
_LAST_STATE: _CaseState | None = None


def _device_key(device: torch.device) -> tuple[str, int | None]:
    return (device.type, device.index)


def _make_uniform_indptr(batch_size: int, stride: int, device: torch.device) -> torch.Tensor:
    return torch.arange(batch_size + 1, dtype=torch.int32, device=device).mul_(stride)


def _get_cached_kv_indices(total_kv_len: int, device: torch.device) -> torch.Tensor:
    key = (total_kv_len, _device_key(device))
    cached = _KV_INDICES_CACHE.get(key)
    if cached is not None:
        return cached

    kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=device)
    _KV_INDICES_CACHE[key] = kv_indices
    return kv_indices


def _build_paged_cache(
    batch_size: int,
    q_seq_len: int,
    kv_seq_len: int,
    total_kv: int,
    device: torch.device,
    page_size: int,
    kv_granularity: int,
) -> _PagedCache:
    page = _PagedCache()
    page.page_size = page_size
    page.kv_granularity = kv_granularity

    if page_size == 1:
        page_count_per_batch = kv_seq_len
        page.kv_indptr_i32 = _make_uniform_indptr(batch_size, page_count_per_batch, device)
        page.kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
        page.kv_indices = _get_cached_kv_indices(total_kv, device)
        page.num_pages = total_kv
    else:
        page_count_per_batch = (kv_seq_len + page_size - 1) // page_size
        last_page_len = kv_seq_len % page_size
        if last_page_len == 0:
            last_page_len = page_size
        page.kv_indptr_i32 = _make_uniform_indptr(batch_size, page_count_per_batch, device)
        page.kv_last_page_len = torch.full((batch_size,), last_page_len, dtype=torch.int32, device=device)
        page.kv_indices = _get_cached_kv_indices(total_kv, device)[::page_size].div(page_size, rounding_mode="floor")
        page.num_pages = batch_size * page_count_per_batch

    if _AITER_AVAILABLE:
        page.meta_bf16_fp8 = _make_mla_decode_metadata(
            batch_size=batch_size,
            max_q_len=q_seq_len,
            q_dtype=torch.bfloat16,
            kv_dtype=FP8_DTYPE,
            qo_indptr=_make_uniform_indptr(batch_size, q_seq_len, device),
            kv_indptr=page.kv_indptr_i32,
            kv_last_page_len=page.kv_last_page_len,
            page_size=page.page_size,
            kv_granularity=page.kv_granularity,
        )
        page.meta_fp8_fp8 = _make_mla_decode_metadata(
            batch_size=batch_size,
            max_q_len=q_seq_len,
            q_dtype=FP8_DTYPE,
            kv_dtype=FP8_DTYPE,
            qo_indptr=_make_uniform_indptr(batch_size, q_seq_len, device),
            kv_indptr=page.kv_indptr_i32,
            kv_last_page_len=page.kv_last_page_len,
            page_size=page.page_size,
            kv_granularity=page.kv_granularity,
        )
    else:  # pragma: no cover - metadata only matters in remote runtime
        page.meta_bf16_fp8 = None
        page.meta_fp8_fp8 = None

    return page


def _make_mla_decode_metadata(
    batch_size: int,
    max_q_len: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    page_size: int,
    kv_granularity: int,
) -> _Meta:
    info = get_mla_metadata_info_v1(
        batch_size,
        max_q_len,
        NUM_HEADS,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=True,
        num_kv_splits=NUM_KV_SPLITS,
        intra_batch_mode=True,
    )
    work = [torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info]

    meta = _Meta()
    (
        meta.work_meta_data,
        meta.work_indptr,
        meta.work_info_set,
        meta.reduce_indptr,
        meta.reduce_final_map,
        meta.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,
        meta.work_meta_data,
        meta.work_info_set,
        meta.work_indptr,
        meta.reduce_indptr,
        meta.reduce_final_map,
        meta.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=True,
        max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )
    return meta


def _new_case_state(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device) -> _CaseState:
    state = _CaseState()
    state.batch_size = batch_size
    state.q_seq_len = q_seq_len
    state.kv_seq_len = kv_seq_len
    state.device = device
    state.device_key = _device_key(device)
    state.total_q = batch_size * q_seq_len
    state.total_kv = batch_size * kv_seq_len
    state.prefer_a8w8 = batch_size >= 256 and kv_seq_len >= 8192
    state.dispatch_mode = 0
    state.dispatch_page_size = 0
    state.a16w8_use_unit_scale = False

    state.qo_indptr_i32 = _make_uniform_indptr(batch_size, q_seq_len, device)
    state.page1 = _build_paged_cache(
        batch_size=batch_size,
        q_seq_len=q_seq_len,
        kv_seq_len=kv_seq_len,
        total_kv=state.total_kv,
        device=device,
        page_size=PAGE_SIZE,
        kv_granularity=KV_GRANULARITY,
    )
    state.page2 = None
    if kv_seq_len % PG2_PAGE_SIZE == 0 and state.total_kv % PG2_PAGE_SIZE == 0:
        state.page2 = _build_paged_cache(
            batch_size=batch_size,
            q_seq_len=q_seq_len,
            kv_seq_len=kv_seq_len,
            total_kv=state.total_kv,
            device=device,
            page_size=PG2_PAGE_SIZE,
            kv_granularity=max(PG2_PAGE_SIZE, PG2_KV_GRANULARITY),
        )
    state.active_page = state.page1
    state.out = torch.empty((state.total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
    state.unit_scale = torch.ones(1, dtype=torch.float32, device=device)

    state.q_raw_ptr = 0
    state.q_bf16 = None
    state.q_fp8_src_ptr = 0
    state.q_fp8 = None
    state.q_scale = None
    state.kv_raw_ptr = 0
    state.kv_fp8 = None
    state.kv_buffer_4d = None
    state.kv_scale_raw_ptr = 0
    state.kv_scale = None
    return state


def _get_case_state(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device) -> _CaseState:
    global _LAST_CASE_KEY, _LAST_STATE

    key = (batch_size, kv_seq_len, _device_key(device))
    state = _LAST_STATE
    if state is not None and _LAST_CASE_KEY == key:
        if state.q_seq_len == q_seq_len:
            return state

    state = _STATE_CACHE.get(key)
    if state is None or state.q_seq_len != q_seq_len:
        state = _new_case_state(batch_size, q_seq_len, kv_seq_len, device)
        _STATE_CACHE[key] = state

    _LAST_CASE_KEY = key
    _LAST_STATE = state
    return state


def _normalize_bf16_tensor(tensor: torch.Tensor, device: torch.device) -> torch.Tensor:
    if tensor.device == device and tensor.dtype == torch.bfloat16 and tensor.is_contiguous():
        return tensor
    return tensor.to(device=device, dtype=torch.bfloat16).contiguous()


def _normalize_tensor(tensor: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
    if tensor.device == device and tensor.dtype == dtype and tensor.is_contiguous():
        return tensor
    return tensor.to(device=device, dtype=dtype).contiguous()


def _normalize_scale(scale: torch.Tensor, device: torch.device) -> torch.Tensor:
    if scale.device == device and scale.dtype == torch.float32 and scale.is_contiguous() and scale.numel() == 1:
        return scale.reshape(1)
    return scale.to(device=device, dtype=torch.float32).reshape(1).contiguous()


def _can_zero_copy_pg2_view(state: _CaseState, kv_fp8: torch.Tensor) -> bool:
    if state.page2 is None:
        return False
    if kv_fp8.dim() != 3 or kv_fp8.shape != (state.total_kv, NUM_KV_HEADS, QK_HEAD_DIM):
        return False
    if kv_fp8.storage_offset() != 0 or not kv_fp8.is_contiguous():
        return False
    try:
        kv_fp8.view(state.page2.num_pages, PG2_PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
    except RuntimeError:
        return False
    return True


def _bind_inputs(state: _CaseState, q: torch.Tensor, kv_fp8: torch.Tensor, kv_scale: torch.Tensor) -> None:
    q_ptr = q.data_ptr()
    if q_ptr != state.q_raw_ptr:
        state.q_raw_ptr = q_ptr
        state.q_bf16 = _normalize_bf16_tensor(q, state.device)
        state.q_fp8_src_ptr = 0
        state.q_fp8 = None
        state.q_scale = None

    kv_ptr = kv_fp8.data_ptr()
    if kv_ptr != state.kv_raw_ptr:
        state.kv_raw_ptr = kv_ptr
        prev_page_size = state.active_page.page_size
        use_pg2 = _can_zero_copy_pg2_view(state, kv_fp8)
        state.active_page = state.page2 if use_pg2 else state.page1
        state.kv_fp8 = _normalize_tensor(kv_fp8, state.device, kv_fp8.dtype)
        state.kv_buffer_4d = state.kv_fp8.view(
            state.active_page.num_pages,
            state.active_page.page_size,
            NUM_KV_HEADS,
            QK_HEAD_DIM,
        )
        if state.active_page.page_size != prev_page_size:
            state.dispatch_mode = 0
            state.dispatch_page_size = 0

    kv_scale_ptr = kv_scale.data_ptr()
    if kv_scale_ptr != state.kv_scale_raw_ptr:
        state.kv_scale_raw_ptr = kv_scale_ptr
        state.kv_scale = _normalize_scale(kv_scale, state.device)


def _quantize_fp8_per_tensor(tensor: torch.Tensor, fp8_dtype: torch.dtype) -> tuple[torch.Tensor, torch.Tensor]:
    finfo = torch.finfo(fp8_dtype)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = (amax / finfo.max).to(torch.float32).reshape(1)
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(fp8_dtype)
    if not fp8_tensor.is_contiguous():
        fp8_tensor = fp8_tensor.contiguous()
    return fp8_tensor, scale


def _ensure_q_fp8(state: _CaseState) -> None:
    src_ptr = state.q_bf16.data_ptr()
    if src_ptr == state.q_fp8_src_ptr:
        return

    state.q_fp8_src_ptr = src_ptr
    state.q_fp8, state.q_scale = _quantize_fp8_per_tensor(state.q_bf16, FP8_DTYPE)


def _call_mla_decode_fwd(*args, **kwargs) -> None:
    try:
        mla_decode_fwd(*args, persistent=True, **kwargs)
        return
    except TypeError as exc:
        if "persistent" not in str(exc):
            raise
    mla_decode_fwd(*args, **kwargs)


def _run_aiter_decode(state: _CaseState, q_input: torch.Tensor, q_scale: torch.Tensor | None, meta: _Meta) -> torch.Tensor:
    page = state.active_page
    _call_mla_decode_fwd(
        q_input,
        state.kv_buffer_4d,
        state.out,
        state.qo_indptr_i32,
        page.kv_indptr_i32,
        page.kv_indices,
        page.kv_last_page_len,
        state.q_seq_len,
        page_size=page.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=state.kv_scale,
        intra_batch_mode=True,
        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,
    )
    return state.out


def _run_a16w8(state: _CaseState) -> torch.Tensor:
    q_scale = state.unit_scale if state.a16w8_use_unit_scale else None
    return _run_aiter_decode(state, state.q_bf16, q_scale, state.active_page.meta_bf16_fp8)


def _run_a8w8(state: _CaseState) -> torch.Tensor:
    _ensure_q_fp8(state)
    return _run_aiter_decode(state, state.q_fp8, state.q_scale, state.active_page.meta_fp8_fp8)


def _resolve_dispatch(state: _CaseState) -> int:
    if not _AITER_AVAILABLE:
        state.dispatch_mode = 3
        state.dispatch_page_size = state.active_page.page_size
        return 3

    if state.prefer_a8w8:
        try:
            _run_a8w8(state)
            state.dispatch_mode = 2
            state.dispatch_page_size = state.active_page.page_size
            return 2
        except Exception:
            pass
        try:
            _run_a16w8(state)
            state.dispatch_mode = 1
            state.dispatch_page_size = state.active_page.page_size
            return 1
        except Exception:
            if not state.a16w8_use_unit_scale:
                try:
                    state.a16w8_use_unit_scale = True
                    _run_a16w8(state)
                    state.dispatch_mode = 1
                    state.dispatch_page_size = state.active_page.page_size
                    return 1
                except Exception:
                    pass
    else:
        try:
            _run_a16w8(state)
            state.dispatch_mode = 1
            state.dispatch_page_size = state.active_page.page_size
            return 1
        except Exception:
            if not state.a16w8_use_unit_scale:
                try:
                    state.a16w8_use_unit_scale = True
                    _run_a16w8(state)
                    state.dispatch_mode = 1
                    state.dispatch_page_size = state.active_page.page_size
                    return 1
                except Exception:
                    pass
        try:
            _run_a8w8(state)
            state.dispatch_mode = 2
            state.dispatch_page_size = state.active_page.page_size
            return 2
        except Exception:
            pass

    state.dispatch_mode = 3
    state.dispatch_page_size = state.active_page.page_size
    return 3


def _unwrap_scaled_mm_result(result):
    if isinstance(result, tuple):
        return result[0]
    return result


def _scaled_mm_scores(
    q_fp8_2d: torch.Tensor,
    k_fp8_t_2d: torch.Tensor,
    q_scale: torch.Tensor,
    k_scale: torch.Tensor,
) -> torch.Tensor:
    scaled_mm = getattr(torch, "_scaled_mm", None)
    if scaled_mm is None:
        raise RuntimeError("torch._scaled_mm is unavailable")

    kwargs_candidates = (
        {
            "scale_a": q_scale,
            "scale_b": k_scale,
            "out_dtype": torch.float32,
            "use_fast_accum": True,
        },
        {
            "scale_a": q_scale,
            "scale_b": k_scale,
            "out_dtype": torch.float32,
        },
        {
            "scale_a": q_scale,
            "scale_b": k_scale,
        },
        {
            "scale_a": q_scale.item(),
            "scale_b": k_scale.item(),
            "out_dtype": torch.float32,
            "use_fast_accum": True,
        },
        {
            "scale_a": q_scale.item(),
            "scale_b": k_scale.item(),
            "out_dtype": torch.float32,
        },
        {
            "scale_a": q_scale.item(),
            "scale_b": k_scale.item(),
        },
    )
    last_error = None
    for kwargs in kwargs_candidates:
        try:
            return _unwrap_scaled_mm_result(scaled_mm(q_fp8_2d, k_fp8_t_2d, **kwargs))
        except Exception as exc:  # pragma: no cover - runtime compatibility path
            last_error = exc
    raise RuntimeError(f"torch._scaled_mm call failed: {last_error}")


def _segment_attention_with_fp8_scores(
    q_bf16: torch.Tensor,
    q_fp8: torch.Tensor,
    q_scale: torch.Tensor,
    kv_fp8: torch.Tensor,
    kv_scale: torch.Tensor,
) -> torch.Tensor:
    seg_q, num_heads, _ = q_bf16.shape
    kv_len = kv_fp8.shape[0]
    if kv_len == 0:
        return torch.zeros((seg_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device=q_bf16.device)

    q_fp8_2d = q_fp8.reshape(seg_q * num_heads, QK_HEAD_DIM)
    if not q_fp8_2d.is_contiguous():
        q_fp8_2d = q_fp8_2d.contiguous()
    k_fp8_t_2d = kv_fp8[:, 0, :QK_HEAD_DIM].transpose(0, 1)
    if not k_fp8_t_2d.is_contiguous():
        k_fp8_t_2d = k_fp8_t_2d.contiguous()

    scores = _scaled_mm_scores(q_fp8_2d, k_fp8_t_2d, q_scale, kv_scale)
    scores = scores.to(torch.float32).mul_(SM_SCALE)
    weights = torch.softmax(scores, dim=-1)
    v_f32 = kv_fp8[:, 0, :V_HEAD_DIM].to(torch.float32).mul_(kv_scale)
    out = weights @ v_f32
    return out.view(seg_q, num_heads, V_HEAD_DIM)


def _segment_attention_fallback(
    q_bf16: torch.Tensor,
    kv_fp8: torch.Tensor,
    kv_scale: torch.Tensor,
) -> torch.Tensor:
    seg_q, num_heads, _ = q_bf16.shape
    kv_len = kv_fp8.shape[0]
    if kv_len == 0:
        return torch.zeros((seg_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device=q_bf16.device)

    k_f32 = kv_fp8[:, 0, :QK_HEAD_DIM].to(torch.float32).mul_(kv_scale)
    v_f32 = kv_fp8[:, 0, :V_HEAD_DIM].to(torch.float32).mul_(kv_scale)
    scores = torch.matmul(q_bf16.to(torch.float32), k_f32.transpose(0, 1)).mul_(SM_SCALE)
    weights = torch.softmax(scores, dim=-1)
    out = torch.matmul(weights, v_f32)
    return out.view(seg_q, num_heads, V_HEAD_DIM)


def _run_torch_fallback(state: _CaseState) -> torch.Tensor:
    _ensure_q_fp8(state)
    out = state.out
    batch_size = state.batch_size
    q_seq_len = state.q_seq_len
    kv_seq_len = state.kv_seq_len
    q_bf16 = state.q_bf16
    q_fp8 = state.q_fp8
    q_scale = state.q_scale
    kv_fp8 = state.kv_fp8
    kv_scale = state.kv_scale

    use_fp8_scores = True
    for batch_idx in range(batch_size):
        q_start = batch_idx * q_seq_len
        q_end = q_start + q_seq_len
        kv_start = batch_idx * kv_seq_len
        kv_end = kv_start + kv_seq_len

        q_seg_bf16 = q_bf16[q_start:q_end]
        kv_seg_fp8 = kv_fp8[kv_start:kv_end]
        if use_fp8_scores:
            try:
                out[q_start:q_end] = _segment_attention_with_fp8_scores(
                    q_seg_bf16,
                    q_fp8[q_start:q_end],
                    q_scale,
                    kv_seg_fp8,
                    kv_scale,
                )
                continue
            except Exception:  # pragma: no cover - runtime fallback path
                use_fp8_scores = False

        out[q_start:q_end] = _segment_attention_fallback(q_seg_bf16, kv_seg_fp8, kv_scale)

    return out


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, _kv_indptr, _config = data
    kv_fp8, kv_scale = kv_data["fp8"]

    batch_size = qo_indptr.numel() - 1
    q_seq_len = q.shape[0] // batch_size
    kv_seq_len = kv_fp8.shape[0] // batch_size

    state = _get_case_state(batch_size, q_seq_len, kv_seq_len, q.device)
    _bind_inputs(state, q, kv_fp8, kv_scale)

    mode = state.dispatch_mode
    if state.dispatch_page_size != state.active_page.page_size:
        mode = 0
    if mode == 0:
        mode = _resolve_dispatch(state)

    if mode == 1:
        return _run_a16w8(state)
    if mode == 2:
        return _run_a8w8(state)
    return _run_torch_fallback(state)
scrolls · 687 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