Skip to content
KernelIndex
Search⌘K

submission 617362

pradeep03071 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-617362?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
94.7µs
#450 of 766
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b165d3cf28a228c0976f1bb33b03aed26197fc3b0d322302a4e78433684b5bc0
license declaredunknown
license concludedunknown
authorspradeep03071
imported2026-08-26

Techniques

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

fp4"mxfp4": (Tensor, Tensor) kv_buffer fp4x2 + fp8_e8m0 — block-32 quantized
persistent-kernelDecode only — persistent mode with get_mla_metadata_v1.

Kernel source

submission.py643 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

# gpumode leaderboard reference
"""
Reference implementation for MLA (Multi-head Latent Attention) decode kernel.

Uses aiter MLA kernels (mla_decode_fwd) as the reference.
DeepSeek R1 forward_absorb MLA: absorbed q (576), compressed kv_buffer (576),
output v_head_dim = kv_lora_rank = 512.

The input provides:
  q:       (total_q, 16, 576) bfloat16 — absorbed query
  kv_data: dict with KV cache in three formats:
    "bf16":  Tensor  (total_kv, 1, 576)  bfloat16          — highest precision
    "fp8":   (Tensor, Tensor)  kv_buffer fp8 + scalar scale — per-tensor quantized
    "mxfp4": (Tensor, Tensor)  kv_buffer fp4x2 + fp8_e8m0  — block-32 quantized
  The reference quantizes Q to fp8 on-the-fly inside ref_kernel.

The reference kernel quantizes Q to fp8 on-the-fly and uses fp8 KV (a8w8 kernel),
which is ~2-3x faster than bf16 on MI355X with negligible accuracy loss.

Decode only — persistent mode with get_mla_metadata_v1.
"""

import torch
import torch.nn.functional as F
import atexit
import os
from task import input_t, output_t
from utils import make_match_reference

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
from aiter.utility.fp4_utils import (
    dynamic_mxfp4_quant,
    mxfp4_to_f32,
    e8m0_to_f32,
)

# ---------------------------------------------------------------------------
# DeepSeek R1 latent MQA constants (forward_absorb path)
# https://huggingface.co/deepseek-ai/DeepSeek-R1-0528/blob/main/config.json
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM   # 576
V_HEAD_DIM = KV_LORA_RANK                        # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

PAGE_SIZE = 1
NUM_KV_SPLITS = 32

# FP8 dtype (platform-specific via aiter)
FP8_DTYPE = aiter_dtypes.fp8

# Reuse persistent scratch to cut per-call allocator overhead.
_META_WORK_CACHE: dict[tuple, list[torch.Tensor]] = {}
_META_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
_KV_INDICES_CACHE: dict[tuple, torch.Tensor] = {}

# Remote submit runners may not forward local env vars, so keep a file-local
# switch for one-off profiling submissions.
PROFILE_FORCE_ENABLE = False
PROFILE_ENABLED = PROFILE_FORCE_ENABLE or os.getenv("MLA_PROFILE", "").lower() in {"1", "true", "yes", "on"}
_PROFILE_STATS: dict[tuple, dict[str, float]] = {}

# Experiment presets.
# Keep run_002 as the default pinned baseline for reproducibility.
# benchmark_history.md has a newer run_004 result, but the exact preset that
# produced it was not preserved here, so ACTIVE_EXPERIMENT should not move
# until that mapping is recovered and named explicitly.
EXPERIMENTS = {
    # Best-known baseline from benchmark history run_002.
    "run_002_baseline": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_token",
        "split_strategy": "heuristic",
        "fast_mode": False,
    },
    # Remove Q quantization overhead while keeping KV compressed.
    "bf16_q_fp8_kv": {
        "q_dtype": "bf16",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_tensor",
        "split_strategy": "heuristic",
        "fast_mode": False,
    },
    # Compare against the simpler per-tensor Q scaling path.
    "fp8_per_tensor_q": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_tensor",
        "split_strategy": "heuristic",
        "fast_mode": False,
    },
    # Try the backend's faster scheduling mode without changing precision.
    "baseline_fast_mode": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_token",
        "split_strategy": "heuristic",
        "fast_mode": True,
    },
    # Combine the two highest-signal low-risk changes.
    "bf16_q_fast_mode": {
        "q_dtype": "bf16",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_tensor",
        "split_strategy": "heuristic",
        "fast_mode": True,
    },
    # Preserve the current tuned-split path for direct comparison.
    "run_003_lookup": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_tensor",
        "split_strategy": "lookup",
        "fast_mode": False,
    },
    # More aggressive split reduction for qseqlen=1 decode.
    "run_005_reduce_tuned": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_token",
        "split_strategy": "lookup_reduce_tuned",
        "fast_mode": False,
    },
    # Combine the run_003 split table with the backend fast scheduler.
    "run_006_lookup_fast": {
        "q_dtype": "fp8",
        "kv_dtype": "fp8",
        "q_fp8_scaling": "per_token",
        "split_strategy": "lookup",
        "fast_mode": True,
    },
}
# Flip this to run neighboring ablations without rewriting the hot path.
ACTIVE_EXPERIMENT = "run_006_lookup_fast"
EXPERIMENT = EXPERIMENTS[ACTIVE_EXPERIMENT]


# ---------------------------------------------------------------------------
# FP8 quantization (sglang style: dynamic per-tensor)
# ---------------------------------------------------------------------------
def quantize_fp8_per_tensor(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).

    Args:
        tensor: bf16 tensor to quantize

    Returns:
        (fp8_tensor, scale) where scale is a scalar float32 tensor.
        Dequantize: fp8_tensor.to(bf16) * scale
    """
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)

def quantize_fp8_per_token(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Dynamic per-token FP8 quantization.

    For q with shape [total_q, num_heads, head_dim], each token (row in total_q)
    gets an independent scale across [num_heads, head_dim].

    Returns:
        (fp8_tensor, scales) where scales has shape [total_q] (float32).
        Dequantize per token i: fp8_tensor[i].to(bf16) * scales[i]
    """
    if tensor.ndim != 3:
        raise ValueError(f"Expected 3D tensor [total_q, nhead, dim], got {tuple(tensor.shape)}")
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax(dim=(1, 2), keepdim=True).clamp(min=1e-12)   # [total_q, 1, 1]
    scales = (amax / finfo.max).to(torch.float32)                          # [total_q, 1, 1]
    fp8_tensor = (tensor / scales).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scales.reshape(-1).contiguous()                     # [total_q]


# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
# Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
# ---------------------------------------------------------------------------

def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.

    Block size = 32. Each block gets an E8M0 scale factor.
    Two FP4 E2M1 values are packed per byte.

    Args:
        tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)

    Returns:
        (fp4_data, scale_e8m0)
        - fp4_data:   shape [B, M, N//2] in aiter_dtypes.fp4x2
        - scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
    """
    orig_shape = tensor.shape  # (B, M, N)
    B, M, N = orig_shape

    # dynamic_mxfp4_quant expects 2D: (B*M, N)
    tensor_2d = tensor.reshape(B * M, N)
    fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)

    # Reshape fp4_data back to 3D: (B, M, N//2)
    fp4_data = fp4_data_2d.view(B, M, N // 2)

    return fp4_data, scale_e8m0


def dequantize_mxfp4(
    fp4_data: torch.Tensor,
    scale_e8m0: torch.Tensor,
    orig_shape: tuple,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    """
    Dequantize MXFP4 tensor using aiter utilities.

    Note: dynamic_mxfp4_quant may pad both row and block dimensions in scale_e8m0.
    We trim scales to match the actual data dimensions.

    Args:
        fp4_data:   packed FP4 data, shape [B, M, N//2] in fp4x2 or uint8
        scale_e8m0: E8M0 block scale factors (possibly padded) in fp8_e8m0
        orig_shape: original (B, M, N) for reshaping
        dtype:      output dtype

    Returns:
        Dequantized tensor of shape orig_shape.
    """
    B, M, N = orig_shape
    num_rows = B * M
    block_size = 32
    num_blocks = N // block_size  # actual blocks needed (e.g. 576/32 = 18)

    # Unpack FP4 to float32: mxfp4_to_f32 expects (..., N//2) -> (..., N)
    fp4_data_2d = fp4_data.reshape(num_rows, N // 2)
    float_vals = mxfp4_to_f32(fp4_data_2d)  # (num_rows, N)

    # Convert E8M0 scales to float32 and trim padded dimensions
    scale_f32 = e8m0_to_f32(scale_e8m0)  # (padded_rows, padded_blocks)
    scale_f32 = scale_f32[:num_rows, :num_blocks]  # (num_rows, num_blocks)

    # Apply block scales
    float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
    scaled = float_vals_blocked * scale_f32.unsqueeze(-1)

    return scaled.view(B, M, N).to(dtype)


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


def _record_profile_sample(key: tuple, q_ms: float, meta_ms: float, kernel_ms: float) -> None:
    stats = _PROFILE_STATS.setdefault(
        key,
        {"count": 0.0, "q_ms": 0.0, "meta_ms": 0.0, "kernel_ms": 0.0},
    )
    stats["count"] += 1
    stats["q_ms"] += q_ms
    stats["meta_ms"] += meta_ms
    stats["kernel_ms"] += kernel_ms


def _print_profile_sample(key: tuple, q_ms: float, meta_ms: float, kernel_ms: float) -> None:
    batch_size, q_seq_len, total_kv_len, experiment = key
    total_ms = q_ms + meta_ms + kernel_ms
    print(
        f"MLA profile sample "
        f"shape(batch={batch_size}, q={q_seq_len}, kv={total_kv_len}, experiment={experiment}): "
        f"q={q_ms:.3f} ms, meta={meta_ms:.3f} ms, kernel={kernel_ms:.3f} ms, total={total_ms:.3f} ms",
        flush=True,
    )


def _print_profile_summary() -> None:
    if not _PROFILE_STATS:
        return
    print("MLA profile summary", flush=True)
    for key in sorted(_PROFILE_STATS):
        stats = _PROFILE_STATS[key]
        count = stats["count"]
        q_ms = stats["q_ms"] / count
        meta_ms = stats["meta_ms"] / count
        kernel_ms = stats["kernel_ms"] / count
        total_ms = q_ms + meta_ms + kernel_ms
        batch_size, q_seq_len, total_kv_len, experiment = key
        print(
            f"  shape(batch={batch_size}, q={q_seq_len}, kv={total_kv_len}, experiment={experiment}): "
            f"q={q_ms:.3f} ms, meta={meta_ms:.3f} ms, kernel={kernel_ms:.3f} ms, total={total_ms:.3f} ms",
            flush=True,
        )


if PROFILE_ENABLED:
    atexit.register(_print_profile_summary)


def _select_num_kv_splits(total_kv_len: int) -> int:
    """
    Heuristic split count for persistent MLA decode.

    Smaller contexts benefit from fewer splits to reduce reduction/setup overhead.
    Larger contexts usually need more splits to keep the kernel fed.
    """
    if total_kv_len <= 1024:
        return 8
    if total_kv_len <= 4096:
        return 16
    if total_kv_len <= 8192:
        return 32
    return 64


def _lookup_num_kv_splits(batch_size: int, total_kv_len: int) -> int:
    """
    Small lookup table for the known qseqlen=1 benchmark family.

    Higher batch sizes already provide more parallel work, so they often benefit
    from fewer KV splits to reduce reduction overhead.
    """
    tuned = {
        (4, 1024): 8,
        (4, 8192): 32,
        (32, 1024): 8,
        (32, 8192): 32,
        (64, 1024): 8,
        (64, 8192): 16,
        (256, 1024): 4,
        (256, 8192): 16,
    }
    return tuned.get((batch_size, total_kv_len), _select_num_kv_splits(total_kv_len))


def _lookup_reduce_tuned_num_kv_splits(batch_size: int, total_kv_len: int) -> int:
    """
    Aggressive split reduction for qseqlen=1 decode.

    The profiler points at mla_reduce_v1 as the dominant named hotspot, so this
    table biases toward fewer splits and relies more on batch-level parallelism.
    """
    tuned = {
        (4, 1024): 4,
        (4, 8192): 16,
        (32, 1024): 4,
        (32, 8192): 16,
        (64, 1024): 4,
        (64, 8192): 8,
        (256, 1024): 2,
        (256, 8192): 8,
    }
    return tuned.get((batch_size, total_kv_len), _select_num_kv_splits(total_kv_len))


def _resolve_num_kv_splits(batch_size: int, total_kv_len: int) -> int:
    strategy = EXPERIMENT["split_strategy"]
    if strategy == "fixed":
        return NUM_KV_SPLITS
    if strategy == "lookup":
        return _lookup_num_kv_splits(batch_size, total_kv_len)
    if strategy == "lookup_reduce_tuned":
        return _lookup_reduce_tuned_num_kv_splits(batch_size, total_kv_len)
    return _select_num_kv_splits(total_kv_len)


def _get_kv_indices(total_kv_len: int, device: torch.device) -> torch.Tensor:
    key = (_device_key(device), total_kv_len)
    kv_indices = _KV_INDICES_CACHE.get(key)
    if kv_indices is None:
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=device)
        _KV_INDICES_CACHE[key] = kv_indices
    return kv_indices


def _get_meta_work_tensors(
    batch_size: int,
    max_q_len: int,
    nhead: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    num_kv_splits: int,
    device: torch.device,
) -> list[torch.Tensor]:
    cache_key = (
        batch_size,
        max_q_len,
        nhead,
        str(q_dtype),
        str(kv_dtype),
        num_kv_splits,
        EXPERIMENT["fast_mode"],
        _device_key(device),
    )
    work = _META_WORK_CACHE.get(cache_key)
    if work is not None:
        return work

    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nhead, q_dtype, kv_dtype,
        is_sparse=False, fast_mode=EXPERIMENT["fast_mode"],
        num_kv_splits=num_kv_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device=device) for s, t in info]
    _META_WORK_CACHE[cache_key] = work
    return work


# ---------------------------------------------------------------------------
# Persistent mode metadata helpers
# ---------------------------------------------------------------------------

def _make_mla_decode_metadata(
    batch_size: int,
    max_q_len: int,
    total_kv_len: int,
    nhead: int,
    nhead_kv: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    num_kv_splits: int = NUM_KV_SPLITS,
):
    """
    Allocate and populate work buffers for persistent mla_decode_fwd.

    For this benchmark family, metadata depends only on decode shape, dtype, and
    split strategy, so cache the populated metadata tensors directly and skip
    rebuilding them on every call.
    """
    meta_cache_key = (
        batch_size,
        max_q_len,
        nhead,
        nhead_kv,
        total_kv_len,
        str(q_dtype),
        str(kv_dtype),
        num_kv_splits,
        EXPERIMENT["fast_mode"],
        _device_key(qo_indptr.device),
    )
    cached_meta = _META_CACHE.get(meta_cache_key)
    if cached_meta is not None:
        return cached_meta

    work = _get_meta_work_tensors(
        batch_size=batch_size,
        max_q_len=max_q_len,
        nhead=nhead,
        q_dtype=q_dtype,
        kv_dtype=kv_dtype,
        num_kv_splits=num_kv_splits,
        device=qo_indptr.device,
    )
    (work_metadata, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map) = work

    # Populate the metadata buffers
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        nhead // nhead_kv,   # num_heads_per_head_k
        nhead_kv,            # num_heads_k
        True,                # is_causal
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=EXPERIMENT["fast_mode"],
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    meta = {
        "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[meta_cache_key] = meta
    return meta


# ---------------------------------------------------------------------------
# Aiter reference kernel (decode only)
# ---------------------------------------------------------------------------

def _aiter_mla_decode(
    q: torch.Tensor,
    kv_buffer: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
    q_scale: torch.Tensor | None = None,
    kv_scale: torch.Tensor | None = None,
    profile_events: dict[str, torch.cuda.Event] | None = None,
) -> torch.Tensor:
    """
    MLA decode attention using aiter persistent-mode kernel.

    Supports multiple Q/KV dtype combinations:
      - Q_DTYPE="fp8":  fp8 Q + fp8 KV (a8w8) — fastest on MI355X
      - Q_DTYPE="bf16": bf16 Q + bf16 KV (a16w16) — highest precision

    q:          (total_q, num_heads, 576)  fp8 or bf16
    kv_buffer:  (total_kv, 1, 576)         fp8 or bf16
    q_scale:    scalar or per-token float32 scale (required for fp8 Q, None for bf16)
    kv_scale:   scalar float32 (required for fp8 KV, None for bf16)
    """
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    # Avoid synchronizing on kv_indptr[-1].item(); the flattened KV buffer shape
    # already gives the total sequence length for page_size=1 decode.
    total_kv_len = kv_buffer.shape[0]

    # Reshape kv_buffer to 4D for aiter: (total_kv, page_size, nhead_kv, dim)
    kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1])

    max_q_len = q_seq_len
    num_kv_splits = _resolve_num_kv_splits(batch_size, total_kv_len)
    kv_indices = _get_kv_indices(total_kv_len, qo_indptr.device)
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    if profile_events is not None:
        profile_events["meta_start"].record()
    meta = _make_mla_decode_metadata(
        batch_size, max_q_len, total_kv_len, nq, nkv,
        q.dtype, kv_buffer.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        num_kv_splits=num_kv_splits,
    )
    if profile_events is not None:
        profile_events["meta_end"].record()

    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device)
    if profile_events is not None:
        profile_events["kernel_start"].record()
    mla_decode_fwd(
        q.view(-1, nq, dq),
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        max_q_len,
        page_size=PAGE_SIZE,
        nhead_kv=nkv,
        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=True,
        **meta,
    )
    if profile_events is not None:
        profile_events["kernel_end"].record()
    return o

def custom_kernel(data: input_t) -> output_t:
    """Reference MLA decode attention using the active experiment preset."""
    q, kv_data, qo_indptr, kv_indptr, config = data
    profile_events = None
    if PROFILE_ENABLED and q.is_cuda:
        profile_events = {
            name: torch.cuda.Event(enable_timing=True)
            for name in (
                "q_start",
                "q_end",
                "meta_start",
                "meta_end",
                "kernel_start",
                "kernel_end",
            )
        }

    # Resolve Q
    if profile_events is not None:
        profile_events["q_start"].record()
    if EXPERIMENT["q_dtype"] == "fp8":
        if EXPERIMENT["q_fp8_scaling"] == "per_token":
            q_input, q_scale = quantize_fp8_per_token(q)
        else:
            q_input, q_scale = quantize_fp8_per_tensor(q)
    else:
        q_input, q_scale = q, None
    if profile_events is not None:
        profile_events["q_end"].record()

    # Resolve KV
    if EXPERIMENT["kv_dtype"] == "fp8":
        kv_buffer_fp8, kv_scale = kv_data["fp8"]
        kv_input = kv_buffer_fp8
    else:
        kv_input, kv_scale = kv_data["bf16"], None

    out = _aiter_mla_decode(
        q_input, kv_input, qo_indptr, kv_indptr, config,
        q_scale=q_scale, kv_scale=kv_scale, profile_events=profile_events,
    )
    if profile_events is not None:
        torch.cuda.synchronize(q.device)
        total_kv_len = int(kv_indptr[-1].item())
        profile_key = (
            config["batch_size"],
            config["q_seq_len"],
            total_kv_len,
            ACTIVE_EXPERIMENT,
        )
        q_ms = profile_events["q_start"].elapsed_time(profile_events["q_end"])
        meta_ms = profile_events["meta_start"].elapsed_time(profile_events["meta_end"])
        kernel_ms = profile_events["kernel_start"].elapsed_time(profile_events["kernel_end"])
        _record_profile_sample(
            profile_key,
            q_ms=q_ms,
            meta_ms=meta_ms,
            kernel_ms=kernel_ms,
        )
        _print_profile_sample(profile_key, q_ms=q_ms, meta_ms=meta_ms, kernel_ms=kernel_ms)
    return out
scrolls · 643 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