1// RUN: mlir-opt --test-transform-dialect-interpreter --canonicalize %s | FileCheck %s 2 3transform.with_pdl_patterns { 4^bb0(%arg0: !pdl.operation): 5 // This implements a 2D multisize tiling with target sizes [3, 10]. 6 transform.sequence %arg0 { 7 ^bb1(%arg1: !pdl.operation): 8 %0 = transform.structured.match ops{["linalg.generic"]} in %arg1 9 %1:3 = transform.structured.multitile_sizes %0 { dimension = 0, target_size = 3} 10 %t:3 = transform.structured.multitile_sizes %0 { dimension = 1, target_size = 10} 11 %2:2 = transform.structured.split %0 after %1#2 { dimension = 0 } 12 %3:2 = transform.structured.tile %2#0 [%1#0] 13 %4:2 = transform.structured.tile %2#1 [%1#1] 14 %5 = merge_handles %3#0, %4#0 15 %tt:3 = replicate num(%5) %t#0, %t#1, %t#2 16 %6:2 = transform.structured.split %5 after %tt#2 { dimension = 1 } 17 transform.structured.tile %6#0 [0, %tt#0] 18 transform.structured.tile %6#1 [0, %tt#1] 19 } 20} 21 22func.func private @elem(%arg0: f32, %arg1: index, %arg2: index) -> f32 23 24// CHECK-LABEL: @two_d 25// CHECK-SAME: %[[IN:.+]]: tensor<10x34xf32>, %[[OUT:.+]]: tensor<10x34xf32> 26func.func @two_d(%arg0: tensor<10x34xf32>, 27 %arg1: tensor<10x34xf32>) -> tensor<10x34xf32> { 28 %0 = linalg.generic { 29 indexing_maps = [affine_map<(i, j) -> (i, j)>, 30 affine_map<(i, j) -> (i, j)>], 31 iterator_types = ["parallel", "parallel"] 32 } 33 ins(%arg0: tensor<10x34xf32>) 34 outs(%arg1: tensor<10x34xf32>) { 35 ^bb0(%0: f32, %1: f32): 36 %i = linalg.index 0 : index 37 %j = linalg.index 1 : index 38 %call_res = func.call @elem(%0, %i, %j) : (f32, index, index) -> f32 39 linalg.yield %call_res : f32 40 } -> tensor<10x34xf32> 41 42 // 2D multi-size tiling should produce for quadrants with sizes 43 // (2, 8), (2, 9), (3, 8), (3, 9) 44 // respectively, and in this order. 45 // Check the full code for the first quadrant, the data flow for the second 46 // quadrant and only the overall code structure for the remaining quadrants. 47 // The canonicalizer is able to recover static shapes of for linalg.generic 48 // instances, use those to differentiate the quadrants. 49 50 // CHECK: %[[SLICE_1:.+]] = tensor.extract_slice %[[OUT]][0, 0] [4, 34] [1, 1] 51 // CHECK: scf.for %[[I1:.+]] = %{{.*}} to %{{.*}} step %{{.*}} iter_args(%[[ITERARG_1:.+]] = %[[SLICE_1]]) 52 // CHECK: %[[INSLICE_1:.+]] = tensor.extract_slice %[[IN]][%[[I1]], 0] [2, 34] [1, 1] 53 // CHECK: %[[OUTSLICE_1:.+]] = tensor.extract_slice %[[ITERARG_1]][%[[I1]], 0] [2, 34] [1, 1] 54 55 // CHECK: %[[SLICE_2:.+]] = tensor.extract_slice %[[OUTSLICE_1]][0, 0] [2, 16] [1, 1] 56 // CHECK: %[[LOOPRES:.+]] = scf.for %[[I2:.+]] = %{{.*}} to %{{.*}} step %{{.*}} iter_args(%[[ITERARG_2:.+]] = %[[SLICE_2]]) 57 // CHECK: %[[INSLICE_2:.+]] = tensor.extract_slice %[[INSLICE_1]][0, %[[I2]]] [2, 8] [1, 1] 58 // CHECK: %[[OUTSLICE_2:.+]] = tensor.extract_slice %[[ITERARG_2]][0, %[[I2]]] [2, 8] [1, 1] 59 // CHECK: %[[RESSLICE_1:.+]] = linalg.generic {{.*}} ins(%[[INSLICE_2]] : tensor<2x8xf32>) outs(%[[OUTSLICE_2]] : tensor<2x8xf32>) 60 // CHECK: %[[RESPARTIAL:.+]] = tensor.insert_slice %[[RESSLICE_1]] into %[[ITERARG_2]] 61 // CHECK: scf.yield %[[RESPARTIAL]] 62 63 // CHECK: %[[INSERTED:.+]] = tensor.insert_slice %[[LOOPRES]] into %[[OUTSLICE_1]][0, 0] [2, 16] [1, 1] 64 // CHECK: %[[OUTSLICE_3:.+]] = tensor.extract_slice %[[INSERTED]][0, 16] [2, 18] [1, 1] 65 // CHECK: scf.for %{{.*}} iter_args(%{{.*}} = %[[OUTSLICE_3]]) 66 // CHECK-COUNT-2: tensor.extract_slice 67 // CHECK: linalg.generic {{.*}} ins(%{{.*}} : tensor<2x9xf32>) 68 // CHECK: tensor.insert_slice 69 // CHECK: scf.yield 70 // CHECK: %[[INSERTED_2:.+]] = tensor.insert_slice %{{.*}} into %[[INSERTED]] 71 // CHECK: %[[INSERTED_3:.+]] = tensor.insert_slice %[[INSERTED_2]] into %[[ITERARG_1]] 72 // CHECK: scf.yield %[[INSERTED_3]] 73 74 // CHECK: tensor.insert_slice 75 // CHECK: tensor.extract_slice 76 // CHECK: scf.for 77 // CHECK-COUNT-3: tensor.extract_slice 78 // CHECK: scf.for 79 // CHECK-COUNT-2: tensor.extract_slice 80 // CHECK: linalg.generic {{.*}} ins(%{{.*}} : tensor<3x8xf32>) 81 // CHECK: tensor.insert_slice 82 // CHECK: scf.yield 83 // CHECK: tensor.insert_slice 84 // CHECK: tensor.extract_slice 85 // CHECK: scf.for 86 // CHECK-COUNT-2: tensor.extract_slice 87 // CHECK: linalg.generic {{.*}} ins(%{{.*}} : tensor<3x9xf32>) 88 // CHECK: tensor.insert_slice 89 // CHECK: scf.yield 90 // CHECK-COUNT-2: tensor.insert_slice 91 // CHECK: scf.yield 92 // CHECK: %[[RESULT:.+]] = tensor.insert_slice 93 // CHECK: return %[[RESULT]] 94 95 return %0 : tensor<10x34xf32> 96} 97