Skip to content

Commit a0fa823

Browse files
authored
Replace memref.load/store with affine.load/store (#300)
This PR modifies to generated MLIR code In ConversionPatterns.hpp. It replaces memref.load/store with affine.load/store, and also updates lit tests
1 parent d3702bd commit a0fa823

23 files changed

Lines changed: 78 additions & 94 deletions

include/triton-shared/Conversion/TritonArithToLinalg/ConversionPatterns.hpp

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -361,10 +361,9 @@ struct LoadConverter : public OpConversionPattern<triton::LoadOp> {
361361
if (!isa<ShapedType>(op.getResult().getType())) {
362362
auto sMemRef = PtrAnalysis::getScalarMemRef(op.getPtr(), adaptor.getPtr(),
363363
loc, rewriter);
364-
auto index =
365-
rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0))
366-
.getResult();
367-
auto loadOp = rewriter.create<memref::LoadOp>(loc, sMemRef, index);
364+
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
365+
auto loadOp = rewriter.create<affine::AffineLoadOp>(
366+
op.getLoc(), sMemRef, zeroMap, std::nullopt);
368367
rewriter.replaceOp(op, loadOp.getResult());
369368
return success();
370369
}
@@ -525,7 +524,9 @@ struct StoreConverter : public OpConversionPattern<triton::StoreOp> {
525524
auto index =
526525
rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0))
527526
.getResult();
528-
rewriter.create<memref::StoreOp>(loc, val, sMemRef, index);
527+
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
528+
rewriter.create<affine::AffineStoreOp>(loc, val, sMemRef, zeroMap,
529+
std::nullopt);
529530
rewriter.eraseOp(op);
530531
return success();
531532
}

lib/Conversion/TritonToLinalgExperimental/TritonToLinalgExperimentalPass.cpp

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,9 @@ class TritonToLinalgExperimentalPass
5151
void runOnOperation() override {
5252
auto moduleOp = getOperation();
5353
PassManager pm(&getContext(), moduleOp.getOperationName());
54-
pm.addPass(createTritonToStructuredPass(enableMakeGatherScatterTensorPtr));
54+
55+
pm.addPass(createTritonToStructuredPass(
56+
enableMakeGatherScatterTensorPtr));
5557

5658
// Erase dead code and fold constants created during lowering
5759
pm.addPass(createCSEPass());

lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -100,10 +100,11 @@ struct ScalarLoadConverter : public OpConversionPattern<tts::GatherOp> {
100100
basePtr, getAsOpFoldResult(loadIndex) /*offset*/,
101101
ArrayRef<OpFoldResult>{rewriter.getIndexAttr(1)} /*sizes*/,
102102
ArrayRef<OpFoldResult>{rewriter.getIndexAttr(1)} /*strides*/);
103-
auto index =
104-
rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0))
105-
.getResult();
106-
auto scalarLoadOp = rewriter.create<memref::LoadOp>(loc, memref, index);
103+
104+
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
105+
106+
auto scalarLoadOp = rewriter.create<affine::AffineLoadOp>(
107+
loc, memref, zeroMap, std::nullopt);
107108

108109
rewriter.replaceOp(gatherOp, scalarLoadOp.getResult());
109110

@@ -146,10 +147,10 @@ struct ScalarStoreConverter : public OpConversionPattern<tts::ScatterOp> {
146147
ArrayRef<OpFoldResult>{rewriter.getIndexAttr(1)} /*strides*/);
147148

148149
auto storeVal = scatterOp.getValue();
149-
auto index =
150-
rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0))
151-
.getResult();
152-
rewriter.create<memref::StoreOp>(loc, storeVal, memref, index);
150+
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
151+
152+
rewriter.create<affine::AffineStoreOp>(loc, storeVal, memref, zeroMap,
153+
std::nullopt);
153154
rewriter.eraseOp(scatterOp);
154155

155156
return success();

test/Conversion/StructuredToMemref/addptr_chain.mlir

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,6 @@ module {
2424

2525
// CHECK-LABEL: func.func @addptr
2626
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xf32>, [[PARAM_1_:%.+]]: memref<*xf32>, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32) {
27-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
2827
// CHECK-DAG: [[CST_1_:%.+]] = arith.constant 1 : i32
2928
// CHECK-DAG: [[CST_10_:%.+]] = arith.constant 10 : i32
3029
// CHECK-DAG: [[CST_2_:%.+]] = arith.constant 2 : i32
@@ -34,14 +33,14 @@ module {
3433
// CHECK-DAG: [[VAR_1_:%.+]] = arith.addi [[I_0_]], [[CST_2_]] : i32
3534
// CHECK: [[VAR_2_:%.+]] = arith.index_cast [[VAR_0_]] : i32 to index
3635
// CHECK: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: {{.}}[[VAR_2_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
37-
// CHECK: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = memref.load [[VAR_reinterpret_cast_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
38-
// CHECK: [[VAR_4_:%.+]] = arith.index_cast [[VAR_1_]] : i32 to index
36+
// CHECK: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = affine.load [[VAR_reinterpret_cast_]][0] : memref<1xf32, strided<[1], offset: ?>>
37+
// CHECK: [[VAR_4_:%.+]] = arith.index_cast [[VAR_1_]] : i32 to index
3938
// CHECK: [[VAR_reinterpret_cast_0_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: {{.}}[[VAR_4_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
40-
// CHECK-DAG: [[LOAD_VAR_reinterpret_cast_0_MEM_:%.+]] = memref.load [[VAR_reinterpret_cast_0_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
39+
// CHECK-DAG: [[LOAD_VAR_reinterpret_cast_0_MEM_:%.+]] = affine.load [[VAR_reinterpret_cast_0_]][0] : memref<1xf32, strided<[1], offset: ?>>
4140
// CHECK-DAG: [[VAR_reinterpret_cast_1_:%.+]] = memref.reinterpret_cast [[PARAM_1_]] to offset: {{.}}[[VAR_2_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
42-
// CHECK: memref.store [[LOAD_VAR_reinterpret_cast_MEM_]], [[VAR_reinterpret_cast_1_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
41+
// CHECK: affine.store [[LOAD_VAR_reinterpret_cast_MEM_]], [[VAR_reinterpret_cast_1_]][0] : memref<1xf32, strided<[1], offset: ?>>
4342
// CHECK: [[VAR_reinterpret_cast_2_:%.+]] = memref.reinterpret_cast [[PARAM_1_]] to offset: {{.}}[[VAR_4_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
44-
// CHECK: memref.store [[LOAD_VAR_reinterpret_cast_0_MEM_]], [[VAR_reinterpret_cast_2_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
43+
// CHECK: affine.store [[LOAD_VAR_reinterpret_cast_0_MEM_]], [[VAR_reinterpret_cast_2_]][0] : memref<1xf32, strided<[1], offset: ?>>
4544
// CHECK: }
4645
// CHECK: return
4746
// CHECK: }

test/Conversion/StructuredToMemref/addptr_scalar_loopback.mlir

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,10 @@ module {
1616

1717
// CHECK-LABEL: func.func @kernel
1818
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xbf16>, [[PARAM_1_:%.+]]: memref<*xbf16>, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32, [[PARAM_8_:%.+]]: i32) {
19-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
2019
// CHECK: [[VAR_0_:%.+]] = arith.index_cast [[PARAM_2_]] : i32 to index
2120
// CHECK-DAG: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: {{.}}[[VAR_0_]]{{.}}, sizes: [1], strides: [1] : memref<*xbf16> to memref<1xbf16, strided<[1], offset: ?>>
2221
// CHECK-DAG: [[VAR_reinterpret_cast_0_:%.+]] = memref.reinterpret_cast [[PARAM_1_]] to offset: {{.}}[[VAR_0_]]{{.}}, sizes: [1], strides: [1] : memref<*xbf16> to memref<1xbf16, strided<[1], offset: ?>>
23-
// CHECK-DAG: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = memref.load [[VAR_reinterpret_cast_]][%[[C0]]] : memref<1xbf16, strided<[1], offset: ?>>
24-
// CHECK: memref.store [[LOAD_VAR_reinterpret_cast_MEM_]], [[VAR_reinterpret_cast_0_]][%[[C0]]] : memref<1xbf16, strided<[1], offset: ?>>
22+
// CHECK-DAG: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = affine.load [[VAR_reinterpret_cast_]][0] : memref<1xbf16, strided<[1], offset: ?>>
23+
// CHECK: affine.store [[LOAD_VAR_reinterpret_cast_MEM_]], [[VAR_reinterpret_cast_0_]][0] : memref<1xbf16, strided<[1], offset: ?>>
2524
// CHECK: return
2625
// CHECK: }

test/Conversion/StructuredToMemref/convert_addi_reduce.mlir

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ module {
1515

1616
// CHECK-LABEL: func.func @addi
1717
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xi32>, [[PARAM_1_:%.+]]: i32, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32) {
18-
// CHECK-DAG: %[[CST_0_:.+]] = arith.constant 0 : index
18+
// CHECK-DAG: [[CST_0_:%.+]] = arith.constant 0 : index
1919
// CHECK-DAG: [[CST_0_1_:%.+]] = arith.constant 0 : i32
2020
// CHECK-DAG: [[VAR_0_:%.+]] = tensor.empty() : tensor<4096xi32>
2121
// CHECK-NOT: separator of consecutive DAGs
@@ -28,7 +28,7 @@ module {
2828
// CHECK: linalg.yield [[VAR_3_]] : i32
2929
// CHECK: }
3030
// CHECK-DAG: [[VAR_extracted_:%.+]] = tensor.extract [[VAR_reduced_]][] : tensor<i32>
31-
// CHECK-DAG: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: [%[[CST_0_]]], sizes: [1], strides: [1] : memref<*xi32> to memref<1xi32, strided<[1], offset: ?>>
32-
// CHECK: memref.store [[VAR_extracted_]], [[VAR_reinterpret_cast_]][%[[CST_0_]]] : memref<1xi32, strided<[1], offset: ?>>
31+
// CHECK-DAG: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: {{.}}[[CST_0_]]{{.}}, sizes: [1], strides: [1] : memref<*xi32> to memref<1xi32, strided<[1], offset: ?>>
32+
// CHECK: affine.store [[VAR_extracted_]], [[VAR_reinterpret_cast_]][0] : memref<1xi32, strided<[1], offset: ?>>
3333
// CHECK: return
3434
// CHECK: }

test/Conversion/StructuredToMemref/convert_argmin_argmax.mlir

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,6 @@ module {
3030
// CHECK-DAG: [[MAP_0_:#.+]] = affine_map<(d0) -> (d0)>
3131
// CHECK-LABEL: func.func @argmax_012
3232
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xf32>, [[PARAM_1_:%.+]]: memref<*xi32>, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32, [[PARAM_8_:%.+]]: i32) {
33-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
3433
// CHECK-DAG: [[CST_minus_1_:%.+]] = arith.constant -1 : i32
3534
// CHECK-DAG: [[CST_0_:%.+]] = arith.constant 0xFF800000 : f32
3635
// CHECK-DAG: [[VAR_0_:%.+]] = arith.muli [[PARAM_6_]], [[PARAM_2_]] : i32
@@ -67,7 +66,7 @@ module {
6766
// CHECK-DAG: [[VAR_extracted_:%.+]] = tensor.extract [[VAR_reduced_]]#1[] : tensor<i32>
6867
// CHECK-DAG: [[VAR_9_:%.+]] = arith.index_cast [[PARAM_6_]] : i32 to index
6968
// CHECK: [[VAR_reinterpret_cast_0_:%.+]] = memref.reinterpret_cast [[PARAM_1_]] to offset: {{.}}[[VAR_9_]]{{.}}, sizes: [1], strides: [1] : memref<*xi32> to memref<1xi32, strided<[1], offset: ?>>
70-
// CHECK: memref.store [[VAR_extracted_]], [[VAR_reinterpret_cast_0_]][%[[C0]]] : memref<1xi32, strided<[1], offset: ?>>
69+
// CHECK: affine.store [[VAR_extracted_]], [[VAR_reinterpret_cast_0_]][0] : memref<1xi32, strided<[1], offset: ?>>
7170
// CHECK: return
7271
// CHECK: }
7372

@@ -103,7 +102,6 @@ module {
103102
// CHECK-DAG: [[MAP_0_:#.+]] = affine_map<(d0) -> (d0)>
104103
// CHECK-LABEL: func.func @argmin_012
105104
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xf32>, [[PARAM_1_:%.+]]: memref<*xi32>, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32, [[PARAM_8_:%.+]]: i32) {
106-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
107105
// CHECK-DAG: [[CST_minus_1_:%.+]] = arith.constant -1 : i32
108106
// CHECK-DAG: [[CST_0_:%.+]] = arith.constant 0x7F800000 : f32
109107
// CHECK-DAG: [[VAR_0_:%.+]] = arith.muli [[PARAM_6_]], [[PARAM_2_]] : i32
@@ -140,6 +138,6 @@ module {
140138
// CHECK-DAG: [[VAR_extracted_:%.+]] = tensor.extract [[VAR_reduced_]]#1[] : tensor<i32>
141139
// CHECK-DAG: [[VAR_9_:%.+]] = arith.index_cast [[PARAM_6_]] : i32 to index
142140
// CHECK: [[VAR_reinterpret_cast_0_:%.+]] = memref.reinterpret_cast [[PARAM_1_]] to offset: {{.}}[[VAR_9_]]{{.}}, sizes: [1], strides: [1] : memref<*xi32> to memref<1xi32, strided<[1], offset: ?>>
143-
// CHECK: memref.store [[VAR_extracted_]], [[VAR_reinterpret_cast_0_]][%[[C0]]] : memref<1xi32, strided<[1], offset: ?>>
141+
// CHECK: affine.store [[VAR_extracted_]], [[VAR_reinterpret_cast_0_]][0] : memref<1xi32, strided<[1], offset: ?>>
144142
// CHECK: return
145143
// CHECK: }

test/Conversion/StructuredToMemref/dot.mlir

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ module {
4949
}
5050
}
5151

52-
// CHECK-DAG: [[MAP_0_:#.+]] = affine_map<(d0, d1) -> (d0, d1)>
52+
// CHECK: [[MAP_:#.+]] = affine_map<(d0, d1) -> (d0, d1)>
5353
// CHECK-LABEL: func.func @kernel
5454
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xbf16>, [[PARAM_1_:%.+]]: memref<*xbf16>, [[PARAM_2_:%.+]]: memref<*xbf16>, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32, [[PARAM_8_:%.+]]: i32) {
55-
// CHECK-DAG: [[CST_0_dot_000000_:%.+]] = arith.constant 0.000000e+00 : bf16
5655
// CHECK-NOT: separator of consecutive DAGs
56+
// CHECK: [[CST_0_:%.+]] = arith.constant 0.000000e+00 : bf16
5757
// CHECK-DAG: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: [0], sizes: [128, 64], strides: [128, 1] : memref<*xbf16> to memref<128x64xbf16, strided<[128, 1]>>
5858
// CHECK-DAG: [[RES_:%.+]] = memref.alloc() : memref<128x64xbf16>
5959
// CHECK: memref.copy [[VAR_reinterpret_cast_]], [[RES_]] : memref<128x64xbf16, strided<[128, 1]>> to memref<128x64xbf16>
@@ -68,14 +68,14 @@ module {
6868
// CHECK-DAG: [[VAR_reinterpret_cast_2_:%.+]] = memref.reinterpret_cast [[PARAM_2_]] to offset: [0], sizes: [128, 256], strides: [256, 1] : memref<*xbf16> to memref<128x256xbf16, strided<[256, 1]>>
6969
// CHECK-DAG: [[RES_2_:%.+]] = memref.alloc() : memref<128x256xbf16>
7070
// CHECK: memref.copy [[VAR_reinterpret_cast_2_]], [[RES_2_]] : memref<128x256xbf16, strided<[256, 1]>> to memref<128x256xbf16>
71-
// CHECK-DAG: [[VAR_3_:%.+]] = bufferization.to_tensor [[RES_2_]] restrict writable : memref<128x256xbf16>
72-
// CHECK-DAG: [[VAR_4_:%.+]] = tensor.empty() : tensor<128x256xbf16>
73-
// CHECK: [[VAR_5_:%.+]] = linalg.fill ins([[CST_0_dot_000000_]] : bf16) outs([[VAR_4_]] : tensor<128x256xbf16>) -> tensor<128x256xbf16>
71+
// CHECK: [[VAR_3_:%.+]] = bufferization.to_tensor [[RES_2_]] restrict writable : memref<128x256xbf16>
72+
// CHECK: [[VAR_4_:%.+]] = tensor.empty() : tensor<128x256xbf16>
73+
// CHECK: [[VAR_5_:%.+]] = linalg.fill ins([[CST_0_]] : bf16) outs([[VAR_4_]] : tensor<128x256xbf16>) -> tensor<128x256xbf16>
7474
// CHECK: [[VAR_6_:%.+]] = linalg.matmul ins([[VAR_0_]], [[VAR_transposed_]] : tensor<128x64xbf16>, tensor<64x256xbf16>) outs([[VAR_5_]] : tensor<128x256xbf16>) -> tensor<128x256xbf16>
75-
// CHECK: [[VAR_7_:%.+]] = linalg.generic {indexing_maps = [#map, #map, #map], iterator_types = ["parallel", "parallel"]} ins([[VAR_3_]], [[VAR_6_]] : tensor<128x256xbf16>, tensor<128x256xbf16>) outs([[VAR_3_]] : tensor<128x256xbf16>) {
76-
// CHECK: ^bb0([[IN_0_:%.+]]: bf16, [[IN_1_:%.+]]: bf16, [[IN_2_:%.+]]: bf16):
77-
// CHECK: [[VAR_8_:%.+]] = arith.addf [[IN_0_]], [[IN_1_]] : bf16
78-
// CHECK: linalg.yield [[VAR_8_]] : bf16
75+
// CHECK: [[VAR_7_:%.+]] = linalg.generic {indexing_maps = [[[MAP_]], [[MAP_]], [[MAP_]]], iterator_types = ["parallel", "parallel"]} ins([[VAR_3_]], [[VAR_6_]] : tensor<128x256xbf16>, tensor<128x256xbf16>) outs([[VAR_3_]] : tensor<128x256xbf16>) {
76+
// CHECK: ^bb0([[VAR_in_1:%.+]]: bf16, [[VAR_in_2:%.+]]: bf16, {{%.+}}: bf16):
77+
// CHECK: [[VAR_8_:%.+]] = arith.addf [[VAR_in_1]], [[VAR_in_2]] : bf16
78+
// CHECK: linalg.yield [[VAR_8_:%.+]] : bf16
7979
// CHECK: } -> tensor<128x256xbf16>
8080
// CHECK: bufferization.materialize_in_destination [[VAR_7_]] in writable [[VAR_reinterpret_cast_2_]] : (tensor<128x256xbf16>, memref<128x256xbf16, strided<[256, 1]>>) -> ()
8181
// CHECK: return

test/Conversion/StructuredToMemref/early_return.mlir

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,15 +35,14 @@ module {
3535
// CHECK-DAG: [[MAP_0_:#.+]] = affine_map<(d0) -> (d0)>
3636
// CHECK-LABEL: func.func @test_1
3737
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xf32>, [[PARAM_1_:%.+]]: memref<*xf32>, [[PARAM_2_:%.+]]: i32, [[PARAM_3_:%.+]]: i32, [[PARAM_4_:%.+]]: i32, [[PARAM_5_:%.+]]: i32, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32) {
38-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
3938
// CHECK-DAG: [[CST_1_:%.+]] = arith.constant 1 : i32
4039
// CHECK-DAG: [[CST_minus_1_dot_000000_:%.+]] = arith.constant -1.000000e+00 : f32
4140
// CHECK-DAG: [[VAR_0_:%.+]] = tensor.empty() : tensor<4xi32>
4241
// CHECK-NOT: separator of consecutive DAGs
4342
// CHECK-DAG: [[VAR_1_:%.+]] = linalg.fill ins([[CST_1_]] : i32) outs([[VAR_0_]] : tensor<4xi32>) -> tensor<4xi32>
4443
// CHECK-DAG: [[VAR_2_:%.+]] = arith.index_cast [[PARAM_5_]] : i32 to index
4544
// CHECK: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_0_]] to offset: {{.}}[[VAR_2_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
46-
// CHECK: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = memref.load [[VAR_reinterpret_cast_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
45+
// CHECK: [[LOAD_VAR_reinterpret_cast_MEM_:%.+]] = affine.load [[VAR_reinterpret_cast_]][0] : memref<1xf32, strided<[1], offset: ?>>
4746
// CHECK: [[VAR_4_:%.+]] = arith.cmpf oeq, [[LOAD_VAR_reinterpret_cast_MEM_]], [[CST_minus_1_dot_000000_]] : f32
4847
// CHECK: cf.cond_br [[VAR_4_]], ^bb1, ^bb2
4948
// CHECK: ^bb1: // pred: ^bb0

test/Conversion/StructuredToMemref/kernel-05-layer-norm-fwd.mlir

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,6 @@ module {
9191
// CHECK-DAG: [[MAP_0_:#.+]] = affine_map<(d0) -> (d0)>
9292
// CHECK-LABEL: func.func @_layer_norm_fwd_fused_0123456789
9393
// CHECK-SAME: ([[PARAM_0_:%.+]]: memref<*xf32>, [[PARAM_1_:%.+]]: memref<*xf32>, [[PARAM_2_:%.+]]: memref<*xf32>, [[PARAM_3_:%.+]]: memref<*xf32>, [[PARAM_4_:%.+]]: memref<*xf32>, [[PARAM_5_:%.+]]: memref<*xf32>, [[PARAM_6_:%.+]]: i32, [[PARAM_7_:%.+]]: i32, [[PARAM_8_:%.+]]: f32, [[PARAM_9_:%.+]]: i32, [[PARAM_10_:%.+]]: i32, [[PARAM_11_:%.+]]: i32, [[PARAM_12_:%.+]]: i32, [[PARAM_13_:%.+]]: i32, [[PARAM_14_:%.+]]: i32) {
94-
// CHECK-DAG: %[[C0:.*]] = arith.constant 0 : index
9594
// CHECK-DAG: [[CST_1_dot_000000_:%.+]] = arith.constant 1.000000e+00 : f32
9695
// CHECK-DAG: [[CST_0_:%.+]] = arith.constant 0 : i32
9796
// CHECK-DAG: [[CST_256_:%.+]] = arith.constant 256 : i32
@@ -213,9 +212,9 @@ module {
213212
// CHECK-DAG: [[VAR_17_:%.+]] = arith.divf [[CST_1_dot_000000_]], [[VAR_16_]] : f32
214213
// CHECK-DAG: [[VAR_18_:%.+]] = arith.index_cast [[PARAM_12_]] : i32 to index
215214
// CHECK: [[VAR_reinterpret_cast_:%.+]] = memref.reinterpret_cast [[PARAM_4_]] to offset: {{.}}[[VAR_18_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
216-
// CHECK: memref.store [[VAR_10_]], [[VAR_reinterpret_cast_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
215+
// CHECK: affine.store [[VAR_10_]], [[VAR_reinterpret_cast_]][0] : memref<1xf32, strided<[1], offset: ?>>
217216
// CHECK: [[VAR_reinterpret_cast_4_:%.+]] = memref.reinterpret_cast [[PARAM_5_]] to offset: {{.}}[[VAR_18_]]{{.}}, sizes: [1], strides: [1] : memref<*xf32> to memref<1xf32, strided<[1], offset: ?>>
218-
// CHECK: memref.store [[VAR_17_]], [[VAR_reinterpret_cast_4_]][%[[C0]]] : memref<1xf32, strided<[1], offset: ?>>
217+
// CHECK: affine.store [[VAR_17_]], [[VAR_reinterpret_cast_4_]][0] : memref<1xf32, strided<[1], offset: ?>>
219218
// CHECK: [[VAR_19_:%.+]] = linalg.fill ins([[VAR_17_]] : f32) outs([[VAR_0_]] : tensor<256xf32>) -> tensor<256xf32>
220219
// CHECK: scf.for [[VAR_arg15_1_:%.+]] = [[CST_0_]] to [[PARAM_7_]] step [[CST_256_]] : i32 {
221220
// CHECK: [[VAR_20_5_:%.+]] = arith.index_cast [[VAR_arg15_1_]] : i32 to index

0 commit comments

Comments
 (0)