vgg8 6.88 vs 7.89
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
ilgeco
2026-07-28 12:46:46 +02:00
parent 4c6fc1173e
commit 78bfb8a9aa
6 changed files with 239 additions and 64 deletions
@@ -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<FloatType>(denseAttr.getElementType()))
return denseAttr.getSplatValue<APFloat>().isZero();
if (isa<IntegerType>(denseAttr.getElementType()))
return denseAttr.getSplatValue<APInt>().isZero();
return false;
}
} // namespace onnx_mlir
@@ -21,6 +21,8 @@ bool isCompileTimeOp(mlir::Operation* op);
mlir::DenseElementsAttr getHostConstDenseElementsAttr(mlir::Value value);
bool isZeroSplatHostConstant(mlir::Value value);
mlir::FailureOr<mlir::DenseElementsAttr> transposeDenseElementsAttr(
mlir::DenseElementsAttr denseAttr, llvm::ArrayRef<int64_t> permutation);
@@ -2899,20 +2899,60 @@ static FailureOr<Value> createPixelMajorConvPatchRow(Value paddedWindow,
.getResult();
}
static FailureOr<Value> createConvOutputTile(Value patchRow,
Value& partialInputScratch,
Value tileWeights,
int64_t patchSize,
int64_t numKSlices,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
auto elementType = cast<RankedTensorType>(patchRow.getType()).getElementType();
static FailureOr<SmallVector<Value>> 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<RankedTensorType>(tileWeights.getType()).getElementType();
auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType);
SmallVector<Value> 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<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelRow),
inputWidthOffset,
rewriter.getIndexAttr(channelOffset)},
SmallVector<OpFoldResult> {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<ReassociationIndices> {{0, 1, 2}, {3}})
.getResult());
}
return inputTiles;
}
FailureOr<Value> 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<Value> createConvOutputTile(Value patchRow,
inputTile = extractStaticSliceOrIdentity(
rewriter,
loc,
patchRow,
*patchRow,
paddedRowType,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)},
@@ -2934,7 +2974,7 @@ static FailureOr<Value> createConvOutputTile(Value patchRow,
Value partial = extractStaticSliceOrIdentity(
rewriter,
loc,
patchRow,
*patchRow,
partialType,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)},
@@ -2949,6 +2989,26 @@ static FailureOr<Value> createConvOutputTile(Value patchRow,
getUnitStrides(rewriter, 2));
inputTile = partialInputScratch;
}
inputTiles.push_back(inputTile);
}
return inputTiles;
}
static FailureOr<Value> createConvOutputTile(ValueRange inputTiles,
Value tileWeights,
int64_t outputChannels,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
auto elementType = cast<RankedTensorType>(inputTiles.front().getType()).getElementType();
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto resultType = RankedTensorType::get({1, outputChannels}, elementType);
auto weightElementType = cast<RankedTensorType>(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<int64_t>(kSlice) * xbarDim;
SmallVector<OpFoldResult> bOffsets {
rewriter.getIndexAttr(kOffset), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)};
@@ -2956,26 +3016,32 @@ static FailureOr<Value> 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<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {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<Value> createConvOutputRow(Value patchRow,
Value& partialInputScratch,
int64_t patchSize,
static FailureOr<Value> 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<RankedTensorType>(patchRow.getType()).getElementType();
auto elementType = cast<RankedTensorType>(inputTiles.front().getType()).getElementType();
auto rowType = RankedTensorType::get({1, outputChannels}, elementType);
auto tileWeightsType =
RankedTensorType::get({paddedK, xbarDim},
@@ -2995,19 +3061,10 @@ static FailureOr<Value> createConvOutputRow(Value patchRow,
if (outputTileCount == 1) {
FailureOr<Value> 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<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {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<Value> createConvOutputRow(Value patchRow,
Value paddedOutput = tensor::EmptyOp::create(rewriter, loc, paddedOutputType.getShape(), elementType);
for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) {
FailureOr<Value> tileResult = createConvOutputTile(
patchRow, partialInputScratch, getTileWeights(outputTile), patchSize, numKSlices, xbarDim, rewriter, loc);
inputTiles, getTileWeights(outputTile), xbarDim, xbarDim, rewriter, loc);
if (failed(tileResult))
return failure();
SmallVector<OpFoldResult> tileOffsets {
@@ -3107,19 +3164,20 @@ static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLow
Value widthIndex,
ValueRange widthIterArgs,
SmallVectorImpl<Value>& widthYielded) {
FailureOr<Value> patchRow =
createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc);
if (failed(patchRow))
return failure();
Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value();
FailureOr<Value> paddedOutputRow = createConvOutputTile(*patchRow,
partialInputScratch,
tileWeights,
patchSize,
numKSlices,
xbarDim,
rewriter,
widthLoc);
FailureOr<SmallVector<Value>> inputTiles = createConvInputTiles(*inputWindow,
state,
widthIndex,
partialInputScratch,
patchSize,
numKSlices,
xbarDim,
rewriter,
widthLoc);
if (failed(inputTiles))
return failure();
FailureOr<Value> paddedOutputRow =
createConvOutputTile(*inputTiles, tileWeights, xbarDim, xbarDim, rewriter, widthLoc);
if (failed(paddedOutputRow))
return failure();
if (state.hasBias)
@@ -3221,19 +3279,23 @@ static FailureOr<Value>
c1,
widthLoopInit,
[&](OpBuilder&, Location widthLoc, Value widthIndex, ValueRange widthIterArgs, SmallVectorImpl<Value>& widthYielded) {
FailureOr<Value> patchRow =
createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc);
if (failed(patchRow))
return failure();
Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value();
FailureOr<Value> outputRow = createConvOutputRow(*patchRow,
partialInputScratch,
patchSize,
FailureOr<SmallVector<Value>> inputTiles = createConvInputTiles(*inputWindow,
state,
widthIndex,
partialInputScratch,
patchSize,
numKSlices,
xbarDim,
rewriter,
widthLoc);
if (failed(inputTiles))
return failure();
FailureOr<Value> 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<Value> createConvOutputFromPixelMajorRowStripFragments(Value ro
c1,
widthLoopInit,
[&](OpBuilder&, Location widthLoc, Value widthIndex, ValueRange widthIterArgs, SmallVectorImpl<Value>& widthYielded) {
FailureOr<Value> patchRow =
createPixelMajorConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc);
if (failed(patchRow))
return failure();
Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value();
FailureOr<Value> outputRow = createConvOutputRow(*patchRow,
partialInputScratch,
patchSize,
FailureOr<SmallVector<Value>> inputTiles = createConvInputTiles(*inputWindow,
state,
widthIndex,
partialInputScratch,
patchSize,
numKSlices,
xbarDim,
rewriter,
widthLoc);
if (failed(inputTiles))
return failure();
FailureOr<Value> 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<ConvLoweringState> 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<ONNXNoneOp>(state.b.getDefiningOp());
state.hasBias =
state.b && !isa<ONNXNoneOp>(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<ConvLoweringState> 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<bool>(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();
@@ -323,7 +323,7 @@ static FailureOr<RankedTensorType> verifyDynamicGemmBiasType(RankedTensorType cT
static bool hasGemmBias(Value c) {
Operation* definingOp = c.getDefiningOp();
return !definingOp || !isa<ONNXNoneOp>(definingOp);
return (!definingOp || !isa<ONNXNoneOp>(definingOp)) && !isZeroSplatHostConstant(c);
}
static Value createScalarTensorConstant(RankedTensorType scalarType,
@@ -131,6 +131,12 @@ static Value getForwardedInputConsumerOutput(OpOperand& use) {
pim::PimVVMaxOp,
pim::PimVVDMulOp>(owner))
return use.getOperandNumber() < 2 ? owner->getOperand(2) : Value();
if (isa<pim::PimVAvgOp,
pim::PimVReluOp,
pim::PimVTanhOp,
pim::PimVSigmOp,
pim::PimVSoftmaxOp>(owner))
return use.getOperandNumber() == 0 ? owner->getOperand(1) : Value();
return {};
}
@@ -532,6 +532,95 @@ struct FoldConstantMemCpPattern final : OpRewritePattern<pim::PimMemCopyOp> {
}
};
static bool isOne(Attribute value) {
if (auto floatValue = dyn_cast<FloatAttr>(value))
return floatValue.getValue().isExactlyValue(1.0);
if (auto integerValue = dyn_cast<IntegerAttr>(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<Attribute>());
int64_t index = 0;
for (Attribute value : source->getValues<Attribute>()) {
if (index >= firstElement && index < endElement && !isOne(value))
return false;
if (++index >= endElement)
break;
}
return true;
}
struct FoldMultiplyByOnePattern final : OpRewritePattern<pim::PimVVMulOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(pim::PimVVMulOp mulOp, PatternRewriter& rewriter) const override {
auto moduleOp = mulOp->getParentOfType<ModuleOp>();
if (!moduleOp)
return failure();
for (auto [mask, input] : {std::pair {mulOp.getLhs(), mulOp.getRhs()},
std::pair {mulOp.getRhs(), mulOp.getLhs()}}) {
auto maskAlloc = mask.getDefiningOp<memref::AllocOp>();
if (!maskAlloc)
continue;
pim::PimMemCopyHostToDevOp copyOp;
for (Operation* user : maskAlloc->getUsers()) {
if (user == mulOp)
continue;
auto candidate = dyn_cast<pim::PimMemCopyHostToDevOp>(user);
if (!candidate || candidate.getDeviceTarget() != mask || copyOp) {
copyOp = {};
break;
}
copyOp = candidate;
}
auto maskType = dyn_cast<MemRefType>(mask.getType());
if (!copyOp || !copyOp.use_empty() || !maskType || !isAllOneHostCopy(copyOp, moduleOp, maskType))
continue;
auto outputAlloc = mulOp.getOutputBuffer().getDefiningOp<memref::AllocOp>();
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