已合并
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
已合并
共 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 | + |
| 377 | def test_ldst_indirect_08(): | 377 | def test_ldst_indirect_08(): |
| 378 | 378 | ||
| 379 | 379 | ||
| 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.constexpr | 382 | + 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 + 1 | 393 | + 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 y0 | 405 | 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.float32 | 414 | DTYPE = torch.float32 |
| @@ -416,7 +416,7 @@ def test_ldst_indirect_08(): | |||
| 416 | N0, N1 = 16, 32 | 416 | N0, N1 = 16, 32 |
| 417 | blocksize = 8 | 417 | blocksize = 8 |
| 418 | lowdimsize = N0 | 418 | 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 | + | ||
| 428 | def test_ldst_indirect_09(): | 429 | def test_ldst_indirect_09(): |
| 429 | 430 | ||
| 430 | 431 | ||
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 %s | 1 | +// 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 | ||
| 3 | tt.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 | tt.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 : i32 | 6 | + %c8_i32 = arith.constant 8 : i32 |
| 6 | - %0 = tt.get_program_id x : i32 | 7 | + %0 = tt.get_program_id x : i32 |
| 7 | - %1 = arith.muli %0, %c8_i32 : i32 | 8 | + %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.return | 37 | + 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.return | 72 | // 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 { | |||
| 44 | using namespace mlir; | 44 | using namespace mlir; |
| 45 | using namespace triton; | 45 | using namespace triton; |
| 46 | 46 | ||
| 47 | +static constexpr unsigned kFuncNameCap = 128; | ||
| 48 | + | ||
| 47 | /* | 49 | /* |
| 48 | Convert `tt.precise_div` operation to `arith.divf` operation. | 50 | Convert `tt.precise_div` operation to `arith.divf` operation. |
| 49 | tensor_x / tensor_y | 51 | tensor_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 | + | ||
| 589 | class IndexSelectSimdConverter : public OpConversionPattern<triton::IndexSelectSimdOp> { | 601 | class IndexSelectSimdConverter : public OpConversionPattern<triton::IndexSelectSimdOp> { |
| 590 | public: | 602 | public: |
| 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 | - | ||
| 2112 | LogicalResult | 2114 | LogicalResult |
| 2113 | IndirectLoadConverter::matchAndRewrite(triton::IndirectLoadOp op, OpAdaptor adaptor, | 2115 | IndirectLoadConverter::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 | + | ||
| 2155 | IndexSelectSimdConverter::IndexSelectSimdConverter(MLIRContext *context) | 2204 | IndexSelectSimdConverter::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::IndirectLoadOp | 85 | + 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 | + | ||
| 89 | TritonTypeConverter::TritonTypeConverter() { | 175 | TritonTypeConverter::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 | - | ||
| 238 | LogicalResult | 244 | LogicalResult |
| 239 | TritonToLinalgPass::convertMultipleBlockControlFlow(Operation *funcOp, | 245 | TritonToLinalgPass::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 | - | ||
| 723 | void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) { | 725 | void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) { |
| 724 | moduleOp.walk([&](triton::FuncOp func) { | 726 | moduleOp.walk([&](triton::FuncOp func) { |
| 725 | // INPUT tensors | 727 | // INPUT tensors |
| @@ -730,6 +732,7 @@ void TritonToLinalgPass::annotateTensorKindForModule(ModuleOp moduleOp) { | |||
| 730 | triton::LoadOp>(func); | 732 | triton::LoadOp>(func); |
| 731 | // OUTPUT tensors | 733 | // 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 tensors | 737 | // 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_OPS | 1607 | #endif // Triton_OPS |