已合并
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
已合并
Pull Request已成功合入, 合并人@ascend-robot
(感谢 candyhong 的贡献)ascend-robot
2025年11月20日 评论:
2025年11月20日 评论:
openLiBingCI
2025年11月20日 评论:
2025年11月20日 评论:
ascend-robot
2025年11月20日 评论:
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 |


2025年11月20日 添加了label:ascend-cla/yes
此处折叠了153条消息 查看更多
2025年11月27日 添加了label:approved
wangzhanpeng5
2025年11月27日 评论:
2025年11月27日 评论:
/lgtm


2025年11月27日 添加了label:lgtm
ascend-robot
2025年11月27日 评论:
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.


2025年11月27日 合入了pull request
描述
新增
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
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 } }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