一键导入
triton-ascend-example-matmul
标准矩阵乘法的完整 Triton Ascend 实现示例。展示 2D 分块(tiling)、K 维循环累加、2D mask 处理、Cube Core 利用等关键模式。当生成 matmul 类算子时可参考此示例的代码结构。
用 Codex 或 Claude 帮你安装 复制这段 Prompt,粘贴到 Codex、Claude 或其他助手里,让它检查 Skill 页面并帮你完成安装。
菜单
标准矩阵乘法的完整 Triton Ascend 实现示例。展示 2D 分块(tiling)、K 维循环累加、2D mask 处理、Cube Core 利用等关键模式。当生成 matmul 类算子时可参考此示例的代码结构。
用 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-matmul |
| description | 标准矩阵乘法的完整 Triton Ascend 实现示例。展示 2D 分块(tiling)、K 维循环累加、2D mask 处理、Cube Core 利用等关键模式。当生成 matmul 类算子时可参考此示例的代码结构。 |
| category | example |
| version | 1.0.0 |
| metadata | {"backend":"ascend","dsl":"triton_ascend","hardware":"Atlas A2, Atlas A3","operator_type":"matmul","framework":"torch"} |
import torch
import triton
import triton.language as tl
@triton.jit
def matmul_kernel(
a_ptr, b_ptr, c_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
CORE_NUM: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_N: tl.constexpr,
):
NUM_BLOCKS_M = tl.cdiv(M, BLOCK_M)
NUM_BLOCKS_N = tl.cdiv(N, BLOCK_N)
NUM_BLOCKS = NUM_BLOCKS_M * NUM_BLOCKS_N
pid = tl.program_id(0)
for block_idx in range(pid, NUM_BLOCKS, CORE_NUM):
bm = block_idx // NUM_BLOCKS_N
bn = block_idx % NUM_BLOCKS_N
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, K, BLOCK_K):
a_off_m = bm * BLOCK_M + tl.arange(0, BLOCK_M)
a_off_k = k + tl.arange(0, BLOCK_K)
a_mask = (a_off_m < M)[:, None] & (a_off_k < K)[None, :]
a = tl.load(a_ptr + a_off_m[:, None] * stride_am
+ a_off_k[None, :] * stride_ak,
mask=a_mask, other=0.0)
b_off_k = k + tl.arange(0, BLOCK_K)
b_off_n = bn * BLOCK_N + tl.arange(0, BLOCK_N)
b_mask = (b_off_k < K)[:, None] & (b_off_n < N)[None, :]
b = tl.load(b_ptr + b_off_k[:, None] * stride_bk
+ b_off_n[None, :] * stride_bn,
mask=b_mask, other=0.0)
acc += tl.dot(a, b)
c_off_m = bm * BLOCK_M + tl.arange(0, BLOCK_M)
c_off_n = bn * BLOCK_N + tl.arange(0, BLOCK_N)
c_mask = (c_off_m < M)[:, None] & (c_off_n < N)[None, :]
tl.store(c_ptr + c_off_m[:, None] * stride_cm
+ c_off_n[None, :] * stride_cn,
acc, mask=c_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.CUBE_CORE_NUM = properties.get("num_aicore", 20)
except:
self.CUBE_CORE_NUM = 20
def forward(self, A, B):
if not A.is_contiguous():
A = A.contiguous()
if not B.is_contiguous():
B = B.contiguous()
M, K = A.shape
_, N = B.shape
C = torch.empty((M, N), dtype=torch.float32, device=A.device)
BLOCK_M, BLOCK_K, BLOCK_N = 128, 256, 128
grid = (self.CUBE_CORE_NUM,)
matmul_kernel[grid](
A, B, C, M, N, K,
A.stride(0), A.stride(1), B.stride(0), B.stride(1),
C.stride(0), C.stride(1),
CORE_NUM=self.CUBE_CORE_NUM,
BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K, BLOCK_N=BLOCK_N)
return C