From 5615907206da4df8aaf77ed870d3b843f9c9897d Mon Sep 17 00:00:00 2001 From: Jannik Lindemann Date: Wed, 22 Jul 2026 14:10:05 +0200 Subject: [PATCH] [OOC] Wire OOC Instructions with New Primitives --- .../ooc/BinaryOOCInstruction.java | 63 ++++++++----------- .../ooc/DataGenOOCInstruction.java | 42 +++++-------- .../ParameterizedBuiltinOOCInstruction.java | 5 +- .../ooc/TernaryOOCInstruction.java | 59 ++++++++++------- .../sysds/test/functions/ooc/SeqTest.java | 5 ++ 5 files changed, 82 insertions(+), 92 deletions(-) diff --git a/src/main/java/org/apache/sysds/runtime/instructions/ooc/BinaryOOCInstruction.java b/src/main/java/org/apache/sysds/runtime/instructions/ooc/BinaryOOCInstruction.java index c252cd3d1ec..89b6d911000 100644 --- a/src/main/java/org/apache/sysds/runtime/instructions/ooc/BinaryOOCInstruction.java +++ b/src/main/java/org/apache/sysds/runtime/instructions/ooc/BinaryOOCInstruction.java @@ -30,6 +30,7 @@ import org.apache.sysds.runtime.matrix.operators.BinaryOperator; import org.apache.sysds.runtime.matrix.operators.Operator; import org.apache.sysds.runtime.matrix.operators.ScalarOperator; +import org.apache.sysds.runtime.ooc.store.CountingLiveness; import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils; public class BinaryOOCInstruction extends ComputationOOCInstruction { @@ -67,8 +68,8 @@ protected void processMatrixMatrixInstruction(ExecutionContext ec) { OOCStream qOut = new SubscribableTaskQueue<>(); ec.getMatrixObject(output).setStreamHandle(qOut); - final boolean known1 = (m1.getNumRows() >= 0 && m1.getNumColumns() >= 0); - final boolean known2 = (m2.getNumRows() >= 0 && m2.getNumColumns() >= 0); + final boolean known1 = m1.getNumRows() >= 0 && m1.getNumColumns() >= 0 && m1.getBlocksize() > 0; + final boolean known2 = m2.getNumRows() >= 0 && m2.getNumColumns() >= 0 && m2.getBlocksize() > 0; // If dimensions are unknown, we cannot safely detect broadcasting. // Fall back to strict key-based join and let downstream operators validate as needed. @@ -94,36 +95,28 @@ protected void processMatrixMatrixInstruction(ExecutionContext ec) { boolean isRowBroadcast = m1.getNumRows() > 1 && m2.getNumRows() == 1; if (isColBroadcast && !isRowBroadcast) { - OOCStream qIn1 = m1.getStreamHandle(); - OOCStream qIn2 = m2.getStreamHandle(); - final long maxProcessesPerBroadcast = (m1.getNumColumns() + m1.getBlocksize() - 1) / m1.getBlocksize(); - - broadcastJoinOOC(qIn1, qIn2, qOut, (tmp1, b) -> { - IndexedMatrixValue tmpOut = new IndexedMatrixValue(); - tmpOut.set(tmp1.getIndexes(), - tmp1.getValue().binaryOperations((BinaryOperator)_optr, b.getValue().getValue(), tmpOut.getValue())); - - if (b.incrProcessCtrAndGet() >= maxProcessesPerBroadcast) - b.release(); - - return tmpOut; - }, tmp -> tmp.getIndexes().getRowIndex()); + int broadcastBlocks = Math.toIntExact(m2.getDataCharacteristics().getNumRowBlocks()); + int usesPerBlock = Math.toIntExact(m1.getDataCharacteristics().getNumColBlocks()); + OOCInstructionUtils.indexedBroadcastMap(m1.getStreamable(), m2.getStreamable(), qOut, + tmp -> Math.toIntExact(tmp.getIndexes().getRowIndex() - 1), + () -> new CountingLiveness(broadcastBlocks, usesPerBlock), (tmp, broadcast) -> { + IndexedMatrixValue tmpOut = new IndexedMatrixValue(); + tmpOut.set(tmp.getIndexes(), tmp.getValue().binaryOperations((BinaryOperator) _optr, + broadcast.getValue(), tmpOut.getValue())); + return tmpOut; + }, getContext()); } else if (isRowBroadcast && !isColBroadcast) { - OOCStream qIn1 = m1.getStreamHandle(); - OOCStream qIn2 = m2.getStreamHandle(); - final long maxProcessesPerBroadcast = (m1.getNumRows() + m1.getBlocksize() - 1) / m1.getBlocksize(); - - broadcastJoinOOC(qIn1, qIn2, qOut, (tmp1, b) -> { - IndexedMatrixValue tmpOut = new IndexedMatrixValue(); - tmpOut.set(tmp1.getIndexes(), - tmp1.getValue().binaryOperations((BinaryOperator)_optr, b.getValue().getValue(), tmpOut.getValue())); - - if (b.incrProcessCtrAndGet() >= maxProcessesPerBroadcast) - b.release(); - - return tmpOut; - }, tmp -> tmp.getIndexes().getColumnIndex()); + int broadcastBlocks = Math.toIntExact(m2.getDataCharacteristics().getNumColBlocks()); + int usesPerBlock = Math.toIntExact(m1.getDataCharacteristics().getNumRowBlocks()); + OOCInstructionUtils.indexedBroadcastMap(m1.getStreamable(), m2.getStreamable(), qOut, + tmp -> Math.toIntExact(tmp.getIndexes().getColumnIndex() - 1), + () -> new CountingLiveness(broadcastBlocks, usesPerBlock), (tmp, broadcast) -> { + IndexedMatrixValue tmpOut = new IndexedMatrixValue(); + tmpOut.set(tmp.getIndexes(), tmp.getValue().binaryOperations((BinaryOperator) _optr, + broadcast.getValue(), tmpOut.getValue())); + return tmpOut; + }, getContext()); } else { if (m1.getNumColumns() != m2.getNumColumns() || m1.getNumRows() != m2.getNumRows()) @@ -144,15 +137,9 @@ protected void processScalarMatrixInstruction(ExecutionContext ec) { //create thread and process binary operation MatrixObject min = ec.getMatrixObject(input1.isMatrix() ? input1 : input2); - OOCStream qIn = min.getStreamHandle(); OOCStream qOut = createWritableStream(); ec.getMatrixObject(output).setStreamHandle(qOut); - - mapOOC(qIn, qOut, tmp -> { - IndexedMatrixValue tmpOut = new IndexedMatrixValue(); - tmpOut.set(tmp.getIndexes(), - tmp.getValue().scalarOperations(sc_op, new MatrixBlock())); - return tmpOut; - }); + OOCInstructionUtils.equiMapBlock(min.getStreamable(), qOut, + block -> block.scalarOperations(sc_op, new MatrixBlock()), getContext()); } } diff --git a/src/main/java/org/apache/sysds/runtime/instructions/ooc/DataGenOOCInstruction.java b/src/main/java/org/apache/sysds/runtime/instructions/ooc/DataGenOOCInstruction.java index c5c2e299cbb..51a0e414a7d 100644 --- a/src/main/java/org/apache/sysds/runtime/instructions/ooc/DataGenOOCInstruction.java +++ b/src/main/java/org/apache/sysds/runtime/instructions/ooc/DataGenOOCInstruction.java @@ -34,10 +34,8 @@ import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue; import org.apache.sysds.runtime.matrix.data.LibMatrixDatagen; import org.apache.sysds.runtime.matrix.data.MatrixBlock; -import org.apache.sysds.runtime.matrix.data.MatrixIndexes; import org.apache.sysds.runtime.matrix.data.RandomMatrixGenerator; import org.apache.sysds.runtime.matrix.operators.UnaryOperator; -import org.apache.sysds.runtime.ooc.stream.StreamContext; import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils; import org.apache.sysds.runtime.util.UtilFunctions; @@ -251,34 +249,22 @@ else if(method == Types.OpOpDG.SEQ) { final int maxK = (int) UtilFunctions.getSeqLength(lfrom, lto, lincr); final double finalLincr = lincr; + ec.getDataCharacteristics(output.getName()).set(maxK, 1, blen, -1); - OOCInstructionUtils.submitOOCTask(() -> { - int k = 0; - double curFrom = lfrom; - double curTo; - MatrixBlock mb; - - while (k < maxK) { - long desiredLen = Math.min(blen, maxK - k); - curTo = curFrom + (desiredLen - 1) * finalLincr; - long actualLen = UtilFunctions.getSeqLength(curFrom, curTo, finalLincr); - - if (actualLen != desiredLen) { - // Then we add / subtract a small correction term - curTo += (actualLen < desiredLen) ? finalLincr / 2 : -finalLincr / 2; - - if (UtilFunctions.getSeqLength(curFrom, curTo, finalLincr) != desiredLen) - throw new DMLRuntimeException("OOC seq could not construct the right number of elements."); - } - - mb = MatrixBlock.seqOperations(curFrom, curTo, finalLincr); - qOut.enqueue(new IndexedMatrixValue(new MatrixIndexes(1 + k / blen, 1), mb)); - curFrom = mb.get(mb.getNumRows() - 1, 0) + finalLincr; - k += blen; + OOCInstructionUtils.dataGen(qOut, idx -> { + long offset = (idx.getRowIndex() - 1) * blen; + long desiredLen = Math.min(blen, maxK - offset); + double curFrom = lfrom + offset * finalLincr; + double curTo = curFrom + (desiredLen - 1) * finalLincr; + long actualLen = UtilFunctions.getSeqLength(curFrom, curTo, finalLincr); + + if(actualLen != desiredLen) { + curTo += actualLen < desiredLen ? finalLincr / 2 : -finalLincr / 2; + if(UtilFunctions.getSeqLength(curFrom, curTo, finalLincr) != desiredLen) + throw new DMLRuntimeException("OOC seq could not construct the right number of elements."); } - - qOut.closeInput(); - }, new StreamContext(_callerId, getExtendedOpcode()).addOutStream(qOut)); + return MatrixBlock.seqOperations(curFrom, curTo, finalLincr); + }, getContext()); } else throw new NotImplementedException(); diff --git a/src/main/java/org/apache/sysds/runtime/instructions/ooc/ParameterizedBuiltinOOCInstruction.java b/src/main/java/org/apache/sysds/runtime/instructions/ooc/ParameterizedBuiltinOOCInstruction.java index 2f71d0e4538..d46c606768a 100644 --- a/src/main/java/org/apache/sysds/runtime/instructions/ooc/ParameterizedBuiltinOOCInstruction.java +++ b/src/main/java/org/apache/sysds/runtime/instructions/ooc/ParameterizedBuiltinOOCInstruction.java @@ -40,6 +40,7 @@ import org.apache.sysds.runtime.matrix.data.MatrixBlock; import org.apache.sysds.runtime.matrix.operators.Operator; import org.apache.sysds.runtime.matrix.operators.SimpleOperator; +import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils; import org.apache.sysds.runtime.util.UtilFunctions; import java.util.ArrayList; @@ -92,13 +93,13 @@ public void processInstruction(ExecutionContext ec) { throw new NotImplementedException(); } else{ MatrixObject targetObj = ec.getMatrixObject(params.get("target")); - OOCStream qIn = targetObj.getStreamHandle(); OOCStream qOut = createWritableStream(); double pattern = Double.parseDouble(params.get("pattern")); double replacement = Double.parseDouble(params.get("replacement")); - mapOOC(qIn, qOut, tmp -> new IndexedMatrixValue(tmp.getIndexes(), tmp.getValue().replaceOperations(new MatrixBlock(), pattern, replacement))); + OOCInstructionUtils.equiMapBlock(targetObj.getStreamable(), qOut, + block -> block.replaceOperations(new MatrixBlock(), pattern, replacement), getContext()); ec.getMatrixObject(output).setStreamHandle(qOut); } diff --git a/src/main/java/org/apache/sysds/runtime/instructions/ooc/TernaryOOCInstruction.java b/src/main/java/org/apache/sysds/runtime/instructions/ooc/TernaryOOCInstruction.java index 6036647cc7f..7b91b16d237 100644 --- a/src/main/java/org/apache/sysds/runtime/instructions/ooc/TernaryOOCInstruction.java +++ b/src/main/java/org/apache/sysds/runtime/instructions/ooc/TernaryOOCInstruction.java @@ -35,6 +35,7 @@ import org.apache.sysds.runtime.matrix.data.MatrixBlock; import org.apache.sysds.runtime.matrix.operators.Operator; import org.apache.sysds.runtime.matrix.operators.TernaryOperator; +import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils; public class TernaryOOCInstruction extends ComputationOOCInstruction { @@ -106,27 +107,20 @@ private void processSingleMatrixInstruction(ExecutionContext ec, int matrixPos) MatrixBlock s2 = input2.isMatrix() ? null : getScalarInputBlock(ec, input2); MatrixBlock s3 = input3.isMatrix() ? null : getScalarInputBlock(ec, input3); - OOCStream qIn = mo.getStreamHandle(); OOCStream qOut = createWritableStream(); ec.getMatrixObject(output).setStreamHandle(qOut); - mapOOC(qIn, qOut, tmp -> { - IndexedMatrixValue outVal = new IndexedMatrixValue(); - MatrixBlock op1 = resolveOperandBlock(1, tmp, null, matrixPos, -1, s1, s2, s3); - MatrixBlock op2 = resolveOperandBlock(2, tmp, null, matrixPos, -1, s1, s2, s3); - MatrixBlock op3 = resolveOperandBlock(3, tmp, null, matrixPos, -1, s1, s2, s3); - outVal.set(tmp.getIndexes(), - op1.ternaryOperations((TernaryOperator)_optr, op2, op3, new MatrixBlock())); - return outVal; - }); + OOCInstructionUtils.equiMapBlock(mo.getStreamable(), qOut, block -> { + MatrixBlock op1 = resolveOperandBlock(1, block, null, matrixPos, -1, s1, s2, s3); + MatrixBlock op2 = resolveOperandBlock(2, block, null, matrixPos, -1, s1, s2, s3); + MatrixBlock op3 = resolveOperandBlock(3, block, null, matrixPos, -1, s1, s2, s3); + return op1.ternaryOperations((TernaryOperator) _optr, op2, op3, new MatrixBlock()); + }, getContext()); } private void processTwoMatrixInstruction(ExecutionContext ec, int leftPos, int rightPos) { MatrixObject left = getMatrixObject(ec, leftPos); MatrixObject right = getMatrixObject(ec, rightPos); - OOCStream leftStream = left.getStreamHandle(); - OOCStream rightStream = right.getStreamHandle(); - MatrixBlock s1 = input1.isMatrix() ? null : getScalarInputBlock(ec, input1); MatrixBlock s2 = input2.isMatrix() ? null : getScalarInputBlock(ec, input2); MatrixBlock s3 = input3.isMatrix() ? null : getScalarInputBlock(ec, input3); @@ -134,15 +128,26 @@ private void processTwoMatrixInstruction(ExecutionContext ec, int leftPos, int r OOCStream qOut = createWritableStream(); ec.getMatrixObject(output).setStreamHandle(qOut); - joinOOC(leftStream, rightStream, qOut, (l, r) -> { - IndexedMatrixValue outVal = new IndexedMatrixValue(); - MatrixBlock op1 = resolveOperandBlock(1, l, r, leftPos, rightPos, s1, s2, s3); - MatrixBlock op2 = resolveOperandBlock(2, l, r, leftPos, rightPos, s1, s2, s3); - MatrixBlock op3 = resolveOperandBlock(3, l, r, leftPos, rightPos, s1, s2, s3); - outVal.set(l.getIndexes(), - op1.ternaryOperations((TernaryOperator)_optr, op2, op3, new MatrixBlock())); - return outVal; - }, IndexedMatrixValue::getIndexes); + if(left.getDataCharacteristics().dimsKnown() && right.getDataCharacteristics().dimsKnown()) { + OOCInstructionUtils.equiJoin(left.getStreamable(), right.getStreamable(), qOut, (l, r) -> { + MatrixBlock op1 = resolveOperandBlock(1, l, r, leftPos, rightPos, s1, s2, s3); + MatrixBlock op2 = resolveOperandBlock(2, l, r, leftPos, rightPos, s1, s2, s3); + MatrixBlock op3 = resolveOperandBlock(3, l, r, leftPos, rightPos, s1, s2, s3); + return op1.ternaryOperations((TernaryOperator) _optr, op2, op3, new MatrixBlock()); + }, getContext()); + } + else { + OOCStream leftStream = left.getStreamHandle(); + OOCStream rightStream = right.getStreamHandle(); + joinOOC(leftStream, rightStream, qOut, (l, r) -> { + IndexedMatrixValue outVal = new IndexedMatrixValue(); + MatrixBlock op1 = resolveOperandBlock(1, l, r, leftPos, rightPos, s1, s2, s3); + MatrixBlock op2 = resolveOperandBlock(2, l, r, leftPos, rightPos, s1, s2, s3); + MatrixBlock op3 = resolveOperandBlock(3, l, r, leftPos, rightPos, s1, s2, s3); + outVal.set(l.getIndexes(), op1.ternaryOperations((TernaryOperator) _optr, op2, op3, new MatrixBlock())); + return outVal; + }, IndexedMatrixValue::getIndexes); + } } private void processThreeMatrixInstruction(ExecutionContext ec) { @@ -185,10 +190,16 @@ private MatrixBlock getScalarInputBlock(ExecutionContext ec, CPOperand operand) private MatrixBlock resolveOperandBlock(int operandPos, IndexedMatrixValue left, IndexedMatrixValue right, int leftPos, int rightPos, MatrixBlock s1, MatrixBlock s2, MatrixBlock s3) { + return resolveOperandBlock(operandPos, left == null ? null : (MatrixBlock) left.getValue(), + right == null ? null : (MatrixBlock) right.getValue(), leftPos, rightPos, s1, s2, s3); + } + + private MatrixBlock resolveOperandBlock(int operandPos, MatrixBlock left, MatrixBlock right, int leftPos, + int rightPos, MatrixBlock s1, MatrixBlock s2, MatrixBlock s3) { if(operandPos == leftPos && left != null) - return (MatrixBlock) left.getValue(); + return left; if(operandPos == rightPos && right != null) - return (MatrixBlock) right.getValue(); + return right; if(operandPos == 1) return s1; diff --git a/src/test/java/org/apache/sysds/test/functions/ooc/SeqTest.java b/src/test/java/org/apache/sysds/test/functions/ooc/SeqTest.java index f7855b93e2d..47761e85c30 100644 --- a/src/test/java/org/apache/sysds/test/functions/ooc/SeqTest.java +++ b/src/test/java/org/apache/sysds/test/functions/ooc/SeqTest.java @@ -53,6 +53,11 @@ public void testSeq2() { runSeqTest(0, 15.9, 0.01); } + @Test + public void testDescendingSeq() { + runSeqTest(10, 0, -0.1); + } + private void runSeqTest(double from, double to, double incr) { Types.ExecMode platformOld = setExecMode(Types.ExecMode.SINGLE_NODE);