已合并
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日
共 7 个文件变更+351-158
@@ -373,42 +373,42 @@ def test_ldst_indirect_07():
373 triton_cal = triton_ldst_indirect_07_func(xr, xc, x2, blocksize, lowdimsize)373 triton_cal = triton_ldst_indirect_07_func(xr, xc, x2, blocksize, lowdimsize)
374 torch.testing.assert_close(triton_cal, torch_ref)374 torch.testing.assert_close(triton_cal, torch_ref)
375 375 
376-@pytest.mark.skip(reason="Indirect store to be supported")376+ 
377def test_ldst_indirect_08():377def test_ldst_indirect_08():
378 378 
379 @triton.jit379 @triton.jit
380 def triton_ldst_indirect_08_kernel(380 def triton_ldst_indirect_08_kernel(
381- out_ptr0, in_ptr1, in_ptr2, in_ptr3, stride_in_r,381+ out_ptr0, in_ptr_xc, in_ptr_x2, stride_in_r,
382- XS: tl.constexpr, RS: tl.constexpr382+ OUT_COLS: tl.constexpr, XS: tl.constexpr, RS: tl.constexpr
383 ):383 ):
384 pid = tl.program_id(0)384 pid = tl.program_id(0)
385- in_idx0 = pid * XS + tl.arange(0, XS)385+ row_idx_full = pid * XS + tl.arange(0, XS)
386- in_idx1 = tl.arange(0, RS)386+ col_pos = tl.arange(0, RS)
387- tmp0 = tl.arange(0, XS)387+ xc_vals = tl.load(in_ptr_xc + col_pos)
388- tmp1 = tl.load(in_ptr1 + in_idx1)388+ row_arange = tl.arange(0, XS)
389- in_idx2 = tmp0[:, None] * stride_in_r + tmp1[None, :]389+ gather_flat = row_arange[:, None] * stride_in_r + xc_vals[None, :]
390- tmp2 = tl.load(in_ptr2 + in_idx2)390+ vals = tl.load(in_ptr_x2 + gather_flat)
391- tmp2 = tl_math.exp(tmp2)391+ vals = tl_math.exp(vals)
392- tmp3 = tl.load(in_ptr3 + in_idx1)392+ out_flat = row_idx_full[:, None] * OUT_COLS + xc_vals[None, :]
393- tmp3 = tmp3 + 1393+ tl.store(out_ptr0 + out_flat, vals)
394- out0_idx = in_idx0[:, None] * RS + tmp3[None, :]
395- tl.store(out_ptr0 + out0_idx, tmp2)
396 394 
397 def triton_ldst_indirect_08_func(xc, x2, xs, rs):395 def triton_ldst_indirect_08_func(xc, x2, xs, rs):
398- nr = x2.size()[0]396+ nr = x2.size(0)
399- nc = xc.numel()397+ out_cols = x2.size(1)
400- stride_in_r = x2.stride()[0]398+ stride_in_r = x2.stride(0)
401 assert nr == xs, "test only single core"399 assert nr == xs, "test only single core"
402- y0 = torch.empty((nr, nc), dtype=x2.dtype, device=x2.device)400+ y0 = torch.zeros((nr, out_cols), dtype=x2.dtype, device=x2.device)
403- xc1 = xc - 1
404 triton_ldst_indirect_08_kernel[nr // xs, 1, 1](401 triton_ldst_indirect_08_kernel[nr // xs, 1, 1](
405- y0, xc, x2, xc1, stride_in_r, XS = xs, RS = rs)402+ y0, xc, x2, stride_in_r,
403+ OUT_COLS=out_cols, XS=xs, RS=rs
404+ )
406 return y0405 return y0
407 406 
408 def torch_ldst_indirect_08_func(xr, xc, x2):407 def torch_ldst_indirect_08_func(xr, xc, x2):
409- flatten_idx = (xr[:, None] * x2.stride()[0] + xc[None, :]).flatten()408+ out = torch.zeros((xr.numel(), x2.size(1)), dtype=x2.dtype, device=x2.device)
410- extracted = x2.flatten()[flatten_idx].reshape([xr.numel(), xc.numel()])409+ gathered = torch.exp(x2[xr[:, None], xc[None, :]])
411- return torch.exp(extracted)410+ out.scatter_(1, xc.expand(xr.numel(), -1), gathered)
411+ return out
412 412 
413 DEV = "npu"413 DEV = "npu"
414 DTYPE = torch.float32414 DTYPE = torch.float32
@@ -416,7 +416,7 @@ def test_ldst_indirect_08():
416 N0, N1 = 16, 32416 N0, N1 = 16, 32
417 blocksize = 8417 blocksize = 8
418 lowdimsize = N0418 lowdimsize = N0
419- assert N1 >= N0+offset, "N1 must be >= N0+offset"419+ assert N1 >= N0 + offset, "N1 must be >= N0+offset"
420 assert N0 == lowdimsize, "N0 must be == lowdimsize"420 assert N0 == lowdimsize, "N0 must be == lowdimsize"
421 xc = offset + torch.arange(0, N0, device=DEV)421 xc = offset + torch.arange(0, N0, device=DEV)
422 xr = torch.arange(0, blocksize, device=DEV)422 xr = torch.arange(0, blocksize, device=DEV)
@@ -425,6 +425,7 @@ def test_ldst_indirect_08():
425 triton_cal = triton_ldst_indirect_08_func(xc, x2, blocksize, lowdimsize)425 triton_cal = triton_ldst_indirect_08_func(xc, x2, blocksize, lowdimsize)
426 torch.testing.assert_close(triton_cal, torch_ref)426 torch.testing.assert_close(triton_cal, torch_ref)
427 427 
428+ 
428def test_ldst_indirect_09():429def test_ldst_indirect_09():
429 430 
430 @triton.jit431 @triton.jit
Rascend/test/Conversion/TritonToUnstructure/indirect_load.mlir→ascend/test/Conversion/TritonToUnstructure/indirect_func.mlir+108-37
@@ -1,41 +1,41 @@
1-// RUN: triton-adapter-opt --triton-to-unstructure=compile-on-910-95=true %s | FileCheck %s1+// RUN: triton-adapter-opt --triton-linearize '--discrete-mask-access-conversion=compile-on-910-95=True force-simt-template=True' --triton-to-annotation '--triton-to-unstructure=compile-on-910-95=True force-simt-template=True' %s --split-input-file | FileCheck %s
2 2 
3+// tt.load -> tt.indirect_load
3tt.func public @triton_ldst_indirect_05_kernel(%arg0: !tt.ptr<f32>, %arg1: !tt.ptr<i64>, %arg2: !tt.ptr<f32>, %arg3: i32) attributes {noinline = false} {4tt.func public @triton_ldst_indirect_05_kernel(%arg0: !tt.ptr<f32>, %arg1: !tt.ptr<i64>, %arg2: !tt.ptr<f32>, %arg3: i32) attributes {noinline = false} {
4- %cst = arith.constant dense<16> : tensor<8x1xi32>5+ %cst = arith.constant dense<16> : tensor<8x1xi32>
5- %c8_i32 = arith.constant 8 : i326+ %c8_i32 = arith.constant 8 : i32
6- %0 = tt.get_program_id x : i327+ %0 = tt.get_program_id x : i32
7- %1 = arith.muli %0, %c8_i32 : i328+ %1 = arith.muli %0, %c8_i32 : i32
8- %2 = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32>9+ %2 = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32>
9- %3 = tt.splat %1 : i32 -> tensor<8xi32>10+ %3 = tt.splat %1 : i32 -> tensor<8xi32>
10- %4 = arith.addi %3, %2 : tensor<8xi32>11+ %4 = arith.addi %3, %2 : tensor<8xi32>
11- %5 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32>12+ %5 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32>
12- %6 = tt.splat %arg1 : !tt.ptr<i64> -> tensor<16x!tt.ptr<i64>>13+ %6 = tt.splat %arg1 : !tt.ptr<i64> -> tensor<16x!tt.ptr<i64>>
13- %7 = tt.addptr %6, %5 : tensor<16x!tt.ptr<i64>>, tensor<16xi32>14+ %7 = tt.addptr %6, %5 : tensor<16x!tt.ptr<i64>>, tensor<16xi32>
14- %8 = tt.load %7 : tensor<16x!tt.ptr<i64>>15+ %8 = tt.load %7 : tensor<16x!tt.ptr<i64>>
15- %9 = tt.expand_dims %2 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>16+ %9 = tt.expand_dims %2 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
16- %10 = tt.splat %arg3 : i32 -> tensor<8x1xi32>17+ %10 = tt.splat %arg3 : i32 -> tensor<8x1xi32>
17- %11 = arith.muli %9, %10 : tensor<8x1xi32>18+ %11 = arith.muli %9, %10 : tensor<8x1xi32>
18- %12 = tt.expand_dims %8 {axis = 0 : i32} : tensor<16xi64> -> tensor<1x16xi64>19+ %12 = tt.expand_dims %8 {axis = 0 : i32} : tensor<16xi64> -> tensor<1x16xi64>
19- %13 = arith.extsi %11 : tensor<8x1xi32> to tensor<8x1xi64>20+ %13 = arith.extsi %11 : tensor<8x1xi32> to tensor<8x1xi64>
20- %14 = tt.broadcast %13 : tensor<8x1xi64> -> tensor<8x16xi64>21+ %14 = tt.broadcast %13 : tensor<8x1xi64> -> tensor<8x16xi64>
21- %15 = tt.broadcast %12 : tensor<1x16xi64> -> tensor<8x16xi64>22+ %15 = tt.broadcast %12 : tensor<1x16xi64> -> tensor<8x16xi64>
22- %16 = arith.addi %14, %15 : tensor<8x16xi64>23+ %16 = arith.addi %14, %15 : tensor<8x16xi64>
23- %17 = tt.splat %arg2 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>24+ %17 = tt.splat %arg2 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>
24- %18 = tt.addptr %17, %16 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi64>25+ %18 = tt.addptr %17, %16 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi64>
25- %19 = tt.load %18 : tensor<8x16x!tt.ptr<f32>>26+ %19 = tt.load %18 : tensor<8x16x!tt.ptr<f32>>
26- %20 = math.exp %19 : tensor<8x16xf32>27+ %20 = math.exp %19 : tensor<8x16xf32>
27- %21 = tt.expand_dims %4 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>28+ %21 = tt.expand_dims %4 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
28- %22 = arith.muli %21, %cst : tensor<8x1xi32>29+ %22 = arith.muli %21, %cst : tensor<8x1xi32>
29- %23 = tt.expand_dims %5 {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32>30+ %23 = tt.expand_dims %5 {axis = 0 : i32} : tensor<16xi32> -> tensor<1x16xi32>
30- %24 = tt.broadcast %22 : tensor<8x1xi32> -> tensor<8x16xi32>31+ %24 = tt.broadcast %22 : tensor<8x1xi32> -> tensor<8x16xi32>
31- %25 = tt.broadcast %23 : tensor<1x16xi32> -> tensor<8x16xi32>32+ %25 = tt.broadcast %23 : tensor<1x16xi32> -> tensor<8x16xi32>
32- %26 = arith.addi %24, %25 : tensor<8x16xi32>33+ %26 = arith.addi %24, %25 : tensor<8x16xi32>
33- %27 = tt.splat %arg0 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>34+ %27 = tt.splat %arg0 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>
34- %28 = tt.addptr %27, %26 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi32>35+ %28 = tt.addptr %27, %26 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi32>
35- tt.store %28, %20 : tensor<8x16x!tt.ptr<f32>>36+ tt.store %28, %20 : tensor<8x16x!tt.ptr<f32>>
36- tt.return37+ tt.return
37- }38+}
38- 
39 39 
40// CHECK-LABEL: tt.func public @triton_ldst_indirect_05_kernel(40// CHECK-LABEL: tt.func public @triton_ldst_indirect_05_kernel(
41// CHECK-SAME: %[[VAL_0:.*]]: !tt.ptr<f32>, %[[VAL_1:.*]]: !tt.ptr<i64>, %[[VAL_2:.*]]: !tt.ptr<f32>, %[[VAL_3:.*]]: i32) attributes {noinline = false} {41// CHECK-SAME: %[[VAL_0:.*]]: !tt.ptr<f32>, %[[VAL_1:.*]]: !tt.ptr<i64>, %[[VAL_2:.*]]: !tt.ptr<f32>, %[[VAL_3:.*]]: i32) attributes {noinline = false} {
@@ -70,4 +70,75 @@ tt.func public @triton_ldst_indirect_05_kernel(%arg0: !tt.ptr<f32>, %arg1: !tt.p
70// CHECK: %[[VAL_32:.*]] = tt.addptr %[[VAL_31]], %[[VAL_30]] : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi32>70// CHECK: %[[VAL_32:.*]] = tt.addptr %[[VAL_31]], %[[VAL_30]] : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi32>
71// CHECK: tt.store %[[VAL_32]], %[[VAL_24]] : tensor<8x16x!tt.ptr<f32>>71// CHECK: tt.store %[[VAL_32]], %[[VAL_24]] : tensor<8x16x!tt.ptr<f32>>
72// CHECK: tt.return72// CHECK: tt.return
73-// CHECK: }73+// CHECK: }
74+ 
75+// -----
76+ 
77+// tt.store -> tt.indirect_store
78+tt.func public @triton_ldst_indirect_08_kernel(%arg0: !tt.ptr<f32>, %arg1: !tt.ptr<i64>, %arg2: !tt.ptr<f32>, %arg3: i32) attributes {noinline = false} {
79+ %cst = arith.constant dense<32> : tensor<8x1xi32>
80+ %c8_i32 = arith.constant 8 : i32
81+ %0 = tt.get_program_id x : i32
82+ %1 = arith.muli %0, %c8_i32 : i32
83+ %2 = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32>
84+ %3 = tt.splat %1 : i32 -> tensor<8xi32>
85+ %4 = arith.addi %3, %2 : tensor<8xi32>
86+ %5 = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32>
87+ %6 = tt.splat %arg1 : !tt.ptr<i64> -> tensor<16x!tt.ptr<i64>>
88+ %7 = tt.addptr %6, %5 : tensor<16x!tt.ptr<i64>>, tensor<16xi32>
89+ %8 = tt.load %7 : tensor<16x!tt.ptr<i64>>
90+ %9 = tt.expand_dims %2 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
91+ %10 = tt.splat %arg3 : i32 -> tensor<8x1xi32>
92+ %11 = arith.muli %9, %10 : tensor<8x1xi32>
93+ %12 = tt.expand_dims %8 {axis = 0 : i32} : tensor<16xi64> -> tensor<1x16xi64>
94+ %13 = arith.extsi %11 : tensor<8x1xi32> to tensor<8x1xi64>
95+ %14 = tt.broadcast %13 : tensor<8x1xi64> -> tensor<8x16xi64>
96+ %15 = tt.broadcast %12 : tensor<1x16xi64> -> tensor<8x16xi64>
97+ %16 = arith.addi %14, %15 : tensor<8x16xi64>
98+ %17 = tt.splat %arg2 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>
99+ %18 = tt.addptr %17, %16 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi64>
100+ %19 = tt.load %18 : tensor<8x16x!tt.ptr<f32>>
101+ %20 = math.exp %19 : tensor<8x16xf32>
102+ %21 = tt.expand_dims %4 {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
103+ %22 = arith.muli %21, %cst : tensor<8x1xi32>
104+ %23 = arith.extsi %22 : tensor<8x1xi32> to tensor<8x1xi64>
105+ %24 = tt.broadcast %23 : tensor<8x1xi64> -> tensor<8x16xi64>
106+ %25 = arith.addi %24, %15 : tensor<8x16xi64>
107+ %26 = tt.splat %arg0 : !tt.ptr<f32> -> tensor<8x16x!tt.ptr<f32>>
108+ %27 = tt.addptr %26, %25 : tensor<8x16x!tt.ptr<f32>>, tensor<8x16xi64>
109+ tt.store %27, %20 : tensor<8x16x!tt.ptr<f32>>
110+ tt.return
111+}
112+ 
113+// CHECK-LABEL: tt.func public @triton_ldst_indirect_08_kernel(
114+// CHECK-SAME: %[[VAL_0:.*]]: !tt.ptr<f32>, %[[VAL_1:.*]]: !tt.ptr<i64>, %[[VAL_2:.*]]: !tt.ptr<f32>, %[[VAL_3:.*]]: i32) attributes {noinline = false} {
115+// CHECK: %[[VAL_4:.*]] = arith.constant dense<32> : tensor<8x1xi32>
116+// CHECK: %[[VAL_5:.*]] = arith.constant 8 : i32
117+// CHECK: %[[VAL_6:.*]] = tt.get_program_id x : i32
118+// CHECK: %[[VAL_7:.*]] = arith.muli %[[VAL_6:.*]], %[[VAL_5:.*]] : i32
119+// CHECK: %[[VAL_8:.*]] = tt.make_range {end = 8 : i32, start = 0 : i32} : tensor<8xi32>
120+// CHECK: %[[VAL_9:.*]] = tt.splat %[[VAL_7:.*]] : i32 -> tensor<8xi32>
121+// CHECK: %[[VAL_10:.*]] = arith.addi %[[VAL_9:.*]], %[[VAL_8:.*]] : tensor<8xi32>
122+// CHECK: %[[VAL_11:.*]] = tt.make_range {end = 16 : i32, start = 0 : i32} : tensor<16xi32>
123+// CHECK: %[[VAL_12:.*]] = tt.splat %[[VAL_1:.*]] : !tt.ptr<i64> -> tensor<16x!tt.ptr<i64>>
124+// CHECK: %[[VAL_13:.*]] = tt.addptr %[[VAL_12:.*]], %[[VAL_11:.*]] : tensor<16x!tt.ptr<i64>>, tensor<16xi32>
125+// CHECK: %[[VAL_14:.*]] = tt.load %[[VAL_13:.*]] : tensor<16x!tt.ptr<i64>>
126+// CHECK: %[[VAL_15:.*]] = tt.expand_dims %[[VAL_8:.*]] {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
127+// CHECK: %[[VAL_16:.*]] = tt.splat %[[VAL_3:.*]] : i32 -> tensor<8x1xi32>
128+// CHECK: %[[VAL_17:.*]] = arith.muli %[[VAL_15:.*]], %[[VAL_16:.*]] : tensor<8x1xi32>
129+// CHECK: %[[VAL_18:.*]] = tt.expand_dims %[[VAL_14:.*]] {axis = 0 : i32} : tensor<16xi64> -> tensor<1x16xi64>
130+// CHECK: %[[VAL_19:.*]] = arith.extsi %[[VAL_17:.*]] : tensor<8x1xi32> to tensor<8x1xi64>
131+// CHECK: %[[VAL_20:.*]] = tt.broadcast %[[VAL_19:.*]] : tensor<8x1xi64> -> tensor<8x16xi64>
132+// CHECK: %[[VAL_21:.*]] = tt.broadcast %[[VAL_18:.*]] : tensor<1x16xi64> -> tensor<8x16xi64>
133+// CHECK: %[[VAL_22:.*]] = arith.addi %[[VAL_20:.*]], %[[VAL_21:.*]] : tensor<8x16xi64>
134+// CHECK: %[[VAL_23:.*]] = tt.indirect_load %[[VAL_2:.*]] : <f32>, %[[VAL_22:.*]] : tensor<8x16xi64> -> tensor<8x16xf32>
135+// CHECK: %[[VAL_24:.*]] = math.exp %[[VAL_23:.*]] : tensor<8x16xf32>
136+// CHECK: %[[VAL_25:.*]] = tt.expand_dims %[[VAL_10:.*]] {axis = 1 : i32} : tensor<8xi32> -> tensor<8x1xi32>
137+// CHECK: %[[VAL_26:.*]] = arith.muli %[[VAL_25:.*]], %[[VAL_4:.*]] : tensor<8x1xi32>
138+// CHECK: %[[VAL_27:.*]] = arith.extsi %[[VAL_26:.*]] : tensor<8x1xi32> to tensor<8x1xi64>
139+// CHECK: %[[VAL_28:.*]] = tt.broadcast %[[VAL_27:.*]] : tensor<8x1xi64> -> tensor<8x16xi64>
140+// CHECK: %[[VAL_29:.*]] = arith.addi %[[VAL_28:.*]], %[[VAL_21:.*]] : tensor<8x16xi64>
141+// CHECK: tt.indirect_store %[[VAL_0:.*]] : <f32>, %[[VAL_29:.*]] : tensor<8x16xi64>, %[[VAL_24:.*]] : tensor<8x16xf32>
142+// CHECK: tt.return
143+// CHECK: }
144+ 
@@ -44,6 +44,8 @@ namespace TTOpConverters {
44using namespace mlir;44using namespace mlir;
45using namespace triton;45using namespace triton;
46 46 
47+static constexpr unsigned kFuncNameCap = 128;
48+ 
47/*49/*
48Convert `tt.precise_div` operation to `arith.divf` operation.50Convert `tt.precise_div` operation to `arith.divf` operation.
49tensor_x / tensor_y51tensor_x / tensor_y
@@ -586,6 +588,16 @@ private:
586 static constexpr llvm::StringRef funcNameBase = "triton_indirect_load";588 static constexpr llvm::StringRef funcNameBase = "triton_indirect_load";
587};589};
588 590 
591+class IndirectStoreConverter : public OpConversionPattern<triton::IndirectStoreOp> {
592+public:
593+ using OpConversionPattern<triton::IndirectStoreOp>::OpConversionPattern;
594+ LogicalResult
595+ matchAndRewrite(triton::IndirectStoreOp op, OpAdaptor adaptor,
596+ ConversionPatternRewriter &rewriter) const override;
597+private:
598+ static constexpr llvm::StringRef funcNameBase = "triton_indirect_store";
599+};
600+ 
589class IndexSelectSimdConverter : public OpConversionPattern<triton::IndexSelectSimdOp> {601class IndexSelectSimdConverter : public OpConversionPattern<triton::IndexSelectSimdOp> {
590public:602public:
591 explicit IndexSelectSimdConverter(MLIRContext *context);603 explicit IndexSelectSimdConverter(MLIRContext *context);
@@ -1473,7 +1473,7 @@ GatherConverter::matchAndRewrite(triton::GatherOp op, OpAdaptor adaptor,
1473 rewriter.setInsertionPoint(moduleOp.getBody(),1473 rewriter.setInsertionPoint(moduleOp.getBody(),
1474 std::prev(moduleOp.getBody()->end()));1474 std::prev(moduleOp.getBody()->end()));
1475 1475 
1476- llvm::SmallString<128> funcName = gatherFuncNameBase;1476+ llvm::SmallString<kFuncNameCap> funcName = gatherFuncNameBase;
1477 int uniqueId = 0;1477 int uniqueId = 0;
1478 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {1478 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {
1479 funcName = gatherFuncNameBase;1479 funcName = gatherFuncNameBase;
@@ -2070,7 +2070,7 @@ EmbeddingGatherConverter::matchAndRewrite(triton::EmbeddingGatherOp op, OpAdapto
2070 rewriter.setInsertionPoint(moduleOp.getBody(),2070 rewriter.setInsertionPoint(moduleOp.getBody(),
2071 std::prev(moduleOp.getBody()->end()));2071 std::prev(moduleOp.getBody()->end()));
2072 2072 
2073- llvm::SmallString<128> funcName = funcNameBase;2073+ llvm::SmallString<kFuncNameCap> funcName = funcNameBase;
2074 int uniqueId = 0;2074 int uniqueId = 0;
2075 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {2075 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {
2076 funcName = funcNameBase;2076 funcName = funcNameBase;
@@ -2087,7 +2087,10 @@ EmbeddingGatherConverter::matchAndRewrite(triton::EmbeddingGatherOp op, OpAdapto
2087 auto resTy = res.getType();2087 auto resTy = res.getType();
2088 2088 
2089 // convert !tt.ptr<f32> to memref<?xf32>2089 // convert !tt.ptr<f32> to memref<?xf32>
2090- auto srcTy = cast<MemRefType>(src.getType());2090+ auto srcTy = dyn_cast<MemRefType>(src.getType());
2091+ if (!srcTy) {
2092+ return rewriter.notifyMatchFailure(op, "expected MemRefType for src");
2093+ }
2091 SmallVector<Type> inputTypes({srcTy, idx.getType(), bound.getType(),2094 SmallVector<Type> inputTypes({srcTy, idx.getType(), bound.getType(),
2092 blksiz.getType()});2095 blksiz.getType()});
2093 inputTypes.append(offsets.getTypes().begin(), offsets.getTypes().end());2096 inputTypes.append(offsets.getTypes().begin(), offsets.getTypes().end());
@@ -2108,17 +2111,17 @@ EmbeddingGatherConverter::matchAndRewrite(triton::EmbeddingGatherOp op, OpAdapto
2108 return success();2111 return success();
2109}2112}
2110 2113 
2111- 
2112LogicalResult2114LogicalResult
2113IndirectLoadConverter::matchAndRewrite(triton::IndirectLoadOp op, OpAdaptor adaptor,2115IndirectLoadConverter::matchAndRewrite(triton::IndirectLoadOp op, OpAdaptor adaptor,
2114- ConversionPatternRewriter &rewriter) const {2116+ ConversionPatternRewriter &rewriter) const
2117+{
2115 auto loc = op.getLoc();2118 auto loc = op.getLoc();
2116 2119 
2117 auto moduleOp = op->getParentOfType<ModuleOp>();2120 auto moduleOp = op->getParentOfType<ModuleOp>();
2118 rewriter.setInsertionPoint(moduleOp.getBody(),2121 rewriter.setInsertionPoint(moduleOp.getBody(),
2119 std::prev(moduleOp.getBody()->end()));2122 std::prev(moduleOp.getBody()->end()));
2120 2123 
2121- llvm::SmallString<128> funcName = funcNameBase;2124+ llvm::SmallString<kFuncNameCap> funcName = funcNameBase;
2122 int uniqueId = 0;2125 int uniqueId = 0;
2123 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {2126 while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {
2124 funcName = funcNameBase;2127 funcName = funcNameBase;
@@ -2133,7 +2136,10 @@ IndirectLoadConverter::matchAndRewrite(triton::IndirectLoadOp op, OpAdaptor adap
2133 auto resTy = res.getType();2136 auto resTy = res.getType();
2134 2137 
2135 // convert !tt.ptr<f32> to memref<?xf32>2138 // convert !tt.ptr<f32> to memref<?xf32>
2136- auto srcTy = cast<MemRefType>(src.getType());2139+ auto srcTy = dyn_cast<MemRefType>(src.getType());
2140+ if (!srcTy) {
2141+ return rewriter.notifyMatchFailure(op, "expected MemRefType for src");
2142+ }
2137 SmallVector<Type> inputTypes({srcTy, offsets.getType()});2143 SmallVector<Type> inputTypes({srcTy, offsets.getType()});
2138 if (mask) inputTypes.push_back(mask.getType());2144 if (mask) inputTypes.push_back(mask.getType());
2139 if (other) inputTypes.push_back(other.getType());2145 if (other) inputTypes.push_back(other.getType());
@@ -2152,6 +2158,49 @@ IndirectLoadConverter::matchAndRewrite(triton::IndirectLoadOp op, OpAdaptor adap
2152 return success();2158 return success();
2153}2159}
2154 2160 
2161+ 
2162+LogicalResult
2163+IndirectStoreConverter::matchAndRewrite(triton::IndirectStoreOp op, OpAdaptor adaptor,
2164+ ConversionPatternRewriter &rewriter) const
2165+{
2166+ auto loc = op.getLoc();
2167+ 
2168+ auto moduleOp = op->getParentOfType<ModuleOp>();
2169+ rewriter.setInsertionPoint(moduleOp.getBody(),
2170+ std::prev(moduleOp.getBody()->end()));
2171+ 
2172+ llvm::SmallString<kFuncNameCap> funcName = funcNameBase;
2173+ int uniqueId = 0;
2174+ while (SymbolTable::lookupSymbolIn(moduleOp, funcName)) {
2175+ funcName = funcNameBase;
2176+ funcName += ("_" + std::to_string(uniqueId++));
2177+ }
2178+ 
2179+ auto src = adaptor.getSrc();
2180+ auto offsets = op.getOffsets();
2181+ auto value = op.getValue();
2182+ auto mask = op.getMask();
2183+ 
2184+ // convert !tt.ptr<f32> to memref<?xf32>
2185+ auto srcTy = dyn_cast<MemRefType>(src.getType());
2186+ if (!srcTy) {
2187+ return rewriter.notifyMatchFailure(op, "expected MemRefType for src");
2188+ }
2189+ SmallVector<Type> inputTypes({srcTy, offsets.getType(), value.getType()});
2190+ if (mask) inputTypes.push_back(mask.getType());
2191+ 
2192+ auto libFnType = rewriter.getFunctionType(inputTypes, {});
2193+ auto funcOp = rewriter.create<func::FuncOp>(loc, funcName.str(), libFnType);
2194+ SymbolTable::setSymbolVisibility(funcOp, SymbolTable::Visibility::Private);
2195+ 
2196+ rewriter.setInsertionPoint(op);
2197+ SmallVector<Value> inputVals({src, offsets, value});
2198+ if (mask) inputVals.push_back(mask);
2199+ rewriter.create<func::CallOp>(loc, funcOp, inputVals);
2200+ rewriter.eraseOp(op);
2201+ return success();
2202+}
2203+ 
2155IndexSelectSimdConverter::IndexSelectSimdConverter(MLIRContext *context)2204IndexSelectSimdConverter::IndexSelectSimdConverter(MLIRContext *context)
2156 : OpConversionPattern<triton::IndexSelectSimdOp>(context) {}2205 : OpConversionPattern<triton::IndexSelectSimdOp>(context) {}
2157 2206 
@@ -82,10 +82,96 @@ inline bool isSIMTOp(Operation *op)
82{82{
83 return isa<83 return isa<
84 triton::EmbeddingGatherOp,84 triton::EmbeddingGatherOp,
85- triton::IndirectLoadOp85+ triton::IndirectLoadOp,
86+ triton::IndirectStoreOp
86 >(op);87 >(op);
87}88}
88 89 
90+template <typename T, typename = void> struct has_getPtr : std::false_type {};
91+template <typename T>
92+struct has_getPtr<T, std::void_t<decltype(std::declval<T>().getPtr())>> : std::true_type {};
93+ 
94+template <typename T, typename = void> struct has_getSrc : std::false_type {};
95+template <typename T>
96+struct has_getSrc<T, std::void_t<decltype(std::declval<T>().getSrc())>> : std::true_type {};
97+ 
98+template <typename T, typename = void> struct has_getBase : std::false_type {};
99+template <typename T>
100+struct has_getBase<T, std::void_t<decltype(std::declval<T>().getBase())>> : std::true_type {};
101+ 
102+template <typename OpTy>
103+static Value extractPointer(OpTy op) {
104+ if constexpr (has_getPtr<OpTy>::value)
105+ return op.getPtr();
106+ else if constexpr (has_getSrc<OpTy>::value)
107+ return op.getSrc();
108+ else if constexpr (has_getBase<OpTy>::value)
109+ return op.getBase();
110+ else {
111+ Operation *raw = op.getOperation();
112+ if (!raw || raw->getNumOperands() == 0)
113+ return Value();
114+ return raw->getOperand(0);
115+ }
116+}
117+ 
118+static void setBlockArgumentAttr(BlockArgument blockArg, triton::FuncOp func, TensorKind tensorKind)
119+{
120+ unsigned argIdx = blockArg.getArgNumber();
121+ auto existingAttr = func.getArgAttrOfType<IntegerAttr>(argIdx, "tt.tensor_kind");
122+ TensorKind oldVal = existingAttr ? static_cast<TensorKind>(existingAttr.getInt()) : TensorKind::NONE;
123+ 
124+ TensorKind finalVal = tensorKind;
125+ if ((oldVal == TensorKind::INPUT && tensorKind == TensorKind::OUTPUT) ||
126+ (oldVal == TensorKind::OUTPUT && tensorKind == TensorKind::INPUT)) {
127+ finalVal = TensorKind::INPUT_OUTPUT;
128+ } else if (oldVal == TensorKind::INPUT_OUTPUT) {
129+ finalVal = oldVal;
130+ }
131+ 
132+ LLVM_DEBUG(llvm::dbgs() << "Setting tensor_kind for argument " << argIdx << ": " << finalVal << "\n";);
133+ 
134+ func.setArgAttr(argIdx, "tt.tensor_kind",
135+ IntegerAttr::get(IntegerType::get(func.getContext(), INT_BIT_WIDTH), static_cast<int>(finalVal)));
136+}
137+ 
138+template <typename OpTy>
139+void TritonToLinalgPass::addTensorKindToArguments(OpTy op, triton::FuncOp func, TensorKind tensorKind)
140+{
141+ Value ptr = extractPointer(op);
142+ if (!ptr)
143+ return;
144+ 
145+ LLVM_DEBUG(llvm::dbgs() << "Processing op: " << *op.getOperation() << "\n";);
146+ 
147+ Value cur = ptr;
148+ llvm::SmallPtrSet<Value, SET_INIT_SIZE> visited;
149+ // Walk back the def-use chain to find originating BlockArgument
150+ while (visited.insert(cur).second) {
151+ // If reach a BlockArgument, set the attribute
152+ if (auto blockArg = dyn_cast<BlockArgument>(cur)) {
153+ if (blockArg.getOwner() == &func.getBody().front()) {
154+ auto type = blockArg.getType();
155+ // Check if it's a triton::PointerType
156+ if (!isa<triton::PointerType>(type))
157+ break;
158+ setBlockArgumentAttr(blockArg, func, tensorKind);
159+ break;
160+ }
161+ }
162+ 
163+ Operation *defOp = cur.getDefiningOp();
164+ if (!defOp)
165+ break;
166+ cur = defOp->getOperand(0);
167+ }
168+}
169+ 
170+template <TensorKind Kind, typename... Ops>
171+void TritonToLinalgPass::walkAndMarkTensorKind(triton::FuncOp func) {
172+ (func.walk([&](Ops op) { this->addTensorKindToArguments(op, func, Kind); }), ...);
173+}
174+ 
89TritonTypeConverter::TritonTypeConverter() {175TritonTypeConverter::TritonTypeConverter() {
90 addConversion([](Type type) { return type; });176 addConversion([](Type type) { return type; });
91 177 
@@ -155,86 +241,6 @@ void TritonToLinalgPass::addProgramInfo(triton::FuncOp func,
155 }241 }
156}242}
157 243 
158-static void setBlockArgumentAttr(BlockArgument blockArg, triton::FuncOp func, TensorKind tensorKind)
159-{
160- unsigned argIdx = blockArg.getArgNumber();
161- auto existingAttr = func.getArgAttrOfType<IntegerAttr>(argIdx, "tt.tensor_kind");
162- TensorKind oldVal = existingAttr ? static_cast<TensorKind>(existingAttr.getInt()) : TensorKind::NONE;
163- 
164- TensorKind finalVal = tensorKind;
165- if ((oldVal == TensorKind::INPUT && tensorKind == TensorKind::OUTPUT) ||
166- (oldVal == TensorKind::OUTPUT && tensorKind == TensorKind::INPUT)) {
167- finalVal = TensorKind::INPUT_OUTPUT;
168- } else if (oldVal == TensorKind::INPUT_OUTPUT) {
169- finalVal = oldVal;
170- }
171- 
172- LLVM_DEBUG(llvm::dbgs() << "Setting tensor_kind for argument " << argIdx << ": " << finalVal << "\n";);
173- 
174- func.setArgAttr(argIdx, "tt.tensor_kind",
175- IntegerAttr::get(IntegerType::get(func.getContext(), INT_BIT_WIDTH), static_cast<int>(finalVal)));
176-}
177- 
178-template <typename T, typename = void> struct has_getPtr : std::false_type {};
179-template <typename T>
180-struct has_getPtr<T, std::void_t<decltype(std::declval<T>().getPtr())>> : std::true_type {};
181- 
182-template <typename T, typename = void> struct has_getSrc : std::false_type {};
183-template <typename T>
184-struct has_getSrc<T, std::void_t<decltype(std::declval<T>().getSrc())>> : std::true_type {};
185- 
186-template <typename T, typename = void> struct has_getBase : std::false_type {};
187-template <typename T>
188-struct has_getBase<T, std::void_t<decltype(std::declval<T>().getBase())>> : std::true_type {};
189- 
190-template <typename OpTy>
191-static Value extractPointer(OpTy op) {
192- if constexpr (has_getPtr<OpTy>::value)
193- return op.getPtr();
194- else if constexpr (has_getSrc<OpTy>::value)
195- return op.getSrc();
196- else if constexpr (has_getBase<OpTy>::value)
197- return op.getBase();
198- else {
199- Operation *raw = op.getOperation();
200- if (!raw || raw->getNumOperands() == 0)
201- return Value();
202- return raw->getOperand(0);
203- }
204-}
205- 
206-template <typename OpTy>
207-void TritonToLinalgPass::addTensorKindToArguments(OpTy op, triton::FuncOp func, TensorKind tensorKind)
208-{
209- Value ptr = extractPointer(op);
210- if (!ptr)
211- return;
212- 
213- LLVM_DEBUG(llvm::dbgs() << "Processing op: " << *op.getOperation() << "\n";);
214- 
215- Value cur = ptr;
216- llvm::SmallPtrSet<Value, SET_INIT_SIZE> visited;
217- // Walk back the def-use chain to find originating BlockArgument
218- while (visited.insert(cur).second) {
219- // If reach a BlockArgument, set the attribute
220- if (auto blockArg = dyn_cast<BlockArgument>(cur)) {
221- if (blockArg.getOwner() == &func.getBody().front()) {
222- auto type = blockArg.getType();
223- // Check if it's a triton::PointerType
224- if (!isa<triton::PointerType>(type))
225- break;
226- setBlockArgumentAttr(blockArg, func, tensorKind);
227- break;
228- }
229- }
230- 
231- Operation *defOp = cur.getDefiningOp();
232- if (!defOp)
233- break;
234- cur = defOp->getOperand(0);
235- }
236-}
237- 
238LogicalResult244LogicalResult
239TritonToLinalgPass::convertMultipleBlockControlFlow(Operation *funcOp,245TritonToLinalgPass::convertMultipleBlockControlFlow(Operation *funcOp,
240 OpBuilder &builder) {246 OpBuilder &builder) {
@@ -667,6 +673,7 @@ void TritonToLinalgPass::populateTritonToLinalgConversionPatterns(
667 patterns.add<TTOpConverters::DotScaledConverter>(patterns.getContext());673 patterns.add<TTOpConverters::DotScaledConverter>(patterns.getContext());
668 patterns.add<TTOpConverters::PtrToIntConverter>(patterns.getContext());674 patterns.add<TTOpConverters::PtrToIntConverter>(patterns.getContext());
669 patterns.add<TTOpConverters::IndirectLoadConverter>(patterns.getContext());675 patterns.add<TTOpConverters::IndirectLoadConverter>(patterns.getContext());
676+ patterns.add<TTOpConverters::IndirectStoreConverter>(patterns.getContext());
670 patterns.add<TTOpConverters::IndexSelectSimdConverter>(patterns.getContext());677 patterns.add<TTOpConverters::IndexSelectSimdConverter>(patterns.getContext());
671 678 
672 if (!this->namedOps) {679 if (!this->namedOps) {
@@ -715,11 +722,6 @@ LogicalResult TritonToLinalgPass::processDescriptorOperations(ModuleOp moduleOp)
715 return success();722 return success();
716}723}
717 724 
718-template <TensorKind Kind, typename... Ops>
719-void TritonToLinalgPass::walkAndMarkTensorKind(triton::FuncOp func) {
720- (func.walk([&](Ops op) { this->addTensorKindToArguments(op, func, Kind); }), ...);
721-}
722- 
723void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) {725void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) {
724 moduleOp.walk([&](triton::FuncOp func) {726 moduleOp.walk([&](triton::FuncOp func) {
725 // INPUT tensors727 // INPUT tensors
@@ -730,6 +732,7 @@ void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) {
730 triton::LoadOp>(func);732 triton::LoadOp>(func);
731 // OUTPUT tensors733 // OUTPUT tensors
732 this->walkAndMarkTensorKind<TensorKind::OUTPUT,734 this->walkAndMarkTensorKind<TensorKind::OUTPUT,
735+ triton::IndirectStoreOp,
733 triton::StoreOp>(func);736 triton::StoreOp>(func);
734 // INPUT_OUTPUT tensors737 // INPUT_OUTPUT tensors
735 this->walkAndMarkTensorKind<TensorKind::INPUT_OUTPUT,738 this->walkAndMarkTensorKind<TensorKind::INPUT_OUTPUT,
@@ -339,16 +339,34 @@ LogicalResult UnstructuredMemAccessConverter<MemAccOpTy>::matchAndRewrite(
339 if (sizeInByte % 32 != 0)339 if (sizeInByte % 32 != 0)
340 ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank());340 ptrOffsetInfo.setUnstructured(ptrOffsetInfo.getRank());
341 341
342- // Fast path on A5: rewrite tl.load to tt.indirect_load directly.342+ // Fast path on A5: rewrite tt.load/store to tt.indirect_load/store directly.
343- if constexpr (std::is_same_v<MemAccOpTy, triton::LoadOp>) {343+ if (compileOn91095Flag && forceSimtTemplateFlag && ptrOffsetInfo.isUnstructured()) {
344- if (compileOn91095Flag && forceSimtTemplateFlag && ptrOffsetInfo.isUnstructured()) {344+ if constexpr (std::is_same_v<MemAccOpTy, triton::LoadOp>) {
345- assert(!isa<RankedTensorType>(srcPtr.getType()) && "src must be ptr type");345+ assert(isa<triton::PointerType>(srcPtr.getType()) && "src must be ptr type");
346 Value mask = op.getMask();346 Value mask = op.getMask();
347 Value other = op.getOther();347 Value other = op.getOther();
348 auto resultType = op.getType();348 auto resultType = op.getType();
349 auto indirect = rewriter.create<triton::IndirectLoadOp>(349 auto indirect = rewriter.create<triton::IndirectLoadOp>(
350 loc, resultType, srcPtr, ptrOffset, mask, other);350 loc, resultType, srcPtr, ptrOffset, mask, other);
351 rewriter.replaceOp(op, indirect.getResult());351 rewriter.replaceOp(op, indirect.getResult());
352+ LLVM_DEBUG({
353+ auto &os = llvm::dbgs();
354+ os << "Rewriting tt.load to tt.indirect_load\n";
355+ os << indirect << "\n";
356+ });
357+ return success();
358+ } else if constexpr (std::is_same_v<MemAccOpTy, triton::StoreOp>) {
359+ assert(isa<triton::PointerType>(srcPtr.getType()) && "src must be ptr type");
360+ Value value = op.getValue();
361+ Value mask = op.getMask();
362+ auto indirect = rewriter.create<triton::IndirectStoreOp>(
363+ loc, srcPtr, ptrOffset, value, mask);
364+ rewriter.eraseOp(op);
365+ LLVM_DEBUG({
366+ auto &os = llvm::dbgs();
367+ os << "Rewriting tt.store to tt.indirect_store\n";
368+ os << indirect << "\n";
369+ });
352 return success();370 return success();
353 }371 }
354 }372 }
@@ -1565,4 +1565,43 @@ def TT_IndirectLoadOp : TT_Op<"indirect_load", [
1565 ];1565 ];
1566}1566}
1567 1567 
1568+ 
1569+//
1570+// Built-in: IndirectStore Op
1571+//
1572+def TT_IndirectStoreOp : TT_Op<"indirect_store", [
1573+ MemoryEffects<[MemWrite<GlobalMemory>]>
1574+]> {
1575+ let summary = "Built-in: indirect store from UB using per-element offsets with optional mask/other";
1576+ 
1577+ let description = [{
1578+ Built-in operation emitted by the compiler for unstructured (discrete) memory
1579+ accesses.These are not written directly in the user IR.
1580+ 
1581+ Store values from UB based to GM on per-element offsets.
1582+ 
1583+ The operation takes:
1584+ - src: Source pointer
1585+ - offsets: Tensor of per-element offsets (relative to `src`) for accessing source memory
1586+ - value: The tensor of elements to be stored
1587+ - mask (optional): If mask[idx] is false, do not store value[idx] at pointer[idx]
1588+ }];
1589+ 
1590+ let arguments = (
1591+ ins TT_Ptr:$src,
1592+ TT_IntTensor:$offsets,
1593+ TT_Type:$value,
1594+ Optional<TT_BoolLike>:$mask
1595+ );
1596+ 
1597+ let assemblyFormat = [{
1598+ $src `:` type($src) `,`
1599+ $offsets `:` type($offsets) `,`
1600+ $value `:` type($value)
1601+ (`,` $mask^ `:` type($mask))?
1602+ attr-dict
1603+ }];
1604+ 
1605+}
1606+ 
1568#endif // Triton_OPS1607#endif // Triton_OPS