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]
