relu_conv_relu Faster on Arch-A
This commit is contained in:
@@ -481,9 +481,12 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
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);
|
||||
FailureOr<RowStripPhysicalValue> physicalValue = describeRowStripPhysicalValue(input, inputType);
|
||||
const bool physicalInput = succeeded(physicalValue);
|
||||
if (!physicalInput && actualInputType != inputType)
|
||||
return failure();
|
||||
const int64_t tilesPerRow = physicalInput ? physicalValue->tilesPerRow : 1;
|
||||
const int64_t tileChannels = physicalInput ? physicalValue->fragmentType.getDimSize(3) : channels;
|
||||
|
||||
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value rowTable = createClampedPoolIndexTable(rewriter,
|
||||
@@ -502,51 +505,78 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
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 inputFragmentType =
|
||||
physicalInput ? physicalValue->fragmentType : getRowStripFragmentType(inputType);
|
||||
auto nchwInputFragmentType = RankedTensorType::get(
|
||||
{1, channels, 1, inputWidth}, inputType.getElementType(), inputType.getEncoding());
|
||||
auto outputFragmentType = RankedTensorType::get(
|
||||
{1, 1, outputWidth, tileChannels}, outputType.getElementType(), outputType.getEncoding());
|
||||
auto outputStorageType =
|
||||
spatial::getGraphBatchPhysicalResultType(outputHeight * tilesPerRow, outputFragmentType);
|
||||
auto tileType = RankedTensorType::get({1, 1, 1, tileChannels}, outputType.getElementType());
|
||||
auto batch = createSpatComputeBatch(
|
||||
rewriter,
|
||||
loc,
|
||||
TypeRange {outputStorageType},
|
||||
outputHeight,
|
||||
outputHeight * tilesPerRow,
|
||||
{},
|
||||
ValueRange {input},
|
||||
[&](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, args.lane, kernelRow, kernelHeight);
|
||||
extractPoolIndex(rewriter, loc, anchorOp, rowTable, outputRow, kernelRow, kernelHeight);
|
||||
if (physicalInput) {
|
||||
inputRows.push_back(
|
||||
extractRowStripFragment(args.inputs.front(), inputType, sourceRow, rewriter, loc));
|
||||
Value sourceSlot = sourceRow;
|
||||
if (tilesPerRow != 1) {
|
||||
sourceSlot = arith::AddIOp::create(
|
||||
rewriter,
|
||||
loc,
|
||||
arith::MulIOp::create(rewriter,
|
||||
loc,
|
||||
sourceRow,
|
||||
getOrCreateIndexConstant(rewriter, anchorOp, tilesPerRow)),
|
||||
channelTile);
|
||||
}
|
||||
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
|
||||
rewriter, loc, args.inputs.front(), sourceSlot, inputFragmentType);
|
||||
if (failed(fragment))
|
||||
return failure();
|
||||
inputRows.push_back(*fragment);
|
||||
}
|
||||
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)));
|
||||
Value nchw = tensor::ExtractSliceOp::create(rewriter,
|
||||
loc,
|
||||
nchwInputFragmentType,
|
||||
args.inputs.front(),
|
||||
offsets,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(channels),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(inputWidth)},
|
||||
getUnitStrides(rewriter, 4));
|
||||
inputRows.push_back(ONNXTransposeOp::create(
|
||||
rewriter, loc, inputFragmentType, nchw, rewriter.getI64ArrayAttr({0, 2, 3, 1})));
|
||||
}
|
||||
}
|
||||
|
||||
auto windowType = RankedTensorType::get(
|
||||
{1, channels, kernelHeight, inputWidth}, inputType.getElementType(), inputType.getEncoding());
|
||||
{1, kernelHeight, inputWidth, tileChannels}, 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),
|
||||
rewriter.getIndexAttr(0)};
|
||||
window = tensor::InsertSliceOp::create(rewriter,
|
||||
loc,
|
||||
@@ -554,9 +584,9 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
window,
|
||||
offsets,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(channels),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(inputWidth)},
|
||||
rewriter.getIndexAttr(inputWidth),
|
||||
rewriter.getIndexAttr(tileChannels)},
|
||||
getUnitStrides(rewriter, 4));
|
||||
}
|
||||
|
||||
@@ -585,43 +615,43 @@ FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
kernelColumn,
|
||||
kernelWidth);
|
||||
SmallVector<OpFoldResult> offsets {
|
||||
rewriter.getIndexAttr(0),
|
||||
rewriter.getIndexAttr(0),
|
||||
rewriter.getIndexAttr(kernelRow),
|
||||
sourceColumn};
|
||||
sourceColumn,
|
||||
rewriter.getIndexAttr(0)};
|
||||
Value point = tensor::ExtractSliceOp::create(rewriter,
|
||||
nestedLoc,
|
||||
tileType,
|
||||
window,
|
||||
offsets,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(channels),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(1)},
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(tileChannels)},
|
||||
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};
|
||||
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), outputColumn, rewriter.getIndexAttr(0)};
|
||||
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)},
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(tileChannels)},
|
||||
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);
|
||||
publishGraphBatchPhysicalFragment(
|
||||
rewriter, loc, outputLoop->results.front(), args.outputs.front(), args.lane);
|
||||
return success();
|
||||
});
|
||||
if (failed(batch))
|
||||
|
||||
Reference in New Issue
Block a user