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 表示本次前向与反向计算没有出现非有限值;

  • 本示例使用单卡,不涉及多卡分布式训练。