1// RUN: mlir-opt %s --test-transform-dialect-interpreter --split-input-file -verify-diagnostics | FileCheck %s
2// RUN: mlir-opt %s --test-transform-dialect-interpreter --canonicalize --split-input-file -verify-diagnostics | FileCheck %s --check-prefix=CANON
3
4transform.with_pdl_patterns {
5^bb0(%arg0: !pdl.operation):
6  pdl.pattern @linalg_generic : benefit(1) {
7    %0 = pdl.operands
8    %1 = pdl.types
9    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
10    pdl.rewrite %2 with "transform.dialect"
11  }
12
13  transform.sequence %arg0 {
14  ^bb1(%arg1: !pdl.operation):
15    %0 = transform.pdl_match @linalg_generic in %arg1
16    %1:2 = transform.structured.split %0 after 42 { dimension = 0 }
17  }
18}
19
20func.func private @elem(%arg0: f32, %arg1: index, %arg2: index) -> f32
21
22// CHECK: #[[$ADD_42_MAP:.+]] = affine_map<(d0) -> (d0 + 42)>
23// CHECK: #[[$ADD_10_MAP:.+]] = affine_map<(d0) -> (d0 + 10)>
24
25// CHECK-LABEL: @one_d_static
26// CHECK-SAME:  %[[IN:.+]]: tensor<100xf32>, %[[OUT:.+]]: tensor<100xf32>
27func.func @one_d_static(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
28  // CHECK: %[[IN_SLICE_LOW:.+]] = tensor.extract_slice %[[IN]][0] [42] [1] : tensor<100xf32> to tensor<42xf32>
29  // CHECK: %[[OUT_SLICE_LOW:.+]] = tensor.extract_slice %[[OUT]][0] [42] [1] : tensor<100xf32> to tensor<42xf32>
30  // CHECK: %[[RES_SLICE_LOW:.+]] = linalg.generic
31  // CHECK:   ins(%[[IN_SLICE_LOW]]
32  // CHECK:   outs(%[[OUT_SLICE_LOW]]
33  // CHECK:   linalg.index 0
34  // CHECK:   func.call @elem
35  // CHECK: %[[RES_PARTIAL:.+]] = tensor.insert_slice %[[RES_SLICE_LOW]] into %[[OUT]][0] [42] [1]
36  //
37  // CHECK: %[[IN_SLICE_HIGH:.+]] = tensor.extract_slice %[[IN]][42] [58] [1] : tensor<100xf32> to tensor<58xf32>
38  // CHECK: %[[OUT_SLICE_HIGH:.+]] = tensor.extract_slice %[[RES_PARTIAL]][42] [58] [1] : tensor<100xf32> to tensor<58xf32>
39  // CHECK: %[[RES_SLICE_HIGH:.+]] = linalg.generic
40  // CHECK:   ins(%[[IN_SLICE_HIGH]]
41  // CHECK:   outs(%[[OUT_SLICE_HIGH]]
42  // CHECK:   %[[IDX:.+]] = linalg.index 0
43  // CHECK:   affine.apply #[[$ADD_42_MAP]](%[[IDX]])
44  // CHECK:   func.call @elem
45  // CHECK: %[[RES:.+]] = tensor.insert_slice %[[RES_SLICE_HIGH]] into %[[RES_PARTIAL]][42] [58] [1]
46  %0 = linalg.generic {
47    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
48    iterator_types = ["parallel"]
49  }
50  ins(%arg0: tensor<100xf32>) outs(%arg1: tensor<100xf32>) {
51  ^bb0(%0: f32, %1: f32):
52    %i = linalg.index 0 : index
53    %call_res = func.call @elem(%0, %i, %i) : (f32, index, index) -> f32
54    linalg.yield %call_res : f32
55  } -> tensor<100xf32>
56
57  // CHECK: return %[[RES]]
58  return %0 : tensor<100xf32>
59}
60
61// CHECK-LABEL: @one_d_static_overflow
62// CHECK-SAME:  %[[IN:.+]]: tensor<10xf32>, %[[OUT:.+]]: tensor<10xf32>
63// CANON-LABEL:  @one_d_static_overflow
64// CANON-SAME:  %[[IN:.+]]: tensor<10xf32>, %[[OUT:.+]]: tensor<10xf32>
65func.func @one_d_static_overflow(%arg0: tensor<10xf32>, %arg1: tensor<10xf32>) -> tensor<10xf32> {
66  // CHECK: %[[IN_SLICE_LOW:.+]] = tensor.extract_slice %[[IN]][0] [10] [1] : tensor<10xf32> to tensor<10xf32>
67  // CHECK: %[[OUT_SLICE_LOW:.+]] = tensor.extract_slice %[[OUT]][0] [10] [1] : tensor<10xf32> to tensor<10xf32>
68  // CHECK: %[[RES_SLICE_LOW:.+]] = linalg.generic
69  // CHECK:   ins(%[[IN_SLICE_LOW]]
70  // CHECK:   outs(%[[OUT_SLICE_LOW]]
71  // CHECK:   linalg.index 0
72  // CHECK:   func.call @elem
73  // CHECK: %[[RES_PARTIAL:.+]] = tensor.insert_slice %[[RES_SLICE_LOW]] into %[[OUT]][0] [10] [1]
74  //
75  // Due to overflow, the first part of the split computes everything and the
76  // insert/extract slices are folded away by the canonicalizer.
77  // CANON: %[[RES_PARTIAL:.+]] = linalg.generic
78  // CANON:   ins(%[[IN]]
79  // CANON:   outs(%[[OUT]]
80  // CANON:   linalg.index 0
81  // CANON:   func.call @elem
82  // The second part operates on zero-sized slices that are not currently
83  // folded away.
84  //
85  // CHECK: %[[IN_SLICE_HIGH:.+]] = tensor.extract_slice %[[IN]][10] [0] [1] : tensor<10xf32> to tensor<0xf32>
86  // CHECK: %[[OUT_SLICE_HIGH:.+]] = tensor.extract_slice %[[RES_PARTIAL]][10] [0] [1] : tensor<10xf32> to tensor<0xf32>
87  // CHECK: %[[RES_SLICE_HIGH:.+]] = linalg.generic
88  // CHECK:   ins(%[[IN_SLICE_HIGH]]
89  // CHECK:   outs(%[[OUT_SLICE_HIGH]]
90  // CHECK:   %[[IDX:.+]] = linalg.index 0
91  // CHECK:   affine.apply #[[$ADD_10_MAP]](%[[IDX]])
92  // CHECK:   func.call @elem
93  // CHECK: %[[RES:.+]] = tensor.insert_slice %[[RES_SLICE_HIGH]] into %[[RES_PARTIAL]][10] [0] [1]
94  %0 = linalg.generic {
95    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
96    iterator_types = ["parallel"]
97  }
98  ins(%arg0: tensor<10xf32>) outs(%arg1: tensor<10xf32>) {
99  ^bb0(%0: f32, %1: f32):
100    %i = linalg.index 0 : index
101    %call_res = func.call @elem(%0, %i, %i) : (f32, index, index) -> f32
102    linalg.yield %call_res : f32
103  } -> tensor<10xf32>
104  return %0 : tensor<10xf32>
105}
106
107// -----
108
109transform.with_pdl_patterns {
110^bb0(%arg0: !pdl.operation):
111  pdl.pattern @func_call : benefit(1) {
112    %0 = pdl.operands
113    %1 = pdl.types
114    %2 = pdl.operation "func.call"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
115    pdl.rewrite %2 with "transform.dialect"
116  }
117  pdl.pattern @linalg_generic : benefit(1) {
118    %0 = pdl.operands
119    %1 = pdl.types
120    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
121    pdl.rewrite %2 with "transform.dialect"
122  }
123
124  transform.sequence %arg0 {
125  ^bb1(%arg1: !pdl.operation):
126    %0 = transform.pdl_match @linalg_generic in %arg1
127    %1 = transform.pdl_match @func_call in %arg1
128    transform.structured.split %0 after %1 { dimension = 0 }
129  }
130}
131
132func.func private @get_size() -> index
133
134// CHECK: #[[$MAP_MIN_100:.+]] = affine_map<()[s0] -> (s0, 100)>
135// CHECK: #[[$MAP_S_MINUS_100:.+]] = affine_map<()[s0] -> (-s0 + 100)>
136
137// CHECK-LABEL: @dynamic
138func.func @dynamic(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
139  // CHECK: %[[SPLIT:.+]] = call @get_size
140  // CHECK: %[[SPLIT_LOW:.+]] = affine.min #[[$MAP_MIN_100]]()[%[[SPLIT]]
141  // CHECK: %[[IN_SLICE_LOW:.+]] = tensor.extract_slice %[[IN:.+]][0] [%[[SPLIT_LOW]]] [1] : tensor<100xf32> to tensor<?xf32>
142  // CHECK: %[[OUT_SLICE_LOW:.+]] = tensor.extract_slice %[[OUT:.+]][0] [%[[SPLIT_LOW]]] [1] : tensor<100xf32> to tensor<?xf32>
143  // CHECK: %[[RES_SLICE_LOW:.+]] = linalg.generic
144  // CHECK:   ins(%[[IN_SLICE_LOW]]
145  // CHECK:   outs(%[[OUT_SLICE_LOW]]
146  // CHECK: %[[PARTIAL:.+]] = tensor.insert_slice %[[RES_SLICE_LOW]] into %[[OUT]][0] [%[[SPLIT_LOW]]] [1]
147  //
148  // CHECK: %[[SPLIT_HIGH_1:.+]] = affine.apply #[[$MAP_S_MINUS_100]]()[%[[SPLIT_LOW]]]
149  // CHECK: %[[SPLIT_HIGH_2:.+]] = affine.apply #[[$MAP_S_MINUS_100]]()[%[[SPLIT_LOW]]]
150  // CHECK: %[[IN_SLICE_HIGH:.+]] = tensor.extract_slice %[[IN:.+]][%[[SPLIT_LOW]]] [%[[SPLIT_HIGH_2]]] [1] : tensor<100xf32> to tensor<?xf32>
151  // CHECK: %[[SPLIT_HIGH_3:.+]] = affine.apply #[[$MAP_S_MINUS_100]]()[%[[SPLIT_LOW]]]
152  // CHECK: %[[OUT_SLICE_HIGH:.+]] = tensor.extract_slice %[[PARTIAL:.+]][%[[SPLIT_LOW]]] [%[[SPLIT_HIGH_3]]] [1] : tensor<100xf32> to tensor<?xf32>
153  // CHECK: %[[RES_SLICE_HIGH:.+]] = linalg.generic
154  // CHECK:   ins(%[[IN_SLICE_HIGH]]
155  // CHECK:   outs(%[[OUT_SLICE_HIGH]]
156  // CHECK: tensor.insert_slice %[[RES_SLICE_HIGH]] into %[[PARTIAL]][%[[SPLIT_LOW]]] [%[[SPLIT_HIGH_3]]] [1]
157  %0 = func.call @get_size() : () -> index
158  %1 = linalg.generic {
159    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
160    iterator_types = ["parallel"]
161  }
162  ins(%arg0: tensor<100xf32>) outs(%arg1: tensor<100xf32>) {
163  ^bb0(%3: f32, %4: f32):
164    %5 = arith.addf %3, %4 : f32
165    linalg.yield %5 : f32
166  } -> tensor<100xf32>
167  return %1 : tensor<100xf32>
168}
169
170// -----
171
172transform.with_pdl_patterns {
173^bb0(%arg0: !pdl.operation):
174  pdl.pattern @linalg_generic : benefit(1) {
175    %0 = pdl.operands
176    %1 = pdl.types
177    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
178    pdl.rewrite %2 with "transform.dialect"
179  }
180
181  transform.sequence %arg0 {
182  ^bb1(%arg1: !pdl.operation):
183    %0 = transform.pdl_match @linalg_generic in %arg1
184    %1:2 = transform.structured.split %0 after 4 { dimension = 0}
185    %2:2 = transform.structured.split %1#1 after 16 { dimension = 1 }
186  }
187}
188
189func.func private @elem(%arg0: f32, %arg1: index, %arg2: index) -> f32
190
191// CHECK-LABEL: @two_d
192func.func @two_d(%arg0: tensor<10x34xf32>,
193                 %arg1: tensor<10x34xf32>) -> tensor<10x34xf32> {
194  // Check the overall structure: split along the dimension 0, and then split
195  // the second half only along the dimension 1.
196  // CHECK:      %[[IN_1:.+]] = tensor.extract_slice %[[IN:.+]][0, 0]
197  // CHECK:      %[[OUT_1:.+]] = tensor.extract_slice %[[OUT:.+]][0, 0]
198  // CHECK:      %[[RES_1:.+]] = linalg.generic
199  // CHECK-SAME:   ins(%[[IN_1]] : tensor<4x34xf32>)
200  // CHECK-SAME:   outs(%[[OUT_1]] : tensor<4x34xf32>)
201  // CHECK:      %[[PARTIAL_1:.+]] = tensor.insert_slice %[[RES_1]] into %[[OUT]]
202  //
203  // CHECK:      %[[IN_2:.+]] = tensor.extract_slice %[[IN]]
204  // CHECK:      %[[OUT_2:.+]] = tensor.extract_slice %[[PARTIAL_1]]
205  // CHECK:      %[[IN_21:.+]] = tensor.extract_slice %[[IN_2]]
206  // CHECK:      %[[OUT_21:.+]] = tensor.extract_slice %[[OUT_2]]
207  // CHECK:      %[[RES_21:.+]] = linalg.generic
208  // CHECK-SAME:   ins(%[[IN_21]] : tensor<6x16xf32>)
209  // CHECK-SAME:   outs(%[[OUT_21]] : tensor<6x16xf32>)
210  // CHECK:      %[[PARTIAL_21:.+]] = tensor.insert_slice %[[RES_21]] into %[[OUT_2]]
211  //
212  // CHECK:      %[[IN_22:.+]] = tensor.extract_slice %[[IN_2]]
213  // CHECK:      %[[OUT_22:.+]] = tensor.extract_slice %[[PARTIAL_21]]
214  // CHECK:      %[[RES_22:.+]] = linalg.generic
215  // CHECK-SAME:   ins(%[[IN_22]] : tensor<6x18xf32>)
216  // CHECK-SAME:   outs(%[[OUT_22]] : tensor<6x18xf32>)
217  // CHECK:      %[[PARTIAL_22:.+]] = tensor.insert_slice %[[RES_22]] into %[[PARTIAL_21]]
218  // CHECK:      %[[PARTIAL_2:.+]] = tensor.insert_slice %[[PARTIAL_22]] into %[[PARTIAL_1]]
219  %0 = linalg.generic {
220    indexing_maps = [affine_map<(i, j) -> (i, j)>,
221                     affine_map<(i, j) -> (i, j)>],
222    iterator_types = ["parallel", "parallel"]
223  }
224  ins(%arg0: tensor<10x34xf32>)
225  outs(%arg1: tensor<10x34xf32>) {
226  ^bb0(%0: f32, %1: f32):
227    %i = linalg.index 0 : index
228    %j = linalg.index 1 : index
229    %call_res = func.call @elem(%0, %i, %j) : (f32, index, index) -> f32
230    linalg.yield %call_res : f32
231  } -> tensor<10x34xf32>
232  return %0 : tensor<10x34xf32>
233}
234
235// -----
236
237transform.sequence {
238^bb1(%arg1: !pdl.operation):
239  // expected-error @below {{expects either a dynamic or a static split point to be provided}}
240  %0:2 = "transform.structured.split"(%arg1) { dimension = 1, static_split_point = -1 } : (!pdl.operation) -> (!pdl.operation, !pdl.operation)
241}
242
243// -----
244
245transform.with_pdl_patterns {
246^bb0(%arg0: !pdl.operation):
247  pdl.pattern @func_call : benefit(1) {
248    %0 = pdl.operands
249    %1 = pdl.types
250    %2 = pdl.operation "func.call"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
251    pdl.rewrite %2 with "transform.dialect"
252  }
253  pdl.pattern @linalg_generic : benefit(1) {
254    %0 = pdl.operands
255    %1 = pdl.types
256    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
257    pdl.rewrite %2 with "transform.dialect"
258  }
259
260  transform.sequence %arg0 {
261  ^bb1(%arg1: !pdl.operation):
262    %0 = transform.pdl_match @linalg_generic in %arg1
263    %1 = transform.pdl_match @func_call in %arg1
264    // expected-error @below {{expected dynamic split point handle to point to a single-result index-typed op}}
265    transform.structured.split %0 after %1 { dimension = 0 }
266  }
267}
268
269func.func private @get_size() -> i64
270
271func.func @dynamic(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
272  // expected-note @below {{dynamic split point}}
273  %0 = func.call @get_size() : () -> i64
274  %1 = linalg.generic {
275    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
276    iterator_types = ["parallel"]
277  }
278  ins(%arg0: tensor<100xf32>) outs(%arg1: tensor<100xf32>) {
279  ^bb0(%3: f32, %4: f32):
280    linalg.yield %3 : f32
281  } -> tensor<100xf32>
282  return %1 : tensor<100xf32>
283}
284
285// -----
286
287transform.with_pdl_patterns {
288^bb0(%arg0: !pdl.operation):
289  pdl.pattern @func_call : benefit(1) {
290    %0 = pdl.operands
291    %1 = pdl.types
292    %2 = pdl.operation "func.call"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
293    pdl.rewrite %2 with "transform.dialect"
294  }
295  pdl.pattern @linalg_generic : benefit(1) {
296    %0 = pdl.operands
297    %1 = pdl.types
298    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
299    pdl.rewrite %2 with "transform.dialect"
300  }
301
302  transform.sequence %arg0 {
303  ^bb1(%arg1: !pdl.operation):
304    %0 = transform.pdl_match @linalg_generic in %arg1
305    %1 = transform.pdl_match @func_call in %arg1
306    // expected-error @below {{expected the dynamic split point handle to point to as many operations (0) as the target handle (1)}}
307    transform.structured.split %0 after %1 { dimension = 0 }
308  }
309}
310
311func.func private @get_size() -> i64
312
313func.func @dynamic(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
314  %1 = linalg.generic {
315    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
316    iterator_types = ["parallel"]
317  }
318  ins(%arg0: tensor<100xf32>) outs(%arg1: tensor<100xf32>) {
319  ^bb0(%3: f32, %4: f32):
320    linalg.yield %3 : f32
321  } -> tensor<100xf32>
322  return %1 : tensor<100xf32>
323}
324
325// -----
326
327transform.with_pdl_patterns {
328^bb0(%arg0: !pdl.operation):
329  pdl.pattern @func_return : benefit(1) {
330    %0 = pdl.operands
331    %1 = pdl.types
332    %2 = pdl.operation "func.return"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
333    pdl.rewrite %2 with "transform.dialect"
334  }
335
336  transform.sequence %arg0 {
337  ^bb1(%arg1: !pdl.operation):
338    %0 = transform.pdl_match @func_return in %arg1
339    // expected-error @below {{only applies to structured ops}}
340    transform.structured.split %0 after 16 { dimension = 1 }
341  }
342}
343
344func.func @noop(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
345  // expected-note @below {{target op}}
346  return %arg0 : tensor<100xf32>
347}
348
349// -----
350
351transform.with_pdl_patterns {
352^bb0(%arg0: !pdl.operation):
353  pdl.pattern @linalg_generic : benefit(1) {
354    %0 = pdl.operands
355    %1 = pdl.types
356    %2 = pdl.operation "linalg.generic"(%0 : !pdl.range<value>) -> (%1 : !pdl.range<type>)
357    pdl.rewrite %2 with "transform.dialect"
358  }
359
360  transform.sequence %arg0 {
361  ^bb1(%arg1: !pdl.operation):
362    %0 = transform.pdl_match @linalg_generic in %arg1
363    // expected-error @below {{dimension 1 does not exist in target op}}
364    transform.structured.split %0 after 16 { dimension = 1 }
365  }
366}
367
368func.func @one_d_static(%arg0: tensor<100xf32>, %arg1: tensor<100xf32>) -> tensor<100xf32> {
369  // expected-note @below {{target op}}
370  %0 = linalg.generic {
371    indexing_maps = [affine_map<(i) -> (i)>, affine_map<(i) -> (i)>],
372    iterator_types = ["parallel"]
373  }
374  ins(%arg0: tensor<100xf32>) outs(%arg1: tensor<100xf32>) {
375  ^bb0(%0: f32, %1: f32):
376    linalg.yield %0 : f32
377  } -> tensor<100xf32>
378  return %0 : tensor<100xf32>
379}
380
381