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