|
19 | 19 |
|
20 | 20 | package org.apache.sysds.test.component.matrixmult; |
21 | 21 |
|
22 | | -import org.apache.sysds.runtime.data.DenseBlock; |
| 22 | +import java.util.Random; |
| 23 | + |
23 | 24 | import org.apache.sysds.runtime.matrix.data.LibMatrixMult; |
24 | 25 | import org.apache.sysds.runtime.matrix.data.LibMatrixReorg; |
25 | 26 | import org.apache.sysds.runtime.matrix.data.MatrixBlock; |
26 | 27 | import org.apache.sysds.test.TestUtils; |
27 | 28 | import org.junit.Test; |
28 | 29 |
|
29 | | -import java.util.Random; |
30 | | - |
31 | 30 | public class MatrixMultTransposedTest { |
32 | 31 |
|
33 | | - // run multiple random scenarios |
34 | 32 | @Test |
35 | | - public void testCaseNoTransATransB() { |
36 | | - for(int i=0; i<10; i++) { |
37 | | - runTest(false, true); |
38 | | - } |
| 33 | + public void testDenseDenseTransA() throws Exception { |
| 34 | + for(int i = 0; i < 10; i++) |
| 35 | + runRandomTest(false, false, true, false); |
39 | 36 | } |
40 | 37 |
|
41 | 38 | @Test |
42 | | - public void testCaseTransANoTransB() { |
43 | | - for(int i=0; i<10; i++) { |
44 | | - runTest(true, false); |
45 | | - } |
| 39 | + public void testDenseDenseTransB() throws Exception { |
| 40 | + for(int i = 0; i < 10; i++) |
| 41 | + runRandomTest(false, false, false, true); |
46 | 42 | } |
47 | 43 |
|
48 | 44 | @Test |
49 | | - public void testCaseTransATransB() { |
50 | | - for(int i=0; i<10; i++) { |
51 | | - runTest(true, true); |
52 | | - } |
| 45 | + public void testDenseDenseTransATransB() throws Exception { |
| 46 | + for(int i = 0; i < 10; i++) |
| 47 | + runRandomTest(false, false, true, true); |
53 | 48 | } |
54 | 49 |
|
55 | | - private void runTest(boolean tA, boolean tB) { |
56 | | - Random rand = new Random(); |
| 50 | + @Test |
| 51 | + public void testSparseDenseTransA() throws Exception { |
| 52 | + for(int i = 0; i < 10; i++) |
| 53 | + runRandomTest(true, false, true, false); |
| 54 | + } |
| 55 | + |
| 56 | + @Test |
| 57 | + public void testSparseDenseTransB() throws Exception { |
| 58 | + for(int i = 0; i < 10; i++) |
| 59 | + runRandomTest(true, false, false, true); |
| 60 | + } |
| 61 | + |
| 62 | + @Test |
| 63 | + public void testSparseDenseTransATransB() throws Exception { |
| 64 | + for(int i = 0; i < 10; i++) |
| 65 | + runRandomTest(true, false, true, true); |
| 66 | + } |
| 67 | + |
| 68 | + @Test |
| 69 | + public void testDenseSparseTransA() throws Exception { |
| 70 | + for(int i = 0; i < 10; i++) |
| 71 | + runRandomTest(false, true, true, false); |
| 72 | + } |
| 73 | + |
| 74 | + @Test |
| 75 | + public void testDenseSparseTransB() throws Exception { |
| 76 | + for(int i = 0; i < 10; i++) |
| 77 | + runRandomTest(false, true, false, true); |
| 78 | + } |
57 | 79 |
|
58 | | - // generate random dimensions between 1 and 300 |
| 80 | + @Test |
| 81 | + public void testDenseSparseTransATransB() throws Exception { |
| 82 | + for(int i = 0; i < 10; i++) |
| 83 | + runRandomTest(false, true, true, true); |
| 84 | + } |
| 85 | + |
| 86 | + private void runRandomTest(boolean sparseA, boolean sparseB, boolean tA, boolean tB) throws Exception { |
| 87 | + Random rand = new Random(); |
59 | 88 | int m = rand.nextInt(300) + 1; |
60 | 89 | int n = rand.nextInt(300) + 1; |
61 | 90 | int k = rand.nextInt(300) + 1; |
62 | 91 |
|
| 92 | + double spA = sparseA ? 0.05 : 1.0; |
| 93 | + double spB = sparseB ? 0.05 : 1.0; |
| 94 | + |
| 95 | + runTest(spA, spB, tA, tB, m, n, k); |
| 96 | + } |
63 | 97 |
|
| 98 | + private void runTest(double spA, double spB, boolean tA, boolean tB, int m, int n, int k) throws Exception { |
64 | 99 | int rowsA = tA ? k : m; |
65 | 100 | int colsA = tA ? m : k; |
66 | 101 | int rowsB = tB ? n : k; |
67 | 102 | int colsB = tB ? k : n; |
68 | 103 |
|
69 | | - MatrixBlock ma = MatrixBlock.randOperations(rowsA, colsA, 1.0, -1, 1, "uniform", 7); |
70 | | - MatrixBlock mb = MatrixBlock.randOperations(rowsB, colsB, 1.0, -1, 1, "uniform", 3); |
71 | | - |
| 104 | + MatrixBlock ma = generateInput(rowsA, colsA, spA, 7); |
| 105 | + MatrixBlock mb = generateInput(rowsB, colsB, spB, 3); |
72 | 106 | MatrixBlock mc = new MatrixBlock(m, n, false); |
73 | 107 | mc.allocateDenseBlock(); |
74 | 108 |
|
75 | | - DenseBlock a = ma.getDenseBlock(); |
76 | | - DenseBlock b = mb.getDenseBlock(); |
77 | | - DenseBlock c = mc.getDenseBlock(); |
| 109 | + runNewKernel(ma, mb, mc, tA, tB); |
78 | 110 |
|
79 | | - LibMatrixMult.matrixMultDenseDenseMM(a, b, c, tA, tB, n, k, 0, m, 0, n); |
| 111 | + MatrixBlock A_in = tA ? LibMatrixReorg.transpose(ma) : ma; |
| 112 | + MatrixBlock B_in = tB ? LibMatrixReorg.transpose(mb) : mb; |
| 113 | + MatrixBlock expected = LibMatrixMult.matrixMult(A_in, B_in); |
80 | 114 |
|
81 | | - mc.recomputeNonZeros(); |
| 115 | + TestUtils.compareMatrices(expected, mc, 1e-8); |
| 116 | + } |
82 | 117 |
|
83 | | - // calc true result with existing methods |
84 | | - MatrixBlock ma_in = tA ? LibMatrixReorg.transpose(ma) : ma; |
85 | | - MatrixBlock mb_in = tB ? LibMatrixReorg.transpose(mb) : mb; |
86 | | - MatrixBlock expected = LibMatrixMult.matrixMult(ma_in, mb_in); |
| 118 | + private MatrixBlock generateInput(int rows, int cols, double sparsity, long seed) { |
| 119 | + MatrixBlock mb = MatrixBlock.randOperations(rows, cols, sparsity, -1, 1, "uniform", seed); |
| 120 | + mb.examSparsity(); |
| 121 | + if (sparsity < 1.0) { |
| 122 | + if (!mb.isInSparseFormat()) |
| 123 | + mb.denseToSparse(true); |
| 124 | + if (mb.getSparseBlock() == null) |
| 125 | + mb.allocateSparseRowsBlock(); |
| 126 | + } |
| 127 | + return mb; |
| 128 | + } |
87 | 129 |
|
88 | | - // compare results |
89 | | - TestUtils.compareMatrices(expected, mc, 1e-8); |
| 130 | + private void runNewKernel(MatrixBlock ma, MatrixBlock mb, MatrixBlock mc, boolean tA, boolean tB) throws Exception { |
| 131 | + mc.reset(); |
| 132 | + mc.allocateDenseBlock(); |
| 133 | + |
| 134 | + boolean sparseA = ma.isInSparseFormat(); |
| 135 | + boolean sparseB = mb.isInSparseFormat(); |
| 136 | + |
| 137 | + int m = tA ? ma.getNumColumns() : ma.getNumRows(); |
| 138 | + int n = tB ? mb.getNumRows() : mb.getNumColumns(); |
| 139 | + int k = tA ? ma.getNumRows() : ma.getNumColumns(); |
| 140 | + |
| 141 | + if (!sparseA && !sparseB) { |
| 142 | + LibMatrixMult.matrixMultDenseDenseMM( |
| 143 | + ma.getDenseBlock(), |
| 144 | + mb.getDenseBlock(), |
| 145 | + mc.getDenseBlock(), |
| 146 | + tA, tB, n, k, 0, m, 0, n |
| 147 | + ); |
| 148 | + } |
| 149 | + else if (sparseA && !sparseB) { |
| 150 | + int cd = tB ? mb.getNumColumns() : mb.getNumRows(); |
| 151 | + long xsp = (long) m * cd / Math.max(1L, ma.getNonZeros()); |
| 152 | + LibMatrixMult.matrixMultSparseDenseMM( |
| 153 | + ma.getSparseBlock(), |
| 154 | + mb.getDenseBlock(), |
| 155 | + mc.getDenseBlock(), |
| 156 | + tA, tB, n, cd, xsp, 0, m |
| 157 | + ); |
| 158 | + } |
| 159 | + else if (!sparseA && sparseB) { |
| 160 | + long xsp = (long) ma.getNumRows() * ma.getNumColumns() / Math.max(1L, ma.getNonZeros()); |
| 161 | + LibMatrixMult.matrixMultDenseSparseMM( |
| 162 | + ma.getDenseBlock(), |
| 163 | + mb.getSparseBlock(), |
| 164 | + mc.getDenseBlock(), |
| 165 | + tA, tB, n, k, xsp, 0, m |
| 166 | + ); |
| 167 | + } |
| 168 | + mc.recomputeNonZeros(); |
90 | 169 | } |
91 | 170 | } |
0 commit comments