Improve bilingual project documentation
中文说明 · 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。
flash-attn
atol=0.04, rtol=0.04
基准是算子级对比:候选后端与本项目 PyTorch reference 对比,不等同于端到端模型生成性能。原始数据、报告和汇总图见 docs/results/。
docs/results/
详细边界见 docs/support-matrix.md。
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。
flash-attn 2.6.3+metax3.5.3.9torch2.8
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
公开接口为:
flash_attn_func
flash_attn_varlen_func
flash_attn_with_kvcache
dense Q/K/V 使用 [batch, sequence, heads, head_dim] 布局,当前支持 head dimension 64 和 128:
[batch, sequence, heads, head_dim]
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 返回会明确报错,不会静默产生错误结果。
RuntimeWarning
get_last_dispatch_info()
默认配置覆盖 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,需要单独的生成实验。
集成目录提供 vLLM 适配桥接和 preflight 配置,目标版本为 vLLM 0.30.0。已检查的 C500 镜像为 vLLM 0.17.0,版本不匹配,因此当前没有声称 vLLM backend 注册或端到端生成成功;SGLang 也尚未安装验证。详见 integrations/vllm/README.md。
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 条件测试
本项目采用 Apache License 2.0。MXMACA runtime、vendor PyTorch、MetaX wheel、vLLM、模型权重和数据集分别遵循各自许可证。
欢迎通过 CONTRIBUTING.md 提交问题和改进。
CONTRIBUTING.md
面向 MXMACA / MetaX C500 的 FlashAttention 兼容算子与推理适配项目,支持 FP16/BF16、GQA/MQA、varlen、paged KV cache,并提供可复现的 C500 benchmark 和 vLLM 适配桥接。
版权所有:中国计算机学会技术支持:开源发展技术委员会 京ICP备13000930号-9 京公网安备 11010802047560号
MXFlashAttn
中文说明 · English
MXFlashAttn 是面向 MXMACA / MetaX C500 的 FlashAttention 兼容前向算子与推理适配项目。项目提供 dense、varlen 和 KV cache API,支持显式后端选择、可追踪的 fallback 行为,以及可复现的 C500 benchmark。
已验证结果
atol=0.04, rtol=0.04)基准是算子级对比:候选后端与本项目 PyTorch reference 对比,不等同于端到端模型生成性能。原始数据、报告和汇总图见
docs/results/。能力矩阵
详细边界见
docs/support-matrix.md。安装
C500 必须先安装与驱动匹配的 MXMACA vendor PyTorch 和 MetaX
flash-attnwheel。当前验证版本为flash-attn 2.6.3+metax3.5.3.9torch2.8,通用 CUDA wheel 不会被当作 MXMACA provider。开发依赖和本地 reference 测试:
如需构建仓库内 ATen 扩展并强制使用它进行硬件验证:
API 示例
公开接口为:
flash_attn_funcflash_attn_varlen_funcflash_attn_with_kvcachedense Q/K/V 使用
[batch, sequence, heads, head_dim]布局,当前支持 head dimension 64 和 128:在没有 MetaX provider 或编译扩展时,调用默认报错。只有主动设置以下变量才会启用 PyTorch reference fallback:
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 上复现实验:
每组记录误差、tokens/s、prefill/decode 延迟、峰值显存、backend、fallback 原因以及 PyTorch/MXMACA/驱动/vLLM 版本。模型级首 token 延迟不属于该算子 benchmark,需要单独的生成实验。
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-InstructQwen/Qwen2.5-7B-Instruct仓库结构
路线图
许可证
本项目采用 Apache License 2.0。MXMACA runtime、vendor PyTorch、MetaX wheel、vLLM、模型权重和数据集分别遵循各自许可证。
欢迎通过
CONTRIBUTING.md提交问题和改进。