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