Resnet is fast
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user