Skip to content
KernelIndex
Search⌘K

submission 738050

pingp2574 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738050?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
176.8µs
#556 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9ce1b6024c19a85928676d5c448e20515f10f616963462cc09da8249135f6b7b
license declaredunknown
license concludedunknown
authorspingp2574
imported2026-08-26

Techniques

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

fp44. _dq_bf16_buf: 反量化缓冲区(MXFP4 → bf16 时使用)
persistent-kernel获取或创建 MLA 解码内核所需的持久模式元数据(persistent-mode metadata)。

Kernel source

submission.py572 lines
"""
Submission 实现:MLA (Multi-head Latent Attention) 解码阶段的注意力计算。

=== 背景知识 ===

这是 DeepSeek R1 模型中使用的 MLA 注意力机制的解码(decode)阶段实现。
MLA 的核心思想是:将传统的 Multi-Head Attention 中的 KV 缓存进行压缩,
使用一个低秩的隐空间表示(latent representation),从而大幅减少 KV 缓存的显存占用。

=== MLA 的 forward_absorb 路径 ===

在 DeepSeek R1 中,MLA 采用了 "forward_absorb"(前向吸收)策略:
  - Q(Query)被投影到一个 576 维的吸收空间(absorbed space)
  - KV 缓存也被压缩到同样的 576 维空间
  - 576 = 512 (kv_lora_rank, 用于 V 输出) + 64 (qk_rope_head_dim, 用于旋转位置编码)

=== 数据流概览 ===

  输入: q (total_q, 16, 576) bf16  +  KV缓存 (total_kv, 1, 576) fp8
          ↓                              ↓
     量化为 FP8                     已预量化为 FP8
          ↓                              ↓
          └──────── mla_decode_fwd ────────┘
                       ↓
  输出: attention output (total_q, 16, 512) bf16

=== 本文件的策略 ===

本实现使用 FP8 精度进行注意力计算(a8w8 模式:Q 和 KV 都用 FP8),
相比 bf16 在 AMD MI355X 上可获得约 2-3 倍的加速,精度损失可忽略。

同时使用动态 KV 分片策略(num_kv_splits),根据 batch_size * kv_seq_len
的总工作量自适应调整分片数量,在保证精度的同时优化性能。
"""

import torch
from task import input_t, output_t

# aiter 是 AMD 的 AI 推理加速库,提供了高性能的 MLA 注意力内核
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 (
    mxfp4_to_f32,
    e8m0_to_f32,
)

# ========================================================================
# DeepSeek R1 MLA 模型超参数(forward_absorb 路径)
# 这些值来自 DeepSeek-R1 的 config.json 配置文件
# ========================================================================

# 注意力头数:16 个 Query 头
NUM_HEADS = 16

# KV 头数:只有 1 个(MLA 本质上是 Multi-Query Attention,所有 Q 头共享一个 KV 头)
# 这是 MLA 节省显存的关键之一
NUM_KV_HEADS = 1

# KV LoRA 秩:KV 压缩后的隐空间维度为 512
# 原始的 KV 维度远大于 512,通过低秩投影压缩到 512 维
KV_LORA_RANK = 512

# QK 旋转位置编码(RoPE)维度:64
# RoPE 是一种相对位置编码,只作用于 QK 计算的一部分维度
QK_ROPE_HEAD_DIM = 64

# QK 头总维度 = KV LoRA 秩 + RoPE 维度 = 512 + 64 = 576
# 576 维中:前 512 维用于 V 输出,全部 576 维用于 QK 匹配
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM

# V 头维度 = KV LoRA 秩 = 512
# 注意:V 只使用前 512 维,不包含 RoPE 的 64 维
V_HEAD_DIM = KV_LORA_RANK

# Softmax 缩放因子:1 / sqrt(d_k) = 1 / sqrt(576)
# 防止点积结果过大导致 softmax 梯度消失
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

# 页大小:PagedAttention 中的页大小为 1(即每个 KV token 独立存储)
PAGE_SIZE = 1

# FP8 数据类型(平台相关,由 aiter 库决定具体是 FP8 E4M3 还是 E5M2)
FP8_DTYPE = aiter_dtypes.fp8

# FP8 可表示的最大值,用于量化时的缩放计算
FP8_MAX = torch.finfo(FP8_DTYPE).max


class _BufferCache:
    """
    GPU 显存缓冲区缓存(Buffer Cache)。

    === 为什么需要这个类? ===

    在 LLM 推理的解码阶段,每次生成一个新 token 都需要调用注意力计算。
    如果每次调用都重新分配 GPU 显存,会带来巨大的开销(CUDA malloc 很慢)。

    这个类通过缓存和复用 GPU 张量来避免重复的显存分配:
    - 当请求的张量形状与之前相同时,直接复用已分配的显存
    - 当形状不同时,才重新分配

    === 缓存的缓冲区类型 ===

    1. _kv_indices: KV 索引数组,用于 PagedAttention
    2. _output_buf:  输出缓冲区,存储注意力计算结果
    3. _meta:        持久模式的元数据(work_metadata 等),用于优化内核调度
    4. _dq_bf16_buf: 反量化缓冲区(MXFP4 → bf16 时使用)
    """

    # __slots__ 禁止动态添加属性,节省内存并提高访问速度
    __slots__ = (
        "_meta_key", "_meta",
        "_kv_indices", "_kv_indices_len",
        "_output_buf", "_output_key",
        "_dq_bf16_buf", "_dq_bf16_key",
    )

    def __init__(self):
        self._meta_key = None
        self._meta = None
        self._kv_indices = None
        self._kv_indices_len = 0
        self._output_buf = None
        self._output_key = None
        self._dq_bf16_buf = None
        self._dq_bf16_key = None

    def get_kv_indices(self, total_kv_len: int) -> torch.Tensor:
        """
        获取 KV 缓存的索引数组 [0, 1, 2, ..., total_kv_len-1]。

        在 PagedAttention 中,kv_indices 用于指定要访问的 KV 缓存页。
        当 PAGE_SIZE=1 时,索引就是简单的连续整数序列。

        缓存策略:如果已有的数组长度 >= 需要的长度,直接切片复用;
        否则重新创建一个更大的数组。
        """
        if self._kv_indices is None or self._kv_indices_len < total_kv_len:
            self._kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
            self._kv_indices_len = total_kv_len
        return self._kv_indices[:total_kv_len]

    def get_output(self, total_q: int, nq: int, dv: int) -> torch.Tensor:
        """
        获取输出缓冲区,形状为 (total_q, nq, dv)。

        参数:
            total_q: 所有 batch 中 Q token 的总数
            nq:      注意力头数 (16)
            dv:      V 头维度 (512)

        返回: 未初始化的 bf16 张量(后续会被 mla_decode_fwd 写入结果)
        """
        key = (total_q, nq, dv)
        if self._output_key != key or self._output_buf is None:
            self._output_buf = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
            self._output_key = key
        return self._output_buf

    def get_metadata(
        self,
        batch_size, max_q_len, nhead, nhead_kv,
        q_dtype, kv_dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        num_kv_splits,
    ):
        """
        获取或创建 MLA 解码内核所需的持久模式元数据(persistent-mode metadata)。

        === 什么是持久模式(Persistent Mode)? ===

        在传统的 GPU 内核调用中,每次启动内核都需要重新计算调度信息。
        而在解码阶段,连续的 token 生成之间,很多调度参数是不变的。
        持久模式通过预先计算并缓存这些调度信息,避免重复计算,提高效率。

        === 元数据包含什么? ===

        - work_metadata:   工作分配元数据(哪个 GPU 线程块处理哪个任务)
        - workIndptr:      工作项的索引指针
        - workInfoSet:     工作项的详细信息
        - reduceIndptr:    归约操作的索引指针(用于 KV 分片后的结果合并)
        - reduceFinalMap:  最终归约的映射表
        - reducePartialMap: 部分归约的映射表

        === 缓存策略 ===

        使用 (batch_size, max_q_len, nhead, nhead_kv, dtype, num_kv_splits) 作为缓存键。
        如果键匹配,只重新填充元数据(不重新分配显存);
        如果键不匹配,需要重新分配元数据缓冲区。
        """
        key = (batch_size, max_q_len, nhead, nhead_kv, str(q_dtype), str(kv_dtype), num_kv_splits)

        if self._meta_key == key and self._meta is not None:
            (work_metadata, work_indptr, work_info_set,
             reduce_indptr, reduce_final_map, reduce_partial_map) = self._meta
            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=max(PAGE_SIZE, 16),
                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=True,
                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,
            }

        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=True,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t 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=max(PAGE_SIZE, 16),
            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=True,
            dtype_q=q_dtype,
            dtype_kv=kv_dtype,
        )

        self._meta = (work_metadata, work_indptr, work_info_set,
                      reduce_indptr, reduce_final_map, reduce_partial_map)
        self._meta_key = key

        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_dequant_bf16(self, total_kv: int, qk_head_dim: int) -> torch.Tensor:
        """
        获取用于 MXFP4 反量化的 bf16 缓冲区。

        当使用 MXFP4 格式的 KV 缓存时,需要先将 FP4 反量化为 bf16,
        这个缓冲区用于存储反量化的中间结果。

        参数:
            total_kv:    所有 batch 中 KV token 的总数
            qk_head_dim: QK 头维度 (576)
        """
        key = (total_kv, qk_head_dim)
        if self._dq_bf16_key != key or self._dq_bf16_buf is None:
            self._dq_bf16_buf = torch.empty(
                (total_kv, NUM_KV_HEADS, qk_head_dim),
                dtype=torch.bfloat16, device="cuda",
            )
            self._dq_bf16_key = key
        return self._dq_bf16_buf


# 创建全局单例缓存对象,在整个推理过程中复用
_cache = _BufferCache()


def _quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    动态逐张量 FP8 量化(dynamic per-tensor FP8 quantization)。

    === 量化原理 ===

    FP8(8-bit 浮点数)只有 8 位来表示数值,范围和精度都有限。
    为了最大化利用 FP8 的表示范围,我们需要:
    1. 找到输入张量的最大绝对值 (amax)
    2. 计算缩放因子 scale = amax / FP8_MAX
    3. 将每个元素除以 scale,映射到 FP8 的完整范围
    4. 裁剪到 FP8 可表示的范围,然后转换类型

    反量化公式:original ≈ fp8_tensor * scale

    参数:
        tensor: 待量化的 bf16 张量

    返回:
        (fp8_tensor, scale):
        - fp8_tensor: 量化后的 FP8 张量
        - scale: 标量缩放因子(float32),形状为 (1,)
    """
    # 找到整个张量的最大绝对值,clamp 确保不为零(避免除零)
    amax = tensor.abs().amax().clamp(min=1e-12)
    # 计算缩放因子:amax 映射到 FP8 的最大可表示值
    scale = amax / FP8_MAX
    # 量化:除以 scale → 裁剪到 FP8 范围 → 转换为 FP8 类型
    fp8_tensor = (tensor / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)


def _dequantize_mxfp4_inplace(
    fp4_data: torch.Tensor,
    scale_e8m0: torch.Tensor,
    out_bf16: torch.Tensor,
) -> torch.Tensor:
    """
    就地反量化 MXFP4 数据为 bf16。

    === MXFP4 格式说明 ===

    MXFP4 (Microscaling FP4) 是一种块级量化格式:
    - 每 32 个连续元素为一个块(block)
    - 每个元素只有 4 bit(FP4 E2M1 格式:2位指数 + 1位尾数 + 1位符号)
    - 每个块共享一个 E8M0 格式的缩放因子(8位纯指数,无尾数)
    - 两个 FP4 值打包(pack)在一个字节(byte)中

    因此对于 N 个元素:
    - fp4_data 的形状为 (rows, N//2),每个字节包含 2 个 FP4 值
    - scale_e8m0 的形状为 (rows, N//32),每 32 个元素共享 1 个缩放因子

    === 反量化过程 ===

    1. 将打包的 FP4 数据解包为 float32(mxfp4_to_f32)
    2. 将 E8M0 缩放因子转换为 float32(e8m0_to_f32)
    3. 将数据按块 reshape,乘以对应的缩放因子
    4. 将结果复制到 bf16 输出缓冲区

    参数:
        fp4_data:   打包的 FP4 数据,形状 (rows, N//2)
        scale_e8m0: E8M0 块缩放因子,形状 (rows, num_blocks)
        out_bf16:   预分配的 bf16 输出缓冲区,形状 (total_kv, 1, N)
    """
    total_kv = fp4_data.shape[0]
    N = out_bf16.shape[-1]           # 原始维度 = 576
    num_rows = total_kv
    block_size = 32                   # MXFP4 的块大小固定为 32
    num_blocks = N // block_size      # 576 / 32 = 18 个块

    # 步骤 1: 将打包的 FP4 数据解包为 float32
    # fp4_data 形状 (rows, N//2=288) → float_vals 形状 (rows, N=576)
    fp4_data_2d = fp4_data.view(num_rows, N // 2)
    float_vals = mxfp4_to_f32(fp4_data_2d)

    # 步骤 2: 将 E8M0 缩放因子转换为 float32
    # scale_e8m0 可能包含填充(padding),需要裁剪到实际需要的行数和块数
    scale_f32 = e8m0_to_f32(scale_e8m0)
    scale_f32 = scale_f32[:num_rows, :num_blocks]

    # 步骤 3: 将数据按块 reshape,然后乘以对应的缩放因子
    # float_vals: (rows, 576) → (rows, 18, 32)
    # scale_f32:  (rows, 18) → (rows, 18, 1)  通过 unsqueeze 广播
    float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
    float_vals_blocked.mul_(scale_f32.unsqueeze(-1))  # 就地乘法,节省显存

    # 步骤 4: 将结果复制到 bf16 输出缓冲区
    # (rows, 18, 32) → (rows, 576) → (rows, 1, 576)
    out_bf16.copy_(float_vals_blocked.view(num_rows, 1, N))
    return out_bf16


def _get_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    """
    根据工作量动态计算 KV 分片数量(num_kv_splits)。

    === 什么是 KV 分片? ===

    在 FlashAttention 等高效注意力算法中,KV 序列可以被分成多个片段(split),
    每个 GPU 线程块独立处理一个片段,最后通过归约(reduction)合并结果。

    分片数量的权衡:
    - 分片太少:GPU 利用率低,性能差
    - 分片太多:归约开销增大,且可能引入精度损失

    === 动态策略 ===

    根据总工作量 total_work = batch_size × kv_seq_len 来决定:
    - 工作量小(≤ 4096):不分片(1),因为数据量小,单线程块就够
    - 工作量中等:逐步增加分片数(2, 4, 8, 16)
    - 工作量大(> 262144):最多 32 个分片

    注意:reference.py 中固定使用 32 个分片,而 submission.py 使用动态策略,
    在工作量较小时可以减少不必要的归约开销。
    """
    total_work = batch_size * kv_seq_len
    if total_work <= 4096:
        return 1
    elif total_work <= 16384:
        return 2
    elif total_work <= 65536:
        return 4
    elif total_work <= 131072:
        return 8
    elif total_work <= 262144:
        return 16
    else:
        return 32


def _run_a8w8(
    q_fp8: torch.Tensor,
    q_scale: torch.Tensor,
    kv_fp8: torch.Tensor,
    kv_scale: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
    num_kv_splits: int,
) -> torch.Tensor:
    """
    使用 A8W8(Activation-8bit, Weight-8bit)模式运行 MLA 解码注意力。

    === A8W8 模式 ===

    "A8W8" 表示 Q(激活)和 KV(权重)都使用 8-bit 浮点数(FP8)。
    这是 aiter 库中最快的注意力计算模式。

    由于 FP8 的范围有限,需要提供缩放因子(scale)来恢复原始数值范围:
    - q_scale:  Q 的缩放因子,用于在计算中恢复 Q 的数值范围
    - kv_scale: KV 的缩放因子,用于恢复 KV 的数值范围

    === 注意力计算流程 ===

    1. 将 KV 缓存 reshape 为 4D 格式(total_kv, page_size, nhead_kv, dim)
    2. 准备 KV 索引和元数据
    3. 调用 mla_decode_fwd 执行核心注意力计算
    4. 返回输出张量

    参数:
        q_fp8:        FP8 量化的 Query,形状 (total_q, num_heads, qk_head_dim)
        q_scale:      Q 的缩放因子,标量 float32
        kv_fp8:       FP8 量化的 KV 缓存,形状 (total_kv, 1, 576)
        kv_scale:     KV 的缩放因子,标量 float32
        qo_indptr:    Query/Output 的索引指针,形状 (batch_size+1,)
                      类似 CSR 格式的行指针,qo_indptr[i]:qo_indptr[i+1] 是第 i 个 batch 的 Q
        kv_indptr:    KV 的索引指针,形状 (batch_size+1,)
                      kv_indptr[i]:kv_indptr[i+1] 是第 i 个 batch 的 KV
        config:       包含注意力参数的配置字典
        num_kv_splits: KV 分片数量
    """
    # 从配置字典中提取参数
    batch_size = config["batch_size"]      # 批大小
    nq = config["num_heads"]               # Q 头数 = 16
    nkv = config["num_kv_heads"]           # KV 头数 = 1
    dq = config["qk_head_dim"]             # QK 维度 = 576
    dv = config["v_head_dim"]              # V 维度 = 512
    q_seq_len = config["q_seq_len"]        # Q 序列长度(解码时通常为 1)
    total_kv_len = int(kv_indptr[-1].item())  # 所有 batch 的 KV token 总数

    # 将 KV 缓存从 3D (total_kv, 1, 576) reshape 为 4D (total_kv, page_size, nhead_kv, 576)
    # 这是 aiter mla_decode_fwd 要求的输入格式
    kv_buffer_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])

    # 获取 KV 索引数组(从缓存中获取,避免重复分配)
    kv_indices = _cache.get_kv_indices(total_kv_len)

    # 计算每个 batch 的最后一个页的有效长度
    # 由于 PAGE_SIZE=1,每个 batch 的 kv_last_page_len 就是该 batch 的 KV 长度
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    max_q_len = q_seq_len

    # 获取或创建持久模式的元数据(从缓存中获取)
    meta = _cache.get_metadata(
        batch_size, max_q_len, nq, nkv,
        q_fp8.dtype, kv_fp8.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        num_kv_splits,
    )

    # 获取输出缓冲区(从缓存中获取,避免重复分配)
    o = _cache.get_output(q_fp8.shape[0], nq, dv)

    # ====================================================================
    # 核心调用:aiter MLA 解码前向内核
    # ====================================================================
    #
    # mla_decode_fwd 是 AMD aiter 库提供的高性能 MLA 解码注意力内核。
    # 它在 GPU 上执行以下计算:
    #
    #   O_i = softmax(Q_i @ K^T / sqrt(d)) @ V
    #
    # 其中:
    #   Q_i: 第 i 个 query token 的吸收表示 (16 heads × 576 dim)
    #   K:   所有 KV 缓存的键部分 (576 dim),前 512 维 + 64 维 RoPE
    #   V:   所有 KV 缓存的值部分 (前 512 维)
    #   O_i: 输出 (16 heads × 512 dim)
    #
    mla_decode_fwd(
        q_fp8.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,
    )
    return o


def custom_kernel(data: input_t) -> output_t:
    """
    主入口函数:接收输入数据,执行 MLA 解码注意力计算。

    === 调用流程 ===

    1. 解包输入数据(Q, KV 缓存, 索引指针, 配置)
    2. 根据工作量计算 KV 分片数量
    3. 将 Q 从 bf16 量化为 FP8
    4. 从 KV 数据中提取 FP8 格式的 KV 缓存(已预量化)
    5. 调用 A8W8 内核执行注意力计算

    === 输入数据结构 ===

    data = (q, kv_data, qo_indptr, kv_indptr, config)

    - q:          (total_q, 16, 576) bf16 — 吸收后的 Query
    - kv_data:    字典,包含三种精度的 KV 缓存:
        - "bf16":  (total_kv, 1, 576) bf16 — 最高精度
        - "fp8":   (kv_fp8, scale) — FP8 量化格式(本实现使用这个)
        - "mxfp4": (fp4_data, scale) — MXFP4 量化格式(备选)
    - qo_indptr:  (batch_size+1,) int32 — Q/Output 的索引指针
    - kv_indptr:  (batch_size+1,) int32 — KV 的索引指针
    - config:     配置字典(batch_size, 维度信息等)

    === 输出 ===

    output: (total_q, 16, 512) bf16 — 注意力计算结果
    """
    q, kv_data, qo_indptr, kv_indptr, config = data

    num_kv_splits = 32

    q_fp8, q_scale = _quantize_fp8(q)
    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    return _run_a8w8(
        q_fp8, q_scale,
        kv_buffer_fp8, kv_scale,
        qo_indptr, kv_indptr,
        config, num_kv_splits,
    )
scrolls · 572 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