From 6bad9a80088b1580109c97aa9cc81c8453361e84 Mon Sep 17 00:00:00 2001 From: ilgeco Date: Mon, 20 Jul 2026 11:34:59 +0200 Subject: [PATCH] Resnet is fast --- src/PIM/Common/IR/TensorSliceUtils.cpp | 11 +- src/PIM/Compiler/PimCodeGen.cpp | 2 +- .../ONNXToSpatial/LowerSpatialPlansPass.cpp | 32 ++ .../ONNXToSpatial/ONNXToSpatialPass.cpp | 3 +- .../ONNXToSpatial/ONNXToSpatialVerifier.cpp | 1 + .../ONNXToSpatial/Patterns/Math/Conv.cpp | 453 ++++++++++++++---- .../ONNXToSpatial/Patterns/NN/Pool.cpp | 235 ++++++++- .../Conversion/ONNXToSpatial/PlanLowering.hpp | 7 + .../SpatialLayoutPlanningPass.cpp | 23 +- .../BatchCoreLoweringPatterns.cpp | 32 +- src/PIM/Dialect/Spatial/Spatial.td | 19 + src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp | 20 + .../Scheduling/ComputeGraph.cpp | 2 +- .../Scheduling/ComputeInstanceUtils.cpp | 64 ++- .../Scheduling/ComputeInstanceUtils.hpp | 6 +- .../Scheduling/PeftScheduler.cpp | 111 +++-- 16 files changed, 883 insertions(+), 138 deletions(-) diff --git a/src/PIM/Common/IR/TensorSliceUtils.cpp b/src/PIM/Common/IR/TensorSliceUtils.cpp index c8991c3..a74f2c0 100644 --- a/src/PIM/Common/IR/TensorSliceUtils.cpp +++ b/src/PIM/Common/IR/TensorSliceUtils.cpp @@ -80,8 +80,17 @@ Value extractMixedSliceOrIdentity(RewriterBase &rewriter, Value insertMixedSlice(OpBuilder &builder, Location loc, Value source, Value dest, const MixedSliceGeometry &geometry) { + SmallVector sizes(geometry.sizes); + auto sourceType = dyn_cast(source.getType()); + auto destType = dyn_cast(dest.getType()); + if (sourceType && destType && sourceType.hasStaticShape() + && sourceType.getRank() == destType.getRank()) { + sizes.clear(); + for (int64_t dimension : sourceType.getShape()) + sizes.push_back(builder.getIndexAttr(dimension)); + } return tensor::InsertSliceOp::create(builder, loc, source, dest, - geometry.offsets, geometry.sizes, + geometry.offsets, sizes, geometry.strides); } diff --git a/src/PIM/Compiler/PimCodeGen.cpp b/src/PIM/Compiler/PimCodeGen.cpp index 4a35aee..322d56f 100644 --- a/src/PIM/Compiler/PimCodeGen.cpp +++ b/src/PIM/Compiler/PimCodeGen.cpp @@ -1590,7 +1590,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std:: } for (auto [slot, fileName] : llvm::enumerate(weightFiles)) { - xbarsPerGroup.push_back(static_cast(slot)); + xbarsPerGroup.push_back(1); std::string sourcePath = outputDirPath + "/weights/" + fileName; std::string targetPath = coreWeightsDirPath + "/crossbar_" + std::to_string(slot) + ".bin"; sys::fs::remove(targetPath); diff --git a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp index e60c2ef..a0e983b 100644 --- a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp @@ -191,6 +191,37 @@ struct LowerSpatialPlansPass final : PassWrapper(&op)) { + auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { + auto blueprint = dyn_cast(user); + return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; + }); + if (outputBlueprint == planOp.getResult().getUsers().end()) { + planOp.emitOpError("selected MaxPool plan requires a row-strip blueprint result"); + signalPassFailure(); + return; + } + + FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); + rewriter.setInsertionPoint(planOp); + FailureOr lowered = lowerSelectedMaxPool2DPlan( + planOp, succeeded(input) ? std::optional {input->storage} : std::nullopt, rewriter); + if (failed(lowered)) { + planOp.emitOpError("failed to lower selected row-strip Spatial MaxPool plan"); + signalPassFailure(); + return; + } + auto blueprint = cast(*outputBlueprint); + FailureOr output = buildRowStripValue(blueprint, *lowered); + if (failed(output)) { + signalPassFailure(); + return; + } + rowStripValues[blueprint.getResult()] = *output; + eraseAfterLowering.insert(planOp); + eraseAfterLowering.insert(blueprint); + continue; + } if (auto planOp = dyn_cast(&op)) { if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) { auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { @@ -386,6 +417,7 @@ struct LowerSpatialPlansPass final : PassWrapper(op) || op->getDialect()->getNamespace() == "onnx") { op->emitOpError("operation must not remain after LowerSpatialPlans"); diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp index 0d5feb2..f9a7637 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp @@ -48,10 +48,11 @@ static void populateEmptyFunction(func::FuncOp funcOp) { SmallVector convPlans(funcOp.getOps()); SmallVector biasAddPlans(funcOp.getOps()); SmallVector reluPlans(funcOp.getOps()); + SmallVector maxPoolPlans(funcOp.getOps()); SmallVector blueprints(funcOp.getOps()); SmallVector materializers(funcOp.getOps()); if (!computes.empty() || !computeBatches.empty() || !convPlans.empty() || !biasAddPlans.empty() || !reluPlans.empty() - || !blueprints.empty() || !materializers.empty()) { + || !maxPoolPlans.empty() || !blueprints.empty() || !materializers.empty()) { return; } diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp index fd220e1..d1230ab 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp @@ -147,6 +147,7 @@ void verifyLogicalTopLevelOps(func::FuncOp funcOp, pim::CappedDiagnosticReporter spatial::SpatConv2DPlanOp, spatial::SpatBiasAddPlanOp, spatial::SpatReluPlanOp, + spatial::SpatMaxPool2DPlanOp, spatial::SpatBlueprintOp, spatial::SpatMaterializeLayoutOp>(&op)) { continue; diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp index 20a5c6c..c5625fc 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp @@ -1890,6 +1890,37 @@ static Value createPaddedInputKTiledWeightConstant(DenseElementsAttr sourceAttr, return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType); } +static Value createPaddedOutputChannelTiledWeightConstant(DenseElementsAttr sourceAttr, + const ConvLoweringState& state, + int64_t paddedK, + int64_t xbarDim, + PatternRewriter& rewriter) { + const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim); + auto paddedType = + RankedTensorType::get({outputTileCount, paddedK, xbarDim}, state.wType.getElementType()); + SmallVector sourceValues(sourceAttr.getValues()); + SmallVector paddedValues( + paddedType.getNumElements(), cast(rewriter.getZeroAttr(paddedType.getElementType()))); + for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel) { + const int64_t outputTile = outChannel / xbarDim; + const int64_t tileChannel = outChannel % xbarDim; + for (int64_t inChannel = 0; inChannel < state.numChannelsIn; ++inChannel) { + for (int64_t kernelH = 0; kernelH < state.wHeight; ++kernelH) { + for (int64_t kernelW = 0; kernelW < state.wWidth; ++kernelW) { + const int64_t sourceFlatIndex = + (((outChannel * state.numChannelsIn) + inChannel) * state.wHeight + kernelH) * state.wWidth + kernelW; + const int64_t patchIndex = ((inChannel * state.wHeight) + kernelH) * state.wWidth + kernelW; + const int64_t destinationFlatIndex = + ((outputTile * paddedK) + patchIndex) * xbarDim + tileChannel; + paddedValues[destinationFlatIndex] = sourceValues[sourceFlatIndex]; + } + } + } + } + auto paddedAttr = DenseElementsAttr::get(paddedType, paddedValues); + return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType); +} + static FailureOr rewriteInputKTiledConv(const ConvLoweringState& state, ArrayRef distributedConsumers, PatternRewriter& rewriter, @@ -2541,10 +2572,16 @@ static Value createHorizontallyPaddedRowStripFragment(Value fragment, const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { - auto paddedType = RankedTensorType::get({1, state.numChannelsIn, 1, state.xWidth + 2}, - state.xType.getElementType(), - state.xType.getEncoding()); - return createZeroPaddedTensor(fragment, paddedType, {0, 0, 0, 1}, {0, 0, 0, 1}, rewriter, loc); + auto paddedType = RankedTensorType::get( + {1, state.numChannelsIn, 1, state.xWidth + state.padWidthBegin + state.padWidthEnd}, + state.xType.getElementType(), + state.xType.getEncoding()); + return createZeroPaddedTensor(fragment, + paddedType, + {0, 0, 0, state.padWidthBegin}, + {0, 0, 0, state.padWidthEnd}, + rewriter, + loc); } static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state, PatternRewriter& rewriter) { @@ -2554,7 +2591,8 @@ static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state, values.reserve(tableType.getNumElements()); for (int64_t outputRow = 0; outputRow < state.outHeight; ++outputRow) { for (int64_t kernelRow = 0; kernelRow < state.wHeight; ++kernelRow) { - int64_t sourceRow = outputRow + kernelRow - state.padHeightBegin; + int64_t sourceRow = + outputRow * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin; sourceRow = std::clamp(sourceRow, int64_t {0}, state.xHeight - 1); values.push_back(rewriter.getIndexAttr(sourceRow)); } @@ -2588,6 +2626,26 @@ static Value extractProjectedRowStripWindowRow(Value rowStripStorage, return extractRowStripFragment(rowStripStorage, state.xType, sourceRow, rewriter, loc); } +static Value extractDenseConvWindowRow(Value denseInput, + Value sourceRowTable, + const ConvLoweringState& state, + Value outputHeight, + Value kernelRow, + PatternRewriter& rewriter, + Location loc) { + Value tableIndex = createRowStripWindowTableIndex(outputHeight, kernelRow, state, rewriter, loc); + Value sourceRow = tensor::ExtractOp::create(rewriter, loc, sourceRowTable, ValueRange {tableIndex}).getResult(); + auto fragmentType = getRowStripFragmentType(state.xType); + SmallVector offsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)}; + SmallVector sizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.numChannelsIn), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.xWidth)}; + return tensor::ExtractSliceOp::create( + rewriter, loc, fragmentType, denseInput, offsets, sizes, getUnitStrides(rewriter, 4)); +} + static FailureOr createRowStripWindowMaskTable(const ConvLoweringState& state, PatternRewriter& rewriter) { auto elementType = state.xType.getElementType(); auto floatType = dyn_cast(elementType); @@ -2604,7 +2662,8 @@ static FailureOr createRowStripWindowMaskTable(const ConvLoweringState& s values.reserve(tableType.getNumElements()); for (int64_t outputRow = 0; outputRow < state.outHeight; ++outputRow) { for (int64_t kernelRow = 0; kernelRow < state.wHeight; ++kernelRow) { - int64_t sourceRow = outputRow + kernelRow - state.padHeightBegin; + int64_t sourceRow = + outputRow * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin; Attribute value = (sourceRow < 0 || sourceRow >= state.xHeight) ? zero : one; for (int64_t channel = 0; channel < state.numChannelsIn; ++channel) for (int64_t width = 0; width < state.xWidth; ++width) @@ -2638,15 +2697,20 @@ static Value extractProjectedRowStripWindowMask(Value maskTable, getUnitStrides(rewriter, 4)); } -static FailureOr createNchwRowStripConvWindow(Value rowStripStorage, - const ConvLoweringState& state, - Value outputHeight, - PatternRewriter& rewriter, - Location loc) { +static FailureOr createConvInputWindow(Value input, + const ConvLoweringState& state, + Value outputHeight, + PatternRewriter& rewriter, + Location loc) { auto fragmentType = getRowStripFragmentType(state.xType); - auto paddedWindowType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.xWidth + 2}, - state.xType.getElementType(), - state.xType.getEncoding()); + auto inputType = dyn_cast(input.getType()); + const bool denseInput = inputType == state.xType; + if (!denseInput && inputType != getRowStripStorageType(state.xType)) + return failure(); + auto paddedWindowType = RankedTensorType::get( + {1, state.numChannelsIn, state.wHeight, state.xWidth + state.padWidthBegin + state.padWidthEnd}, + state.xType.getElementType(), + state.xType.getEncoding()); Value sourceRowTable = createRowStripWindowSourceRowTable(state, rewriter); FailureOr maskTable = createRowStripWindowMaskTable(state, rewriter); if (failed(maskTable)) @@ -2657,8 +2721,10 @@ static FailureOr createNchwRowStripConvWindow(Value rowStripStorage, Value window = initWindow; for (int64_t kernelRowIndex = 0; kernelRowIndex < state.wHeight; ++kernelRowIndex) { Value kernelRow = getOrCreateIndexConstant(rewriter, anchorOp, kernelRowIndex); - Value sourceRow = - extractProjectedRowStripWindowRow(rowStripStorage, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc); + Value sourceRow = denseInput + ? extractDenseConvWindowRow(input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc) + : extractProjectedRowStripWindowRow( + input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc); Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, loc); Value semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, sourceRow, mask).getResult(); Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, loc); @@ -2673,7 +2739,9 @@ static FailureOr createNchwRowStripConvWindow(Value rowStripStorage, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsIn), rewriter.getIndexAttr(1), - rewriter.getIndexAttr(state.xWidth + 2)}, + rewriter.getIndexAttr( + state.xWidth + state.padWidthBegin + + state.padWidthEnd)}, getUnitStrides(rewriter, 4)); } return window; @@ -2691,12 +2759,13 @@ static FailureOr createNchwRowStripConvPatchRow(Value paddedWindow, state.xType.getEncoding()); auto rowType = RankedTensorType::get({1, patchSize}, state.xType.getElementType(), state.xType.getEncoding()); Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); + Value inputWidthOffset = affineMulConst(rewriter, loc, outputWidth, state.strideWidth, anchorOp); Value patch = createConvInputPatch(paddedWindow, patchType, c0, c0, c0, - outputWidth, + inputWidthOffset, state.dilationHeight, state.dilationWidth, rewriter, @@ -2706,6 +2775,57 @@ static FailureOr createNchwRowStripConvPatchRow(Value paddedWindow, .getResult(); } +static FailureOr createPaddedConvOutputTile(Value paddedPatchRow, + Value tileWeights, + int64_t numKSlices, + int64_t xbarDim, + PatternRewriter& rewriter, + Location loc) { + auto elementType = cast(paddedPatchRow.getType()).getElementType(); + auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType); + auto weightElementType = cast(tileWeights.getType()).getElementType(); + auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType); + + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); + Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); + Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices); + Value cXbar = getOrCreateIndexConstant(rewriter, anchorOp, xbarDim); + auto createPiece = [&](Value kSlice, Location pieceLoc) -> Value { + Value kOffset = arith::MulIOp::create(rewriter, pieceLoc, kSlice, cXbar); + SmallVector aOffsets {rewriter.getIndexAttr(0), kOffset}; + SmallVector aSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)}; + Value aTile = extractStaticSliceOrIdentity( + rewriter, pieceLoc, paddedPatchRow, paddedRowType, aOffsets, aSizes, getUnitStrides(rewriter, 2)); + SmallVector bOffsets {kOffset, rewriter.getIndexAttr(0)}; + SmallVector bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)}; + Value bTile = extractStaticSliceOrIdentity( + rewriter, pieceLoc, tileWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2)); + return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, bTile, aTile).getResult(); + }; + + Value tileResult = createPiece(c0, loc); + if (numKSlices == 1) + return tileResult; + + auto kLoop = buildNormalizedScfFor( + rewriter, + loc, + c1, + cNumKSlices, + c1, + ValueRange {tileResult}, + [&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange reduceIterArgs, SmallVectorImpl& reduceYielded) { + Value piece = createPiece(kSlice, reduceLoc); + reduceYielded.push_back( + spatial::SpatVAddOp::create(rewriter, reduceLoc, paddedRowType, reduceIterArgs.front(), piece).getResult()); + return success(); + }); + if (failed(kLoop)) + return failure(); + return kLoop->results.front(); +} + static FailureOr createPaddedConvOutputRow(Value patchRow, const ConvLoweringState& state, Value paddedWeights, @@ -2720,67 +2840,241 @@ static FailureOr createPaddedConvOutputRow(Value patchRow, auto rowType = RankedTensorType::get({1, state.numChannelsOut}, elementType); auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType); auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType); - auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType()); + auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, state.wType.getElementType()); + const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim); Value paddedPatchRow = patchRow; if (patchSize != paddedK) paddedPatchRow = createZeroPaddedTensor( paddedPatchRow, paddedPatchRowType, {0, 0}, {0, paddedK - patchSize}, rewriter, loc); - Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); - Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); - Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); - Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices); - Value cXbar = getOrCreateIndexConstant(rewriter, anchorOp, xbarDim); - auto createPiece = [&](Value kSlice, Location pieceLoc) -> Value { - Value kOffset = arith::MulIOp::create(rewriter, pieceLoc, kSlice, cXbar); - SmallVector aOffsets {rewriter.getIndexAttr(0), kOffset}; - SmallVector aSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)}; - Value aTile = extractStaticSliceOrIdentity( - rewriter, pieceLoc, paddedPatchRow, paddedRowType, aOffsets, aSizes, getUnitStrides(rewriter, 2)); - SmallVector bOffsets {kOffset, rewriter.getIndexAttr(0)}; - SmallVector bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)}; - Value bTile = extractStaticSliceOrIdentity( - rewriter, pieceLoc, paddedWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2)); - return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, bTile, aTile).getResult(); + auto getTileWeights = [&](int64_t outputTile) { + if (outputTileCount == 1) + return paddedWeights; + SmallVector offsets { + rewriter.getIndexAttr(outputTile), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; + SmallVector sizes { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)}; + return extractStaticSliceOrIdentity( + rewriter, loc, paddedWeights, tileWeightsType, offsets, sizes, getUnitStrides(rewriter, 3)); }; - Value rowResult = createPiece(c0, loc); - if (numKSlices > 1) { - auto kLoop = buildNormalizedScfFor( - rewriter, - loc, - c1, - cNumKSlices, - c1, - ValueRange {rowResult}, - [&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange reduceIterArgs, SmallVectorImpl& reduceYielded) { - Value piece = createPiece(kSlice, reduceLoc); - reduceYielded.push_back( - spatial::SpatVAddOp::create(rewriter, reduceLoc, paddedRowType, reduceIterArgs.front(), piece).getResult()); - return success(); - }); - if (failed(kLoop)) + if (outputTileCount == 1) { + FailureOr rowResult = createPaddedConvOutputTile( + paddedPatchRow, getTileWeights(0), numKSlices, xbarDim, rewriter, loc); + if (failed(rowResult)) return failure(); - rowResult = kLoop->results.front(); + if (paddedBias) + rowResult = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, *rowResult, paddedBias).getResult(); + if (state.numChannelsOut == xbarDim) + return *rowResult; + + SmallVector outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; + SmallVector outputSizes { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)}; + return tensor::ExtractSliceOp::create( + rewriter, loc, rowType, *rowResult, outputOffsets, outputSizes, getUnitStrides(rewriter, 2)) + .getResult(); } - if (paddedBias) - rowResult = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, rowResult, paddedBias).getResult(); - if (state.numChannelsOut == xbarDim) - return rowResult; + const int64_t paddedOutputChannels = outputTileCount * xbarDim; + auto paddedOutputType = RankedTensorType::get({1, paddedOutputChannels}, elementType); + Value paddedOutput = tensor::EmptyOp::create(rewriter, loc, paddedOutputType.getShape(), elementType); + for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) { + FailureOr tileResult = createPaddedConvOutputTile( + paddedPatchRow, getTileWeights(outputTile), numKSlices, xbarDim, rewriter, loc); + if (failed(tileResult)) + return failure(); + SmallVector tileOffsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(outputTile * xbarDim)}; + SmallVector tileSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)}; + paddedOutput = tensor::InsertSliceOp::create( + rewriter, loc, *tileResult, paddedOutput, tileOffsets, tileSizes, getUnitStrides(rewriter, 2)); + } SmallVector outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; SmallVector outputSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)}; return tensor::ExtractSliceOp::create( - rewriter, loc, rowType, rowResult, outputOffsets, outputSizes, getUnitStrides(rewriter, 2)) + rewriter, loc, rowType, paddedOutput, outputOffsets, outputSizes, getUnitStrides(rewriter, 2)) .getResult(); } +static bool rowStripOutputFitsOneCore(const ConvGeometry& geometry) { + const int64_t inputTileCount = ceilIntegerDivide(geometry.k, geometry.xbarSize); + const int64_t outputTileCount = ceilIntegerDivide(geometry.c, geometry.xbarSize); + return inputTileCount * outputTileCount <= static_cast(crossbarCountInCore.getValue()); +} + +static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) { + return ceilIntegerDivide(geometry.k, geometry.xbarSize) + <= static_cast(crossbarCountInCore.getValue()); +} + +static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLoweringState& state, + Value paddedWeights, + int64_t paddedK, + int64_t numKSlices, + int64_t xbarDim, + PatternRewriter& rewriter, + Location loc) { + const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim); + const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + auto elementType = state.outType.getElementType(); + auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType); + auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, state.wType.getElementType()); + SmallVector outputTiles; + outputTiles.reserve(outputTileCount); + + for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) { + const int64_t channelOffset = outputTile * xbarDim; + const int64_t tileChannels = std::min(xbarDim, state.numChannelsOut - channelOffset); + auto tileRowType = RankedTensorType::get({1, tileChannels}, elementType); + auto tilePixelType = RankedTensorType::get({1, tileChannels, 1, 1}, elementType); + auto tileFragmentType = RankedTensorType::get({1, tileChannels, 1, state.outWidth}, elementType); + auto tileStorageType = spatial::getGraphBatchPhysicalResultType(state.outHeight, tileFragmentType); + SmallVector weightOffsets { + rewriter.getIndexAttr(outputTile), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; + SmallVector weightSizes { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)}; + Value tileWeights = extractStaticSliceOrIdentity( + rewriter, loc, paddedWeights, tileWeightsType, weightOffsets, weightSizes, getUnitStrides(rewriter, 3)); + + auto tileBatch = createSpatComputeBatch( + rewriter, + loc, + TypeRange {tileStorageType}, + state.outHeight, + ValueRange {tileWeights}, + ValueRange {state.x}, + [&](detail::SpatComputeBatchBodyArgs args) { + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); + Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); + Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth); + FailureOr inputWindow = + createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc); + if (failed(inputWindow)) + return failure(); + Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType); + auto widthLoop = buildNormalizedScfFor( + rewriter, + loc, + c0, + cOutWidth, + c1, + ValueRange {fragmentInit}, + [&](OpBuilder&, + Location widthLoc, + Value widthIndex, + ValueRange widthIterArgs, + SmallVectorImpl& widthYielded) { + FailureOr patchRow = + createNchwRowStripConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc); + if (failed(patchRow)) + return failure(); + Value paddedPatchRow = *patchRow; + if (patchSize != paddedK) + paddedPatchRow = createZeroPaddedTensor( + paddedPatchRow, paddedPatchRowType, {0, 0}, {0, paddedK - patchSize}, rewriter, widthLoc); + FailureOr paddedOutputRow = createPaddedConvOutputTile( + paddedPatchRow, args.weights.front(), numKSlices, xbarDim, rewriter, widthLoc); + if (failed(paddedOutputRow)) + return failure(); + Value outputRow = *paddedOutputRow; + if (tileChannels != xbarDim) { + SmallVector rowOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; + SmallVector rowSizes { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels)}; + outputRow = tensor::ExtractSliceOp::create(rewriter, + widthLoc, + tileRowType, + outputRow, + rowOffsets, + rowSizes, + getUnitStrides(rewriter, 2)); + } + Value outputPixel = tensor::ExpandShapeOp::create( + rewriter, widthLoc, tilePixelType, outputRow, SmallVector {{0}, {1, 2, 3}}); + SmallVector rowOffsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), widthIndex}; + SmallVector rowSizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1)}; + Value nextFragment = tensor::InsertSliceOp::create(rewriter, + widthLoc, + outputPixel, + widthIterArgs.front(), + rowOffsets, + rowSizes, + getUnitStrides(rewriter, 4)); + widthYielded.push_back(nextFragment); + return success(); + }); + if (failed(widthLoop)) + return failure(); + publishGraphBatchPhysicalFragment( + rewriter, loc, widthLoop->results.front(), args.outputs.front(), args.lane); + return success(); + }); + if (failed(tileBatch)) + return failure(); + outputTiles.push_back(tileBatch->getResult(0)); + } + + auto fragmentType = getRowStripFragmentType(state.outType); + auto outputStorageType = getRowStripStorageType(state.outType); + auto assemblyBatch = createSpatComputeBatch(rewriter, + loc, + TypeRange {outputStorageType}, + state.outHeight, + {}, + ValueRange(outputTiles), + [&](detail::SpatComputeBatchBodyArgs args) { + Value fragment = tensor::EmptyOp::create( + rewriter, loc, fragmentType.getShape(), elementType); + for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) { + const int64_t channelOffset = outputTile * xbarDim; + const int64_t tileChannels = + std::min(xbarDim, state.numChannelsOut - channelOffset); + auto tileFragmentType = RankedTensorType::get( + {1, tileChannels, 1, state.outWidth}, elementType); + FailureOr tileFragment = extractGraphBatchPhysicalFragment( + rewriter, loc, args.inputs[outputTile], args.lane, tileFragmentType); + if (failed(tileFragment)) + return failure(); + SmallVector offsets {rewriter.getIndexAttr(0), + rewriter.getIndexAttr(channelOffset), + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(0)}; + SmallVector sizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.outWidth)}; + fragment = tensor::InsertSliceOp::create(rewriter, + loc, + *tileFragment, + fragment, + offsets, + sizes, + getUnitStrides(rewriter, 4)); + } + insertRowStripFragment( + fragment, args.outputs.front(), state.outType, args.lane, rewriter, loc); + return success(); + }); + if (failed(assemblyBatch)) + return failure(); + Value output = assemblyBatch->getResult(0); + if (state.hasBias) + return applyRowStripBiasAdd(output, state.outType, state.b, rewriter, loc); + return output; +} + static FailureOr createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { ConvGeometry geometry = buildConvGeometry(state); - if (state.group != 1 || state.batchSize != 1 || geometry.c > geometry.xbarSize) + if (state.group != 1 || state.batchSize != 1 || !rowStripOutputTileFitsOneCore(geometry)) return failure(); auto weightDenseAttr = getHostConstDenseElementsAttr(state.w); @@ -2796,12 +3090,17 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe auto elementType = state.outType.getElementType(); auto fragmentType = getRowStripFragmentType(state.outType); auto outputPixelType = RankedTensorType::get({1, state.numChannelsOut, 1, 1}, elementType); - auto patchType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.wWidth}, state.xType.getElementType()); - auto patchRowType = RankedTensorType::get({1, patchSize}, state.xType.getElementType()); auto outputStorageType = getRowStripStorageType(state.outType); - PreparedConvInput preparedInput = standard::prepareInputForIm2Col(state, rewriter, loc); - Value paddedWeights = standard::createPaddedInputKTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter); + Value paddedWeights = state.numChannelsOut <= xbarDim + ? standard::createPaddedInputKTiledWeightConstant( + weightDenseAttr, state, paddedK, xbarDim, rewriter) + : standard::createPaddedOutputChannelTiledWeightConstant( + weightDenseAttr, state, paddedK, xbarDim, rewriter); + if (!rowStripOutputFitsOneCore(geometry)) + return createOutputChannelTiledRowStripConvOutput( + state, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc); + FailureOr paddedBias = failure(); if (state.hasBias) paddedBias = createPaddedBiasRowConstant(state, xbarDim, rewriter); @@ -2814,13 +3113,16 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe TypeRange {outputStorageType}, state.outHeight, ValueRange {paddedWeights}, - state.hasBias ? ValueRange {preparedInput.value, *paddedBias} : ValueRange {preparedInput.value}, + state.hasBias ? ValueRange {state.x, *paddedBias} : ValueRange {state.x}, [&](detail::SpatComputeBatchBodyArgs args) { Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth); - Value inputHeightOffset = affineMulConst(rewriter, loc, args.lane, state.strideHeight, anchorOp); + FailureOr inputWindow = + createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc); + if (failed(inputWindow)) + return failure(); Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType); auto widthLoop = buildNormalizedScfFor( rewriter, @@ -2830,20 +3132,11 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe c1, ValueRange {fragmentInit}, [&](OpBuilder&, Location widthLoc, Value widthIndex, ValueRange widthIterArgs, SmallVectorImpl& widthYielded) { - Value inputWidthOffset = affineMulConst(rewriter, widthLoc, widthIndex, state.strideWidth, anchorOp); - Value patch = createConvInputPatch(args.inputs.front(), - patchType, - c0, - c0, - inputHeightOffset, - inputWidthOffset, - state.dilationHeight, - state.dilationWidth, - rewriter, - widthLoc); - Value patchRow = tensor::CollapseShapeOp::create( - rewriter, widthLoc, patchRowType, patch, SmallVector {{0}, {1, 2, 3}}); - FailureOr outputRow = createPaddedConvOutputRow(patchRow, + FailureOr patchRow = + createNchwRowStripConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc); + if (failed(patchRow)) + return failure(); + FailureOr outputRow = createPaddedConvOutputRow(*patchRow, state, args.weights.front(), state.hasBias ? args.inputs[1] : Value(), @@ -2924,7 +3217,7 @@ static FailureOr createConvOutputFromNchwRowStripFragments(Value rowStrip Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth); auto fragmentType = getRowStripFragmentType(state.outType); - FailureOr inputWindow = createNchwRowStripConvWindow(args.inputs.front(), state, args.lane, rewriter, loc); + FailureOr inputWindow = createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc); if (failed(inputWindow)) return failure(); Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType); @@ -3811,7 +4104,7 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) { analysis.barrierKind = DistributedConvBarrierKind::UnsupportedConsumer; analysis.barrierDetail = "selected row-strip layout"; ConvGeometry geometry = buildConvGeometry(*state); - if (geometry.c > geometry.xbarSize) + if (!rowStripOutputTileFitsOneCore(geometry)) return failure(); ConvLoweringDecision decision = chooseConvLoweringStrategy(geometry, *requestedStrategy, analysis); if (decision.strategy == PimConvLoweringDepthwise && !depthwise::canUseStructuredRewrite(*state) diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp index 596523c..abce781 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp @@ -16,6 +16,8 @@ #include "src/Accelerators/PIM/Common/PimCommon.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Dialect/ONNX/ONNXOps.hpp" @@ -24,7 +26,7 @@ using namespace mlir; namespace onnx_mlir { namespace { -static Value materializeTileTensor(ConversionPatternRewriter& rewriter, Location loc, Value tile) { +static Value materializeTileTensor(PatternRewriter& rewriter, Location loc, Value tile) { auto tileType = cast(tile.getType()); Value empty = tensor::EmptyOp::create(rewriter, loc, tileType.getShape(), tileType.getElementType()); return insertStaticSlice(rewriter, loc, tile, empty, getZeroOffsets(rewriter, tileType.getRank())); @@ -228,6 +230,23 @@ struct PoolToSpatialComputeBase : public OpConversionPattern { } } + if constexpr (std::is_same_v) { + if (batchSize == 1) { + auto plan = spatial::SpatMaxPool2DPlanOp::create( + rewriter, + loc, + outType, + x, + rewriter.getDenseI64ArrayAttr({kernelHeight, kernelWidth}), + rewriter.getDenseI64ArrayAttr({padTop, padLeft, padBottom, padRight}), + rewriter.getDenseI64ArrayAttr({strideHeight, strideWidth}), + rewriter.getDenseI64ArrayAttr({dilationHeight, dilationWidth}), + rewriter.getStringAttr("nchw")); + rewriter.replaceOp(poolOp, plan.getResult()); + return success(); + } + } + const int64_t xbarSize = static_cast(crossbarSize.getValue()); const int64_t channelTileCount = (channels + xbarSize - 1) / xbarSize; const int64_t outputPatchCount = batchSize * outputHeight * outputWidth; @@ -396,6 +415,220 @@ struct PoolToSpatialCompute } // namespace +LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) { + auto inputType = dyn_cast(planOp.getInput().getType()); + auto outputType = dyn_cast(planOp.getOutput().getType()); + if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape()) + return failure(); + if (inputType.getRank() != 4 || outputType.getRank() != 4 || inputType.getDimSize(0) != 1 + || outputType.getDimSize(0) != 1 || inputType.getDimSize(1) != outputType.getDimSize(1)) + return failure(); + if (llvm::any_of(planOp.getKernelShape(), [](int64_t value) { return value <= 0; }) + || llvm::any_of(planOp.getStrides(), [](int64_t value) { return value <= 0; }) + || llvm::any_of(planOp.getDilations(), [](int64_t value) { return value <= 0; })) + return failure(); + return success(); +} + +static Value createClampedPoolIndexTable(PatternRewriter& rewriter, + Operation* anchorOp, + int64_t outputSize, + int64_t kernelSize, + int64_t stride, + int64_t dilation, + int64_t padBegin, + int64_t inputSize) { + auto tableType = RankedTensorType::get({outputSize * kernelSize}, rewriter.getIndexType()); + SmallVector values; + values.reserve(tableType.getNumElements()); + for (int64_t output = 0; output < outputSize; ++output) + for (int64_t kernel = 0; kernel < kernelSize; ++kernel) + values.push_back(rewriter.getIndexAttr( + std::clamp(output * stride + kernel * dilation - padBegin, int64_t {0}, inputSize - 1))); + return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType); +} + +static Value extractPoolIndex(PatternRewriter& rewriter, + Location loc, + Operation* anchorOp, + Value table, + Value outputIndex, + int64_t kernelIndex, + int64_t kernelSize) { + Value tableIndex = arith::MulIOp::create( + rewriter, loc, outputIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelSize)); + if (kernelIndex != 0) + tableIndex = arith::AddIOp::create( + rewriter, loc, tableIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelIndex)); + return tensor::ExtractOp::create(rewriter, loc, table, tableIndex); +} + +FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, + std::optional rowStripInput, + PatternRewriter& rewriter) { + if (failed(canLowerMaxPoolPlanToRowStrip(planOp))) + return failure(); + + Location loc = planOp.getLoc(); + auto inputType = cast(planOp.getInput().getType()); + auto outputType = cast(planOp.getOutput().getType()); + const int64_t channels = inputType.getDimSize(1); + const int64_t inputHeight = inputType.getDimSize(2); + const int64_t inputWidth = inputType.getDimSize(3); + const int64_t outputHeight = outputType.getDimSize(2); + const int64_t outputWidth = outputType.getDimSize(3); + const int64_t kernelHeight = planOp.getKernelShape()[0]; + const int64_t kernelWidth = planOp.getKernelShape()[1]; + Value input = rowStripInput.value_or(planOp.getInput()); + auto actualInputType = dyn_cast(input.getType()); + const bool physicalInput = actualInputType == getRowStripStorageType(inputType); + if (!physicalInput && actualInputType != inputType) + return failure(); + + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + Value rowTable = createClampedPoolIndexTable(rewriter, + anchorOp, + outputHeight, + kernelHeight, + planOp.getStrides()[0], + planOp.getDilations()[0], + planOp.getPads()[0], + inputHeight); + Value columnTable = createClampedPoolIndexTable(rewriter, + anchorOp, + outputWidth, + kernelWidth, + planOp.getStrides()[1], + planOp.getDilations()[1], + planOp.getPads()[1], + inputWidth); + auto inputFragmentType = getRowStripFragmentType(inputType); + auto outputFragmentType = getRowStripFragmentType(outputType); + auto outputStorageType = getRowStripStorageType(outputType); + auto tileType = RankedTensorType::get({1, channels, 1, 1}, outputType.getElementType()); + auto batch = createSpatComputeBatch( + rewriter, + loc, + TypeRange {outputStorageType}, + outputHeight, + {}, + ValueRange {input}, + [&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult { + SmallVector inputRows; + inputRows.reserve(kernelHeight); + for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) { + Value sourceRow = + extractPoolIndex(rewriter, loc, anchorOp, rowTable, args.lane, kernelRow, kernelHeight); + if (physicalInput) { + inputRows.push_back( + extractRowStripFragment(args.inputs.front(), inputType, sourceRow, rewriter, loc)); + } + else { + SmallVector offsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)}; + inputRows.push_back(tensor::ExtractSliceOp::create(rewriter, + loc, + inputFragmentType, + args.inputs.front(), + offsets, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(channels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(inputWidth)}, + getUnitStrides(rewriter, 4))); + } + } + + auto windowType = RankedTensorType::get( + {1, channels, kernelHeight, inputWidth}, inputType.getElementType(), inputType.getEncoding()); + Value window = tensor::EmptyOp::create( + rewriter, loc, windowType.getShape(), windowType.getElementType()); + for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) { + SmallVector offsets {rewriter.getIndexAttr(0), + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(kernelRow), + rewriter.getIndexAttr(0)}; + window = tensor::InsertSliceOp::create(rewriter, + loc, + inputRows[kernelRow], + window, + offsets, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(channels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(inputWidth)}, + getUnitStrides(rewriter, 4)); + } + + Value outputInit = tensor::EmptyOp::create( + rewriter, loc, outputFragmentType.getShape(), outputFragmentType.getElementType()); + Operation* bodyAnchor = rewriter.getInsertionBlock()->getParentOp(); + Value c0 = getOrCreateIndexConstant(rewriter, bodyAnchor, 0); + Value c1 = getOrCreateIndexConstant(rewriter, bodyAnchor, 1); + Value cOutputWidth = getOrCreateIndexConstant(rewriter, bodyAnchor, outputWidth); + auto outputLoop = buildNormalizedScfFor( + rewriter, + loc, + c0, + cOutputWidth, + c1, + ValueRange {outputInit}, + [&](OpBuilder&, Location nestedLoc, Value outputColumn, ValueRange iterArgs, SmallVectorImpl& yielded) { + Value reduced; + for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) { + for (int64_t kernelColumn = 0; kernelColumn < kernelWidth; ++kernelColumn) { + Value sourceColumn = extractPoolIndex(rewriter, + nestedLoc, + bodyAnchor, + columnTable, + outputColumn, + kernelColumn, + kernelWidth); + SmallVector offsets { + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(kernelRow), + sourceColumn}; + Value point = tensor::ExtractSliceOp::create(rewriter, + nestedLoc, + tileType, + window, + offsets, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(channels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + reduced = reduced ? spatial::SpatVMaxOp::create(rewriter, nestedLoc, tileType, reduced, point).getResult() + : materializeTileTensor(rewriter, nestedLoc, point); + } + } + SmallVector outputOffsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), outputColumn}; + Value updated = tensor::InsertSliceOp::create(rewriter, + nestedLoc, + reduced, + iterArgs.front(), + outputOffsets, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(channels), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + yielded.push_back(updated); + return success(); + }); + if (failed(outputLoop)) + return failure(); + insertRowStripFragment( + outputLoop->results.front(), args.outputs.front(), outputType, args.lane, rewriter, loc); + return success(); + }); + if (failed(batch)) + return failure(); + return batch->getResult(0); +} + void populatePoolPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.insert>(ctx); patterns.insert>(ctx); diff --git a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp index dd5a722..100500f 100644 --- a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp @@ -18,4 +18,11 @@ lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp, mlir::LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp); mlir::LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp); +mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp); + +mlir::FailureOr +lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, + std::optional rowStripInput, + mlir::PatternRewriter& rewriter); + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp index d3de0d6..5c8df36 100644 --- a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp @@ -38,6 +38,8 @@ static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap(user)) return getSelectedLayout(layouts, convPlan.getResult()) == SelectedLayout::NchwRowStrip; + if (auto maxPoolPlan = dyn_cast(user)) + return getSelectedLayout(layouts, maxPoolPlan.getResult()) == SelectedLayout::NchwRowStrip; return false; } @@ -60,6 +62,8 @@ static bool canConsumeRowStripAsUser(Operation* user) { } if (auto convPlan = dyn_cast(user)) return succeeded(canConsumeAndProduceRowStrip(convPlan)); + if (auto maxPoolPlan = dyn_cast(user)) + return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan)); return false; } @@ -70,7 +74,6 @@ static bool hasRowStripConsumer(Value value) { return false; } - static bool canSelectConvRowStrip(spatial::SpatConv2DPlanOp convPlan, llvm::DenseMap& layouts) { SelectedLayout inputLayout = getSelectedLayout(layouts, convPlan.getInput()); @@ -83,9 +86,6 @@ static SelectedLayout chooseConvLayout(spatial::SpatConv2DPlanOp convPlan, llvm::DenseMap& layouts) { if (!canSelectConvRowStrip(convPlan, layouts)) return SelectedLayout::DenseNchw; - if (getSelectedLayout(layouts, convPlan.getInput()) != SelectedLayout::NchwRowStrip - && !hasRowStripConsumer(convPlan.getResult())) - return SelectedLayout::DenseNchw; if (!allUsersCanHandleRowStrip(convPlan.getResult(), layouts)) return SelectedLayout::DenseNchw; return SelectedLayout::NchwRowStrip; @@ -116,6 +116,11 @@ static SelectedLayout chooseBiasAddLayout(spatial::SpatBiasAddPlanOp biasAddPlan return SelectedLayout::NchwRowStrip; } +static SelectedLayout chooseMaxPoolLayout(spatial::SpatMaxPool2DPlanOp maxPoolPlan) { + return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan)) ? SelectedLayout::NchwRowStrip + : SelectedLayout::DenseNchw; +} + static spatial::SpatBlueprintOp insertRowStripBlueprint(IRRewriter& rewriter, Value value) { auto outputType = cast(value.getType()); auto [offsets, sizes] = buildRowStripMetadata(outputType); @@ -208,6 +213,14 @@ struct SpatialLayoutPlanningPass final : PassWrapper(&op)) { + SelectedLayout selected = chooseMaxPoolLayout(maxPoolPlan); + if (layouts[maxPoolPlan.getResult()] != selected) { + layouts[maxPoolPlan.getResult()] = selected; + changed = true; + } + continue; + } } } @@ -219,6 +232,8 @@ struct SpatialLayoutPlanningPass final : PassWrapper(&op)) producedValue = reluPlan.getResult(); + else if (auto maxPoolPlan = dyn_cast(&op)) + producedValue = maxPoolPlan.getResult(); else continue; diff --git a/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp b/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp index 5461e63..b31e3b4 100644 --- a/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp +++ b/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp @@ -29,6 +29,11 @@ static bool isUsedOnlyAsExplicitHostOperand(Value value) { }); } +static bool isUsedOnlyByExtractSlices(Value value) { + return !value.use_empty() + && llvm::all_of(value.getUsers(), [](Operation* user) { return isa(user); }); +} + static FailureOr getDirectReturnOperandIndex(OpResult result) { if (!result.hasOneUse()) return failure(); @@ -357,6 +362,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul rewriter.createBlock(&coreBatchOp.getBody(), coreBatchOp.getBody().end(), TypeRange(blockArgTypes), blockArgLocs); IRMapping mapper; + SmallPtrSet hostResidentTensors; rewriter.setInsertionPointToStart(newBlock); auto oldLaneArg = computeBatchOp.getLaneArgument(); if (!oldLaneArg) @@ -523,8 +529,10 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul if (isa_and_present(toTensorOp.getBuffer().getDefiningOp())) { Operation* cloned = rewriter.clone(op, mapper); auto clonedTensor = cloned->getResult(0); - if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult())) { + if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult()) + || isUsedOnlyByExtractSlices(toTensorOp.getResult())) { mapper.map(toTensorOp.getResult(), clonedTensor); + hostResidentTensors.insert(toTensorOp.getResult()); continue; } auto clonedType = cast(clonedTensor.getType()); @@ -542,6 +550,28 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul } } + if (auto extractSlice = dyn_cast(op); + extractSlice && hostResidentTensors.contains(extractSlice.getSource())) { + Operation* cloned = rewriter.clone(op, mapper); + Value hostSlice = cloned->getResult(0); + auto outputBuffer = createEmptyTensorFromShaped(rewriter, loc, cast(hostSlice.getType())); + Value zeroOffset = getOrCreateIndexConstant(rewriter, coreBatchOp.getOperation(), 0); + auto sizeAttr = getTensorSizeInBytesAttr(rewriter, coreBatchOp.getOperation(), hostSlice); + if (failed(sizeAttr)) + return failure(); + auto copied = pim::PimMemCopyHostToDevOp::create(rewriter, + loc, + outputBuffer.getType(), + zeroOffset, + zeroOffset, + outputBuffer, + hostSlice, + *sizeAttr) + .getOutput(); + mapper.map(extractSlice.getResult(), copied); + continue; + } + for (auto [operandIndex, operand] : llvm::enumerate(op.getOperands())) { if (!isa(operand.getType()) || mapper.contains(operand)) continue; diff --git a/src/PIM/Dialect/Spatial/Spatial.td b/src/PIM/Dialect/Spatial/Spatial.td index eb4c84a..5c8b585 100644 --- a/src/PIM/Dialect/Spatial/Spatial.td +++ b/src/PIM/Dialect/Spatial/Spatial.td @@ -270,6 +270,25 @@ def SpatReluPlanOp : SpatOp<"relu_plan", []> { let hasVerifier = 1; } +def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> { + let summary = "Layout-aware 2D NCHW MaxPool planning op"; + + let arguments = (ins + SpatTensor:$input, + DenseI64ArrayAttr:$kernelShape, + DenseI64ArrayAttr:$pads, + DenseI64ArrayAttr:$strides, + DenseI64ArrayAttr:$dilations, + StrAttr:$logicalLayout + ); + + let results = (outs + SpatTensor:$output + ); + + let hasVerifier = 1; +} + def SpatBiasAddPlanOp : SpatOp<"bias_add_plan", []> { let summary = "Layout-aware Conv-style bias add planning op"; diff --git a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp index 758b015..94af059 100644 --- a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp +++ b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp @@ -486,6 +486,26 @@ LogicalResult SpatReluPlanOp::verify() { return success(); } +LogicalResult SpatMaxPool2DPlanOp::verify() { + if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.max_pool2d_plan"))) + return failure(); + auto inputType = dyn_cast(getInput().getType()); + auto outputType = dyn_cast(getOutput().getType()); + if (!inputType.hasStaticShape() || !outputType.hasStaticShape() || inputType.getRank() != 4 + || outputType.getRank() != 4) + return emitError("requires static rank-4 input and output tensors"); + if (getLogicalLayout() != "nchw") + return emitError("requires logical layout \"nchw\""); + if (getKernelShape().size() != 2 || getStrides().size() != 2 || getDilations().size() != 2) + return emitError("requires two kernel, stride, and dilation values"); + if (getPads().size() != 4) + return emitError("requires four pad values"); + if (inputType.getDimSize(0) != outputType.getDimSize(0) + || inputType.getDimSize(1) != outputType.getDimSize(1)) + return emitError("requires matching input/output batch and channel dimensions"); + return success(); +} + LogicalResult SpatBiasAddPlanOp::verify() { if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.bias_add_plan"))) return failure(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp index 36796e4..25fbefd 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp @@ -780,7 +780,7 @@ ComputeGraph buildComputeGraph(Operation* entryOp) { if (auto batch = dyn_cast(&op)) { if (isUsedAsWeightOnly(batch.getOperation())) continue; - size_t chunkCount = getBatchChunkTargetCount(batch.getLaneCount()); + size_t chunkCount = getBatchChunkTargetCount(batch); for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) { ComputeInstance instance = getBatchChunkForIndex(batch, chunkIndex); size_t index = graph.nodes.size(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp index b526828..98a885e 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp @@ -1,9 +1,11 @@ #include "mlir/Dialect/Arith/IR/Arith.h" #include "mlir/Dialect/Tensor/IR/Tensor.h" +#include #include #include +#include "ComputeGraph.hpp" #include "ComputeInstanceUtils.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" @@ -12,20 +14,16 @@ using namespace mlir; namespace onnx_mlir { namespace spatial { +static constexpr llvm::StringLiteral kMergeChunkCountAttr = "spat.merge_chunk_count"; + size_t getSchedulingCpuBudget() { if (coresCount.getValue() > 0) return static_cast(coresCount.getValue()); return std::numeric_limits::max(); } -size_t getBatchChunkTargetCount(int32_t laneCount) { +static BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkCount, size_t chunkIndex) { assert(laneCount > 0 && "laneCount must be positive"); - return std::min(static_cast(laneCount), getSchedulingCpuBudget()); -} - -BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) { - assert(laneCount > 0 && "laneCount must be positive"); - size_t chunkCount = getBatchChunkTargetCount(laneCount); assert(chunkIndex < chunkCount && "chunkIndex out of range"); size_t laneCountSize = static_cast(laneCount); @@ -38,11 +36,51 @@ BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) { return {static_cast(start), static_cast(count)}; } -size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) { +static bool batchChunksFit(SpatComputeBatch batch, size_t chunkCount, size_t crossbarCapacity) { + for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) { + BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkCount, chunkIndex); + ComputeInstance instance {batch.getOperation(), chunk.laneStart, chunk.laneCount}; + if (getComputeInstanceCrossbarUsage(instance).size() > crossbarCapacity) + return false; + } + return true; +} + +size_t getBatchChunkTargetCount(SpatComputeBatch batch) { + if (auto chunkCount = batch->getAttrOfType(kMergeChunkCountAttr)) + return static_cast(chunkCount.getInt()); + + int32_t laneCount = batch.getLaneCount(); + assert(laneCount > 0 && "laneCount must be positive"); + size_t maxChunkCount = std::min(static_cast(laneCount), getSchedulingCpuBudget()); + size_t crossbarCapacity = crossbarCountInCore.getValue(); + CrossbarUsage fullUsage = collectDistinctCrossbarWeights(batch.getOperation()); + if (fullUsage.empty() || crossbarCapacity == 0) { + batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount)); + return maxChunkCount; + } + + size_t chunkCount = std::max(1, (fullUsage.size() + crossbarCapacity - 1) / crossbarCapacity); + for (; chunkCount <= maxChunkCount; ++chunkCount) { + if (batchChunksFit(batch, chunkCount, crossbarCapacity)) { + batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), chunkCount)); + return chunkCount; + } + } + batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount)); + return maxChunkCount; +} + +BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex) { + return getBatchChunkRange(batch.getLaneCount(), getBatchChunkTargetCount(batch), chunkIndex); +} + +size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane) { + int32_t laneCount = batch.getLaneCount(); assert(laneCount > 0 && "laneCount must be positive"); assert(lane < static_cast(laneCount) && "lane out of range"); - size_t chunkCount = getBatchChunkTargetCount(laneCount); + size_t chunkCount = getBatchChunkTargetCount(batch); size_t laneCountSize = static_cast(laneCount); size_t baseChunkSize = laneCountSize / chunkCount; size_t remainder = laneCountSize % chunkCount; @@ -56,12 +94,12 @@ size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) { } ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex) { - BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkIndex); + BatchChunkRange chunk = getBatchChunkRange(batch, chunkIndex); return {batch.getOperation(), chunk.laneStart, chunk.laneCount}; } ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane) { - return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch.getLaneCount(), lane)); + return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch, lane)); } llvm::SmallVector @@ -74,8 +112,8 @@ getBatchChunksForRange(SpatComputeBatch batch, uint32_t laneStart, uint32_t lane assert(laneEnd >= laneStart && "lane range overflow"); assert(laneEnd <= static_cast(batch.getLaneCount()) && "lane range out of bounds"); - size_t firstChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneStart); - size_t lastChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneEnd - 1); + size_t firstChunk = getBatchChunkIndexForLane(batch, laneStart); + size_t lastChunk = getBatchChunkIndexForLane(batch, laneEnd - 1); chunks.reserve(lastChunk - firstChunk + 1); for (size_t chunkIndex = firstChunk; chunkIndex <= lastChunk; ++chunkIndex) chunks.push_back(getBatchChunkForIndex(batch, chunkIndex)); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.hpp index 75702e3..ad556f0 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.hpp @@ -27,9 +27,9 @@ struct BatchChunkRange { }; size_t getSchedulingCpuBudget(); -size_t getBatchChunkTargetCount(int32_t laneCount); -BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex); -size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane); +size_t getBatchChunkTargetCount(SpatComputeBatch batch); +BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex); +size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane); ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex); ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane); llvm::SmallVector diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp index ca48f10..06693c1 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp @@ -6,7 +6,9 @@ #include #include +#include #include +#include #include #include "PeftScheduler.hpp" @@ -133,6 +135,55 @@ void verifyOctTableSize(size_t nodeCount, size_t processorCount) { } } +std::vector planCrossbarResidency(const ComputeGraph& graph, + size_t processorCount, + size_t crossbarCapacity, + const MeshModel& mesh) { + std::vector weightedTasks; + for (size_t task = 0; task < graph.nodes.size(); ++task) + if (!graph.nodes[task].crossbarUsage.empty()) + weightedTasks.push_back(task); + llvm::sort(weightedTasks, [&](size_t lhs, size_t rhs) { + if (graph.nodes[lhs].crossbarUsage.size() != graph.nodes[rhs].crossbarUsage.size()) + return graph.nodes[lhs].crossbarUsage.size() > graph.nodes[rhs].crossbarUsage.size(); + return graph.nodes[lhs].originalOrder < graph.nodes[rhs].originalOrder; + }); + + std::vector residency(processorCount); + for (size_t task : weightedTasks) { + size_t bestProcessor = std::numeric_limits::max(); + using ResidencyScore = std::tuple; + std::optional bestScore; + for (size_t processor = 0; processor < processorCount; ++processor) { + size_t crossbarUnion = getCrossbarUnionSize(residency[processor], graph.nodes[task].crossbarUsage); + if (crossbarUnion > crossbarCapacity) + continue; + size_t addedCrossbars = crossbarUnion - residency[processor].size(); + ResidencyScore score {addedCrossbars, + crossbarCapacity - crossbarUnion, + mesh.getCenterDistance(processor), + processor}; + if (!bestScore || score < *bestScore) { + bestProcessor = processor; + bestScore = score; + } + } + if (bestProcessor == std::numeric_limits::max()) { + std::string message = + llvm::formatv("PEFT residency planner: cannot place task {0} with {1} distinct weights in {2} " + "processors of capacity {3}", + graph.nodes[task].originalOrder, + graph.nodes[task].crossbarUsage.size(), + processorCount, + crossbarCapacity) + .str(); + llvm::report_fatal_error(llvm::StringRef(message)); + } + insertCrossbarWeights(residency[bestProcessor], graph.nodes[task].crossbarUsage); + } + return residency; +} + } // namespace Time getPeftTransferTime(Time transferCost, size_t sourceProcessor, size_t targetProcessor, size_t processorCount) { @@ -145,6 +196,8 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu if (processorCount == 0) llvm::report_fatal_error("PEFT scheduler: processor count must be positive"); MeshModel mesh = MeshModel::infer(processorCount); + std::vector plannedResidency = + planCrossbarResidency(graph, processorCount, options.crossbarCapacity, mesh); verifyOctTableSize(nodeCount, processorCount); std::vector> reverseLevels = buildReverseLevels(graph); @@ -237,20 +290,23 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu size_t bestProcessor = std::numeric_limits::max(); Time bestEst = 0; Time bestEft = 0; - Time bestOeft = std::numeric_limits