flash-linear-attention
在单张昇腾 NPU 上安装 flash-linear-attention,并运行 GatedDeltaNet 的
前向与反向计算。完成后,可以确认 Triton-Ascend 能识别 NPU,且模型输出
和输入梯度均正常。
前置条件
硬件
Atlas 800T / 900 A2 训练系列;
至少一张可用的 Ascend 910B NPU;
物理机或容器已正确配置驱动和设备。
基础软件
在运行本文档之前,机器上需要已经安装并可用:
Linux aarch64 操作系统;
可用的 Python 环境;
可用的 CANN toolkit 和驱动;
npu-smi能正常显示 NPU 设备。
CANN 安装可参考快速安装昇腾环境。Torch、Torch-NPU、torchvision 和 Triton-Ascend 由目标 release 的 [npu] extra 安装。
本文档示例使用的版本
以下是本示例使用的依赖环境,镜像为
swr.cn-south-1.myhuaweicloud.com/ascendhub/cann:9.0.0-910b-ubuntu22.04-py3.11。
组件 |
版本 |
|---|---|
操作系统 |
Ubuntu 22.04,Linux aarch64 |
Python |
3.11 |
CANN |
9.0.0 |
flash-linear-attention |
xxx |
torch |
2.7.1+cpu |
torch_npu |
2.7.1.post4 |
torchvision |
0.22.1 |
triton-ascend |
3.2.1(分发包版本) |
NPU |
Ascend 910B × 1 |
Note
不同 release 的 NPU 依赖可能不同。安装时以所选 release 的 [npu] 依赖为准,
不要混用 main 分支的版本要求;如需升级 CANN,请先核对相应版本的兼容性。
检查前置是否满足
source /usr/local/Ascend/ascend-toolkit/set_env.sh
export PATH=/usr/local/sbin:$PATH
test -n "$ASCEND_HOME_PATH"
command -v npu-smi >/dev/null
printf 'CANN ready\n'
输出结果如下:
CANN ready
安装 flash-linear-attention
安装命令会读取所选 release 的 [npu] 依赖,安装匹配的 Torch、Torch-NPU、
torchvision 与 Triton-Ascend。
test ! -e flash-linear-attention
git clone https://github.com/fla-org/flash-linear-attention.git
cd flash-linear-attention
git checkout <UPSTREAM_REF>
python -m pip install -q -U pip setuptools wheel
python -m pip install -q pybind11 cmake attrs sympy pyyaml scipy decorator einops
python -m pip install -q ".[npu]" \
--extra-index-url https://triton-ascend.osinfra.cn/pypi/simple
python -c "import fla; print('fla', fla.__version__)"
输出结果如下:
fla xxx
Note
xxx 表示最新的版本号。
<UPSTREAM_REF> 替换为 flash-linear-attention 当前最新 release 的标签,可从
Releases 获取。
验证 Ascend NPU backend
flash-linear-attention 通过 Triton runtime 识别 npu backend,并将 fla.utils.IS_NPU 设置为 True。
以下代码用 Python 执行:
from importlib.metadata import version
import torch
import torch_npu
import triton
from fla.utils import IS_NPU, device_platform
assert torch.npu.is_available()
assert IS_NPU
assert device_platform == "npu"
print("torch", torch.__version__)
print("torch_npu", torch_npu.__version__)
print("triton-ascend", version("triton-ascend"))
print("device_platform", device_platform)
print("npu_available", torch.npu.is_available())
输出结果如下:
torch 2.7.1+cpu
torch_npu 2.7.1.post4
triton-ascend 3.2.1
device_platform npu
npu_available True
使用样例
GatedDeltaNet 前向与反向
创建一个较小的 GatedDeltaNet 模型,在 NPU 上完成前向计算和反向传播。
示例会检查输出形状,并确认输出和输入梯度均为有限值。
以下代码用 Python 执行:
import torch
import torch_npu
from fla.layers import GatedDeltaNet
from fla.utils import IS_NPU, device_platform
assert IS_NPU
assert device_platform == "npu"
torch.manual_seed(42)
layer = GatedDeltaNet(
hidden_size=512,
head_dim=64,
num_heads=6,
expand_v=2,
mode="chunk",
).to(device="npu", dtype=torch.bfloat16).train()
x = torch.randn(
1,
128,
512,
device="npu",
dtype=torch.bfloat16,
requires_grad=True,
)
y = layer(x)[0]
loss = y.float().square().mean()
loss.backward()
torch.npu.synchronize()
assert y.shape == (1, 128, 512)
assert torch.isfinite(y).all()
assert x.grad is not None
assert torch.isfinite(x.grad).all()
print("device", y.device.type)
print("output_shape", tuple(y.shape))
print("forward_finite", torch.isfinite(y).all().item())
print("backward_finite", torch.isfinite(x.grad).all().item())
输出结果如下:
device npu
output_shape (1, 128, 512)
forward_finite True
backward_finite True
说明
device npu表示输出位于昇腾 NPU;forward_finite True和backward_finite True表示本次前向与反向计算没有出现非有限值;本示例使用单卡,不涉及多卡分布式训练。