219 lines · plain
1// NOTE: this test requires gpu-sm80 and cusparselt2//3// DEFINE: %{compile} = mlir-opt --convert-vector-to-scf --convert-scf-to-cf -convert-cf-to-llvm --convert-vector-to-llvm \4// DEFINE: --convert-arith-to-llvm --gpu-to-llvm --reconcile-unrealized-casts \5// DEFINE: %s6// DEFINE: %{run} = mlir-runner \7// DEFINE: --shared-libs=%mlir_cuda_runtime \8// DEFINE: --shared-libs=%mlir_c_runner_utils \9// DEFINE: --e main --entry-point-result=void \10// DEFINE: | FileCheck %s11//12// RUN: %{compile} | %{run}13 14module {15 llvm.func @mgpuCreateSparseLtEnv()16 llvm.func @mgpuDestroySparseLtEnv()17 18 // cuSparselt version for matmul coded by hand.19 func.func @matmul24(%a : memref<16x32xf16>,20 %b : memref<32x16xf16>,21 %c : memref<16x16xf16>) {22 %c0 = arith.constant 0.0 : f1623 %c1 = arith.constant 1 : index24 %c2 = arith.constant 2 : index25 %c8 = arith.constant 8 : index26 %c16 = arith.constant 16 : index27 %c32 = arith.constant 32 : index28 %c1048576 = arith.constant 1048576 : index29 %token0 = gpu.wait async30 %d_a, %token1 = gpu.alloc async [%token0] () : memref<16x32xf16>31 %d_b, %token2 = gpu.alloc async [%token1] () : memref<32x16xf16>32 %d_c, %token3 = gpu.alloc async [%token2] () : memref<16x16xf16>33 %token4 = gpu.memcpy async [%token3] %d_a, %a : memref<16x32xf16>, memref<16x32xf16>34 %token5 = gpu.memcpy async [%token4] %d_b, %b : memref<32x16xf16>, memref<32x16xf16>35 %token6 = gpu.memcpy async [%token5] %d_c, %c : memref<16x16xf16>, memref<16x16xf16>36 %spmat, %token8 = gpu.create_2to4_spmat async [%token6]{PRUNE_AND_CHECK} %c16, %c32, %d_a: memref<16x32xf16>37 %dnmat, %token9 = gpu.create_dn_tensor async [%token8] %d_b, %c32, %c16: index, index into memref<32x16xf16>38 %dnmat2, %token10 = gpu.create_dn_tensor async [%token9] %d_c, %c16, %c16: index, index into memref<16x16xf16>39 %bufferSz0, %bufferSz1, %bufferSz2, %token11 = gpu.spmm_buffer_size async [%token10] %spmat{NON_TRANSPOSE}, %dnmat{NON_TRANSPOSE}, %dnmat2 : index, index,index into f1640 %mem1, %token12 = gpu.alloc async [%token11] (%bufferSz0) : memref<?xf16>41 %mem2, %token13 = gpu.alloc async [%token12] (%bufferSz1) : memref<?xf16>42 %mem3, %token14 = gpu.alloc async [%token13] (%bufferSz2) : memref<?xf16>43 %token15 = gpu.spmm async [%token14] %spmat{NON_TRANSPOSE}, %dnmat{NON_TRANSPOSE}, %dnmat2, %mem1, %mem2, %mem3 : memref<?xf16>, memref<?xf16>,memref<?xf16> into f1644 %token16 = gpu.destroy_sp_mat async [%token15] %spmat45 %token17 = gpu.destroy_dn_tensor async [%token16] %dnmat46 %token18 = gpu.destroy_dn_tensor async [%token17] %dnmat247 %token19 = gpu.memcpy async [%token18] %c, %d_c : memref<16x16xf16>, memref<16x16xf16>48 %token20 = gpu.dealloc async [%token19] %d_c : memref<16x16xf16>49 %token21 = gpu.dealloc async [%token20] %d_b : memref<32x16xf16>50 %token22 = gpu.dealloc async [%token21] %d_a : memref<16x32xf16>51 %token23 = gpu.dealloc async [%token22] %mem3 : memref<?xf16>52 %token24 = gpu.dealloc async [%token23] %mem2 : memref<?xf16>53 %token25 = gpu.dealloc async [%token24] %mem1 : memref<?xf16>54 gpu.wait [%token25]55 return56 }57 58 //59 // This test performs a matrix multiplication60 // C = A x B61 // using NVidia 2:4 structured sparsity for A.62 //63 func.func @main() {64 llvm.call @mgpuCreateSparseLtEnv() : () -> ()65 %f0 = arith.constant 0.0 : f1666 %c0 = arith.constant 0 : index67 %c1 = arith.constant 1 : index68 %c2 = arith.constant 2 : index69 %c8 = arith.constant 8 : index70 %c16 = arith.constant 16 : index71 %c32 = arith.constant 32 : index72 %c64 = arith.constant 64 : index73 74 // Matrices A, B, C (16x32, 32x16, 16x16).75 %a = memref.alloc() : memref<16x32xf16> // 16x32 with 2:4, row-major76 %b = memref.alloc() : memref<32x16xf16> // regular dense column-major77 %c = memref.alloc() : memref<16x16xf16> // accumulator row-major78 79 //80 // Setup matrix A.81 //82 scf.for %ai = %c0 to %c16 step %c1 {83 scf.for %aj = %c0 to %c16 step %c1 {84 %cf0 = arith.constant 0.0: f1685 %a0 = arith.addi %ai, %aj : index86 %a1 = arith.addi %a0, %c1 : index87 %a2 = arith.index_cast %a1 : index to i3288 %a3 = arith.sitofp %a2 : i32 to f1689 %ajj = arith.muli %aj, %c2 : index90 %ajj2 = arith.addi %ajj, %c1 : index91 memref.store %a3, %a[%ai, %ajj] : memref<16x32xf16>92 memref.store %cf0, %a[%ai, %ajj2] : memref<16x32xf16>93 }94 }95 96 //97 // Setup matrix B.98 //99 scf.for %bi = %c0 to %c8 step %c1 {100 scf.for %bj = %c0 to %c32 step %c1 {101 %b0 = arith.subi %bi, %bj : index102 %b1 = arith.index_cast %b0 : index to i32103 %b2 = arith.sitofp %b1 : i32 to f16104 %bii = arith.addi %bi, %c8 : index105 memref.store %b2, %b[%bj, %bi] : memref<32x16xf16>106 memref.store %b2, %b[%bj, %bii] : memref<32x16xf16>107 }108 }109 110 //111 // Reset matrix C.112 //113 scf.for %ci = %c0 to %c16 step %c1 {114 scf.for %cj = %c0 to %c16 step %c1 {115 memref.store %f0, %c[%ci, %cj] : memref<16x16xf16>116 }117 }118 119 //120 // Sanity check on 16x32 full 2:4 input matrix A.121 //122 //123 // CHECK: ( 1, 0, 2, 0, 3, 0, 4, 0, 5, 0, 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0 )124 // CHECK-NEXT: ( 2, 0, 3, 0, 4, 0, 5, 0, 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0 )125 // CHECK-NEXT: ( 3, 0, 4, 0, 5, 0, 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0 )126 // CHECK-NEXT: ( 4, 0, 5, 0, 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0 )127 // CHECK-NEXT: ( 5, 0, 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0 )128 // CHECK-NEXT: ( 6, 0, 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0 )129 // CHECK-NEXT: ( 7, 0, 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0 )130 // CHECK-NEXT: ( 8, 0, 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0 )131 // CHECK-NEXT: ( 9, 0, 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0 )132 // CHECK-NEXT: ( 10, 0, 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0 )133 // CHECK-NEXT: ( 11, 0, 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0 )134 // CHECK-NEXT: ( 12, 0, 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0, 27, 0 )135 // CHECK-NEXT: ( 13, 0, 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0, 27, 0, 28, 0 )136 // CHECK-NEXT: ( 14, 0, 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0, 27, 0, 28, 0, 29, 0 )137 // CHECK-NEXT: ( 15, 0, 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0, 27, 0, 28, 0, 29, 0, 30, 0 )138 // CHECK-NEXT: ( 16, 0, 17, 0, 18, 0, 19, 0, 20, 0, 21, 0, 22, 0, 23, 0, 24, 0, 25, 0, 26, 0, 27, 0, 28, 0, 29, 0, 30, 0, 31, 0 )139 //140 scf.for %pai = %c0 to %c16 step %c1 {141 %pa0 = vector.transfer_read %a[%pai, %c0], %f0 : memref<16x32xf16>, vector<32xf16>142 vector.print %pa0 : vector<32xf16>143 }144 145 //146 // Sanity check on input matrix 32x16 B.147 //148 // CHECK-NEXT: ( 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7 )149 // CHECK-NEXT: ( -1, 0, 1, 2, 3, 4, 5, 6, -1, 0, 1, 2, 3, 4, 5, 6 )150 // CHECK-NEXT: ( -2, -1, 0, 1, 2, 3, 4, 5, -2, -1, 0, 1, 2, 3, 4, 5 )151 // CHECK-NEXT: ( -3, -2, -1, 0, 1, 2, 3, 4, -3, -2, -1, 0, 1, 2, 3, 4 )152 // CHECK-NEXT: ( -4, -3, -2, -1, 0, 1, 2, 3, -4, -3, -2, -1, 0, 1, 2, 3 )153 // CHECK-NEXT: ( -5, -4, -3, -2, -1, 0, 1, 2, -5, -4, -3, -2, -1, 0, 1, 2 )154 // CHECK-NEXT: ( -6, -5, -4, -3, -2, -1, 0, 1, -6, -5, -4, -3, -2, -1, 0, 1 )155 // CHECK-NEXT: ( -7, -6, -5, -4, -3, -2, -1, 0, -7, -6, -5, -4, -3, -2, -1, 0 )156 // CHECK-NEXT: ( -8, -7, -6, -5, -4, -3, -2, -1, -8, -7, -6, -5, -4, -3, -2, -1 )157 // CHECK-NEXT: ( -9, -8, -7, -6, -5, -4, -3, -2, -9, -8, -7, -6, -5, -4, -3, -2 )158 // CHECK-NEXT: ( -10, -9, -8, -7, -6, -5, -4, -3, -10, -9, -8, -7, -6, -5, -4, -3 )159 // CHECK-NEXT: ( -11, -10, -9, -8, -7, -6, -5, -4, -11, -10, -9, -8, -7, -6, -5, -4 )160 // CHECK-NEXT: ( -12, -11, -10, -9, -8, -7, -6, -5, -12, -11, -10, -9, -8, -7, -6, -5 )161 // CHECK-NEXT: ( -13, -12, -11, -10, -9, -8, -7, -6, -13, -12, -11, -10, -9, -8, -7, -6 )162 // CHECK-NEXT: ( -14, -13, -12, -11, -10, -9, -8, -7, -14, -13, -12, -11, -10, -9, -8, -7 )163 // CHECK-NEXT: ( -15, -14, -13, -12, -11, -10, -9, -8, -15, -14, -13, -12, -11, -10, -9, -8 )164 // CHECK-NEXT: ( -16, -15, -14, -13, -12, -11, -10, -9, -16, -15, -14, -13, -12, -11, -10, -9 )165 // CHECK-NEXT: ( -17, -16, -15, -14, -13, -12, -11, -10, -17, -16, -15, -14, -13, -12, -11, -10 )166 // CHECK-NEXT: ( -18, -17, -16, -15, -14, -13, -12, -11, -18, -17, -16, -15, -14, -13, -12, -11 )167 // CHECK-NEXT: ( -19, -18, -17, -16, -15, -14, -13, -12, -19, -18, -17, -16, -15, -14, -13, -12 )168 // CHECK-NEXT: ( -20, -19, -18, -17, -16, -15, -14, -13, -20, -19, -18, -17, -16, -15, -14, -13 )169 // CHECK-NEXT: ( -21, -20, -19, -18, -17, -16, -15, -14, -21, -20, -19, -18, -17, -16, -15, -14 )170 // CHECK-NEXT: ( -22, -21, -20, -19, -18, -17, -16, -15, -22, -21, -20, -19, -18, -17, -16, -15 )171 // CHECK-NEXT: ( -23, -22, -21, -20, -19, -18, -17, -16, -23, -22, -21, -20, -19, -18, -17, -16 )172 // CHECK-NEXT: ( -24, -23, -22, -21, -20, -19, -18, -17, -24, -23, -22, -21, -20, -19, -18, -17 )173 // CHECK-NEXT: ( -25, -24, -23, -22, -21, -20, -19, -18, -25, -24, -23, -22, -21, -20, -19, -18 )174 // CHECK-NEXT: ( -26, -25, -24, -23, -22, -21, -20, -19, -26, -25, -24, -23, -22, -21, -20, -19 )175 // CHECK-NEXT: ( -27, -26, -25, -24, -23, -22, -21, -20, -27, -26, -25, -24, -23, -22, -21, -20 )176 // CHECK-NEXT: ( -28, -27, -26, -25, -24, -23, -22, -21, -28, -27, -26, -25, -24, -23, -22, -21 )177 // CHECK-NEXT: ( -29, -28, -27, -26, -25, -24, -23, -22, -29, -28, -27, -26, -25, -24, -23, -22 )178 // CHECK-NEXT: ( -30, -29, -28, -27, -26, -25, -24, -23, -30, -29, -28, -27, -26, -25, -24, -23 )179 // CHECK-NEXT: ( -31, -30, -29, -28, -27, -26, -25, -24, -31, -30, -29, -28, -27, -26, -25, -24 )180 //181 //182 scf.for %pbi = %c0 to %c32 step %c1 {183 %pb0 = vector.transfer_read %b[%pbi, %c0], %f0 : memref<32x16xf16>, vector<16xf16>184 vector.print %pb0 : vector<16xf16>185 }186 187 // Call the kernel.188 call @matmul24(%a, %b, %c): (memref<16x32xf16>, memref<32x16xf16>, memref<16x16xf16>) -> ()189 190 //191 // Verify computed matrix C.192 //193 // CHECK-NEXT: ( -2720, -2584, -2448, -2312, -2176, -2040, -1904, -1768, -2720, -2584, -2448, -2312, -2176, -2040, -1904, -1768 )194 // CHECK-NEXT: ( -2960, -2808, -2656, -2504, -2352, -2200, -2048, -1896, -2960, -2808, -2656, -2504, -2352, -2200, -2048, -1896 )195 // CHECK-NEXT: ( -3200, -3032, -2864, -2696, -2528, -2360, -2192, -2024, -3200, -3032, -2864, -2696, -2528, -2360, -2192, -2024 )196 // CHECK-NEXT: ( -3440, -3256, -3072, -2888, -2704, -2520, -2336, -2152, -3440, -3256, -3072, -2888, -2704, -2520, -2336, -2152 )197 // CHECK-NEXT: ( -3680, -3480, -3280, -3080, -2880, -2680, -2480, -2280, -3680, -3480, -3280, -3080, -2880, -2680, -2480, -2280 )198 // CHECK-NEXT: ( -3920, -3704, -3488, -3272, -3056, -2840, -2624, -2408, -3920, -3704, -3488, -3272, -3056, -2840, -2624, -2408 )199 // CHECK-NEXT: ( -4160, -3928, -3696, -3464, -3232, -3000, -2768, -2536, -4160, -3928, -3696, -3464, -3232, -3000, -2768, -2536 )200 // CHECK-NEXT: ( -4400, -4152, -3904, -3656, -3408, -3160, -2912, -2664, -4400, -4152, -3904, -3656, -3408, -3160, -2912, -2664 )201 // CHECK-NEXT: ( -4640, -4376, -4112, -3848, -3584, -3320, -3056, -2792, -4640, -4376, -4112, -3848, -3584, -3320, -3056, -2792 )202 // CHECK-NEXT: ( -4880, -4600, -4320, -4040, -3760, -3480, -3200, -2920, -4880, -4600, -4320, -4040, -3760, -3480, -3200, -2920 )203 // CHECK-NEXT: ( -5120, -4824, -4528, -4232, -3936, -3640, -3344, -3048, -5120, -4824, -4528, -4232, -3936, -3640, -3344, -3048 )204 // CHECK-NEXT: ( -5360, -5048, -4736, -4424, -4112, -3800, -3488, -3176, -5360, -5048, -4736, -4424, -4112, -3800, -3488, -3176 )205 // CHECK-NEXT: ( -5600, -5272, -4944, -4616, -4288, -3960, -3632, -3304, -5600, -5272, -4944, -4616, -4288, -3960, -3632, -3304 )206 // CHECK-NEXT: ( -5840, -5496, -5152, -4808, -4464, -4120, -3776, -3432, -5840, -5496, -5152, -4808, -4464, -4120, -3776, -3432 )207 // CHECK-NEXT: ( -6080, -5720, -5360, -5000, -4640, -4280, -3920, -3560, -6080, -5720, -5360, -5000, -4640, -4280, -3920, -3560 )208 // CHECK-NEXT: ( -6320, -5944, -5568, -5192, -4816, -4440, -4064, -3688, -6320, -5944, -5568, -5192, -4816, -4440, -4064, -3688 )209 //210 scf.for %pci = %c0 to %c16 step %c1 {211 %pc0 = vector.transfer_read %c[%pci, %c0], %f0 : memref<16x16xf16>, vector<16xf16>212 vector.print %pc0 : vector<16xf16>213 }214 215 llvm.call @mgpuDestroySparseLtEnv() : () -> ()216 return217 }218}219