Resnet is fast

This commit is contained in:
ilgeco
2026-07-20 11:34:59 +02:00
parent 5f42da36ae
commit 6bad9a8008
16 changed files with 883 additions and 138 deletions
@@ -16,6 +16,8 @@
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Dialect/ONNX/ONNXOps.hpp"
@@ -24,7 +26,7 @@ using namespace mlir;
namespace onnx_mlir {
namespace {
static Value materializeTileTensor(ConversionPatternRewriter& rewriter, Location loc, Value tile) {
static Value materializeTileTensor(PatternRewriter& rewriter, Location loc, Value tile) {
auto tileType = cast<RankedTensorType>(tile.getType());
Value empty = tensor::EmptyOp::create(rewriter, loc, tileType.getShape(), tileType.getElementType());
return insertStaticSlice(rewriter, loc, tile, empty, getZeroOffsets(rewriter, tileType.getRank()));
@@ -228,6 +230,23 @@ struct PoolToSpatialComputeBase : public OpConversionPattern<PoolOp> {
}
}
if constexpr (std::is_same_v<PoolOp, ONNXMaxPoolSingleOutOp>) {
if (batchSize == 1) {
auto plan = spatial::SpatMaxPool2DPlanOp::create(
rewriter,
loc,
outType,
x,
rewriter.getDenseI64ArrayAttr({kernelHeight, kernelWidth}),
rewriter.getDenseI64ArrayAttr({padTop, padLeft, padBottom, padRight}),
rewriter.getDenseI64ArrayAttr({strideHeight, strideWidth}),
rewriter.getDenseI64ArrayAttr({dilationHeight, dilationWidth}),
rewriter.getStringAttr("nchw"));
rewriter.replaceOp(poolOp, plan.getResult());
return success();
}
}
const int64_t xbarSize = static_cast<int64_t>(crossbarSize.getValue());
const int64_t channelTileCount = (channels + xbarSize - 1) / xbarSize;
const int64_t outputPatchCount = batchSize * outputHeight * outputWidth;
@@ -396,6 +415,220 @@ struct PoolToSpatialCompute<ONNXAveragePoolOp>
} // namespace
LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) {
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape())
return failure();
if (inputType.getRank() != 4 || outputType.getRank() != 4 || inputType.getDimSize(0) != 1
|| outputType.getDimSize(0) != 1 || inputType.getDimSize(1) != outputType.getDimSize(1))
return failure();
if (llvm::any_of(planOp.getKernelShape(), [](int64_t value) { return value <= 0; })
|| llvm::any_of(planOp.getStrides(), [](int64_t value) { return value <= 0; })
|| llvm::any_of(planOp.getDilations(), [](int64_t value) { return value <= 0; }))
return failure();
return success();
}
static Value createClampedPoolIndexTable(PatternRewriter& rewriter,
Operation* anchorOp,
int64_t outputSize,
int64_t kernelSize,
int64_t stride,
int64_t dilation,
int64_t padBegin,
int64_t inputSize) {
auto tableType = RankedTensorType::get({outputSize * kernelSize}, rewriter.getIndexType());
SmallVector<Attribute> values;
values.reserve(tableType.getNumElements());
for (int64_t output = 0; output < outputSize; ++output)
for (int64_t kernel = 0; kernel < kernelSize; ++kernel)
values.push_back(rewriter.getIndexAttr(
std::clamp(output * stride + kernel * dilation - padBegin, int64_t {0}, inputSize - 1)));
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value extractPoolIndex(PatternRewriter& rewriter,
Location loc,
Operation* anchorOp,
Value table,
Value outputIndex,
int64_t kernelIndex,
int64_t kernelSize) {
Value tableIndex = arith::MulIOp::create(
rewriter, loc, outputIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelSize));
if (kernelIndex != 0)
tableIndex = arith::AddIOp::create(
rewriter, loc, tableIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelIndex));
return tensor::ExtractOp::create(rewriter, loc, table, tableIndex);
}
FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
std::optional<Value> rowStripInput,
PatternRewriter& rewriter) {
if (failed(canLowerMaxPoolPlanToRowStrip(planOp)))
return failure();
Location loc = planOp.getLoc();
auto inputType = cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = cast<RankedTensorType>(planOp.getOutput().getType());
const int64_t channels = inputType.getDimSize(1);
const int64_t inputHeight = inputType.getDimSize(2);
const int64_t inputWidth = inputType.getDimSize(3);
const int64_t outputHeight = outputType.getDimSize(2);
const int64_t outputWidth = outputType.getDimSize(3);
const int64_t kernelHeight = planOp.getKernelShape()[0];
const int64_t kernelWidth = planOp.getKernelShape()[1];
Value input = rowStripInput.value_or(planOp.getInput());
auto actualInputType = dyn_cast<RankedTensorType>(input.getType());
const bool physicalInput = actualInputType == getRowStripStorageType(inputType);
if (!physicalInput && actualInputType != inputType)
return failure();
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value rowTable = createClampedPoolIndexTable(rewriter,
anchorOp,
outputHeight,
kernelHeight,
planOp.getStrides()[0],
planOp.getDilations()[0],
planOp.getPads()[0],
inputHeight);
Value columnTable = createClampedPoolIndexTable(rewriter,
anchorOp,
outputWidth,
kernelWidth,
planOp.getStrides()[1],
planOp.getDilations()[1],
planOp.getPads()[1],
inputWidth);
auto inputFragmentType = getRowStripFragmentType(inputType);
auto outputFragmentType = getRowStripFragmentType(outputType);
auto outputStorageType = getRowStripStorageType(outputType);
auto tileType = RankedTensorType::get({1, channels, 1, 1}, outputType.getElementType());
auto batch = createSpatComputeBatch(
rewriter,
loc,
TypeRange {outputStorageType},
outputHeight,
{},
ValueRange {input},
[&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult {
SmallVector<Value> inputRows;
inputRows.reserve(kernelHeight);
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
Value sourceRow =
extractPoolIndex(rewriter, loc, anchorOp, rowTable, args.lane, kernelRow, kernelHeight);
if (physicalInput) {
inputRows.push_back(
extractRowStripFragment(args.inputs.front(), inputType, sourceRow, rewriter, loc));
}
else {
SmallVector<OpFoldResult> offsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)};
inputRows.push_back(tensor::ExtractSliceOp::create(rewriter,
loc,
inputFragmentType,
args.inputs.front(),
offsets,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(channels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(inputWidth)},
getUnitStrides(rewriter, 4)));
}
}
auto windowType = RankedTensorType::get(
{1, channels, kernelHeight, inputWidth}, inputType.getElementType(), inputType.getEncoding());
Value window = tensor::EmptyOp::create(
rewriter, loc, windowType.getShape(), windowType.getElementType());
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelRow),
rewriter.getIndexAttr(0)};
window = tensor::InsertSliceOp::create(rewriter,
loc,
inputRows[kernelRow],
window,
offsets,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(channels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(inputWidth)},
getUnitStrides(rewriter, 4));
}
Value outputInit = tensor::EmptyOp::create(
rewriter, loc, outputFragmentType.getShape(), outputFragmentType.getElementType());
Operation* bodyAnchor = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, bodyAnchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, bodyAnchor, 1);
Value cOutputWidth = getOrCreateIndexConstant(rewriter, bodyAnchor, outputWidth);
auto outputLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cOutputWidth,
c1,
ValueRange {outputInit},
[&](OpBuilder&, Location nestedLoc, Value outputColumn, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value reduced;
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
for (int64_t kernelColumn = 0; kernelColumn < kernelWidth; ++kernelColumn) {
Value sourceColumn = extractPoolIndex(rewriter,
nestedLoc,
bodyAnchor,
columnTable,
outputColumn,
kernelColumn,
kernelWidth);
SmallVector<OpFoldResult> offsets {
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelRow),
sourceColumn};
Value point = tensor::ExtractSliceOp::create(rewriter,
nestedLoc,
tileType,
window,
offsets,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(channels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1)},
getUnitStrides(rewriter, 4));
reduced = reduced ? spatial::SpatVMaxOp::create(rewriter, nestedLoc, tileType, reduced, point).getResult()
: materializeTileTensor(rewriter, nestedLoc, point);
}
}
SmallVector<OpFoldResult> outputOffsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), outputColumn};
Value updated = tensor::InsertSliceOp::create(rewriter,
nestedLoc,
reduced,
iterArgs.front(),
outputOffsets,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(channels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1)},
getUnitStrides(rewriter, 4));
yielded.push_back(updated);
return success();
});
if (failed(outputLoop))
return failure();
insertRowStripFragment(
outputLoop->results.front(), args.outputs.front(), outputType, args.lane, rewriter, loc);
return success();
});
if (failed(batch))
return failure();
return batch->getResult(0);
}
void populatePoolPatterns(RewritePatternSet& patterns, MLIRContext* ctx) {
patterns.insert<PoolToSpatialCompute<ONNXMaxPoolSingleOutOp>>(ctx);
patterns.insert<PoolToSpatialCompute<ONNXAveragePoolOp>>(ctx);