目录

MXFlashAttn

CI License Python

中文说明 · English

MXFlashAttn 是面向 MXMACA / MetaX C500 的 FlashAttention 兼容前向算子与推理适配项目。项目提供 dense、varlen 和 KV cache API,支持显式后端选择、可追踪的 fallback 行为,以及可复现的 C500 benchmark。

项目定位:先把可运行、可验证、可复现的国产算力适配交付出来。仓库中的 C++ 扩展是 PyTorch ATen bring-up/correctness 路径,不冒充项目自研的融合 FlashAttention kernel;C500 优化执行依赖匹配的 MetaX flash-attn wheel。

已验证结果

指标 结果
C500 benchmark 384/384 组完成,未发生 reference fallback
Decode 延迟 中位下降 51.64%,185/192 组达到 20% 目标
全矩阵延迟 中位下降 53.77%
最大绝对误差 0.0078125(测试阈值 atol=0.04, rtol=0.04)
C500 环境 32 GiB sGPU 配额;MXMACA 3.5.3.20;驱动 3.8.30;PyTorch 2.8.0
本地 CI 31 passed,8 个 C500 专属测试在无硬件环境跳过

基准是算子级对比:候选后端与本项目 PyTorch reference 对比,不等同于端到端模型生成性能。原始数据、报告和汇总图见 docs/results/。

能力矩阵

能力 状态
FP16 / BF16 forward C500 已验证
Head dimension 64 / 128 C500 已验证
GQA / MQA、causal C500 已验证
varlen、dense KV cache、paged KV cache API 已实现;paged KV append 已在 C500 验证
PyTorch reference fallback 显式开启后可用
backward、FP8、多卡 首个版本不支持
vLLM backend 注册与端到端生成 尚未验证

详细边界见 docs/support-matrix.md。

安装

C500 必须先安装与驱动匹配的 MXMACA vendor PyTorch 和 MetaX flash-attn wheel。当前验证版本为 flash-attn 2.6.3+metax3.5.3.9torch2.8,通用 CUDA wheel 不会被当作 MXMACA provider。

python -m pip install -r requirements-c500.txt
python -m pip install --no-build-isolation -e .

开发依赖和本地 reference 测试:

python -m pip install -e ".[dev]"
python -m pytest

如需构建仓库内 ATen 扩展并强制使用它进行硬件验证:

MXFLASHATTN_BUILD_NATIVE=1 python -m pip install --no-build-isolation -e .
MXFLASHATTN_BACKEND=aten python -m pytest -q

API 示例

公开接口为:

  • flash_attn_func
  • flash_attn_varlen_func
  • flash_attn_with_kvcache

dense Q/K/V 使用 [batch, sequence, heads, head_dim] 布局,当前支持 head dimension 64 和 128:

import torch
from mxflashattn import flash_attn_func, get_last_dispatch_info

q = torch.randn(1, 32, 8, 64, dtype=torch.float16, device="cuda")
k = torch.randn(1, 128, 2, 64, dtype=q.dtype, device=q.device)
v = torch.randn_like(k)
out = flash_attn_func(q, k, v, causal=True)
print(get_last_dispatch_info())

在没有 MetaX provider 或编译扩展时,调用默认报错。只有主动设置以下变量才会启用 PyTorch reference fallback:

# Windows cmd
set MXFLASHATTN_ALLOW_FALLBACK=1
# Windows PowerShell
$env:MXFLASHATTN_ALLOW_FALLBACK="1"
# Linux / macOS
export MXFLASHATTN_ALLOW_FALLBACK=1

fallback 会发出 RuntimeWarning,并通过 get_last_dispatch_info() 和 benchmark 报告记录原因。dropout、局部窗口、ALiBi、soft-capping、rotary embeddings、backward 和 softmax-LSE 返回会明确报错,不会静默产生错误结果。

Benchmark

默认配置覆盖 384 组 batch size、query length、KV length、head dimension、GQA 比例、cache layout、dtype 和 causal 组合。C500 上复现实验:

python -m benchmarks.run \
  --config benchmarks/configs/c500.yaml \
  --device cuda --warmup 1 --repeats 3 --limit 384 \
  --output artifacts/c500.json
python -m benchmarks.report artifacts/c500.json --output artifacts/c500.md
python -m pip install -e ".[reports]"
python -m benchmarks.plot_summary artifacts/c500.json --output artifacts/c500-summary.png

每组记录误差、tokens/s、prefill/decode 延迟、峰值显存、backend、fallback 原因以及 PyTorch/MXMACA/驱动/vLLM 版本。模型级首 token 延迟不属于该算子 benchmark,需要单独的生成实验。

C500 benchmark summary

vLLM 与模型演示

集成目录提供 vLLM 适配桥接和 preflight 配置,目标版本为 vLLM 0.30.0。已检查的 C500 镜像为 vLLM 0.17.0,版本不匹配,因此当前没有声称 vLLM backend 注册或端到端生成成功;SGLang 也尚未安装验证。详见 integrations/vllm/README.md。

计划使用以下模型做后续 smoke test,权重不会提交仓库:

  • Qwen/Qwen2.5-1.5B-Instruct
  • Qwen/Qwen2.5-7B-Instruct

仓库结构

mxflashattn/       Python API、校验、reference 与 dispatch
csrc/mxmac/        C++/ATen bring-up extension
integrations/vllm/ vLLM 适配桥接与 preflight
benchmarks/        固定 YAML 配置、运行器和报告工具
docs/results/      C500 原始结果、报告和图表
tests/             CPU/reference 与 C500 条件测试

路线图

  1. 针对 C500 兼容版本完成 vLLM 或 SGLang 的真实模型生成 smoke test。
  2. 优化剩余未达到 20% decode 延迟下降的场景。
  3. 在独立里程碑中评估 backward、FP8 和多卡能力。

许可证

本项目采用 Apache License 2.0。MXMACA runtime、vendor PyTorch、MetaX wheel、vLLM、模型权重和数据集分别遵循各自许可证。

欢迎通过 CONTRIBUTING.md 提交问题和改进。

关于

面向 MXMACA / MetaX C500 的 FlashAttention 兼容算子与推理适配项目,支持 FP16/BF16、GQA/MQA、varlen、paged KV cache,并提供可复现的 C500 benchmark 和 vLLM 适配桥接。

218.0 KB
邀请码
    Gitlink(确实开源)
  • 加入我们
  • 官网邮箱:gitlink@ccf.org.cn
  • QQ群
  • QQ群
  • 公众号
  • 公众号

版权所有:中国计算机学会技术支持:开源发展技术委员会
京ICP备13000930号-9 京公网安备 11010802047560号