一键导入
triton-ascend-example-softmax
Softmax 归约算子的完整 Triton Ascend 实现示例。展示三阶段归约模式(求 max → 求 sum(exp) → 归一化)、分块累加、标量累加器精度提升等技巧。当生成 reduce 类算子时可参考此示例的代码结构。
用 Codex 或 Claude 帮你安装 复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它检查 Skill 页面并帮你完成安装。
菜单
Softmax 归约算子的完整 Triton Ascend 实现示例。展示三阶段归约模式(求 max → 求 sum(exp) → 归一化)、分块累加、标量累加器精度提升等技巧。当生成 reduce 类算子时可参考此示例的代码结构。
用 Codex 或 Claude 帮你安装 复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它检查 Skill 页面并帮你完成安装。
基于 SOC 职业分类
AscendC direct-invoke 崩溃/挂起修复索引:Kernel timeout、hang、Segmentation Fault、aic error、buffer 死锁、plog/memcheck 调试。
AscendC direct-invoke 精度失败修复索引:输出全 0/随机值、DataCopy 对齐、EnQue/DeQue 同步、FP16/FP32 精度、Cast RoundMode、DumpTensor 分段定位。
AscendC direct-invoke 工程契约:WA 使用 kernel.py + ascendc_op/,ModelNew 调用 torch.ops.npu.*,adapter 负责复制工程、CMake 构建和 npu-arch patch。适用于 dsl=ascendc 的 autoresearch 任务。
把注册式 AscendC 算子迁移为 direct-invoke 工程的保真原则:kernel 算法和 tiling 公式不乱改,只替换注册框架胶水、入口 ABI、host launch 与 PyTorch extension。
CATLASS TileShape 与 on-chip 缓存容量约束:L1/L0A/L0B/L0C 预算公式、fp16/fp32 Pingpong 双缓冲、512B 对齐与排布对 Tile 选型的影响。调参前必读。
CATLASS Gemm 性能调优:DispatchPolicy、Tile 与分核负载均衡、Swizzle、何时用 padding/Split-K/Preload。面向 AR 修改 catlass_kernel.asc 中的类型别名。
| name | triton-ascend-example-softmax |
| description | Softmax 归约算子的完整 Triton Ascend 实现示例。展示三阶段归约模式(求 max → 求 sum(exp) → 归一化)、分块累加、标量累加器精度提升等技巧。当生成 reduce 类算子时可参考此示例的代码结构。 |
| category | example |
| version | 1.0.0 |
| metadata | {"backend":"ascend","dsl":"triton_ascend","hardware":"Atlas A2, Atlas A3","operator_type":"reduce","framework":"torch"} |
import torch
import triton
import triton.language as tl
@triton.jit
def softmax_kernel(
X_ptr, Y_ptr,
B: tl.constexpr, N: tl.constexpr,
stride_xb: tl.constexpr, stride_xn: tl.constexpr,
stride_yb: tl.constexpr, stride_yn: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr, CORE_NUM: tl.constexpr,
):
pid = tl.program_id(0)
for b in range(pid, B, CORE_NUM):
# Phase 1: max
max_val = -float('inf')
for off in range(0, N, BLOCK_SIZE_N):
n_off = off + tl.arange(0, BLOCK_SIZE_N)
mask = n_off < N
x = tl.load(X_ptr + b * stride_xb + n_off * stride_xn,
mask=mask, other=-float('inf'))
max_val = tl.maximum(max_val, tl.max(x, axis=0))
# Phase 2: sum(exp(x - max))
sum_val = 0.0
for off in range(0, N, BLOCK_SIZE_N):
n_off = off + tl.arange(0, BLOCK_SIZE_N)
mask = n_off < N
x = tl.load(X_ptr + b * stride_xb + n_off * stride_xn,
mask=mask, other=0.0)
exp_x = tl.math.exp(x - max_val)
sum_val += tl.sum(exp_x, axis=0).to(tl.float32)
# Phase 3: normalize
for off in range(0, N, BLOCK_SIZE_N):
n_off = off + tl.arange(0, BLOCK_SIZE_N)
mask = n_off < N
x = tl.load(X_ptr + b * stride_xb + n_off * stride_xn,
mask=mask, other=0.0)
result = tl.math.exp(x - max_val) / sum_val
tl.store(Y_ptr + b * stride_yb + n_off * stride_yn,
result, mask=mask)
class ModelNew(torch.nn.Module):
def __init__(self):
super().__init__()
try:
import torch
import triton
device = torch.npu.current_device()
properties = triton.runtime.driver.active.utils.get_device_properties(device)
self.VEC_CORE_NUM = properties.get("num_vectorcore", 40)
except:
self.VEC_CORE_NUM = 40
def forward(self, x):
if not x.is_contiguous():
x = x.contiguous()
B, N = x.shape
y = torch.empty_like(x)
BLOCK_SIZE_N = 4096
grid = (self.VEC_CORE_NUM,)
softmax_kernel[grid](
x, y, B, N,
x.stride(0), x.stride(1), y.stride(0), y.stride(1),
BLOCK_SIZE_N=BLOCK_SIZE_N, CORE_NUM=self.VEC_CORE_NUM)
return y