作者:昇腾实战派
知识地图链接:Triton Ascend知识地图

背景概述

在深度学习模型推理与训练过程中,张量规约操作(如均值计算)是高频且关键的算子之一。随着算子对确定性、高性能和跨平台兼容性要求的提升,基于Triton的自定义算子成为优化核心路径。本文以mean_kernel算子为切入点,系统阐述其在NPU平台上的适配流程与性能优化策略,涵盖从原始GPU代码迁移、核心问题定位到多维度调优的完整实践,为开发者提供可复用的技术范式。


1. mean_kernel源码解析

1.1 完整代码(run_triton.py)

本文所用代码源自开源项目中的批次不变算子实现,经简化与注释后用于演示。核心功能为在指定维度上计算张量均值,等价于PyTorch的torch.mean(input, dim=dim, keepdim=keepdim)

import torch
import triton
import triton.language as tl

@triton.jit
def mean_kernel(
    input_ptr,
    output_ptr,
    input_stride0,
    input_stride1,
    input_stride2,
    output_stride0,
    output_stride1,
    M,  # size before reduction dim
    N,  # size of reduction dim
    K,  # size after reduction dim
    BLOCK_SIZE: tl.constexpr,
):
    """
    Kernel for computing mean along a single dimension.
    Input is viewed as (M, N, K) where N is the dimension being reduced.
    """
    pid = tl.program_id(0)

    m_idx = pid // K
    k_idx = pid % K

    if m_idx >= M or k_idx >= K:
        return

    acc = 0.0
    for n_start in range(0, N, BLOCK_SIZE):
        n_offsets = n_start + tl.arange(0, BLOCK_SIZE)
        mask = n_offsets < N

        input_idx = (
            m_idx * input_stride0 + n_offsets * input_stride1 + k_idx * input_stride2
        )

        vals = tl.load(input_ptr + input_idx, mask=mask, other=0.0)
        acc += tl.sum(vals)

    mean_val = acc / N
    output_idx = m_idx * output_stride0 + k_idx * output_stride1
    tl.store(output_ptr + output_idx, mean_val)

def mean_dim(
    input: torch.Tensor,
    dim: int,
    keepdim: bool = False,
    dtype: torch.dtype | None = None,
) -> torch.Tensor:
    assert input.is_npu, "Input must be a npu tensor"
    assert (-input.ndim <= dim < input.ndim), f"Invalid dimension {dim}"

    if dim < 0:
        dim = dim + input.ndim

    if dtype is None:
        if input.dtype in [torch.int8, torch.int16, torch.int32, torch.int64]:
            dtype = torch.float32
        else:
            dtype = input.dtype

    if input.dtype != dtype:
        input = input.to(dtype)

    shape = list(input.shape)
    M = 1
    for i in range(dim):
        M *= shape[i]
    N = shape[dim]
    K = 1
    for i in range(dim + 1, len(shape)):
        K *= shape[i]

    input_3d = input.reshape(M, N, K)
    if keepdim:
        output_shape = shape.copy()
        output_shape[dim] = 1
    else:
        output_shape = shape[:dim] + shape[dim + 1:]

    output = torch.empty(output_shape, dtype=dtype, device=input.device)
    if keepdim:
        output_2d = output.reshape(M, 1, K).squeeze(1)
    else:
        output_2d = output.reshape(M, K)

    grid = (M * K,)
    BLOCK_SIZE = 1024

    mean_kernel[grid](
        input_3d,
        output_2d,
        input_3d.stride(0),
        input_3d.stride(1),
        input_3d.stride(2),
        output_2d.stride(0),
        output_2d.stride(1) if output_2d.ndim > 1 else 0,
        M,
        N,
        K,
        BLOCK_SIZE,
    )

    return output

def mean_batch_invariant(input, dim, keepdim=False, dtype: torch.dtype | None = None):
    assert dtype is None or dtype == torch.float32, f"unsupported dtype: {dtype}"
    if len(dim) == 1:
        return mean_dim(input, dim[0], keepdim=keepdim)
    else:
        assert input.dtype in {torch.float16, torch.bfloat16, torch.float32}, "only float types supported"
        n_elems = 1
        for d in dim:
            n_elems *= input.shape[d]
        return torch.sum(input, dim=dim, keepdim=keepdim, dtype=torch.float32) / n_elems

def test_mean_batch_invariant():
    torch.manual_seed(42)
    test_cases = [
        (1024,),
        (512, 512),
        (2048, 4096, 64),
    ]
    atol = 1e-3
    rtol = 1e-3

    for shape in test_cases:
        input_tensor = torch.randn(*shape, dtype=torch.float32, device='npu:0')
        input_dim = input_tensor.ndim

        for dim in range(input_dim):
            output_torch = torch.mean(input_tensor, dim=dim, keepdim=True)
            output_triton = mean_batch_invariant(input_tensor, dim=[dim], keepdim=True)
            assert torch.allclose(output_torch, output_triton, atol=atol, rtol=rtol), \
                f"Test failed for shape {shape}, dim {dim}"

            t_torch = triton.testing.do_bench(lambda: torch.mean(input_tensor, dim=dim, keepdim=True))
            t_triton = triton.testing.do_bench(lambda: mean_batch_invariant(input_tensor, dim=[dim], keepdim=True))
            print(f"Dim {dim}: torch={t_torch:.2f}ms, triton={t_triton:.2f}ms, ratio={t_triton/t_torch:.2f}x")

        print(f"Test passed for shape {shape}.")

if __name__ == "__main__":
    test_mean_batch_invariant()

1.2 功能说明

该算子实现任意维度张量的均值计算,支持keepdim选项,并具备批次不变性(Batch-Invariant),确保在不同执行环境下结果一致。其核心思想是将多维张量统一重构成三维形式 (M, N, K),其中:

  • M:规约维度前所有维度的乘积(前置维度)
  • N:待规约的维度(即均值计算维度)
  • K:规约维度后所有维度的乘积(后置维度)

通过此映射,可将任意维度的均值问题转化为统一的三维规约问题。

1.3 数学公式

对于张量在维度 dim 上的均值计算,其数学表达为:

mean i = 1 N ∑ j = 0 N − 1 X i , j \text{mean}_i = \frac{1}{N} \sum_{j=0}^{N-1} X_{i,j} meani=N1j=0N1Xi,j

其中:

  • N N N 为规约维度的长度
  • X i , j X_{i,j} Xi,j 为沿该维度的第 j j j 个元素

1.4 代码执行流程

整体流程如下:

  1. Host侧函数mean_dim):解析输入张量,计算M、N、K,重塑为3D张量,并配置kernel执行参数。
  2. Kernel启动:设置grid = (M * K,),每个线程块负责计算一个输出元素。
  3. Kernel执行mean_kernel):每个线程块独立完成对N维度的累加与均值计算。

1.5 核间并行策略

  • Grid维度M × K,每个线程块处理一个输出元素。
  • 并行粒度:按输出元素划分,实现核间并行。
  • 任务分配:线程块ID(pid)映射为 (m_idx, k_idx),对应输出位置。

1.6 核内并行策略

  • 分块加载:沿N维度以BLOCK_SIZE为单位分块,避免UB(Unified Buffer)溢出。
  • 边界处理:使用mask机制处理非对齐尾部数据。
  • 索引计算:利用张量的stride信息,精准计算内存偏移。

2. 原始实现存在的问题

在NPU平台迁移过程中,原始实现暴露出两个关键问题:

  1. coreDim超限:当M × K > 65536时,grid = (M * K,)超出NPU物理核数限制,导致启动失败。
  2. 访存效率低:当K较大时,input_stride1input_stride2差异显著,导致内存访问不连续,影响带宽利用率。

3. NPU适配与性能优化

3.1 coreDim超限问题解决

针对M × K过大导致的coreDim超限问题,采用分核处理策略:

  • Grid设置:改为 (num_core,),其中num_core为NPU物理向量核数量。
  • 任务分摊:每个核处理num_mean = (M * K + num_core - 1) // num_core个输出元素。
  • 外层循环:在kernel中引入for output_idx_ in range(start_pid, end_pid),实现核间任务分发。

修改后的kernel代码如下:

@triton.jit
def mean_kernel(
    input_ptr,
    output_ptr,
    input_stride0,
    input_stride1,
    input_stride2,
    output_stride0,
    output_stride1,
    M,
    N,
    K,
    num_mean,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    start_pid = pid * num_mean
    end_pid = tl.minimum((pid + 1) * num_mean, M * K)

    for output_idx_ in range(start_pid, end_pid):
        m_idx = output_idx_ // K
        k_idx = output_idx_ % K

        acc = 0.0
        for n_start in range(0, N, BLOCK_SIZE):
            n_offsets = n_start + tl.arange(0, BLOCK_SIZE)
            mask = n_offsets < N

            input_idx = m_idx * input_stride0 + n_offsets * input_stride1 + k_idx * input_stride2
            vals = tl.load(input_ptr + input_idx, mask=mask, other=0.0)
            acc += tl.sum(vals)

        mean_val = acc / N
        output_idx = m_idx * output_stride0 + k_idx * output_stride1
        tl.store(output_ptr + output_idx, mean_val)

3.2 完整适配代码(run_triton_npu.py)

import torch
import triton
import triton.language as tl
import triton.runtime.driver as driver

@triton.jit
def mean_kernel(
    input_ptr,
    output_ptr,
    input_stride0,
    input_stride1,
    input_stride2,
    output_stride0,
    output_stride1,
    M,
    N,
    K,
    num_mean,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    start_pid = pid * num_mean
    end_pid = tl.minimum((pid + 1) * num_mean, M * K)

    for output_idx_ in range(start_pid, end_pid):
        m_idx = output_idx_ // K
        k_idx = output_idx_ % K

        acc = 0.0
        for n_start in range(0, N, BLOCK_SIZE):
            n_offsets = n_start + tl.arange(0, BLOCK_SIZE)
            mask = n_offsets < N

            input_idx = m_idx * input_stride0 + n_offsets * input_stride1 + k_idx * input_stride2
            vals = tl.load(input_ptr + input_idx, mask=mask, other=0.0)
            acc += tl.sum(vals)

        mean_val = acc / N
        output_idx = m_idx * output_stride0 + k_idx * output_stride1
        tl.store(output_ptr + output_idx, mean_val)

def mean_dim(
    input: torch.Tensor,
    dim: int,
    keepdim: bool = False,
    dtype: torch.dtype | None = None,
) -> torch.Tensor:
    assert input.is_npu, "Input must be a npu tensor"
    assert (-input.ndim <= dim < input.ndim), f"Invalid dimension {dim}"

    if dim < 0:
        dim = dim + input.ndim

    if dtype is None:
        if input.dtype in [torch.int8, torch.int16, torch.int32, torch.int64]:
            dtype = torch.float32
        else:
            dtype = input.dtype

    if input.dtype != dtype:
        input = input.to(dtype)

    shape = list(input.shape)
    M = 1
    for i in range(dim):
        M *= shape[i]
    N = shape[dim]
    K = 1
    for i in range(dim + 1, len(shape)):
        K *= shape[i]

    input_3d = input.reshape(M, N, K)
    if keepdim:
        output_shape = shape.copy()
        output_shape[dim] = 1
    else:
        output_shape = shape[:dim] + shape[dim + 1:]

    output = torch.empty(output_shape, dtype=dtype, device=input.device)
    if keepdim:
        output_2d = output.reshape(M, 1, K).squeeze(1)
    else:
        output_2d = output.reshape(M, K)

    num_core = get_npu_properties()["num_vectorcore"]
    grid = (num_core,)
    num_mean = (M * K + num_core - 1) // num_core
    BLOCK_SIZE = 2048

    mean_kernel[grid](
        input_3d,
        output_2d,
        input_3d.stride(0),
        input_3d.stride(1),
        input_3d.stride(2),
        output_2d.stride(0),
        output_2d.stride(1) if output_2d.ndim > 1 else 0,
        M,
        N,
        K,
        num_mean,
        BLOCK_SIZE,
    )

    return output

def get_npu_properties():
    device = torch.npu.current_device()
    return driver.active.utils.get_device_properties(device)

def mean_batch_invariant(input, dim, keepdim=False, dtype: torch.dtype | None = None):
    assert dtype is None or dtype == torch.float32, f"unsupported dtype: {dtype}"
    if len(dim) == 1:
        return mean_dim(input, dim[0], keepdim=keepdim)
    else:
        assert input.dtype in {torch.float16, torch.bfloat16, torch.float32}, "only float types supported"
        n_elems = 1
        for d in dim:
            n_elems *= input.shape[d]
        return torch.sum(input, dim=dim, keepdim=keepdim, dtype=torch.float32) / n_elems

def test_mean_batch_invariant():
    torch.manual_seed(42)
    test_cases = [(2048, 4096, 64)]
    atol = 1e-3
    rtol = 1e-3

    for shape in test_cases:
        input_tensor = torch.randn(*shape, dtype=torch.float32, device='npu:0')
        input_dim = input_tensor.ndim

        for dim in range(input_dim):
            output_torch = torch.mean(input_tensor, dim=dim, keepdim=True)
            output_triton = mean_batch_invariant(input_tensor, dim=[dim], keepdim=True)
            assert torch.allclose(output_torch, output_triton, atol=atol, rtol=rtol), \
                f"Test failed for shape {shape}, dim {dim}"

            t_torch = triton.testing.do_bench(lambda: torch.mean(input_tensor, dim=dim, keepdim=True))
            t_triton = triton.testing.do_bench(lambda: mean_batch_invariant(input_tensor, dim=[dim], keepdim=True))
            print(f"Dim {dim}: torch={t_torch:.2f}ms, triton={t_triton:.2f}ms, ratio={t_triton/t_torch:.2f}x")

        print(f"Test passed for shape {shape}.")

if __name__ == "__main__":
    test_mean_batch_invariant()

3.3 自动调优(AutoTune)优化

为提升性能,引入Triton自动调优机制,对BLOCK_SIZE进行最优搜索。

  • 配置:在kernel上方添加@triton.heuristics@triton.jit配置。
  • 参数:仅保留BLOCK_SIZE为可调参数。
  • 启用:设置环境变量TRITON_PRINT_AUTOTUNING=1以查看调优过程。
dim BLOCK_SIZE 原始耗时 AutoTune推荐 优化后耗时
0 1024 224ms 1024 232ms
1 1024 76ms 2048 76ms
2 1024 64ms 128 43ms

结果显示,dim=2场景下,通过调优将性能提升约33%。

3.4 完整优化代码(run_triton_npu_autotune.py)

import torch
import triton
import triton.language as tl
import triton.runtime.driver as driver
import os
os.environ['TRITON_PRINT_AUTOTUNING'] = '1'

@triton.heuristics({
    'BLOCK_SIZE': lambda args: 2 ** triton.next_power_of_2(args[0])
})
@triton.jit
def mean_kernel(
    input_ptr,
    output_ptr,
    input_stride0,
    input_stride1,
    input_stride2,
    output_stride0,
    output_stride1,
    M,
    N,
    K,
    num_mean,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    start_pid = pid * num_mean
    end_pid = tl.minimum((pid + 1) * num_mean, M * K)

    for output_idx_ in range(start_pid, end_pid):
        m_idx = output_idx_ // K
        k_idx = output_idx_ % K

        acc = 0.0
        for n_start in range(0, N, BLOCK_SIZE):
            n_offsets = n_start + tl.arange(0, BLOCK_SIZE)
            mask = n_offsets < N

            input_idx = m_idx * input_stride0 + n_offsets * input_stride1 + k_idx * input_stride2
            vals = tl.load(input_ptr + input_idx, mask=mask, other=0.0)
            acc += tl.sum(vals)

        mean_val = acc / N
        output_idx = m_idx * output_stride0 + k_idx * output_stride1
        tl.store(output_ptr + output_idx, mean_val)

# ...(其余代码同上,省略)

4. 访存优化建议(进阶)

针对shape=(2048, 4096, 64)dim=0场景,原始实现中K=4096×64=262144,导致input_stride1极大,访问不连续。

优化方案:将张量重塑为(M, K, N),使规约维度位于尾部,提升数据局部性。

  • Host侧input_3d = input.reshape(M, K, N)
  • Kernel内:交换stride1stride2,并调整索引计算逻辑。

该优化可使dim=0场景耗时从224ms降至9.4ms,性能提升超20倍。


总结

本文系统梳理了mean_kernel算子从GPU到NPU的迁移路径,涵盖:

  • 核心逻辑解析与数学建模
  • 核间并行重构以解决coreDim超限
  • 自动调优提升BLOCK_SIZE效率
  • 访存优化实现极致性能

最终实现的算子在保证精度与确定性的前提下,具备良好的可扩展性与高性能表现,为复杂算子的NPU适配提供了可复用的工程范式。# Triton NPU算子优化实践:基于访存模式的性能提升

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐