Skip to content

Commit 4ca22c1

Browse files
committed
[SYSTEMDS-3168] dense-sparse transpose kernels and component tests
1 parent 624b287 commit 4ca22c1

2 files changed

Lines changed: 281 additions & 33 deletions

File tree

src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixMult.java

Lines changed: 169 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1544,6 +1544,175 @@ private static void matrixMultDenseDenseOutSparseVector(MatrixBlock m1, MatrixBl
15441544
}
15451545
}
15461546

1547+
public static void matrixMultDenseSparseMM(DenseBlock a, SparseBlock b, DenseBlock c,
1548+
boolean transA, boolean transB, int n, int cd, long xsp, int rl, int ru)
1549+
{
1550+
if(!transA && !transB){
1551+
// dispatcher parameter mismatch necessitates dummy wrappers
1552+
MatrixBlock m1 = new MatrixBlock(c.numRows(), cd, false);
1553+
m1.setDenseBlock(a);
1554+
1555+
MatrixBlock m2 = new MatrixBlock();
1556+
m2.sparseBlock = b;
1557+
1558+
MatrixBlock ret = new MatrixBlock();
1559+
ret.setDenseBlock(c);
1560+
1561+
matrixMultDenseSparseOutDense(m1, m2, ret, false, rl, ru);
1562+
}
1563+
else if(transA && !transB)
1564+
multDenseSparseTransA(a, b, c, n, cd, xsp, rl, ru);
1565+
else if(!transA && transB)
1566+
multDenseSparseTransB(a, b, c, n, cd, xsp, rl, ru);
1567+
else
1568+
multDenseSparseTransATransB(a, b, c, n, cd, xsp, rl, ru);
1569+
}
1570+
1571+
private static void multDenseSparseTransA(DenseBlock a, SparseBlock b, DenseBlock c,
1572+
int n, int cd, long xsp, int rl, int ru)
1573+
{
1574+
final int blocksizeJ = 1024;
1575+
1576+
for(int k = 0; k < cd; k++) {
1577+
1578+
if(b.isEmpty(k))
1579+
continue;
1580+
1581+
final int bpos = b.pos(k);
1582+
final int blen = b.size(k);
1583+
final int[] bix = b.indexes(k);
1584+
final double[] bvals = b.values(k);
1585+
1586+
final double[] arow = a.values(k);
1587+
final int apos = a.pos(k);
1588+
1589+
for(int bj = 0; bj < n; bj += blocksizeJ) {
1590+
1591+
int p1 = (bj == 0) ? bpos :
1592+
((b.posFIndexGTE(k, bj) >= 0) ?
1593+
bpos + b.posFIndexGTE(k, bj) :
1594+
bpos + blen);
1595+
1596+
int p2 =
1597+
((b.posFIndexGTE(k, bj + blocksizeJ) >= 0) ?
1598+
bpos + b.posFIndexGTE(k, bj + blocksizeJ) :
1599+
bpos + blen);
1600+
1601+
if(p1 >= p2)
1602+
continue;
1603+
1604+
for(int i = rl; i < ru; i++) {
1605+
final double aval = arow[apos + i];
1606+
1607+
if(aval == 0)
1608+
continue;
1609+
1610+
final double[] cvals = c.values(i);
1611+
final int cix = c.pos(i);
1612+
1613+
vectMultiplyAdd(
1614+
aval,
1615+
bvals,
1616+
cvals,
1617+
bix,
1618+
p1,
1619+
cix,
1620+
p2 - p1);
1621+
}
1622+
}
1623+
}
1624+
}
1625+
1626+
private static void multDenseSparseTransB(DenseBlock a, SparseBlock b, DenseBlock c,
1627+
int n, int cd, long xsp, int rl, int ru)
1628+
{
1629+
if( a.isContiguous() ) {
1630+
final double[] adata = a.values(0);
1631+
for(int i = rl; i < ru; i++) {
1632+
final int apos = i * cd;
1633+
final double[] cvals = c.values(i);
1634+
final int cix = c.pos(i);
1635+
for(int j = 0; j < n; j++) {
1636+
if(b.isEmpty(j))
1637+
continue;
1638+
final int bpos = b.pos(j);
1639+
final int blen = b.size(j);
1640+
final int[] bix = b.indexes(j);
1641+
final double[] bvals = b.values(j);
1642+
cvals[cix + j] = dotProduct(bvals, adata, bix, bpos, apos, blen);
1643+
}
1644+
}
1645+
}
1646+
else {
1647+
for(int i = rl; i < ru; i++) {
1648+
final double[] arow = a.values(i);
1649+
final int apos = a.pos(i);
1650+
final double[] cvals = c.values(i);
1651+
final int cix = c.pos(i);
1652+
for(int j = 0; j < n; j++) {
1653+
if(b.isEmpty(j))
1654+
continue;
1655+
final int bpos = b.pos(j);
1656+
final int blen = b.size(j);
1657+
final int[] bix = b.indexes(j);
1658+
final double[] bvals = b.values(j);
1659+
cvals[cix + j] = dotProduct(bvals, arow, bix, bpos, apos, blen);
1660+
}
1661+
}
1662+
}
1663+
}
1664+
1665+
private static void multDenseSparseTransATransB(DenseBlock a, SparseBlock b, DenseBlock c,
1666+
int n, int cd, long xsp, int rl, int ru)
1667+
{
1668+
final int m = a.numCols();
1669+
if( a.isContiguous() && c.isContiguous() ) {
1670+
final double[] adata = a.values(0);
1671+
final double[] cvals = c.values(0);
1672+
for(int j = 0; j < n; j++) {
1673+
if(b.isEmpty(j))
1674+
continue;
1675+
final int bpos = b.pos(j);
1676+
final int blen = b.size(j);
1677+
final int[] bix = b.indexes(j);
1678+
final double[] bvals = b.values(j);
1679+
for(int p = bpos; p < bpos + blen; p++) {
1680+
final int k = bix[p];
1681+
final double bval = bvals[p];
1682+
if (bval == 0)
1683+
continue;
1684+
final int apos = k * m;
1685+
for(int i = rl; i < ru; i++) {
1686+
cvals[i * n + j] += bval * adata[apos + i];
1687+
}
1688+
}
1689+
}
1690+
}
1691+
else {
1692+
for(int j = 0; j < n; j++) {
1693+
if(b.isEmpty(j))
1694+
continue;
1695+
final int bpos = b.pos(j);
1696+
final int blen = b.size(j);
1697+
final int[] bix = b.indexes(j);
1698+
final double[] bvals = b.values(j);
1699+
for(int p = bpos; p < bpos + blen; p++) {
1700+
final int k = bix[p];
1701+
final double bval = bvals[p];
1702+
if (bval == 0)
1703+
continue;
1704+
final double[] arow = a.values(k);
1705+
final int apos = a.pos(k);
1706+
for(int i = rl; i < ru; i++) {
1707+
final double[] cvals = c.values(i);
1708+
final int cix = c.pos(i);
1709+
cvals[cix + j] += bval * arow[apos + i];
1710+
}
1711+
}
1712+
}
1713+
}
1714+
}
1715+
15471716
private static void matrixMultDenseSparseOutSparse(MatrixBlock m1, MatrixBlock m2, MatrixBlock ret, boolean pm2,
15481717
int rl, int ru) {
15491718
final DenseBlock a = m1.getDenseBlock();

src/test/java/org/apache/sysds/test/component/matrixmult/MatrixMultTransposedTest.java

Lines changed: 112 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -19,73 +19,152 @@
1919

2020
package org.apache.sysds.test.component.matrixmult;
2121

22-
import org.apache.sysds.runtime.data.DenseBlock;
22+
import java.util.Random;
23+
2324
import org.apache.sysds.runtime.matrix.data.LibMatrixMult;
2425
import org.apache.sysds.runtime.matrix.data.LibMatrixReorg;
2526
import org.apache.sysds.runtime.matrix.data.MatrixBlock;
2627
import org.apache.sysds.test.TestUtils;
2728
import org.junit.Test;
2829

29-
import java.util.Random;
30-
3130
public class MatrixMultTransposedTest {
3231

33-
// run multiple random scenarios
3432
@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);
3936
}
4037

4138
@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);
4642
}
4743

4844
@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);
5348
}
5449

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+
}
5779

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();
5988
int m = rand.nextInt(300) + 1;
6089
int n = rand.nextInt(300) + 1;
6190
int k = rand.nextInt(300) + 1;
6291

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+
}
6397

98+
private void runTest(double spA, double spB, boolean tA, boolean tB, int m, int n, int k) throws Exception {
6499
int rowsA = tA ? k : m;
65100
int colsA = tA ? m : k;
66101
int rowsB = tB ? n : k;
67102
int colsB = tB ? k : n;
68103

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);
72106
MatrixBlock mc = new MatrixBlock(m, n, false);
73107
mc.allocateDenseBlock();
74108

75-
DenseBlock a = ma.getDenseBlock();
76-
DenseBlock b = mb.getDenseBlock();
77-
DenseBlock c = mc.getDenseBlock();
109+
runNewKernel(ma, mb, mc, tA, tB);
78110

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);
80114

81-
mc.recomputeNonZeros();
115+
TestUtils.compareMatrices(expected, mc, 1e-8);
116+
}
82117

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+
}
87129

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();
90169
}
91170
}

0 commit comments

Comments
 (0)