Previous commit was broken in this one 9ms vs 7ms for vgg8 Arch-A

This commit is contained in:
ilgeco
2026-07-28 12:09:59 +02:00
parent 87bd7b726d
commit 2b899b62a8
8 changed files with 518 additions and 210 deletions
@@ -448,6 +448,29 @@ static Value createClampedPoolIndexTable(PatternRewriter& rewriter,
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value createClampedPoolRowSlotTable(PatternRewriter& rewriter,
Operation* anchorOp,
int64_t outputHeight,
int64_t kernelHeight,
int64_t stride,
int64_t dilation,
int64_t padBegin,
int64_t inputHeight,
int64_t tilesPerRow) {
auto tableType =
RankedTensorType::get({outputHeight * tilesPerRow * kernelHeight}, rewriter.getIndexType());
SmallVector<Attribute> values;
values.reserve(tableType.getNumElements());
for (int64_t outputRow = 0; outputRow < outputHeight; ++outputRow)
for (int64_t tile = 0; tile < tilesPerRow; ++tile)
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
const int64_t sourceRow =
std::clamp(outputRow * stride + kernelRow * dilation - padBegin, int64_t {0}, inputHeight - 1);
values.push_back(rewriter.getIndexAttr(sourceRow * tilesPerRow + tile));
}
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value extractPoolIndex(PatternRewriter& rewriter,
Location loc,
Operation* anchorOp,
@@ -497,6 +520,15 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
planOp.getDilations()[0],
planOp.getPads()[0],
inputHeight);
Value rowSlotTable = createClampedPoolRowSlotTable(rewriter,
anchorOp,
outputHeight,
kernelHeight,
planOp.getStrides()[0],
planOp.getDilations()[0],
planOp.getPads()[0],
inputHeight,
tilesPerRow);
Value columnTable = createClampedPoolIndexTable(rewriter,
anchorOp,
outputWidth,
@@ -524,27 +556,10 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
[&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult {
SmallVector<Value> inputRows;
inputRows.reserve(kernelHeight);
Value outputRow = tilesPerRow == 1
? args.lane
: affineFloorDivConst(rewriter, loc, args.lane, tilesPerRow, anchorOp);
Value channelTile = tilesPerRow == 1
? getOrCreateIndexConstant(rewriter, anchorOp, 0)
: affineModConst(rewriter, loc, args.lane, tilesPerRow, anchorOp);
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
Value sourceRow =
extractPoolIndex(rewriter, loc, anchorOp, rowTable, outputRow, kernelRow, kernelHeight);
if (physicalInput) {
Value sourceSlot = sourceRow;
if (tilesPerRow != 1) {
sourceSlot = arith::AddIOp::create(
rewriter,
loc,
arith::MulIOp::create(rewriter,
loc,
sourceRow,
getOrCreateIndexConstant(rewriter, anchorOp, tilesPerRow)),
channelTile);
}
Value sourceSlot =
extractPoolIndex(rewriter, loc, anchorOp, rowSlotTable, args.lane, kernelRow, kernelHeight);
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs.front(), sourceSlot, inputFragmentType);
if (failed(fragment))
@@ -552,6 +567,8 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
inputRows.push_back(*fragment);
}
else {
Value sourceRow =
extractPoolIndex(rewriter, loc, anchorOp, rowTable, args.lane, kernelRow, kernelHeight);
SmallVector<OpFoldResult> offsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)};
Value nchw = tensor::ExtractSliceOp::create(rewriter,