From 78bfb8a9aac3135276602dfaf1fe9a6b7a391383 Mon Sep 17 00:00:00 2001 From: ilgeco Date: Tue, 28 Jul 2026 12:46:46 +0200 Subject: [PATCH] vgg8 6.88 vs 7.89 --- .../Conversion/ONNXToSpatial/CompileTime.cpp | 11 + .../Conversion/ONNXToSpatial/CompileTime.hpp | 2 + .../ONNXToSpatial/Patterns/Math/Conv.cpp | 190 ++++++++++++------ .../ONNXToSpatial/Patterns/Math/Gemm.cpp | 2 +- .../Bufferization/PimBufferizationPass.cpp | 6 + .../HostConstantFolding/Patterns/Constant.cpp | 92 ++++++++- 6 files changed, 239 insertions(+), 64 deletions(-) diff --git a/src/PIM/Conversion/ONNXToSpatial/CompileTime.cpp b/src/PIM/Conversion/ONNXToSpatial/CompileTime.cpp index 5fa76d7..3d4074f 100644 --- a/src/PIM/Conversion/ONNXToSpatial/CompileTime.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/CompileTime.cpp @@ -336,4 +336,15 @@ DenseElementsAttr getHostConstDenseElementsAttr(Value value) { return getHostConstantDenseElementsAttrImpl(value, visited); } +bool isZeroSplatHostConstant(Value value) { + auto denseAttr = getHostConstDenseElementsAttr(value); + if (!denseAttr || !denseAttr.isSplat()) + return false; + if (isa(denseAttr.getElementType())) + return denseAttr.getSplatValue().isZero(); + if (isa(denseAttr.getElementType())) + return denseAttr.getSplatValue().isZero(); + return false; +} + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/CompileTime.hpp b/src/PIM/Conversion/ONNXToSpatial/CompileTime.hpp index b5e6795..29faa71 100644 --- a/src/PIM/Conversion/ONNXToSpatial/CompileTime.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/CompileTime.hpp @@ -21,6 +21,8 @@ bool isCompileTimeOp(mlir::Operation* op); mlir::DenseElementsAttr getHostConstDenseElementsAttr(mlir::Value value); +bool isZeroSplatHostConstant(mlir::Value value); + mlir::FailureOr transposeDenseElementsAttr( mlir::DenseElementsAttr denseAttr, llvm::ArrayRef permutation); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp index 48952ad..8b9c076 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp @@ -2899,20 +2899,60 @@ static FailureOr createPixelMajorConvPatchRow(Value paddedWindow, .getResult(); } -static FailureOr createConvOutputTile(Value patchRow, - Value& partialInputScratch, - Value tileWeights, - int64_t patchSize, - int64_t numKSlices, - int64_t xbarDim, - PatternRewriter& rewriter, - Location loc) { - auto elementType = cast(patchRow.getType()).getElementType(); +static FailureOr> createConvInputTiles(Value paddedWindow, + const ConvLoweringState& state, + Value outputWidth, + Value& partialInputScratch, + int64_t patchSize, + int64_t numKSlices, + int64_t xbarDim, + PatternRewriter& rewriter, + Location loc) { + auto elementType = state.xType.getElementType(); auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType); - auto weightElementType = cast(tileWeights.getType()).getElementType(); - auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType); + SmallVector inputTiles; + inputTiles.reserve(numKSlices); + + if (state.numChannelsIn % xbarDim == 0) { + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + auto inputTileType = RankedTensorType::get( + {1, 1, 1, xbarDim}, elementType, state.xType.getEncoding()); + for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) { + const int64_t linearOffset = kSlice * xbarDim; + const int64_t kernelPixel = linearOffset / state.numChannelsIn; + const int64_t kernelRow = kernelPixel / state.wWidth; + const int64_t kernelColumn = kernelPixel % state.wWidth; + const int64_t channelOffset = linearOffset % state.numChannelsIn; + Value inputWidthOffset = + affineMulConst(rewriter, loc, outputWidth, state.strideWidth, anchorOp); + inputWidthOffset = affineAddConst( + rewriter, loc, inputWidthOffset, kernelColumn * state.dilationWidth, anchorOp); + Value inputTile = tensor::ExtractSliceOp::create( + rewriter, + loc, + inputTileType, + paddedWindow, + SmallVector {rewriter.getIndexAttr(0), + rewriter.getIndexAttr(kernelRow), + inputWidthOffset, + rewriter.getIndexAttr(channelOffset)}, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(xbarDim)}, + getUnitStrides(rewriter, 4)); + inputTiles.push_back(tensor::CollapseShapeOp::create( + rewriter, loc, paddedRowType, inputTile, SmallVector {{0, 1, 2}, {3}}) + .getResult()); + } + return inputTiles; + } + + FailureOr patchRow = + createPixelMajorConvPatchRow(paddedWindow, state, outputWidth, rewriter, loc); + if (failed(patchRow)) + return failure(); - Value tileResult; for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) { const int64_t kOffset = kSlice * xbarDim; const int64_t sliceSize = std::min(xbarDim, patchSize - kOffset); @@ -2921,7 +2961,7 @@ static FailureOr createConvOutputTile(Value patchRow, inputTile = extractStaticSliceOrIdentity( rewriter, loc, - patchRow, + *patchRow, paddedRowType, SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)}, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)}, @@ -2934,7 +2974,7 @@ static FailureOr createConvOutputTile(Value patchRow, Value partial = extractStaticSliceOrIdentity( rewriter, loc, - patchRow, + *patchRow, partialType, SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)}, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)}, @@ -2949,6 +2989,26 @@ static FailureOr createConvOutputTile(Value patchRow, getUnitStrides(rewriter, 2)); inputTile = partialInputScratch; } + inputTiles.push_back(inputTile); + } + return inputTiles; +} + +static FailureOr createConvOutputTile(ValueRange inputTiles, + Value tileWeights, + int64_t outputChannels, + int64_t xbarDim, + PatternRewriter& rewriter, + Location loc) { + auto elementType = cast(inputTiles.front().getType()).getElementType(); + auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType); + auto resultType = RankedTensorType::get({1, outputChannels}, elementType); + auto weightElementType = cast(tileWeights.getType()).getElementType(); + auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType); + + Value tileResult; + for (auto [kSlice, inputTile] : llvm::enumerate(inputTiles)) { + const int64_t kOffset = static_cast(kSlice) * xbarDim; SmallVector bOffsets { rewriter.getIndexAttr(kOffset), rewriter.getIndexAttr(0)}; SmallVector bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)}; @@ -2956,26 +3016,32 @@ static FailureOr createConvOutputTile(Value patchRow, rewriter, loc, tileWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2)); Value piece = spatial::SpatVMMOp::create( rewriter, loc, paddedRowType, bTile, inputTile).getResult(); + if (outputChannels != xbarDim) + piece = tensor::ExtractSliceOp::create( + rewriter, + loc, + resultType, + piece, + SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, + SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(outputChannels)}, + getUnitStrides(rewriter, 2)); tileResult = tileResult ? spatial::SpatVAddOp::create( - rewriter, loc, paddedRowType, tileResult, piece).getResult() + rewriter, loc, resultType, tileResult, piece).getResult() : piece; } return tileResult; } -static FailureOr createConvOutputRow(Value patchRow, - Value& partialInputScratch, - int64_t patchSize, +static FailureOr createConvOutputRow(ValueRange inputTiles, int64_t paddedK, int64_t outputChannels, Value paddedWeights, Value bias, - int64_t numKSlices, int64_t xbarDim, PatternRewriter& rewriter, Location loc) { - auto elementType = cast(patchRow.getType()).getElementType(); + auto elementType = cast(inputTiles.front().getType()).getElementType(); auto rowType = RankedTensorType::get({1, outputChannels}, elementType); auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, @@ -2995,19 +3061,10 @@ static FailureOr createConvOutputRow(Value patchRow, if (outputTileCount == 1) { FailureOr rowResult = createConvOutputTile( - patchRow, partialInputScratch, getTileWeights(0), patchSize, numKSlices, xbarDim, rewriter, loc); + inputTiles, getTileWeights(0), outputChannels, xbarDim, rewriter, loc); if (failed(rowResult)) return failure(); Value validRow = *rowResult; - if (outputChannels != xbarDim) - validRow = tensor::ExtractSliceOp::create( - rewriter, - loc, - rowType, - validRow, - SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, - SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(outputChannels)}, - getUnitStrides(rewriter, 2)); if (bias) validRow = spatial::SpatVAddOp::create(rewriter, loc, rowType, validRow, bias).getResult(); return validRow; @@ -3018,7 +3075,7 @@ static FailureOr createConvOutputRow(Value patchRow, Value paddedOutput = tensor::EmptyOp::create(rewriter, loc, paddedOutputType.getShape(), elementType); for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) { FailureOr tileResult = createConvOutputTile( - patchRow, partialInputScratch, getTileWeights(outputTile), patchSize, numKSlices, xbarDim, rewriter, loc); + inputTiles, getTileWeights(outputTile), xbarDim, xbarDim, rewriter, loc); if (failed(tileResult)) return failure(); SmallVector tileOffsets { @@ -3107,19 +3164,20 @@ static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLow 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); + FailureOr> inputTiles = createConvInputTiles(*inputWindow, + state, + widthIndex, + partialInputScratch, + patchSize, + numKSlices, + xbarDim, + rewriter, + widthLoc); + if (failed(inputTiles)) + return failure(); + FailureOr paddedOutputRow = + createConvOutputTile(*inputTiles, tileWeights, xbarDim, xbarDim, rewriter, widthLoc); if (failed(paddedOutputRow)) return failure(); if (state.hasBias) @@ -3221,19 +3279,23 @@ static FailureOr 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 outputRow = createConvOutputRow(*patchRow, - partialInputScratch, - patchSize, + FailureOr> inputTiles = createConvInputTiles(*inputWindow, + state, + widthIndex, + partialInputScratch, + patchSize, + numKSlices, + xbarDim, + rewriter, + widthLoc); + if (failed(inputTiles)) + return failure(); + FailureOr outputRow = createConvOutputRow(*inputTiles, paddedK, state.numChannelsOut, args.weights.front(), state.hasBias ? args.inputs[1] : Value(), - numKSlices, xbarDim, rewriter, widthLoc); @@ -3334,20 +3396,23 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro 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 outputRow = createConvOutputRow(*patchRow, - partialInputScratch, - patchSize, + FailureOr> inputTiles = createConvInputTiles(*inputWindow, + state, + widthIndex, + partialInputScratch, + patchSize, + numKSlices, + xbarDim, + rewriter, + widthLoc); + if (failed(inputTiles)) + return failure(); + FailureOr outputRow = createConvOutputRow(*inputTiles, paddedK, state.numChannelsOut, args.weights.front(), state.hasBias ? args.inputs[1] : Value(), - numKSlices, xbarDim, rewriter, widthLoc); @@ -3746,7 +3811,8 @@ static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, state.wWidth = state.wType.getDimSize(3); state.outHeight = state.outType.getDimSize(2); state.outWidth = state.outType.getDimSize(3); - state.hasBias = state.b && !isa(state.b.getDefiningOp()); + state.hasBias = + state.b && !isa(state.b.getDefiningOp()) && !isZeroSplatHostConstant(state.b); if (state.numChannelsIn % state.group != 0) { convOp.emitOpError() << "requires input channels " << state.numChannelsIn << " to be divisible by group " @@ -3872,7 +3938,7 @@ static FailureOr analyzeConvLoweringState(spatial::SpatConv2D state.wWidth = state.wType.getDimSize(3); state.outHeight = state.outType.getDimSize(2); state.outWidth = state.outType.getDimSize(3); - state.hasBias = static_cast(planOp.getBias()); + state.hasBias = planOp.getBias() && !isZeroSplatHostConstant(planOp.getBias()); if (state.numChannelsIn % state.group != 0 || state.numChannelsOut % state.group != 0) return planOp.emitOpError("requires input and output channels divisible by group"), failure(); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp index 3c2082a..98bf1e1 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp @@ -323,7 +323,7 @@ static FailureOr verifyDynamicGemmBiasType(RankedTensorType cT static bool hasGemmBias(Value c) { Operation* definingOp = c.getDefiningOp(); - return !definingOp || !isa(definingOp); + return (!definingOp || !isa(definingOp)) && !isZeroSplatHostConstant(c); } static Value createScalarTensorConstant(RankedTensorType scalarType, diff --git a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp index 1629485..b7977b1 100644 --- a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp @@ -131,6 +131,12 @@ static Value getForwardedInputConsumerOutput(OpOperand& use) { pim::PimVVMaxOp, pim::PimVVDMulOp>(owner)) return use.getOperandNumber() < 2 ? owner->getOperand(2) : Value(); + if (isa(owner)) + return use.getOperandNumber() == 0 ? owner->getOperand(1) : Value(); return {}; } diff --git a/src/PIM/Dialect/Pim/Transforms/HostConstantFolding/Patterns/Constant.cpp b/src/PIM/Dialect/Pim/Transforms/HostConstantFolding/Patterns/Constant.cpp index b964097..ad557d3 100644 --- a/src/PIM/Dialect/Pim/Transforms/HostConstantFolding/Patterns/Constant.cpp +++ b/src/PIM/Dialect/Pim/Transforms/HostConstantFolding/Patterns/Constant.cpp @@ -532,6 +532,95 @@ struct FoldConstantMemCpPattern final : OpRewritePattern { } }; +static bool isOne(Attribute value) { + if (auto floatValue = dyn_cast(value)) + return floatValue.getValue().isExactlyValue(1.0); + if (auto integerValue = dyn_cast(value)) + return integerValue.getValue() == 1; + return false; +} + +static bool isAllOneHostCopy(pim::PimMemCopyHostToDevOp copyOp, ModuleOp moduleOp, MemRefType copiedType) { + auto targetOffset = resolveIndexValue(copyOp.getDeviceTargetOffset()); + auto sourceOffset = resolveIndexValue(copyOp.getHostSourceOffset()); + if (failed(targetOffset) || failed(sourceOffset) || *targetOffset != 0) + return false; + + Type elementType = copiedType.getElementType(); + if (!elementType.isIntOrFloat()) + return false; + unsigned bitWidth = elementType.getIntOrFloatBitWidth(); + if (bitWidth == 0 || bitWidth % 8 != 0) + return false; + + int64_t elementBytes = bitWidth / 8; + int64_t copiedElements = copiedType.getNumElements(); + if (*sourceOffset % elementBytes != 0 || copyOp.getSize() != copiedElements * elementBytes) + return false; + + auto source = getDenseGlobalValue(moduleOp, copyOp.getHostSource()); + if (failed(source) || source->getElementType() != elementType) + return false; + + int64_t firstElement = *sourceOffset / elementBytes; + int64_t endElement = firstElement + copiedElements; + if (firstElement < 0 || endElement > source->getNumElements()) + return false; + if (source->isSplat()) + return isOne(source->getSplatValue()); + + int64_t index = 0; + for (Attribute value : source->getValues()) { + if (index >= firstElement && index < endElement && !isOne(value)) + return false; + if (++index >= endElement) + break; + } + return true; +} + +struct FoldMultiplyByOnePattern final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(pim::PimVVMulOp mulOp, PatternRewriter& rewriter) const override { + auto moduleOp = mulOp->getParentOfType(); + if (!moduleOp) + return failure(); + + for (auto [mask, input] : {std::pair {mulOp.getLhs(), mulOp.getRhs()}, + std::pair {mulOp.getRhs(), mulOp.getLhs()}}) { + auto maskAlloc = mask.getDefiningOp(); + if (!maskAlloc) + continue; + + pim::PimMemCopyHostToDevOp copyOp; + for (Operation* user : maskAlloc->getUsers()) { + if (user == mulOp) + continue; + auto candidate = dyn_cast(user); + if (!candidate || candidate.getDeviceTarget() != mask || copyOp) { + copyOp = {}; + break; + } + copyOp = candidate; + } + auto maskType = dyn_cast(mask.getType()); + if (!copyOp || !copyOp.use_empty() || !maskType || !isAllOneHostCopy(copyOp, moduleOp, maskType)) + continue; + + auto outputAlloc = mulOp.getOutputBuffer().getDefiningOp(); + rewriter.replaceOp(mulOp, input); + rewriter.eraseOp(copyOp); + if (maskAlloc.use_empty()) + rewriter.eraseOp(maskAlloc); + if (outputAlloc && outputAlloc.use_empty()) + rewriter.eraseOp(outputAlloc); + return success(); + } + return failure(); + } +}; + } // namespace void populateConstantFoldingConstantPatterns(RewritePatternSet& patterns) { @@ -539,7 +628,8 @@ void populateConstantFoldingConstantPatterns(RewritePatternSet& patterns) { FoldConstantAllocPattern, FoldConstantCoreMapPattern, FoldConstantHostCopyPattern, - FoldConstantMemCpPattern>(patterns.getContext()); + FoldConstantMemCpPattern, + FoldMultiplyByOnePattern>(patterns.getContext()); } } // namespace onnx_mlir