From 2b899b62a87be600b8efa08476602523972f73e9 Mon Sep 17 00:00:00 2001 From: ilgeco Date: Tue, 28 Jul 2026 12:09:59 +0200 Subject: [PATCH] Previous commit was broken in this one 9ms vs 7ms for vgg8 Arch-A --- .../ONNXToSpatial/LowerSpatialPlansPass.cpp | 20 +- .../ONNXToSpatial/Patterns/Math/Conv.cpp | 428 +++++++++++------- .../ONNXToSpatial/Patterns/NN/Pool.cpp | 55 ++- .../ONNXToSpatial/Patterns/Tensor/Flatten.cpp | 111 +++++ .../Conversion/ONNXToSpatial/PlanLowering.hpp | 8 + .../SpatialLayoutPlanningPass.cpp | 2 + .../Scheduling/ComputeGraph.cpp | 71 ++- .../Scheduling/ComputeInstanceUtils.cpp | 33 +- 8 files changed, 518 insertions(+), 210 deletions(-) diff --git a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp index 2264a7e..12f6727 100644 --- a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp @@ -169,7 +169,10 @@ struct LowerSpatialPlansPass final : PassWrappergetType(); signalPassFailure(); return; } @@ -336,6 +339,21 @@ struct LowerSpatialPlansPass final : PassWrapper(&op)) { + if (flattenOp.getInputs().size() == 1) { + FailureOr input = + getRowStripValue(rowStripValues, flattenOp.getInputs().front()); + if (succeeded(input) && succeeded(canLowerFlattenFromRowStrip(flattenOp))) { + rewriter.setInsertionPoint(flattenOp); + if (failed(lowerFlattenFromRowStrip(*input, flattenOp, rewriter))) { + flattenOp.emitOpError("failed to preserve row-strip layout through Flatten"); + signalPassFailure(); + return; + } + continue; + } + } + } if (auto materializeOp = dyn_cast(&op)) { if (materializeOp.getSourcePhysicalLayout() == kDenseLayout && materializeOp.getTargetPhysicalLayout() == kDenseLayout) { diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp index 9a67ebf..48952ad 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp @@ -2491,6 +2491,11 @@ static Value createZeroGemmBias(RankedTensorType gemmResultType, PatternRewriter return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), zeroAttr, gemmResultType); } +static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) { + return ceilIntegerDivide(geometry.k, geometry.xbarSize) + <= static_cast(crossbarCountInCore.getValue()); +} + static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state, StringRef& failureReason) { if (state.batchSize != 1) { failureReason = "batch_not_one"; @@ -2516,7 +2521,6 @@ static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state failureReason = "dilation_not_one"; return false; } - ConvGeometry geometry = buildConvGeometry(state); const bool pointwise = state.xHeight == 1 && state.xWidth == 1 && state.outHeight == 1 && state.outWidth == 1 && state.wHeight == 1 && state.wWidth == 1 && state.padHeightBegin == 0 && state.padHeightEnd == 0 && state.padWidthBegin == 0 && state.padWidthEnd == 0; @@ -2547,12 +2551,12 @@ static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state failureReason = "non_constant_weight"; return false; } - if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)) { - failureReason = "unsupported_bias"; + if (!rowStripOutputTileFitsOneCore(buildConvGeometry(state))) { + failureReason = "output_row_does_not_fit_one_core"; return false; } - if (geometry.c > geometry.xbarSize) { - failureReason = "output_channels_exceed_crossbar"; + if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)) { + failureReason = "unsupported_bias"; return false; } return true; @@ -2631,6 +2635,25 @@ static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state, return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType); } +static Value createRowStripWindowSourceSlotTable(const ConvLoweringState& state, + int64_t tilesPerRow, + PatternRewriter& rewriter) { + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + auto tableType = + RankedTensorType::get({state.outHeight * state.wHeight * tilesPerRow}, rewriter.getIndexType()); + SmallVector values; + 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 * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin; + sourceRow = std::clamp(sourceRow, int64_t {0}, state.xHeight - 1); + for (int64_t tile = 0; tile < tilesPerRow; ++tile) + values.push_back(rewriter.getIndexAttr(sourceRow * tilesPerRow + tile)); + } + return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType); +} + static Value createRowStripWindowTableIndex(Value outputHeight, Value kernelRow, const ConvLoweringState& state, @@ -2644,16 +2667,72 @@ static Value createRowStripWindowTableIndex(Value outputHeight, rewriter, loc, outputRowExpr * state.wHeight + kernelRowExpr, ValueRange {outputHeight, kernelRow}, anchorOp); } -static Value extractProjectedRowStripWindowRow(Value rowStripStorage, - Value sourceRowTable, - const ConvLoweringState& state, - Value outputHeight, - Value kernelRow, - PatternRewriter& rewriter, - Location loc) { +static FailureOr extractProjectedRowStripWindowRow(Value rowStripStorage, + Value sourceSlotTable, + const ConvLoweringState& state, + Value outputHeight, + Value kernelRow, + PatternRewriter& rewriter, + Location loc) { + FailureOr physical = describeRowStripPhysicalValue(rowStripStorage, state.xType); + if (failed(physical)) + return failure(); Value tableIndex = createRowStripWindowTableIndex(outputHeight, kernelRow, state, rewriter, loc); - Value sourceRow = tensor::ExtractOp::create(rewriter, loc, sourceRowTable, ValueRange {tableIndex}).getResult(); - return extractRowStripFragment(rowStripStorage, state.xType, sourceRow, rewriter, loc); + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + Value tileTableIndex = + affineMulConst(rewriter, loc, tableIndex, physical->tilesPerRow, anchorOp); + if (physical->tilesPerRow == 1) { + Value sourceSlot = + tensor::ExtractOp::create(rewriter, loc, sourceSlotTable, ValueRange {tileTableIndex}).getResult(); + return extractGraphBatchPhysicalFragment( + rewriter, loc, rowStripStorage, sourceSlot, physical->fragmentType); + } + + auto fullFragmentType = getRowStripFragmentType(state.xType); + Value fullFragment = tensor::EmptyOp::create( + rewriter, loc, fullFragmentType.getShape(), fullFragmentType.getElementType()); + const int64_t tileChannels = physical->fragmentType.getDimSize(3); + for (int64_t tile = 0; tile < physical->tilesPerRow; ++tile) { + Value slotTableIndex = affineAddConst(rewriter, loc, tileTableIndex, tile, anchorOp); + Value tileSlot = + tensor::ExtractOp::create(rewriter, loc, sourceSlotTable, ValueRange {slotTableIndex}).getResult(); + FailureOr fragment = extractGraphBatchPhysicalFragment( + rewriter, loc, rowStripStorage, tileSlot, physical->fragmentType); + if (failed(fragment)) + return failure(); + const int64_t channelOffset = tile * tileChannels; + const int64_t validChannels = std::min(tileChannels, state.numChannelsIn - channelOffset); + auto validType = RankedTensorType::get( + {1, 1, state.xWidth, validChannels}, state.xType.getElementType(), state.xType.getEncoding()); + Value validFragment = *fragment; + if (validChannels != tileChannels) + validFragment = tensor::ExtractSliceOp::create( + rewriter, + loc, + validType, + *fragment, + SmallVector(4, rewriter.getIndexAttr(0)), + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.xWidth), + rewriter.getIndexAttr(validChannels)}, + getUnitStrides(rewriter, 4)); + fullFragment = tensor::InsertSliceOp::create( + rewriter, + loc, + validFragment, + fullFragment, + SmallVector {rewriter.getIndexAttr(0), + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(0), + rewriter.getIndexAttr(channelOffset)}, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.xWidth), + rewriter.getIndexAttr(validChannels)}, + getUnitStrides(rewriter, 4)); + } + return fullFragment; } static Value extractDenseConvWindowRow(Value denseInput, @@ -2739,13 +2818,17 @@ static FailureOr createConvInputWindow(Value input, auto fragmentType = getRowStripFragmentType(state.xType); auto inputType = dyn_cast(input.getType()); const bool denseInput = inputType == state.xType; - if (!denseInput && inputType != getRowStripStorageType(state.xType)) + if (!denseInput && failed(describeRowStripPhysicalValue(input, state.xType))) return failure(); auto paddedWindowType = RankedTensorType::get( {1, state.wHeight, state.xWidth + state.padWidthBegin + state.padWidthEnd, state.numChannelsIn}, state.xType.getElementType(), state.xType.getEncoding()); - Value sourceRowTable = createRowStripWindowSourceRowTable(state, rewriter); + FailureOr physicalInput = + denseInput ? FailureOr(failure()) : describeRowStripPhysicalValue(input, state.xType); + Value sourceIndexTable = + denseInput ? createRowStripWindowSourceRowTable(state, rewriter) + : createRowStripWindowSourceSlotTable(state, physicalInput->tilesPerRow, rewriter); FailureOr maskTable = createRowStripWindowMaskTable(state, rewriter); if (failed(maskTable)) return failure(); @@ -2755,14 +2838,17 @@ static FailureOr createConvInputWindow(Value input, Value window = initWindow; for (int64_t kernelRowIndex = 0; kernelRowIndex < state.wHeight; ++kernelRowIndex) { Value kernelRow = getOrCreateIndexConstant(rewriter, anchorOp, kernelRowIndex); - Value sourceRow = denseInput - ? extractDenseConvWindowRow(input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc) - : extractProjectedRowStripWindowRow( - input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc); - Value semanticRow = sourceRow; + FailureOr sourceRow = + denseInput + ? FailureOr( + extractDenseConvWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, loc)) + : extractProjectedRowStripWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, loc); + if (failed(sourceRow)) + return failure(); + Value semanticRow = *sourceRow; if (state.padHeightBegin != 0 || state.padHeightEnd != 0) { Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, loc); - semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, sourceRow, mask).getResult(); + semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, semanticRow, mask).getResult(); } Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, loc); window = tensor::InsertSliceOp::create(rewriter, @@ -2814,7 +2900,7 @@ static FailureOr createPixelMajorConvPatchRow(Value paddedWindow, } static FailureOr createConvOutputTile(Value patchRow, - Value partialInputScratch, + Value& partialInputScratch, Value tileWeights, int64_t patchSize, int64_t numKSlices, @@ -2853,7 +2939,7 @@ static FailureOr createConvOutputTile(Value patchRow, SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)}, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)}, getUnitStrides(rewriter, 2)); - inputTile = tensor::InsertSliceOp::create( + partialInputScratch = tensor::InsertSliceOp::create( rewriter, loc, partial, @@ -2861,6 +2947,7 @@ static FailureOr createConvOutputTile(Value patchRow, SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)}, getUnitStrides(rewriter, 2)); + inputTile = partialInputScratch; } SmallVector bOffsets { rewriter.getIndexAttr(kOffset), rewriter.getIndexAttr(0)}; @@ -2878,7 +2965,7 @@ static FailureOr createConvOutputTile(Value patchRow, } static FailureOr createConvOutputRow(Value patchRow, - Value partialInputScratch, + Value& partialInputScratch, int64_t patchSize, int64_t paddedK, int64_t outputChannels, @@ -2943,17 +3030,15 @@ static FailureOr createConvOutputRow(Value patchRow, SmallVector outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; SmallVector outputSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(outputChannels)}; - return tensor::ExtractSliceOp::create( - rewriter, loc, rowType, paddedOutput, outputOffsets, outputSizes, getUnitStrides(rewriter, 2)) - .getResult(); -} - -static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) { - return ceilIntegerDivide(geometry.k, geometry.xbarSize) - <= static_cast(crossbarCountInCore.getValue()); + Value validRow = tensor::ExtractSliceOp::create( + rewriter, loc, rowType, paddedOutput, outputOffsets, outputSizes, getUnitStrides(rewriter, 2)); + if (bias) + validRow = spatial::SpatVAddOp::create(rewriter, loc, rowType, validRow, bias).getResult(); + return validRow; } static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLoweringState& state, + Value input, Value paddedWeights, int64_t paddedK, int64_t numKSlices, @@ -2962,8 +3047,9 @@ static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLow Location loc) { const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim); const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + const bool hasPartialInputTile = patchSize % xbarDim != 0; auto elementType = state.outType.getElementType(); - auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType); + auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType); auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType); auto tilePixelType = RankedTensorType::get({1, 1, 1, xbarDim}, elementType); auto tileFragmentType = RankedTensorType::get({1, 1, state.outWidth, xbarDim}, elementType); @@ -2975,101 +3061,102 @@ static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLow paddedBias = createPaddedBiasTileConstant(state, xbarDim, rewriter); if (state.hasBias && failed(paddedBias)) return failure(); + auto tileBatch = createSpatComputeBatch( - rewriter, loc, TypeRange {tileStorageType}, laneCount, ValueRange {paddedWeights}, - state.hasBias ? ValueRange {state.x, *paddedBias} : ValueRange {state.x}, + rewriter, + loc, + TypeRange {tileStorageType}, + laneCount, + ValueRange {paddedWeights}, + state.hasBias ? ValueRange {input, *paddedBias} : ValueRange {input}, [&](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 outputRow = affineFloorDivConst(rewriter, loc, args.lane, outputTileCount, anchorOp); - Value outputTile = affineModConst(rewriter, loc, args.lane, outputTileCount, anchorOp); - SmallVector weightOffsets { - outputTile, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; - SmallVector weightSizes { - rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)}; - Value tileWeights = tensor::ExtractSliceOp::create( - rewriter, loc, tileWeightsType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3)); - FailureOr biasTile = failure(); - if (state.hasBias) - biasTile = extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[1], outputTile, paddedRowType); - if (state.hasBias && failed(biasTile)) - return failure(); - FailureOr inputWindow = - createConvInputWindow(args.inputs.front(), state, outputRow, rewriter, loc); - if (failed(inputWindow)) - return failure(); - Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType); - SmallVector widthLoopInit {fragmentInit}; - if (patchSize != paddedK) - widthLoopInit.push_back(createZeroTensorConstant(paddedPatchRowType, rewriter)); - auto widthLoop = buildNormalizedScfFor( - rewriter, - loc, - c0, - cOutWidth, - c1, - widthLoopInit, - [&](OpBuilder&, - Location widthLoc, - Value widthIndex, - ValueRange widthIterArgs, - SmallVectorImpl& widthYielded) { - FailureOr patchRow = - createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc); - if (failed(patchRow)) - return failure(); - Value paddedPatchRow = *patchRow; - if (patchSize != paddedK) - paddedPatchRow = tensor::InsertSliceOp::create( - rewriter, - widthLoc, - paddedPatchRow, - widthIterArgs[1], - SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, - SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(patchSize)}, - getUnitStrides(rewriter, 2)); - FailureOr paddedOutputRow = createPaddedConvOutputTile( - paddedPatchRow, tileWeights, numKSlices, xbarDim, rewriter, widthLoc); - if (failed(paddedOutputRow)) - return failure(); - if (state.hasBias) - paddedOutputRow = spatial::SpatVAddOp::create( - rewriter, widthLoc, paddedRowType, *paddedOutputRow, *biasTile).getResult(); - Value outputPixel = tensor::ExpandShapeOp::create( - rewriter, widthLoc, tilePixelType, *paddedOutputRow, SmallVector {{0, 1, 2}, {3}}); - SmallVector rowOffsets { - rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), widthIndex, rewriter.getIndexAttr(0)}; - SmallVector rowSizes {rewriter.getIndexAttr(1), - rewriter.getIndexAttr(1), - rewriter.getIndexAttr(1), - rewriter.getIndexAttr(xbarDim)}; - Value nextFragment = tensor::InsertSliceOp::create(rewriter, - widthLoc, - outputPixel, - widthIterArgs.front(), - rowOffsets, - rowSizes, - getUnitStrides(rewriter, 4)); - widthYielded.push_back(nextFragment); - if (patchSize != paddedK) - widthYielded.push_back(paddedPatchRow); - return success(); - }); - if (failed(widthLoop)) - return failure(); - publishGraphBatchPhysicalFragment( - rewriter, loc, widthLoop->results.front(), args.outputs.front(), args.lane); - return success(); - }); + 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 outputRow = affineFloorDivConst(rewriter, loc, args.lane, outputTileCount, anchorOp); + Value outputTile = affineModConst(rewriter, loc, args.lane, outputTileCount, anchorOp); + SmallVector weightOffsets { + outputTile, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; + SmallVector weightSizes { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)}; + Value tileWeights = tensor::ExtractSliceOp::create( + rewriter, loc, tileWeightsType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3)); + FailureOr biasTile = failure(); + if (state.hasBias) + biasTile = extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[1], outputTile, paddedRowType); + if (state.hasBias && failed(biasTile)) + return failure(); + FailureOr inputWindow = + createConvInputWindow(args.inputs.front(), state, outputRow, rewriter, loc); + if (failed(inputWindow)) + return failure(); + Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType); + SmallVector widthLoopInit {fragmentInit}; + if (hasPartialInputTile) + widthLoopInit.push_back(createZeroTensorConstant(partialInputScratchType, rewriter)); + auto widthLoop = buildNormalizedScfFor( + rewriter, + loc, + c0, + cOutWidth, + c1, + widthLoopInit, + [&](OpBuilder&, + Location widthLoc, + Value widthIndex, + ValueRange widthIterArgs, + SmallVectorImpl& widthYielded) { + FailureOr patchRow = + createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc); + if (failed(patchRow)) + return failure(); + Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value(); + FailureOr paddedOutputRow = createConvOutputTile(*patchRow, + partialInputScratch, + tileWeights, + patchSize, + numKSlices, + xbarDim, + rewriter, + widthLoc); + if (failed(paddedOutputRow)) + return failure(); + if (state.hasBias) + paddedOutputRow = + spatial::SpatVAddOp::create(rewriter, widthLoc, paddedRowType, *paddedOutputRow, *biasTile).getResult(); + Value outputPixel = tensor::ExpandShapeOp::create( + rewriter, widthLoc, tilePixelType, *paddedOutputRow, SmallVector {{0, 1, 2}, {3}}); + SmallVector rowOffsets { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), widthIndex, rewriter.getIndexAttr(0)}; + SmallVector rowSizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(xbarDim)}; + Value nextFragment = tensor::InsertSliceOp::create(rewriter, + widthLoc, + outputPixel, + widthIterArgs.front(), + rowOffsets, + rowSizes, + getUnitStrides(rewriter, 4)); + widthYielded.push_back(nextFragment); + if (hasPartialInputTile) + widthYielded.push_back(partialInputScratch); + 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(); return tileBatch->getResult(0); } static FailureOr -createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { + createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { ConvGeometry geometry = buildConvGeometry(state); if (state.group != 1 || state.batchSize != 1 || !rowStripOutputTileFitsOneCore(geometry)) return failure(); @@ -3084,8 +3171,9 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; const int64_t numKSlices = ceilIntegerDivide(patchSize, xbarDim); const int64_t paddedK = numKSlices * xbarDim; + const bool hasPartialInputTile = patchSize % xbarDim != 0; auto elementType = state.outType.getElementType(); - auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType); + auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType); auto fragmentType = getRowStripFragmentType(state.outType); auto outputPixelType = RankedTensorType::get({1, 1, 1, state.numChannelsOut}, elementType); auto outputStorageType = getRowStripStorageType(state.outType); @@ -3097,7 +3185,7 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe weightDenseAttr, state, paddedK, xbarDim, rewriter); if (state.numChannelsOut > xbarDim) return createOutputChannelTiledRowStripConvOutput( - state, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc); + state, state.x, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc); FailureOr bias = failure(); if (state.hasBias) @@ -3123,8 +3211,8 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe return failure(); Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType); SmallVector widthLoopInit {fragmentInit}; - if (patchSize != paddedK) - widthLoopInit.push_back(createZeroTensorConstant(paddedPatchRowType, rewriter)); + if (hasPartialInputTile) + widthLoopInit.push_back(createZeroTensorConstant(partialInputScratchType, rewriter)); auto widthLoop = buildNormalizedScfFor( rewriter, loc, @@ -3137,24 +3225,18 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc); if (failed(patchRow)) return failure(); - Value paddedPatchRow = *patchRow; - if (patchSize != paddedK) - paddedPatchRow = tensor::InsertSliceOp::create( - rewriter, - widthLoc, - paddedPatchRow, - widthIterArgs[1], - SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, - SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(patchSize)}, - getUnitStrides(rewriter, 2)); - FailureOr outputRow = createPaddedConvOutputRow(paddedPatchRow, - state.numChannelsOut, - args.weights.front(), - state.hasBias ? args.inputs[1] : Value(), - numKSlices, - xbarDim, - rewriter, - widthLoc); + Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value(); + FailureOr outputRow = createConvOutputRow(*patchRow, + partialInputScratch, + patchSize, + paddedK, + state.numChannelsOut, + args.weights.front(), + state.hasBias ? args.inputs[1] : Value(), + numKSlices, + xbarDim, + rewriter, + widthLoc); if (failed(outputRow)) return failure(); @@ -3171,8 +3253,8 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe Value nextFragment = tensor::InsertSliceOp::create( rewriter, widthLoc, outputFragment, widthIterArgs.front(), rowOffsets, rowSizes, getUnitStrides(rewriter, 4)); widthYielded.push_back(nextFragment); - if (patchSize != paddedK) - widthYielded.push_back(paddedPatchRow); + if (hasPartialInputTile) + widthYielded.push_back(partialInputScratch); return success(); }); if (failed(widthLoop)) @@ -3190,8 +3272,7 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { - auto inputType = dyn_cast(rowStripStorage.getType()); - if (!inputType || inputType != getRowStripStorageType(state.xType)) + if (failed(describeRowStripPhysicalValue(rowStripStorage, state.xType))) return failure(); StringRef failureReason; @@ -3203,15 +3284,22 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; const int64_t numKSlices = ceilIntegerDivide(patchSize, xbarDim); const int64_t paddedK = numKSlices * xbarDim; + const bool hasPartialInputTile = patchSize % xbarDim != 0; auto elementType = state.outType.getElementType(); - auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType); + auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType); auto outputPixelType = RankedTensorType::get({1, 1, 1, state.numChannelsOut}, elementType); auto outputStorageType = getRowStripStorageType(state.outType); auto weightDenseAttr = getHostConstDenseElementsAttr(state.w); if (!weightDenseAttr) return failure(); - Value paddedWeights = - standard::createPaddedPixelMajorWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter); + Value paddedWeights = state.numChannelsOut <= xbarDim + ? standard::createPaddedPixelMajorWeightConstant( + weightDenseAttr, state, paddedK, xbarDim, rewriter) + : standard::createPaddedOutputChannelTiledWeightConstant( + weightDenseAttr, state, paddedK, xbarDim, rewriter); + if (state.numChannelsOut > xbarDim) + return createOutputChannelTiledRowStripConvOutput( + state, rowStripStorage, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc); FailureOr bias = failure(); if (state.hasBias) bias = createBiasRowConstant(state, rewriter); @@ -3236,8 +3324,8 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro return failure(); Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType); SmallVector widthLoopInit {fragmentInit}; - if (patchSize != paddedK) - widthLoopInit.push_back(createZeroTensorConstant(paddedPatchRowType, rewriter)); + if (hasPartialInputTile) + widthLoopInit.push_back(createZeroTensorConstant(partialInputScratchType, rewriter)); auto widthLoop = buildNormalizedScfFor( rewriter, loc, @@ -3251,24 +3339,18 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro if (failed(patchRow)) return failure(); - Value paddedPatchRow = *patchRow; - if (patchSize != paddedK) - paddedPatchRow = tensor::InsertSliceOp::create( - rewriter, - widthLoc, - paddedPatchRow, - widthIterArgs[1], - SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, - SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(patchSize)}, - getUnitStrides(rewriter, 2)); - FailureOr outputRow = createPaddedConvOutputRow(paddedPatchRow, - state.numChannelsOut, - args.weights.front(), - state.hasBias ? args.inputs[1] : Value(), - numKSlices, - xbarDim, - rewriter, - widthLoc); + Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value(); + FailureOr outputRow = createConvOutputRow(*patchRow, + partialInputScratch, + patchSize, + paddedK, + state.numChannelsOut, + args.weights.front(), + state.hasBias ? args.inputs[1] : Value(), + numKSlices, + xbarDim, + rewriter, + widthLoc); if (failed(outputRow)) return failure(); @@ -3285,8 +3367,8 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro Value nextFragment = tensor::InsertSliceOp::create( rewriter, widthLoc, outputFragment, widthIterArgs.front(), rowOffsets, rowSizes, getUnitStrides(rewriter, 4)); widthYielded.push_back(nextFragment); - if (patchSize != paddedK) - widthYielded.push_back(paddedPatchRow); + if (hasPartialInputTile) + widthYielded.push_back(partialInputScratch); return success(); }); if (failed(widthLoop)) @@ -4243,6 +4325,7 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) { switch (decision.strategy) { case PimConvLoweringLegacy: + case PimConvLoweringDepthwise: case PimConvLoweringPackedIm2Col: case PimConvLoweringStreamedPatch: case PimConvLoweringOutputChannelTiled: @@ -4250,7 +4333,6 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) { case PimConvLoweringStreamedPacked: return success(); case PimConvLoweringAuto: - case PimConvLoweringDepthwise: case PimConvLoweringInputKTiled: return failure(); } diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp index 3c41c2e..1876149 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp @@ -448,6 +448,29 @@ static Value createClampedPoolIndexTable(PatternRewriter& rewriter, return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType); } +static Value createClampedPoolRowSlotTable(PatternRewriter& rewriter, + Operation* anchorOp, + int64_t outputHeight, + int64_t kernelHeight, + int64_t stride, + int64_t dilation, + int64_t padBegin, + int64_t inputHeight, + int64_t tilesPerRow) { + auto tableType = + RankedTensorType::get({outputHeight * tilesPerRow * kernelHeight}, rewriter.getIndexType()); + SmallVector values; + values.reserve(tableType.getNumElements()); + for (int64_t outputRow = 0; outputRow < outputHeight; ++outputRow) + for (int64_t tile = 0; tile < tilesPerRow; ++tile) + for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) { + const int64_t sourceRow = + std::clamp(outputRow * stride + kernelRow * dilation - padBegin, int64_t {0}, inputHeight - 1); + values.push_back(rewriter.getIndexAttr(sourceRow * tilesPerRow + tile)); + } + return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType); +} + static Value extractPoolIndex(PatternRewriter& rewriter, Location loc, Operation* anchorOp, @@ -497,6 +520,15 @@ FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, planOp.getDilations()[0], planOp.getPads()[0], inputHeight); + Value rowSlotTable = createClampedPoolRowSlotTable(rewriter, + anchorOp, + outputHeight, + kernelHeight, + planOp.getStrides()[0], + planOp.getDilations()[0], + planOp.getPads()[0], + inputHeight, + tilesPerRow); Value columnTable = createClampedPoolIndexTable(rewriter, anchorOp, outputWidth, @@ -524,27 +556,10 @@ FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, [&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult { SmallVector inputRows; inputRows.reserve(kernelHeight); - Value outputRow = tilesPerRow == 1 - ? args.lane - : affineFloorDivConst(rewriter, loc, args.lane, tilesPerRow, anchorOp); - Value channelTile = tilesPerRow == 1 - ? getOrCreateIndexConstant(rewriter, anchorOp, 0) - : affineModConst(rewriter, loc, args.lane, tilesPerRow, anchorOp); for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) { - Value sourceRow = - extractPoolIndex(rewriter, loc, anchorOp, rowTable, outputRow, kernelRow, kernelHeight); if (physicalInput) { - Value sourceSlot = sourceRow; - if (tilesPerRow != 1) { - sourceSlot = arith::AddIOp::create( - rewriter, - loc, - arith::MulIOp::create(rewriter, - loc, - sourceRow, - getOrCreateIndexConstant(rewriter, anchorOp, tilesPerRow)), - channelTile); - } + Value sourceSlot = + extractPoolIndex(rewriter, loc, anchorOp, rowSlotTable, args.lane, kernelRow, kernelHeight); FailureOr fragment = extractGraphBatchPhysicalFragment( rewriter, loc, args.inputs.front(), sourceSlot, inputFragmentType); if (failed(fragment)) @@ -552,6 +567,8 @@ FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, inputRows.push_back(*fragment); } else { + Value sourceRow = + extractPoolIndex(rewriter, loc, anchorOp, rowTable, args.lane, kernelRow, kernelHeight); SmallVector offsets { rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)}; Value nchw = tensor::ExtractSliceOp::create(rewriter, diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp index 621558d..9fdb01b 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp @@ -3,7 +3,12 @@ #include "llvm/ADT/SmallVector.h" +#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp" +#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Dialect/ONNX/ONNXOps.hpp" @@ -105,8 +110,114 @@ struct Flatten : OpConversionPattern { } }; +struct RowStripFlattenAnalysis { + spatial::SpatGraphComputeBatch consumer; + tensor::CollapseShapeOp collapse; + RankedTensorType sourceType; + RankedTensorType resultType; + RankedTensorType weightType; + DenseElementsAttr weight; +}; + +static FailureOr analyzeRowStripFlatten(spatial::SpatGraphCompute flattenOp) { + if (flattenOp.getWeights().size() != 0 || flattenOp.getInputs().size() != 1 + || flattenOp.getOutputs().size() != 1) + return failure(); + auto sourceType = dyn_cast(flattenOp.getInputs().front().getType()); + auto resultType = dyn_cast(flattenOp.getOutputs().front().getType()); + if (!sourceType || !resultType || !sourceType.hasStaticShape() || !resultType.hasStaticShape() + || sourceType.getRank() != 4 || resultType.getRank() != 2 || sourceType.getDimSize(0) != 1 + || resultType.getDimSize(0) != 1 || resultType.getDimSize(1) != sourceType.getNumElements()) + return failure(); + const int64_t channels = sourceType.getDimSize(1); + const int64_t xbarDim = static_cast(crossbarSize.getValue()); + if (channels > xbarDim && channels % xbarDim != 0) + return failure(); + + auto yieldOp = dyn_cast(flattenOp.getBody().front().getTerminator()); + if (!yieldOp || yieldOp.getOutputs().size() != 1) + return failure(); + auto collapse = yieldOp.getOutputs().front().getDefiningOp(); + if (!collapse || collapse.getSrc() != *flattenOp.getInputArgument(0)) + return failure(); + + if (!flattenOp.getResult(0).hasOneUse()) + return failure(); + auto consumer = dyn_cast(*flattenOp.getResult(0).getUsers().begin()); + if (!consumer || consumer.getInputs().size() != 1 || consumer.getInputs().front() != flattenOp.getResult(0) + || consumer.getWeights().size() != 1) + return failure(); + auto weightType = dyn_cast(consumer.getWeights().front().getType()); + DenseElementsAttr weight = getHostConstDenseElementsAttr(consumer.getWeights().front()); + if (!weightType || !weight || !weightType.hasStaticShape() || weightType.getRank() != 2 + || weightType.getDimSize(0) != resultType.getDimSize(1)) + return failure(); + if (llvm::none_of(consumer.getBody().getOps(), + [](spatial::SpatVMMOp) { return true; })) + return failure(); + return RowStripFlattenAnalysis {consumer, collapse, sourceType, resultType, weightType, weight}; +} + } // namespace void populateFlattenPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.add(ctx); } +LogicalResult canLowerFlattenFromRowStrip(spatial::SpatGraphCompute flattenOp) { + return succeeded(analyzeRowStripFlatten(flattenOp)) ? success() : failure(); +} + +LogicalResult lowerFlattenFromRowStrip(const RowStripPhysicalValue& input, + spatial::SpatGraphCompute flattenOp, + PatternRewriter& rewriter) { + FailureOr analysis = analyzeRowStripFlatten(flattenOp); + if (failed(analysis)) + return failure(); + auto storageType = dyn_cast(input.storage.getType()); + if (!storageType || storageType.getNumElements() != analysis->resultType.getNumElements()) + return failure(); + + const int64_t channels = input.logicalType.getDimSize(1); + const int64_t height = input.logicalType.getDimSize(2); + const int64_t width = input.logicalType.getDimSize(3); + const int64_t tileChannels = input.fragmentType.getDimSize(3); + const int64_t outputColumns = analysis->weightType.getDimSize(1); + SmallVector sourceValues(analysis->weight.getValues()); + SmallVector reorderedValues(sourceValues.size()); + for (int64_t row = 0; row < height; ++row) + for (int64_t tile = 0; tile < input.tilesPerRow; ++tile) + for (int64_t column = 0; column < width; ++column) + for (int64_t channelInTile = 0; channelInTile < tileChannels; ++channelInTile) { + const int64_t channel = tile * tileChannels + channelInTile; + if (channel >= channels) + return failure(); + const int64_t physicalRow = + ((row * input.tilesPerRow + tile) * width + column) * tileChannels + channelInTile; + const int64_t logicalRow = (channel * height + row) * width + column; + for (int64_t output = 0; output < outputColumns; ++output) + reorderedValues[physicalRow * outputColumns + output] = + sourceValues[logicalRow * outputColumns + output]; + } + + Value reorderedWeight = getOrCreateConstant(rewriter, + rewriter.getInsertionBlock()->getParentOp(), + DenseElementsAttr::get(analysis->weightType, reorderedValues), + analysis->weightType); + analysis->consumer->setOperand(0, reorderedWeight); + + BlockArgument flattenInput = *flattenOp.getInputArgument(0); + flattenOp.getInputsMutable().assign(input.storage); + flattenInput.setType(storageType); + + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(analysis->collapse); + auto flatType = RankedTensorType::get( + {storageType.getNumElements()}, storageType.getElementType(), storageType.getEncoding()); + Value flat = tensor::CollapseShapeOp::create( + rewriter, flattenOp.getLoc(), flatType, flattenInput, getCollapseTo1DReassociation(storageType.getRank())); + Value logicalInput = tensor::ExpandShapeOp::create( + rewriter, flattenOp.getLoc(), analysis->resultType, flat, getExpandFrom1DReassociation(2)); + rewriter.replaceOp(analysis->collapse, logicalInput); + return success(); +} + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp index 100500f..f24b421 100644 --- a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp @@ -9,6 +9,8 @@ namespace onnx_mlir { +struct RowStripPhysicalValue; + mlir::FailureOr lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp, std::optional rowStripInput, @@ -25,4 +27,10 @@ lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, std::optional rowStripInput, mlir::PatternRewriter& rewriter); +mlir::LogicalResult canLowerFlattenFromRowStrip(spatial::SpatGraphCompute flattenOp); + +mlir::LogicalResult lowerFlattenFromRowStrip(const RowStripPhysicalValue& input, + spatial::SpatGraphCompute flattenOp, + mlir::PatternRewriter& rewriter); + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp index 1e2443b..57d75e4 100644 --- a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp @@ -40,6 +40,8 @@ static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap(user)) return getSelectedLayout(layouts, maxPoolPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; + if (auto flattenCompute = dyn_cast(user)) + return succeeded(canLowerFlattenFromRowStrip(flattenCompute)); return false; } diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp index 5f7dc58..66635a7 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.cpp @@ -379,6 +379,53 @@ std::optional getBatchProjectedInputTransferCost(SpatComputeBatch batch, V return projectedCost; } +static std::optional> +collectProjectedProducerValueRefs(SpatComputeBatch producer, + Value input, + const ComputeInstance& consumerInstance) { + auto consumer = dyn_cast(consumerInstance.op); + if (!consumer) + return std::nullopt; + auto inputIt = llvm::find(consumer.getInputs(), input); + if (inputIt == consumer.getInputs().end()) + return std::nullopt; + size_t inputIndex = std::distance(consumer.getInputs().begin(), inputIt); + std::optional inputArg = consumer.getInputArgument(inputIndex); + std::optional laneArg = consumer.getLaneArgument(); + if (!inputArg || !laneArg) + return std::nullopt; + + SmallVector producers; + DenseMap bindings; + for (Operation* user : inputArg->getUsers()) { + auto extract = dyn_cast(user); + if (!extract || extract.getSource() != *inputArg || extract.getMixedOffsets().empty() + || extract.getMixedSizes().empty() || extract.getMixedStrides().empty()) + return std::nullopt; + for (uint32_t lane = consumerInstance.laneStart; + lane < consumerInstance.laneStart + consumerInstance.laneCount; + ++lane) { + FailureOr offset = + evaluateIndexLike(extract.getMixedOffsets().front(), bindings, lane, *laneArg); + FailureOr size = + evaluateIndexLike(extract.getMixedSizes().front(), bindings, lane, *laneArg); + FailureOr stride = + evaluateIndexLike(extract.getMixedStrides().front(), bindings, lane, *laneArg); + if (failed(offset) || failed(size) || failed(stride) || *offset < 0 || *size <= 0 || *stride <= 0) + return std::nullopt; + for (int64_t index = 0; index < *size; ++index) { + int64_t producerLane = *offset + index * *stride; + if (producerLane < 0 || producerLane >= producer.getLaneCount()) + return std::nullopt; + ComputeInstance instance = getBatchChunkForLane(producer, static_cast(producerLane)); + if (llvm::none_of(producers, [&](const ProducerValueRef& ref) { return ref.instance == instance; })) + producers.push_back({instance, 0}); + } + } + } + return producers; +} + Cost getInputTransferCost(const ComputeInstance& consumerInstance, Value input) { auto inputType = cast(input.getType()); if (auto batch = dyn_cast(consumerInstance.op)) @@ -419,13 +466,9 @@ SmallVector collectProducerValueRefs(Value value, const Com return producers; } - if (isa(consumerInstance.op)) - for (ComputeInstance instance : - getBatchChunksForRange(batch, consumerInstance.laneStart, consumerInstance.laneCount)) - producers.push_back({instance, 0}); - else - for (ComputeInstance instance : getBatchChunksForRange(batch, 0, static_cast(batch.getLaneCount()))) - producers.push_back({instance, 0}); + for (ComputeInstance instance : + getBatchChunksForRange(batch, 0, static_cast(batch.getLaneCount()))) + producers.push_back({instance, 0}); return producers; } @@ -445,10 +488,13 @@ SmallVector collectProducerValueRefs(Value value, const Com if (auto batch = dyn_cast(op)) { if (batch.getNumResults() != 0) { - uint32_t laneStart = isa(consumerInstance.op) ? consumerInstance.laneStart : 0; - uint32_t laneCount = isa(consumerInstance.op) ? consumerInstance.laneCount - : static_cast(batch.getLaneCount()); - for (ComputeInstance instance : getBatchChunksForRange(batch, laneStart, laneCount)) + if (auto projected = collectProjectedProducerValueRefs(batch, value, consumerInstance)) + return *projected; + std::optional producer = getProducerValueRef(value, &consumerInstance); + if (!producer) + return producers; + for (ComputeInstance instance : + getBatchChunksForRange(batch, producer->instance.laneStart, producer->instance.laneCount)) producers.push_back({instance, 0}); return producers; } @@ -473,8 +519,7 @@ Cost getProducerTransferCost(Value input, if (auto consumerBatch = dyn_cast(consumerInstance.op)) { if (std::optional projectedCost = getBatchProjectedInputTransferCost(consumerBatch, input)) { uint32_t overlapLaneCount = getLaneOverlapCount(consumerInstance, producerRef.instance); - assert(overlapLaneCount > 0 && "projected batch edge must overlap consumer lanes"); - return checkedMultiply(*projectedCost, static_cast(overlapLaneCount)); + return checkedMultiply(*projectedCost, static_cast(std::max(1, overlapLaneCount))); } } diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp index 9012c16..1e3906e 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeInstanceUtils.cpp @@ -110,15 +110,40 @@ static std::optional getConstantExtractLane(tensor::ExtractSliceOp ext return std::nullopt; } +static bool hasNonLaneAlignedBatchProjection(SpatComputeBatch consumer, Value input) { + auto inputIt = llvm::find(consumer.getInputs(), input); + if (inputIt == consumer.getInputs().end()) + return false; + size_t inputIndex = std::distance(consumer.getInputs().begin(), inputIt); + std::optional inputArg = consumer.getInputArgument(inputIndex); + std::optional laneArg = consumer.getLaneArgument(); + if (!inputArg || !laneArg) + return true; + for (Operation* user : inputArg->getUsers()) { + auto extract = dyn_cast(user); + if (!extract || extract.getSource() != *inputArg || extract.getMixedOffsets().empty()) + return true; + auto offset = dyn_cast(extract.getMixedOffsets().front()); + if (!offset || offset != *laneArg) + return true; + } + return false; +} + static std::optional getResultfulBatchProducerValueRef(SpatComputeBatch batch, + Value value, const ComputeInstance* consumerInstance) { if (!consumerInstance || !isa(consumerInstance->op)) return ProducerValueRef { {batch.getOperation(), 0, static_cast(batch.getLaneCount())}, 0 }; - if (consumerInstance->laneStart + consumerInstance->laneCount > static_cast(batch.getLaneCount())) - return std::nullopt; + auto consumer = cast(consumerInstance->op); + if (consumer.getLaneCount() != batch.getLaneCount() || hasNonLaneAlignedBatchProjection(consumer, value)) + return ProducerValueRef { + {batch.getOperation(), 0, static_cast(batch.getLaneCount())}, + 0 + }; return ProducerValueRef { {batch.getOperation(), consumerInstance->laneStart, consumerInstance->laneCount}, 0 @@ -142,7 +167,7 @@ std::optional getProducerValueRef(Value value, const ComputeIn 0 }; } - return getResultfulBatchProducerValueRef(batch, consumerInstance); + return getResultfulBatchProducerValueRef(batch, source, consumerInstance); } value = source; @@ -160,7 +185,7 @@ std::optional getProducerValueRef(Value value, const ComputeIn if (auto batch = dyn_cast(op)) { if (batch.getNumResults() != 0) - return getResultfulBatchProducerValueRef(batch, consumerInstance); + return getResultfulBatchProducerValueRef(batch, value, consumerInstance); uint32_t lane = cast(value).getResultNumber(); ComputeInstance instance = getBatchChunkForLane(batch, lane); size_t resultIndex = lane - instance.laneStart;