relu_conv_relu Faster on Arch-A

This commit is contained in:
ilgeco
2026-07-28 11:16:28 +02:00
parent e76c278d1c
commit e4d8fba579
9 changed files with 446 additions and 257 deletions
@@ -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))