triton-ascend实操:mean_kernel算子适配与优化
作者:昇腾实战派
知识地图链接: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=0∑N−1Xi,j
其中:
- N N N 为规约维度的长度
- X i , j X_{i,j} Xi,j 为沿该维度的第 j j j 个元素
1.4 代码执行流程
整体流程如下:
- Host侧函数(
mean_dim):解析输入张量,计算M、N、K,重塑为3D张量,并配置kernel执行参数。 - Kernel启动:设置
grid = (M * K,),每个线程块负责计算一个输出元素。 - 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平台迁移过程中,原始实现暴露出两个关键问题:
- coreDim超限:当
M × K > 65536时,grid = (M * K,)超出NPU物理核数限制,导致启动失败。 - 访存效率低:当K较大时,
input_stride1与input_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内:交换
stride1与stride2,并调整索引计算逻辑。
该优化可使dim=0场景耗时从224ms降至9.4ms,性能提升超20倍。
总结
本文系统梳理了mean_kernel算子从GPU到NPU的迁移路径,涵盖:
- 核心逻辑解析与数学建模
- 核间并行重构以解决coreDim超限
- 自动调优提升
BLOCK_SIZE效率 - 访存优化实现极致性能
最终实现的算子在保证精度与确定性的前提下,具备良好的可扩展性与高性能表现,为复杂算子的NPU适配提供了可复用的工程范式。# Triton NPU算子优化实践:基于访存模式的性能提升
鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。
更多推荐


所有评论(0)