| name | triton-ascend-case-elemwise-cast |
| description | 大shape类型转换(int8→fp16)优化:通过二次切分(BLOCK_SIZE+TILE_SIZE)提高UB利用率,核数2048时性能最优,适用于shape较大(百万级元素)的elementwise类型转换场景 |
| category | case |
| version | 1.0.0 |
| metadata | {"backend":"ascend","dsl":"triton_ascend","hardware":"Atlas A2, Atlas A3"} |
Int8 到 FP16 类型转换优化案例
任务特征
- 操作类型:Elementwise,类型转换操作
- 数据尺寸:(128, 1024, 1024),shape较大
- 数据类型:输入int8,输出fp16
- 任务特点:可以按照轴的顺序(可flatten为一根轴),外层并行,内层向量化,若UB存不下,可考虑多次切分
优化:二次切分 + 用满UB
configs = [
triton.Config({"BLOCK_SIZE": 65536, "TILE_SIZE": 65536}),
triton.Config({"BLOCK_SIZE": 65536, "TILE_SIZE": 32768}),
triton.Config({"BLOCK_SIZE": 2097152, "TILE_SIZE": 65536}),
triton.Config({"BLOCK_SIZE": 4194304, "TILE_SIZE": 65536}),
]
block_start = pid * BLOCK_SIZE
for i in range(0, BLOCK_SIZE, TILE_SIZE):
offsets = block_start + tl.arange(0, TILE_SIZE)
mask = offsets < n_elements
input_data = tl.load(input_ptr + offsets, mask=mask)
output_data = tl.cast(input_data, tl.float16)
tl.store(output_ptr + offsets, output_data, mask=mask)
优化内容
- triton 内核部分使用for循环,尝试进行二次切分,每次搬运TILE_SIZE大小的数据,提高UB的利用率
- 在一定范围内提高核数,并尝试用满UB
- 核内没有二次切分时性能最优(BLOCK_SIZE = TILE_SIZE = 65536)
总结
- 当数据的shape较大时,为了获得更佳的性能,切分值设置尽量能被shape的大小整除
- 对于单纯的Elementwise操作,将多根轴的元素展开为一根轴,然后在这根轴上进行切分
- 将block分配给每个线程块,若UB存不下,可考虑多次切分(二次切分)