From 76bfe5face220d147ba5ae227f4a62c6295384b2 Mon Sep 17 00:00:00 2001 From: mike0609king Date: Fri, 17 Jul 2026 18:32:22 +0200 Subject: [PATCH 01/10] tests: Add new matrix type for ifelse test Adapt the R tests to work with column and row vectors (internally those are matrices). --- .../functions/ternary/FullIfElseTest.java | 98 +++++++++++-------- .../scripts/functions/ternary/TernaryIfElse.R | 8 +- 2 files changed, 66 insertions(+), 40 deletions(-) diff --git a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java index 09306eb5fd8..2c13f0685cc 100644 --- a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java +++ b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java @@ -44,6 +44,10 @@ public class FullIfElseTest extends AutomatedTestBase private final static double sparsity1 = 0.6; private final static double sparsity2 = 0.1; + private enum MatType { + MATRIX, COL, ROW, SCALAR + } + @Override public void setUp() { TestUtils.clearAssertionInformation(); @@ -52,167 +56,167 @@ public void setUp() { @Test public void testScalarScalarScalarDenseCP() { - runIfElseTest(false, false, false, false, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); } @Test public void testMatrixScalarScalarDenseCP() { - runIfElseTest(true, false, false, false, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); } @Test public void testScalarMatrixScalarDenseCP() { - runIfElseTest(false, true, false, false, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); } @Test public void testMatrixMatrixScalarDenseCP() { - runIfElseTest(true, true, false, false, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); } @Test public void testScalarScalarMatrixDenseCP() { - runIfElseTest(false, false, true, false, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); } @Test public void testMatrixScalarMatrixDenseCP() { - runIfElseTest(true, false, true, false, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); } @Test public void testScalarMatrixMatrixDenseCP() { - runIfElseTest(false, true, true, false, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); } @Test public void testMatrixMatrixMatrixDenseCP() { - runIfElseTest(true, true, true, false, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); } @Test public void testScalarScalarScalarSparseCP() { - runIfElseTest(false, false, false, true, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); } @Test public void testMatrixScalarScalarSparseCP() { - runIfElseTest(true, false, false, true, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); } @Test public void testScalarMatrixScalarSparseCP() { - runIfElseTest(false, true, false, true, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); } @Test public void testMatrixMatrixScalarSparseCP() { - runIfElseTest(true, true, false, true, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); } @Test public void testScalarScalarMatrixSparseCP() { - runIfElseTest(false, false, true, true, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); } @Test public void testMatrixScalarMatrixSparseCP() { - runIfElseTest(true, false, true, true, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); } @Test public void testScalarMatrixMatrixSparseCP() { - runIfElseTest(false, true, true, true, ExecType.CP); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); } @Test public void testMatrixMatrixMatrixSparseCP() { - runIfElseTest(true, true, true, true, ExecType.CP); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); } //SPARK @Test public void testScalarScalarScalarDenseSP() { - runIfElseTest(false, false, false, false, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, false, ExecType.SPARK); } @Test public void testMatrixScalarScalarDenseSP() { - runIfElseTest(true, false, false, false, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, false, ExecType.SPARK); } @Test public void testScalarMatrixScalarDenseSP() { - runIfElseTest(false, true, false, false, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, false, ExecType.SPARK); } @Test public void testMatrixMatrixScalarDenseSP() { - runIfElseTest(true, true, false, false, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, false, ExecType.SPARK); } @Test public void testScalarScalarMatrixDenseSP() { - runIfElseTest(false, false, true, false, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, false, ExecType.SPARK); } @Test public void testMatrixScalarMatrixDenseSP() { - runIfElseTest(true, false, true, false, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, false, ExecType.SPARK); } @Test public void testScalarMatrixMatrixDenseSP() { - runIfElseTest(false, true, true, false, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, false, ExecType.SPARK); } @Test public void testMatrixMatrixMatrixDenseSP() { - runIfElseTest(true, true, true, false, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, false, ExecType.SPARK); } @Test public void testScalarScalarScalarSparseSP() { - runIfElseTest(false, false, false, true, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, true, ExecType.SPARK); } @Test public void testMatrixScalarScalarSparseSP() { - runIfElseTest(true, false, false, true, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, true, ExecType.SPARK); } @Test public void testScalarMatrixScalarSparseSP() { - runIfElseTest(false, true, false, true, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, true, ExecType.SPARK); } @Test public void testMatrixMatrixScalarSparseSP() { - runIfElseTest(true, true, false, true, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, true, ExecType.SPARK); } @Test public void testScalarScalarMatrixSparseSP() { - runIfElseTest(false, false, true, true, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, true, ExecType.SPARK); } @Test public void testMatrixScalarMatrixSparseSP() { - runIfElseTest(true, false, true, true, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, true, ExecType.SPARK); } @Test public void testScalarMatrixMatrixSparseSP() { - runIfElseTest(false, true, true, true, ExecType.SPARK); + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, true, ExecType.SPARK); } @Test public void testMatrixMatrixMatrixSparseSP() { - runIfElseTest(true, true, true, true, ExecType.SPARK); + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, true, ExecType.SPARK); } - private void runIfElseTest(boolean matrix1, boolean matrix2, boolean matrix3, boolean sparse, ExecType et){ + private void runIfElseTest(MatType mtype1, MatType mtype2, MatType mtype3, boolean sparse, ExecType et){ setOutputBuffering(true); //rtplatform for MR ExecMode platformOld = rtplatform; @@ -239,12 +243,11 @@ private void runIfElseTest(boolean matrix1, boolean matrix2, boolean matrix3, bo rCmd = "Rscript" + " " + fullRScriptName + " " + inputDir() + " " + expectedDir(); //generate actual datasets (matrices and scalars) - double sparsity = sparse ? sparsity2 : sparsity1; - double[][] A = matrix1 ? getRandomMatrix(rows, cols, 0, 1, sparsity, 1) : getScalar(1); + double[][] A = getMatrixOfType(mtype1, sparse, 1); writeInputMatrixWithMTD("A", A, true); - double[][] B = matrix2 ? getRandomMatrix(rows, cols, 0, 1, sparsity, 2) : getScalar(2); + double[][] B = getMatrixOfType(mtype2, sparse, 2); writeInputMatrixWithMTD("B", B, true); - double[][] C = matrix3 ? getRandomMatrix(rows, cols, 0, 1, sparsity, 3) : getScalar(3); + double[][] C = getMatrixOfType(mtype2, sparse, 3); writeInputMatrixWithMTD("C", C, true); //run test cases @@ -263,7 +266,24 @@ private void runIfElseTest(boolean matrix1, boolean matrix2, boolean matrix3, bo } } - private static double[][] getScalar(int input) { - return new double[][]{{7d*input}}; + private double[][] getMatrixOfType(MatType mtype, boolean sparse, long seed) { + double[][] ret = null; + double sparsity = sparse ? sparsity2 : sparsity1; + switch(mtype) { + case SCALAR: + ret = getRandomMatrix(1, 1, 0, 1, sparsity, seed); + break; + case MATRIX: + ret = getRandomMatrix(rows, cols, 0, 1, sparsity, seed); + break; + case COL: + ret = getRandomMatrix(1, cols, 0, 1, sparsity, seed); + break; + case ROW: + ret = getRandomMatrix(rows, 1, 0, 1, sparsity, seed); + break; + default: + } + return ret; } } diff --git a/src/test/scripts/functions/ternary/TernaryIfElse.R b/src/test/scripts/functions/ternary/TernaryIfElse.R index a1e8a030021..d0e2f49c634 100644 --- a/src/test/scripts/functions/ternary/TernaryIfElse.R +++ b/src/test/scripts/functions/ternary/TernaryIfElse.R @@ -31,14 +31,20 @@ m = max(max(nrow(A), nrow(B)), nrow(C)) n = max(max(ncol(A), ncol(B)), ncol(C)) if( nrow(A)==1 ) { + A = matrix(A, m, n, byrow=TRUE); +} else if ( ncol(A) == 1 ) { A = matrix(A, m, n); } if( nrow(B)==1 ) { + B = matrix(B, m, n, byrow=TRUE); +} else if ( ncol(B)==1 ) { B = matrix(B, m, n); } if( nrow(C)==1 ) { + C = matrix(C, m, n, byrow=TRUE); +} else if( ncol(C)==1 ) { C = matrix(C, m, n); -} +} R = matrix(ifelse(as.vector(A), as.vector(B), as.vector(C)), m, n); From 1b32aefa3d12d50264acb270d693e512d887c89b Mon Sep 17 00:00:00 2001 From: mike0609king Date: Fri, 17 Jul 2026 19:22:15 +0200 Subject: [PATCH 02/10] tests: Add test for all combination of matrix types --- .../functions/ternary/FullIfElseTest.java | 641 ++++++++++++++++-- 1 file changed, 601 insertions(+), 40 deletions(-) diff --git a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java index 2c13f0685cc..173af94bfc8 100644 --- a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java +++ b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java @@ -55,85 +55,646 @@ public void setUp() { } @Test - public void testScalarScalarScalarDenseCP() { - runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + public void testScalarScalarScalarSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); } - + @Test - public void testMatrixScalarScalarDenseCP() { - runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + public void testScalarScalarColSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.COL, true, ExecType.CP); } - + @Test - public void testScalarMatrixScalarDenseCP() { - runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + public void testScalarScalarRowSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.ROW, true, ExecType.CP); } - + @Test - public void testMatrixMatrixScalarDenseCP() { - runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + public void testScalarScalarMatrixSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); } - + @Test - public void testScalarScalarMatrixDenseCP() { - runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + public void testScalarColScalarSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.SCALAR, true, ExecType.CP); } - + @Test - public void testMatrixScalarMatrixDenseCP() { - runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + public void testScalarColColSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.COL, true, ExecType.CP); } - + @Test - public void testScalarMatrixMatrixDenseCP() { - runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + public void testScalarColRowSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.ROW, true, ExecType.CP); } - + @Test - public void testMatrixMatrixMatrixDenseCP() { - runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + public void testScalarColMatrixSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.MATRIX, true, ExecType.CP); } @Test - public void testScalarScalarScalarSparseCP() { - runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); + public void testScalarRowScalarSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.SCALAR, true, ExecType.CP); } - + @Test - public void testMatrixScalarScalarSparseCP() { - runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); + public void testScalarRowColSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.COL, true, ExecType.CP); } - + + @Test + public void testScalarRowRowSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testScalarRowMatrixSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.MATRIX, true, ExecType.CP); + } + @Test public void testScalarMatrixScalarSparseCP() { runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); } - + @Test - public void testMatrixMatrixScalarSparseCP() { - runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); + public void testScalarMatrixColSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.COL, true, ExecType.CP); } - + @Test - public void testScalarScalarMatrixSparseCP() { - runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); + public void testScalarMatrixRowSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.ROW, true, ExecType.CP); } - + + @Test + public void testScalarMatrixMatrixSparseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testColScalarScalarSparseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testColScalarColSparseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.COL, true, ExecType.CP); + } + + @Test + public void testColScalarRowSparseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testColScalarMatrixSparseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testColColScalarSparseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testColColColSparseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.COL, true, ExecType.CP); + } + + @Test + public void testColColRowSparseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testColColMatrixSparseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testColRowScalarSparseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testColRowColSparseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.COL, true, ExecType.CP); + } + + @Test + public void testColRowRowSparseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testColRowMatrixSparseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testColMatrixScalarSparseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testColMatrixColSparseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.COL, true, ExecType.CP); + } + + @Test + public void testColMatrixRowSparseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testColMatrixMatrixSparseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testRowScalarScalarSparseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testRowScalarColSparseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.COL, true, ExecType.CP); + } + + @Test + public void testRowScalarRowSparseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testRowScalarMatrixSparseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testRowColScalarSparseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testRowColColSparseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.COL, true, ExecType.CP); + } + + @Test + public void testRowColRowSparseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testRowColMatrixSparseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testRowRowScalarSparseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testRowRowColSparseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.COL, true, ExecType.CP); + } + + @Test + public void testRowRowRowSparseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testRowRowMatrixSparseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testRowMatrixScalarSparseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testRowMatrixColSparseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.COL, true, ExecType.CP); + } + + @Test + public void testRowMatrixRowSparseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testRowMatrixMatrixSparseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testMatrixScalarScalarSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testMatrixScalarColSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.COL, true, ExecType.CP); + } + + @Test + public void testMatrixScalarRowSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.ROW, true, ExecType.CP); + } + @Test public void testMatrixScalarMatrixSparseCP() { runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, true, ExecType.CP); } - + @Test - public void testScalarMatrixMatrixSparseCP() { - runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); + public void testMatrixColScalarSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.SCALAR, true, ExecType.CP); } - + + @Test + public void testMatrixColColSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.COL, true, ExecType.CP); + } + + @Test + public void testMatrixColRowSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testMatrixColMatrixSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testMatrixRowScalarSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testMatrixRowColSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.COL, true, ExecType.CP); + } + + @Test + public void testMatrixRowRowSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.ROW, true, ExecType.CP); + } + + @Test + public void testMatrixRowMatrixSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.MATRIX, true, ExecType.CP); + } + + @Test + public void testMatrixMatrixScalarSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, true, ExecType.CP); + } + + @Test + public void testMatrixMatrixColSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.COL, true, ExecType.CP); + } + + @Test + public void testMatrixMatrixRowSparseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.ROW, true, ExecType.CP); + } + @Test public void testMatrixMatrixMatrixSparseCP() { runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, true, ExecType.CP); } + @Test + public void testScalarScalarScalarDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testScalarScalarColDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.COL, false, ExecType.CP); + } + + @Test + public void testScalarScalarRowDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testScalarScalarMatrixDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testScalarColScalarDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testScalarColColDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.COL, false, ExecType.CP); + } + + @Test + public void testScalarColRowDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testScalarColMatrixDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.COL, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testScalarRowScalarDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testScalarRowColDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.COL, false, ExecType.CP); + } + + @Test + public void testScalarRowRowDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testScalarRowMatrixDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.ROW, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testScalarMatrixScalarDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testScalarMatrixColDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.COL, false, ExecType.CP); + } + + @Test + public void testScalarMatrixRowDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testScalarMatrixMatrixDenseCP() { + runIfElseTest(MatType.SCALAR, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testColScalarScalarDenseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testColScalarColDenseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.COL, false, ExecType.CP); + } + + @Test + public void testColScalarRowDenseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testColScalarMatrixDenseCP() { + runIfElseTest(MatType.COL, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testColColScalarDenseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testColColColDenseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.COL, false, ExecType.CP); + } + + @Test + public void testColColRowDenseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testColColMatrixDenseCP() { + runIfElseTest(MatType.COL, MatType.COL, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testColRowScalarDenseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testColRowColDenseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.COL, false, ExecType.CP); + } + + @Test + public void testColRowRowDenseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testColRowMatrixDenseCP() { + runIfElseTest(MatType.COL, MatType.ROW, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testColMatrixScalarDenseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testColMatrixColDenseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.COL, false, ExecType.CP); + } + + @Test + public void testColMatrixRowDenseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testColMatrixMatrixDenseCP() { + runIfElseTest(MatType.COL, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testRowScalarScalarDenseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testRowScalarColDenseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.COL, false, ExecType.CP); + } + + @Test + public void testRowScalarRowDenseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testRowScalarMatrixDenseCP() { + runIfElseTest(MatType.ROW, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testRowColScalarDenseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testRowColColDenseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.COL, false, ExecType.CP); + } + + @Test + public void testRowColRowDenseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testRowColMatrixDenseCP() { + runIfElseTest(MatType.ROW, MatType.COL, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testRowRowScalarDenseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testRowRowColDenseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.COL, false, ExecType.CP); + } + + @Test + public void testRowRowRowDenseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testRowRowMatrixDenseCP() { + runIfElseTest(MatType.ROW, MatType.ROW, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testRowMatrixScalarDenseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testRowMatrixColDenseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.COL, false, ExecType.CP); + } + + @Test + public void testRowMatrixRowDenseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testRowMatrixMatrixDenseCP() { + runIfElseTest(MatType.ROW, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testMatrixScalarScalarDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testMatrixScalarColDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.COL, false, ExecType.CP); + } + + @Test + public void testMatrixScalarRowDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testMatrixScalarMatrixDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.SCALAR, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testMatrixColScalarDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testMatrixColColDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.COL, false, ExecType.CP); + } + + @Test + public void testMatrixColRowDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testMatrixColMatrixDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.COL, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testMatrixRowScalarDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testMatrixRowColDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.COL, false, ExecType.CP); + } + + @Test + public void testMatrixRowRowDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testMatrixRowMatrixDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.ROW, MatType.MATRIX, false, ExecType.CP); + } + + @Test + public void testMatrixMatrixScalarDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.SCALAR, false, ExecType.CP); + } + + @Test + public void testMatrixMatrixColDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.COL, false, ExecType.CP); + } + + @Test + public void testMatrixMatrixRowDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.ROW, false, ExecType.CP); + } + + @Test + public void testMatrixMatrixMatrixDenseCP() { + runIfElseTest(MatType.MATRIX, MatType.MATRIX, MatType.MATRIX, false, ExecType.CP); + } + + //SPARK @Test From f83f3c291f78b01629b3306594b4b7092e739a78 Mon Sep 17 00:00:00 2001 From: mike0609king Date: Fri, 17 Jul 2026 20:42:38 +0200 Subject: [PATCH 03/10] fix: Typos --- .../apache/sysds/test/functions/ternary/FullIfElseTest.java | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java index 173af94bfc8..d28bea2dc9f 100644 --- a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java +++ b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java @@ -808,7 +808,7 @@ private void runIfElseTest(MatType mtype1, MatType mtype2, MatType mtype3, boole writeInputMatrixWithMTD("A", A, true); double[][] B = getMatrixOfType(mtype2, sparse, 2); writeInputMatrixWithMTD("B", B, true); - double[][] C = getMatrixOfType(mtype2, sparse, 3); + double[][] C = getMatrixOfType(mtype3, sparse, 3); writeInputMatrixWithMTD("C", C, true); //run test cases @@ -838,10 +838,10 @@ private double[][] getMatrixOfType(MatType mtype, boolean sparse, long seed) { ret = getRandomMatrix(rows, cols, 0, 1, sparsity, seed); break; case COL: - ret = getRandomMatrix(1, cols, 0, 1, sparsity, seed); + ret = getRandomMatrix(rows, 1, 0, 1, sparsity, seed); break; case ROW: - ret = getRandomMatrix(rows, 1, 0, 1, sparsity, seed); + ret = getRandomMatrix(1, cols, 0, 1, sparsity, seed); break; default: } From ad8ef68ed7909e3f686a12ebb3be5bd32f44e9ef Mon Sep 17 00:00:00 2001 From: mike0609king Date: Sat, 18 Jul 2026 01:28:08 +0200 Subject: [PATCH 04/10] tests: Adapt dml script to work with new MatTypes --- .../functions/ternary/TernaryIfElse.dml | 20 +++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/src/test/scripts/functions/ternary/TernaryIfElse.dml b/src/test/scripts/functions/ternary/TernaryIfElse.dml index 12a11cc1253..9dca7c3d7f5 100644 --- a/src/test/scripts/functions/ternary/TernaryIfElse.dml +++ b/src/test/scripts/functions/ternary/TernaryIfElse.dml @@ -23,21 +23,25 @@ A = read($1); B = read($2); C = read($3); -if( nrow(A)==1 & nrow(B)==1 & nrow(C)==1 ) +isscalar = function(matrix[double] A) return (boolean C) { + C = nrow(A)==1 & ncol(A)==1 +} + +if( isscalar(A) & isscalar(B) & isscalar(C) ) R = as.matrix(ifelse(as.scalar(A), as.scalar(B), as.scalar(C))); -else if( nrow(A)>1 & nrow(B)==1 & nrow(C)==1 ) +else if( !isscalar(A) & isscalar(B) & isscalar(C)) R = ifelse(A, as.scalar(B), as.scalar(C)); -else if( nrow(A)==1 & nrow(B)>1 & nrow(C)==1 ) +else if( isscalar(A) & !isscalar(B) & isscalar(C) ) R = ifelse(as.scalar(A), B, as.scalar(C)); -else if( nrow(A)>1 & nrow(B)>1 & nrow(C)==1 ) +else if( !isscalar(A) & !isscalar(B) & isscalar(C) ) R = ifelse(A, B, as.scalar(C)); -else if( nrow(A)==1 & nrow(B)==1 & nrow(C)>1 ) +else if( isscalar(A)==1 & isscalar(B) & !isscalar(C) ) R = ifelse(as.scalar(A), as.scalar(B), C); -else if( nrow(A)>1 & nrow(B)==1 & nrow(C)>1 ) +else if( !isscalar(A) & isscalar(B) & !isscalar(C) ) R = ifelse(A, as.scalar(B), C); -else if( nrow(A)==1 & nrow(B)>1 & nrow(C)>1 ) +else if( isscalar(A)==1 & !isscalar(B) & !isscalar(C) ) R = ifelse(as.scalar(A), B, C); -else if( nrow(A)>1 & nrow(B)>1 & nrow(C)>1 ) +else if( !isscalar(A) & !isscalar(B) & !isscalar(C) ) R = ifelse(A, B, C); write(R, $4); From 61f93ca8e6a194043eae24af641f533910609cc6 Mon Sep 17 00:00:00 2001 From: mike0609king Date: Sat, 18 Jul 2026 18:32:08 +0200 Subject: [PATCH 05/10] feat: Relax constraints for ifelse for matrix checking --- .../parser/BuiltinFunctionExpression.java | 36 +++++++++++++------ .../runtime/compress/lib/CLALibTernaryOp.java | 2 +- .../runtime/matrix/data/MatrixBlock.java | 23 +++++++++--- 3 files changed, 44 insertions(+), 17 deletions(-) diff --git a/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java b/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java index ab0c7993b4e..0de335ed465 100644 --- a/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java +++ b/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java @@ -2170,19 +2170,33 @@ private void setBinaryOutputProperties(DataIdentifier output) { } private void setTernaryOutputProperties(DataIdentifier output, boolean conditional) { - DataType dt1 = getFirstExpr().getOutput().getDataType(); - DataType dt2 = getSecondExpr().getOutput().getDataType(); - DataType dt3 = getThirdExpr().getOutput().getDataType(); - DataType dtOut = (dt1.isMatrix() || dt2.isMatrix() || dt3.isMatrix()) ? - DataType.MATRIX : DataType.SCALAR; - if( dt1==DataType.MATRIX && dt2==DataType.MATRIX ) - checkMatchingDimensions(getFirstExpr(), getSecondExpr(), false, conditional); - if( dt1==DataType.MATRIX && dt3==DataType.MATRIX ) - checkMatchingDimensions(getFirstExpr(), getThirdExpr(), false, conditional); - if( dt2==DataType.MATRIX && dt3==DataType.MATRIX ) - checkMatchingDimensions(getSecondExpr(), getThirdExpr(), false, conditional); + Expression expr1 = getFirstExpr(); + Expression expr2 = getSecondExpr(); + Expression expr3 = getThirdExpr(); + DataType dt1 = expr1.getOutput().getDataType(); + DataType dt2 = expr2.getOutput().getDataType(); + DataType dt3 = expr3.getOutput().getDataType(); + final long r1 = expr1.getOutput().getDim1(); + final long r2 = expr2.getOutput().getDim1(); + final long r3 = expr3.getOutput().getDim1(); + final long c1 = expr1.getOutput().getDim2(); + final long c2 = expr2.getOutput().getDim2(); + final long c3 = expr3.getOutput().getDim2(); + final long m = Math.max(Math.max(r1, r2), r3); + final long n = Math.max(Math.max(c1, c2), c3); + + boolean unknownDim = (r1 == -1 || r2 == -1 || r3 == -1 || c1 == -1 || c2 == -1 || c3 == -1); + if (!unknownDim && ((r1 != 1 && r1 != m) || (r2 != 1 && r2 != m) + || (r3 != 1 && r3 != m) || (c1 != 1 && c1 != n) + || (c2 != 1 && c2 != n) || (c3 != 1 && c3 != n))) { + raiseValidateError("Mismatch in matrix dimensions of parameters for function " + + this.getOpCode(), conditional, LanguageErrorCodes.INVALID_PARAMETERS); + } + MatrixCharacteristics dims1 = getBinaryMatrixCharacteristics(getFirstExpr(), getSecondExpr()); MatrixCharacteristics dims2 = getBinaryMatrixCharacteristics(getSecondExpr(), getThirdExpr()); + DataType dtOut = (dt1.isMatrix() || dt2.isMatrix() || dt3.isMatrix()) ? + DataType.MATRIX : DataType.SCALAR; output.setDataType(dtOut); output.setValueType(dtOut==DataType.MATRIX ? ValueType.FP64 : computeValueType(getSecondExpr(), getThirdExpr(), true)); diff --git a/src/main/java/org/apache/sysds/runtime/compress/lib/CLALibTernaryOp.java b/src/main/java/org/apache/sysds/runtime/compress/lib/CLALibTernaryOp.java index 8dae24df79c..3576484406e 100644 --- a/src/main/java/org/apache/sysds/runtime/compress/lib/CLALibTernaryOp.java +++ b/src/main/java/org/apache/sysds/runtime/compress/lib/CLALibTernaryOp.java @@ -59,7 +59,7 @@ public static MatrixBlock ternaryOperations(TernaryOperator op, MatrixBlock m1, final int n = Math.max(Math.max(c1, c2), c3); // double check that the dimensions are valid. - MatrixBlock.ternaryOperationCheck(s1, s2, s3, m, r1, r2, r3, n, c1, c2, c3); + MatrixBlock.ternaryOperationCheck(op, s1, s2, s3, m, r1, r2, r3, n, c1, c2, c3); MatrixBlock ret; diff --git a/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java b/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java index 361d190bd02..1835cd4218a 100644 --- a/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java +++ b/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java @@ -3080,7 +3080,7 @@ public MatrixBlock ternaryOperations(TernaryOperator op, MatrixBlock m2, MatrixB final int n = Math.max(Math.max(c1, c2), c3); final long nnz = nonZeros; - ternaryOperationCheck(s1, s2, s3, m, r1, r2, r3, n, c1, c2, c3); + ternaryOperationCheck(op, s1, s2, s3, m, r1, r2, r3, n, c1, c2, c3); //prepare result if( op.fn instanceof IfElse && (s1 || nnz==0 || nnz==(long)m*n) ) { @@ -3141,11 +3141,24 @@ else if (s2 != s3 && (op.fn instanceof PlusMultiply || op.fn instanceof MinusMul return ret; } - public static void ternaryOperationCheck(boolean s1, boolean s2, boolean s3, int m, int r1, int r2, int r3, int n, int c1, int c2, int c3){ + public static void ternaryOperationCheck(TernaryOperator op, boolean s1, boolean s2, boolean s3, int m, int r1, int r2, int r3, int n, int c1, int c2, int c3){ //error handling - if( (!s1 && (r1 != m || c1 != n)) - || (!s2 && (r2 != m || c2 != n)) - || (!s3 && (r3 != m || c3 != n)) ) { + boolean error = false; + if (op.fn instanceof IfElse) { + error = ((r1 != 1 && r1 != m) + || (r2 != 1 && r2 != m) + || (r3 != 1 && r3 != m) + || (c1 != 1 && c1 != n) + || (c2 != 1 && c2 != n) + || (c3 != 1 && c3 != n)); + } + else { + error = ((!s1 && (r1 != m || c1 != n)) + || (!s2 && (r2 != m || c2 != n)) + || (!s3 && (r3 != m || c3 != n))); + } + + if (error) { throw new DMLRuntimeException("Block sizes are not matched for ternary cell operations: " + r1 + "x" + c1 + " vs " + r2 + "x" + c2 + " vs " + r3 + "x" + c3); } From 81ce25addd8990712c9ff1daf535f743d48100e5 Mon Sep 17 00:00:00 2001 From: mike0609king Date: Sat, 18 Jul 2026 18:32:56 +0200 Subject: [PATCH 06/10] feat: implement broadcasting for cellwise ternary operation --- .../sysds/runtime/matrix/data/LibMatrixTercell.java | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java b/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java index 72aa51ed5c6..14c2d36cfc2 100644 --- a/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java +++ b/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java @@ -92,11 +92,18 @@ private static long unsafeTernary(MatrixBlock m1, MatrixBlock m2, MatrixBlock m3 //basic ternary operations (all combinations sparse/dense) int n = ret.clen; long lnnz = 0; + final int r1 = m1.getNumRows(); + final int r2 = m2.getNumRows(); + final int r3 = m3.getNumRows(); + final int c1 = m1.getNumColumns(); + final int c2 = m2.getNumColumns(); + final int c3 = m3.getNumColumns(); + for( int i=rl; i Date: Sun, 19 Jul 2026 00:17:29 +0200 Subject: [PATCH 07/10] fix: only matrices need to be checked for dimensions --- .../apache/sysds/parser/BuiltinFunctionExpression.java | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java b/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java index 0de335ed465..98362dc46ba 100644 --- a/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java +++ b/src/main/java/org/apache/sysds/parser/BuiltinFunctionExpression.java @@ -2186,9 +2186,12 @@ private void setTernaryOutputProperties(DataIdentifier output, boolean condition final long n = Math.max(Math.max(c1, c2), c3); boolean unknownDim = (r1 == -1 || r2 == -1 || r3 == -1 || c1 == -1 || c2 == -1 || c3 == -1); - if (!unknownDim && ((r1 != 1 && r1 != m) || (r2 != 1 && r2 != m) - || (r3 != 1 && r3 != m) || (c1 != 1 && c1 != n) - || (c2 != 1 && c2 != n) || (c3 != 1 && c3 != n))) { + if (!unknownDim && ((dt1 == DataType.MATRIX && r1 != 1 && r1 != m) + || (dt2 == DataType.MATRIX && r2 != 1 && r2 != m) + || (dt3 == DataType.MATRIX && r3 != 1 && r3 != m) + || (dt1 == DataType.MATRIX && c1 != 1 && c1 != n) + || (dt2 == DataType.MATRIX && c2 != 1 && c2 != n) + || (dt3 == DataType.MATRIX && c3 != 1 && c3 != n))) { raiseValidateError("Mismatch in matrix dimensions of parameters for function " + this.getOpCode(), conditional, LanguageErrorCodes.INVALID_PARAMETERS); } From 2d59c196e860d4bd79edee21313bd36e0e9717de Mon Sep 17 00:00:00 2001 From: mike0609king Date: Sun, 19 Jul 2026 00:42:51 +0200 Subject: [PATCH 08/10] feat: Add broadcasting to the ifelse operator optimization --- .../runtime/matrix/data/MatrixBlock.java | 83 ++++++++++++++----- 1 file changed, 64 insertions(+), 19 deletions(-) diff --git a/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java b/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java index 1835cd4218a..cc79123c644 100644 --- a/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java +++ b/src/main/java/org/apache/sysds/runtime/matrix/data/MatrixBlock.java @@ -3084,25 +3084,7 @@ public MatrixBlock ternaryOperations(TernaryOperator op, MatrixBlock m2, MatrixB //prepare result if( op.fn instanceof IfElse && (s1 || nnz==0 || nnz==(long)m*n) ) { - - ret.reset(m, n, false); - - //SPECIAL CASE for shallow-copy if-else - boolean expr = s1 ? (d1 != 0) : (nnz==(long)m*n); - MatrixBlock tmp = expr ? m2 : m3; - if( tmp.rlen==m && tmp.clen==n ) { - //shallow copy incl meta data - ret.copyShallow(tmp); - } - else { - //fill output with given scalar value - double tmpVal = tmp.get(0, 0); - if( tmpVal != 0 ) { - ret.allocateDenseBlock(); - ret.denseBlock.set(tmpVal); - ret.nonZeros = (long)m * n; - } - } + ternaryIfElseCopy(s1, m, n, d1, nnz, m2, m3, ret); } else{ final boolean PM_Or_MM = (op.fn instanceof PlusMultiply || op.fn instanceof MinusMultiply); @@ -3141,6 +3123,69 @@ else if (s2 != s3 && (op.fn instanceof PlusMultiply || op.fn instanceof MinusMul return ret; } + /** + * Copies and applies broadcasting to either second or third input to IfElse operation. + * + * If the first value of the IfElse-operation meets certain conditions, then + * either the m2 or m3 matrix inputs of IfElse is the result of the operation. + * This methods handles the copying if the conditions are met. + * + * @param s1 Flag, whether the first matrix is a scalar. + * @param m Rows of the resulting matrix. + * @param n Columns of the resulting matrix + * @param d1 Value of the entry 0,0 of the first matrix input. + * @param nnz Non-zero entries of the first matrix + * @param m2 Second matrix input of the IfElse operation + * @param m3 Third matrix input of the IfElse operation + * @param ret Result of the operation, where either m2 or m3 is copied into + */ + private void ternaryIfElseCopy(boolean s1, int m, int n, double d1, long nnz, MatrixBlock m2, MatrixBlock m3, MatrixBlock ret) { + ret.reset(m, n, false); + + boolean expr = s1 ? (d1 != 0) : (nnz==(long)m*n); + MatrixBlock tmp = expr ? m2 : m3; + if (tmp.rlen==m && tmp.clen==n) { + //shallow copy incl meta data + ret.copyShallow(tmp); + } + else if (tmp.rlen==m && tmp.clen==1) { + ret.allocateDenseBlock(); + ret.nonZeros = 0; + for (int i = 0; i < m; i++) { + double tmpVal = tmp.get(i, 0); + if (tmpVal != 0) { + ret.denseBlock.fillRow(i, tmp.get(i, 0)); + ret.nonZeros += n; + } + } + ret.examSparsity(); + } + else if (tmp.rlen==1 && tmp.clen==n) { + if (tmp.nonZeros != 0) { + ret.allocateDenseBlock(); + ret.nonZeros = 0; + double[] tmpArr = new double[n]; + for (int i = 0; i < n; i++) { + tmpArr[i] = tmp.get(0, i); + } + for (int i = 0; i < m; i++) { + ret.denseBlock.set(i, tmpArr); + ret.nonZeros += tmp.nonZeros; + } + ret.examSparsity(); + } + } + else { + //fill output with given scalar value + double tmpVal = tmp.get(0, 0); + if (tmpVal != 0) { + ret.allocateDenseBlock(); + ret.denseBlock.set(tmpVal); + ret.nonZeros = (long)m * n; + } + } + } + public static void ternaryOperationCheck(TernaryOperator op, boolean s1, boolean s2, boolean s3, int m, int r1, int r2, int r3, int n, int c1, int c2, int c3){ //error handling boolean error = false; From eaf2eb60a72c208f6a0a6075878f59bfc1a0a26e Mon Sep 17 00:00:00 2001 From: mike0609king Date: Sun, 19 Jul 2026 17:31:13 +0200 Subject: [PATCH 09/10] feat: Optimizations for row and col ifelse --- .../runtime/matrix/data/LibMatrixTercell.java | 70 +++++++++++++++++-- 1 file changed, 66 insertions(+), 4 deletions(-) diff --git a/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java b/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java index 14c2d36cfc2..89da76aa880 100644 --- a/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java +++ b/src/main/java/org/apache/sysds/runtime/matrix/data/LibMatrixTercell.java @@ -27,6 +27,7 @@ import java.util.concurrent.Future; import org.apache.sysds.runtime.DMLRuntimeException; +import org.apache.sysds.runtime.functionobjects.IfElse; import org.apache.sysds.runtime.matrix.operators.TernaryOperator; import org.apache.sysds.runtime.util.CommonThreadPool; import org.apache.sysds.runtime.util.UtilFunctions; @@ -88,6 +89,63 @@ public static void tercellOp(MatrixBlock m1, MatrixBlock m2, MatrixBlock m3, Mat private static long unsafeTernary(MatrixBlock m1, MatrixBlock m2, MatrixBlock m3, MatrixBlock ret, TernaryOperator op, boolean s1, boolean s2, boolean s3, double d1, double d2, double d3, int rl, int ru) + { + if(op.fn instanceof IfElse) { + return unsafeTernaryIfElse(m1, m2, m3, ret, op, s1, s2, + s3, d1, d2, d3, rl, ru); + } + else { + return unsafeTernaryDefault(m1, m2, m3, ret, op, s1, s2, + s3, d1, d2, d3, rl, ru); + } + } + + private static long unsafeTernaryIfElse(MatrixBlock m1, MatrixBlock m2, MatrixBlock m3, MatrixBlock ret, + TernaryOperator op, boolean s1, boolean s2, boolean s3, double d1, double d2, double d3, int rl, int ru) + { + // IfElse specific optimizations are applied here + int n = ret.clen; + long lnnz = 0; + final int r1 = m1.getNumRows(); + final int c1 = m1.getNumColumns(); + if (c1 == 1) { + for( int i=rl; i Date: Sun, 19 Jul 2026 17:35:53 +0200 Subject: [PATCH 10/10] tests: Add eps for matrix comparison --- .../sysds/test/functions/ternary/FullIfElseTest.java | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java index d28bea2dc9f..d0c9b955e1d 100644 --- a/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java +++ b/src/test/java/org/apache/sysds/test/functions/ternary/FullIfElseTest.java @@ -37,9 +37,10 @@ public class FullIfElseTest extends AutomatedTestBase private final static String TEST_DIR = "functions/ternary/"; private final static String TEST_CLASS_DIR = TEST_DIR + FullIfElseTest.class.getSimpleName() + "/"; + private final static double eps = 1e-10; private final static int rows = 2111; - private final static int cols = 30; + private final static int cols = 300; private final static double sparsity1 = 0.6; private final static double sparsity2 = 0.1; @@ -799,7 +800,7 @@ private void runIfElseTest(MatType mtype1, MatType mtype2, MatType mtype3, boole String HOME = SCRIPT_DIR + TEST_DIR; fullDMLScriptName = HOME + TEST_NAME1 + ".dml"; - programArgs = new String[]{"-explain","-args", input("A"), input("B"), input("C"), output("R")}; + programArgs = new String[]{"-explain", "-stats", "-args", input("A"), input("B"), input("C"), output("R")}; fullRScriptName = HOME + TEST_NAME1 + ".R"; rCmd = "Rscript" + " " + fullRScriptName + " " + inputDir() + " " + expectedDir(); @@ -818,7 +819,7 @@ private void runIfElseTest(MatType mtype1, MatType mtype2, MatType mtype3, boole //compare output matrices HashMap dmlfile = readDMLMatrixFromOutputDir("R"); HashMap rfile = readRMatrixFromExpectedDir("R"); - TestUtils.compareMatrices(dmlfile, rfile, 0, "Stat-DML", "Stat-R"); + TestUtils.compareMatrices(dmlfile, rfile, eps, "Stat-DML", "Stat-R"); } finally { rtplatform = platformOld;