1// RUN: mlir-opt %s \
2// RUN:   --sparsification --sparse-tensor-conversion \
3// RUN:   --convert-vector-to-scf --convert-scf-to-std \
4// RUN:   --func-bufferize --tensor-constant-bufferize --tensor-bufferize \
5// RUN:   --std-bufferize --finalizing-bufferize  \
6// RUN:   --convert-vector-to-llvm --convert-memref-to-llvm --convert-std-to-llvm | \
7// RUN: mlir-cpu-runner \
8// RUN:  -e entry -entry-point-result=void  \
9// RUN:  -shared-libs=%mlir_integration_test_dir/libmlir_c_runner_utils%shlibext | \
10// RUN: FileCheck %s
11
12#Tensor1  = #sparse_tensor.encoding<{
13  dimLevelType = [ "compressed", "compressed", "compressed" ],
14  dimOrdering = affine_map<(i,j,k) -> (i,j,k)>
15}>
16
17#Tensor2  = #sparse_tensor.encoding<{
18  dimLevelType = [ "compressed", "compressed", "compressed" ],
19  dimOrdering = affine_map<(i,j,k) -> (j,k,i)>
20}>
21
22#Tensor3  = #sparse_tensor.encoding<{
23  dimLevelType = [ "compressed", "compressed", "compressed" ],
24  dimOrdering = affine_map<(i,j,k) -> (k,i,j)>
25}>
26
27//
28// Integration test that tests conversions between sparse tensors.
29//
30module {
31  func private @exit(index) -> ()
32
33  //
34  // Verify utilities.
35  //
36  func @checkf64(%arg0: memref<?xf64>, %arg1: memref<?xf64>) {
37    %c0 = constant 0 : index
38    %c1 = constant 1 : index
39    // Same lengths?
40    %0 = memref.dim %arg0, %c0 : memref<?xf64>
41    %1 = memref.dim %arg1, %c0 : memref<?xf64>
42    %2 = cmpi ne, %0, %1 : index
43    scf.if %2 {
44      call @exit(%c1) : (index) -> ()
45    }
46    // Same content?
47    scf.for %i = %c0 to %0 step %c1 {
48      %a = memref.load %arg0[%i] : memref<?xf64>
49      %b = memref.load %arg1[%i] : memref<?xf64>
50      %c = cmpf une, %a, %b : f64
51      scf.if %c {
52        call @exit(%c1) : (index) -> ()
53      }
54    }
55    return
56  }
57  func @check(%arg0: memref<?xindex>, %arg1: memref<?xindex>) {
58    %c0 = constant 0 : index
59    %c1 = constant 1 : index
60    // Same lengths?
61    %0 = memref.dim %arg0, %c0 : memref<?xindex>
62    %1 = memref.dim %arg1, %c0 : memref<?xindex>
63    %2 = cmpi ne, %0, %1 : index
64    scf.if %2 {
65      call @exit(%c1) : (index) -> ()
66    }
67    // Same content?
68    scf.for %i = %c0 to %0 step %c1 {
69      %a = memref.load %arg0[%i] : memref<?xindex>
70      %b = memref.load %arg1[%i] : memref<?xindex>
71      %c = cmpi ne, %a, %b : index
72      scf.if %c {
73        call @exit(%c1) : (index) -> ()
74      }
75    }
76    return
77  }
78
79  //
80  // Output utility.
81  //
82  func @dumpf64(%arg0: memref<?xf64>) {
83    %c0 = constant 0 : index
84    %d0 = constant 0.0 : f64
85    %0 = vector.transfer_read %arg0[%c0], %d0: memref<?xf64>, vector<24xf64>
86    vector.print %0 : vector<24xf64>
87    return
88  }
89
90  //
91  // Main driver.
92  //
93  func @entry() {
94    %c0 = constant 0 : index
95    %c1 = constant 1 : index
96    %c2 = constant 2 : index
97
98    //
99    // Initialize a 3-dim dense tensor.
100    //
101    %t = constant dense<[
102       [  [  1.0,  2.0,  3.0,  4.0 ],
103          [  5.0,  6.0,  7.0,  8.0 ],
104          [  9.0, 10.0, 11.0, 12.0 ] ],
105       [  [ 13.0, 14.0, 15.0, 16.0 ],
106          [ 17.0, 18.0, 19.0, 20.0 ],
107          [ 21.0, 22.0, 23.0, 24.0 ] ]
108    ]> : tensor<2x3x4xf64>
109
110    //
111    // Convert dense tensor directly to various sparse tensors.
112    //    tensor1: stored as 2x3x4
113    //    tensor2: stored as 3x4x2
114    //    tensor3: stored as 4x2x3
115    //
116    %1 = sparse_tensor.convert %t : tensor<2x3x4xf64> to tensor<2x3x4xf64, #Tensor1>
117    %2 = sparse_tensor.convert %t : tensor<2x3x4xf64> to tensor<2x3x4xf64, #Tensor2>
118    %3 = sparse_tensor.convert %t : tensor<2x3x4xf64> to tensor<2x3x4xf64, #Tensor3>
119
120    //
121    // Convert sparse tensor to various sparse tensors. Note that the result
122    // should always correspond to the direct conversion, since the sparse
123    // tensor formats have the ability to restore into the original ordering.
124    //
125    %a = sparse_tensor.convert %1 : tensor<2x3x4xf64, #Tensor1> to tensor<2x3x4xf64, #Tensor1>
126    %b = sparse_tensor.convert %2 : tensor<2x3x4xf64, #Tensor2> to tensor<2x3x4xf64, #Tensor1>
127    %c = sparse_tensor.convert %3 : tensor<2x3x4xf64, #Tensor3> to tensor<2x3x4xf64, #Tensor1>
128    %d = sparse_tensor.convert %1 : tensor<2x3x4xf64, #Tensor1> to tensor<2x3x4xf64, #Tensor2>
129    %e = sparse_tensor.convert %2 : tensor<2x3x4xf64, #Tensor2> to tensor<2x3x4xf64, #Tensor2>
130    %f = sparse_tensor.convert %3 : tensor<2x3x4xf64, #Tensor3> to tensor<2x3x4xf64, #Tensor2>
131    %g = sparse_tensor.convert %1 : tensor<2x3x4xf64, #Tensor1> to tensor<2x3x4xf64, #Tensor3>
132    %h = sparse_tensor.convert %2 : tensor<2x3x4xf64, #Tensor2> to tensor<2x3x4xf64, #Tensor3>
133    %i = sparse_tensor.convert %3 : tensor<2x3x4xf64, #Tensor3> to tensor<2x3x4xf64, #Tensor3>
134
135    //
136    // Check values equality.
137    //
138
139    %v1 = sparse_tensor.values %1 : tensor<2x3x4xf64, #Tensor1> to memref<?xf64>
140    %v2 = sparse_tensor.values %2 : tensor<2x3x4xf64, #Tensor2> to memref<?xf64>
141    %v3 = sparse_tensor.values %3 : tensor<2x3x4xf64, #Tensor3> to memref<?xf64>
142
143    %av = sparse_tensor.values %a : tensor<2x3x4xf64, #Tensor1> to memref<?xf64>
144    %bv = sparse_tensor.values %b : tensor<2x3x4xf64, #Tensor1> to memref<?xf64>
145    %cv = sparse_tensor.values %c : tensor<2x3x4xf64, #Tensor1> to memref<?xf64>
146    %dv = sparse_tensor.values %d : tensor<2x3x4xf64, #Tensor2> to memref<?xf64>
147    %ev = sparse_tensor.values %e : tensor<2x3x4xf64, #Tensor2> to memref<?xf64>
148    %fv = sparse_tensor.values %f : tensor<2x3x4xf64, #Tensor2> to memref<?xf64>
149    %gv = sparse_tensor.values %g : tensor<2x3x4xf64, #Tensor3> to memref<?xf64>
150    %hv = sparse_tensor.values %h : tensor<2x3x4xf64, #Tensor3> to memref<?xf64>
151    %iv = sparse_tensor.values %i : tensor<2x3x4xf64, #Tensor3> to memref<?xf64>
152
153    call @checkf64(%v1, %av) : (memref<?xf64>, memref<?xf64>) -> ()
154    call @checkf64(%v1, %bv) : (memref<?xf64>, memref<?xf64>) -> ()
155    call @checkf64(%v1, %cv) : (memref<?xf64>, memref<?xf64>) -> ()
156    call @checkf64(%v2, %dv) : (memref<?xf64>, memref<?xf64>) -> ()
157    call @checkf64(%v2, %ev) : (memref<?xf64>, memref<?xf64>) -> ()
158    call @checkf64(%v2, %fv) : (memref<?xf64>, memref<?xf64>) -> ()
159    call @checkf64(%v3, %gv) : (memref<?xf64>, memref<?xf64>) -> ()
160    call @checkf64(%v3, %hv) : (memref<?xf64>, memref<?xf64>) -> ()
161    call @checkf64(%v3, %iv) : (memref<?xf64>, memref<?xf64>) -> ()
162
163    //
164    // Check index equality.
165    //
166
167    %v10 = sparse_tensor.indices %1, %c0 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
168    %v11 = sparse_tensor.indices %1, %c1 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
169    %v12 = sparse_tensor.indices %1, %c2 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
170    %v20 = sparse_tensor.indices %2, %c0 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
171    %v21 = sparse_tensor.indices %2, %c1 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
172    %v22 = sparse_tensor.indices %2, %c2 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
173    %v30 = sparse_tensor.indices %3, %c0 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
174    %v31 = sparse_tensor.indices %3, %c1 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
175    %v32 = sparse_tensor.indices %3, %c2 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
176
177    %a10 = sparse_tensor.indices %a, %c0 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
178    %a11 = sparse_tensor.indices %a, %c1 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
179    %a12 = sparse_tensor.indices %a, %c2 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
180    %b10 = sparse_tensor.indices %b, %c0 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
181    %b11 = sparse_tensor.indices %b, %c1 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
182    %b12 = sparse_tensor.indices %b, %c2 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
183    %c10 = sparse_tensor.indices %c, %c0 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
184    %c11 = sparse_tensor.indices %c, %c1 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
185    %c12 = sparse_tensor.indices %c, %c2 : tensor<2x3x4xf64, #Tensor1> to memref<?xindex>
186
187    %d10 = sparse_tensor.indices %d, %c0 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
188    %d11 = sparse_tensor.indices %d, %c1 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
189    %d12 = sparse_tensor.indices %d, %c2 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
190    %e10 = sparse_tensor.indices %e, %c0 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
191    %e11 = sparse_tensor.indices %e, %c1 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
192    %e12 = sparse_tensor.indices %e, %c2 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
193    %f10 = sparse_tensor.indices %f, %c0 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
194    %f11 = sparse_tensor.indices %f, %c1 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
195    %f12 = sparse_tensor.indices %f, %c2 : tensor<2x3x4xf64, #Tensor2> to memref<?xindex>
196
197    %g10 = sparse_tensor.indices %g, %c0 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
198    %g11 = sparse_tensor.indices %g, %c1 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
199    %g12 = sparse_tensor.indices %g, %c2 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
200    %h10 = sparse_tensor.indices %h, %c0 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
201    %h11 = sparse_tensor.indices %h, %c1 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
202    %h12 = sparse_tensor.indices %h, %c2 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
203    %i10 = sparse_tensor.indices %i, %c0 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
204    %i11 = sparse_tensor.indices %i, %c1 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
205    %i12 = sparse_tensor.indices %i, %c2 : tensor<2x3x4xf64, #Tensor3> to memref<?xindex>
206
207    call @check(%v10, %a10) : (memref<?xindex>, memref<?xindex>) -> ()
208    call @check(%v11, %a11) : (memref<?xindex>, memref<?xindex>) -> ()
209    call @check(%v12, %a12) : (memref<?xindex>, memref<?xindex>) -> ()
210    call @check(%v10, %b10) : (memref<?xindex>, memref<?xindex>) -> ()
211    call @check(%v11, %b11) : (memref<?xindex>, memref<?xindex>) -> ()
212    call @check(%v12, %b12) : (memref<?xindex>, memref<?xindex>) -> ()
213    call @check(%v10, %c10) : (memref<?xindex>, memref<?xindex>) -> ()
214    call @check(%v11, %c11) : (memref<?xindex>, memref<?xindex>) -> ()
215    call @check(%v12, %c12) : (memref<?xindex>, memref<?xindex>) -> ()
216
217    call @check(%v20, %d10) : (memref<?xindex>, memref<?xindex>) -> ()
218    call @check(%v21, %d11) : (memref<?xindex>, memref<?xindex>) -> ()
219    call @check(%v22, %d12) : (memref<?xindex>, memref<?xindex>) -> ()
220    call @check(%v20, %e10) : (memref<?xindex>, memref<?xindex>) -> ()
221    call @check(%v21, %e11) : (memref<?xindex>, memref<?xindex>) -> ()
222    call @check(%v22, %e12) : (memref<?xindex>, memref<?xindex>) -> ()
223    call @check(%v20, %f10) : (memref<?xindex>, memref<?xindex>) -> ()
224    call @check(%v21, %f11) : (memref<?xindex>, memref<?xindex>) -> ()
225    call @check(%v22, %f12) : (memref<?xindex>, memref<?xindex>) -> ()
226
227    call @check(%v30, %g10) : (memref<?xindex>, memref<?xindex>) -> ()
228    call @check(%v31, %g11) : (memref<?xindex>, memref<?xindex>) -> ()
229    call @check(%v32, %g12) : (memref<?xindex>, memref<?xindex>) -> ()
230    call @check(%v30, %h10) : (memref<?xindex>, memref<?xindex>) -> ()
231    call @check(%v31, %h11) : (memref<?xindex>, memref<?xindex>) -> ()
232    call @check(%v32, %h12) : (memref<?xindex>, memref<?xindex>) -> ()
233    call @check(%v30, %i10) : (memref<?xindex>, memref<?xindex>) -> ()
234    call @check(%v31, %i11) : (memref<?xindex>, memref<?xindex>) -> ()
235    call @check(%v32, %i12) : (memref<?xindex>, memref<?xindex>) -> ()
236
237    //
238    // Sanity check direct results.
239    //
240    // CHECK:      ( 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24 )
241    // CHECK-NEXT: ( 1, 13, 2, 14, 3, 15, 4, 16, 5, 17, 6, 18, 7, 19, 8, 20, 9, 21, 10, 22, 11, 23, 12, 24 )
242    // CHECK-NEXT: ( 1, 5, 9, 13, 17, 21, 2, 6, 10, 14, 18, 22, 3, 7, 11, 15, 19, 23, 4, 8, 12, 16, 20, 24 )
243    //
244    call @dumpf64(%v1) : (memref<?xf64>) -> ()
245    call @dumpf64(%v2) : (memref<?xf64>) -> ()
246    call @dumpf64(%v3) : (memref<?xf64>) -> ()
247
248    return
249  }
250}
251
252