已合并
feat(op): support (built-in)indirect_store for unstructure store #782
candyhong创建于 2025年11月20日
feat(op): support (built-in)indirect_store for unstructure store #782
已合并
candyhong创建于 2025年11月20日
candyhong成员
2025年11月20日

描述

新增 TT_IndirectStoreOp 内置方言,当且仅当 compileOn91095Flag && forceSimtTemplateFlag && ptrOffsetInfo.isUnstructured() 满足时,离散访存的 tt.store 转换为 tt.indirect_store 调用底层 simt 模板函数。

测试

测试脚本:pytest -sv ascend/examples/pytest_ut/test_ldst.py::test_ldst_indirect_08

  • compileOn91095Flag && forceSimtTemplateFlag 时,Adapter IR 中出现 call @triton_indirect_store 调用,直接下发接口完成间接存储
module {
  func.func private @triton_indirect_load(memref<?xf32>, tensor<8x16xi64>) -> tensor<8x16xf32>
  func.func private @triton_indirect_store(memref<?xf32>, tensor<8x16xi64>, tensor<8x16xf32>)
  func.func @triton_ldst_indirect_08_kernel(%arg0: memref<?xi8>, %arg1: memref<?xi8>, %arg2: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 1 : i32}, %arg3: memref<?xi64> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32}, %arg4: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32}, %arg5: i32 {tt.divisibility = 16 : i32}, %arg6: i32, %arg7: i32, %arg8: i32, %arg9: i32, %arg10: i32, %arg11: i32) attributes {SyncBlockLockArgIdx = 0 : i64, WorkspaceArgIdx = 1 : i64, global_kernel = "local", mix_mode = "aiv", parallel_mode = "mix_simd_simt"} {
    %c8_i32 = arith.constant {Undefined} 8 : i32
    %c32_i32 = arith.constant 32 : i32
    %0 = tensor.empty() : tensor<8x1xi32>
    %1 = linalg.fill ins(%c32_i32 : i32) outs(%0 : tensor<8x1xi32>) -> tensor<8x1xi32>
    %2 = arith.muli %arg9, %c8_i32 {Undefined} : i32
    %3 = tensor.empty() : tensor<8xi32>
    %4 = linalg.generic {indexing_maps = [#map], iterator_types = ["parallel"]} outs(%3 : tensor<8xi32>) {
    ^bb0(%out: i32):
      %18 = linalg.index 0 : index
      %19 = arith.index_cast %18 : index to i32
      linalg.yield %19 : i32
    } -> tensor<8xi32>
    %5 = linalg.fill ins(%2 : i32) outs(%3 : tensor<8xi32>) -> tensor<8xi32>
    %6 = arith.addi %5, %4 {Undefined} : tensor<8xi32>
    %reinterpret_cast = memref.reinterpret_cast %arg3 to offset: [0], sizes: [16], strides: [1] : memref<?xi64> to memref<16xi64, strided<[1]>>
    %alloc = memref.alloc() : memref<16xi64>
    memref.copy %reinterpret_cast, %alloc : memref<16xi64, strided<[1]>> to memref<16xi64>
    %7 = bufferization.to_tensor %alloc restrict writable : memref<16xi64>
    %expanded = tensor.expand_shape %4 [[0, 1]] output_shape [8, 1] : tensor<8xi32> into tensor<8x1xi32>
    %8 = linalg.fill ins(%arg5 : i32) outs(%0 : tensor<8x1xi32>) -> tensor<8x1xi32>
    %9 = arith.muli %expanded, %8 {Undefined} : tensor<8x1xi32>
    %10 = arith.extsi %9 {Undefined} : tensor<8x1xi32> to tensor<8x1xi64>
    %11 = tensor.empty() : tensor<8x16xi64>
    %collapsed = tensor.collapse_shape %10 [[0, 1]] : tensor<8x1xi64> into tensor<8xi64>
    %broadcasted = linalg.broadcast ins(%collapsed : tensor<8xi64>) outs(%11 : tensor<8x16xi64>) dimensions = [1] 
    %broadcasted_0 = linalg.broadcast ins(%7 : tensor<16xi64>) outs(%11 : tensor<8x16xi64>) dimensions = [0] 
    %12 = arith.addi %broadcasted, %broadcasted_0 {Undefined} : tensor<8x16xi64>
    %13 = call @triton_indirect_load(%arg4, %12) : (memref<?xf32>, tensor<8x16xi64>) -> tensor<8x16xf32>
    %14 = math.exp %13 {Undefined} : tensor<8x16xf32>
    %expanded_1 = tensor.expand_shape %6 [[0, 1]] output_shape [8, 1] : tensor<8xi32> into tensor<8x1xi32>
    %15 = arith.muli %expanded_1, %1 {Undefined} : tensor<8x1xi32>
    %16 = arith.extsi %15 {Undefined} : tensor<8x1xi32> to tensor<8x1xi64>
    %collapsed_2 = tensor.collapse_shape %16 [[0, 1]] : tensor<8x1xi64> into tensor<8xi64>
    %broadcasted_3 = linalg.broadcast ins(%collapsed_2 : tensor<8xi64>) outs(%11 : tensor<8x16xi64>) dimensions = [1] 
    %17 = arith.addi %broadcasted_3, %broadcasted_0 {Undefined} : tensor<8x16xi64>
    call @triton_indirect_store(%arg2, %17, %14) : (memref<?xf32>, tensor<8x16xi64>, tensor<8x16xf32> ) -> ()
    return
  }
}
  • 上述条件不满足时,AdapterIR 仍然使用原有的间接访存实现,通过 scf.for 循环嵌套和 tensor.extract 完成数据提取和存储。
module {
  func.func @triton_ldst_indirect_08_kernel(%arg0: memref<?xi8>, %arg1: memref<?xi8>, %arg2: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 1 : i32}, %arg3: memref<?xi64> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32}, %arg4: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32}, %arg5: i32 {tt.divisibility = 16 : i32}, %arg6: i32, %arg7: i32, %arg8: i32, %arg9: i32, %arg10: i32, %arg11: i32) attributes {SyncBlockLockArgIdx = 0 : i64, WorkspaceArgIdx = 1 : i64, global_kernel = "local", mix_mode = "aiv", parallel_mode = "simd"} {
    %c8_i32 = arith.constant 8 : i32
    %c0 = arith.constant 0 : index
    %c1 = arith.constant 1 : index
    %c8 = arith.constant 8 : index
    %c16 = arith.constant 16 : index
    %c32_i32 = arith.constant 32 : i32
    %0 = arith.muli %arg9, %c8_i32 : i32
    %reinterpret_cast = memref.reinterpret_cast %arg3 to offset: [0], sizes: [16], strides: [1] : memref<?xi64> to memref<16xi64, strided<[1]>>
    %alloc = memref.alloc() : memref<16xi64>
    memref.copy %reinterpret_cast, %alloc : memref<16xi64, strided<[1]>> to memref<16xi64>
    %1 = bufferization.to_tensor %alloc restrict writable : memref<16xi64>
    %2 = tensor.empty() : tensor<8x16xf32>
    %3 = scf.for %arg12 = %c0 to %c8 step %c1 iter_args(%arg13 = %2) -> (tensor<8x16xf32>) {
      %5 = scf.for %arg14 = %c0 to %c16 step %c1 iter_args(%arg15 = %arg13) -> (tensor<8x16xf32>) {
        %6 = arith.index_cast %arg12 : index to i32
        %7 = arith.muli %6, %arg5 : i32
        %8 = arith.extsi %7 : i32 to i64
        %extracted = tensor.extract %1[%arg14] {DiscreteMemAccess} : tensor<16xi64>
        %9 = arith.addi %8, %extracted : i64
        %10 = arith.index_cast %9 : i64 to index
        %reinterpret_cast_0 = memref.reinterpret_cast %arg4 to offset: [%10], sizes: [1], strides: [1] : memref<?xf32> to memref<1xf32, strided<[1], offset: ?>>
        %11 = memref.load %reinterpret_cast_0[%c0] : memref<1xf32, strided<[1], offset: ?>>
        %inserted = tensor.insert %11 into %arg15[%arg12, %arg14] : tensor<8x16xf32>
        scf.yield {DiscreteMemAccess} %inserted : tensor<8x16xf32>
      } {ExtractedLoadOrStore}
      scf.yield %5 : tensor<8x16xf32>
    } {ExtractedLoadOrStore}
    %4 = math.exp %3 : tensor<8x16xf32>
    scf.for %arg12 = %c0 to %c8 step %c1 {
      scf.for %arg13 = %c0 to %c16 step %c1 {
        %5 = arith.index_cast %arg12 : index to i32
        %6 = arith.addi %0, %5 : i32
        %7 = arith.muli %6, %c32_i32 : i32
        %8 = arith.extsi %7 : i32 to i64
        %extracted = tensor.extract %1[%arg13] {DiscreteMemAccess} : tensor<16xi64>
        %9 = arith.addi %8, %extracted : i64
        %10 = arith.index_cast %9 : i64 to index
        %reinterpret_cast_0 = memref.reinterpret_cast %arg2 to offset: [%10], sizes: [1], strides: [1] : memref<?xf32> to memref<1xf32, strided<[1], offset: ?>>
        %extracted_1 = tensor.extract %4[%arg12, %arg13] {DiscreteMemAccess} : tensor<8x16xf32>
        %11 = tensor.empty() : tensor<1xf32>
        %inserted = tensor.insert %extracted_1 into %11[%c0] : tensor<1xf32>
        bufferization.materialize_in_destination %inserted in writable %reinterpret_cast_0 : (tensor<1xf32>, memref<1xf32, strided<[1], offset: ?>>) -> ()
      } {ExtractedLoadOrStore}
    } {ExtractedLoadOrStore}
    return
  }
}

checklist

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 candyhong 的贡献)
Ccandyhong成员
2025年11月20日 创建了 pull request,commit 61b48825
ascend-robot
ascend-robot成员
2025年11月20日 评论:

Thank your for your pull-request.

The full list of commands accepted by me can be found at here.

You can get sig-info at here

likedislike
openLiBingCI成员
2025年11月20日 评论:
ascend-robot
ascend-robot成员
2025年11月20日 评论:

以下是根据您提交的修改文件推荐的Reviewer和Committer序列,需各模块评审通过后方可合入

Module List Reviewers Committers
repo-Ascend/triton-ascend wcleungaj, zhucehw, wangzhanpeng5, leopold0801, chnjz233 rusyaev-roman, liuzhuheng, kpxing, zhaozhijie, zhang-chunli01
likedislike
ascend-robotascend-robot成员
2025年11月20日 添加了label:ascend-cla/yes
此处折叠了153条消息 查看更多
ascend-robotascend-robot成员
2025年11月27日 添加了label:approved
wangzhanpeng5
2025年11月27日 评论:

/lgtm

likedislike
ascend-robotascend-robot成员
2025年11月27日 添加了label:lgtm
ascend-robot
ascend-robot成员
2025年11月27日 评论:

Review Guide

This Pull-Request Passes Review.
Committers who writed a comment of /approve are: zhang-chunli01.
Reviewers who writed a comment of /lgtm are: wangzhanpeng5, zhang-chunli01.

likedislike
ascend-robotascend-robot成员
2025年11月27日 合入了pull request