oj教程更新

This commit is contained in:
xiao 2026-06-18 10:21:54 +08:00
parent 6b2e4a709c
commit d4ac7f5a2f
25 changed files with 1783 additions and 90 deletions

View File

@ -1,20 +1,26 @@
# FlashAttention Baseline 入门从环境验证到 KV-Cache Benchmark 
# Flashattention 迁移 Baseline 实战从性能基线到 XPU-OJ 评测
## 一、教程定位
本教程是参赛训练课程的 FlashAttention Baseline 入门模块主要帮助用户快速跑通 FlashAttention  KV-Cache 推理性能基准测试最小可运行流程
本教程是参赛训练课程的 **FlashAttention Baseline 入门与评测提交衔接** 模块主要帮助用户跑通 FlashAttention paged KVcache 推理核函数 `flash_attn_with_kvcache` 的基准测试流程理解 baseline 的输入输出、性能指标和评测含义并基于 XPUOJ 题包完成一个最小正确版 `run_kernel` 的实现与提交
完成本教程后,用户应能够:
需要特别说明本教程中的 Baseline 主要用于帮助参赛者理解目标算子的调用方式、输入输出结构和性能基线Baseline 不是最终提交物。最终评测以 XPUOJ 题包为准参赛者需要根据题包中的接口约定实现自己的 `run_kernel`并在输出结果对齐 baseline 参考结果的前提下提升性能。
* 完成环境验证和依赖检查
完成本教程后,学员应能够:
* 完成环境验证与依赖检查
* 理解并配置基准测试参数
* 理解并配置 KVCache Benchmark 的核心参数
* 明确Baseline的解读理解为什么KV-Cache是性能瓶颈
* 明确 Baseline 的含义理解 KVCache 为何成为推理性能瓶颈
* 运行 KV-Cache Benchmark 测试
* 运行 `flash_attn_with_kvcache`  Benchmark 测试并获取性能数据
* 输出一份 baseline 性能结果记录表为后续算子优化提供对比基准
* 输出一份 Baseline 性能结果记录表为后续算子优化提供对比基准
* 理解 XPUOJ 评测的 `run_kernel` 接口规范与精度要求
* 基于 OJ 题包接口实现一个最小正确版 `run_kernel`,通过正确性校验
> Baseline 解读为什么本教程基于 KV Cache 做性能基线
@ -125,17 +131,20 @@
完成本模块后,你将能够:
1. 理解 FlashAttention `flash_attn_with_kvcache` 核函数的基本作用与应用场景;
1. **理解 FlashAttention Paged KVCache 算子的作用**
明白 `flash_attn_with_kvcache`  LLM 推理 decode 阶段如何高效利用分页 KV Cache减少显存碎片并提升吞吐。
2. 利用预装专属镜像完成沐曦 GPU 硬件环境的快速验证
2. **完成环境准备与 Benchmark 运行**
安装所需依赖,运行 `benchmark_kvcache.py`生成包含执行时间与显存带宽的 CSV 性能记录。
3. 深入理解 Baseline 的概念掌握性能测试与正确性测试的联系与区别
3. **读懂 Baseline 并定位性能瓶颈**
分析不同 `batch_size`、`seq_len_kv` 下的带宽曲线理解显存带宽对 decode 阶段的影响。
4. 跑通 KV-Cache Benchmark 基准测试脚本
4. **理清 Baseline  OJ 题包的关系**
明确基准脚本用于性能参照OJ 题包定义最终提交接口、数据范围和精度校验标准。
5. 完成多种 `batch_size × seq_len_kv` 组合的性能测试;
6. 输出带宽性能结果 CSV 文件并进行结果分析。
5. **实现并提交一个最小正确版** `**run_kernel**`
根据题包中的接口约定编写 CUDA 算子通过 OJ 正确性校验并记录首次提交耗时。
---
@ -163,37 +172,41 @@
---
## 四、前置准备
开始实战前,请确认你已经完成以下准备:
### 环境准备
* **设置领取与兑换算力券**
* 前往沐曦开发者社区注册账号并完成邮箱验证申请并获取 MACA 算力代金券兑换码。
1. 前往沐曦开发者社区注册账号并完成邮箱验证申请并获取 MACA 算力代金券兑换码。
链接:[https://developer.metax-tech.com/activities/6](https://developer.metax-tech.com/activities/6)
链接https://developer.metax-tech.com/activities/6
* 登录 模力方舟平台 (Gitee AI),在“费用中心 -> 算力券”页面输入兑换码完成充值。
链接:[https://ai.gitee.com/](https://ai.gitee.com/)
2. 登录 模力方舟平台 (Gitee AI),在“费用中心 -> 算力券”页面输入兑换码完成充值。 链接https://ai.gitee.com/
* **创建并启动实例**
* 进入 算力市场筛选“沐曦”芯片厂商选择合适的 GPU 规格推荐 曦云 C500 节点
* **关键配置:** 在预装镜像处,务必选择专属开发镜像(`PyTorch Agent/2.8.0/Python 3.12/maca 3.7.2.1`)。
* 创建完成后,进入算力容器,点击“工具-lab”即可打开 JupyterLab 终端开始项目创作。
**说明:**由于本次使用的是预装的专属镜像环境中已经默认安装并配置好了 PyTorch、FlashAttention、einops 等所有依赖包。因此在启动实例后无需再进行繁琐的依赖库版本验证即可直接进入测试环节。
1. 进入 算力市场筛选“沐曦”芯片厂商选择合适的 GPU 规格推荐 曦云 C500 节点
2. 关键配置 在预装镜像处务必选择专属开发镜像PyTorch Agent/2.8.0/Python 3.12/maca 3.7.2.1)。
3. 创建完成后,进入算力容器,点击“工具-lab”即可打开 JupyterLab 终端开始项目创作。
\*\*说明:\*\*由于本次使用的是预装的专属镜像环境中已经默认安装并配置好了 PyTorch、FlashAttention、einops 等所有依赖包。因此在启动实例后无需再进行繁琐的依赖库版本验证即可直接进入测试环节。
### 代码准备
* 已获取基准测试脚本 `benchmark_kvcache.py`
* 获取目标源码(包含 `benchmark_kvcache.py`  OJ 题包
* 已进入项目目录 `/data/flashattn_baseline`
* 准备 Benchmark 脚本与 OJ 测试脚本
---
@ -201,69 +214,63 @@
### 关键术语
* **KV-Cache**缓存历史 Token  Key/Value 向量避免 Transformer 推理时重复计算。
* **KV-Cache**缓存历史 Token  Key/Value 向量避免 Transformer 推理时重复计算。
* **Paged KV-Cache** KV-Cache 分页管理减少显存碎片提高利用率。
* **Paged KV-Cache**将 KV-Cache 分页管理减少显存碎片提高利用率。
* **Batch Size**一次处理的样本数,越大并行度越高,但显存占用越大。
* **Batch Size**一次处理的样本数,越大并行度越高,但显存占用越大。
* **seq\_len\_kv**KV-Cache 中已缓存的历史 Token 数量。
* **seq\_len\_kv**KV-Cache 中已缓存的历史 Token 数量。
* **headdim**注意力头维度常见为 64/128/256。
* **headdim**注意力头维度常见为 64/128/256。
### 核心知识
* **正确性测试 :**验证算子输出结果的数学精度是否与标准实现一致,这是**绝对底线**
* **正确性测试 :**验证算子输出结果的数学精度是否与标准实现一致,这是绝对底线。
* **性能测试 :**在正确的前提下,测算速度与吞吐。通过建立 **Baseline基线**才能量化后续每次代码修改带来的真实收益(加速比)。
* **性能测试 :**在正确的前提下测算速度与吞吐。通过建立Baseline基线才能量化后续每次代码修改带来的真实收益加速比
* **Benchmark 基准测试**在固定条件下反复运行同一任务,获取可重复的性能指标,用于建立基线、量化优化效果和定位瓶颈。
### 关键指标
* **XPUOJ**:比赛官方在线评测平台,最终评测会调用参赛者提交代码中的 `run_kernel`
* **Kernel 执行时间**GPU 核函数运行耗时ms使用 GPU 端同步计时获得。
### 关键指标
* **Kernel 执行时间**GPU 核函数运行耗时ms使用 GPU 端同步计时获得。
* **有效带宽:**`数据传输量 (GB) ÷ Kernel 时间 (s)`,越接近理论峰值说明显存带宽利用越充分。
* **有效带宽:**数据传输量 (GB) ÷ Kernel 时间 (s),越接近理论峰值说明显存带宽利用越充分。
### 其他要点
* **Warmup**预热若干次不记录使 GPU 进入稳定状态。
* **Warmup**预热若干次不记录使 GPU 进入稳定状态。
* **Repeat**正式运行多次,取平均值或中位数以消除波动。
* **Repeat**正式运行多次,取平均值或中位数以消除波动。
* **同步**调用 `torch.cuda.synchronize()` 确保精确计时。
* **同步**调用 torch.cuda.synchronize() 确保精确计时。
* **数据类型**本教程使用 `bfloat16`,在精度和性能取得平衡。
* **数据类型**本教程使用 bfloat16在精度和性能取得平衡。
* **显存占用估算**`KV-Cache  batch × seq_len_kv × num_heads_k × headdim × 2(K+V) × 字节数`
* **显存占用估算**KV-Cache  batch × seq\_len\_kv × num\_heads\_k × headdim × 2(K+V) × 字节数。
* **OOM 应对**减小 batch/seq\_len\_kv、使用更小 dtype 或释放中间变量。
* **OOM 应对**减小 batch/seq\_len\_kv、使用更小 dtype 或释放中间变量。
* **Tensor Core**现代 GPU含沐曦 C500的矩阵乘法专用单元要求维度对齐为 8  16 的倍数。
* **Tensor Core**现代 GPU含沐曦 C500的矩阵乘法专用单元要求维度对齐为 8  16 的倍数。
---
## 六、项目实践FlashAttention KV-Cache Benchmark
## 六、项目实践FlashAttention KVCache Benchmark & OJ 评测
### 项目目标
 FlashAttention  paged KV-cache 推理核函数`flash_attn_with_kvcache`)进行自动化性能基准测试,覆盖多种 `batch_size × seq_len_kv` 组合,输出执行时间和有效显存带宽。
基准测试脚本 `benchmark_kvcache.py` 的最小闭环包括:
1. 准备 paged KV-cache 张量和 block table。
2. 通过 `torch.profiler`  kernel 进行计时。
3. 计算有效显存带宽GB/s
4. 将结果写入带时间戳的 CSV 文件。
* 对 FlashAttention  paged KV-cache 推理核函数`flash_attn_with_kvcache`)进行自动化性能基准测试,覆盖多种 `batch_size × seq_len_kv` 组合,输出执行时间和有效显存带宽。
---
* 理解 XPUOJ 题包的 `run_kernel` 接口实现一个最小正确版 CUDA 算子通过所有 OJ 正确性测试用例。
### 步骤 0进入创建的实例环境
@ -271,7 +278,7 @@
选择工具-lab进入实例环境
![lab enter instance environment](https://origin.picgo.net/2026/06/04/lab-enter-instance-environment67995d391603a686.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/b5d7d783-3106-4a4d-97f5-c7ef4d7fa537.png)
### 步骤 1检查运行环境
@ -279,7 +286,7 @@
在JupyterLab Terminal中检查运行环境的配置。
![jupyterlab terminal check](https://origin.picgo.net/2026/06/04/jupyterlab-terminal-check143d7d42426453cf.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/339ac3f8-e31d-43c6-992b-07e20e35ef95.png)
**操作:** 检查 GPU 状态、Python 版本和依赖版本。
@ -307,22 +314,22 @@ python -c "import einops; print('einops OK')"
* `mx-smi` 显示沐曦 GPU 信息
![result mx smi](https://origin.picgo.net/2026/06/04/result-mx-smif86a3bed6681382e.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/016119b5-960b-478d-9119-af27b9f3c727.png)
* Python 版本 = 3.8
![result python](https://origin.picgo.net/2026/06/04/result-pythond15f856ddb84c649.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/568886e4-85a7-41c6-a94c-89a09913e0f7.png)
* `torch.cuda.is_available()` 返回 `True`
![result gpu available](https://origin.picgo.net/2026/06/04/result-gpu-available44e049addf5638fb.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/d1c97030-d40f-416f-89b3-b81caff88966.png)
* 所有依赖版本符合要求
![result dependency version](https://origin.picgo.net/2026/06/04/result-dependency-versionc17845898ee89434.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/352b41ee-8a87-405e-991b-1c3714bdeff0.png)
**常见问题:**
@ -336,8 +343,45 @@ python -c "import einops; print('einops OK')"
### 步骤 2进入项目目录
**目标:** 进入本模块所需的项目目录。
目标:进入本模块所需的源码目录。
<<<<<<< HEAD
1. 克隆代码仓库
```Bash
git clone https://gitlink.org.cn/metax-maca/op_optimization.git
```
2. 准备flashattn\_baseline
在仓库目录 `op_optimization/基于AI Agent开发范式的国产GPU大模型推理算子库优化` 下,找到 `flashattn_baseline` 文件夹。可以将 `flashattn_baseline` 整个目录复制到工作目录 `data/` 下。
**下一步操作:** 切换到 Flashattn\_Baseline 项目目录。
**命令示例:**
```bash
cd flashattn_baseline/Flashattn_Baselinels -la
```
**预期结果:**
```Plain
total xx
drwxr-xr-x 2 root root 4096 Jun 1 09:00 __MACOSX
- rw-r--r-- 1 root root 5232 Jun 1 09:00 benchmark_kvcache.py
- rw-r--r-- 1 root root 1440 Jun 1 09:00 benchmark_kvcache_20260526_150953.csv
...
```
3. 将Flashattn\_Baseline文件加入JupyterLab。
=======
1. 克隆代码仓库
```Bash
@ -347,9 +391,12 @@ python -c "import einops; print('einops OK')"
2. 准备flashattn_baseline
在仓库目录 `op_optimization/基于AI Agent开发范式的国产GPU大模型推理算子库优化` 下,找到 `flashattn_baseline` 文件夹。可以将 `flashattn_baseline` 整个目录复制到工作目录 `data/` 下。
>>>>>>> upstream/master
![flashattn baseline extract](https://origin.picgo.net/2026/06/04/flashattn-baseline-extract66d40632f4253973.png)
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/YdgOk2bRrmLe7q4B/img/3eb4d4fe-e948-405e-8d11-6301b4d8f16e.png)
<<<<<<< HEAD
=======
**操作:** 切换到基准测试脚本所在目录。
**命令示例:**
@ -374,6 +421,7 @@ drwxr-xr-x 2 root root 4096 Jun 1 09:00 __MACOSX
...
```
>>>>>>> upstream/master
---
### 步骤 3配置基准测试参数
@ -550,40 +598,434 @@ repeat = 200
ms = run_with_profiler(run_fn, warmup=warmup, reps=repeat, print_result=True, target_kernels=["flash"])
```
---
## 七、Agent 使用说明
### 步骤 7:从 Baseline  XPU-OJ 提交
在本模块中Agent 可用于以下场景
Baseline benchmark 用于理解目标算子的调用方式、输入输出 shape 和性能基线;XPU-OJ 题包用于定义最终评测接口、数据范围、参考输出和精度要求。
### Prompt 模板
跑完 baseline 后,选手需要完成以下转换:
**环境验证:**
1. 从 benchmark 脚本中理解目标 API,本任务对应的是 `flash_attn.flash_attn_interface` 中的 `flash_attn_with_kvcache`(paged KV cache 布局);
2. 在对应 OJ 题包中查看 `run_kernel(...)` 接口;
3. 对照题包中的输入 shape、数据范围和精度要求;
4. 编写自己的 `run_kernel(...)`;
5. 提交 OJ,先通过正确性;
6. 正确性通过后,再对比 baseline / OJ 耗时继续优化。
```Plain
请帮我验证当前环境是否满足 FlashAttention KV-Cache Benchmark 的运行要求,包括:
1. 沐曦 GPU 是否可见
2. PyTorch 版本和 CUDA 支持
3. flash-attn 和 einops 是否已安装
### 题目说明(FlashAttention KV Cache Decode)
> 注意:每个子题的接口参数、数据范围和精度要求可能不同,正式要求以对应 XPU-OJ 题包为准。本节以 **FlashAttention KV Cache Decode** 为例,演示从 baseline benchmark  XPU-OJ 提交的完整流程。
* **对应 baseline 脚本**:`flashattn_baseline/baseline/benchmark_kvcache.py`
* **对应 FlashAttention API**:`flash_attn_with_kvcache`(paged KV cache 版本)
* **算子说明**:实现 paged KV cache 下的 decode 注意力,每个 batch 只有 1  query token,KV cache  page 存储,长度由 `seqlen_k` 决定。
* **对应 OJ 题包**:XPU-OJ  `FlashAttention KV Cache Decode` 题(题号 20005)。
### 步骤 8:理解 XPU-OJ 评测接口与精度要求
**目标**:明确 Baseline 与最终评测提交之间的关系,理解选手需要实现什么。
完成 baseline benchmark 后,需要注意:baseline 脚本主要用于建立性能基线,**不是最终提交物**。最终评测以 XPU-OJ 题包为准,评测程序会调用选手提交代码中的 `run_kernel`,并将输出结果与 baseline 参考结果进行比较。
下面以 **FlashAttention KV Cache Decode** 题为例,题包目录中通常包含以下文件:
* `zh_CN/00_题目描述.md`:说明需要实现的算子功能;
* `zh_CN/01_接口约定.cuda.md`:说明必须实现的 `run_kernel` 函数签名;
* `zh_CN/05_数据范围与提示.md`:说明测试范围和精度要求;
* `testcase_config.py`:定义测试数据生成、baseline 参考实现和正确性校验方式。
#### 1. 必须实现的接口
选手需要在提交的 CUDA 源码中提供如下 C 符号,函数名、参数类型、顺序必须完全一致,并使用 `extern "C"` 防止 name mangling:
```cpp
#include <stdint.h>
#include <cuda_bf16.h>
extern "C" void run_kernel(
const __nv_bfloat16* q,
const __nv_bfloat16* k_cache_paged,
const __nv_bfloat16* v_cache_paged,
__nv_bfloat16* output,
const int32_t* cache_seqlens,
const int32_t* block_table,
int64_t batch_size,
int64_t seqlen_k,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t num_blocks,
int64_t causal
);
```
**参数配置建议:**
#### 2. 参数说明
```Plain
我需要测试 headdim=128 和 headdim=256 的性能差异,请帮我推荐合适的 batch_sizes 和 seq_lens_kv 扫描范围。
| 参数 | 说明 |
| --- | --- |
| `q` | decode query tensor,shape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16` |
| `k_cache_paged` | paged key cache,shape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16` |
| `v_cache_paged` | paged value cache,shape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16` |
| `output` | 输出缓冲区,shape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16` |
| `cache_seqlens` | 每个 batch  KV 长度,shape `(batch_size)`,连续 `int32` |
| `block_table` | 每个 batch  page 映射表,shape `(batch_size, num_blocks / batch_size)`,连续 `int32` |
| `seqlen_q` | query 长度,评测中固定为 `1` |
| `page_block_size` | page size,评测中固定为 `16` |
| `causal` | 是否启用 causal mask,评测中固定为 `0` |
`run_kernel` 内部需要自行计算合适的 launch 配置并启动 CUDA kernel。为保证计时准确,**不建议在** `**run_kernel**` **内部做** `**cudaDeviceSynchronize()**` **或显式同步**
#### 3. KV cache 布局
KV cache layout 固定为 `flash_attn_with_kvcache`  paged cache 布局:`(num_blocks, page_block_size, num_heads_k, headdim)`。
第 `t`  KV token 位于 `block_table[batch_idx, t / page_block_size]` 指向的物理 page 中,page 内偏移为 `t % page_block_size`
例如 `batch_size = 1`、`seqlen_k = 512`、`page_block_size = 16` 时,每个序列需要访问 `32` 个有效 page。
#### 4. 评测数据范围
`testcase_config.py` 中定义了本题的测试配置(摘自题包):
```python
HEAD_DIMS = [128]
BATCH_SIZES = [1, 4, 16]
SEQ_LENS_KV = [1024, 4096, 8192, 16384]
SEQ_LEN_Q = 1
NUM_HEADS = 8
NUM_HEADS_K = 8
PAGE_BLOCK_SIZE = 16
CAUSAL = 0
```
**结果分析:**
`num_blocks = max(1024, ceil(seqlen_k / page_block_size) * batch_size * 3)`
```Plain
请帮我分析这份 benchmark 结果 CSV 文件,找出峰值带宽配置和 OOM 边界。
#### 5. 精度要求
当前 `FlashAttention KV Cache Decode` 题的校验方式为:
```python
torch.allclose(output_t.float(), output_ref.float(), rtol=1e-2, atol=1e-2)
```
---
## 八、常见问题
也就是说,选手实现的输出需要在上述容差范围内与 baseline 输出一致。不同算子的容差可能不同,**正式精度要求以对应 OJ 题包说明为准**。
### 步骤 9登录 XPU-OJ 并进入题目页面
使用组委会统一发放的账号登录 XPU-OJ并进入对应赛题页面。
1.打开 XPU-OJ 平台https://xpuoj.com/
2.使用组委会统一发放的账号和初始密码登录【后续发布】;
3.登录后进入比赛 / 题目列表页面;
4.找到对应题目,例如 `20005 FlashAttention KV Cache Decode`
5.点击进入题目详情页,查看题目描述、接口约定、数据范围和提交入口。
### 步骤 10在Agent的帮助下提交 OJ 冒烟代码
**目标:**完成一次最小提交确认 OJ 提交链路、语言环境和 `run_kernel(...)` 接口可用。
**操作:**
1. 在语言下拉框中选择本题支持的提交语言,例如 `MXMACA C++`、`TileLang` 、 `Triton`
2. 将实现了题目要求接口的代码复制到提交框中;
如果你还没有 run\_kernel应该从哪里开始
OJ 最终评测不会直接运行 baseline 脚本而是调用你提交代码中的 `run_kernel(...)`。   本赛题鼓励参赛者使用 AI Agent 辅助完成代码阅读、接口理解、初版实现、错误定位和性能优化。  如果你还没有自己的 `run_kernel`可以先让 Agent 阅读题包并生成一个最小正确版实现思路。
3. 借助 Agent 从题包生成 run\_kernel 初版
### 10.1 在镜像终端中安装并启动 OpenCode
1. 首先返回镜像JupyterLab Terminal中, 打开容器内的 **Terminal**(终端)
2. 执行下面的命令安装 OpenCode:
```bash
curl -fsSL https://opencode.ai/install | bash
```
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/Lk3lbmbEQX2A6Om9/img/b4d06ad0-183b-47ce-a673-27a6a69ef348.png)
3. 安装完成后,`cd` 进入赛题文件夹(根据实际目录调整):
```bash
cd xpuoj_problem/
```
4. 在该目录下直接输入 `opencode` 并回车,即可进入 OpenCode  Agent 界面:
```bash
opencode
```
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/Lk3lbmbEQX2A6Om9/img/60afdbc5-c334-4c78-8df9-80433091ddaf.png)
参考prompt
```python
本题是 FlashAttention paged KV cache decode 的 CUDA C++ 前向算子。
输入输出规范如下:
- q: [batch, 1, num_heads, head_dim] (float)
- k_cache_paged: [num_blocks, num_heads, page_size, head_dim] (float)
- v_cache_paged: [num_blocks, num_heads, page_size, head_dim] (float)
- cache_seqlen: [batch] (int)
- block_table: [batch, max_num_blocks] (int)
- output: [batch, 1, num_heads, head_dim]
需调用 run_kernel(q, k_cache_paged, v_cache_paged, cache_seqlen, block_table, output)
你需要根据 cache_seqlen 和 block_table 从 paged cache 中取出对应的 K, V计算 attention 结果,存入 output。
参考out = flash_attn_with_kvcache(q, k_cache_paged, v_cache_paged, cache_seqlen=..., block_table=...)
要求你的实现与这个 API 的计算结果一致误差允许1e-5
请生成一个 run_kernel 函数,优先保证正确性。
```
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/Lk3lbmbEQX2A6Om9/img/2efeb3c0-77b0-407b-9bff-815670bcbf00.png)
输出:
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/Lk3lbmbEQX2A6Om9/img/095af1c8-aa7e-44e3-ad0e-620206aaba1e.png)
为方便参赛者先跑通完整流程,这里直接提供一份完整的冒烟代码,可直接复制粘贴到右侧编辑器,用于验证提交链路是否正常:
```python
#include <stdint.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#define PAGE_SIZE 16
#define HEAD_DIM 128
__global__ void paged_attention_kernel(
const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k_cache_paged,
const __nv_bfloat16* __restrict__ v_cache_paged,
__nv_bfloat16* __restrict__ output,
const int32_t* __restrict__ cache_seqlens,
const int32_t* __restrict__ block_table,
int64_t batch_size,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t blocks_per_batch)
{
int batch_idx = blockIdx.x / num_heads;
int head_idx = blockIdx.x % num_heads;
if (batch_idx >= batch_size || head_idx >= num_heads) return;
int seqlen = cache_seqlens[batch_idx];
if (seqlen <= 0) {
// 无有效 KV输出 0
int64_t out_base = ((batch_idx * seqlen_q + 0) * num_heads + head_idx) * headdim;
for (int i = threadIdx.x; i < headdim; i += blockDim.x)
output[out_base + i] = __float2bfloat16(0.0f);
return;
}
int tid = threadIdx.x;
// 加载对应 head 的 query 元素(每个线程负责一个维度)
int64_t q_offset = ((batch_idx * seqlen_q + 0) * num_heads + head_idx) * headdim;
float q_val = __bfloat162float(q[q_offset + tid]);
// 全局 softmax 状态(每个线程维护自己维度的累加器)
float max_val = -1e38f;
float sum_exp = 0.0f;
float out_acc = 0.0f;
float scale = rsqrtf(static_cast<float>(headdim));
// 共享内存布局:
// K_tile[PAGE_SIZE][HEAD_DIM] (bf16)
// V_tile[PAGE_SIZE][HEAD_DIM] (bf16)
// partial_scores[PAGE_SIZE][HEAD_DIM] (float, 用于归约点积)
__shared__ __nv_bfloat16 K_tile[PAGE_SIZE][HEAD_DIM];
__shared__ __nv_bfloat16 V_tile[PAGE_SIZE][HEAD_DIM];
__shared__ float partial_scores[PAGE_SIZE][HEAD_DIM];
int total_pages = (seqlen + PAGE_SIZE - 1) / PAGE_SIZE;
for (int page = 0; page < total_pages; ++page) {
int physical_block = block_table[batch_idx * blocks_per_batch + page];
int tokens_this_page = min(seqlen - page * PAGE_SIZE, PAGE_SIZE);
// 1. 将当前 page 的 K 和 V 从全局显存加载到共享内存
// 每个线程负责加载所有 token 的同一个 head 维度
const int64_t kv_stride = num_heads_k * headdim; // 每个 (block, offset) 的 stride
for (int j = 0; j < tokens_this_page; ++j) {
int64_t offset = (physical_block * PAGE_SIZE + j) * kv_stride + head_idx * headdim + tid;
K_tile[j][tid] = k_cache_paged[offset];
V_tile[j][tid] = v_cache_paged[offset];
}
__syncthreads();
// 2. 计算该 page 内每个 token 与 Q 的部分点积,存入 partial_scores
for (int j = 0; j < tokens_this_page; ++j) {
float k_val = __bfloat162float(K_tile[j][tid]);
partial_scores[j][tid] = q_val * k_val;
}
__syncthreads();
// 3. 对 partial_scores 做 tree reduction得到每个 token 的完整点积
#pragma unroll
for (int stride = HEAD_DIM / 2; stride > 0; stride >>= 1) {
if (tid < stride) {
#pragma unroll
for (int j = 0; j < tokens_this_page; ++j) {
partial_scores[j][tid] += partial_scores[j][tid + stride];
}
}
__syncthreads();
}
// 4. 在线 safe softmax 更新 + V 累加
// 4.1 找出本 page 内点积的最大值,结合全局 max 得到 new_max
float local_max = -1e38f;
#pragma unroll
for (int j = 0; j < tokens_this_page; ++j) {
local_max = fmaxf(local_max, partial_scores[j][0]);
}
float new_max = fmaxf(max_val, local_max);
// 4.2 用旧的 max 对全局状态进行重缩放
float rescale = expf(max_val - new_max);
sum_exp *= rescale;
out_acc *= rescale;
max_val = new_max;
// 4.3 累加 V并更新 sum_exp
#pragma unroll
for (int j = 0; j < tokens_this_page; ++j) {
float score = partial_scores[j][0] * scale;
float weight = expf(score - new_max);
sum_exp += weight;
float v_val = __bfloat162float(V_tile[j][tid]);
out_acc += weight * v_val;
}
__syncthreads(); // 准备下一个 page 的共享内存加载
}
// 5. 最终归一化并写回
if (sum_exp > 0.0f) {
out_acc /= sum_exp;
} else {
out_acc = 0.0f;
}
int64_t out_offset = ((batch_idx * seqlen_q + 0) * num_heads + head_idx) * headdim + tid;
output[out_offset] = __float2bfloat16(out_acc);
}
extern "C" void run_kernel(
const __nv_bfloat16* q,
const __nv_bfloat16* k_cache_paged,
const __nv_bfloat16* v_cache_paged,
__nv_bfloat16* output,
const int32_t* cache_seqlens,
const int32_t* block_table,
int64_t batch_size,
int64_t seqlen_k,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t num_blocks,
int64_t causal)
{
int64_t blocks_per_batch = num_blocks / batch_size;
dim3 grid(batch_size * num_heads);
dim3 block(HEAD_DIM);
paged_attention_kernel<<<grid, block>>>(
q, k_cache_paged, v_cache_paged, output,
cache_seqlens, block_table,
batch_size, seqlen_q, num_heads, num_heads_k, headdim,
page_block_size, blocks_per_batch
);
}
```
以上代码仅用于说明接口结构,不代表最优实现,也不作为评分参考
5. 点击提交,等待评测结果返回;
评测时间与题目测试点数量、队列状态和平台负载有关,通常需要等待数十秒到数分钟。以平台实际返回为准。
![0e9d0dccd68a1ddf1419973bc0c4e4bb.png](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/Yvenve5yZMWEwloy/img/95a36602-dcd7-48ea-9010-0dc4e7bb45d6.png)
OJ 对每次提交大致会走这个流程
```plaintext
XPU-OJ 的一次评测大致流程如下:
1. 选手提交代码;
2. 平台按所选语言编译或加载提交代码;
3. 评测程序构造测试输入;
4. 调用选手代码中的 `run_kernel(...)`
5. 调用 `testcase_config.py` 中的 `baseline()` / 参考实现生成 `output_ref`
6. 将 `run_kernel(...)` 的输出与 `output_ref` 做正确性校验;
7. 正确性通过后,统计运行耗时或性能指标;
8. 根据题目评分规则换算该题得分;
9. 更新该题历史最好成绩;
10. 汇总各题最好成绩,得到排行榜总分。
```
6. 查看结果
提交详情会显示状态、得分、时间、内存、编译信息以及各测试点结果。
![cd119ff85c024846b91d69b9ca9a1e00.png](https://origin.picgo.net/2026/06/17/cd119ff85c024846b91d69b9ca9a1e006456033e499779c6.png)
## 从 Baseline 到参赛作品的路径回顾
1. 跑通 baseline benchmark记录原始性能
2. 提交冒烟代码确认 OJ 链路正常
3. 使用 Agent 阅读题包理解输入输出、数据范围和精度要求
4. 在 `run_kernel(...)` 中实现最小正确版算子;
5. 提交 OJ先通过正确性
6. 正确性通过后再让 Agent 辅助分析性能瓶颈
7. 围绕访存、softmax、线程划分、head\_dim 特化、K/V 复用等方向迭代优化
8. 保存每轮 Agent Prompt、代码改动、OJ 结果和性能变化形成可复现的 Agent/Skill 优化流程。
## 七、常见问题
### Q1: 运行时提示 `ModuleNotFoundError: No module named &#39;flash_attn &#39;`
@ -601,9 +1043,23 @@ ms = run_with_profiler(run_fn, warmup=warmup, reps=repeat, print_result=True, ta
**原因:** 可能是 warmup 不足或 GPU 未达到稳态。 **解决:** 增加 `warmup` 次数,如 `warmup = 20`
---
### Q5: 评测状态显示 `**Compile Error**`或提示 `**Undefined reference to run_kernel**`
## 九、下一步学习建议
**原因:**C++函数名被修饰Name Mangling或参数类型/顺序与接口约定不符。 **解决:**在 run\_kernel 前添加 extern "C",并严格逐字核对参数的类型和修饰符。
#### Q6: 评测状态显示 `**Wrong Answer**`提示 torch.allclose校验失败
**原因:**bfloat16精度截断溢出、线程同步缺失或无效TokenPadding区域处理错误。 **解决:**累加和 Softmax 强制转为 float32 计算检查 \_\_syncthreads() 逻辑增加 seqlen 的边界判空。
#### Q7: 评测状态显示 `**Runtime Error**`  (非法内存访问/段错误)
**原因:**Paged KV 地址映射索引错误、尾部 Page 越界读取或线程块维度超限。 **解决:**仔细核对物理块寻址公式增加当前 Page 有效 Token 数量的越界判断检查单 Block 线程数配置。
#### Q8: 评测状态显示`**Time Limit Exceeded**` (评测超时)
**原因:**发散分支内的 \_\_syncthreads() 导致内核死锁、误加主机端同步指令或并行度划分错误导致串行。 **解决:**确保同步指令在所有线程必经路径上移除主机端多余的 cudaDeviceSynchronize();优化 <<<grid, block>>> 参数以提升并行度。
## 八、下一步学习建议
完成本模块后,建议继续学习以下内容:
@ -615,5 +1071,6 @@ ms = run_with_profiler(run_fn, warmup=warmup, reps=repeat, print_result=True, ta
4. **性能对比分析**   baseline 结果与优化后结果进行对比
---
5. **OJ 题包深度解析** — 学习阅读题包中的 `testcase_config.py`掌握本地构造边界用例与独立 Debug 的能力
6. **评测打榜与极限优化**  在通过 OJ 正确性校验的基础上挑战排行榜Leaderboard不断逼近硬件理论带宽极限

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,160,0.0580,45.25
2,512,8,160,0.0611,85.84
4,512,8,160,0.0656,159.96
8,512,8,160,0.0699,300.40
16,512,8,160,0.1321,317.92
32,512,8,160,0.2002,419.52
64,512,8,160,0.3383,496.43
128,512,8,160,0.6669,503.61
1,1024,8,160,0.1129,46.45
2,1024,8,160,0.1190,88.18
4,1024,8,160,0.1224,171.49
8,1024,8,160,0.1287,326.07
16,1024,8,160,0.2479,338.49
32,1024,8,160,0.3767,445.54
64,1024,8,160,0.6419,523.01
128,1024,8,160,1.2804,524.37
1,2048,8,160,0.2270,46.20
2,2048,8,160,0.2299,91.22
4,2048,8,160,0.2349,178.63
8,2048,8,160,0.2447,342.96
16,2048,8,160,0.4773,351.60
32,2048,8,160,0.7279,461.07
64,2048,8,160,1.2559,534.49
128,2048,8,160,2.5613,524.15
1,4096,8,160,0.4460,47.02
2,4096,8,160,0.4513,92.94
4,4096,8,160,0.4593,182.64
8,4096,8,160,0.4813,348.64
16,4096,8,160,0.9363,358.43
32,4096,8,160,1.4552,461.21
64,4096,8,160,2.5615,524.05
128,4096,8,160,5.1420,522.11
1,8192,8,160,0.8847,47.41
2,8192,8,160,0.8944,93.80
4,8192,8,160,0.9094,184.51
8,8192,8,160,0.9625,348.64
16,8192,8,160,1.8550,361.80
32,8192,8,160,2.9567,453.97
64,8192,8,160,5.1398,522.30
128,8192,8,160,10.2972,521.41
1,16384,8,160,1.7608,47.64
2,16384,8,160,1.7786,94.33
4,16384,8,160,1.8143,184.95
8,16384,8,160,1.9317,347.42
16,16384,8,160,3.7301,359.83
32,16384,8,160,5.9216,453.33
64,16384,8,160,10.2668,522.94
128,16384,8,160,20.6062,521.09
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 160 0.0580 45.25
3 2 512 8 160 0.0611 85.84
4 4 512 8 160 0.0656 159.96
5 8 512 8 160 0.0699 300.40
6 16 512 8 160 0.1321 317.92
7 32 512 8 160 0.2002 419.52
8 64 512 8 160 0.3383 496.43
9 128 512 8 160 0.6669 503.61
10 1 1024 8 160 0.1129 46.45
11 2 1024 8 160 0.1190 88.18
12 4 1024 8 160 0.1224 171.49
13 8 1024 8 160 0.1287 326.07
14 16 1024 8 160 0.2479 338.49
15 32 1024 8 160 0.3767 445.54
16 64 1024 8 160 0.6419 523.01
17 128 1024 8 160 1.2804 524.37
18 1 2048 8 160 0.2270 46.20
19 2 2048 8 160 0.2299 91.22
20 4 2048 8 160 0.2349 178.63
21 8 2048 8 160 0.2447 342.96
22 16 2048 8 160 0.4773 351.60
23 32 2048 8 160 0.7279 461.07
24 64 2048 8 160 1.2559 534.49
25 128 2048 8 160 2.5613 524.15
26 1 4096 8 160 0.4460 47.02
27 2 4096 8 160 0.4513 92.94
28 4 4096 8 160 0.4593 182.64
29 8 4096 8 160 0.4813 348.64
30 16 4096 8 160 0.9363 358.43
31 32 4096 8 160 1.4552 461.21
32 64 4096 8 160 2.5615 524.05
33 128 4096 8 160 5.1420 522.11
34 1 8192 8 160 0.8847 47.41
35 2 8192 8 160 0.8944 93.80
36 4 8192 8 160 0.9094 184.51
37 8 8192 8 160 0.9625 348.64
38 16 8192 8 160 1.8550 361.80
39 32 8192 8 160 2.9567 453.97
40 64 8192 8 160 5.1398 522.30
41 128 8192 8 160 10.2972 521.41
42 1 16384 8 160 1.7608 47.64
43 2 16384 8 160 1.7786 94.33
44 4 16384 8 160 1.8143 184.95
45 8 16384 8 160 1.9317 347.42
46 16 16384 8 160 3.7301 359.83
47 32 16384 8 160 5.9216 453.33
48 64 16384 8 160 10.2668 522.94
49 128 16384 8 160 20.6062 521.09

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,192,0.0458,68.82
2,512,8,192,0.0515,122.32
4,512,8,192,0.0574,219.28
8,512,8,192,0.0607,414.80
16,512,8,192,0.1147,439.27
32,512,8,192,0.1763,571.40
64,512,8,192,0.2978,676.79
128,512,8,192,0.5874,686.11
1,1024,8,192,0.0946,66.55
2,1024,8,192,0.1033,121.85
4,1024,8,192,0.1073,234.66
8,1024,8,192,0.1131,445.23
16,1024,8,192,0.2165,465.24
32,1024,8,192,0.3347,601.80
64,1024,8,192,0.5701,706.63
128,1024,8,192,1.1302,712.88
1,2048,8,192,0.1943,64.79
2,2048,8,192,0.1992,126.38
4,2048,8,192,0.2059,244.52
8,2048,8,192,0.2174,463.13
16,2048,8,192,0.4202,479.24
32,2048,8,192,0.6503,619.36
64,2048,8,192,1.1158,721.93
128,2048,8,192,2.2250,724.05
1,4096,8,192,0.3834,65.65
2,4096,8,192,0.3904,128.95
4,4096,8,192,0.4043,249.04
8,4096,8,192,0.4267,471.92
16,4096,8,192,0.8271,486.90
32,4096,8,192,1.2840,627.28
64,4096,8,192,2.2148,727.29
128,4096,8,192,4.3819,735.21
1,8192,8,192,0.7566,66.52
2,8192,8,192,0.7712,130.54
4,8192,8,192,0.7974,252.49
8,8192,8,192,0.8433,477.47
16,8192,8,192,1.6433,490.09
32,8192,8,192,2.5573,629.84
64,8192,8,192,4.3785,735.73
128,8192,8,192,8.7303,737.99
1,16384,8,192,1.5068,66.81
2,16384,8,192,1.5350,131.16
4,16384,8,192,1.5868,253.76
8,16384,8,192,1.6778,479.99
16,16384,8,192,3.2750,491.81
32,16384,8,192,5.0659,635.88
64,16384,8,192,8.7435,736.85
128,16384,8,192,17.5040,736.13
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 192 0.0458 68.82
3 2 512 8 192 0.0515 122.32
4 4 512 8 192 0.0574 219.28
5 8 512 8 192 0.0607 414.80
6 16 512 8 192 0.1147 439.27
7 32 512 8 192 0.1763 571.40
8 64 512 8 192 0.2978 676.79
9 128 512 8 192 0.5874 686.11
10 1 1024 8 192 0.0946 66.55
11 2 1024 8 192 0.1033 121.85
12 4 1024 8 192 0.1073 234.66
13 8 1024 8 192 0.1131 445.23
14 16 1024 8 192 0.2165 465.24
15 32 1024 8 192 0.3347 601.80
16 64 1024 8 192 0.5701 706.63
17 128 1024 8 192 1.1302 712.88
18 1 2048 8 192 0.1943 64.79
19 2 2048 8 192 0.1992 126.38
20 4 2048 8 192 0.2059 244.52
21 8 2048 8 192 0.2174 463.13
22 16 2048 8 192 0.4202 479.24
23 32 2048 8 192 0.6503 619.36
24 64 2048 8 192 1.1158 721.93
25 128 2048 8 192 2.2250 724.05
26 1 4096 8 192 0.3834 65.65
27 2 4096 8 192 0.3904 128.95
28 4 4096 8 192 0.4043 249.04
29 8 4096 8 192 0.4267 471.92
30 16 4096 8 192 0.8271 486.90
31 32 4096 8 192 1.2840 627.28
32 64 4096 8 192 2.2148 727.29
33 128 4096 8 192 4.3819 735.21
34 1 8192 8 192 0.7566 66.52
35 2 8192 8 192 0.7712 130.54
36 4 8192 8 192 0.7974 252.49
37 8 8192 8 192 0.8433 477.47
38 16 8192 8 192 1.6433 490.09
39 32 8192 8 192 2.5573 629.84
40 64 8192 8 192 4.3785 735.73
41 128 8192 8 192 8.7303 737.99
42 1 16384 8 192 1.5068 66.81
43 2 16384 8 192 1.5350 131.16
44 4 16384 8 192 1.5868 253.76
45 8 16384 8 192 1.6778 479.99
46 16 16384 8 192 3.2750 491.81
47 32 16384 8 192 5.0659 635.88
48 64 16384 8 192 8.7435 736.85
49 128 16384 8 192 17.5040 736.13

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,224,0.1254,29.29
2,512,8,224,0.1412,52.05
4,512,8,224,0.1497,98.13
8,512,8,224,0.1533,191.70
16,512,8,224,0.1913,307.17
32,512,8,224,0.3292,357.08
64,512,8,224,0.5187,453.28
128,512,8,224,0.9522,493.84
1,1024,8,224,0.2727,26.93
2,1024,8,224,0.2836,51.78
4,1024,8,224,0.2890,101.63
8,1024,8,224,0.2959,198.55
16,1024,8,224,0.3696,317.93
32,1024,8,224,0.6408,366.75
64,1024,8,224,1.0081,466.21
128,1024,8,224,1.8548,506.78
1,2048,8,224,0.5515,26.63
2,2048,8,224,0.5575,52.67
4,2048,8,224,0.5666,103.65
8,2048,8,224,0.5803,202.42
16,2048,8,224,0.7250,324.05
32,2048,8,224,1.2593,373.14
64,2048,8,224,1.9890,472.48
128,2048,8,224,3.6905,509.28
1,4096,8,224,1.0939,26.84
2,4096,8,224,1.1044,53.18
4,4096,8,224,1.1219,104.69
8,4096,8,224,1.1500,204.26
16,4096,8,224,1.4390,326.48
32,4096,8,224,2.4992,375.97
64,4096,8,224,4.0082,468.86
128,4096,8,224,7.3372,512.26
1,8192,8,224,2.1775,26.97
2,8192,8,224,2.1989,53.41
4,8192,8,224,2.2338,105.15
8,8192,8,224,2.3268,201.90
16,8192,8,224,2.8806,326.18
32,8192,8,224,5.0187,374.43
64,8192,8,224,8.0323,467.90
128,8192,8,224,14.6300,513.78
1,16384,8,224,4.3360,27.09
2,16384,8,224,4.3820,53.60
4,16384,8,224,4.5006,104.38
8,16384,8,224,4.6987,199.96
16,16384,8,224,5.7361,327.59
32,16384,8,224,10.1291,371.03
64,16384,8,224,16.0745,467.60
128,16384,8,224,OOM,OOM
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 224 0.1254 29.29
3 2 512 8 224 0.1412 52.05
4 4 512 8 224 0.1497 98.13
5 8 512 8 224 0.1533 191.70
6 16 512 8 224 0.1913 307.17
7 32 512 8 224 0.3292 357.08
8 64 512 8 224 0.5187 453.28
9 128 512 8 224 0.9522 493.84
10 1 1024 8 224 0.2727 26.93
11 2 1024 8 224 0.2836 51.78
12 4 1024 8 224 0.2890 101.63
13 8 1024 8 224 0.2959 198.55
14 16 1024 8 224 0.3696 317.93
15 32 1024 8 224 0.6408 366.75
16 64 1024 8 224 1.0081 466.21
17 128 1024 8 224 1.8548 506.78
18 1 2048 8 224 0.5515 26.63
19 2 2048 8 224 0.5575 52.67
20 4 2048 8 224 0.5666 103.65
21 8 2048 8 224 0.5803 202.42
22 16 2048 8 224 0.7250 324.05
23 32 2048 8 224 1.2593 373.14
24 64 2048 8 224 1.9890 472.48
25 128 2048 8 224 3.6905 509.28
26 1 4096 8 224 1.0939 26.84
27 2 4096 8 224 1.1044 53.18
28 4 4096 8 224 1.1219 104.69
29 8 4096 8 224 1.1500 204.26
30 16 4096 8 224 1.4390 326.48
31 32 4096 8 224 2.4992 375.97
32 64 4096 8 224 4.0082 468.86
33 128 4096 8 224 7.3372 512.26
34 1 8192 8 224 2.1775 26.97
35 2 8192 8 224 2.1989 53.41
36 4 8192 8 224 2.2338 105.15
37 8 8192 8 224 2.3268 201.90
38 16 8192 8 224 2.8806 326.18
39 32 8192 8 224 5.0187 374.43
40 64 8192 8 224 8.0323 467.90
41 128 8192 8 224 14.6300 513.78
42 1 16384 8 224 4.3360 27.09
43 2 16384 8 224 4.3820 53.60
44 4 16384 8 224 4.5006 104.38
45 8 16384 8 224 4.6987 199.96
46 16 16384 8 224 5.7361 327.59
47 32 16384 8 224 10.1291 371.03
48 64 16384 8 224 16.0745 467.60
49 128 16384 8 224 OOM OOM

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,256,0.0877,47.89
2,512,8,256,0.0921,91.17
4,512,8,256,0.0940,178.74
8,512,8,256,0.0964,348.52
16,512,8,256,0.1450,463.27
32,512,8,256,0.2250,597.21
64,512,8,256,0.3609,744.43
128,512,8,256,0.6932,775.25
1,1024,8,256,0.1747,48.04
2,1024,8,256,0.1762,95.27
4,1024,8,256,0.1784,188.22
8,1024,8,256,0.1817,369.53
16,1024,8,256,0.2796,480.25
32,1024,8,256,0.4339,619.00
64,1024,8,256,0.6960,771.73
128,1024,8,256,1.3439,799.36
1,2048,8,256,0.3410,49.21
2,2048,8,256,0.3439,97.60
4,2048,8,256,0.3469,193.52
8,2048,8,256,0.3533,379.94
16,2048,8,256,0.5461,491.67
32,2048,8,256,0.8493,632.28
64,2048,8,256,1.3667,785.82
128,2048,8,256,2.6465,811.64
1,4096,8,256,0.6742,49.77
2,4096,8,256,0.6777,99.03
4,4096,8,256,0.6836,196.36
8,4096,8,256,0.6950,386.31
16,4096,8,256,1.0803,497.02
32,4096,8,256,1.6794,639.44
64,4096,8,256,2.7101,792.50
128,4096,8,256,5.2543,817.52
1,8192,8,256,1.3375,50.18
2,8192,8,256,1.3448,99.81
4,8192,8,256,1.3564,197.91
8,8192,8,256,1.3799,389.08
16,8192,8,256,2.1465,500.25
32,8192,8,256,3.3342,644.12
64,8192,8,256,5.3983,795.67
128,8192,8,256,10.4691,820.55
1,16384,8,256,2.6697,50.28
2,16384,8,256,2.6817,100.10
4,16384,8,256,2.7049,198.49
8,16384,8,256,2.7533,390.00
16,16384,8,256,4.2789,501.89
32,16384,8,256,6.6476,646.11
64,16384,8,256,10.7723,797.43
128,16384,8,256,OOM,OOM
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 256 0.0877 47.89
3 2 512 8 256 0.0921 91.17
4 4 512 8 256 0.0940 178.74
5 8 512 8 256 0.0964 348.52
6 16 512 8 256 0.1450 463.27
7 32 512 8 256 0.2250 597.21
8 64 512 8 256 0.3609 744.43
9 128 512 8 256 0.6932 775.25
10 1 1024 8 256 0.1747 48.04
11 2 1024 8 256 0.1762 95.27
12 4 1024 8 256 0.1784 188.22
13 8 1024 8 256 0.1817 369.53
14 16 1024 8 256 0.2796 480.25
15 32 1024 8 256 0.4339 619.00
16 64 1024 8 256 0.6960 771.73
17 128 1024 8 256 1.3439 799.36
18 1 2048 8 256 0.3410 49.21
19 2 2048 8 256 0.3439 97.60
20 4 2048 8 256 0.3469 193.52
21 8 2048 8 256 0.3533 379.94
22 16 2048 8 256 0.5461 491.67
23 32 2048 8 256 0.8493 632.28
24 64 2048 8 256 1.3667 785.82
25 128 2048 8 256 2.6465 811.64
26 1 4096 8 256 0.6742 49.77
27 2 4096 8 256 0.6777 99.03
28 4 4096 8 256 0.6836 196.36
29 8 4096 8 256 0.6950 386.31
30 16 4096 8 256 1.0803 497.02
31 32 4096 8 256 1.6794 639.44
32 64 4096 8 256 2.7101 792.50
33 128 4096 8 256 5.2543 817.52
34 1 8192 8 256 1.3375 50.18
35 2 8192 8 256 1.3448 99.81
36 4 8192 8 256 1.3564 197.91
37 8 8192 8 256 1.3799 389.08
38 16 8192 8 256 2.1465 500.25
39 32 8192 8 256 3.3342 644.12
40 64 8192 8 256 5.3983 795.67
41 128 8192 8 256 10.4691 820.55
42 1 16384 8 256 2.6697 50.28
43 2 16384 8 256 2.6817 100.10
44 4 16384 8 256 2.7049 198.49
45 8 16384 8 256 2.7533 390.00
46 16 16384 8 256 4.2789 501.89
47 32 16384 8 256 6.6476 646.11
48 64 16384 8 256 10.7723 797.43
49 128 16384 8 256 OOM OOM

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,32,0.0257,20.45
2,512,8,32,0.0256,41.02
4,512,8,32,0.0258,81.28
8,512,8,32,0.0265,158.45
16,512,8,32,0.0396,212.30
32,512,8,32,0.0516,325.43
64,512,8,32,0.0721,465.83
128,512,8,32,0.1270,529.03
1,1024,8,32,0.0461,22.75
2,1024,8,32,0.0465,45.15
4,1024,8,32,0.0477,88.04
8,1024,8,32,0.0548,153.23
16,1024,8,32,0.0734,228.71
32,1024,8,32,0.0958,350.42
64,1024,8,32,0.1334,503.15
128,1024,8,32,0.2381,564.04
1,2048,8,32,0.0872,24.06
2,2048,8,32,0.0904,46.42
4,2048,8,32,0.1028,81.59
8,2048,8,32,0.1067,157.25
16,2048,8,32,0.1428,235.10
32,2048,8,32,0.1818,369.13
64,2048,8,32,0.2554,525.57
128,2048,8,32,0.4622,580.86
1,4096,8,32,0.1730,24.25
2,4096,8,32,0.1955,42.91
4,4096,8,32,0.2020,83.05
8,4096,8,32,0.2140,156.83
16,4096,8,32,0.2777,241.65
32,4096,8,32,0.3542,378.99
64,4096,8,32,0.4990,538.05
128,4096,8,32,0.9099,590.13
1,8192,8,32,0.3820,21.96
2,8192,8,32,0.3913,42.88
4,8192,8,32,0.4127,81.31
8,8192,8,32,0.4224,158.88
16,8192,8,32,0.5490,244.51
32,8192,8,32,0.6960,385.70
64,8192,8,32,0.9870,543.98
128,8192,8,32,1.8100,593.25
1,16384,8,32,0.7655,21.92
2,16384,8,32,0.8067,41.59
4,16384,8,32,0.8228,81.56
8,16384,8,32,0.8397,159.85
16,16384,8,32,1.0910,246.04
32,16384,8,32,1.3824,388.37
64,16384,8,32,1.9663,546.08
128,16384,8,32,3.6107,594.78
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 32 0.0257 20.45
3 2 512 8 32 0.0256 41.02
4 4 512 8 32 0.0258 81.28
5 8 512 8 32 0.0265 158.45
6 16 512 8 32 0.0396 212.30
7 32 512 8 32 0.0516 325.43
8 64 512 8 32 0.0721 465.83
9 128 512 8 32 0.1270 529.03
10 1 1024 8 32 0.0461 22.75
11 2 1024 8 32 0.0465 45.15
12 4 1024 8 32 0.0477 88.04
13 8 1024 8 32 0.0548 153.23
14 16 1024 8 32 0.0734 228.71
15 32 1024 8 32 0.0958 350.42
16 64 1024 8 32 0.1334 503.15
17 128 1024 8 32 0.2381 564.04
18 1 2048 8 32 0.0872 24.06
19 2 2048 8 32 0.0904 46.42
20 4 2048 8 32 0.1028 81.59
21 8 2048 8 32 0.1067 157.25
22 16 2048 8 32 0.1428 235.10
23 32 2048 8 32 0.1818 369.13
24 64 2048 8 32 0.2554 525.57
25 128 2048 8 32 0.4622 580.86
26 1 4096 8 32 0.1730 24.25
27 2 4096 8 32 0.1955 42.91
28 4 4096 8 32 0.2020 83.05
29 8 4096 8 32 0.2140 156.83
30 16 4096 8 32 0.2777 241.65
31 32 4096 8 32 0.3542 378.99
32 64 4096 8 32 0.4990 538.05
33 128 4096 8 32 0.9099 590.13
34 1 8192 8 32 0.3820 21.96
35 2 8192 8 32 0.3913 42.88
36 4 8192 8 32 0.4127 81.31
37 8 8192 8 32 0.4224 158.88
38 16 8192 8 32 0.5490 244.51
39 32 8192 8 32 0.6960 385.70
40 64 8192 8 32 0.9870 543.98
41 128 8192 8 32 1.8100 593.25
42 1 16384 8 32 0.7655 21.92
43 2 16384 8 32 0.8067 41.59
44 4 16384 8 32 0.8228 81.56
45 8 16384 8 32 0.8397 159.85
46 16 16384 8 32 1.0910 246.04
47 32 16384 8 32 1.3824 388.37
48 64 16384 8 32 1.9663 546.08
49 128 16384 8 32 3.6107 594.78

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,512,0.3588,23.40
2,512,8,512,0.3651,46.00
4,512,8,512,0.3736,89.89
8,512,8,512,0.3856,174.22
16,512,8,512,0.7472,179.80
32,512,8,512,1.1447,234.72
64,512,8,512,1.9549,274.89
128,512,8,512,3.8962,275.85
1,1024,8,512,0.7261,23.12
2,1024,8,512,0.7354,45.65
4,1024,8,512,0.7496,89.57
8,1024,8,512,0.7746,173.35
16,1024,8,512,1.5049,178.46
32,1024,8,512,2.3111,232.42
64,1024,8,512,3.9538,271.70
128,1024,8,512,7.8811,272.62
1,2048,8,512,1.4636,22.93
2,2048,8,512,1.4826,45.27
4,2048,8,512,1.5109,88.86
8,2048,8,512,1.5549,172.68
16,2048,8,512,3.0237,177.60
32,2048,8,512,4.6439,231.27
64,2048,8,512,7.9560,269.99
128,2048,8,512,15.8741,270.63
1,4096,8,512,2.9312,22.90
2,4096,8,512,2.9675,45.24
4,4096,8,512,3.0243,88.77
8,4096,8,512,3.1127,172.50
16,4096,8,512,6.0753,176.76
32,4096,8,512,9.3182,230.49
64,4096,8,512,15.9642,269.07
128,4096,8,512,31.8313,269.89
1,8192,8,512,5.8843,22.81
2,8192,8,512,5.9344,45.24
4,8192,8,512,6.0465,88.80
8,8192,8,512,6.2334,172.27
16,8192,8,512,12.1594,176.62
32,8192,8,512,18.6826,229.90
64,8192,8,512,32.0055,268.41
128,8192,8,512,OOM,OOM
1,16384,8,512,11.8153,22.72
2,16384,8,512,11.9237,45.03
4,16384,8,512,12.1671,88.25
8,16384,8,512,12.4948,171.88
16,16384,8,512,24.3414,176.45
32,16384,8,512,37.3907,229.74
64,16384,8,512,OOM,OOM
128,16384,8,512,OOM,OOM
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 512 0.3588 23.40
3 2 512 8 512 0.3651 46.00
4 4 512 8 512 0.3736 89.89
5 8 512 8 512 0.3856 174.22
6 16 512 8 512 0.7472 179.80
7 32 512 8 512 1.1447 234.72
8 64 512 8 512 1.9549 274.89
9 128 512 8 512 3.8962 275.85
10 1 1024 8 512 0.7261 23.12
11 2 1024 8 512 0.7354 45.65
12 4 1024 8 512 0.7496 89.57
13 8 1024 8 512 0.7746 173.35
14 16 1024 8 512 1.5049 178.46
15 32 1024 8 512 2.3111 232.42
16 64 1024 8 512 3.9538 271.70
17 128 1024 8 512 7.8811 272.62
18 1 2048 8 512 1.4636 22.93
19 2 2048 8 512 1.4826 45.27
20 4 2048 8 512 1.5109 88.86
21 8 2048 8 512 1.5549 172.68
22 16 2048 8 512 3.0237 177.60
23 32 2048 8 512 4.6439 231.27
24 64 2048 8 512 7.9560 269.99
25 128 2048 8 512 15.8741 270.63
26 1 4096 8 512 2.9312 22.90
27 2 4096 8 512 2.9675 45.24
28 4 4096 8 512 3.0243 88.77
29 8 4096 8 512 3.1127 172.50
30 16 4096 8 512 6.0753 176.76
31 32 4096 8 512 9.3182 230.49
32 64 4096 8 512 15.9642 269.07
33 128 4096 8 512 31.8313 269.89
34 1 8192 8 512 5.8843 22.81
35 2 8192 8 512 5.9344 45.24
36 4 8192 8 512 6.0465 88.80
37 8 8192 8 512 6.2334 172.27
38 16 8192 8 512 12.1594 176.62
39 32 8192 8 512 18.6826 229.90
40 64 8192 8 512 32.0055 268.41
41 128 8192 8 512 OOM OOM
42 1 16384 8 512 11.8153 22.72
43 2 16384 8 512 11.9237 45.03
44 4 16384 8 512 12.1671 88.25
45 8 16384 8 512 12.4948 171.88
46 16 16384 8 512 24.3414 176.45
47 32 16384 8 512 37.3907 229.74
48 64 16384 8 512 OOM OOM
49 128 16384 8 512 OOM OOM

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,64,0.0404,25.99
2,512,8,64,0.0399,52.60
4,512,8,64,0.0413,101.69
8,512,8,64,0.0482,174.25
16,512,8,64,0.0540,310.86
32,512,8,64,0.0629,533.75
64,512,8,64,0.0833,806.14
128,512,8,64,0.1104,1216.59
1,1024,8,64,0.0747,28.08
2,1024,8,64,0.0766,54.77
4,1024,8,64,0.0891,94.17
8,1024,8,64,0.0918,182.94
16,1024,8,64,0.1044,321.41
32,1024,8,64,0.1179,569.43
64,1024,8,64,0.1566,857.28
128,1024,8,64,0.2078,1292.17
1,2048,8,64,0.1455,28.84
2,2048,8,64,0.1684,49.82
4,2048,8,64,0.1730,97.01
8,2048,8,64,0.1850,181.39
16,2048,8,64,0.2009,334.18
32,2048,8,64,0.2268,592.01
64,2048,8,64,0.3002,894.44
128,2048,8,64,0.4027,1333.64
1,4096,8,64,0.3265,25.69
2,4096,8,64,0.3322,50.51
4,4096,8,64,0.3522,95.27
8,4096,8,64,0.3632,184.79
16,4096,8,64,0.3942,340.56
32,4096,8,64,0.4456,602.47
64,4096,8,64,0.5927,905.94
128,4096,8,64,0.7938,1352.87
1,8192,8,64,0.6508,25.78
2,8192,8,64,0.6879,48.78
4,8192,8,64,0.7008,95.77
8,8192,8,64,0.7199,186.44
16,8192,8,64,0.7786,344.79
32,8192,8,64,0.8798,610.25
64,8192,8,64,1.1745,914.30
128,8192,8,64,1.5728,1365.50
1,16384,8,64,1.3524,24.81
2,16384,8,64,1.3698,48.99
4,16384,8,64,1.3923,96.40
8,16384,8,64,1.4267,188.16
16,16384,8,64,1.5451,347.47
32,16384,8,64,1.7622,609.32
64,16384,8,64,2.3392,918.09
128,16384,8,64,3.1332,1370.84
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 64 0.0404 25.99
3 2 512 8 64 0.0399 52.60
4 4 512 8 64 0.0413 101.69
5 8 512 8 64 0.0482 174.25
6 16 512 8 64 0.0540 310.86
7 32 512 8 64 0.0629 533.75
8 64 512 8 64 0.0833 806.14
9 128 512 8 64 0.1104 1216.59
10 1 1024 8 64 0.0747 28.08
11 2 1024 8 64 0.0766 54.77
12 4 1024 8 64 0.0891 94.17
13 8 1024 8 64 0.0918 182.94
14 16 1024 8 64 0.1044 321.41
15 32 1024 8 64 0.1179 569.43
16 64 1024 8 64 0.1566 857.28
17 128 1024 8 64 0.2078 1292.17
18 1 2048 8 64 0.1455 28.84
19 2 2048 8 64 0.1684 49.82
20 4 2048 8 64 0.1730 97.01
21 8 2048 8 64 0.1850 181.39
22 16 2048 8 64 0.2009 334.18
23 32 2048 8 64 0.2268 592.01
24 64 2048 8 64 0.3002 894.44
25 128 2048 8 64 0.4027 1333.64
26 1 4096 8 64 0.3265 25.69
27 2 4096 8 64 0.3322 50.51
28 4 4096 8 64 0.3522 95.27
29 8 4096 8 64 0.3632 184.79
30 16 4096 8 64 0.3942 340.56
31 32 4096 8 64 0.4456 602.47
32 64 4096 8 64 0.5927 905.94
33 128 4096 8 64 0.7938 1352.87
34 1 8192 8 64 0.6508 25.78
35 2 8192 8 64 0.6879 48.78
36 4 8192 8 64 0.7008 95.77
37 8 8192 8 64 0.7199 186.44
38 16 8192 8 64 0.7786 344.79
39 32 8192 8 64 0.8798 610.25
40 64 8192 8 64 1.1745 914.30
41 128 8192 8 64 1.5728 1365.50
42 1 16384 8 64 1.3524 24.81
43 2 16384 8 64 1.3698 48.99
44 4 16384 8 64 1.3923 96.40
45 8 16384 8 64 1.4267 188.16
46 16 16384 8 64 1.5451 347.47
47 32 16384 8 64 1.7622 609.32
48 64 16384 8 64 2.3392 918.09
49 128 16384 8 64 3.1332 1370.84

View File

@ -0,0 +1,49 @@
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
1,512,8,96,0.0407,38.67
2,512,8,96,0.0398,79.02
4,512,8,96,0.0431,146.08
8,512,8,96,0.0495,254.61
16,512,8,96,0.0698,360.64
32,512,8,96,0.1117,450.87
64,512,8,96,0.1780,566.16
128,512,8,96,0.3329,605.28
1,1024,8,96,0.0732,43.01
2,1024,8,96,0.0794,79.29
4,1024,8,96,0.0871,144.54
8,1024,8,96,0.0934,269.52
16,1024,8,96,0.1297,388.14
32,1024,8,96,0.2114,476.36
64,1024,8,96,0.3379,596.08
128,1024,8,96,0.6327,636.68
1,2048,8,96,0.1505,41.80
2,2048,8,96,0.1619,77.76
4,2048,8,96,0.1713,146.94
8,2048,8,96,0.1780,282.84
16,2048,8,96,0.2492,404.09
32,2048,8,96,0.4088,492.55
64,2048,8,96,0.6575,612.55
128,2048,8,96,1.2457,646.61
1,4096,8,96,0.3099,40.61
2,4096,8,96,0.3259,77.23
4,4096,8,96,0.3346,150.42
8,4096,8,96,0.3467,290.41
16,4096,8,96,0.4888,411.94
32,4096,8,96,0.8055,499.94
64,4096,8,96,1.3209,609.72
128,4096,8,96,2.4810,649.25
1,8192,8,96,0.6343,39.68
2,8192,8,96,0.6437,78.20
4,8192,8,96,0.6601,152.50
8,8192,8,96,0.6826,294.97
16,8192,8,96,0.9688,415.64
32,8192,8,96,1.6057,501.55
64,8192,8,96,2.6527,607.19
128,8192,8,96,4.9464,651.27
1,16384,8,96,1.2581,40.01
2,16384,8,96,1.2812,78.57
4,16384,8,96,1.3112,153.55
8,16384,8,96,1.3653,294.92
16,16384,8,96,1.9351,416.16
32,16384,8,96,3.2277,499.01
64,16384,8,96,5.3192,605.60
128,16384,8,96,9.8747,652.44
1 batch_size seq_len_kv heads headdim time_ms bandwidth_GB_s
2 1 512 8 96 0.0407 38.67
3 2 512 8 96 0.0398 79.02
4 4 512 8 96 0.0431 146.08
5 8 512 8 96 0.0495 254.61
6 16 512 8 96 0.0698 360.64
7 32 512 8 96 0.1117 450.87
8 64 512 8 96 0.1780 566.16
9 128 512 8 96 0.3329 605.28
10 1 1024 8 96 0.0732 43.01
11 2 1024 8 96 0.0794 79.29
12 4 1024 8 96 0.0871 144.54
13 8 1024 8 96 0.0934 269.52
14 16 1024 8 96 0.1297 388.14
15 32 1024 8 96 0.2114 476.36
16 64 1024 8 96 0.3379 596.08
17 128 1024 8 96 0.6327 636.68
18 1 2048 8 96 0.1505 41.80
19 2 2048 8 96 0.1619 77.76
20 4 2048 8 96 0.1713 146.94
21 8 2048 8 96 0.1780 282.84
22 16 2048 8 96 0.2492 404.09
23 32 2048 8 96 0.4088 492.55
24 64 2048 8 96 0.6575 612.55
25 128 2048 8 96 1.2457 646.61
26 1 4096 8 96 0.3099 40.61
27 2 4096 8 96 0.3259 77.23
28 4 4096 8 96 0.3346 150.42
29 8 4096 8 96 0.3467 290.41
30 16 4096 8 96 0.4888 411.94
31 32 4096 8 96 0.8055 499.94
32 64 4096 8 96 1.3209 609.72
33 128 4096 8 96 2.4810 649.25
34 1 8192 8 96 0.6343 39.68
35 2 8192 8 96 0.6437 78.20
36 4 8192 8 96 0.6601 152.50
37 8 8192 8 96 0.6826 294.97
38 16 8192 8 96 0.9688 415.64
39 32 8192 8 96 1.6057 501.55
40 64 8192 8 96 2.6527 607.19
41 128 8192 8 96 4.9464 651.27
42 1 16384 8 96 1.2581 40.01
43 2 16384 8 96 1.2812 78.57
44 4 16384 8 96 1.3112 153.55
45 8 16384 8 96 1.3653 294.92
46 16 16384 8 96 1.9351 416.16
47 32 16384 8 96 3.2277 499.01
48 64 16384 8 96 5.3192 605.60
49 128 16384 8 96 9.8747 652.44

Binary file not shown.
1 ����Mac OS X ���� ���2���ª������Ü��������������������������������������ATTR�������Ü���”���H������������������”���H��com.apple.macl����Ѩ7ÉGÌ®Á*–Æïe ������������������������������������������������������

View File

@ -0,0 +1,330 @@
# guide
# FlashAttention KV Cache Decode - 参赛指南
## 一、登录 XPU-OJ 平台
打开浏览器,访问:\*\*https://xpuoj.com/\*\*
1. 等待组委会统一发放 XPU-OJ 账号;
2. 使用分配的用户名和初始密码登录平台
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/a3af510d-2b4a-4fa0-8e0f-5137ec2cc2fe.png)
点击 **"登录"** 进入平台。
## 二、进入比赛
### 2.1 点击导航栏"比赛"
登录成功后,会跳转到 XPUOJ 平台首页。
找到页面顶部的导航栏,点击 **"比赛"** 标签,进入比赛列表页。
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/d3c48202-7e3f-48a1-8a9e-e96192e51256.png)
### 2.2 找到目标比赛
在比赛列表中,根据比赛状态(全部 / 未开始 / 进行中 / 已结束)找到目标比赛。
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/bff83a88-21ed-4766-9450-6db58a83df43.png)
## 三、进入题目
### 3.1 在题目列表中找到目标题目
进入比赛后,页面会展示该比赛的**题目列表**。每道题都有编号(1 ~ N)和标题。
在列表中找到 **FlashAttention KV Cache Decode** 这道题(前缀为 `FlashAttention`),点击标题即可进入题目详情页。
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/cc304e39-753f-443f-9267-6b14c3ec99da.png)
### 3.2 查看题目要求
进入题目详情页后,可以看到以下几个区域:
* **左侧**:题目描述、接口约定、参数说明
* **右侧**:代码编辑器,用于编写并提交代码
请仔细阅读左侧的 **题目描述** 和 **接口约定**,重点关注:
* 入口函数名(本题为 `run_kernel`)
* 必传的参数列表及其类型、顺序
* 编译/运行环境(语言选择,目标硬件 C500)
### 3.3 编写并提交代码
在右侧的代码编辑器中,按照题目要求填入完整代码(可点击右上角"重置"恢复初始模板)。
#### 最小正确性代码(可先复制跑通)
为方便参赛者先跑通完整流程,这里提供一份 **最小正确性代码**,可直接复制粘贴到右侧编辑器,用于验证提交链路是否正常:
```cpp
#include <stdint.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#define HEAD_DIM 128
__global__ void paged_attention_kernel(
const __nv_bfloat16* q,
const __nv_bfloat16* k_cache_paged,
const __nv_bfloat16* v_cache_paged,
__nv_bfloat16* output,
const int32_t* cache_seqlens,
const int32_t* block_table,
int64_t batch_size,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t blocks_per_batch)
{
int batch_idx = blockIdx.x / num_heads;
int head_idx = blockIdx.x % num_heads;
if (batch_idx >= batch_size || head_idx >= num_heads) return;
int seqlen = cache_seqlens[batch_idx];
int tid = threadIdx.x;
// 加载对应 head 的 query 元素
int64_t q_offset = ((batch_idx * seqlen_q + 0) * num_heads + head_idx) * headdim;
float q_val = __bfloat162float(q[q_offset + tid]);
// Online safe softmax 状态
float max_val = -1e38f;
float sum_exp = 0.0f;
float out_acc = 0.0f;
float scale = 1.0f / sqrtf(static_cast<float>(headdim));
// 静态共享内存,避免动态分配可能带来的兼容性问题
__shared__ float s_score[HEAD_DIM];
for (int token = 0; token < seqlen; ++token) {
int page_idx = token / page_block_size;
int page_offset = token % page_block_size;
int physical_block = block_table[batch_idx * blocks_per_batch + page_idx];
// 读取 key 元素
const __nv_bfloat16* k_ptr = k_cache_paged
+ (physical_block * page_block_size + page_offset) * (num_heads_k * headdim)
+ head_idx * headdim;
float k_val = __bfloat162float(k_ptr[tid]);
// 点积 -> 共享内存归约
s_score[tid] = q_val * k_val;
__syncthreads();
for (int stride = HEAD_DIM >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
s_score[tid] += s_score[tid + stride];
}
__syncthreads();
}
float score = s_score[0] * scale;
// 更新 softmax 状态
float new_max = fmaxf(max_val, score);
float rescale = expf(max_val - new_max);
sum_exp = sum_exp * rescale + expf(score - new_max);
out_acc = out_acc * rescale;
max_val = new_max;
// 读取 value 元素,并累加(用最新 max 的权重)
const __nv_bfloat16* v_ptr = v_cache_paged
+ (physical_block * page_block_size + page_offset) * (num_heads_k * headdim)
+ head_idx * headdim;
float v_val = __bfloat162float(v_ptr[tid]);
out_acc += expf(score - max_val) * v_val;
__syncthreads(); // 确保下次迭代共享内存可安全复用
}
if (seqlen > 0) {
out_acc /= sum_exp;
} else {
out_acc = 0.0f;
}
int64_t out_offset = ((batch_idx * seqlen_q + 0) * num_heads + head_idx) * headdim + tid;
output[out_offset] = __float2bfloat16(out_acc);
}
extern "C" void run_kernel(
const __nv_bfloat16* q,
const __nv_bfloat16* k_cache_paged,
const __nv_bfloat16* v_cache_paged,
__nv_bfloat16* output,
const int32_t* cache_seqlens,
const int32_t* block_table,
int64_t batch_size,
int64_t seqlen_k,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t num_blocks,
int64_t causal)
{
int64_t blocks_per_batch = num_blocks / batch_size;
dim3 grid(batch_size * num_heads);
dim3 block(HEAD_DIM);
paged_attention_kernel<<<grid, block>>>(
q, k_cache_paged, v_cache_paged, output,
cache_seqlens, block_table,
batch_size, seqlen_q, num_heads, num_heads_k, headdim,
page_block_size, blocks_per_batch
);
}
```
完成后:
1. 在编辑器下方 **语言** 选项中,根据代码实际情况选择对应语言(可选 `Triton` / `CUDA Maca` / `MXMACA C++` / `TileLang` 等)
2. 选择 **目标硬件** 为 `C500`
3. 点击右上角 **"提交"** 按钮
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/d4dba895-a1fd-414c-832d-86111734378b.png)
## 四、查看提交结果
点击 **"提交"** 后,系统会自动跳转到提交记录页,展示本次提交的详细信息。
页面顶部会显示一行汇总信息,包括:
* **状态** —— 评测结果(如 `Accepted` 表示通过)
* **分数** —— 本次提交获得的分数
* **题目** —— 对应的题目名称
* **用时** —— 程序运行耗时
* **内存** —— 占用内存大小
* **答案** —— 提交所用的语言/硬件
* **提交时间** —— 提交的时刻
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/14eedaa5-8072-4247-a85b-5a12b284d23c.png)
页面下方会展开 **编译信息** 和 **各测试点**(样例、测试点 #1 ~ #N)的结果,逐个显示:
![image](https://alidocs.oss-cn-zhangjiakou.aliyuncs.com/res/4jKqm0bXGBkPvnw1/img/f1a3c988-c26a-4c30-b32b-f4e10d7270e2.png)

View File

@ -0,0 +1,16 @@
{
"id": 197,
"displayId": 20005,
"type": "Traditional",
"isPublic": false,
"locales": [
"zh_CN"
],
"samples": [
{
"inputData": "1\n",
"outputData": ""
}
],
"problemTagIds": []
}

View File

@ -0,0 +1,298 @@
from __future__ import annotations
HEAD_DIMS = [128]
BATCH_SIZES = [1, 4, 16]
SEQ_LENS_KV = [1024, 4096, 8192, 16384]
SEQ_LEN_Q = 1
NUM_HEADS = 8
NUM_HEADS_K = 8
PAGE_BLOCK_SIZE = 16
CAUSAL = 0
def _build_cases():
cases = []
for headdim in HEAD_DIMS:
for seqlen_k in SEQ_LENS_KV:
for batch_size in BATCH_SIZES:
cases.append(
(
batch_size,
seqlen_k,
SEQ_LEN_Q,
NUM_HEADS,
NUM_HEADS_K,
headdim,
PAGE_BLOCK_SIZE,
CAUSAL,
)
)
return cases
TESTCASES = _build_cases()
def getNumOfTestcases() -> int:
return len(TESTCASES)
try:
from pathlib import Path
from typing import List, Tuple, Union
import math
import sys
import torch
KernelArg = Union[torch.Tensor, int, float]
CURRENT_CASE = None
def _ensure_flashattn_importable():
try:
from flash_attn.flash_attn_interface import flash_attn_with_kvcache # noqa: F401
return
except ImportError:
pass
here = Path(__file__).resolve()
for parent in here.parents:
candidate = parent / "flashattn"
if (candidate / "flash_attn").is_dir():
sys.path.insert(0, str(candidate))
return
def _get_testcase_index() -> int:
try:
raw = input().strip()
except EOFError:
return 0
if raw == "":
return 0
try:
testcase_id = int(raw.split()[0])
except ValueError:
return 0
if 1 <= testcase_id <= len(TESTCASES):
return testcase_id - 1
if 0 <= testcase_id < len(TESTCASES):
return testcase_id
return 0
def _compute_reps(batch_size: int, seq_len: int, head_dim: int, base_reps: int = 100) -> int:
workload = batch_size * seq_len * head_dim
if workload < 1e5:
return base_reps
if workload < 1e6:
return base_reps // 2
if workload < 1e7:
return base_reps // 4
if workload < 1e8:
return base_reps // 8
if workload < 1e9:
return base_reps // 16
return base_reps // 32
def _get_num_blocks(batch_size: int, seqlen_k: int, page_block_size: int) -> int:
num_blocks = math.ceil(seqlen_k / page_block_size) * batch_size * 3
return max(1024, num_blocks)
def getTestCaseSize() -> Tuple[List[Tuple[int, ...]], Tuple[int, int]]:
testcase_id = _get_testcase_index()
global CURRENT_CASE
(
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
causal,
) = TESTCASES[testcase_id]
num_blocks = _get_num_blocks(batch_size, seqlen_k, page_block_size)
blocks_per_batch = num_blocks // batch_size
CURRENT_CASE = (
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
num_blocks,
causal,
20260720 + testcase_id,
)
warmup = 3
iters = max(1, _compute_reps(batch_size, seqlen_k, headdim))
return [
(batch_size, seqlen_q, num_heads, headdim),
(num_blocks, page_block_size, num_heads_k, headdim),
(num_blocks, page_block_size, num_heads_k, headdim),
(batch_size, seqlen_q, num_heads, headdim),
(batch_size,),
(batch_size, blocks_per_batch),
(), (), (), (), (), (), (), (), (),
], (warmup, iters)
def genTestCase(testcase_sizes, device: str = "cuda") -> List[KernelArg]:
del testcase_sizes
(
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
num_blocks,
causal,
seed,
) = CURRENT_CASE
gen = torch.Generator(device=device)
gen.manual_seed(seed)
dtype = torch.bfloat16
blocks_per_batch = num_blocks // batch_size
q = torch.randn(
batch_size,
seqlen_q,
num_heads,
headdim,
dtype=dtype,
device=device,
generator=gen,
).contiguous()
k_cache_paged = torch.randn(
num_blocks,
page_block_size,
num_heads_k,
headdim,
dtype=dtype,
device=device,
generator=gen,
).contiguous()
v_cache_paged = torch.randn(
num_blocks,
page_block_size,
num_heads_k,
headdim,
dtype=dtype,
device=device,
generator=gen,
).contiguous()
output = torch.empty(
batch_size,
seqlen_q,
num_heads,
headdim,
dtype=dtype,
device=device,
)
cache_seqlens = torch.full((batch_size,), seqlen_k, dtype=torch.int32, device=device)
block_table = torch.randperm(num_blocks, dtype=torch.int32, device=device, generator=gen).reshape(
batch_size,
blocks_per_batch,
)
return [
q,
k_cache_paged,
v_cache_paged,
output,
cache_seqlens,
block_table,
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
num_blocks,
causal,
]
def baseline(
q,
k_cache_paged,
v_cache_paged,
output,
cache_seqlens,
block_table,
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
num_blocks,
causal,
):
_ensure_flashattn_importable()
from flash_attn.flash_attn_interface import flash_attn_with_kvcache
out = flash_attn_with_kvcache(
q,
k_cache_paged,
v_cache_paged,
None,
None,
cache_seqlens=cache_seqlens,
cache_batch_idx=None,
block_table=block_table,
causal=bool(causal),
window_size=(-1, -1),
rotary_interleaved=False,
alibi_slopes=None,
num_splits=1,
)
output.copy_(out)
return [
q,
k_cache_paged,
v_cache_paged,
output,
cache_seqlens,
block_table,
batch_size,
seqlen_k,
seqlen_q,
num_heads,
num_heads_k,
headdim,
page_block_size,
num_blocks,
causal,
]
def check(
testcase_sizes,
original_input_tensors,
target_kernel_input_tensors,
baseline_input_tensors,
rtol=1e-2,
atol=1e-2,
) -> bool:
del testcase_sizes, original_input_tensors
output_t = target_kernel_input_tensors[3]
output_ref = baseline_input_tensors[3]
if output_t.shape != output_ref.shape:
print(f"[FAIL] shape mismatch: target {output_t.shape}, ref {output_ref.shape}", file=sys.stderr)
return False
if output_t.dtype != output_ref.dtype:
print(f"[FAIL] dtype mismatch: target {output_t.dtype}, ref {output_ref.dtype}", file=sys.stderr)
return False
if not torch.allclose(output_t.float(), output_ref.float(), rtol=rtol, atol=atol):
diff = (output_t.float() - output_ref.float()).abs()
print(
f"[FAIL] allclose failed: max_abs_diff={float(diff.max().item()):.6f}, "
f"mean_abs_diff={float(diff.mean().item()):.6f} (rtol={rtol}, atol={atol})",
file=sys.stderr,
)
return False
return True
except Exception:
pass

View File

@ -0,0 +1,28 @@
---
sectionTitle: "题目描述"
type: "Text"
---
你需要实现 FlashAttention paged KV cache decode 的 CUDA C++ 前向算子。
本题输入采用 `flash_attn_with_kvcache``flashattn/benchmarks/benchmark_kvcache.py` 中使用的 paged KV cache 配置。每个 batch 只有 1 个 query tokenKV cache 长度为 `seqlen_k`K/V cache 按 page 存储。
评测程序会调用你提交代码中的 `run_kernel` 函数。你需要根据 `cache_seqlens``block_table` 读取 paged KV cache并将结果写入 `output`
baseline 使用 benchmark 中的 FlashAttention Python API
```python
out = flash_attn_with_kvcache(
q, k_cache_paged, v_cache_paged, None, None,
cache_seqlens=cache_seqlens,
cache_batch_idx=None,
block_table=block_table,
causal=False,
window_size=(-1, -1),
rotary_interleaved=False,
alibi_slopes=None,
num_splits=1,
)
output.copy_(out)
```
如何提交代码详见[评测指南](/d/2)。

View File

@ -0,0 +1,43 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "cuda"
---
你必须在提交的 CUDA 源码中提供如下 **C 符号**,函数名、参数类型、顺序必须完全一致,并使用 `extern "C"` 防止 name mangling
```cpp
#include <stdint.h>
#include <cuda_bf16.h>
extern "C" void run_kernel(
const __nv_bfloat16* q,
const __nv_bfloat16* k_cache_paged,
const __nv_bfloat16* v_cache_paged,
__nv_bfloat16* output,
const int32_t* cache_seqlens,
const int32_t* block_table,
int64_t batch_size,
int64_t seqlen_k,
int64_t seqlen_q,
int64_t num_heads,
int64_t num_heads_k,
int64_t headdim,
int64_t page_block_size,
int64_t num_blocks,
int64_t causal
);
```
### 参数说明
* `q`decode query tensorshape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16`
* `k_cache_paged`paged key cacheshape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16`
* `v_cache_paged`paged value cacheshape `(num_blocks, page_block_size, num_heads_k, headdim)`,连续 `bf16`
* `output`输出缓冲区shape `(batch_size, seqlen_q, num_heads, headdim)`,连续 `bf16`
* `cache_seqlens`:每个 batch 的 KV 长度shape `(batch_size)`,连续 `int32`
* `block_table`:每个 batch 的 page 映射表shape `(batch_size, num_blocks / batch_size)`,连续 `int32`
* `seqlen_q`query 长度,评测中固定为 `1`
* `page_block_size`page size评测中固定为 `16`
* `causal`:是否启用 causal mask评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 launch 配置并启动 CUDA kernel。为保证计时准确不建议在 `run_kernel` 内部做 `cudaDeviceSynchronize()` 或显式同步。

View File

@ -0,0 +1,57 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "tilelang"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import tilelang
import tilelang.language as T
from tilelang import jit
real_kernel = None
@jit
def build_kernel(*args):
@T.prim_func
def kernel(*args):
...
return kernel
def run_kernel(
q, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
k_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
v_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
output, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
cache_seqlens, # Tensor[int32], shape (batch_size)
block_table, # Tensor[int32], shape (batch_size, num_blocks / batch_size)
batch_size, # int64
seqlen_k, # int64
seqlen_q, # int64
num_heads, # int64
num_heads_k, # int64
headdim, # int64
page_block_size, # int64
num_blocks, # int64
causal, # int64
):
global real_kernel
if real_kernel is None:
real_kernel = build_kernel(...)
real_kernel(q, k_cache_paged, v_cache_paged, output,
cache_seqlens, block_table,
batch_size, seqlen_k, seqlen_q, num_heads,
num_heads_k, headdim, page_block_size, num_blocks, causal)
```
### 参数说明
* `q`decode query tensor连续 `bfloat16`
* `k_cache_paged/v_cache_paged`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `cache_seqlens/block_table`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 TileLang kernel。

View File

@ -0,0 +1,45 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "triton"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import triton
import triton.language as tl
@triton.jit
def your_kernel(...):
...
def run_kernel(
q, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
k_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
v_cache_paged, # Tensor[bf16], shape (num_blocks, page_block_size, num_heads_k, headdim)
output, # Tensor[bf16], shape (batch_size, seqlen_q, num_heads, headdim)
cache_seqlens, # Tensor[int32], shape (batch_size)
block_table, # Tensor[int32], shape (batch_size, num_blocks / batch_size)
batch_size, # int64
seqlen_k, # int64
seqlen_q, # int64
num_heads, # int64
num_heads_k, # int64
headdim, # int64
page_block_size, # int64
num_blocks, # int64
causal, # int64
):
...
```
### 参数说明
* `q`decode query tensor连续 `bfloat16`
* `k_cache_paged/v_cache_paged`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `cache_seqlens/block_table`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 Triton kernel。

View File

@ -0,0 +1,9 @@
---
sectionTitle: "输入格式"
type: "Text"
---
本题输入由评测程序在 GPU 上构造,并按接口约定中的顺序传入 `run_kernel`
`q/k_cache_paged/v_cache_paged/output` 均为连续 `torch.bfloat16` CUDA tensor`cache_seqlens/block_table` 均为连续 `torch.int32` CUDA tensor。
KV cache layout 固定为 `flash_attn_with_kvcache` 的 paged cache 布局:`(num_blocks, page_block_size, num_heads_k, headdim)`。

View File

@ -0,0 +1,5 @@
---
sectionTitle: "输出格式"
type: "Text"
---
输出写入 `output`shape 为 `(batch_size, 1, num_heads, headdim)`,类型为 `bfloat16`

View File

@ -0,0 +1,12 @@
---
sectionTitle: "样例"
type: "Text"
---
`batch_size = 1`、`seqlen_k = 512`、`page_block_size = 16`,则每个序列需要访问 `32` 个有效 page
```text
cache_seqlens = [512]
block_table.shape = (1, num_blocks)
```
`t` 个 KV token 位于 `block_table[0, t / 16]` 指向的物理 page 中page 内偏移为 `t % 16`