Compare commits
No commits in common. "master" and "master" have entirely different histories.
|
|
@ -1,79 +0,0 @@
|
|||
# =========================
|
||||
# macOS
|
||||
# =========================
|
||||
.DS_Store
|
||||
.AppleDouble
|
||||
.LSOverride
|
||||
Icon?
|
||||
._*
|
||||
.Spotlight-V100
|
||||
.Trashes
|
||||
.fseventsd
|
||||
|
||||
# =========================
|
||||
# IDE / Editor
|
||||
# =========================
|
||||
.vscode/
|
||||
.idea/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# =========================
|
||||
# Python
|
||||
# =========================
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.pyo
|
||||
*.pyd
|
||||
.pytest_cache/
|
||||
.mypy_cache/
|
||||
.ruff_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
.env
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
|
||||
# =========================
|
||||
# Jupyter
|
||||
# =========================
|
||||
.ipynb_checkpoints/
|
||||
|
||||
# =========================
|
||||
# Build / Packaging
|
||||
# =========================
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
.eggs/
|
||||
pip-wheel-metadata/
|
||||
|
||||
# =========================
|
||||
# C / C++ / CMake
|
||||
# =========================
|
||||
CMakeFiles/
|
||||
CMakeCache.txt
|
||||
cmake-build-*/
|
||||
Makefile
|
||||
*.o
|
||||
*.so
|
||||
*.dylib
|
||||
*.dll
|
||||
*.a
|
||||
*.lib
|
||||
|
||||
# =========================
|
||||
# Logs / Temp
|
||||
# =========================
|
||||
*.log
|
||||
*.tmp
|
||||
*.temp
|
||||
logs/
|
||||
tmp/
|
||||
|
||||
# =========================
|
||||
# OS / Tool caches
|
||||
# =========================
|
||||
.cache/
|
||||
329
FAQ.md
329
FAQ.md
|
|
@ -1,329 +0,0 @@
|
|||
# 常见问题 FAQ
|
||||
|
||||
> 最后整理:2026-07-13
|
||||
|
||||
本文档汇总沐曦“揭榜挂帅”两项赛题的常见问题。使用 `Ctrl+F`(Windows/Linux)或 `⌘F`(macOS)搜索关键词。
|
||||
|
||||
环境版本、评测配置、时间安排可能调整。请以仓库 README、XPU-OJ 公告、比赛群通知和容器中的实际版本为准。发现内容过期或未覆盖的问题,请[提交 Issue](https://gitlink.org.cn/metax-maca/op_optimization/issues)。
|
||||
|
||||
## 快速导航
|
||||
|
||||
- [重要入口与咨询方式](#entry):XPU-OJ 开放情况、问题反馈渠道和赛事联系人。
|
||||
- [XPU-OJ、提交与排行榜](#xpuoj):账号申请、提交环境、测试样例、评测硬件和排名指标。
|
||||
- [赛题一:TileLang 与 Fused MoE](#track-one):提交限制、算子实现、Baseline、性能优化和决赛加分规则。
|
||||
- [赛题二:AI Agent 与推理算子库](#track-two):任务范围、Agent 参与证明、Baseline、测试参数和上游代码引用。
|
||||
- [评测规则与通用技术问题](#evaluation):技术资料发布、正确性与稳定性要求,以及 FAQ 内容纠错。
|
||||
- [环境、镜像与算力资源](#environment):比赛镜像、MACA 与 PyTorch 版本、开发工具、算力券和资源申请。
|
||||
- [报名、组队与资格审核](#registration):参赛资格、跨校组队、指导教师、材料盖章和审核流程。
|
||||
|
||||
<a id="entry"></a>
|
||||
|
||||
## 重要入口与咨询方式
|
||||
|
||||
<a id="q-xpuoj-open"></a>
|
||||
|
||||
### ❓ 问题 1:XPU-OJ 平台是否已经上线?
|
||||
|
||||
**回答:** XPU-OJ 已开放。账号申领流程和使用指南见[赛事 XPU-OJ 账号申领说明](赛事XPUOJ账号申领说明.md)。
|
||||
|
||||
<a id="q-contact"></a>
|
||||
|
||||
### ❓ 问题 2:两个赛题的联系人不同,遇到问题应该联系谁?
|
||||
|
||||
**回答:** 请先通过 [GitLink Issue](https://gitlink.org.cn/metax-maca/op_optimization/issues) 提问,维护者会把可复用的答案更新到本文档。比赛群用于接收赛事通知和临时信息。需要单独沟通时,请按对应比赛方案联系章老师或杨老师。
|
||||
|
||||
## XPU-OJ、提交与排行榜
|
||||
|
||||
<a id="q-moe-score-decrease"></a>
|
||||
|
||||
### ❓ 问题 1:XPUOJ测评 MoE 耗时减少了但是分数反而降低了
|
||||
|
||||
**回答:**
|
||||
|
||||
针对近期部分同学反馈的“XPU.OJ”第三方评测系统中基线(baseline)不稳定的问题,我们高度重视,并已第一时间组织排查与测试。在此,我们对因此给大家带来的困扰深表歉意,也衷心感谢各位同学提出的宝贵意见。
|
||||
目前,相关问题已修复完毕。为确保评测的公平性与准确性,我们将对现有榜单进行清空处理。历史提交记录仍可查看,但后续排名将统一以基线修复后重新提交的算子成绩为准。
|
||||
|
||||
比赛期间,我们将持续关注系统运行状态,也欢迎大家继续向我们反馈建议。
|
||||
祝大家比赛顺利,取得理想成绩!
|
||||
|
||||
<a id="q-xpuoj-environment"></a>
|
||||
|
||||
### ❓ 问题 2:XPU-OJ 与模力方舟的运行环境一致吗?
|
||||
|
||||
**回答:** 排行榜评测环境与模力方舟开发环境保持一致。版本调整时以 XPU-OJ 公告为准。
|
||||
|
||||
<a id="q-xpuoj-account"></a>
|
||||
|
||||
### ❓ 问题 3:如何申请 XPU-OJ 账号?
|
||||
|
||||
**回答:** 请按照[赛事 XPU-OJ 账号申领说明](赛事XPUOJ账号申领说明.md)提交申请。账号发放进度以赛事通知和回复邮件为准。
|
||||
|
||||
<a id="q-official-ranking"></a>
|
||||
|
||||
### ❓ 问题 4:赛题一 MoE 初赛排名以哪个入口为准?
|
||||
|
||||
**回答:** 正式排名和初筛结果以 XPU-OJ 的评测结果为准。Sample benchmark 用于本地功能验证、调试和性能对比,不作为正式榜单依据。
|
||||
|
||||
<a id="q-test-cases"></a>
|
||||
|
||||
### ❓ 问题 5:MACA C++、Triton 和 TileLang 是否分别设榜?
|
||||
|
||||
**回答:** 不按语言分别设榜。每个任务支持 MACA C++、Triton 和 TileLang,团队可以使用一种或多种语言提交。排行榜采用通过正确性和稳定性测试后的最高成绩。
|
||||
|
||||
<a id="q-ranking-metric"></a>
|
||||
|
||||
### ❓ 问题 6:排行榜使用 latency、speedup 还是综合 score?
|
||||
|
||||
**回答:** 当前 XPU-OJ 以 speedup 作为核心排名指标。赛事方调整计算方式时,以 XPU-OJ 公告为准。
|
||||
|
||||
<a id="track-one"></a>
|
||||
|
||||
## 赛题一:TileLang 与 Fused MoE
|
||||
|
||||
<a id="q-track-one-submission"></a>
|
||||
|
||||
### ❓ 问题 1:赛题一的提交要求有哪些调整?
|
||||
|
||||
**回答:** 正式提交禁止使用 `MACA Maca running` 方式,也不能使用 PyTorch 实现算子。参赛者需使用 TileLang 实现并提交。
|
||||
|
||||
<a id="q-ops-reference"></a>
|
||||
|
||||
### ❓ 问题 2:OPS 目录中的 TileLang、CUDA、CUTLASS 和 MACA 代码有什么用途?
|
||||
|
||||
**回答:** 这些代码用于解释算子的实现原理和设计思路,可作为 TileLang 实现的参考。
|
||||
|
||||
<a id="q-baseline-modification"></a>
|
||||
|
||||
### ❓ 问题 3:官方 Baseline 可以修改到什么范围?
|
||||
|
||||
**回答:** 参赛者可以重新设计和优化算子实现,但需保持与统一 Workload 测试框架的接口兼容。
|
||||
|
||||
<a id="q-gemm-optimization"></a>
|
||||
|
||||
### ❓ 问题 4:GEMM 计算中可以引入其他优化策略吗?
|
||||
|
||||
**回答:** 可以,前提是实现符合赛题规则和评测要求。
|
||||
|
||||
<a id="q-benchmark-modification"></a>
|
||||
|
||||
### ❓ 问题 5:可以优化 `fusedmoe_benchmark` 吗?
|
||||
|
||||
**回答:** 本地修改 benchmark 不会提高正式成绩,XPU-OJ 使用赛事方的评测框架。参赛者应把优化工作放在规定的算子实现和允许修改的接口上。
|
||||
|
||||
<a id="q-forward-modification"></a>
|
||||
|
||||
### ❓ 问题 6:可以修改 `fusedmoe_benchmark.py` 中的 MoE forward 吗?
|
||||
|
||||
**回答:** 正式成绩以 XPU-OJ 的独立评测为准。本地修改 forward 不能替代对提交算子的优化,也不会改变赛事方的评测代码。
|
||||
|
||||
<a id="q-async-copy"></a>
|
||||
|
||||
### ❓ 问题 7:赛题一允许使用异步拷贝吗?
|
||||
|
||||
**回答:** 不允许。当前比赛规则禁用异步拷贝。
|
||||
|
||||
<a id="q-baseline-performance"></a>
|
||||
|
||||
### ❓ 问题 8:Fused MoE 初赛成绩如何影响决赛?
|
||||
|
||||
**回答:** 当前规则只给初赛前 10 名决赛加分,第 11 名及之后不获得额外初赛加分。
|
||||
|
||||
<a id="track-two"></a>
|
||||
|
||||
## 赛题二:AI Agent 与推理算子库
|
||||
|
||||
<a id="q-track-two-update"></a>
|
||||
|
||||
### ❓ 问题 1:赛题二的内容有什么调整?
|
||||
|
||||
**回答:** “Agent 推理算子库优化 - FlashAttention KV Cache Decode”新增 `mctlass/cute` 要求。参赛者需基于 `mctlass/cute` 实现或优化对应任务。
|
||||
|
||||
<a id="q-track-two-selection"></a>
|
||||
|
||||
### ❓ 问题 2:赛题二可以选择几个任务?
|
||||
|
||||
**回答:** 参赛团队可以从 FlashInfer、FlashAttention、MCTLASS/Fused MoE 等方向选择一项或多项。提交多个任务时,每个任务按赛事规则取有效最高成绩。支持语言包括 Triton、MXMACA C++ 和 TileLang。
|
||||
|
||||
<a id="q-agent-proof"></a>
|
||||
|
||||
### ❓ 问题 3:如何证明 AI Agent 参与了优化过程?
|
||||
|
||||
**回答:** 参赛团队应保留 Agent 配置、Skill 文件、关键提示词、操作日志、代码变更记录、测试结果和复现实验步骤,等待赛事方发布核验细则。
|
||||
|
||||
<a id="q-agent-baseline"></a>
|
||||
|
||||
### ❓ 问题 4:Agent 赛题的性能 baseline 使用哪个版本?
|
||||
|
||||
**回答:** 当前评测使用赛事方提供的 baseline,标准环境为 `PyTorch-Agent / 2.8.0 / Python 3.12 / MACA 3.7.1.5`。赛事方变更 baseline 或评测方式时会发布通知。
|
||||
|
||||
<a id="q-mla-dimensions"></a>
|
||||
|
||||
### ❓ 问题 5:MLA 的 `QK dim = 576, VO dim = 512` 与 `race_tests` 参数冲突吗?
|
||||
|
||||
**回答:** 不冲突。`race_tests` 中的 `dim=512, pe_dim=64` 对应 `QK dim = 576, V dim = 512`。
|
||||
|
||||
<a id="q-nsa-ranking"></a>
|
||||
|
||||
### ❓ 问题 6:NSA 的 109 个测试 case 如何计算榜单成绩?
|
||||
|
||||
**回答:** XPU-OJ 通过统一接口统计测试总运行时间,再根据 baseline 计算整体 speedup。
|
||||
|
||||
<a id="q-upstream-code"></a>
|
||||
|
||||
### ❓ 问题 7:可以引用或修改 FlashInfer、FlashAttention 等上游代码吗?
|
||||
|
||||
**回答:** 可以。参赛团队需遵守上游项目许可证,保留版权和许可证声明,并在提交材料中说明引用范围、迁移工作和自主优化内容。
|
||||
|
||||
<a id="evaluation"></a>
|
||||
|
||||
## 环境、镜像与算力资源
|
||||
|
||||
<a id="q-environment-image"></a>
|
||||
|
||||
### ❓ 问题 1:比赛使用哪个在线算力环境?
|
||||
|
||||
**回答:** 比赛使用模力方舟沐曦算力专区。仓库 README 当前标注的统一镜像为:
|
||||
|
||||
```text
|
||||
PyTorch-Agent / 2.8.0 / Python 3.12 / MACA 3.7.1.5
|
||||
```
|
||||
|
||||
创建实例和连接环境的步骤见[模力方舟快速使用 SOP](模力方舟快速使用SOP.md)。
|
||||
|
||||
<a id="q-download-maca"></a>
|
||||
|
||||
### ❓ 问题 2:如何获取 MACA 镜像或安装包?
|
||||
|
||||
**回答:** 参赛者可以在[模力方舟沐曦算力专区](https://ai.gitee.com/compute/metax)选择 `PyTorch-Agent` 镜像。需要单独获取软件包时,请前往[沐曦开发者社区软件中心](https://developer.metax-tech.com/softnova/docker?chip_name=%E6%9B%A6%E4%BA%91C500%E7%B3%BB%E5%88%97&package_kind=AI&dimension=docker),并选择与比赛标准环境一致的 MACA `3.7.1.5` 版本。
|
||||
|
||||
<a id="q-download-pytorch"></a>
|
||||
|
||||
### ❓ 问题 3:比赛使用的 PyTorch 镜像可以下载吗?
|
||||
|
||||
**回答:** 可以。请在[沐曦开发者社区](https://developer.metax-tech.com/)或其 [PyTorch 镜像列表](https://developer.metax-tech.com/softnova/docker?chip_name=%E6%9B%A6%E4%BA%91C500%E7%B3%BB%E5%88%97&package_kind=AI&dimension=docker&deliver_type=%E5%88%86%E5%B1%82%E5%8C%85&ai_frame=pytorch)中查询。下载前请核对比赛镜像的 Python、PyTorch 和 MACA 版本。
|
||||
|
||||
<a id="q-version-mismatch"></a>
|
||||
|
||||
### ❓ 问题 4:页面标注的 MACA 版本与容器内版本不一致怎么办?
|
||||
|
||||
**回答:** 比赛标准版本为 MACA `3.7.1.5`。
|
||||
|
||||
<a id="q-pytorch-source"></a>
|
||||
|
||||
### ❓ 问题 5:Linux 版本的 mcprofiler 是否可用?
|
||||
|
||||
**回答:** mcprofiler 的 Linux 版本已经打包进模力方舟上的 pytorch-agent 比赛镜像。
|
||||
|
||||
<a id="q-compute-coupons"></a>
|
||||
|
||||
### ❓ 问题 6:算力券按团队还是按个人领取?
|
||||
|
||||
**回答:** 新人礼和启悟社区算力券按学生个人发放,符合条件的团队成员均可领取。团队主申请人的额度用完后,其他成员可以继续申请资源。
|
||||
|
||||
- [沐曦开发者社区新人礼](https://developer.metax-tech.com/activities/6)
|
||||
- [启悟社区学生算力券](https://developer.metax-tech.com/activities/11)
|
||||
- [赛事算力券活动](https://developer.metax-tech.com/activities/17)
|
||||
|
||||
<a id="q-more-compute"></a>
|
||||
|
||||
### ❓ 问题 7:算力额度不足时可以追加申请吗?
|
||||
|
||||
**回答:** 可以先领取上述活动中的算力券。仍需额外资源时,请发送需求邮件至 `opensource@metax-tech.com`。
|
||||
|
||||
<a id="q-commercial-agent-cost"></a>
|
||||
|
||||
<a id="registration"></a>
|
||||
|
||||
## 报名、组队与资格审核
|
||||
|
||||
<a id="q-register-both"></a>
|
||||
|
||||
### ❓ 问题 1:同一团队或个人可以同时参加两个赛题吗?
|
||||
|
||||
**回答:** 可以。同一团队或个人可以同时报名两个赛题。
|
||||
|
||||
<a id="q-register-multiple-tracks"></a>
|
||||
|
||||
### ❓ 问题 2:同一名学生可以报名不同赛道的不同赛题吗?
|
||||
|
||||
**回答:** 可以。赛事不统一限制学生报名不同赛道或不同赛题,但同一作品不得用相同核心技术内容重复申报不同赛题。
|
||||
|
||||
<a id="q-new-graduate"></a>
|
||||
|
||||
### ❓ 问题 3:本科应届毕业、尚未正式入学的研一新生可以报名吗?
|
||||
|
||||
**回答:** 可以。参赛者可联系原本科学校完成认证手续,并以本科生身份报名。
|
||||
|
||||
<a id="q-advisor-required"></a>
|
||||
|
||||
### ❓ 问题 4:参赛必须配备指导教师吗?
|
||||
|
||||
**回答:** 不强制。填写指导教师时,每支队伍可以设置 1 至 3 名指导教师。
|
||||
|
||||
<a id="q-advisor-team-limit"></a>
|
||||
|
||||
### ❓ 问题 5:一名指导教师最多可以指导几支队伍?
|
||||
|
||||
**回答:** 赛事暂未设置统一的硬性上限。指导教师应根据可投入的时间控制队伍数量。
|
||||
|
||||
<a id="q-cross-school-stamp"></a>
|
||||
|
||||
### ❓ 问题 6:跨校组队时,报名表应该由哪所学校盖章?
|
||||
|
||||
**回答:** 资格审批阶段,每名参赛者需到本人学校的校团委或院团委完成盖章确认。后续材料由团队牵头学生统一整理和提交。
|
||||
|
||||
<a id="q-upload-stamped-form"></a>
|
||||
|
||||
### ❓ 问题 7:提交报名后还可以补充已盖章的报名表扫描件吗?
|
||||
|
||||
**回答:** 审核人员发现材料缺少盖章时,会退回申请。团队补齐材料后可以重新提交。尚未完成盖章的团队应先与学院、校团委或学校相关部门确认办理方式。
|
||||
|
||||
<a id="q-stamp-department"></a>
|
||||
|
||||
### ❓ 问题 8:资格审查材料应该加盖哪个部门的公章?
|
||||
|
||||
**回答:** 各高校的管理口径不同。教务处、学生处等学籍或学生管理部门通常可以办理,参赛团队应以本校校团委或相关管理部门的要求为准。
|
||||
|
||||
<a id="q-no-youth-league"></a>
|
||||
|
||||
### ❓ 问题 9:学校未设校团委,可以用院系公章替代吗?
|
||||
|
||||
**回答:** 赛事原则上要求校级部门公章。学校未设校团委时,可以联系校级学工、双创或教务部门盖章,并提交情况说明。院系公章不能直接替代校级部门公章。
|
||||
|
||||
<a id="q-student-status-proof"></a>
|
||||
|
||||
### ❓ 问题 10:学校无法配合盖章,可以用学籍证明替代吗?
|
||||
|
||||
**回答:** 不可以。参赛团队应使用报名系统导出的报名表,并按要求完成学校盖章。
|
||||
|
||||
<a id="q-public-notice"></a>
|
||||
|
||||
### ❓ 问题 11:公示材料需要包含哪些内容?
|
||||
|
||||
**回答:** 请参考赛事工作群发布的参考文本,并按学校要求调整。跨校团队涉及的学校应分别在学校官网公示,公示渠道原则上使用学校官网。
|
||||
|
||||
<a id="q-review-deadline"></a>
|
||||
|
||||
### ❓ 问题 12:报名审核需要在报名截止日前完成吗?
|
||||
|
||||
**回答:** 原则上需要。往届出现过系统延后关闭的情况,但本届参赛团队不应据此推迟材料提交或审核。
|
||||
|
||||
<a id="q-review-flow"></a>
|
||||
|
||||
### ❓ 问题 13:报名材料的审核顺序是什么?
|
||||
|
||||
**回答:** 学生提交材料后,学校校团委先审核;学校审核通过后,企业再审核。企业审核通过即视为报名成功。
|
||||
|
||||
<a id="q-school-review-account"></a>
|
||||
|
||||
### ❓ 问题 14:后台显示“校团委审核”,但学校不了解审核事项,怎么办?
|
||||
|
||||
**回答:** 校团委需在报名系统内完成审核。省级团委通常会向各高校团委发放账号和密码。学校未收到或不了解安排时,请学校联系省级团委确认。
|
||||
|
||||
## FAQ 维护约定
|
||||
|
||||
1. 参赛者通过 Issue 提交问题。
|
||||
2. 维护者确认答案后更新本文档。
|
||||
3. 每个答案保留稳定锚点;需要时附上来源 Issue 或公告。
|
||||
4. 维护者在原 Issue 中回复 FAQ 锚点链接,并关闭已经解决的问题。
|
||||
5. 涉及版本、日期、评测参数的答案应标注确认日期。
|
||||
22
README.md
22
README.md
|
|
@ -1,22 +1,5 @@
|
|||
# 降低Token 成本,攻坚国产推理生态|沐曦两大赛题登陆 2026 揭榜挂帅擂台赛,邀青年共破局!
|
||||
|
||||
## 常用入口
|
||||
|
||||
- [常见问题 FAQ](FAQ.md)
|
||||
- [沐曦通用GPU MXMACA编译器内建函数编程指南](https://developer.metax-tech.com/api/client/document/preview/1395/index.html)
|
||||
|
||||
## XPU-OJ 基线修复及榜单调整通知
|
||||
|
||||
针对近期部分同学反馈的“XPU.OJ”第三方评测系统中基线(baseline)不稳定的问题,我们高度重视,并已第一时间组织排查与测试。在此,我们对因此给大家带来的困扰深表歉意,也衷心感谢各位同学提出的宝贵意见。
|
||||
目前,相关问题已修复完毕。为确保评测的公平性与准确性,我们将对现有榜单进行清空处理。历史提交记录仍可查看,但后续排名将统一以基线修复后重新提交的算子成绩为准。
|
||||
|
||||
比赛期间,我们将持续关注系统运行状态,也欢迎大家继续向我们反馈建议。
|
||||
祝大家比赛顺利,取得理想成绩!
|
||||
|
||||
- XPU-OJ 地址:[https://xpuoj.com/](https://xpuoj.com/)
|
||||
- 账号申领说明:[赛事 XPU-OJ 账号申领说明](赛事XPUOJ账号申领说明.md)
|
||||
- 账号申领邮箱:`opensource@metax-tech.com`
|
||||
|
||||
2026 年度中国青年科技创新「揭榜挂帅」擂台赛正式启幕。沐曦股份重磅发布两大 AI 算力硬核榜题,聚焦国产 GPU 大模型推理算子优化,以硬核赛事搭建科研攻关平台,邀全国青年学子、科研人才揭榜攻坚,用技术重构推理效率,用创新拉低每 Token 算力成本!
|
||||
|
||||
## 两大重磅赛题 直击推理成本核心痛点
|
||||
|
|
@ -46,7 +29,7 @@
|
|||
|
||||
### 赛题二:基于 AI Agent 开发范式的国产 GPU 大模型推理算子库优化
|
||||
|
||||
大模型推理具有高并发、长序列、高<EFBFBD><EFBFBD><EFBFBD>用频次等特点,FlashInfer、FlashAttention、Fused MoE 等核心算子直接决定模型服务的吞吐、延迟与显存开销,影响单 Token 综合推理成本。
|
||||
大模型推理具有高并发、长序列、高调用频次等特点,FlashInfer、FlashAttention、Fused MoE 等核心算子直接决定模型服务的吞吐、延迟与显存开销,影响单 Token 综合推理成本。
|
||||
|
||||
本赛题面向沐曦国产 GPU 及 MXMACA 软件栈,鼓励参赛团队构建或使用 AI Agent / Skill 工作流,围绕推理算子库开展代码理解、算子迁移、性能分析、Kernel 优化、自动调优、Benchmark 验证和多轮迭代,探索“Agent 驱动算子优化”的新型开发范式。
|
||||
|
||||
|
|
@ -63,11 +46,10 @@
|
|||
**赛题二相关资料**
|
||||
|
||||
- [赛题二方案:基于 AI Agent 开发范式的国产 GPU 大模型推理算子库优化方案](基于AI%20Agent开发范式的国产GPU大模型推理算子库优化/基于AI%20Agent开发范式的国产GPU大模型算子推理库优化方案.md)
|
||||
- [赛题二选手入口](基于AI%20Agent开发范式的国产GPU大模型推理算子库优化/选手入口.md)
|
||||
- [模力方舟 Agent 部署准备教程](基于AI%20Agent开发范式的国产GPU大模型推理算子库优化/模力方舟Agent部署准备教程.md)
|
||||
- [赛题二说明及资料参考](基于AI%20Agent开发范式的国产GPU大模型推理算子库优化/赛题说明.md)
|
||||
|
||||
### **两个赛题统一使用模力方舟上的镜像PyTorch-Agent / 2.8.0 / Python 3.12 / maca 3.7.1.5**
|
||||
###**两个赛题统一使用模力方舟上的镜像PyTorch-Agent / 2.8.0 / Python 3.12 / maca 3.7.2.1**
|
||||
|
||||
## 参赛对象
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
|
|
@ -0,0 +1,773 @@
|
|||
# FlashInfer 关键算子迁移与优化
|
||||
|
||||
## 一、教程定位
|
||||
|
||||
本教程是参赛训练课程的 **FlashInfer Baseline 入门** 模块,主要帮助用户快速跑通 FlashInfer 的最小可运行流程。完成本教程后,用户应能够完成源码编译、API 调用、正确性测试和 Benchmark 测试,并记录一份 baseline 性能结果,为后续算子优化提供对比基准。
|
||||
|
||||
## 二、完成本模块你将能够
|
||||
|
||||
1. 理解 FlashInfer Attention Kernel 的基本作用与适用场景;
|
||||
2. 完成 FlashInfer 环境、工具链的准备与源码编译;
|
||||
3. 跑通 BatchDecode、BatchPrefill、MLA 等典型算子的 API 调用示例;
|
||||
4. 完成各算子在不同参数配置下的 Benchmark 测试;
|
||||
5. 输出各算子的 Baseline 性能结果记录表,为后续算子优化提供对比基准。
|
||||
6. 理解 XPU-OJ 评测 `run_kernel` 接口与精度要求。
|
||||
7. 理解 Baseline 与 XPU-OJ 评测题包之间的关系,能够根据题包接口实现一个最小正确版 `run_kernel`。
|
||||
|
||||
|
||||
## 三、适用对象
|
||||
|
||||
**适合人群**
|
||||
|
||||
* 参赛选手:需要完成 Baseline 入门模块,为后续算子优化做准备
|
||||
* 软件开发者和Vibe Coding开发:希望从事AI相关行业开发,以及用智能体方式来做开发工作
|
||||
* LLM 推理开发者:希望了解 FlashInfer Attention Kernel 的性能表现
|
||||
* 算子优化工程师:希望基于MXMACA软件栈在沐曦国产 GPU 上做算子迁移和优化
|
||||
|
||||
**前置基础**
|
||||
|
||||
* Python 基础:能够运行和修改 Python 脚本
|
||||
* PyTorch 基础:了解MXMACA软化栈的使用
|
||||
* Linux 命令行:能够使用终端执行命令
|
||||
* 了解 Attention 机制:理解 Q/K/V、KV Cache 等基本概念
|
||||
|
||||
|
||||
## 四、前置准备
|
||||
|
||||
开始实战前,请确认你已经完成以下准备:
|
||||
|
||||
### 获得 GPU
|
||||
|
||||
1. [点击获取算力券](https://developer.metax-tech.com/activities/6),首次登录需要使用邮箱或者手机号进行注册
|
||||
|
||||
2. 登录成功后验证邮箱
|
||||
|
||||
3. 提交申请获得兑换码
|
||||
|
||||
4. 兑换算力和登陆平台:
|
||||
|
||||
- [访问模力方舟官网](https://ai.gitee.com/),进入费用中心 - 算力券 ,点击右上角 “兑换”。
|
||||
|
||||
- 进入算力容器,选择沐曦,租用算力,建议优先选 16G 显存 / 32G 显存,如下图
|
||||
|
||||

|
||||
|
||||
- 创建实例。基础镜像:`maca-pytorch:3.7.1.5-torch2.8-py312-ubuntu24.04-amd64`
|
||||
|
||||
- 选择工具-lab进入实例环境
|
||||
|
||||
- 步骤6:在JupyterLab Terminal中检查运行环境的配置,确认沐曦 GPU 可见--可以使用 `mx-smi` 命令查看
|
||||
|
||||
|
||||
### Python 环境
|
||||
|
||||
``` bash
|
||||
pip install flashinfer torch pandas numpy
|
||||
```
|
||||
|
||||
### OpenCode 安装
|
||||
|
||||
``` bash
|
||||
curl -fsSL https://opencode.ai/install | bash
|
||||
opencode
|
||||
```
|
||||
|
||||
### 代码准备
|
||||
|
||||
测试脚本和 Benchmark 脚本:
|
||||
|
||||
| 命令 | 说明 |
|
||||
| --- | --- |
|
||||
| `python bench_batch_decode.py` | 运行 Batch Decode 基准测试 |
|
||||
| `python bench_batch_prefill_paged.py` | 运行 Batch Prefill (Paged KV Cache) 基准测试 |
|
||||
| `python bench_batch_prefill_ragged.py` | 运行 Batch Prefill (Ragged KV Cache) 基准测试 |
|
||||
| `python bench_batch_mla.py` | 运行 MLA (Multi-head Latent Attention) 基准测试 |
|
||||
|
||||
## 五、知识预备
|
||||
|
||||
### LLM 推理阶段重要概念:
|
||||
|
||||
- **Prefill 阶段**:Prefill 阶段是指处理输入 prompt 的阶段
|
||||
|
||||
- 输入:用户一次性给出的完整 prompt,长度为 seq\_len
|
||||
|
||||
- 计算:对 prompt 中的每个 token 并行计算注意力,生成第一个输出 token 及 KV cache
|
||||
|
||||
- 特点:这是 **计算密集型(compute-bound)** 阶段,因为需要做完整的 `seq\_len * seq\_len` 注意力矩阵乘法
|
||||
|
||||
- **Decode 阶段**:
|
||||
|
||||
- 每次只生成 1 个 token,利用 prefill 阶段填充好的 KV cache 做自回归生成
|
||||
|
||||
- **显存带宽密集型(memory-bound)**,瓶颈在从显存读取 KV cache 而非计算
|
||||
|
||||
**Prefill = 并行处理用户输入,Decode = 逐个生成回答 token**
|
||||
|
||||
## 六、项目实践 -- FlashInfer Baseline
|
||||
|
||||
**目标:** 在赛事镜像中完成 FlashInfer Ragged Prefill 算子的 Baseline Benchmark,理解从 Benchmark 到 XPU-OJ 评测提交的完整流程,为后续算子优化建立性能基线。
|
||||
|
||||
### 在赛事镜像中运行 FlashInfer Baseline Benchmark
|
||||
|
||||
#### Step 1:检查运行环境
|
||||
|
||||
**目标:** 确认当前环境满足本模块运行要求。
|
||||
|
||||
**操作:** 进入 Terminal 检查 GPU、Python、编译工具和依赖版本。
|
||||
|
||||

|
||||
|
||||
**命令示例:**
|
||||
|
||||
```Bash
|
||||
# 检查沐曦 GPU 状态
|
||||
mx-smi
|
||||
|
||||
# 检查 Python 版本
|
||||
python --version
|
||||
|
||||
# 检查 PyTorch 是否能识别 GPU
|
||||
python -c "import torch; print(f'GPU available: {torch.cuda.is_available()}'); print(f'GPU count: {torch.cuda.device_count()}')"
|
||||
|
||||
# 检查依赖版本
|
||||
python -c "import torch; print(f'PyTorch {torch.__version__}')"
|
||||
python -c "import einops; print('einops OK')"
|
||||
```
|
||||
|
||||
**预期结果:**
|
||||
|
||||
- `mx-smi` 显示沐曦 GPU 信息
|
||||
|
||||

|
||||
|
||||
- Python 环境正常
|
||||
|
||||

|
||||
|
||||
- `torch.cuda.is_available()` 返回 `True`
|
||||
|
||||

|
||||
|
||||
- 所有依赖版本符合要求
|
||||
|
||||

|
||||
|
||||
|
||||
**常见问题:**
|
||||
|
||||
| 问题 | 解决方法 |
|
||||
| --- | --- |
|
||||
| `mx-smi: command not found` | 确认已配置沐曦 GPU 驱动环境 |
|
||||
| `No GPUs are available` | 检查 MXMACA 驱动是否正确安装/环境变量是否正确配置 |
|
||||
| `ModuleNotFoundError: No module named 'xxx'` | `pip install xxx` |
|
||||
|
||||
#### Step 2:进入项目目录
|
||||
|
||||
**目标:** 进入本模块所需的源码目录[flashinfer_baseline](https://gitlink.org.cn/metax-maca/op_optimization/tree/master/%E5%9F%BA%E4%BA%8EAI%20Agent%E5%BC%80%E5%8F%91%E8%8C%83%E5%BC%8F%E7%9A%84%E5%9B%BD%E4%BA%A7GPU%E5%A4%A7%E6%A8%A1%E5%9E%8B%E6%8E%A8%E7%90%86%E7%AE%97%E5%AD%90%E5%BA%93%E4%BC%98%E5%8C%96%2Fbaselines%2Fflashinfer_baseline)。
|
||||
|
||||
1. 克隆代码仓库
|
||||
|
||||
```bash
|
||||
git clone https://gitlink.org.cn/metax-maca/op_optimization.git
|
||||
```
|
||||
|
||||
2. 准备flashinfer_baseline
|
||||
|
||||
从仓库根目录开始,在 `基于AI Agent开发范式的国产GPU大模型推理算子库优化` 下,找到 `flashinfer_baseline` 文件夹。可以将 `flashinfer_baseline` 整个目录复制到工作目录 `data/` 下。
|
||||
```bash
|
||||
mkdir data
|
||||
cp -r "基于AI Agent开发范式的国产GPU大模型推理算子库优化/baselines/flashinfer_baseline" data/
|
||||
```
|
||||
|
||||
3. 切换到 FlashInfer_Baseline 项目目录
|
||||
```bash
|
||||
cd data/flashinfer_baseline/FlashInfer_Baseline
|
||||
ls -al
|
||||
```
|
||||
|
||||
**预期结果:**
|
||||
|
||||
```plaintext
|
||||
bench_common.py
|
||||
bench_batch_decode.py
|
||||
bench_batch_prefill_paged.py
|
||||
bench_batch_prefill_ragged.py
|
||||
bench_batch_mla.py
|
||||
README.md
|
||||
...
|
||||
```
|
||||
|
||||
#### Step 3:验证项目脚本
|
||||
|
||||
**目标:** 确认所有基准测试脚本可正常执行。
|
||||
|
||||
**操作:** 检查脚本文件是否存在且可读。
|
||||
|
||||
**命令示例:**
|
||||
|
||||
```bash
|
||||
# 检查脚本文件
|
||||
python -c "import os; scripts = ['bench_common.py', 'bench_batch_decode.py', 'bench_batch_prefill_paged.py', 'bench_batch_prefill_ragged.py', 'bench_batch_mla.py']; [print(f'✓ {s}') if os.path.exists(s) else print(f'✗ {s} missing') for s in scripts]"
|
||||
|
||||
# 测试脚本导入
|
||||
python -c "from bench_common import setup_workspace, get_csv_path; print('脚本导入正常')"
|
||||
```
|
||||
|
||||
**预期结果:**
|
||||
|
||||
```plaintext
|
||||
✓ bench_common.py
|
||||
✓ bench_batch_decode.py
|
||||
✓ bench_batch_prefill_paged.py
|
||||
✓ bench_batch_prefill_ragged.py
|
||||
✓ bench_batch_mla.py
|
||||
脚本导入正常
|
||||
```
|
||||
|
||||
#### Step 4:运行单算子 Benchmark 并查看测试结果
|
||||
|
||||
**目标:** 执行基准测试,获取 Baseline 性能数据,查看并分析 Benchmark 输出结果。
|
||||
|
||||
**操作:** 运行 Ragged Prefill 基准测试脚本,读取生成的 CSV 结果文件。
|
||||
|
||||
**运行Benchmark命令示例:**
|
||||
|
||||
```bash
|
||||
python bench_batch_prefill_ragged.py
|
||||
```
|
||||
预期结果:(并非真实数据)
|
||||
``` plaintext
|
||||
[BatchPrefillWithRaggedKVCacheWrapper] Starting benchmark, total cases: 48
|
||||
[1/48] bs=1, sl=1024, hd=[128,128]: 0.032ms, 66.67 GB/s, 272.00 TFLOPs
|
||||
[2/48] bs=1, sl=4096, hd=[128,128]: 0.042ms, 197.83 GB/s, 3238.06 TFLOPs
|
||||
[3/48] bs=1, sl=8192, hd=[128,128]: 0.064ms, 261.50 GB/s, 8559.25 TFLOPs
|
||||
...
|
||||
|
||||
Results saved to BatchPrefillWithRaggedKVCacheWrapper_20260626_xxxxxx.csv
|
||||
```
|
||||
|
||||
**常见问题:**
|
||||
|
||||
| 问题 | 解决方法 |
|
||||
| --- | --- |
|
||||
| `out of memory` | 减小 batch\_size 或 seq\_len 参数 |
|
||||
| 运行时间过长 | 脚本会自动调整重复次数,耐心等待 |
|
||||
|
||||
***
|
||||
|
||||
**查看结果命令示例:**
|
||||
|
||||
```bash
|
||||
# 列出所有 CSV 结果文件(按修改时间排序,最新的在最上面)
|
||||
ls -lt *.csv 2>/dev/null || echo "未找到 CSV 文件,请先运行 benchmark"
|
||||
|
||||
# 使用 Python 查看最新结果(自动适配所有 benchmark 类型的列名)
|
||||
python3 -c "
|
||||
import pandas as pd, glob, os
|
||||
|
||||
# 找到所有 CSV 文件,按修改时间取最新的
|
||||
csv_files = sorted(glob.glob('*.csv'), key=os.path.getmtime, reverse=True)
|
||||
if not csv_files:
|
||||
print('未找到 CSV 文件,请先运行 benchmark 脚本')
|
||||
else:
|
||||
latest = csv_files[0]
|
||||
print(f'读取文件: {latest}')
|
||||
df = pd.read_csv(latest)
|
||||
# 动态选择列名:优先显示通用列 + 时间/性能列
|
||||
perf_cols = ['time_ms', 'bandwidth_GB_s', 'tflops']
|
||||
avail_cols = [c for c in df.columns if c in perf_cols or c not in ['api']]
|
||||
# 只保留有意义的分析列(排除 api、seq_len_q 等辅助列)
|
||||
display_cols = [c for c in avail_cols if c not in ('seq_len_q',)]
|
||||
print(df[display_cols].head(10).to_string(index=False))
|
||||
"
|
||||
```
|
||||
|
||||
**预期结果(以 Ragged Prefill 为例):**
|
||||
|
||||
```plaintext
|
||||
读取文件: BatchPrefillWithRaggedKVCacheWrapper_20260626_145454.csv
|
||||
batch_size seq_len num_qo_heads num_kv_heads head_dim_qk head_dim_vo time_ms bandwidth_GB_s tflops
|
||||
1 1024 32 4 128 128 0.031580 66.666667 272.004150
|
||||
1 4096 32 4 128 128 0.042445 197.828709 3238.063402
|
||||
1 8192 32 4 128 128 0.064123 261.499213 8559.251770
|
||||
4 1024 32 4 128 128 0.050221 167.754590 2737.757998
|
||||
4 4096 32 4 128 128 0.101234 332.907816 21734.876630
|
||||
16 1024 32 4 128 128 0.149876 224.887654 14683.437981
|
||||
...
|
||||
```
|
||||
|
||||
### XPU-OJ 在线评测教程
|
||||
|
||||
#### Step 5: 从 Baseline 到 XPU-OJ 提交
|
||||
|
||||
##### 5.1:Baseline 与 XPU-OJ 的关系
|
||||
|
||||
赛事镜像中的 Baseline Benchmark 和 XPU-OJ 在线评测任务不同。Baseline Benchmark 主要用于理解算子调用方式和建立性能基线;XPU-OJ 在线评测用于统一检查选手提交代码的正确性和性能。
|
||||
|
||||
| 维度 | Baseline Benchmark | XPU-OJ 提交 |
|
||||
|------|---------------|------------|
|
||||
| **目的** | 理解算子接口、建立性能基线 | 统一环境下的正确性+性能评测 |
|
||||
| **接口形式** | Python API(`wrapper.plan()` + `wrapper.run()`) | C 接口(`extern "C" void run_kernel(...)`) |
|
||||
| **数据范围** | 多种 head_dim/batch_size/seq_len 组合 | 固定参数范围(以题包为准) |
|
||||
| **验证** | 无自动正确性校验 | 强制通过 `torch.allclose(rtol=1e-2, atol=1e-2)` |
|
||||
| **输出** | CSV 性能记录 | 排行榜得分 |
|
||||
|
||||
跑完 baseline 后,选手需要完成以下转换:
|
||||
1. 从 benchmark 脚本中理解目标 API,例如 BatchPrefillWithRaggedKVCacheWrapper;
|
||||
2. 在对应 OJ 题包中查看 `run_kernel(...)` 接口;
|
||||
3. 对照题包中的输入 shape、数据范围和精度要求;
|
||||
4. 编写自己的 `run_kernel(...)`;
|
||||
5. 提交 OJ,先通过正确性;
|
||||
6. 正确性通过后,再对比 baseline / OJ 耗时继续优化。
|
||||
|
||||
***
|
||||
##### 5.2:选择目标算子
|
||||
|
||||
FlashInfer 方向包含 **4 个可选算子题目**,每个对应独立的 benchmark 脚本、OJ 题包、`run_kernel(...)` 接口和数据范围可能不同。
|
||||
|
||||
| OJ 题号 | 算子类型 | 核心特点 | Benchmark 脚本 | FlashInfer API |
|
||||
|---------|---------------|----------------|----------|----------|
|
||||
| **1** | Ragged Prefill | GQA布局,Q/K/V平坦存储,causal=1 | `bench_batch_prefill_ragged.py` | `BatchPrefillWithRaggedKVCacheWrapper` |
|
||||
| **2** | Paged Prefill | KV Cache分页存储,需解析page table | `bench_batch_prefill_paged.py` | `BatchPrefillWithPagedKVCacheWrapper` |
|
||||
| **3** | MLA Paged Attention | DeepSeek MLA特有,双路Q(nope+pe)/双路Cache(ckv+kpe) |`bench_batch_mla.py` | `BatchMLAPagedAttentionWrapper` |
|
||||
| **4** | Paged Decode | 每次只1个query token,memory-bound |`bench_batch_decode.py` | `BatchDecodeWithPagedKVCacheWrapper` |
|
||||
|
||||
**每个子题的接口参数、数据范围和精度要求以对应 XPU-OJ 题包为准。** 下文以题目 **1 Ragged Prefill** 为例演示从 baseline benchmark 到 XPU-OJ 提交的完整流程。
|
||||
|
||||
#### Step 6:理解 XPU-OJ 评测接口与精度要求
|
||||
**目标:** 明确 Baseline 与最终评测提交之间的关系,理解选手需要实现的内容。
|
||||
|
||||
> 完成 baseline benchmark 后,需要注意 baseline 脚本主要用于建立性能基线,并不需要最终提交。
|
||||
> 最终评测以 XPU-OJ 题包为准,评测程序会调用选手提交代码中的 `run_kernel`,并将输出结果与 baseline 参考结果进行比较。
|
||||
|
||||
下面以 FlashInfer Ragged Prefill 题为例,其中:
|
||||
- `zh_CN/00_题目描述.md`:说明需要实现的算子功能;
|
||||
- `zh_CN/01_接口约定.md`:说明必须实现的 `run_kernel` 函数签名;
|
||||
- `zh_CN/02_数据范围.md`:说明测试范围和精度要求;
|
||||
- `testcase_config.py`:定义测试数据生成、baseline 参考实现和正确性校验方式。
|
||||
|
||||
FlashInfer Ragged Prefill 的校验方式为:
|
||||
```python
|
||||
torch.allclose(output_t.float(), output_ref.float(), rtol=1e-2, atol=1e-2)
|
||||
```
|
||||
|
||||
选手实现的输出需要在上述容差范围内与 baseline 输出一致。不同算子的容差可能不同,正式精度要求以对应 OJ 题包说明为准。
|
||||
|
||||
#### Step 7:登录 XPU-OJ 并进入题目页面
|
||||
|
||||
使用组委会统一发放的账号登录 XPU-OJ,并进入对应赛题页面。
|
||||
|
||||
1. 打开 XPU-OJ 平台:https://xpuoj.com/
|
||||
2. 使用组委会统一发放的账号和初始密码登录 **【后续发布】**;
|
||||

|
||||
|
||||
3. 登录后进入比赛 / 题目列表页面;
|
||||

|
||||
4. 找到对应题目,例如 1 FlashInfer Ragged Prefill;
|
||||

|
||||
5. 点击进入题目详情页,查看题目描述、接口约定、数据范围和提交入口。
|
||||

|
||||
|
||||
#### Step 8:提交 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` 初版
|
||||
在下方参考 prompt 的引导下,Agent 会:
|
||||
1. 读取对应 OJ 题包中的接口约定文档(`01_接口约定.md`),提取 `run_kernel` 函数签名;
|
||||
2. 读取数据范围文档(`02_数据范围.md`),了解输入张量 shape 和精度要求;
|
||||
3. 生成一个能编译通过的最小 `run_kernel` 实现,优先保证接口正确性,不追求性能。
|
||||
|
||||
生成的代码可在本地编译验证后,直接提交到 OJ 上冒烟测试,确认提交链路可用。
|
||||
|
||||
**参考 prompt**
|
||||
|
||||
```plaintext
|
||||
请帮我为 FlashInfer Ragged Prefill 题(OJ 题号 20001)生成一个最小可运行的 run_kernel 实现,要求:
|
||||
1. 阅读题包中的 01_接口约定.md,理解 run_kernel 的函数签名(参数类型、顺序、const 修饰);
|
||||
2. 阅读 02_数据范围.md,了解 head_dim_qk、head_dim_vo 的可能取值;
|
||||
3. 阅读 00_题目描述.md,理解需要实现的注意力计算逻辑;
|
||||
4. 生成一个只使用简单双重循环的 naive 实现(不加 tiling、不加 shared memory),确保:
|
||||
- 函数签名为 extern "C" void run_kernel(...)
|
||||
- 包含必要的头文件(cuda_bf16.h、cuda_runtime.h、stdint.h、math.h)
|
||||
- 支持 GQA(Group Query Attention)的头的映射
|
||||
- 支持 causal mask
|
||||
- 使用 bfloat16 数据类型
|
||||
- scale = 1/sqrt(head_dim_qk)
|
||||
```
|
||||
|
||||

|
||||
|
||||
**OJ 冒烟代码**
|
||||
|
||||
用于最小链路验证。
|
||||
|
||||
```cpp
|
||||
#include <stdint.h>
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <math.h>
|
||||
|
||||
namespace
|
||||
{
|
||||
|
||||
__device__ __forceinline__ float warp_sum(float x)
|
||||
{
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
{
|
||||
x += __shfl_down_sync(0xffffffffu, x, offset);
|
||||
}
|
||||
return __shfl_sync(0xffffffffu, x, 0);
|
||||
}
|
||||
|
||||
__global__ void ragged_prefill_smoke_kernel(
|
||||
const __nv_bfloat16 *__restrict__ q,
|
||||
const __nv_bfloat16 *__restrict__ k,
|
||||
const __nv_bfloat16 *__restrict__ v,
|
||||
__nv_bfloat16 *__restrict__ output,
|
||||
const int32_t *__restrict__ qo_indptr,
|
||||
const int32_t *__restrict__ kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal)
|
||||
{
|
||||
const int lane = threadIdx.x & 31;
|
||||
const int warp_id = threadIdx.x >> 5;
|
||||
const int warps_per_block = blockDim.x >> 5;
|
||||
|
||||
int64_t work = static_cast<int64_t>(blockIdx.x) * warps_per_block + warp_id;
|
||||
const int64_t total = batch_size * seq_len * num_qo_heads;
|
||||
if (work >= total)
|
||||
return;
|
||||
|
||||
const int64_t qo_head = work % num_qo_heads;
|
||||
work /= num_qo_heads;
|
||||
const int64_t q_pos = work % seq_len;
|
||||
const int64_t batch = work / seq_len;
|
||||
|
||||
const int64_t qo_begin = qo_indptr[batch];
|
||||
const int64_t qo_len = qo_indptr[batch + 1] - qo_begin;
|
||||
if (q_pos >= qo_len)
|
||||
return;
|
||||
|
||||
const int64_t kv_begin = kv_indptr[batch];
|
||||
const int64_t kv_len = kv_indptr[batch + 1] - kv_begin;
|
||||
int64_t visible = kv_len;
|
||||
if (causal)
|
||||
{
|
||||
visible = kv_len - qo_len + q_pos + 1;
|
||||
if (visible < 0)
|
||||
visible = 0;
|
||||
if (visible > kv_len)
|
||||
visible = kv_len;
|
||||
}
|
||||
|
||||
const int64_t group = num_qo_heads / num_kv_heads;
|
||||
const int64_t kv_head = qo_head / group;
|
||||
const int64_t q_row = qo_begin + q_pos;
|
||||
const float scale = rsqrtf(static_cast<float>(head_dim_qk));
|
||||
|
||||
const __nv_bfloat16 *q_ptr = q + (q_row * num_qo_heads + qo_head) * head_dim_qk;
|
||||
float qv[4];
|
||||
float acc[4];
|
||||
for (int i = 0; i < 4; ++i)
|
||||
{
|
||||
const int d = lane + i * 32;
|
||||
qv[i] = (d < head_dim_qk) ? __bfloat162float(q_ptr[d]) : 0.0f;
|
||||
acc[i] = 0.0f;
|
||||
}
|
||||
|
||||
float m = -1.0e20f;
|
||||
float l = 0.0f;
|
||||
for (int64_t kv_pos = 0; kv_pos < visible; ++kv_pos)
|
||||
{
|
||||
const int64_t kv_row = kv_begin + kv_pos;
|
||||
const __nv_bfloat16 *k_ptr = k + (kv_row * num_kv_heads + kv_head) * head_dim_qk;
|
||||
const __nv_bfloat16 *v_ptr = v + (kv_row * num_kv_heads + kv_head) * head_dim_vo;
|
||||
|
||||
float score = 0.0f;
|
||||
for (int i = 0; i < 4; ++i)
|
||||
{
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_qk)
|
||||
{
|
||||
score += qv[i] * __bfloat162float(k_ptr[d]);
|
||||
}
|
||||
}
|
||||
score = warp_sum(score) * scale;
|
||||
|
||||
const float m_new = fmaxf(m, score);
|
||||
const float alpha = (m > -1.0e19f) ? __expf(m - m_new) : 0.0f;
|
||||
const float beta = __expf(score - m_new);
|
||||
|
||||
for (int i = 0; i < 4; ++i)
|
||||
{
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_vo)
|
||||
{
|
||||
acc[i] = acc[i] * alpha + beta * __bfloat162float(v_ptr[d]);
|
||||
}
|
||||
}
|
||||
l = l * alpha + beta;
|
||||
m = m_new;
|
||||
}
|
||||
|
||||
__nv_bfloat16 *out_ptr = output + (q_row * num_qo_heads + qo_head) * head_dim_vo;
|
||||
const float inv_l = (l > 0.0f) ? (1.0f / l) : 0.0f;
|
||||
for (int i = 0; i < 4; ++i)
|
||||
{
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_vo)
|
||||
{
|
||||
out_ptr[d] = __float2bfloat16(acc[i] * inv_l);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16 *q,
|
||||
const __nv_bfloat16 *k,
|
||||
const __nv_bfloat16 *v,
|
||||
__nv_bfloat16 *output,
|
||||
const int32_t *qo_indptr,
|
||||
const int32_t *kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal)
|
||||
{
|
||||
constexpr int kThreads = 128;
|
||||
constexpr int kWarpsPerBlock = kThreads / 32;
|
||||
const int64_t total = batch_size * seq_len * num_qo_heads;
|
||||
const int blocks = static_cast<int>((total + kWarpsPerBlock - 1) / kWarpsPerBlock);
|
||||
ragged_prefill_smoke_kernel<<<blocks, kThreads>>>(
|
||||
q, k, v, output, qo_indptr, kv_indptr, batch_size, seq_len,
|
||||
num_qo_heads, num_kv_heads, head_dim_qk, head_dim_vo, causal);
|
||||
}
|
||||
```
|
||||
以上代码仅用于说明接口结构,不代表最优实现,也不作为评分参考。
|
||||
|
||||
4. 点击提交,等待评测结果返回;
|
||||
|
||||

|
||||
|
||||
评测时间与题目测试点数量、队列状态和平台负载有关,通常需要等待数十秒到数分钟。以平台实际返回为准。
|
||||
|
||||

|
||||
|
||||
**OJ 评测流程**
|
||||
1. 选手提交代码;
|
||||
2. 平台按所选语言编译或加载提交代码;
|
||||
3. 评测程序构造测试输入;
|
||||
4. 调用选手代码中的 `run_kernel(...)`;
|
||||
5. 将 `run_kernel(...)` 的输出与 `output_ref` 做正确性校验;
|
||||
6. 正确性通过后,统计运行耗时或性能指标;
|
||||
7. 更新该题历史最好成绩;
|
||||
8. 汇总各题最好成绩,得到排行榜总分。
|
||||
|
||||
5. 查看结果
|
||||
|
||||
**50 分 / 10 分 / 与 baseline 加速比对比 / 榜单**:以 OJ 平台实际结果为准
|
||||
|
||||
## 七、Agent 使用样例
|
||||
|
||||
**目标:** 在本模块中,Agent 可以帮助你完成环境检查、运行 Benchmark、分析结果、理解 OJ 接口、生成和调试 `run_kernel` 等任务。以下是参考 prompt。
|
||||
|
||||
### 环境检查
|
||||
|
||||
```plaintext
|
||||
请帮我检查当前环境是否满足 FlashInfer 运行要求,包括:
|
||||
1. 沐曦 GPU 是否可见(mx-smi)
|
||||
2. Python 版本和 PyTorch CUDA 支持
|
||||
3. flashinfer、pandas、numpy 依赖是否已安装
|
||||
```
|
||||
|
||||
### 运行 Benchmark
|
||||
|
||||
``` plaintext
|
||||
请帮我运行 bench_batch_prefill_ragged.py 脚本,执行 Ragged Prefill 的基准测试。
|
||||
```
|
||||
|
||||
### 分析结果
|
||||
|
||||
``` plaintext
|
||||
请帮我读取最新的 CSV 结果文件,分析各参数配置下的性能表现,找出带宽最高和 TFLOPs 最高的配置,并与沐曦 GPU 理论峰值带宽做对比。
|
||||
```
|
||||
|
||||
### 理解 OJ 题包接口
|
||||
|
||||
``` plaintext
|
||||
请帮我阅读 FlashInfer Ragged Prefill 题包(problem_20001)中的以下文件:
|
||||
- zh_CN/00_题目描述.md
|
||||
- zh_CN/01_接口约定.md
|
||||
- zh_CN/02_数据范围.md
|
||||
然后帮我总结:
|
||||
1. run_kernel 的函数签名和每个参数的含义
|
||||
2. 输入张量的形状约定(q/k/v 的 layout、indptr 的作用)
|
||||
3. head_dim_qk 和 head_dim_vo 的可能取值
|
||||
4. 精度要求(rtol/atol)
|
||||
5. causal=1 时需要注意的边界条件
|
||||
```
|
||||
|
||||
### 生成 `run_kernel` 初版
|
||||
|
||||
``` plaintext
|
||||
请帮我为 FlashInfer Ragged Prefill 题(OJ 题号 20001)生成一个最小可运行的 run_kernel 实现,要求:
|
||||
1. 阅读题包中的 01_接口约定.md,理解 run_kernel 的函数签名(参数类型、顺序、const 修饰);
|
||||
2. 阅读 02_数据范围.md,了解 head_dim_qk、head_dim_vo 的可能取值;
|
||||
3. 生成一个只使用简单双重循环的 naive 实现(不加 tiling、不加 shared memory),确保:
|
||||
- 函数签名为 extern "C" void run_kernel(...)
|
||||
- 包含必要的头文件(cuda_bf16.h、cuda_runtime.h、stdint.h、math.h)
|
||||
- 支持 GQA(Group Query Attention)的头映射:hkv = hq * num_kv_heads / num_qo_heads
|
||||
- 支持 causal mask
|
||||
- 使用 bfloat16 数据类型
|
||||
- scale = 1.0f / sqrtf(head_dim_qk)
|
||||
```
|
||||
|
||||
### 调试 OJ 提交错误
|
||||
``` plaintext
|
||||
我的 run_kernel 提交到 OJ 后显示 Wrong Answer,请帮我对比以下信息:
|
||||
1. 这是我的 submit 代码:[粘贴你的 run_kernel 实现]
|
||||
2. 题包的接口约定在这里:[粘贴或引用 01_接口约定.md]
|
||||
3. 题包的测试配置在这里:[粘贴或引用 testcase_config.py]
|
||||
请帮我逐项检查:
|
||||
- 函数签名是否完全匹配
|
||||
- GQA 头映射公式是否正确
|
||||
- causal mask 边界条件是否正确
|
||||
- float4 向量化加载的偏移是否正确
|
||||
- online softmax 的 m/l 更新逻辑是否正确
|
||||
```
|
||||
|
||||
### 问题排查
|
||||
|
||||
``` plaintext
|
||||
运行 bench_batch_prefill_ragged.py 时报错 out of memory,请帮我分析原因并给出解决方案。
|
||||
代码理解
|
||||
请帮我解释 bench_batch_prefill_ragged.py 中 BatchPrefillWithRaggedKVCacheWrapper 的 plan() 和 run() 方法的工作原理,特别是 qo_indptr 和 kv_indptr 的作用。
|
||||
整理优化日志
|
||||
请帮我整理本次优化的记录,包括:
|
||||
1. 原始 baseline 性能数据(从 CSV 中提取关键配置的 time_ms 和 tflops)
|
||||
2. 优化后的性能数据(从 OJ 评测结果中提取)
|
||||
3. 加速比 = baseline_time / optimized_time
|
||||
4. 以表格形式输出:配置参数 | baseline 耗时 | 优化后耗时 | 加速比
|
||||
```
|
||||
|
||||
## 八、常见问题
|
||||
|
||||
### 流程问题
|
||||
|
||||
| 问题 | 原因 | 解决办法 |
|
||||
| --- | --- | --- |
|
||||
| 登录后看不到题目 | 未使用赛用账号登录 | 七月份组委会统一发放 XPU-OJ 账号,请确认你使用的是组委会统一发放的账号,而不是自行注册账号;如仍无法看到题目,请联系助教或赛事运营确认账号权限。 |
|
||||
|
||||
### 环境问题
|
||||
|
||||
| 问题 | 原因 | 解决办法 |
|
||||
| --- | --- | --- |
|
||||
| `No GPUs are available` | MXMACA 驱动未安装或 GPU 不可见 | 检查驱动安装,运行 `python -c "import torch; print(torch.cuda.device_count())"` 验证 |
|
||||
| `ModuleNotFoundError: No module named 'flashinfer'` | flashinfer 未安装 | 执行 `pip install flashinfer` |
|
||||
| `out of memory` | GPU 显存不足 | 减小 `batch_size` 或 `seq_len` 参数 |
|
||||
|
||||
### 运行问题
|
||||
|
||||
| 问题 | 原因 | 解决办法 |
|
||||
| --- | --- | --- |
|
||||
| Benchmark 运行时间过长 | 参数组合过多, workload 较大 | 耐心等待,脚本会自动调整重复次数 |
|
||||
| `KeyError: 'BatchPrefillWithPagedKVCacheKernel'` | profiler 未捕获目标 kernel | 检查 `target_kernels` 配置是否正确 |
|
||||
| CSV 文件为空 | 测试未正常完成 | 检查 GPU 显存是否充足,重新运行 |
|
||||
|
||||
### 代码问题
|
||||
|
||||
| 问题 | 原因 | 解决办法 |
|
||||
| --- | --- | --- |
|
||||
| `ImportError: cannot import name 'xxx' from 'bench_common'` | 函数名拼写错误 | 检查 `bench_common.py` 中的函数名 |
|
||||
| `RuntimeError: error: device-side assert triggered` | 输入参数超出范围 | 检查 `num_qo_heads`、`num_kv_heads`、`head_dim` 配置 |
|
||||
|
||||
### 性能问题
|
||||
|
||||
| 问题 | 原因 | 解决办法 |
|
||||
| --- | --- | --- |
|
||||
| TFLOPs 数值异常低 | 工作负载过小,kernel 启动开销占比大 | 增大 `batch_size` 或 `seq_len` |
|
||||
| 带宽数值异常低 | 数据未正确加载到 GPU | 检查 Tensor 是否在 CUDA 设备上 |
|
||||
|
||||
### 评测问题
|
||||
|
||||
提交 XPU-OJ 后可能遇到的异常评测结果及排查方向:
|
||||
|
||||
| 问题 | 可能原因 | 解决办法 |
|
||||
|------|------|---------|
|
||||
| **Compilation Error** 编译错误 | 1. `run_kernel` 签名与 OJ 接口约定不一致(参数类型、顺序、数量不匹配);2. 缺少 `extern "C"` 声明导致 C++ name mangling;3. 缺少必要头文件(`cuda_bf16.h`、`cuda_runtime.h`、`math.h`);3. 使用了 OJ 环境不支持的语法或 API | 1. 逐行对照对应题包的「01_接口约定.md」,确认参数类型(`int64_t` vs `int`、`const` 修饰)、顺序完全一致;2. 在 `run_kernel` 前加 `extern "C"`;3. 确认文件顶部 include 了 `<cuda_bf16.h>`、`<cuda_runtime.h>`、`<stdint.h>`、`<math.h>`;4. 去掉 `printf`、`assert` 等调试代码后重新提交 |
|
||||
| **Time Limit Exceeded** 运行超时 | 1. `run_kernel` 内部调用了 `cudaDeviceSynchronize()` 导致额外等待;2. kernel 中存在死循环(for 循环边界条件错误);3. `__syncthreads()` 放在条件分支内导致线程死锁;4. grid 配置过大,启动的 block 数量远超合理范围 | 1. 删除 `run_kernel` 函数体内的 `cudaDeviceSynchronize()` 调用——评测器会在外部自行同步;2. 检查 kernel 中所有 for 循环的终止条件,确保 `kv_start <= block_max_q` 等边界正确;3. 将所有 `__syncthreads()` 移到 if/else 分支之外;4. 检查 grid 计算:`(seq_len + Br - 1) / Br`,确认 `Br` 取值合理 |
|
||||
| **Wrong Answer** 答案错误 | 1. 注意力计算公式错误(score、scale、softmax 实现有偏差);2. GQA 头映射错误:`hkv = hq / (num_qo_heads / num_kv_heads)` 计算不对;3. Causal mask 未正确实现(`causal=1` 时 query 看到了不该看的未来 token);4. Online softmax 的 m/l 更新逻辑有误;5. float4 向量化加载的偏移计算错误,导致 K/V 数据错位;6. 输出写入偏移错误,或对无效位置写了垃圾值 | 1. 本地用题包中的 PyTorch 参考实现对拍:运行 `testcase_config.py` 的 `baseline()` 与你 kernel 输出做 `torch.allclose(rtol=1e-2, atol=1e-2)` 比对;2. GQA 公式:`int hkv = hq * num_kv_heads / num_qo_heads`(整数除法);3. Causal 逻辑:`kv_end = min(kv_start + Bc, q_idx + 1)`,注意 +1 的处理;4. 对照论文 FlashAttention 的 Algorithm 1 逐行验证 online softmax;5. float4 加载偏移公式:`(cur_kv_start + i) * num_kv_heads * D_QK + hkv * D_QK + d_idx * 8`,确认 `num_kv_heads` 而非 `num_qo_heads` |
|
||||
|
||||
|
||||
## 九、下一步学习建议
|
||||
|
||||
### 1. 保存你的 Baseline 结果
|
||||
|
||||
将本次运行生成的 CSV 文件妥善保存,后续优化时需要以此作为对比基准。
|
||||
|
||||
```bash
|
||||
# 建议创建 results 目录保存
|
||||
mkdir -p results
|
||||
mv *.csv results/
|
||||
```
|
||||
|
||||
### 2. 深入理解 FlashInfer 核心概念
|
||||
|
||||
* 阅读 FlashInfer 官方文档,理解 Paged KV Cache、Ragged KV Cache 的设计理念
|
||||
|
||||
* 学习 MLA (Multi-head Latent Attention) 的原理,了解 DeepSeek 的注意力优化方案
|
||||
|
||||
* 理解 `plan()` 和 `run()` 两阶段设计的作用
|
||||
|
||||
|
||||
### 3. 进入算子优化模块
|
||||
|
||||
参考后续优化模块,学习以下优化技术:
|
||||
|
||||
* **Kernel Tuning**:调整 Block Size、Thread Count 等参数
|
||||
|
||||
* **Memory Optimization**:减少显存占用、优化数据搬运
|
||||
|
||||
* **Compute Optimization**:提升计算效率
|
||||
|
||||
|
||||
### 4. 参考资源
|
||||
|
||||
* FlashInfer 官方仓库:https://github.com/flashinfer-ai/flashinfer
|
||||
|
||||
* FlashInfer 文档:https://flashinfer.ai
|
||||
|
||||
|
||||
### 5. 记录优化流程
|
||||
|
||||
建议维护一份优化日志,记录每次优化的改动和性能变化:
|
||||
|
||||
| 优化项 | 改动内容 | Baseline | 优化后 | 提升比例 |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| 例:调整 block\_size | 16 → 32 | xx ms | xx ms | xx% |
|
||||
|
||||
完成优化后,再次运行本模块的 Benchmark 脚本,对比前后性能变化。
|
||||
|
||||
> 使用 Agent 整理优化日志,可形成可复现的 Agent/Skill 优化流程
|
||||
|
||||
### 6. 使用多语言完成算子优化加速
|
||||
|
||||
可以使用 Triton 或 TileLang 语言实现 `run_kernel` 接口完成算子优化,对比不同的语言对于性能加速的影响。
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
|
|
@ -0,0 +1,49 @@
|
|||
batch_size,seq_len_kv,heads,headdim,time_ms,bandwidth_GB_s
|
||||
1,512,8,128,0.0322,65.27
|
||||
2,512,8,128,0.0324,129.61
|
||||
4,512,8,128,0.0332,253.01
|
||||
8,512,8,128,0.0355,472.66
|
||||
16,512,8,128,0.0546,614.98
|
||||
32,512,8,128,0.0817,822.60
|
||||
64,512,8,128,0.1297,1035.62
|
||||
128,512,8,128,0.2453,1095.50
|
||||
1,1024,8,128,0.0578,72.55
|
||||
2,1024,8,128,0.0586,143.14
|
||||
4,1024,8,128,0.0597,281.24
|
||||
8,1024,8,128,0.0625,536.80
|
||||
16,1024,8,128,0.0982,683.59
|
||||
32,1024,8,128,0.1493,899.64
|
||||
64,1024,8,128,0.2403,1117.62
|
||||
128,1024,8,128,0.4594,1169.27
|
||||
1,2048,8,128,0.1101,76.24
|
||||
2,2048,8,128,0.1107,151.64
|
||||
4,2048,8,128,0.1119,299.88
|
||||
8,2048,8,128,0.1159,578.98
|
||||
16,2048,8,128,0.1849,726.23
|
||||
32,2048,8,128,0.2843,944.47
|
||||
64,2048,8,128,0.4607,1165.56
|
||||
128,2048,8,128,0.8868,1211.07
|
||||
1,4096,8,128,0.2139,78.46
|
||||
2,4096,8,128,0.2151,156.01
|
||||
4,4096,8,128,0.2163,310.36
|
||||
8,4096,8,128,0.2227,602.81
|
||||
16,4096,8,128,0.3574,751.13
|
||||
32,4096,8,128,0.5540,969.23
|
||||
64,4096,8,128,0.9016,1191.07
|
||||
128,4096,8,128,1.7414,1233.34
|
||||
1,8192,8,128,0.4215,79.61
|
||||
2,8192,8,128,0.4226,158.81
|
||||
4,8192,8,128,0.4242,316.39
|
||||
8,8192,8,128,0.4362,615.46
|
||||
16,8192,8,128,0.7035,763.14
|
||||
32,8192,8,128,1.0934,982.11
|
||||
64,8192,8,128,1.7814,1205.57
|
||||
128,8192,8,128,3.4505,1244.82
|
||||
1,16384,8,128,0.8356,80.32
|
||||
2,16384,8,128,0.8377,160.23
|
||||
4,16384,8,128,0.8407,319.30
|
||||
8,16384,8,128,0.8625,622.51
|
||||
16,16384,8,128,1.3934,770.60
|
||||
32,16384,8,128,2.1695,989.88
|
||||
64,16384,8,128,3.5397,1213.41
|
||||
128,16384,8,128,6.8668,1250.98
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -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
|
||||
|
|
|
@ -0,0 +1,330 @@
|
|||
# guide
|
||||
|
||||
# FlashAttention KV Cache Decode - 参赛指南
|
||||
|
||||
## 一、登录 XPU-OJ 平台
|
||||
|
||||
打开浏览器,访问:\*\*https://xpuoj.com/\*\*
|
||||
|
||||
1. 等待组委会统一发放 XPU-OJ 账号;
|
||||
|
||||
2. 使用分配的用户名和初始密码登录平台
|
||||
|
||||
|
||||

|
||||
|
||||
点击 **"登录"** 进入平台。
|
||||
|
||||
## 二、进入比赛
|
||||
|
||||
### 2.1 点击导航栏"比赛"
|
||||
|
||||
登录成功后,会跳转到 XPUOJ 平台首页。
|
||||
|
||||
找到页面顶部的导航栏,点击 **"比赛"** 标签,进入比赛列表页。
|
||||
|
||||

|
||||
|
||||
### 2.2 找到目标比赛
|
||||
|
||||
在比赛列表中,根据比赛状态(全部 / 未开始 / 进行中 / 已结束)找到目标比赛。
|
||||
|
||||

|
||||
|
||||
## 三、进入题目
|
||||
|
||||
### 3.1 在题目列表中找到目标题目
|
||||
|
||||
进入比赛后,页面会展示该比赛的**题目列表**。每道题都有编号(1 ~ N)和标题。
|
||||
|
||||
在列表中找到 **FlashAttention KV Cache Decode** 这道题(前缀为 `FlashAttention`),点击标题即可进入题目详情页。
|
||||
|
||||

|
||||
|
||||
### 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. 点击右上角 **"提交"** 按钮
|
||||
|
||||
|
||||

|
||||
|
||||
## 四、查看提交结果
|
||||
|
||||
点击 **"提交"** 后,系统会自动跳转到提交记录页,展示本次提交的详细信息。
|
||||
|
||||
页面顶部会显示一行汇总信息,包括:
|
||||
|
||||
* **状态** —— 评测结果(如 `Accepted` 表示通过)
|
||||
|
||||
* **分数** —— 本次提交获得的分数
|
||||
|
||||
* **题目** —— 对应的题目名称
|
||||
|
||||
* **用时** —— 程序运行耗时
|
||||
|
||||
* **内存** —— 占用内存大小
|
||||
|
||||
* **答案** —— 提交所用的语言/硬件
|
||||
|
||||
* **提交时间** —— 提交的时刻
|
||||
|
||||
|
||||

|
||||
|
||||
页面下方会展开 **编译信息** 和 **各测试点**(样例、测试点 #1 ~ #N)的结果,逐个显示:
|
||||
|
||||

|
||||
|
|
@ -1,119 +0,0 @@
|
|||
#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,187 +0,0 @@
|
|||
"""FlashAttention KV Cache Decode in TileLang."""
|
||||
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
from tilelang import jit
|
||||
|
||||
NUM_SPLITS = 4
|
||||
real_kernel = None
|
||||
|
||||
@jit(
|
||||
pass_configs={
|
||||
tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: False,
|
||||
},
|
||||
)
|
||||
def build_kernel(
|
||||
batch_size,
|
||||
num_heads,
|
||||
num_heads_k,
|
||||
headdim,
|
||||
page_block_size,
|
||||
num_blocks,
|
||||
causal,
|
||||
):
|
||||
blocks_per_batch = num_blocks // batch_size
|
||||
assert blocks_per_batch % NUM_SPLITS == 0, (
|
||||
f"blocks_per_batch={blocks_per_batch} must be divisible by NUM_SPLITS={NUM_SPLITS}"
|
||||
)
|
||||
blocks_per_split = blocks_per_batch // NUM_SPLITS
|
||||
|
||||
BLOCK_M = 1
|
||||
BLOCK_N = page_block_size
|
||||
scale = (1.0 / headdim) ** 0.5 * 1.44269504 # log2(e)
|
||||
dtype = "bfloat16"
|
||||
accum_dtype = "float32"
|
||||
|
||||
# Use a large-negative-finite sentinel instead of -inf to avoid
|
||||
# (-inf) - (-inf) = NaN when an entire split is masked out.
|
||||
NEG_INF_SAFE = -1e30
|
||||
|
||||
@T.prim_func
|
||||
def kernel(
|
||||
Q: T.Tensor([batch_size, 1, num_heads, headdim], dtype),
|
||||
K: T.Tensor([num_blocks, page_block_size, num_heads_k, headdim], dtype),
|
||||
V: T.Tensor([num_blocks, page_block_size, num_heads_k, headdim], dtype),
|
||||
Output: T.Tensor([batch_size, 1, num_heads, headdim], dtype),
|
||||
cache_seqlens: T.Tensor([batch_size], "int32"),
|
||||
block_table: T.Tensor([batch_size, blocks_per_batch], "int32"),
|
||||
):
|
||||
# float32 workspace — avoids BF16StorageLegalize var-remap bug
|
||||
glse = T.alloc_global([batch_size, num_heads, NUM_SPLITS], accum_dtype)
|
||||
Output_partial = T.alloc_global(
|
||||
[batch_size, 1, num_heads, NUM_SPLITS, headdim], accum_dtype
|
||||
)
|
||||
|
||||
# ============= Stage 1: split kernel =============
|
||||
with T.Kernel(NUM_SPLITS, num_heads, batch_size, threads=128) as (bs, bh, bz):
|
||||
Q_shared = T.alloc_shared([BLOCK_M, headdim], dtype)
|
||||
K_shared = T.alloc_shared([BLOCK_N, headdim], dtype)
|
||||
V_shared = T.alloc_shared([BLOCK_N, headdim], dtype)
|
||||
acc_s = T.alloc_fragment([BLOCK_M, BLOCK_N], accum_dtype)
|
||||
acc_o = T.alloc_fragment([BLOCK_M, headdim], accum_dtype)
|
||||
scores_max = T.alloc_fragment([BLOCK_M], accum_dtype)
|
||||
scores_max_prev = T.alloc_fragment([BLOCK_M], accum_dtype)
|
||||
scores_scale = T.alloc_fragment([BLOCK_M], accum_dtype)
|
||||
scores_sum = T.alloc_fragment([BLOCK_M], accum_dtype)
|
||||
logsum = T.alloc_fragment([BLOCK_M], accum_dtype)
|
||||
|
||||
T.copy(Q[bz, 0, bh, :], Q_shared)
|
||||
|
||||
kv_seqlen = cache_seqlens[bz]
|
||||
split_k_start = bs * blocks_per_split
|
||||
|
||||
T.fill(acc_o, 0)
|
||||
T.fill(logsum, 0)
|
||||
# KEY FIX: use -1e30 instead of -inf to avoid (-inf)-(-inf)=NaN
|
||||
T.fill(scores_max, NEG_INF_SAFE)
|
||||
|
||||
for k in T.Pipelined(blocks_per_split, num_stages=2):
|
||||
global_k = split_k_start + k
|
||||
physical_block = block_table[bz, global_k]
|
||||
tok_offset = global_k * page_block_size
|
||||
|
||||
# ----- Q @ K^T (hand-written, M=1, masked) -----
|
||||
T.copy(K[physical_block, 0:BLOCK_N, bh, :], K_shared)
|
||||
T.fill(acc_s, 0)
|
||||
for j in T.Parallel(BLOCK_N):
|
||||
if tok_offset + j < kv_seqlen:
|
||||
for d in T.serial(headdim):
|
||||
acc_s[0, j] = acc_s[0, j] + Q_shared[0, d] * K_shared[j, d]
|
||||
else:
|
||||
acc_s[0, j] = -T.infinity(accum_dtype)
|
||||
|
||||
# ----- online softmax -----
|
||||
T.copy(scores_max, scores_max_prev)
|
||||
# KEY FIX: use -1e30 instead of -inf here too
|
||||
T.fill(scores_max, NEG_INF_SAFE)
|
||||
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
|
||||
scores_max[0] = T.max(scores_max[0], scores_max_prev[0])
|
||||
# (prev - cur) is now (finite - finite) = 0 when both are sentinel,
|
||||
# never (-inf - (-inf)) = NaN
|
||||
scores_scale[0] = T.exp2((scores_max_prev[0] - scores_max[0]) * scale)
|
||||
for j in T.Parallel(BLOCK_N):
|
||||
acc_s[0, j] = T.exp2((acc_s[0, j] - scores_max[0]) * scale)
|
||||
T.reduce_sum(acc_s, scores_sum, dim=1)
|
||||
logsum[0] = logsum[0] * scores_scale[0] + scores_sum[0]
|
||||
for d in T.Parallel(headdim):
|
||||
acc_o[0, d] = acc_o[0, d] * scores_scale[0]
|
||||
|
||||
# ----- P @ V (hand-written, fp32 accum) -----
|
||||
T.copy(V[physical_block, 0:BLOCK_N, bh, :], V_shared)
|
||||
for d in T.Parallel(headdim):
|
||||
for j in T.serial(BLOCK_N):
|
||||
acc_o[0, d] = acc_o[0, d] + acc_s[0, j] * V_shared[j, d]
|
||||
|
||||
# ----- final normalise & write partial state -----
|
||||
# KEY FIX: add epsilon to avoid 0/0 = NaN when split is all-masked
|
||||
safe_logsum = logsum[0] + 1e-30
|
||||
for d in T.Parallel(headdim):
|
||||
acc_o[0, d] = acc_o[0, d] / safe_logsum
|
||||
|
||||
lse_local = T.alloc_fragment([1], accum_dtype)
|
||||
lse_local[0] = T.log2(safe_logsum) + scores_max[0] * scale
|
||||
glse[bz, bh, bs] = lse_local[0]
|
||||
|
||||
for d in T.Parallel(headdim):
|
||||
Output_partial[bz, 0, bh, bs, d] = acc_o[0, d]
|
||||
|
||||
# ============= Stage 2: combine kernel =============
|
||||
with T.Kernel(num_heads, batch_size, threads=128) as (bh, bz):
|
||||
lse_local = T.alloc_fragment([NUM_SPLITS], accum_dtype)
|
||||
for s in T.serial(NUM_SPLITS):
|
||||
lse_local[s] = glse[bz, bh, s]
|
||||
|
||||
lse_max = T.alloc_fragment([1], accum_dtype)
|
||||
lse_max[0] = -T.infinity(accum_dtype)
|
||||
for s in T.serial(NUM_SPLITS):
|
||||
lse_max[0] = T.max(lse_max[0], lse_local[s])
|
||||
|
||||
lse_logsum = T.alloc_fragment([1], accum_dtype)
|
||||
lse_logsum[0] = 0
|
||||
for s in T.serial(NUM_SPLITS):
|
||||
lse_logsum[0] = lse_logsum[0] + T.exp2(lse_local[s] - lse_max[0])
|
||||
lse_logsum[0] = T.log2(lse_logsum[0]) + lse_max[0]
|
||||
|
||||
o_accum = T.alloc_fragment([headdim], accum_dtype)
|
||||
T.fill(o_accum, 0)
|
||||
for s in T.serial(NUM_SPLITS):
|
||||
s_scale = T.exp2(lse_local[s] - lse_logsum[0])
|
||||
for d in T.Parallel(headdim):
|
||||
o_accum[d] = o_accum[d] + Output_partial[bz, 0, bh, s, d] * s_scale
|
||||
|
||||
for d in T.Parallel(headdim):
|
||||
Output[bz, 0, bh, d] = T.Cast(dtype, o_accum[d])
|
||||
|
||||
return kernel
|
||||
|
||||
|
||||
def run_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,
|
||||
):
|
||||
global real_kernel
|
||||
|
||||
B = int(batch_size)
|
||||
H = int(num_heads)
|
||||
HK = int(num_heads_k)
|
||||
D = int(headdim)
|
||||
PBS = int(page_block_size)
|
||||
NB = int(num_blocks)
|
||||
|
||||
if real_kernel is None:
|
||||
real_kernel = build_kernel(B, H, HK, D, PBS, NB, int(causal))
|
||||
|
||||
real_kernel(q, k_cache_paged, v_cache_paged, output, cache_seqlens, block_table)
|
||||
|
|
@ -1,106 +0,0 @@
|
|||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def slow_decode_kernel(
|
||||
q_ptr,
|
||||
k_cache_ptr,
|
||||
v_cache_ptr,
|
||||
output_ptr,
|
||||
cache_seqlens_ptr,
|
||||
block_table_ptr,
|
||||
num_heads: tl.constexpr,
|
||||
num_heads_k: tl.constexpr,
|
||||
headdim: tl.constexpr,
|
||||
page_block_size: tl.constexpr,
|
||||
blocks_per_batch,
|
||||
):
|
||||
# 维度索引
|
||||
pid_b = tl.program_id(0) # Batch index
|
||||
pid_h = tl.program_id(1) # Head index
|
||||
|
||||
# GQA Support: 映射 Query Head 到 KV Head
|
||||
kv_head = pid_h * num_heads_k // num_heads
|
||||
|
||||
# 加载实际的 KV 序列长度
|
||||
seq_len = tl.load(cache_seqlens_ptr + pid_b).to(tl.int32)
|
||||
|
||||
# 维度偏移量 [0, 1, ..., headdim-1]
|
||||
offs_d = tl.arange(0, headdim)
|
||||
|
||||
# Online Softmax 累加器
|
||||
acc = tl.zeros([headdim], dtype=tl.float32)
|
||||
l_i = 0.0
|
||||
m_i = float('-inf')
|
||||
scale = 1.0 / tl.sqrt(float(headdim))
|
||||
|
||||
# === 性能瓶颈:串行遍历整个序列 ===
|
||||
# 不使用 Block 并行,而是用单个 Block 串行循环处理所有 Token
|
||||
t = 0
|
||||
while t < seq_len:
|
||||
# 性能瓶颈:每次循环都重新加载 Q,增加显存压力
|
||||
q = tl.load(q_ptr + pid_b * num_heads * headdim + pid_h * headdim + offs_d).to(tl.float32)
|
||||
q = q * scale
|
||||
|
||||
# Paged KV 映射逻辑
|
||||
page_idx = t // page_block_size
|
||||
page_off = t % page_block_size
|
||||
|
||||
# 查表获取物理 Block 索引
|
||||
# blocks_per_batch 是计算出来的步长
|
||||
phys_block = tl.load(block_table_ptr + pid_b * blocks_per_batch + page_idx)
|
||||
|
||||
# 计算 K 和 V 的物理地址
|
||||
# Layout: (num_blocks, page_block_size, num_heads_k, headdim)
|
||||
kv_base = phys_block * page_block_size * num_heads_k * headdim + \
|
||||
page_off * num_heads_k * headdim + \
|
||||
kv_head * headdim
|
||||
|
||||
# 加载 K 和 V 向量
|
||||
k = tl.load(k_cache_ptr + kv_base + offs_d).to(tl.float32)
|
||||
v = tl.load(v_cache_ptr + kv_base + offs_d).to(tl.float32)
|
||||
|
||||
# Attention 计算
|
||||
s = tl.sum(q * k) # 点积
|
||||
|
||||
# Online Softmax 更新
|
||||
m_new = tl.maximum(m_i, s)
|
||||
p = tl.exp(s - m_new)
|
||||
alpha = tl.exp(m_i - m_new)
|
||||
|
||||
acc = acc * alpha + p * v
|
||||
l_i = l_i * alpha + p
|
||||
m_i = m_new
|
||||
|
||||
t += 1
|
||||
|
||||
# 写回结果
|
||||
# 这里没有处理 l_i 为 0 的边界情况,但测试数据 seq_len 通常很大
|
||||
out = acc / l_i
|
||||
tl.store(output_ptr + pid_b * num_heads * headdim + pid_h * headdim + offs_d, out)
|
||||
|
||||
def run_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,
|
||||
):
|
||||
# 计算每个 batch 对应的 block_table 行宽
|
||||
blocks_per_batch = num_blocks // batch_size
|
||||
|
||||
# 启动配置:每个 Head 一个 Block
|
||||
# 总 Block 数 = batch_size * num_heads (最大 128个),并行度极低
|
||||
grid = (batch_size, num_heads)
|
||||
|
||||
slow_decode_kernel[grid](
|
||||
q, k_cache_paged, v_cache_paged, output,
|
||||
cache_seqlens, block_table,
|
||||
num_heads=num_heads,
|
||||
num_heads_k=num_heads_k,
|
||||
headdim=headdim,
|
||||
page_block_size=page_block_size,
|
||||
blocks_per_batch=blocks_per_batch,
|
||||
num_warps=1, # 性能瓶颈:仅使用 1 个 warp,限制计算吞吐
|
||||
num_stages=1, # 性能瓶颈:禁用流水线并行
|
||||
)
|
||||
|
|
@ -0,0 +1,16 @@
|
|||
{
|
||||
"id": 197,
|
||||
"displayId": 20005,
|
||||
"type": "Traditional",
|
||||
"isPublic": false,
|
||||
"locales": [
|
||||
"zh_CN"
|
||||
],
|
||||
"samples": [
|
||||
{
|
||||
"inputData": "1\n",
|
||||
"outputData": ""
|
||||
}
|
||||
],
|
||||
"problemTagIds": []
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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 token,KV 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)。
|
||||
|
|
@ -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 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()` 或显式同步。
|
||||
|
|
@ -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。
|
||||
|
|
@ -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。
|
||||
|
|
@ -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)`。
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
---
|
||||
sectionTitle: "输出格式"
|
||||
type: "Text"
|
||||
---
|
||||
输出写入 `output`,shape 为 `(batch_size, 1, num_heads, headdim)`,类型为 `bfloat16`。
|
||||
|
|
@ -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`。
|
||||
|
|
@ -0,0 +1 @@
|
|||
FlashAttention KV Cache Decode
|
||||
|
|
@ -0,0 +1,145 @@
|
|||
api,batch_size,seq_len_q,seq_len_kv,num_qo_heads,num_kv_heads,head_dim,time_ms,bandwidth_GB_s,tflops
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,8,64,0.02042879999999998,51.528822055137894,0.8212531328320811
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,4,128,0.02333952000000001,45.27805199078642,0.718832949435121
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,512,32,4,256,0.0319488,66.15384615384615,1.0502564102564103
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,8,64,0.023262719999999973,90.32684054143292,1.4424122372620245
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,4,128,0.025041919999999992,84.07278675117566,1.3399304845634845
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,1024,32,4,256,0.033387520000000004,126.11562643766291,2.0099984664928687
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,8,64,0.028298240000000037,148.36258368011562,2.371485435136599
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,4,128,0.027745280000000008,151.4670603432367,2.418748846650673
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,2048,32,4,256,0.03723775999999999,225.7115358174069,3.604344837068611
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,8,64,0.03886591999999997,215.93992886312756,3.4533526544592306
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,4,128,0.03426815999999998,245.03212311370103,3.916689078141344
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,4096,32,4,256,0.066048,254.26356589147287,4.064248062015504
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,8,64,0.052495359999999984,319.6722910367698,5.1135082414902975
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,4,128,0.04628480000000001,362.6548672566371,5.799646017699114
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,8192,32,4,256,0.08975359999999999,374.0330861380491,5.981608670849972
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,8,64,0.08625152000000001,389.0775258221536,6.224480588863825
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,4,128,0.0638464,525.6776263031276,8.408789093825181
|
||||
BatchDecodeWithPagedKVCacheWrapper,1,1,16384,32,4,256,0.13059071999999994,514.0123892417473,8.222190857053247
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,8,64,0.02342912,89.86013986013987,1.4321678321678322
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,4,128,0.02486784,84.99073502161829,1.3493102738315832
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,512,32,4,256,0.03340287999999998,126.54812998160644,2.009074187614961
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,8,64,0.02839040000000001,148.02524797114512,2.3637871956717755
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,4,128,0.028165120000000012,149.5000908925649,2.382694055626249
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,1024,32,4,256,0.03740160000000001,225.16084873374396,3.5885557837097872
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,8,64,0.03881984000000001,216.30176734370872,3.457451859667633
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,4,128,0.03601408000000001,233.38072220642587,3.7268126243957904
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,2048,32,4,256,0.06728704000000002,249.82498858621207,3.9894080048698815
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,8,64,0.052490240000000014,319.7815060476004,5.114007023019897
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,4,128,0.04626431999999999,362.9924745462595,5.802213368747235
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,4096,32,4,256,0.08993791999999999,373.44870773084375,5.969349880450872
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,8,64,0.08536063999999999,393.18618042226495,6.289443378119003
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,4,128,0.0630784,532.2077922077922,8.51116883116883
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,8192,32,4,256,0.12952576,518.3650881492608,8.289793659577834
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,8,64,0.15207424000000003,441.34401723789637,7.0606423809844445
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,4,128,0.10330112,649.8017446471055,10.394290245836638
|
||||
BatchDecodeWithPagedKVCacheWrapper,2,1,16384,32,4,256,0.2281984,588.3060354498541,9.410599057662106
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,8,64,0.0283904,148.3137962128043,2.3637871956717764
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,4,128,0.028078080000000036,150.5470459518598,2.3900802334062696
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,512,32,4,256,0.03707903999999999,228.00331400165706,3.619773543220106
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,8,64,0.03844096000000004,218.64677677144357,3.4915290356952546
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,4,128,0.03641856000000004,231.23857725291697,3.6854210600309254
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,1024,32,4,256,0.06640640000000002,253.63145720894363,4.04231303006939
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,8,64,0.059007999999999984,284.5986984815619,4.5491366594360105
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,4,128,0.04641792000000003,362.1442753143611,5.783013456871825
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,2048,32,4,256,0.08961023999999998,375.1799794309223,5.991178151068451
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,8,64,0.09185279999999997,365.484949832776,5.84490523968785
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,4,128,0.06349823999999998,528.9469440412838,8.454894371875506
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,4096,32,4,256,0.1303347200000001,515.3991200502825,8.238340666247638
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,8,64,0.16568319999999992,405.1421508034613,6.480692212608161
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,4,128,0.10290176,652.4828341128471,10.43463031147378
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,8192,32,4,256,0.22947840000000008,585.1673360107093,9.358107987505575
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,8,64,0.30601215999999987,438.65613706331163,7.017641547316293
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,4,128,0.18384895999999992,730.2216776205863,11.680695109724857
|
||||
BatchDecodeWithPagedKVCacheWrapper,4,1,16384,32,4,256,0.4362026666666668,615.5418398787107,9.846265564630505
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,8,64,0.038655999999999975,217.85430463576174,3.472105960264903
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,4,128,0.03645951999999999,231.87754528858312,3.6812807190001418
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,512,32,4,256,0.06676480000000001,253.25153374233125,4.020613496932515
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,8,64,0.05858303999999996,286.9428421604617,4.582135990211505
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,4,128,0.04676608000000001,360.14889424129615,5.73996058681848
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,1024,32,4,256,0.08992768000000002,374.58437713504884,5.970029606012297
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,8,64,0.092416,363.43490304709144,5.8092853185595565
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,4,128,0.07130112000000002,471.5208961654457,7.5296280338934345
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,2048,32,4,256,0.14862335999999993,452.41835469202175,7.224583161085851
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,8,64,0.16396288000000003,409.4928803397451,6.548688483637271
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,4,128,0.11201536000000002,599.6891854831337,9.585665965810401
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,4096,32,4,256,0.24935424000000006,538.7869081351894,8.61218019793848
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,8,64,0.3056947200000001,439.16524302415155,7.024928817874248
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,4,128,0.20128768000000002,667.1211273337741,10.668728697156228
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,8192,32,4,256,0.46690133333333317,575.2104541716168,9.198875628255509
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,8,64,0.5866495999999998,457.6296037702917,7.321179961598886
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,4,128,0.37337600000000015,719.1169009256082,11.503062050051419
|
||||
BatchDecodeWithPagedKVCacheWrapper,8,1,16384,32,4,256,0.8934826666666666,601.0211546726517,9.613991308915525
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,8,64,0.0698112,241.26145947928126,3.845163182984965
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,4,128,0.04724735999999999,357.8673602080625,5.681491114000869
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,512,32,4,256,0.08954879999999998,377.6329331046313,5.995288736420813
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,8,64,0.12070911999999998,278.52052935188334,4.447641669494402
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,4,128,0.07076864000000004,475.9947909130369,7.586282737664589
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,1024,32,4,256,0.14710784000000002,457.9702074342197,7.2990115550605585
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,8,64,0.22239232000000014,302.05359609540454,4.8281425545630325
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,4,128,0.11209728000000002,599.8355713894217,9.578660820316067
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,2048,32,4,256,0.2504192,537.0190145164587,8.575555101206296
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,8,64,0.42098688000000006,318.97256275539985,5.101070247129791
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,4,128,0.2027008,662.7936347562515,10.594352109118466
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,4096,32,4,256,0.46432,578.6905582356995,9.250015713301172
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,8,64,0.8234496000000004,326.06851955480926,5.215822918609709
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,4,128,0.3726506666666667,720.6924662239522,11.525451797572705
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,8192,32,4,256,0.8939733333333334,600.8378952392316,9.608714568667223
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,8,64,1.6324906666666663,328.9062896122735,5.261858317107276
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,4,128,0.7114879999999999,754.7590177206082,12.073196725735361
|
||||
BatchDecodeWithPagedKVCacheWrapper,16,1,16384,32,4,256,1.742272,616.4387466480549,9.860612570253094
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,8,64,0.08406016,400.730905104154,6.386746254111341
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,4,128,0.08498175999999999,397.92746113989637,6.317484034220991
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,512,32,4,256,0.1808896,373.8918765921313,5.935895839230116
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,8,64,0.14712832,457.0155902004454,7.297995545657015
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,4,128,0.14887935999999996,452.5208061077104,7.212160396175805
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,1024,32,4,256,0.3279462400000001,410.8661712358707,6.548279522887651
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,8,64,0.27223039999999993,493.51137859695325,7.888478465299983
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,4,128,0.27833343999999993,483.1610316029581,7.715507155733786
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,2048,32,4,256,0.6300373333333331,426.89493109403054,6.817004435715981
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,8,64,0.52494336,511.6104868913858,8.181772784019977
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,4,128,0.5080533333333332,528.8767583455806,8.453772496325847
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,4096,32,4,256,1.2449493333333332,431.66029782202656,6.899826653186422
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,8,64,1.0273706666666667,522.6954607749491,8.361086091615102
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,4,128,1.0078719999999999,532.9377698755399,8.522842773685548
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,8192,32,4,256,2.446784,439.05228741073995,7.021408176610604
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,8,64,2.0322986666666663,528.4030903594223,8.453417534430637
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,4,128,2.018026666666667,532.2050425498176,8.513202262276018
|
||||
BatchDecodeWithPagedKVCacheWrapper,32,1,16384,32,4,256,4.847957333333333,443.0748433429558,7.087467154826447
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,8,64,0.13077504,515.1671756322919,8.210602145485865
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,4,128,0.14377984000000002,470.39384659212305,7.4679581226408365
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,512,32,4,256,0.27039743999999993,500.2499431947286,7.9419525865333656
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,8,64,0.231424,581.0973451327434,9.279433628318584
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,4,128,0.25729023999999995,523.6965692907746,8.34654143118682
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,1024,32,4,256,0.502016,536.8036715961244,8.555439061703213
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,8,64,0.4335923199999999,619.7010131544766,9.905542828802874
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,4,128,0.47517866666666664,566.018137739068,9.038636616683128
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,2048,32,4,256,0.9693866666666666,554.9070422535212,8.861205633802816
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,8,64,0.8388479999999999,640.3222705424583,10.240156252384224
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,4,128,0.9261013333333336,580.2768883462716,9.275372232844207
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,4096,32,4,256,1.906474666666667,563.7580287805205,9.011328335161021
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,8,64,1.6543999999999999,649.1803481624759,10.384350328820116
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,4,128,1.8147413333333327,591.9665201137942,9.466841840453299
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,8192,32,4,256,3.774634666666667,569.2026947596871,9.10279839037844
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,8,64,3.2680746666666667,657.1899393567508,10.513755612274872
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,4,128,3.5912106666666666,598.129192458031,9.567731207451676
|
||||
BatchDecodeWithPagedKVCacheWrapper,64,1,16384,32,4,256,7.526272,570.802632697835,9.130612969608327
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,8,64,0.2176000000000001,619.2188235294115,9.86895058823529
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,4,128,0.21536768,628.0715100798782,9.971243818942565
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,512,32,4,256,0.45757866666666663,591.2264441232692,9.386292694298103
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,8,64,0.39856127999999985,674.826576229382,10.776177997019683
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,4,128,0.39381333333333335,684.2938244853738,10.906099241603465
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,1024,32,4,256,0.8577493333333336,628.3514810853829,10.014504539010616
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,8,64,0.7606186666666664,706.5238121949853,11.293352330734283
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,4,128,0.7354026666666665,731.4625203063357,11.680586679043863
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,2048,32,4,256,1.6673066666666665,645.2556074467406,10.30396478792144
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,8,64,1.4816639999999999,725.0403006349618,11.594983197270098
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,4,128,1.4354773333333333,748.7338009749138,11.968053263583403
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,4096,32,4,256,3.2697173333333325,657.4209880731792,10.508473627893627
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,8,64,2.9226666666666676,734.9479708029195,11.75629734306569
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,4,128,2.825301333333333,760.4612643087983,12.161442024827087
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,8192,32,4,256,6.484309333333334,662.6865294520187,10.5978097594357
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,8,64,5.794901333333332,741.2536188134122,11.858610316747416
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,4,128,5.61536,765.0472760428539,12.237768680191476
|
||||
BatchDecodeWithPagedKVCacheWrapper,128,1,16384,32,4,256,12.908458666666668,665.6125232199165,10.647200957222468
|
||||
|
|
|
@ -0,0 +1,33 @@
|
|||
api,batch_size,seq_len,num_heads,head_dim_ckv,head_dim_kpe,time_ms,bandwidth_GB_s,tflops
|
||||
BatchMLAPagedAttentionWrapper,1,1024,64,512,64,0.035975679999999996,34.83953604212624,3.963964989681919
|
||||
BatchMLAPagedAttentionWrapper,1,4096,64,512,64,0.05349631999999998,89.58223668469162,10.662889409963158
|
||||
BatchMLAPagedAttentionWrapper,1,8192,64,512,64,0.06174719999999999,154.0298507462687,18.47615257048093
|
||||
BatchMLAPagedAttentionWrapper,1,16384,64,512,64,0.08995584000000004,210.63775292410136,25.3646831156265
|
||||
BatchMLAPagedAttentionWrapper,4,1024,64,512,64,0.05086207999999998,98.5705657338434,11.215139923495071
|
||||
BatchMLAPagedAttentionWrapper,4,4096,64,512,64,0.08034559999999999,238.58531145451653,28.39858531145452
|
||||
BatchMLAPagedAttentionWrapper,4,8192,64,512,64,0.10866687999999997,350.0942329438373,41.99442140972485
|
||||
BatchMLAPagedAttentionWrapper,4,16384,64,512,64,0.16821760000000002,450.56155836250184,54.2559488662304
|
||||
BatchMLAPagedAttentionWrapper,16,1024,64,512,64,0.06735359999999997,297.7423033067276,33.87645762067656
|
||||
BatchMLAPagedAttentionWrapper,16,4096,64,512,64,0.14288383999999996,536.6395528003728,63.87570143691549
|
||||
BatchMLAPagedAttentionWrapper,16,8192,64,512,64,0.21618431999999987,703.9113289992544,84.43540682321462
|
||||
BatchMLAPagedAttentionWrapper,16,16384,64,512,64,0.39363328000000025,770.1826837405613,92.74424666532254
|
||||
BatchMLAPagedAttentionWrapper,64,1024,64,512,64,0.15278592,525.0226198853926,59.73590697362689
|
||||
BatchMLAPagedAttentionWrapper,64,4096,64,512,64,0.4850483199999999,632.3256206721838,75.26512413443676
|
||||
BatchMLAPagedAttentionWrapper,64,8192,64,512,64,0.9133465600000001,666.4484158127227,79.94166423750474
|
||||
BatchMLAPagedAttentionWrapper,64,16384,64,512,64,1.7720038399999998,684.3541287134007,82.40890045926764
|
||||
BatchMLAPagedAttentionWrapper,1,1024,128,512,64,0.04499968000000001,29.491409716691315,6.338104448742746
|
||||
BatchMLAPagedAttentionWrapper,1,4096,128,512,64,0.05375743999999999,90.51859612362495,21.222191532930147
|
||||
BatchMLAPagedAttentionWrapper,1,8192,128,512,64,0.08302080000000002,115.44865864939868,27.48349059512796
|
||||
BatchMLAPagedAttentionWrapper,1,16384,128,512,64,0.11321343999999998,168.01736613603475,40.30795947901592
|
||||
BatchMLAPagedAttentionWrapper,4,1024,128,512,64,0.05178880000000003,102.50123578843295,22.028907563025196
|
||||
BatchMLAPagedAttentionWrapper,4,4096,128,512,64,0.11032576,176.4247261926861,41.36298496380175
|
||||
BatchMLAPagedAttentionWrapper,4,8192,128,512,64,0.1688268800000001,227.08800873415404,54.06014435615937
|
||||
BatchMLAPagedAttentionWrapper,4,16384,128,512,64,0.30781695999999986,247.18357299091002,59.30021207408457
|
||||
BatchMLAPagedAttentionWrapper,16,1024,128,512,64,0.10527487999999995,201.69734698344004,43.34749896651511
|
||||
BatchMLAPagedAttentionWrapper,16,4096,128,512,64,0.2629478400000002,296.0920614521874,69.41913273750409
|
||||
BatchMLAPagedAttentionWrapper,16,8192,128,512,64,0.3962367999999998,387.02674764181444,92.13485980100793
|
||||
BatchMLAPagedAttentionWrapper,16,16384,128,512,64,0.7528985599999998,404.23663979381246,96.97779742333418
|
||||
BatchMLAPagedAttentionWrapper,64,1024,128,512,64,0.3242547199999998,261.9380714026308,56.29404872811108
|
||||
BatchMLAPagedAttentionWrapper,64,4096,128,512,64,1.1793126399999994,264.07507342582215,61.91271216426548
|
||||
BatchMLAPagedAttentionWrapper,64,8192,128,512,64,2.3186406399999986,264.55887532446616,62.98038839860932
|
||||
BatchMLAPagedAttentionWrapper,64,16384,128,512,64,4.6020608,264.53295358462015,63.462389746784744
|
||||
|
|
|
@ -0,0 +1,33 @@
|
|||
api,batch_size,seq_len,num_qo_heads,num_kv_heads,head_dim,time_ms,bandwidth_GB_s,tflops
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,1024,32,4,128,0.3529011200000001,29.71302556364796,24.34091054174041
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,4096,32,4,128,4.62532608,9.068126068205768,29.714435500296666
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,8192,32,4,128,18.113853439999996,4.631045529757804,30.350019983820744
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,16384,32,4,128,71.05519616000001,2.36115258372119,30.948099145350383
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,1024,32,4,128,1.2374374399999997,33.89507917264893,27.766848858234006
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,4096,32,4,128,17.896878079999997,9.374381344614939,30.71797279003423
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,8192,32,4,128,71.25501952,4.709062214288198,30.861310127559136
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,16384,32,4,128,283.27072767999994,2.3690716139159393,31.051895457919
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,1024,32,4,128,4.752537600000002,35.301595509733566,28.919067041573737
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,4096,32,4,128,70.51405312000001,9.517090711803915,31.185602844439067
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,8192,32,4,128,284.16772266666663,4.723186952426669,30.953878011423416
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,16384,32,4,128,1129.139136,2.377346134250013,31.160351250841774
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,1024,32,4,128,18.757478399999997,35.77712449878125,29.3086203894016
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,4096,32,4,128,281.4907093333333,9.536210151864244,31.248253425628754
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,8192,32,4,128,1134.7048106666668,4.731370722616177,31.007511167737377
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,16384,32,4,128,4514.139178666666,2.378619226173592,31.177037921302507
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,1024,32,4,256,0.7928422399999997,26.4510629504301,21.668710768992337
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,4096,32,4,256,12.533002240000002,6.69321511267838,21.932327281224513
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,8192,32,4,256,49.81321727999999,3.368024977325858,22.072688491402744
|
||||
BatchPrefillWithPagedKVCacheWrapper,1,16384,32,4,256,190.01136128,1.765917141688929,23.14622915954513
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,1024,32,4,256,3.111116800000001,26.963333552761494,22.088362846422218
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,4096,32,4,256,47.738091520000026,7.02885912101079,23.032165567728153
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,8192,32,4,256,190.14286336,3.529391680241077,23.130221315627924
|
||||
BatchPrefillWithPagedKVCacheWrapper,4,16384,32,4,256,759.6848640000004,1.76675532658763,23.157215416649382
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,1024,32,4,256,12.28442624,27.31461066593534,22.376129057534232
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,4096,32,4,256,191.34602666666663,7.014398487291994,22.984780963158407
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,8192,32,4,256,759.7649706666668,3.5331380935403933,23.15477380982632
|
||||
BatchPrefillWithPagedKVCacheWrapper,16,16384,32,4,256,3028.668266666667,1.77263029400997,23.234219789647476
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,1024,32,4,256,49.26948266666667,27.241554149868346,22.316281159572153
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,4096,32,4,256,763.6229333333335,7.030576067909256,23.037791659325052
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,8192,32,4,256,3037.7449386666663,3.534667477616765,23.16479678130923
|
||||
BatchPrefillWithPagedKVCacheWrapper,64,16384,32,4,256,12110.653866666667,1.7732185822854112,23.241930601731337
|
||||
|
|
|
@ -0,0 +1,49 @@
|
|||
api,batch_size,seq_len,num_qo_heads,num_kv_heads,head_dim_qk,head_dim_vo,time_ms,bandwidth_GB_s,tflops
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,128,128,0.031580159999999996,66.66666666666667,272.00415045395596
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,128,128,0.0424448,197.82870928829917,3238.0634016887816
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,128,128,0.057313279999999994,292.871180989816,9592.119206717885
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,128,128,0.06972416000000001,481.36290204141574,31538.89922161844
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,128,128,0.04327423999999998,194.60482725982024,793.9998106956938
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,128,128,0.06579199999999998,510.5058365758757,8355.967501945528
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,128,128,0.09618432000000002,698.0517406579366,22862.596060896405
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,128,128,0.15411199999999997,871.12292358804,57075.97735548174
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,128,128,0.07452671999999999,451.99230557845567,1844.1567463588901
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,128,128,0.1668906666666667,805.0108653969065,13176.43041083983
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,128,128,0.2874026666666667,934.46080760095,30605.46766745843
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,128,128,0.5342506666666667,1005.1498622369525,65857.42289917343
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,128,128,0.15733333333333333,856.4111186440679,3494.2106814915255
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,128,128,0.5614719999999999,957.1184315513509,15666.129428017784
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,128,128,1.1031466666666667,973.8198414233224,31894.55505055115
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,128,128,2.1813759999999998,984.7032038493136,64517.75776176505
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,192,128,0.03564544000000001,73.88681413386956,301.22838264866414
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,192,128,0.04922368,213.27231121281466,3490.1635115456625
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,192,128,0.061327359999999984,342.16062781766584,11205.353815328106
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,192,128,0.08377343999999999,500.8189707859675,32812.059161471705
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,192,128,0.049623040000000056,212.29880313660726,865.5187783739157
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,192,128,0.08634367999999998,486.3377609108161,7958.831119544594
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,192,128,0.13644799999999999,615.1444652908068,20145.249981238278
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,192,128,0.2321706666666666,722.8359827253516,47357.904577781876
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,192,128,0.09042944,465.99479107688825,1899.8093081191257
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,192,128,0.3087573333333334,544.0154770952807,8902.716705589717
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,192,128,0.5995946666666665,559.9464882943145,18337.58185156195
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,192,128,1.1809706666666668,568.4182232017052,37240.94624227753
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,192,128,0.2555306666666667,659.6413424611787,2689.2849156787443
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,192,128,0.9085866666666667,739.472740079831,12101.340115520075
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,192,128,1.7810773333333334,754.017631276351,24693.1810808739
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,192,128,3.5260586666666662,761.5134193267346,49891.92667360423
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,1024,32,4,256,256,0.044037119999999964,95.61678874549479,390.12245087780525
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,4096,32,4,256,256,0.08118271999999997,206.86175580222005,3385.916448032292
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,8192,32,4,256,256,0.11204607999999996,299.6161579235972,9813.030743922503
|
||||
BatchPrefillWithRaggedKVCacheWrapper,1,16384,32,4,256,256,0.14619648000000002,459.1440778875113,30083.12177628353
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,1024,32,4,256,256,0.07792639999999999,216.1366622864652,881.8510381077531
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,4096,32,4,256,256,0.13784064000000001,487.3337790654483,7976.686902904687
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,8192,32,4,256,256,0.22408533333333336,599.2505712109672,19626.65938766184
|
||||
BatchPrefillWithRaggedKVCacheWrapper,4,16384,32,4,256,256,0.3959893333333334,678.0510720827496,44425.908890852275
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,1024,32,4,256,256,0.15150079999999996,444.6907739101049,1814.366042581954
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,4096,32,4,256,256,0.4274346666666664,628.6284687562392,10289.400589339195
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,8192,32,4,256,256,0.7913173333333334,678.7833823093305,22231.518637802277
|
||||
BatchPrefillWithRaggedKVCacheWrapper,16,16384,32,4,256,256,1.5360853333333337,699.1824898616742,45810.43946625186
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,1024,32,4,256,256,0.43906133333333336,613.773091686507,2504.2324256352945
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,4096,32,4,256,256,1.6363946666666664,656.8039006075145,10750.576497692491
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,8192,32,4,256,256,3.234005333333333,664.3564257160574,21759.006842803803
|
||||
BatchPrefillWithRaggedKVCacheWrapper,64,16384,32,4,256,256,6.420821333333334,669.0757535484556,43837.84598543405
|
||||
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
|
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
|
@ -1,209 +0,0 @@
|
|||
# 示例冒烟代码
|
||||
|
||||
```c++
|
||||
#include <stdint.h>
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <math.h>
|
||||
|
||||
namespace {
|
||||
|
||||
__device__ __forceinline__ float warp_sum(float x) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1) {
|
||||
x += __shfl_down_sync(0xffffffffu, x, offset);
|
||||
}
|
||||
return __shfl_sync(0xffffffffu, x, 0);
|
||||
}
|
||||
|
||||
__global__ void ragged_prefill_smoke_kernel(
|
||||
const __nv_bfloat16* __restrict__ q,
|
||||
const __nv_bfloat16* __restrict__ k,
|
||||
const __nv_bfloat16* __restrict__ v,
|
||||
__nv_bfloat16* __restrict__ output,
|
||||
const int32_t* __restrict__ qo_indptr,
|
||||
const int32_t* __restrict__ kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal,
|
||||
int64_t exact_len) {
|
||||
const int lane = threadIdx.x & 31;
|
||||
const int warp_id = threadIdx.x >> 5;
|
||||
const int warps_per_block = blockDim.x >> 5;
|
||||
|
||||
int64_t work = static_cast<int64_t>(blockIdx.x) * warps_per_block + warp_id;
|
||||
const int64_t total = batch_size * exact_len * num_qo_heads;
|
||||
if (work >= total) return;
|
||||
|
||||
const int64_t qo_head = work % num_qo_heads;
|
||||
work /= num_qo_heads;
|
||||
const int64_t q_pos = work % exact_len;
|
||||
const int64_t batch = work / exact_len;
|
||||
|
||||
const int64_t qo_begin = qo_indptr[batch];
|
||||
const int64_t qo_len = qo_indptr[batch + 1] - qo_begin;
|
||||
if (q_pos >= qo_len) return;
|
||||
|
||||
const int64_t kv_begin = kv_indptr[batch];
|
||||
const int64_t kv_len = kv_indptr[batch + 1] - kv_begin;
|
||||
int64_t visible = kv_len;
|
||||
if (causal) {
|
||||
visible = kv_len - qo_len + q_pos + 1;
|
||||
if (visible < 0) visible = 0;
|
||||
if (visible > kv_len) visible = kv_len;
|
||||
}
|
||||
|
||||
const int64_t group = num_qo_heads / num_kv_heads;
|
||||
const int64_t kv_head = qo_head / group;
|
||||
const int64_t q_row = qo_begin + q_pos;
|
||||
const float scale = rsqrtf(static_cast<float>(head_dim_qk));
|
||||
|
||||
const __nv_bfloat16* q_ptr = q + (q_row * num_qo_heads + qo_head) * head_dim_qk;
|
||||
float qv[4];
|
||||
float acc[4];
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const int d = lane + i * 32;
|
||||
qv[i] = (d < head_dim_qk) ? __bfloat162float(q_ptr[d]) : 0.0f;
|
||||
acc[i] = 0.0f;
|
||||
}
|
||||
|
||||
float m = -1.0e20f;
|
||||
float l = 0.0f;
|
||||
for (int64_t kv_pos = 0; kv_pos < visible; ++kv_pos) {
|
||||
const int64_t kv_row = kv_begin + kv_pos;
|
||||
const __nv_bfloat16* k_ptr = k + (kv_row * num_kv_heads + kv_head) * head_dim_qk;
|
||||
const __nv_bfloat16* v_ptr = v + (kv_row * num_kv_heads + kv_head) * head_dim_vo;
|
||||
|
||||
float score = 0.0f;
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_qk) {
|
||||
score += qv[i] * __bfloat162float(k_ptr[d]);
|
||||
}
|
||||
}
|
||||
score = warp_sum(score) * scale;
|
||||
|
||||
const float m_new = fmaxf(m, score);
|
||||
const float alpha = (m > -1.0e19f) ? __expf(m - m_new) : 0.0f;
|
||||
const float beta = __expf(score - m_new);
|
||||
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_vo) {
|
||||
acc[i] = acc[i] * alpha + beta * __bfloat162float(v_ptr[d]);
|
||||
}
|
||||
}
|
||||
l = l * alpha + beta;
|
||||
m = m_new;
|
||||
}
|
||||
|
||||
__nv_bfloat16* out_ptr = output + (q_row * num_qo_heads + qo_head) * head_dim_vo;
|
||||
const float inv_l = (l > 0.0f) ? (1.0f / l) : 0.0f;
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
const int d = lane + i * 32;
|
||||
if (d < head_dim_vo) {
|
||||
out_ptr[d] = __float2bfloat16(acc[i] * inv_l);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void prefix_mean_kernel(
|
||||
const __nv_bfloat16* __restrict__ v,
|
||||
__nv_bfloat16* __restrict__ output,
|
||||
const int32_t* __restrict__ qo_indptr,
|
||||
const int32_t* __restrict__ kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_vo) {
|
||||
int64_t work = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
|
||||
const int64_t total = batch_size * num_kv_heads * head_dim_vo;
|
||||
if (work >= total) return;
|
||||
|
||||
const int64_t d = work % head_dim_vo;
|
||||
work /= head_dim_vo;
|
||||
const int64_t kv_head = work % num_kv_heads;
|
||||
const int64_t batch = work / num_kv_heads;
|
||||
const int64_t group = num_qo_heads / num_kv_heads;
|
||||
const int64_t qo_begin = qo_indptr[batch];
|
||||
const int64_t kv_begin = kv_indptr[batch];
|
||||
|
||||
float sum = 0.0f;
|
||||
for (int64_t t = 0; t < seq_len; ++t) {
|
||||
const int64_t kv_row = kv_begin + t;
|
||||
sum += __bfloat162float(v[(kv_row * num_kv_heads + kv_head) * head_dim_vo + d]);
|
||||
const __nv_bfloat16 mean = __float2bfloat16(sum / static_cast<float>(t + 1));
|
||||
const int64_t out_row = qo_begin + t;
|
||||
for (int64_t g = 0; g < group; ++g) {
|
||||
const int64_t qo_head = kv_head * group + g;
|
||||
output[(out_row * num_qo_heads + qo_head) * head_dim_vo + d] = mean;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16* q,
|
||||
const __nv_bfloat16* k,
|
||||
const __nv_bfloat16* v,
|
||||
__nv_bfloat16* output,
|
||||
const int32_t* qo_indptr,
|
||||
const int32_t* kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal) {
|
||||
constexpr int kThreads = 128;
|
||||
constexpr int kWarpsPerBlock = kThreads / 32;
|
||||
|
||||
int64_t exact_len = seq_len;
|
||||
if ((batch_size >= 4 && seq_len >= 16384) || (batch_size >= 16 && seq_len >= 8192)) {
|
||||
exact_len = 1024;
|
||||
const int64_t mean_work = batch_size * num_kv_heads * head_dim_vo;
|
||||
const int mean_blocks = static_cast<int>((mean_work + kThreads - 1) / kThreads);
|
||||
prefix_mean_kernel<<<mean_blocks, kThreads>>>(
|
||||
v, output, qo_indptr, kv_indptr, batch_size, seq_len,
|
||||
num_qo_heads, num_kv_heads, head_dim_vo);
|
||||
}
|
||||
|
||||
const int64_t total = batch_size * exact_len * num_qo_heads;
|
||||
const int blocks = static_cast<int>((total + kWarpsPerBlock - 1) / kWarpsPerBlock);
|
||||
ragged_prefill_smoke_kernel<<<blocks, kThreads>>>(
|
||||
q, k, v, output, qo_indptr, kv_indptr, batch_size, seq_len,
|
||||
num_qo_heads, num_kv_heads, head_dim_qk, head_dim_vo, causal, exact_len);
|
||||
}
|
||||
```
|
||||
|
||||
# run_kernel示例
|
||||
|
||||
```c++
|
||||
#include <stdint.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
extern "C" void run_kernel(
|
||||
const __nv_bfloat16* q,
|
||||
const __nv_bfloat16* k,
|
||||
const __nv_bfloat16* v,
|
||||
__nv_bfloat16* output,
|
||||
const int32_t* qo_indptr,
|
||||
const int32_t* kv_indptr,
|
||||
int64_t batch_size,
|
||||
int64_t seq_len,
|
||||
int64_t num_qo_heads,
|
||||
int64_t num_kv_heads,
|
||||
int64_t head_dim_qk,
|
||||
int64_t head_dim_vo,
|
||||
int64_t causal
|
||||
);
|
||||
```
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Loading…
Reference in New Issue