已开启
【Feature】triton算子(kda_gate_fwd_kernel)迁移mindspeed-ops仓 #9
LinShua创建于  5月18日
LinShua成员
5月18日 创建

任务描述

基于fla-org/flash-linear-attention开源仓库的triton源码,迁移优化triton算子在NPU上跑通及性能优化。
本期任务的具体信息如下:
(1)算子名称:kda_gate_fwd_kernel
(2) fla-org/flash-linear-attention开放仓库的算子源码链接:
https://github.com/fla-org/flash-linear-attention/blob/v0.4.1/fla/ops/kda/gate.py
(3) 迁移指导文档:https://gitcode.com/Ascend/triton-ascend/blob/main/docs/zh/programming_guide.md
参考 skill 地址:https://gitcode.com/Ascend/agent-skills/tree/master/skills/simple-vector-triton-gpu-to-npu
(4) 开发平台:Atlas 800T A2或者A3

验收标准

一、任务交付件
本期任务为基于fla-org/flash-linear-attention开放仓库代码进行功能扩展,请合入开发代码。主要开发点如下:
(1) triton算子针对NPU的适配修改优化后的代码
(2) 针对该triton算子的测试用例UT
二、验收标准:
1)精度/性能要求
精度:
对比GPU相同输入算子精度误差小于0.1%
若无GPU标杆,与CPU小算子对齐,精度误差小于0.1%
性能:
补充说明:具体性能对比数据呈现,按 CV类算子达到0.7x竞品, VV类算子达到0.9x竞品
以下以竞品A100为例说明计算具体的性能对比要求:

  • 1.先计算当前运行NPU设备与对应竞品GPU的算力换算比例
硬件规格 910B2 A100 910B2/A100
Cube算力 354 312 1.13x
Vector算力 11.06 19.5 0.56x
带宽 1800 2039 0.88x
  • 2.根据算力换算比例计算对应算子类在不同bound场景的实际性能要求

910B2 VS A100

bound场景 CV类算子(0.7X) VV类算子(0.9x)
Cube算力bound场景 0.8x 1.0X
Vector算力bound场景 0.4x 0.5X
Memory bound场景(HBM) 0.6x 0.8x

计算式子:以CV类算子的Cube算力bound场景为例:1.13 * 0.7 =0.8
2)实践文档:
triton算子介绍1篇
3)任务完成标准
本次任务完成标准为:
精度/性能(根据实际要求)达标,PR完成合入,实践文档提交到仓库issue。

PR 合入

本地完成测试验证后,向MindSpeed-Ops的master分支发起PR。

对接人

LinSHua

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
LLinShua成员
5月18日 修改标题为 “【Feature】triton算子(kda_gate_fwd_kernel)迁移mindspeed-ops仓”,原标题为“【Feature】triton算子迁移mindspeed-ops仓”
LLinShua成员
6月4日 添加了label:feature
Ronald1995成员
6月11日 评论:

/label add good-first-issue

likedislike
ascend-robotascend-robot成员
6月11日 添加了label:good-first-issue
wangx700
wangx700成员
6月12日 评论:

/label add triage-review

likedislike
ascend-robotascend-robot成员
6月12日 添加了label:triage-review
咋哈
咋哈成员
6月18日 评论:

认领这个任务@LinShua

likedislike
Ronald1995成员
6月25日 评论:

认领这个任务@LinShua

@meleys

欢迎认领任务,请参考https://gitcode.com/Ascend/MindSpeed-Ops/issues/3 社区任务池明确该任务的:

  • 完成的截止日期
  • 开发进展反馈
  • 微信答疑群
  • 任务交付注意事项

等信息。如果您同时认领了多项任务,但无法都能进行投入,可以在部分任务中回复退出.

麻烦您加入到对应微信群,群备注名修改为"社区任务+您的gitcode账号", 后续有相关消息和问题都可以在微信群咨询答疑。 等您加入到微信群后,我这边会在社区任务池里面登记任务责任人。

likedislike
Ronald1995成员
6月25日 评论:

认领这个任务@LinShua

@meleys

关于任务进展反馈的要求,您这边的开始时间可以从今天开始算,因为之前您认领任务时,这块细节还没敲定。

likedislike
gcw_MTeIQk9o
6月29日 评论:

认领这个任务

likedislike
yuhanBai成员
6月30日 评论:

认领这个任务

@gcw_MTeIQk9o

你好,感谢您对昇腾平台的支持,当前一个任务只分配一个开发,根据评论认领时间确定,您可以后续观察该任务是否被释放退出,或者认领其他任务 https://gitcode.com/org/Ascend/discussions/4

likedislike
zeshengzongzeshengzong成员
7月2日 关联了看板:@zeshengzong的看板 20260702
ascend-robotascend-robot成员
7月3日 关联了看板:MindStudio ISSUE管理
gcw_K5CvmS79成员
7月3日 评论:

认领这个任务

likedislike
咋哈
咋哈成员
7月6日 评论:

已经完成环境的配置,正在分析算子源码,两周内完善详设

likedislike
yuhanBai成员
7月7日 评论:

认领这个任务

@gcw_K5CvmS79

你好,感谢您对昇腾平台的支持,当前一个任务只分配一个开发,根据评论认领时间确定。请关注https://gitcode.com/Ascend/MindSpeed-Ops/issues/3任务是否释放或者认领其他任务 https://gitcode.com/org/Ascend/discussions/4

likedislike
Lliuzhexu成员
7月11日 关联了里程碑:MindSpeed 26.2.0
咋哈
咋哈成员
7月12日 评论:
咋哈
咋哈成员
7月15日 评论:

环境搭好,基础流程已经打通

likedislike
咋哈
咋哈成员
7月19日 评论:

好的,目前已经初步跑通流程正在进一步优化

likedislike
咋哈
咋哈成员
7月27日 评论:

性能还未达标正在进一步优化

likedislike
咋哈
咋哈成员
8月3日 评论:

性能还未达标正在进一步优化

likedislike
8月3日 评论:

我要认领这个任务

likedislike
LLinShua成员
8月7日 修改了issue 的描述
咋哈
咋哈成员
8月11日 评论:

待根据标准在gpu完成具体性能对比

likedislike
咋哈
咋哈成员
8月17日 评论:

您好,由于上周较忙以及资源受限,部分代码在重新打开环境后遗失,我将于这周重新验证

likedislike
咋哈
咋哈成员
8月24日 评论:

已完成准备推送代码

likedislike
咋哈
咋哈成员
8月30日 评论:

KDA_GATE 算子分析与迁移指南

算子概述

算子名称:kda_gate_fwd (KDA Gate Forward)

来源:Flash Linear Attention (FLA) 项目 v0.4.1

功能:计算 KDA (Kronecker Delta Attention) 注意力机制的门控激活

算子类型:Element-wise 激活函数,Memory Bound


🔢 输入输出分析

输入张量

参数 形状 数据类型 说明
g [..., H, D] bfloat16/float16/float32 门控输入,展平后为 [T, H, D]
A_log [H] float32 每个头的对数衰减参数
dt_bias [H * D] (可选) float32 时间步偏置
output_dtype - - 输出数据类型

维度说明:

  • T = numel / (H * D): 展平后的序列长度
  • H: 注意力头数 (Heads)
  • D: 每个头的特征维度 (Dimension)

输出张量

参数 形状 数据类型 说明
yg [..., H, D] output_dtype 门控激活后的输出

核心计算逻辑

数学公式

yg = -exp(A_log[h]) * softplus(g + dt_bias)

其中:

  • softplus(x) = log(1 + exp(x)),数值稳定版本为:
    softplus(x) = x if x > 20 else log(1 + exp(x))
    
  • A_log[h] 对每个头进行广播
  • dt_bias reshape 为 [H, D] 后广播到 [T, H, D]

计算流程

输入 g [T, H, D]
    ↓
加偏置: g_biased = g + dt_bias.view(H, D)
    ↓
激活: g_activated = softplus(g_biased)
    ↓
门控: yg = -exp(A_log).view(H, 1) * g_activated
    ↓
输出 yg [T, H, D]

算子特征

  1. 纯元素级操作 (Element-wise)

    • 无矩阵乘法、无原子操作
    • 无序列维度依赖(与累积和无关)
    • 可完全并行化
  2. Memory Bound

    • 计算量:每元素 3-4 次浮点运算(exp + log + mul)
    • 内存访问:读取 g、A_log、dt_bias,写入 yg
    • 计算强度低,受内存带宽限制
  3. 数值稳定性

    • Softplus 在 x > 20 时避免 exp 溢出
    • 内部使用 float32 计算,输出转换为指定 dtype

迁移实现步骤

1. 理解原始实现

GPU 源码位置:

  • API 层:fla/ops/kda/gate.py
  • Kernel:同文件中的 kda_gate_fwd_kernel

关键特性:

  • 使用 @triton.autotune 自动调优 BT(T 块大小)
  • 36 种配置组合(BT × num_warps × num_stages)
  • 支持 AMD/NVIDIA 不同架构

2. NPU 适配要点

a. 移除 autotune,手动分档

GPU 代码:

@triton.autotune(
    configs=[
        triton.Config({"BT": BT}, num_warps=num_warps, num_stages=num_stages)
        for BT in [32, 64, 128]
        for num_warps in [2, 4, 8, 16]
        for num_stages in [2, 3]
    ],
    key=["H", "D"],
)

NPU 代码:

def _pick_bt(t: int, bd: int) -> int:
    """手动选择 T 块大小"""
    budget = 8192  # UB 预算(fp32 元素数)
    if t <= 32:
        bt = 32
    elif t <= 128:
        bt = 64
    else:
        bt = 128
    # 根据 D 维度调整,避免 UB 溢出
    cap = max(1, budget // bd)
    cap_pow2 = 1 << (cap.bit_length() - 1)
    return max(1, min(bt, cap_pow2))

优势:

  • 推理时无 autotune 预热开销
  • 根据 NPU UB (Unified Buffer) 限制优化

b. 处理 grid 维度限制

问题:NPU grid 维度上限 65535,长序列会超限

解决方案:分段 launch

_MAX_GRID_DIM = 65535

for t_offset in range(0, n_t_blocks, _MAX_GRID_DIM):
    blocks_this_launch = min(_MAX_GRID_DIM, n_t_blocks - t_offset)
    grid = (blocks_this_launch, h)
    kda_gate_fwd_kernel[grid](
        ...,
        T_OFFSET=t_offset,  # 偏移量传入 kernel
    )

Kernel 中使用偏移:

i_t = tl.program_id(0) + T_OFFSET

c. Softplus 数值稳定实现

内联实现(FLA 的 softplus helper 在 NPU 不可用):

@triton.jit
def _softplus(x):
    # x > 20 时 log(1 + exp(x)) ≈ x,避免溢出
    return tl.where(x > 20.0, x, tl.log(1.0 + tl.exp(x)))

3. Kernel 核心实现

完整 Kernel:

@triton.jit
def kda_gate_fwd_kernel(
    g,           # [T, H, D] input
    A_log,       # [H] per-head log decay
    dt_bias,     # [H * D] optional bias
    yg,          # [T, H, D] output
    T,           # sequence length
    T_OFFSET,    # grid segmentation offset
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,  # T block size
    BD: tl.constexpr,  # D block size (power of 2)
    HAS_BIAS: tl.constexpr,
):
    # 计算当前线程块处理的位置
    i_t = tl.program_id(0) + T_OFFSET  # T 块索引
    i_h = tl.program_id(1)              # Head 索引
    
    # 加载 per-head 参数
    b_A = tl.load(A_log + i_h).to(tl.float32)
    
    # 构造块指针(Triton block pointer API)
    p_g = tl.make_block_ptr(
        g + i_h * D,        # 基地址:第 i_h 个头的起始位置
        (T, D),             # 形状
        (H * D, 1),         # 步长:行跨 H*D,列跨 1
        (i_t * BT, 0),      # 块偏移
        (BT, BD),           # 块大小
        (1, 0)              # 块顺序
    )
    p_yg = tl.make_block_ptr(yg + i_h * D, (T, D), (H * D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
    
    # 加载输入块 [BT, BD]
    b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32)
    
    # 加偏置(如果有)
    if HAS_BIAS:
        p_b = tl.make_block_ptr(dt_bias, (H * D,), (1,), (i_h * D,), (BD,), (0,))
        b_g = b_g + tl.load(p_b, boundary_check=(0,)).to(tl.float32)
    
    # 计算门控输出
    b_yg = -tl.exp(b_A) * _softplus(b_g)
    
    # 写回结果
    tl.store(p_yg, b_yg.to(p_yg.dtype.element_ty), boundary_check=(0, 1))

关键点:

  1. Block pointer API:简化多维内存访问
  2. Boundary check:处理 T、D 维度不整除块大小的情况
  3. Float32 计算:保证精度,结果转换为 output_dtype
  4. Grid 分段:通过 T_OFFSET 支持超长序列

4. Python 封装

Launcher 函数:

def kda_gate_fwd(
    g: torch.Tensor,
    A_log: torch.Tensor,
    dt_bias: torch.Tensor | None = None,
    output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
    """启动 KDA gate kernel"""
    h, d = g.shape[-2:]
    t = g.numel() // (h * d)
    
    # 分配输出
    yg = torch.empty_like(g, dtype=output_dtype)
    
    # 选择块大小
    bd = triton.next_power_of_2(d)
    bt = _pick_bt(t, bd)
    n_t_blocks = triton.cdiv(t, bt)
    
    has_bias = dt_bias is not None
    
    # 分段启动(避免 grid 维度超限)
    for t_offset in range(0, n_t_blocks, _MAX_GRID_DIM):
        blocks_this_launch = min(_MAX_GRID_DIM, n_t_blocks - t_offset)
        grid = (blocks_this_launch, h)
        kda_gate_fwd_kernel[grid](
            g=g,
            A_log=A_log,
            dt_bias=dt_bias,
            yg=yg,
            T=t,
            T_OFFSET=t_offset,
            H=h,
            D=d,
            BT=bt,
            BD=bd,
            HAS_BIAS=has_bias,
        )
    return yg

5. API 层封装

用户接口 (mindspeed_ops/api/triton/kda_gate.py):

def fused_kda_gate(
    g: torch.Tensor,
    A_log: torch.Tensor,
    dt_bias: torch.Tensor | None = None,
    output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
    """KDA gate 算子用户 API"""
    # 参数验证
    assert g.ndim >= 2
    h, d = g.shape[-2:]
    assert A_log.shape == (h,)
    if dt_bias is not None:
        assert dt_bias.numel() == h * d
    
    # 根据架构选择实现
    from mindspeed_ops.utils import is_arch35
    if is_arch35():
        from mindspeed_ops.arch35.triton.kda.gate import kda_gate_fwd
    else:
        from mindspeed_ops.arch32.triton.kda.gate import kda_gate_fwd
    
    return kda_gate_fwd(g, A_log, dt_bias, output_dtype)

6. 参考实现

Naive PyTorch 实现(用于精度验证):

def naive_kda_gate(
    g: torch.Tensor,
    A_log: torch.Tensor,
    dt_bias: torch.Tensor | None = None,
    output_dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
    """PyTorch 参考实现(CPU/GPU 通用)"""
    h, _ = g.shape[-2:]
    g = g.float()
    
    # 加偏置
    if dt_bias is not None:
        g = g + dt_bias.view(h, -1)
    
    # 计算:-exp(A_log) * softplus(g)
    g = (-A_log.view(h, 1).float().exp() * F.softplus(g.float())).to(output_dtype)
    
    return g

测试与验证

精度测试

测试配置:

# 13 种 shape × 2 种 bias 场景 = 26 个测试用例
SHAPES = [
    (128, 12, 96),
    (1024, 8, 128),
    (2048, 16, 128),
    (4096, 16, 256),
    ...
]
DTYPES = [torch.float32, torch.float16, torch.bfloat16]

测试方法:

def test_kda_gate_npu(t, h, k, dtype, with_bias):
    device = 'npu'
    
    # 生成测试数据
    g = torch.randn(t, h, k, device=device, dtype=dtype)
    A_log = torch.randn(h, device=device, dtype=torch.float32)
    dt_bias = torch.randn(h * k, device=device) if with_bias else None
    
    # 参考实现(float32 CPU)
    g_cpu = g.float().cpu()
    A_log_cpu = A_log.cpu()
    dt_bias_cpu = dt_bias.cpu() if dt_bias is not None else None
    ref = naive_kda_gate(g_cpu, A_log_cpu, dt_bias_cpu, dtype).to(device)
    
    # Triton 实现(NPU)
    out = fused_kda_gate(g, A_log, dt_bias, dtype)
    
    # 精度验证
    abs_diff = torch.abs(ref - out)
    max_abs = abs_diff.max().item()
    
    assert max_abs < 1e-3, f"Max error: {max_abs}"

运行测试:

pytest tests/unit_tests/triton/test_kda_gate.py -v

精度测试全部通过

屏幕截图 2026-08-30 195345.png

======================== test session starts ========================
platform linux -- Python 3.10.21, pytest-8.3.2, pluggy-1.6.0 -- /root/miniconda3/envs/mindspeed_test/bin/python3.10
cachedir: .pytest_cache
rootdir: /home/MindSpeed-Ops
configfile: pyproject.toml
plugins: xdist-3.6.1
collected 26 items                                                

tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T1-H4-K64-float32] PASSED [  3%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T1-H4-K64-float16] PASSED [  7%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T1-H4-K64-bfloat16] PASSED [ 11%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T16-H8-K128-float32] PASSED [ 15%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T16-H8-K128-bfloat16] PASSED [ 19%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T64-H16-K64-float32] PASSED [ 23%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T64-H16-K64-float16] PASSED [ 26%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T128-H12-K96-float32] PASSED [ 30%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T128-H12-K96-bfloat16] PASSED [ 34%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T1024-H8-K128-float32] PASSED [ 38%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T1024-H8-K128-bfloat16] PASSED [ 42%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T4096-H16-K256-float32] PASSED [ 46%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[nobias-T4096-H16-K256-bfloat16] PASSED [ 50%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T1-H4-K64-float32] PASSED [ 53%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T1-H4-K64-float16] PASSED [ 57%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T1-H4-K64-bfloat16] PASSED [ 61%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T16-H8-K128-float32] PASSED [ 65%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T16-H8-K128-bfloat16] PASSED [ 69%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T64-H16-K64-float32] PASSED [ 73%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T64-H16-K64-float16] PASSED [ 76%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T128-H12-K96-float32] PASSED [ 80%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T128-H12-K96-bfloat16] PASSED [ 84%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T1024-H8-K128-float32] PASSED [ 88%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T1024-H8-K128-bfloat16] PASSED [ 92%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T4096-H16-K256-float32] PASSED [ 96%]
tests/unit_tests/triton/test_kda_gate.py::TestKdaGateOperator::test_kda_gate_npu[bias-T4096-H16-K256-bfloat16] PASSED [100%]

======================== 26 passed in 51.79s ========================
(mindspeed_test) [root@1aad0a108927 MindSpeed-Ops]# 

性能测试

NPU 测试方法

工具:msprof(Ascend 官方性能分析工具)

原理:从 AI Core 硬件计数器直接采集每个 kernel 的执行时间

命令:

msprof --output=./results \
       --application="python test.py" \
       --ai-core=on --task-time=on

结果:task_time.csv 记录每个 kernel 的执行时间(微秒级)

GPU 测试方法

工具:CUDA Events(NVIDIA 官方计时 API)

原理:在 GPU 流中插入硬件时间戳,测量操作执行时间

代码:

start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)

start.record()
kernel_call()
end.record()
torch.cuda.synchronize()

time_ms = start.elapsed_time(end)

对比说明

两种方法都是硬件级测量,排除了 Python 和框架开销,可直接对比 kernel 性能。

测试配置:T=2048, H=16, K=128, dtype=bfloat16, warmup=10, repeat=100

测试结果如下

gpu侧(A100)

屏幕截图 2026-08-30 215153.png

npu侧(910B)

屏幕截图 2026-08-30 213933.png

likedislike
咋哈
咋哈成员
8月30日 评论:

@LinSHua
您好,辛苦您帮忙检视一下这个方案是否达到标准,确认是否可以通过。如果没问题的话,我这边就准备推送代码并提 PR 了。期待您的反馈,谢谢!

likedislike
咋哈
咋哈成员
28 天前 评论:

@Ronald1995
您好,辛苦您帮忙检视一下这个方案是否达到标准,确认是否可以通过。如果没问题的话,我这边就准备推送代码并提 PR 了。期待您的反馈,谢谢!

likedislike
LinShua成员
22 天前 评论:

@meleys

您好,该方案性能达到标准,精度方面麻烦补充下参考实现(PyTorch)与GPU的精度对比情况,以便验证参考实现(PyTorch)的准确性,即可提交PR检视合入。

likedislike
咋哈
咋哈成员
19 天前 评论:

GPU 精度验证

为验证 KDA Gate Triton Kernel 与 PyTorch 参考实现之间的计算一致性,在 GPU 环境下进行精度测试。测试以 PyTorch 实现作为 Reference,使用 FP32 进行中间计算,并与 Triton Kernel 的输出结果进行逐元素比较。

测试方法

首先构造不同规模的输入张量 g,其形状为 [T, H, D],同时生成 A_log 以及可选的 dt_bias。Triton Kernel 按照与实际 KDA Gate 算子一致的计算流程执行:

g′=g+dt_biasg' = g + dt\_bias

yg=−exp⁡(A_log)×Softplus(g′)yg=-\exp(A\_log)\times Softplus(g')

其中,当不使用 dt_bias 时直接对 g 进行 Softplus 计算。

PyTorch Reference 同样采用 FP32 进行中间计算,完成计算后再转换为测试输出类型。通过比较 Triton Kernel 输出与 PyTorch Reference 输出,评估 Kernel 的数值一致性。

测试流程如下:

  1. 根据测试参数生成输入 g、A_log 和 dt_bias;
  2. 分别执行 PyTorch Reference 和 Triton Kernel;
  3. 将两侧输出转换为 FP32;
  4. 计算最大绝对误差、平均绝对误差和最大相对误差;
  5. 使用 torch.allclose 进行最终精度判定。

测试范围

测试覆盖不同的序列长度 T、Head 数 H、特征维度 D、数据类型以及是否使用 dt_bias 等场景。

测试维度 测试范围
T 1、16、31、32、33、63、64、65、127、128、129、512、1024、2048、4096
H 4、8、12、16、32
D 64、96、128、256
数据类型 FP32、FP16、BF16
dt_bias 有、无
BF16 + bias 暂不测试

其中,31/32/33、63/64/65、127/128/129 等测试用例用于覆盖不同 Block Size 边界附近的场景,验证 Kernel 在尾块以及非整块输入情况下的正确性。

在最终测试配置下,共执行 75 个测试 Case,即 15 组输入规模 × 5 种有效的 dtype/bias 组合。

精度判定标准

采用 torch.allclose 进行精度判断:

∣out−ref∣≤atol+rtol×∣ref∣|out-ref|\leq atol+rtol\times |ref|

测试阈值设置为:

  • atol = 1e-3
  • rtol = 1e-3

同时记录以下误差指标:

  • Max Absolute Error:所有元素中的最大绝对误差;
  • Mean Absolute Error:所有元素的平均绝对误差;
  • Max Relative Error:所有元素中的最大相对误差。

其中最大相对误差计算时对 Reference 的绝对值设置最小值 1e-8,避免 Reference 接近 0 时产生异常大的相对误差。

测试结果

GPU 精度测试共覆盖 75 个 Case,包括 FP32、FP16 以及 BF16 无 Bias 场景,并覆盖有无 dt_bias、不同 T/H/D 规模以及 Block 边界等情况。

测试过程中 Triton Kernel 输出与 PyTorch Reference 保持数值一致性,所有 Case 均满足:

torch.allclose(out,ref,atol=1e−3,rtol=1e−3)torch.allclose(out, ref, atol=1e-3, rtol=1e-3)

GPU Triton Kernel 精度测试全部通过,验证了 KDA Gate Kernel 核心计算逻辑及不同输入规模下的数值正确性。

屏幕截图 2026-09-11 171315.png
屏幕截图 2026-09-11 194231.png
屏幕截图 2026-09-11 194243.png
屏幕截图 2026-09-11 194252.png

likedislike
咋哈
咋哈成员
19 天前 评论:

@LinShua
您好,我已经补充参考实现(PyTorch)与GPU的精度对比情况,以验证参考实现(PyTorch)的准确性

likedislike
LinShua成员
16 天前 评论:

好的,那精度和性能都达到要求,可以提交相关PR并关联该issue,后续进行检视合入即可

likedislike
Hhhhzhuyizhi成员
15 天前 issue类型由 Bug-Report 改变为 任务
梵高的呐喊
梵高的呐喊
13 天前 评论:

认领该任务

likedislike
Xxmz成员
7 天前 关联了里程碑:MindSpeed 26.3.0
此处折叠了5条事件消息 查看更多
咋哈咋哈成员
2 天前 关联了pull request:feat:triton算子(kda_gate_fwd_kernel)迁移
咋哈
咋哈成员
1 天前 评论:

您好,我已经提交相关PR并关联该issue,请帮忙检视

likedislike