Previous commit was broken in this one 9ms vs 7ms for vgg8 Arch-A
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user