| name | triton-ascend-case-reduction-weighted-swiglu |
| description | 3D融合算子(Weighted SwiGLU Backward)优化:Reshape降维将前两维合并简化并行策略,行二次切分避免超UB,在优先占满UB前提下为reduce轴分配较大切分尺寸,grid数较大时可能性能更优,适用于3D张量逐元素+reduce融合的场景 |
| category | case |
| version | 1.0.0 |
| metadata | {"backend":"ascend","dsl":"triton_ascend","hardware":"Atlas A2, Atlas A3"} |
Weighted SwiGLU Backward 融合算子优化
任务特征
- 数据尺寸:(16, 1024, 2048) × 3,3D融合算子
- 特点:先逐元素操作,再reduce最后一根轴
优化 1:Reshape 降维
x_reshaped = x.reshape(BM, N)
weight_reshaped = weight.reshape(BM, N)
grad_reshaped = grad.reshape(BM, N)
weighted_x_reshaped = weighted_x.reshape(BM, N)
grad_weight_reshaped = grad_weight.reshape(BM, N)
grad_x_reshaped = grad_x.reshape(BM)
优势:简化并行策略,优化内存访问模式,提高内核执行效率。
优化 2:行二次切分
for bm_start in range(0, BLOCK_SIZE_BM, SUB_BLOCK_SIZE_BM):
bm_offsets = pid * BLOCK_SIZE_BM + bm_start + tl.arange(0, SUB_BLOCK_SIZE_BM)
bm_mask = bm_offsets < BM
Autotune 配置
triton.Config({'BLOCK_SIZE_BM': 32, 'SUB_BLOCK_SIZE_BM': 32, 'BLOCK_SIZE_N': 128})
triton.Config({'BLOCK_SIZE_BM': 16, 'SUB_BLOCK_SIZE_BM': 16, 'BLOCK_SIZE_N': 256})
triton.Config({'BLOCK_SIZE_BM': 8, 'SUB_BLOCK_SIZE_BM': 8, 'BLOCK_SIZE_N': 512})
triton.Config({'BLOCK_SIZE_BM': 512, 'SUB_BLOCK_SIZE_BM': 8, 'BLOCK_SIZE_N': 512})
triton.Config({'BLOCK_SIZE_BM': 416, 'SUB_BLOCK_SIZE_BM': 8, 'BLOCK_SIZE_N': 512})
总结
- Reshape降维可简化并行策略,优化内存访问
- 在优先占满UB前提下为reduce轴分配较大切分尺寸
- Grid数较大时,可能性能更优(配置3)