brintos

brintos / llvm-project-archived public Read only

0
0
Text · 9.2 KiB · 84a8e10 Raw
215 lines · plain
1// RUN: mlir-opt %s -lower-affine -convert-vector-to-scf -convert-scf-to-cf -convert-vector-to-llvm="enable-arm-sve" -finalize-memref-to-llvm -convert-func-to-llvm -convert-arith-to-llvm -convert-cf-to-llvm -canonicalize | \2// RUN: %mcr_aarch64_cmd -e=entry -entry-point-result=void --march=aarch64 --mattr="+sve" -shared-libs=%native_mlir_c_runner_utils | \3// RUN: FileCheck %s4 5// Note: To run this test, your CPU must support SVE6 7// VLA memcopy8func.func @kernel_copy(%src : memref<?xi64>, %dst : memref<?xi64>, %size : index) {9  %c0 = arith.constant 0 : index10  %c2 = arith.constant 2 : index11  %vs = vector.vscale12  %step = arith.muli %c2, %vs : index13  scf.for %i0 = %c0 to %size step %step {14    %0 = vector.load %src[%i0] : memref<?xi64>, vector<[2]xi64>15    vector.store %0, %dst[%i0] : memref<?xi64>, vector<[2]xi64>16  }17 18  return19}20 21// VLA multiply and add22func.func @kernel_muladd(%a : memref<?xi64>,23                    %b : memref<?xi64>,24                    %c : memref<?xi64>,25                    %size : index) {26  %c0 = arith.constant 0 : index27  %c2 = arith.constant 2 : index28  %vs = vector.vscale29  %step = arith.muli %c2, %vs : index30  scf.for %i0 = %c0 to %size step %step {31    %0 = vector.load %a[%i0] : memref<?xi64>, vector<[2]xi64>32    %1 = vector.load %b[%i0] : memref<?xi64>, vector<[2]xi64>33    %2 = vector.load %c[%i0] : memref<?xi64>, vector<[2]xi64>34    %3 = arith.muli %0, %1 : vector<[2]xi64>35    %4 = arith.addi %3, %2 : vector<[2]xi64>36    vector.store %4, %c[%i0] : memref<?xi64>, vector<[2]xi64>37  }38  return39}40 41// SVE-based absolute difference42func.func @kernel_absdiff(%a : memref<?xi64>,43                     %b : memref<?xi64>,44                     %c : memref<?xi64>,45                     %size : index) {46  %c0 = arith.constant 0 : index47  %c2 = arith.constant 2 : index48  %vs = vector.vscale49  %step = arith.muli %c2, %vs : index50  scf.for %i0 = %c0 to %size step %step {51    %0 = vector.load %a[%i0] : memref<?xi64>, vector<[2]xi64>52    %1 = vector.load %b[%i0] : memref<?xi64>, vector<[2]xi64>53    %agb = arith.cmpi sge, %0, %1 : vector<[2]xi64>54    %bga = arith.cmpi slt, %0, %1 : vector<[2]xi64>55    %10 = arm_sve.masked.subi %agb, %0, %1 : vector<[2]xi1>,56                                             vector<[2]xi64>57    %01 = arm_sve.masked.subi %bga, %1, %0 : vector<[2]xi1>,58                                             vector<[2]xi64>59    vector.maskedstore %c[%i0], %agb, %10 : memref<?xi64>,60                                            vector<[2]xi1>,61                                            vector<[2]xi64>62    vector.maskedstore %c[%i0], %bga, %01 : memref<?xi64>,63                                            vector<[2]xi1>,64                                            vector<[2]xi64>65  }66  return67}68 69// VLA unknown bounds vector addition70func.func @kernel_addition(%a : memref<?xf32>,71                      %b : memref<?xf32>,72                      %c : memref<?xf32>,73                      %N : index) {74  %c0 = arith.constant 0 : index75  %c4 = arith.constant 2 : index76  %v0f = arith.constant dense<0.0> : vector<[4]xf32>77  %vs = vector.vscale78  %s = arith.muli %c4, %vs : index79  scf.for %i0 = %c0 to %N step %s {80    %sub = affine.min affine_map<(d0, d1)[s0] -> (s0, d0 - d1)>(%N, %i0)[%s]81    %mask = vector.create_mask %sub : vector<[4]xi1>82    %la = vector.maskedload %a[%i0], %mask, %v0f : memref<?xf32>, vector<[4]xi1>, vector<[4]xf32> into vector<[4]xf32>83    %lb = vector.maskedload %b[%i0], %mask, %v0f : memref<?xf32>, vector<[4]xi1>, vector<[4]xf32> into vector<[4]xf32>84    %lc = arith.addf %la, %lb : vector<[4]xf32>85    vector.maskedstore %c[%i0], %mask, %lc : memref<?xf32>, vector<[4]xi1>, vector<[4]xf32>86  }87  return88}89 90func.func @entry() {91  %i0 = arith.constant 0: i6492  %i1 = arith.constant 1: i6493  %f0 = arith.constant 0.0: f3294  %c0 = arith.constant 0: index95  %c1 = arith.constant 1: index96  %c2 = arith.constant 2: index97  %c4 = arith.constant 4: index98  %c8 = arith.constant 8: index99  %c32 = arith.constant 32: index100  %c33 = arith.constant 33: index101  %c-1 = arith.constant -1.0 : f32102 103  // Set up memory.104  %a = memref.alloc()      : memref<32xi64>105  %a_copy = memref.alloc() : memref<32xi64>106  %b = memref.alloc()      : memref<32xi64>107  %c = memref.alloc()      : memref<32xi64>108  %d = memref.alloc()      : memref<32xi64>109  %e = memref.alloc()      : memref<33xf32>110  %f = memref.alloc()      : memref<33xf32>111  %g = memref.alloc()      : memref<36xf32>112 113  %a_data = arith.constant dense<[1 , 2,  3 , 4 , 5,  6,  7,  8,114                                9, 10, 11, 12, 13, 14, 15, 16,115                                17, 18, 19, 20, 21, 22, 23, 24,116                                25, 26, 27, 28, 29, 30, 31, 32]> : vector<32xi64>117  vector.transfer_write %a_data, %a[%c0] : vector<32xi64>, memref<32xi64>118  %b_data = arith.constant dense<[33, 34, 35, 36, 37, 38, 39, 40,119                                41, 42, 43, 44, 45, 46, 47, 48,120                                49, 50, 51, 52, 53, 54, 55, 56,121                                57, 58, 59, 60, 61, 62, 63, 64]> : vector<32xi64>122  vector.transfer_write %b_data, %b[%c0] : vector<32xi64>, memref<32xi64>123  %d_data = arith.constant dense<[-9, 76, -7, 78, -5, 80, -3, 82,124                                -1, 84, 1, 86, 3, 88, 5, 90,125                                7, 92, 9, 94, 11, 96, 13, 98,126                                15, 100, 17, 102, 19, 104, 21, 106]> : vector<32xi64>127  vector.transfer_write %d_data, %d[%c0] : vector<32xi64>, memref<32xi64>128  %zero_data = vector.broadcast %i0 : i64 to vector<32xi64>129  vector.transfer_write %zero_data, %a_copy[%c0] : vector<32xi64>, memref<32xi64>130  %one_data = vector.broadcast %i1 : i64 to vector<32xi64>131  vector.transfer_write %one_data, %c[%c0] : vector<32xi64>, memref<32xi64>132 133  %e_data = arith.constant dense<[1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5,134                                9.5, 10.5, 11.5, 12.5, 13.5, 14.5, 15.5, 16.5,135                                17.5, 18.5, 19.5, 20.5, 21.5, 22.5, 23.5, 24.5,136                                25.5, 26.5, 27.5, 28.5, 29.5, 30.5, 31.5, 32.5,137                                33.5]> : vector<33xf32>138  vector.transfer_write %e_data, %e[%c0] : vector<33xf32>, memref<33xf32>139  %f_data = arith.constant dense<[40.5, 39.5, 38.5, 37.5, 36.5, 35.5, 34.5, 33.5,140                                32.5, 31.5, 30.5, 29.5, 28.5, 27.5, 26.5, 25.5,141                                24.5, 23.5, 22.5, 21.5, 20.5, 19.5, 18.5, 17.5,142                                16.5, 15.5, 14.5, 13.5, 12.5, 11.5, 10.5, 9.5,143                                8.5]> : vector<33xf32>144  vector.transfer_write %f_data, %f[%c0] : vector<33xf32>, memref<33xf32>145  %minus1_data = vector.broadcast %c-1 : f32 to vector<36xf32>146  vector.transfer_write %minus1_data, %g[%c0] : vector<36xf32>, memref<36xf32>147 148  // Call kernel.149  %0 = memref.cast %a : memref<32xi64> to memref<?xi64>150  %1 = memref.cast %a_copy : memref<32xi64> to memref<?xi64>151  call @kernel_copy(%0, %1, %c32) : (memref<?xi64>, memref<?xi64>, index) -> ()152 153  // Print and verify.154  //155  // CHECK:      ( 1, 2, 3, 4 )156  // CHECK-NEXT: ( 5, 6, 7, 8 )157  scf.for %i = %c0 to %c32 step %c4 {158    %cv = vector.transfer_read %a_copy[%i], %i0: memref<32xi64>, vector<4xi64>159    vector.print %cv : vector<4xi64>160  }161 162  %2 = memref.cast %a : memref<32xi64> to memref<?xi64>163  %3 = memref.cast %b : memref<32xi64> to memref<?xi64>164  %4 = memref.cast %c : memref<32xi64> to memref<?xi64>165  call @kernel_muladd(%2, %3, %4, %c32) : (memref<?xi64>, memref<?xi64>, memref<?xi64>, index) -> ()166 167  // CHECK:      ( 34, 69, 106, 145 )168  // CHECK-NEXT: ( 186, 229, 274, 321 )169  scf.for %i = %c0 to %c32 step %c4 {170    %macv = vector.transfer_read %c[%i], %i0: memref<32xi64>, vector<4xi64>171    vector.print %macv : vector<4xi64>172  }173 174  %5 = memref.cast %b : memref<32xi64> to memref<?xi64>175  %6 = memref.cast %d : memref<32xi64> to memref<?xi64>176  %7 = memref.cast %c : memref<32xi64> to memref<?xi64>177  call @kernel_absdiff(%5, %6, %7, %c32) : (memref<?xi64>, memref<?xi64>, memref<?xi64>, index) -> ()178 179  // CHECK:      ( 42, 42, 42, 42 )180  // CHECK-NEXT: ( 42, 42, 42, 42 )181  scf.for %i = %c0 to %c32 step %c4 {182    %abdv = vector.transfer_read %c[%i], %i0: memref<32xi64>, vector<4xi64>183    vector.print %abdv : vector<4xi64>184  }185 186  %ee = memref.cast %e : memref<33xf32> to memref<?xf32>187  %ff = memref.cast %f : memref<33xf32> to memref<?xf32>188  %gg = memref.cast %g : memref<36xf32> to memref<?xf32>189  call @kernel_addition(%ee, %ff, %gg, %c33) : (memref<?xf32>, memref<?xf32>, memref<?xf32>, index) -> ()190 191  // CHECK:      ( 42, 42, 42, 42, 42, 42, 42, 42 )192  // CHECK-NEXT: ( 42, 42, 42, 42, 42, 42, 42, 42 )193  // CHECK-NEXT: ( 42, 42, 42, 42, 42, 42, 42, 42 )194  // CHECK-NEXT: ( 42, 42, 42, 42, 42, 42, 42, 42 )195  // CHECK-NEXT: ( 42, -1, -1, -1 )196  scf.for %i = %c0 to %c32 step %c8 {197    %addv = vector.transfer_read %g[%i], %f0: memref<36xf32>, vector<8xf32>198    vector.print %addv : vector<8xf32>199  }200  %remv = vector.transfer_read %g[%c32], %f0: memref<36xf32>, vector<4xf32>201  vector.print %remv : vector<4xf32>202 203  // Release resources.204  memref.dealloc %a      : memref<32xi64>205  memref.dealloc %a_copy : memref<32xi64>206  memref.dealloc %b      : memref<32xi64>207  memref.dealloc %c      : memref<32xi64>208  memref.dealloc %d      : memref<32xi64>209  memref.dealloc %e      : memref<33xf32>210  memref.dealloc %f      : memref<33xf32>211  memref.dealloc %g      : memref<36xf32>212 213  return214}215