1// RUN: mlir-opt %s \ 2// RUN: --sparsification --sparse-tensor-conversion \ 3// RUN: --convert-linalg-to-loops --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-std-to-llvm | \ 7// RUN: TENSOR0="%mlir_integration_test_dir/data/mttkrp_b.tns" \ 8// RUN: mlir-cpu-runner \ 9// RUN: -e entry -entry-point-result=void \ 10// RUN: -shared-libs=%mlir_integration_test_dir/libmlir_c_runner_utils%shlibext | \ 11// RUN: FileCheck %s 12 13!Filename = type !llvm.ptr<i8> 14 15#SparseMatrix = #sparse_tensor.encoding<{ 16 dimLevelType = [ "compressed", "compressed", "compressed" ] 17}> 18 19#mttkrp = { 20 indexing_maps = [ 21 affine_map<(i,j,k,l) -> (i,k,l)>, // B 22 affine_map<(i,j,k,l) -> (k,j)>, // C 23 affine_map<(i,j,k,l) -> (l,j)>, // D 24 affine_map<(i,j,k,l) -> (i,j)> // A (out) 25 ], 26 iterator_types = ["parallel", "parallel", "reduction", "reduction"], 27 doc = "A(i,j) += B(i,k,l) * D(l,j) * C(k,j)" 28} 29 30// 31// Integration test that lowers a kernel annotated as sparse to 32// actual sparse code, initializes a matching sparse storage scheme 33// from file, and runs the resulting code with the JIT compiler. 34// 35module { 36 // 37 // Computes Matricized Tensor Times Khatri-Rao Product (MTTKRP) kernel. See 38 // http://tensor-compiler.org/docs/data_analytics/index.html. 39 // 40 func @kernel_mttkrp(%argb: tensor<?x?x?xf64, #SparseMatrix>, 41 %argc: tensor<?x?xf64>, 42 %argd: tensor<?x?xf64>, 43 %arga: tensor<?x?xf64>) -> tensor<?x?xf64> { 44 %0 = linalg.generic #mttkrp 45 ins(%argb, %argc, %argd: 46 tensor<?x?x?xf64, #SparseMatrix>, tensor<?x?xf64>, tensor<?x?xf64>) 47 outs(%arga: tensor<?x?xf64>) { 48 ^bb(%b: f64, %c: f64, %d: f64, %a: f64): 49 %0 = mulf %b, %c : f64 50 %1 = mulf %d, %0 : f64 51 %2 = addf %a, %1 : f64 52 linalg.yield %2 : f64 53 } -> tensor<?x?xf64> 54 return %0 : tensor<?x?xf64> 55 } 56 57 func private @getTensorFilename(index) -> (!Filename) 58 59 // 60 // Main driver that reads matrix from file and calls the sparse kernel. 61 // 62 func @entry() { 63 %i0 = constant 0. : f64 64 %c0 = constant 0 : index 65 %c1 = constant 1 : index 66 %c2 = constant 2 : index 67 %c3 = constant 3 : index 68 %c4 = constant 4 : index 69 %c5 = constant 5 : index 70 %c256 = constant 256 : index 71 72 // Read the sparse B input from a file. 73 %fileName = call @getTensorFilename(%c0) : (index) -> (!Filename) 74 %b = sparse_tensor.new %fileName 75 : !llvm.ptr<i8> to tensor<?x?x?xf64, #SparseMatrix> 76 77 // Initialize dense C and D inputs and dense output A. 78 %cdata = memref.alloc(%c3, %c5) : memref<?x?xf64> 79 scf.for %i = %c0 to %c3 step %c1 { 80 scf.for %j = %c0 to %c5 step %c1 { 81 %k0 = muli %i, %c5 : index 82 %k1 = addi %k0, %j : index 83 %k2 = index_cast %k1 : index to i32 84 %k = sitofp %k2 : i32 to f64 85 memref.store %k, %cdata[%i, %j] : memref<?x?xf64> 86 } 87 } 88 %c = memref.tensor_load %cdata : memref<?x?xf64> 89 90 %ddata = memref.alloc(%c4, %c5) : memref<?x?xf64> 91 scf.for %i = %c0 to %c4 step %c1 { 92 scf.for %j = %c0 to %c5 step %c1 { 93 %k0 = muli %i, %c5 : index 94 %k1 = addi %k0, %j : index 95 %k2 = index_cast %k1 : index to i32 96 %k = sitofp %k2 : i32 to f64 97 memref.store %k, %ddata[%i, %j] : memref<?x?xf64> 98 } 99 } 100 %d = memref.tensor_load %ddata : memref<?x?xf64> 101 102 %adata = memref.alloc(%c2, %c5) : memref<?x?xf64> 103 scf.for %i = %c0 to %c2 step %c1 { 104 scf.for %j = %c0 to %c5 step %c1 { 105 memref.store %i0, %adata[%i, %j] : memref<?x?xf64> 106 } 107 } 108 %a = memref.tensor_load %adata : memref<?x?xf64> 109 110 // Call kernel. 111 %0 = call @kernel_mttkrp(%b, %c, %d, %a) 112 : (tensor<?x?x?xf64, #SparseMatrix>, 113 tensor<?x?xf64>, tensor<?x?xf64>, tensor<?x?xf64>) -> tensor<?x?xf64> 114 115 // Print the result for verification. 116 // 117 // CHECK: ( ( 16075, 21930, 28505, 35800, 43815 ), 118 // CHECK: ( 10000, 14225, 19180, 24865, 31280 ) ) 119 // 120 %m = memref.buffer_cast %0 : memref<?x?xf64> 121 %v = vector.transfer_read %m[%c0, %c0], %i0 122 : memref<?x?xf64>, vector<2x5xf64> 123 vector.print %v : vector<2x5xf64> 124 125 // Release the resources. 126 memref.dealloc %adata : memref<?x?xf64> 127 memref.dealloc %cdata : memref<?x?xf64> 128 memref.dealloc %ddata : memref<?x?xf64> 129 130 return 131 } 132} 133