finally faster on all pimcomp models
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-07-30 14:07:01 +02:00
parent 1b4f070bef
commit c12f69133d
4 changed files with 140 additions and 111 deletions
@@ -9,6 +9,8 @@
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Dialect/ONNX/ONNXOps.hpp"
#include <numeric>
using namespace mlir;
namespace onnx_mlir {
@@ -255,10 +257,18 @@ FailureOr<Value> applyRowStripAdd(const RowStripPhysicalValue& lhs,
const RowStripPhysicalValue& rhs,
PatternRewriter& rewriter,
Location loc) {
if (lhs.logicalType != rhs.logicalType || lhs.fragmentType != rhs.fragmentType
|| lhs.storage.getType() != rhs.storage.getType() || lhs.tilesPerRow != rhs.tilesPerRow)
if (lhs.logicalType != rhs.logicalType)
return failure();
auto storageType = cast<RankedTensorType>(lhs.storage.getType());
const int64_t fragmentChannels =
std::gcd(lhs.fragmentType.getDimSize(3), rhs.fragmentType.getDimSize(3));
auto fragmentType = RankedTensorType::get(
{1, 1, lhs.logicalType.getDimSize(3), fragmentChannels},
lhs.logicalType.getElementType(),
lhs.logicalType.getEncoding());
const int64_t tilesPerRow = ceilIntegerDivide(lhs.logicalType.getDimSize(1), fragmentChannels);
auto storageType =
spatial::getGraphBatchPhysicalResultType(lhs.logicalType.getDimSize(2) * tilesPerRow, fragmentType);
const int64_t laneCount = storageType.getDimSize(0);
auto batch = createSpatComputeBatch(
rewriter,
@@ -268,13 +278,43 @@ FailureOr<Value> applyRowStripAdd(const RowStripPhysicalValue& lhs,
{},
ValueRange {lhs.storage, rhs.storage},
[&](detail::SpatComputeBatchBodyArgs args) {
FailureOr<Value> lhsFragment =
extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[0], args.lane, lhs.fragmentType);
FailureOr<Value> rhsFragment =
extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[1], args.lane, rhs.fragmentType);
auto extractFragment = [&](Value storage, const RowStripPhysicalValue& input) -> FailureOr<Value> {
Operation* anchor = rewriter.getInsertionBlock()->getParentOp();
MLIRContext* context = rewriter.getContext();
AffineExpr lane = getAffineDimExpr(0, context);
AffineExpr outputTile = lane % tilesPerRow;
AffineExpr channelOffset = outputTile * fragmentChannels;
Value sourceSlot = createOrFoldAffineApply(
rewriter,
loc,
lane.floorDiv(tilesPerRow) * input.tilesPerRow
+ channelOffset.floorDiv(input.fragmentType.getDimSize(3)),
ValueRange {args.lane},
anchor);
FailureOr<Value> source =
extractGraphBatchPhysicalFragment(rewriter, loc, storage, sourceSlot, input.fragmentType);
if (failed(source) || input.fragmentType == fragmentType)
return source;
Value sourceOffset = createOrFoldAffineApply(
rewriter, loc, channelOffset % input.fragmentType.getDimSize(3), ValueRange {args.lane}, anchor);
return tensor::ExtractSliceOp::create(
rewriter,
loc,
fragmentType,
*source,
SmallVector<OpFoldResult> {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(lhs.logicalType.getDimSize(3)),
rewriter.getIndexAttr(fragmentChannels)},
getUnitStrides(rewriter, 4)).getResult();
};
FailureOr<Value> lhsFragment = extractFragment(args.inputs[0], lhs);
FailureOr<Value> rhsFragment = extractFragment(args.inputs[1], rhs);
if (failed(lhsFragment) || failed(rhsFragment))
return failure();
Value added = spatial::SpatVAddOp::create(rewriter, loc, lhs.fragmentType, *lhsFragment, *rhsFragment);
Value added = spatial::SpatVAddOp::create(rewriter, loc, fragmentType, *lhsFragment, *rhsFragment);
publishGraphBatchPhysicalFragment(rewriter, loc, added, args.outputs.front(), args.lane);
return success();
});
@@ -2511,6 +2511,11 @@ static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) {
<= static_cast<int64_t>(crossbarCountInCore.getValue());
}
static bool rowStripOutputChannelTileFitsOneCore(const ConvGeometry& geometry) {
return ceilIntegerDivide(geometry.k, geometry.xbarSize)
<= static_cast<int64_t>(crossbarCountInCore.getValue());
}
static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state, StringRef& failureReason) {
if (state.batchSize != 1) {
failureReason = "batch_not_one";
@@ -2528,24 +2533,16 @@ static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state
failureReason = "non_float_input";
return false;
}
if (state.strideHeight != 1 || state.strideWidth != 1) {
failureReason = "stride_not_one";
return false;
}
if (state.dilationHeight != 1 || state.dilationWidth != 1) {
failureReason = "dilation_not_one";
return false;
}
if (state.outHeight != state.xHeight || state.outWidth != state.xWidth) {
failureReason = "not_same_spatial_shape";
return false;
}
if (!getHostConstDenseElementsAttr(state.w)) {
failureReason = "non_constant_weight";
return false;
}
if (!rowStripOutputTileFitsOneCore(buildConvGeometry(state))) {
failureReason = "output_row_does_not_fit_one_core";
if (!rowStripOutputChannelTileFitsOneCore(buildConvGeometry(state))) {
failureReason = "output_channel_tile_does_not_fit_one_core";
return false;
}
if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)) {
@@ -3306,7 +3303,7 @@ static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLow
static FailureOr<Value>
createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) {
ConvGeometry geometry = buildConvGeometry(state);
if (state.group != 1 || state.batchSize != 1 || !rowStripOutputTileFitsOneCore(geometry))
if (state.group != 1 || state.batchSize != 1 || !rowStripOutputChannelTileFitsOneCore(geometry))
return failure();
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
@@ -3320,11 +3317,16 @@ static FailureOr<Value>
const int64_t numKSlices = ceilIntegerDivide(patchSize, xbarDim);
const int64_t paddedK = numKSlices * xbarDim;
const int64_t paddedOutputChannels =
ceilIntegerDivide(state.numChannelsOut, xbarDim) * xbarDim;
Value paddedWeights = standard::createPaddedPixelMajorWeightConstant(
weightDenseAttr, state, paddedK, paddedOutputChannels, rewriter);
if (!rowStripOutputTileFitsOneCore(geometry)) {
Value tiledWeights =
standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter);
return createOutputChannelTiledRowStripConvOutput(
state, state.x, tiledWeights, paddedK, numKSlices, xbarDim, rewriter, loc);
}
const int64_t paddedOutputChannels = ceilIntegerDivide(state.numChannelsOut, xbarDim) * xbarDim;
Value paddedWeights =
standard::createPaddedPixelMajorWeightConstant(weightDenseAttr, state, paddedK, paddedOutputChannels, rewriter);
FailureOr<Value> bias = failure();
if (state.hasBias)
bias = createBiasRowConstant(state, rewriter);
@@ -3355,6 +3357,13 @@ static FailureOr<Value> createConvOutputFromPixelMajorRowStripFragments(Value ro
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!weightDenseAttr)
return failure();
if (!rowStripOutputTileFitsOneCore(geometry)) {
Value tiledWeights =
standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter);
return createOutputChannelTiledRowStripConvOutput(
state, rowStripStorage, tiledWeights, paddedK, numKSlices, xbarDim, rewriter, loc);
}
const int64_t paddedOutputChannels =
ceilIntegerDivide(state.numChannelsOut, xbarDim) * xbarDim;
Value paddedWeights = standard::createPaddedPixelMajorWeightConstant(
@@ -4298,7 +4307,7 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) {
analysis.barrierKind = DistributedConvBarrierKind::UnsupportedConsumer;
analysis.barrierDetail = "selected row-strip layout";
ConvGeometry geometry = buildConvGeometry(*state);
if (!rowStripOutputTileFitsOneCore(geometry))
if (!rowStripOutputChannelTileFitsOneCore(geometry))
return failure();
ConvLoweringDecision decision = chooseConvLoweringStrategy(geometry, *requestedStrategy, analysis);
if (decision.strategy == PimConvLoweringDepthwise && !depthwise::canUseStructuredRewrite(*state)
@@ -501,7 +501,9 @@ static Value extractReductionPiece(Value partialPiecesArg,
SmallVector<OpFoldResult> unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)};
SmallVector<OpFoldResult> pieceSizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> pieceOffsets {
createPartialGroupOffset(hSlice, kSlice, numKSlices, numOutRows, rewriter, loc), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
createPartialGroupOffset(hSlice, kSlice, numKSlices, numOutRows, rewriter, loc),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0)};
return extractMixedSliceOrIdentity(
rewriter, loc, partialPiecesArg, pieceType,
{pieceOffsets, pieceSizes, unitStrides});
@@ -535,94 +537,72 @@ static Value reducePartialPiecesForHSlice(Value partialPiecesArg,
return activePieces.front();
}
static FailureOr<spatial::SpatCompute> createReductionCompute(Value partialPieces,
Value bias,
RankedTensorType partialPiecesType,
RankedTensorType outType,
RankedTensorType paddedOutType,
int64_t numKSlices,
ConversionPatternRewriter& rewriter,
Location loc) {
static FailureOr<Value> createReductionOutput(Value partialPieces,
Value bias,
RankedTensorType partialPiecesType,
RankedTensorType outType,
RankedTensorType paddedOutType,
int64_t numKSlices,
ConversionPatternRewriter& rewriter,
Location loc) {
const int64_t numOutRows = outType.getDimSize(0);
const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(1), crossbarSize.getValue());
auto pieceType = RankedTensorType::get({numOutRows, static_cast<int64_t>(crossbarSize.getValue())},
partialPiecesType.getElementType());
if (bias && cast<RankedTensorType>(bias.getType()) != paddedOutType)
bias = createZeroPaddedTensor(bias, paddedOutType, rewriter, loc);
SmallVector<Value> inputs {partialPieces};
if (bias)
inputs.push_back(bias);
auto computeOp =
createSpatCompute(rewriter, loc, TypeRange {outType}, {}, inputs, [&](ValueRange blockArgs) -> LogicalResult {
Value partialPiecesArg = blockArgs[0];
Value biasArg = bias ? blockArgs[1] : Value();
if (biasArg && cast<RankedTensorType>(biasArg.getType()) != paddedOutType)
biasArg = createZeroPaddedTensor(biasArg, paddedOutType, rewriter, loc);
const int64_t numOutRows = outType.getDimSize(0);
const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(1), crossbarSize.getValue());
auto pieceType = RankedTensorType::get({numOutRows, static_cast<int64_t>(crossbarSize.getValue())},
partialPiecesType.getElementType());
Value outputInit =
tensor::EmptyOp::create(rewriter, loc, paddedOutType.getShape(), paddedOutType.getElementType()).getResult();
SmallVector<OpFoldResult> unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)};
SmallVector<OpFoldResult> pieceSizes {rewriter.getIndexAttr(numOutRows),
rewriter.getIndexAttr(crossbarSize.getValue())};
auto buildOutputSlice = [&](Value outputAcc, Value hSlice) -> Value {
Value reduced =
reducePartialPiecesForHSlice(partialPiecesArg, hSlice, pieceType, numKSlices, numOutRows, rewriter, loc);
Value hOffset = onnx_mlir::affineMulConst(
rewriter, loc, hSlice, crossbarSize.getValue(), rewriter.getInsertionBlock()->getParentOp());
if (biasArg) {
SmallVector<OpFoldResult> biasOffsets {rewriter.getIndexAttr(0), hOffset};
Value biasSlice =
tensor::ExtractSliceOp::create(rewriter, loc, pieceType, biasArg, biasOffsets, pieceSizes, unitStrides)
.getResult();
reduced = spatial::SpatVAddOp::create(rewriter, loc, pieceType, reduced, biasSlice).getResult();
}
SmallVector<OpFoldResult> outputOffsets {rewriter.getIndexAttr(0), hOffset};
return tensor::InsertSliceOp::create(rewriter, loc, reduced, outputAcc, outputOffsets, pieceSizes, unitStrides)
.getResult();
};
Value paddedOutput = outputInit;
if (numOutHSlices == 1) {
Value hSlice = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0);
paddedOutput = buildOutputSlice(outputInit, hSlice);
}
else {
Value c0 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0);
Value c1 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 1);
Value cOutHSlices =
getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), numOutHSlices);
auto hLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cOutHSlices,
c1,
ValueRange {outputInit},
[&](OpBuilder&, Location, Value hSlice, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
yielded.push_back(buildOutputSlice(iterArgs.front(), hSlice));
return success();
});
if (failed(hLoop))
return failure();
paddedOutput = hLoop->results.front();
}
Value result = paddedOutput;
if (paddedOutType != outType) {
SmallVector<OpFoldResult> outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> outputSizes {rewriter.getIndexAttr(outType.getDimSize(0)),
rewriter.getIndexAttr(outType.getDimSize(1))};
result =
tensor::ExtractSliceOp::create(rewriter, loc, outType, paddedOutput, outputOffsets, outputSizes, unitStrides)
SmallVector<Value> outputSlices;
outputSlices.reserve(numOutHSlices);
for (int64_t hSlice = 0; hSlice < numOutHSlices; ++hSlice) {
const int64_t columnOffset = hSlice * crossbarSize.getValue();
const int64_t columns =
std::min(static_cast<int64_t>(crossbarSize.getValue()), outType.getDimSize(1) - columnOffset);
auto outputSliceType = RankedTensorType::get({numOutRows, columns}, outType.getElementType());
auto computeOp = createSpatCompute(
rewriter, loc, TypeRange {outputSliceType}, {}, inputs, [&](ValueRange blockArgs) -> LogicalResult {
Value hSliceValue =
getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), hSlice);
Value reduced = reducePartialPiecesForHSlice(
blockArgs[0], hSliceValue, pieceType, numKSlices, numOutRows, rewriter, loc);
if (bias) {
SmallVector<OpFoldResult> biasOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(columnOffset)};
SmallVector<OpFoldResult> pieceSizes {rewriter.getIndexAttr(numOutRows),
rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)};
Value biasSlice =
tensor::ExtractSliceOp::create(rewriter, loc, pieceType, blockArgs[1], biasOffsets, pieceSizes, unitStrides)
.getResult();
reduced = spatial::SpatVAddOp::create(rewriter, loc, pieceType, reduced, biasSlice).getResult();
}
spatial::SpatYieldOp::create(rewriter, loc, result);
if (outputSliceType != pieceType) {
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(columns)};
reduced = tensor::ExtractSliceOp::create(
rewriter, loc, outputSliceType, reduced, offsets, sizes, getUnitStrides(rewriter, 2));
}
spatial::SpatYieldOp::create(rewriter, loc, reduced);
return success();
});
if (failed(computeOp))
return failure();
outputSlices.push_back(computeOp->getResult(0));
}
if (outputSlices.size() == 1)
return outputSlices.front();
return computeOp;
auto concatCompute =
createSpatCompute(rewriter, loc, TypeRange {outType}, {}, outputSlices, [&](ValueRange blockArgs) {
Value result =
spatial::SpatConcatOp::create(rewriter, loc, outType, rewriter.getI64IntegerAttr(1), blockArgs).getOutput();
spatial::SpatYieldOp::create(rewriter, loc, result);
});
return concatCompute.getResult(0);
}
struct GemmToSpatialComputes : OpConversionPattern<ONNXGemmOp> {
@@ -800,12 +780,12 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp,
createVmmBatch(a, b, aType, paddedBType, partialPiecesType, numOutRows, numKSlices, numOutHSlices, rewriter, loc);
if (failed(batchOp))
return failure();
auto reductionCompute = createReductionCompute(
auto reductionOutput = createReductionOutput(
batchOp->getResult(0), bias, partialPiecesType, outType, paddedOutType, numKSlices, rewriter, loc);
if (failed(reductionCompute))
if (failed(reductionOutput))
return failure();
rewriter.replaceOp(gemmOp, reductionCompute->getResults());
rewriter.replaceOp(gemmOp, *reductionOutput);
return success();
}
@@ -1,5 +1,5 @@
model,raptor_latency_ms,pimcomp_latency_ms,raptor_energy_pj,pimcomp_energy_pj,faster_compiler,speedup
vgg8,2.451664,7.985074,647608825.120001,1597904071.120000,raptor,3.26
resnet18,113.483357,58.853733,23204566467.119949,13983168748.119974,pimcomp,1.93
resnet34,120.655292,91.607980,27614808203.679939,22722922369.680016,pimcomp,1.32
googlenet,22.936609,62.923463,7798109358.239990,14547526780.240000,raptor,2.74
vgg8,2.463417,7.985074,649524564.120001,1597904071.120000,raptor,3.24
resnet18,31.449811,58.853733,9335506856.119984,13983168748.119974,raptor,1.87
resnet34,58.071522,91.607980,17027195415.679953,22722922369.680016,raptor,1.58
googlenet,24.247113,62.923463,8019605936.239990,14547526780.240000,raptor,2.60
1 model raptor_latency_ms pimcomp_latency_ms raptor_energy_pj pimcomp_energy_pj faster_compiler speedup
2 vgg8 2.451664 2.463417 7.985074 647608825.120001 649524564.120001 1597904071.120000 raptor 3.26 3.24
3 resnet18 113.483357 31.449811 58.853733 23204566467.119949 9335506856.119984 13983168748.119974 pimcomp raptor 1.93 1.87
4 resnet34 120.655292 58.071522 91.607980 27614808203.679939 17027195415.679953 22722922369.680016 pimcomp raptor 1.32 1.58
5 googlenet 22.936609 24.247113 62.923463 7798109358.239990 8019605936.239990 14547526780.240000 raptor 2.74 2.60