Baunsgaard commented on code in PR #2619:
URL: https://github.com/apache/systemds/pull/2619#discussion_r4066970320


##########
src/main/java/org/apache/sysds/runtime/ooc/primitives/ReshapeOOCPrimitive.java:
##########
@@ -0,0 +1,796 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.sysds.runtime.ooc.primitives;
+
+import org.apache.sysds.runtime.DMLRuntimeException;
+import org.apache.sysds.runtime.data.DenseBlockFP64;
+import org.apache.sysds.runtime.instructions.ooc.CachingStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStreamable;
+import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue;
+import org.apache.sysds.runtime.matrix.data.MatrixBlock;
+import org.apache.sysds.runtime.matrix.data.MatrixIndexes;
+import org.apache.sysds.runtime.meta.DataCharacteristics;
+import org.apache.sysds.runtime.ooc.cache.OOCCacheManager;
+import org.apache.sysds.runtime.ooc.cache.OOCFuture;
+import org.apache.sysds.runtime.ooc.memory.ManagedPayload;
+import org.apache.sysds.runtime.ooc.memory.ReservationBudget;
+import org.apache.sysds.runtime.ooc.planning.OOCAccessPattern;
+import org.apache.sysds.runtime.ooc.store.StateTable;
+import org.apache.sysds.runtime.ooc.store.StoreLease;
+import org.apache.sysds.runtime.ooc.stream.AllocatedOOCStream;
+import org.apache.sysds.runtime.ooc.stream.StreamContext;
+import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils;
+import org.apache.sysds.runtime.ooc.util.OOCUtils;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+
+public class ReshapeOOCPrimitive extends OOCPrimitive {
+       private final OOCStreamable<IndexedMatrixValue> _input;
+       private final OOCStreamable<IndexedMatrixValue> _output;
+       private final boolean _byRow;
+       private final long _rows;
+       private final long _cols;
+       private long _rlen;
+       private long _clen;
+       private int _blen;
+
+       public ReshapeOOCPrimitive(OOCStreamable<IndexedMatrixValue> input, 
OOCStreamable<IndexedMatrixValue> output,
+               long rows, long cols, boolean byRow, StreamContext context) {
+               super(context, input);
+               _input = input;
+               _output = output;
+               _byRow = byRow;
+               _rows = rows;
+               _cols = cols;
+               _pattern = byRow ? OOCAccessPattern.ROW_MAJOR : 
OOCAccessPattern.COL_MAJOR;
+       }
+
+       @Override
+       protected void inferPatternsInternal() {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+               inferParentPatterns();
+       }
+
+       @Override
+       protected void requestPatternInternal(OOCAccessPattern accessPattern) {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+       }
+
+       @Override
+       protected void startExecution() {
+               DataCharacteristics inputDc = _input.getDataCharacteristics();
+               if(inputDc == null || !inputDc.dimsKnown() || 
inputDc.getBlocksize() <= 0)
+                       throw new DMLRuntimeException("Reshape OOC reduction 
requires known input dimensions and block size.");
+
+               OOCStream<IndexedMatrixValue> input = getInputReadStream(0);
+               OOCStream<IndexedMatrixValue> output = _output.getWriteStream();
+               getContext().addOutStream(output);
+
+               _rlen = inputDc.getRows();
+               _clen = inputDc.getCols();
+               _blen = Math.toIntExact(inputDc.getBlocksize());
+
+               if(_rlen * _clen != _rows * _cols) {
+                       onComplete();
+                       throw new DMLRuntimeException("Reshape matrix requires 
consistent numbers of input/output cells (" + _rlen
+                               + ":" + _clen + ", " + _rows + ":" + _cols + 
").");
+               }
+
+               if(_rlen == _rows) {
+                       OOCInstructionUtils
+                               .submitAdmittedOOCTasks(input, output,
+                                       value -> new 
IndexedMatrixValue(value.getIndexes(), value.getValue()), _allowance, 
getContext())
+                               .thenRun(this::onComplete);
+                       return;

Review Comment:
   Add a comment per special case to highlight what it is doing. This one is 
for a single block  right?
   



##########
src/main/java/org/apache/sysds/runtime/ooc/primitives/ReshapeOOCPrimitive.java:
##########
@@ -0,0 +1,796 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.sysds.runtime.ooc.primitives;
+
+import org.apache.sysds.runtime.DMLRuntimeException;
+import org.apache.sysds.runtime.data.DenseBlockFP64;
+import org.apache.sysds.runtime.instructions.ooc.CachingStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStreamable;
+import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue;
+import org.apache.sysds.runtime.matrix.data.MatrixBlock;
+import org.apache.sysds.runtime.matrix.data.MatrixIndexes;
+import org.apache.sysds.runtime.meta.DataCharacteristics;
+import org.apache.sysds.runtime.ooc.cache.OOCCacheManager;
+import org.apache.sysds.runtime.ooc.cache.OOCFuture;
+import org.apache.sysds.runtime.ooc.memory.ManagedPayload;
+import org.apache.sysds.runtime.ooc.memory.ReservationBudget;
+import org.apache.sysds.runtime.ooc.planning.OOCAccessPattern;
+import org.apache.sysds.runtime.ooc.store.StateTable;
+import org.apache.sysds.runtime.ooc.store.StoreLease;
+import org.apache.sysds.runtime.ooc.stream.AllocatedOOCStream;
+import org.apache.sysds.runtime.ooc.stream.StreamContext;
+import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils;
+import org.apache.sysds.runtime.ooc.util.OOCUtils;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+
+public class ReshapeOOCPrimitive extends OOCPrimitive {
+       private final OOCStreamable<IndexedMatrixValue> _input;
+       private final OOCStreamable<IndexedMatrixValue> _output;
+       private final boolean _byRow;
+       private final long _rows;
+       private final long _cols;
+       private long _rlen;
+       private long _clen;
+       private int _blen;

Review Comment:
   Consider if you can move more variables out as class variables, This might 
make it cleaner in the end.
   
   



##########
src/main/java/org/apache/sysds/runtime/ooc/primitives/ReshapeOOCPrimitive.java:
##########
@@ -0,0 +1,796 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.sysds.runtime.ooc.primitives;
+
+import org.apache.sysds.runtime.DMLRuntimeException;
+import org.apache.sysds.runtime.data.DenseBlockFP64;
+import org.apache.sysds.runtime.instructions.ooc.CachingStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStreamable;
+import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue;
+import org.apache.sysds.runtime.matrix.data.MatrixBlock;
+import org.apache.sysds.runtime.matrix.data.MatrixIndexes;
+import org.apache.sysds.runtime.meta.DataCharacteristics;
+import org.apache.sysds.runtime.ooc.cache.OOCCacheManager;
+import org.apache.sysds.runtime.ooc.cache.OOCFuture;
+import org.apache.sysds.runtime.ooc.memory.ManagedPayload;
+import org.apache.sysds.runtime.ooc.memory.ReservationBudget;
+import org.apache.sysds.runtime.ooc.planning.OOCAccessPattern;
+import org.apache.sysds.runtime.ooc.store.StateTable;
+import org.apache.sysds.runtime.ooc.store.StoreLease;
+import org.apache.sysds.runtime.ooc.stream.AllocatedOOCStream;
+import org.apache.sysds.runtime.ooc.stream.StreamContext;
+import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils;
+import org.apache.sysds.runtime.ooc.util.OOCUtils;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+
+public class ReshapeOOCPrimitive extends OOCPrimitive {
+       private final OOCStreamable<IndexedMatrixValue> _input;
+       private final OOCStreamable<IndexedMatrixValue> _output;
+       private final boolean _byRow;
+       private final long _rows;
+       private final long _cols;
+       private long _rlen;
+       private long _clen;
+       private int _blen;
+
+       public ReshapeOOCPrimitive(OOCStreamable<IndexedMatrixValue> input, 
OOCStreamable<IndexedMatrixValue> output,
+               long rows, long cols, boolean byRow, StreamContext context) {
+               super(context, input);
+               _input = input;
+               _output = output;
+               _byRow = byRow;
+               _rows = rows;
+               _cols = cols;
+               _pattern = byRow ? OOCAccessPattern.ROW_MAJOR : 
OOCAccessPattern.COL_MAJOR;
+       }
+
+       @Override
+       protected void inferPatternsInternal() {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+               inferParentPatterns();
+       }
+
+       @Override
+       protected void requestPatternInternal(OOCAccessPattern accessPattern) {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+       }
+
+       @Override
+       protected void startExecution() {
+               DataCharacteristics inputDc = _input.getDataCharacteristics();
+               if(inputDc == null || !inputDc.dimsKnown() || 
inputDc.getBlocksize() <= 0)
+                       throw new DMLRuntimeException("Reshape OOC reduction 
requires known input dimensions and block size.");
+
+               OOCStream<IndexedMatrixValue> input = getInputReadStream(0);
+               OOCStream<IndexedMatrixValue> output = _output.getWriteStream();
+               getContext().addOutStream(output);
+
+               _rlen = inputDc.getRows();
+               _clen = inputDc.getCols();
+               _blen = Math.toIntExact(inputDc.getBlocksize());
+
+               if(_rlen * _clen != _rows * _cols) {
+                       onComplete();
+                       throw new DMLRuntimeException("Reshape matrix requires 
consistent numbers of input/output cells (" + _rlen
+                               + ":" + _clen + ", " + _rows + ":" + _cols + 
").");
+               }
+
+               if(_rlen == _rows) {
+                       OOCInstructionUtils
+                               .submitAdmittedOOCTasks(input, output,
+                                       value -> new 
IndexedMatrixValue(value.getIndexes(), value.getValue()), _allowance, 
getContext())
+                               .thenRun(this::onComplete);
+                       return;
+               }
+
+               if(_clen <= _blen && _rlen <= _blen && _cols <= _blen && _rows 
<= _blen) {
+                       OOCInstructionUtils.submitAdmittedOOCTasks(input, 
output,
+                               value -> new 
IndexedMatrixValue(value.getIndexes(),
+                                       ((MatrixBlock) 
value.getValue()).reshape((int) _rows, (int) _cols, _byRow)),
+                               _allowance, 
getContext()).thenRun(this::onComplete);
+                       return;

Review Comment:
   add comment for special case.



##########
src/main/java/org/apache/sysds/runtime/ooc/cache/packed/OOCPackedCache.java:
##########
@@ -47,7 +47,7 @@
 
 public final class OOCPackedCache implements OOCCache {
        private static final long PACKED_STREAM_ID = 
CachingStream._streamSeq.getNextID();
-       private static final long DEFAULT_PACK_THRESHOLD_BYTES = 1L << 18;
+       private static final long DEFAULT_PACK_THRESHOLD_BYTES = 1;

Review Comment:
   maybe add a doc string?



##########
src/main/java/org/apache/sysds/runtime/ooc/primitives/ReshapeOOCPrimitive.java:
##########
@@ -0,0 +1,796 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.sysds.runtime.ooc.primitives;
+
+import org.apache.sysds.runtime.DMLRuntimeException;
+import org.apache.sysds.runtime.data.DenseBlockFP64;
+import org.apache.sysds.runtime.instructions.ooc.CachingStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStreamable;
+import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue;
+import org.apache.sysds.runtime.matrix.data.MatrixBlock;
+import org.apache.sysds.runtime.matrix.data.MatrixIndexes;
+import org.apache.sysds.runtime.meta.DataCharacteristics;
+import org.apache.sysds.runtime.ooc.cache.OOCCacheManager;
+import org.apache.sysds.runtime.ooc.cache.OOCFuture;
+import org.apache.sysds.runtime.ooc.memory.ManagedPayload;
+import org.apache.sysds.runtime.ooc.memory.ReservationBudget;
+import org.apache.sysds.runtime.ooc.planning.OOCAccessPattern;
+import org.apache.sysds.runtime.ooc.store.StateTable;
+import org.apache.sysds.runtime.ooc.store.StoreLease;
+import org.apache.sysds.runtime.ooc.stream.AllocatedOOCStream;
+import org.apache.sysds.runtime.ooc.stream.StreamContext;
+import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils;
+import org.apache.sysds.runtime.ooc.util.OOCUtils;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+
+public class ReshapeOOCPrimitive extends OOCPrimitive {
+       private final OOCStreamable<IndexedMatrixValue> _input;
+       private final OOCStreamable<IndexedMatrixValue> _output;
+       private final boolean _byRow;
+       private final long _rows;
+       private final long _cols;
+       private long _rlen;
+       private long _clen;
+       private int _blen;
+
+       public ReshapeOOCPrimitive(OOCStreamable<IndexedMatrixValue> input, 
OOCStreamable<IndexedMatrixValue> output,
+               long rows, long cols, boolean byRow, StreamContext context) {
+               super(context, input);
+               _input = input;
+               _output = output;
+               _byRow = byRow;
+               _rows = rows;
+               _cols = cols;
+               _pattern = byRow ? OOCAccessPattern.ROW_MAJOR : 
OOCAccessPattern.COL_MAJOR;
+       }
+
+       @Override
+       protected void inferPatternsInternal() {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+               inferParentPatterns();
+       }
+
+       @Override
+       protected void requestPatternInternal(OOCAccessPattern accessPattern) {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+       }
+
+       @Override
+       protected void startExecution() {
+               DataCharacteristics inputDc = _input.getDataCharacteristics();
+               if(inputDc == null || !inputDc.dimsKnown() || 
inputDc.getBlocksize() <= 0)
+                       throw new DMLRuntimeException("Reshape OOC reduction 
requires known input dimensions and block size.");
+
+               OOCStream<IndexedMatrixValue> input = getInputReadStream(0);
+               OOCStream<IndexedMatrixValue> output = _output.getWriteStream();
+               getContext().addOutStream(output);
+
+               _rlen = inputDc.getRows();
+               _clen = inputDc.getCols();
+               _blen = Math.toIntExact(inputDc.getBlocksize());
+
+               if(_rlen * _clen != _rows * _cols) {
+                       onComplete();
+                       throw new DMLRuntimeException("Reshape matrix requires 
consistent numbers of input/output cells (" + _rlen
+                               + ":" + _clen + ", " + _rows + ":" + _cols + 
").");
+               }
+
+               if(_rlen == _rows) {
+                       OOCInstructionUtils
+                               .submitAdmittedOOCTasks(input, output,
+                                       value -> new 
IndexedMatrixValue(value.getIndexes(), value.getValue()), _allowance, 
getContext())
+                               .thenRun(this::onComplete);
+                       return;
+               }
+
+               if(_clen <= _blen && _rlen <= _blen && _cols <= _blen && _rows 
<= _blen) {
+                       OOCInstructionUtils.submitAdmittedOOCTasks(input, 
output,
+                               value -> new 
IndexedMatrixValue(value.getIndexes(),
+                                       ((MatrixBlock) 
value.getValue()).reshape((int) _rows, (int) _cols, _byRow)),
+                               _allowance, 
getContext()).thenRun(this::onComplete);
+                       return;
+               }
+
+               int numColBlocksIn = Math.toIntExact(inputDc.getNumColBlocks());
+               int numRowBlocksIn = Math.toIntExact(inputDc.getNumRowBlocks());
+               int numColBlocksOut = (int) Math.ceil((double) _cols / _blen);
+               int numRowBlocksOut = (int) Math.ceil((double) _rows / _blen);
+
+               long sliceBytes = 
(OOCUtils.estimateFullTileBytes(input.getDataCharacteristics()) +
+                       (_blen - 1) * MatrixBlock.getHeaderSize()) / _blen;
+
+               StateTable<IndexedMatrixValue> table = new 
StateTable<>(OOCCacheManager.getGlobalCache(),
+                       CachingStream._streamSeq.getNextID());
+
+               if(_byRow) {
+                       if(_clen % _blen == 0 && _cols % _blen == 0) {
+                               // singleRowBlocks do not need to be split
+                               if(_rows == 1) {
+                                       // result is one single row
+                                       submitSingleRowColTask(input, output, 
numColBlocksIn, numRowBlocksIn);
+                               }
+                               else {
+                                       CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                                       f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(
+                                               () -> 
reshapeFullColBlocks(table, output, numColBlocksOut, numRowBlocksOut, 
sliceBytes),
+                                               getContext()));
+                               }
+                       }
+                       else {
+                               CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                               f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(() -> reshapePartialColBlocks(table, output,
+                                       numColBlocksIn, numColBlocksOut, 
numRowBlocksOut, sliceBytes), getContext()));
+                       }
+               }
+               else {
+                       if(_rlen % _blen == 0 && _rows % _blen == 0) {
+                               // singleColBlocks do not need to be split
+                               if(_cols == 1) {
+                                       // result is one single col
+                                       submitSingleRowColTask(input, output, 
numColBlocksIn, numRowBlocksIn);
+                               }
+                               else {
+                                       CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                                       f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(
+                                               () -> 
reshapeFullRowBlocks(table, output, numColBlocksOut, numRowBlocksOut, 
sliceBytes),
+                                               getContext()));
+                               }
+                       }
+                       else {
+                               CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                               f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(() -> reshapePartialRowBlocks(table, output,
+                                       numRowBlocksIn, numColBlocksOut, 
numRowBlocksOut, sliceBytes), getContext()));
+                       }
+               }
+       }
+
+       private CompletableFuture<Void> 
splitIntoTable(OOCStream<IndexedMatrixValue> in,
+               StateTable<IndexedMatrixValue> table, int numColBlocks, int 
numRowBlocks) {
+
+               long blockBytes = 
OOCUtils.estimateFullTileBytes(in.getDataCharacteristics()) +
+                       (_blen - 1) * MatrixBlock.getHeaderSize();
+               long singleSliceBytes = blockBytes / _blen;
+
+               AllocatedOOCStream<IndexedMatrixValue> allocated = new 
AllocatedOOCStream<>(in, _allowance, ignored -> blockBytes);
+
+               return OOCInstructionUtils.submitOOCTasks(allocated, callback 
-> {
+                       try(ReservationBudget budget = 
AllocatedOOCStream.detachBudget(callback)) {
+                               if(budget == null)
+                                       throw new DMLRuntimeException("Missing 
admitted output budget");
+
+                               IndexedMatrixValue imv = callback.get();
+                               MatrixBlock blk = (MatrixBlock) imv.getValue();
+                               long r = imv.getIndexes().getRowIndex();
+                               long c = imv.getIndexes().getColumnIndex();
+                               long rIdx;
+                               long cIdx;
+
+                               int n = _byRow ? blk.getNumRows() : 
blk.getNumColumns();
+                               for(int i = 0; i < n; i++) {
+                                       MatrixBlock slice;
+                                       if(_byRow) {
+                                               slice = blk.slice(i, i);
+                                               rIdx = (r - 1) * _blen + i + 1;
+                                               cIdx = c;
+                                       }
+                                       else {
+                                               slice = blk.slice(0, 
blk.getNumRows() - 1, i, i);
+                                               cIdx = (c - 1) * _blen + i + 1;
+                                               rIdx = r;
+                                       }
+
+                                       long targetIdx = _byRow ? (rIdx - 1) * 
numColBlocks + c - 1 : (cIdx - 1) * numRowBlocks + r - 1;
+                                       IndexedMatrixValue sliceImv = new 
IndexedMatrixValue(new MatrixIndexes(rIdx, cIdx), slice);
+                                       
budget.reserveBlocking(singleSliceBytes);
+                                       table.put((int) targetIdx, new 
ManagedPayload<>(sliceImv, singleSliceBytes, budget));
+                               }
+                       }
+                       catch(IllegalStateException e) {
+                               throw new DMLRuntimeException(e);
+                       }
+               }, getContext());
+       }
+
+       private void submitSingleRowColTask(OOCStream<IndexedMatrixValue> in, 
OOCStream<IndexedMatrixValue> out,
+               int numColBlocks, int numRowBlocks) {
+
+               // one input block is split into blen output blocks
+               long outputBytes = _blen * 
OOCUtils.estimateOutputTileBytes(out.getDataCharacteristics());
+
+               AllocatedOOCStream<IndexedMatrixValue> allocated = new 
AllocatedOOCStream<>(in, _allowance, ignored -> outputBytes);
+
+               OOCInstructionUtils.submitOOCTasks(allocated, callback -> {
+                       try(ReservationBudget budget = 
AllocatedOOCStream.detachBudget(callback)) {
+                               if(budget == null)
+                                       throw new DMLRuntimeException("Missing 
admitted output budget");
+
+                               IndexedMatrixValue imv = callback.get();
+                               MatrixBlock blk = (MatrixBlock) imv.getValue();
+                               long r = imv.getIndexes().getRowIndex();
+                               long c = imv.getIndexes().getColumnIndex();
+                               long rIdx;
+                               long cIdx;
+
+                               int n = _byRow ? blk.getNumRows() : 
blk.getNumColumns();
+                               for(int i = 0; i < n; i++) {
+                                       MatrixBlock slice;
+                                       if(_byRow) {
+                                               // split and adjust idx
+                                               slice = blk.slice(i, i);
+                                               // total row, 1 based
+                                               rIdx = (r - 1) * _blen + i + 1;
+                                               // all in single row
+                                               cIdx = (rIdx - 1) * 
numColBlocks + c;
+                                               rIdx = 1;
+                                       }
+                                       else {
+                                               // split and adjust idx
+                                               slice = blk.slice(0, 
blk.getNumRows() - 1, i, i);
+                                               cIdx = (c - 1) * _blen + i + 1;
+                                               // all in single col
+                                               rIdx = (cIdx - 1) * 
numRowBlocks + r;
+                                               cIdx = 1;
+                                       }
+
+                                       IndexedMatrixValue sliceImv = new 
IndexedMatrixValue(new MatrixIndexes(rIdx, cIdx), slice);
+                                       OOCUtils.enqueueExact(out, sliceImv, 
budget, false);
+                               }
+                       }
+                       catch(IllegalStateException e) {
+                               throw new DMLRuntimeException(e);
+                       }
+               }, 
getContext()).thenRun(this::onComplete).thenRun(out::closeInput).exceptionally(error
 -> {
+                       out.propagateFailure(DMLRuntimeException.of(error));
+                       return null;
+               });
+       }
+
+       private void reshapeFullColBlocks(StateTable<IndexedMatrixValue> table, 
OOCStream<IndexedMatrixValue> out,
+               int numColBlocksOut, int numRowBlocksOut, long rowBytes) {
+
+               List<OOCFuture<StoreLease<IndexedMatrixValue>>> futures = new 
ArrayList<>();
+               ReservationBudget budget = null;
+
+               long numRowsBlockOut = 
Math.min(out.getDataCharacteristics().getRows(), _blen);
+               long blockBytes = 
OOCUtils.estimateOutputTileBytes(out.getDataCharacteristics());
+               long outputBytes = numRowsBlockOut * rowBytes + blockBytes;
+
+               // totalRowIdx corresponds to index of row block when all 
aligned in one row
+               // br * numColBlocksOut * blen + b + r * numColBlocksOut;
+               // with numColBlocksOut * blen = cols
+               long totalIdx = -_cols - 1 - numColBlocksOut;
+
+               try {
+                       // iterate through rows of output blocks
+                       for(int br = 0; br < numRowBlocksOut; br++) {
+                               totalIdx += _cols;
+                               long tmp = totalIdx;
+                               int localRows = (br == numRowBlocksOut - 1 && 
_rows % _blen != 0) ? (int) _rows % _blen : _blen;
+                               // for each block in row
+                               for(int b = 0; b < numColBlocksOut; b++) {
+                                       totalIdx += 1;
+                                       long tmp2 = totalIdx;
+                                       budget = 
OOCUtils.reserveBudget(_allowance, outputBytes);
+                                       // for each row in block
+                                       for(int r = 0; r < _blen && r < 
localRows; r++) {
+                                               totalIdx += numColBlocksOut;
+                                               
OOCFuture<StoreLease<IndexedMatrixValue>> rowFuture = table.take((int) 
totalIdx, budget);
+                                               futures.add(rowFuture);
+                                       }
+                                       totalIdx = tmp2;
+                                       
OOCFuture<List<StoreLease<IndexedMatrixValue>>> future = 
OOCFuture.allOf(futures, StoreLease::close);
+                                       MatrixIndexes idx = new 
MatrixIndexes(br + 1, b + 1);
+
+                                       ReservationBudget finalBudget = budget;
+                                       future.whenComplete((leases, error) -> {
+                                               MatrixBlock block = new 
MatrixBlock(localRows, _blen, false);
+                                               for(int r = 0; r < 
leases.size(); r++) {
+                                                       
StoreLease<IndexedMatrixValue> lease = leases.get(r);
+                                                       MatrixBlock row = 
(MatrixBlock) lease.value().getValue();
+                                                       block.setRow(r, 
row.getDenseBlockValues());
+                                                       lease.close();
+                                               }
+                                               block.recomputeNonZeros();
+                                               OOCUtils.enqueueExact(out, new 
IndexedMatrixValue(idx, block), finalBudget, true);
+                                               futures.clear();
+                                       });
+                                       budget = null;
+                               }
+                               totalIdx = tmp;
+                       }
+               }
+               catch(IllegalStateException e) {
+                       throw new DMLRuntimeException(e);
+               }
+               finally {
+                       if(budget != null) {
+                               budget.close();
+                       }
+                       try {
+                               table.close();
+                               onComplete();
+                       }
+                       finally {
+                               out.closeInput();
+                       }
+               }

Review Comment:
   This finally block is duplicated across 3-4 places. Maybe make a helper 
function.



##########
src/main/java/org/apache/sysds/runtime/ooc/primitives/ReshapeOOCPrimitive.java:
##########
@@ -0,0 +1,796 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.sysds.runtime.ooc.primitives;
+
+import org.apache.sysds.runtime.DMLRuntimeException;
+import org.apache.sysds.runtime.data.DenseBlockFP64;
+import org.apache.sysds.runtime.instructions.ooc.CachingStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStream;
+import org.apache.sysds.runtime.instructions.ooc.OOCStreamable;
+import org.apache.sysds.runtime.instructions.spark.data.IndexedMatrixValue;
+import org.apache.sysds.runtime.matrix.data.MatrixBlock;
+import org.apache.sysds.runtime.matrix.data.MatrixIndexes;
+import org.apache.sysds.runtime.meta.DataCharacteristics;
+import org.apache.sysds.runtime.ooc.cache.OOCCacheManager;
+import org.apache.sysds.runtime.ooc.cache.OOCFuture;
+import org.apache.sysds.runtime.ooc.memory.ManagedPayload;
+import org.apache.sysds.runtime.ooc.memory.ReservationBudget;
+import org.apache.sysds.runtime.ooc.planning.OOCAccessPattern;
+import org.apache.sysds.runtime.ooc.store.StateTable;
+import org.apache.sysds.runtime.ooc.store.StoreLease;
+import org.apache.sysds.runtime.ooc.stream.AllocatedOOCStream;
+import org.apache.sysds.runtime.ooc.stream.StreamContext;
+import org.apache.sysds.runtime.ooc.util.OOCInstructionUtils;
+import org.apache.sysds.runtime.ooc.util.OOCUtils;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+
+public class ReshapeOOCPrimitive extends OOCPrimitive {
+       private final OOCStreamable<IndexedMatrixValue> _input;
+       private final OOCStreamable<IndexedMatrixValue> _output;
+       private final boolean _byRow;
+       private final long _rows;
+       private final long _cols;
+       private long _rlen;
+       private long _clen;
+       private int _blen;
+
+       public ReshapeOOCPrimitive(OOCStreamable<IndexedMatrixValue> input, 
OOCStreamable<IndexedMatrixValue> output,
+               long rows, long cols, boolean byRow, StreamContext context) {
+               super(context, input);
+               _input = input;
+               _output = output;
+               _byRow = byRow;
+               _rows = rows;
+               _cols = cols;
+               _pattern = byRow ? OOCAccessPattern.ROW_MAJOR : 
OOCAccessPattern.COL_MAJOR;
+       }
+
+       @Override
+       protected void inferPatternsInternal() {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+               inferParentPatterns();
+       }
+
+       @Override
+       protected void requestPatternInternal(OOCAccessPattern accessPattern) {
+               for(OOCPrimitive child : getChildren())
+                       child.requestPattern(_pattern);
+       }
+
+       @Override
+       protected void startExecution() {
+               DataCharacteristics inputDc = _input.getDataCharacteristics();
+               if(inputDc == null || !inputDc.dimsKnown() || 
inputDc.getBlocksize() <= 0)
+                       throw new DMLRuntimeException("Reshape OOC reduction 
requires known input dimensions and block size.");
+
+               OOCStream<IndexedMatrixValue> input = getInputReadStream(0);
+               OOCStream<IndexedMatrixValue> output = _output.getWriteStream();
+               getContext().addOutStream(output);
+
+               _rlen = inputDc.getRows();
+               _clen = inputDc.getCols();
+               _blen = Math.toIntExact(inputDc.getBlocksize());
+
+               if(_rlen * _clen != _rows * _cols) {
+                       onComplete();
+                       throw new DMLRuntimeException("Reshape matrix requires 
consistent numbers of input/output cells (" + _rlen
+                               + ":" + _clen + ", " + _rows + ":" + _cols + 
").");
+               }
+
+               if(_rlen == _rows) {
+                       OOCInstructionUtils
+                               .submitAdmittedOOCTasks(input, output,
+                                       value -> new 
IndexedMatrixValue(value.getIndexes(), value.getValue()), _allowance, 
getContext())
+                               .thenRun(this::onComplete);
+                       return;
+               }
+
+               if(_clen <= _blen && _rlen <= _blen && _cols <= _blen && _rows 
<= _blen) {
+                       OOCInstructionUtils.submitAdmittedOOCTasks(input, 
output,
+                               value -> new 
IndexedMatrixValue(value.getIndexes(),
+                                       ((MatrixBlock) 
value.getValue()).reshape((int) _rows, (int) _cols, _byRow)),
+                               _allowance, 
getContext()).thenRun(this::onComplete);
+                       return;
+               }
+
+               int numColBlocksIn = Math.toIntExact(inputDc.getNumColBlocks());
+               int numRowBlocksIn = Math.toIntExact(inputDc.getNumRowBlocks());
+               int numColBlocksOut = (int) Math.ceil((double) _cols / _blen);
+               int numRowBlocksOut = (int) Math.ceil((double) _rows / _blen);
+
+               long sliceBytes = 
(OOCUtils.estimateFullTileBytes(input.getDataCharacteristics()) +
+                       (_blen - 1) * MatrixBlock.getHeaderSize()) / _blen;
+
+               StateTable<IndexedMatrixValue> table = new 
StateTable<>(OOCCacheManager.getGlobalCache(),
+                       CachingStream._streamSeq.getNextID());
+
+               if(_byRow) {
+                       if(_clen % _blen == 0 && _cols % _blen == 0) {
+                               // singleRowBlocks do not need to be split
+                               if(_rows == 1) {
+                                       // result is one single row
+                                       submitSingleRowColTask(input, output, 
numColBlocksIn, numRowBlocksIn);
+                               }
+                               else {
+                                       CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                                       f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(
+                                               () -> 
reshapeFullColBlocks(table, output, numColBlocksOut, numRowBlocksOut, 
sliceBytes),
+                                               getContext()));
+                               }
+                       }
+                       else {
+                               CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                               f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(() -> reshapePartialColBlocks(table, output,
+                                       numColBlocksIn, numColBlocksOut, 
numRowBlocksOut, sliceBytes), getContext()));
+                       }
+               }
+               else {
+                       if(_rlen % _blen == 0 && _rows % _blen == 0) {
+                               // singleColBlocks do not need to be split
+                               if(_cols == 1) {
+                                       // result is one single col
+                                       submitSingleRowColTask(input, output, 
numColBlocksIn, numRowBlocksIn);
+                               }
+                               else {
+                                       CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                                       f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(
+                                               () -> 
reshapeFullRowBlocks(table, output, numColBlocksOut, numRowBlocksOut, 
sliceBytes),
+                                               getContext()));
+                               }
+                       }
+                       else {
+                               CompletableFuture<Void> f = 
splitIntoTable(input, table, numColBlocksIn, numRowBlocksIn);
+                               f.thenRun(() -> 
OOCInstructionUtils.submitOOCTask(() -> reshapePartialRowBlocks(table, output,
+                                       numRowBlocksIn, numColBlocksOut, 
numRowBlocksOut, sliceBytes), getContext()));
+                       }
+               }
+       }
+
+       private CompletableFuture<Void> 
splitIntoTable(OOCStream<IndexedMatrixValue> in,
+               StateTable<IndexedMatrixValue> table, int numColBlocks, int 
numRowBlocks) {
+
+               long blockBytes = 
OOCUtils.estimateFullTileBytes(in.getDataCharacteristics()) +
+                       (_blen - 1) * MatrixBlock.getHeaderSize();
+               long singleSliceBytes = blockBytes / _blen;
+
+               AllocatedOOCStream<IndexedMatrixValue> allocated = new 
AllocatedOOCStream<>(in, _allowance, ignored -> blockBytes);
+
+               return OOCInstructionUtils.submitOOCTasks(allocated, callback 
-> {
+                       try(ReservationBudget budget = 
AllocatedOOCStream.detachBudget(callback)) {
+                               if(budget == null)
+                                       throw new DMLRuntimeException("Missing 
admitted output budget");
+
+                               IndexedMatrixValue imv = callback.get();
+                               MatrixBlock blk = (MatrixBlock) imv.getValue();
+                               long r = imv.getIndexes().getRowIndex();
+                               long c = imv.getIndexes().getColumnIndex();
+                               long rIdx;
+                               long cIdx;
+
+                               int n = _byRow ? blk.getNumRows() : 
blk.getNumColumns();
+                               for(int i = 0; i < n; i++) {
+                                       MatrixBlock slice;
+                                       if(_byRow) {
+                                               slice = blk.slice(i, i);
+                                               rIdx = (r - 1) * _blen + i + 1;
+                                               cIdx = c;
+                                       }
+                                       else {
+                                               slice = blk.slice(0, 
blk.getNumRows() - 1, i, i);
+                                               cIdx = (c - 1) * _blen + i + 1;
+                                               rIdx = r;
+                                       }
+
+                                       long targetIdx = _byRow ? (rIdx - 1) * 
numColBlocks + c - 1 : (cIdx - 1) * numRowBlocks + r - 1;
+                                       IndexedMatrixValue sliceImv = new 
IndexedMatrixValue(new MatrixIndexes(rIdx, cIdx), slice);
+                                       
budget.reserveBlocking(singleSliceBytes);
+                                       table.put((int) targetIdx, new 
ManagedPayload<>(sliceImv, singleSliceBytes, budget));
+                               }
+                       }
+                       catch(IllegalStateException e) {
+                               throw new DMLRuntimeException(e);
+                       }

Review Comment:
   I would suggest to try to move these lambda functions out to simplify the 
methods. Try to reduce them to 10-20 lines max per function.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to