fix vgg16 diamonds
Validate Operations / validate-operations (push) Waiting to run

updat ops validations
This commit is contained in:
NiccoloN
2026-07-22 17:54:09 +02:00
parent b491ff77b1
commit 45578ef4c4
29 changed files with 734 additions and 421 deletions
@@ -2326,54 +2326,25 @@ static Value maybeUnpackChunkRows(Value gemmRows,
return unpackCompute.getResult(0);
}
static Value createChunkedConvRows(const ConvLoweringState& state,
const PreparedConvInput& preparedInput,
Value weightMatrix,
Value biasMatrix,
DenseElementsAttr wDenseAttr,
DenseElementsAttr biasDenseAttr,
int64_t forcedPackFactor,
uint64_t chunkPositions,
PatternRewriter& rewriter,
Location loc) {
SmallVector<Value> chunkRows;
static Value createStreamedConvRows(const ConvLoweringState& state,
const PreparedConvInput& preparedInput,
Value weightMatrix,
Value biasMatrix,
DenseElementsAttr wDenseAttr,
DenseElementsAttr biasDenseAttr,
int64_t forcedPackFactor,
PatternRewriter& rewriter,
Location loc) {
const int64_t totalPatches = state.batchSize * state.outHeight * state.outWidth;
for (int64_t chunkStart = 0; chunkStart < totalPatches; chunkStart += static_cast<int64_t>(chunkPositions)) {
const int64_t chunkNumPatches = std::min<int64_t>(static_cast<int64_t>(chunkPositions), totalPatches - chunkStart);
ConvGemmPlan chunkPlan = buildConvGemmPlan(state,
static_cast<bool>(wDenseAttr),
!state.hasBias || static_cast<bool>(biasDenseAttr),
chunkStart,
chunkNumPatches,
forcedPackFactor);
Value chunkInputRows = createIm2colRows(state, preparedInput, chunkPlan, rewriter, loc);
Value chunkB = buildPackedWeights(wDenseAttr, weightMatrix, state, chunkPlan, rewriter, loc);
Value gemmBias = createZeroGemmBias(chunkPlan.gemmOutputRowsType, rewriter);
if (state.hasBias)
gemmBias = state.b;
Value chunkC = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, chunkPlan, rewriter, loc);
Value chunkGemmRows = ONNXGemmOp::create(rewriter,
loc,
chunkPlan.gemmOutputRowsType,
chunkInputRows,
chunkB,
chunkC,
APFloat(1.0f),
APFloat(1.0f),
/*transA=*/0,
/*transB=*/0)
.getY();
chunkRows.push_back(maybeUnpackChunkRows(chunkGemmRows, chunkPlan, rewriter, loc));
}
if (chunkRows.size() == 1)
return chunkRows.front();
auto rowType = RankedTensorType::get({totalPatches, state.numChannelsOut}, state.outType.getElementType());
auto collectRows = createSpatCompute(rewriter, loc, TypeRange {rowType}, {}, chunkRows, [&](ValueRange rows) {
spatial::SpatYieldOp::create(rewriter, loc, createSpatConcat(rewriter, loc, /*axis=*/0, rows));
});
return collectRows.getResult(0);
ConvGemmPlan plan = buildConvGemmPlan(state, static_cast<bool>(wDenseAttr),
!state.hasBias || static_cast<bool>(biasDenseAttr), 0, totalPatches, forcedPackFactor);
Value inputRows = createIm2colRows(state, preparedInput, plan, rewriter, loc);
Value packedWeights = buildPackedWeights(wDenseAttr, weightMatrix, state, plan, rewriter, loc);
Value gemmBias = state.hasBias ? state.b : createZeroGemmBias(plan.gemmOutputRowsType, rewriter);
Value packedBias = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, plan, rewriter, loc);
Value gemmRows = ONNXGemmOp::create(rewriter, loc, plan.gemmOutputRowsType, inputRows,
packedWeights, packedBias, APFloat(1.0f), APFloat(1.0f), 0, 0).getY();
return maybeUnpackChunkRows(gemmRows, plan, rewriter, loc);
}
static Value rewritePackedIm2ColConv(const ConvLoweringState& state,
@@ -2444,16 +2415,13 @@ static Value rewriteStreamedConv(const ConvLoweringState& state,
ConvGemmPlan seedPlan = buildConvGemmPlan(
state, static_cast<bool>(wDenseAttr), !state.hasBias || static_cast<bool>(biasDenseAttr), 0, 1, forcedPackFactor);
Value weightMatrix = createWeightMatrix(state.w, seedPlan, rewriter, loc);
ConvGeometry geo = buildConvGeometry(state);
uint64_t chunkPositions = chooseStreamChunkPositions(geo, forcedPackFactor);
Value collectedRows = createChunkedConvRows(state,
Value collectedRows = createStreamedConvRows(state,
preparedInput,
weightMatrix,
biasMatrix,
wDenseAttr,
biasDenseAttr,
forcedPackFactor,
chunkPositions,
rewriter,
loc);
auto gemmOutType = cast<RankedTensorType>(collectedRows.getType());
@@ -2524,6 +2492,21 @@ static bool canConsumeNchwRowStripFragments(const ConvLoweringState& state, Stri
failureReason = "dilation_not_one";
return false;
}
ConvGeometry geometry = buildConvGeometry(state);
const bool pointwise = state.xHeight == 1 && state.xWidth == 1 && state.outHeight == 1 && state.outWidth == 1
&& state.wHeight == 1 && state.wWidth == 1 && state.padHeightBegin == 0
&& state.padHeightEnd == 0 && state.padWidthBegin == 0 && state.padWidthEnd == 0;
if (pointwise) {
if (!getHostConstDenseElementsAttr(state.w)) {
failureReason = "non_constant_weight";
return false;
}
if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)) {
failureReason = "unsupported_bias";
return false;
}
return true;
}
if (state.wHeight != 3 || state.wWidth != 3) {
failureReason = "kernel_not_3x3";
return false;
@@ -2544,7 +2527,6 @@ static bool canConsumeNchwRowStripFragments(const ConvLoweringState& state, Stri
failureReason = "unsupported_bias";
return false;
}
ConvGeometry geometry = buildConvGeometry(state);
if (geometry.c > geometry.xbarSize) {
failureReason = "output_channels_exceed_crossbar";
return false;
@@ -2575,6 +2557,25 @@ static FailureOr<Value> createPaddedBiasRowConstant(const ConvLoweringState& sta
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), biasAttr, biasType);
}
static FailureOr<Value> createPaddedBiasTileConstant(const ConvLoweringState& state,
int64_t tileChannels,
PatternRewriter& rewriter) {
DenseElementsAttr denseAttr;
if (!isSupportedBiasAddValue(state.b, state.outType, &denseAttr))
return failure();
FailureOr<SmallVector<Attribute>> channelValues = getBiasChannelValues(denseAttr, state.outType);
if (failed(channelValues))
return failure();
const int64_t tileCount = ceilIntegerDivide(state.numChannelsOut, tileChannels);
auto tileType = RankedTensorType::get({tileCount, 1, tileChannels}, state.outType.getElementType());
SmallVector<Attribute> values(
tileType.getNumElements(), cast<Attribute>(rewriter.getZeroAttr(tileType.getElementType())));
for (int64_t channel = 0; channel < state.numChannelsOut; ++channel)
values[channel] = (*channelValues)[channel];
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(),
DenseElementsAttr::get(tileType, values), tileType);
}
static Value createHorizontallyPaddedRowStripFragment(Value fragment,
const ConvLoweringState& state,
PatternRewriter& rewriter,
@@ -2732,8 +2733,11 @@ static FailureOr<Value> createConvInputWindow(Value input,
? extractDenseConvWindowRow(input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc)
: extractProjectedRowStripWindowRow(
input, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc);
Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, loc);
Value semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, sourceRow, mask).getResult();
Value semanticRow = sourceRow;
if (state.padHeightBegin != 0 || state.padHeightEnd != 0) {
Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, loc);
semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, sourceRow, mask).getResult();
}
Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, loc);
window = tensor::InsertSliceOp::create(rewriter,
loc,
@@ -2906,12 +2910,6 @@ static FailureOr<Value> createPaddedConvOutputRow(Value patchRow,
.getResult();
}
static bool rowStripOutputFitsOneCore(const ConvGeometry& geometry) {
const int64_t inputTileCount = ceilIntegerDivide(geometry.k, geometry.xbarSize);
const int64_t outputTileCount = ceilIntegerDivide(geometry.c, geometry.xbarSize);
return inputTileCount * outputTileCount <= static_cast<int64_t>(crossbarCountInCore.getValue());
}
static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) {
return ceilIntegerDivide(geometry.k, geometry.xbarSize)
<= static_cast<int64_t>(crossbarCountInCore.getValue());
@@ -2928,38 +2926,40 @@ static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLow
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
auto elementType = state.outType.getElementType();
auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType);
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto tilePixelType = RankedTensorType::get({1, xbarDim, 1, 1}, elementType);
auto tileFragmentType = RankedTensorType::get({1, xbarDim, 1, state.outWidth}, elementType);
auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, state.wType.getElementType());
SmallVector<Value> outputTiles;
outputTiles.reserve(outputTileCount);
for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) {
const int64_t channelOffset = outputTile * xbarDim;
const int64_t tileChannels = std::min(xbarDim, state.numChannelsOut - channelOffset);
auto tileRowType = RankedTensorType::get({1, tileChannels}, elementType);
auto tilePixelType = RankedTensorType::get({1, tileChannels, 1, 1}, elementType);
auto tileFragmentType = RankedTensorType::get({1, tileChannels, 1, state.outWidth}, elementType);
auto tileStorageType = spatial::getGraphBatchPhysicalResultType(state.outHeight, tileFragmentType);
SmallVector<OpFoldResult> weightOffsets {
rewriter.getIndexAttr(outputTile), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> weightSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)};
Value tileWeights = extractStaticSliceOrIdentity(
rewriter, loc, paddedWeights, tileWeightsType, weightOffsets, weightSizes, getUnitStrides(rewriter, 3));
auto tileBatch = createSpatComputeBatch(
rewriter,
loc,
TypeRange {tileStorageType},
state.outHeight,
ValueRange {tileWeights},
ValueRange {state.x},
[&](detail::SpatComputeBatchBodyArgs args) {
const int64_t laneCount = state.outHeight * outputTileCount;
auto tileStorageType = spatial::getGraphBatchPhysicalResultType(laneCount, tileFragmentType);
FailureOr<Value> paddedBias = failure();
if (state.hasBias)
paddedBias = createPaddedBiasTileConstant(state, xbarDim, rewriter);
if (state.hasBias && failed(paddedBias))
return failure();
auto tileBatch = createSpatComputeBatch(
rewriter, loc, TypeRange {tileStorageType}, laneCount, ValueRange {paddedWeights},
state.hasBias ? ValueRange {state.x, *paddedBias} : ValueRange {state.x},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth);
Value outputRow = affineFloorDivConst(rewriter, loc, args.lane, outputTileCount, anchorOp);
Value outputTile = affineModConst(rewriter, loc, args.lane, outputTileCount, anchorOp);
SmallVector<OpFoldResult> weightOffsets {
outputTile, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> weightSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)};
Value tileWeights = tensor::ExtractSliceOp::create(
rewriter, loc, tileWeightsType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3));
FailureOr<Value> biasTile = failure();
if (state.hasBias)
biasTile = extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[1], outputTile, paddedRowType);
if (state.hasBias && failed(biasTile))
return failure();
FailureOr<Value> inputWindow =
createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc);
createConvInputWindow(args.inputs.front(), state, outputRow, rewriter, loc);
if (failed(inputWindow))
return failure();
Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType);
@@ -2984,28 +2984,18 @@ static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLow
paddedPatchRow = createZeroPaddedTensor(
paddedPatchRow, paddedPatchRowType, {0, 0}, {0, paddedK - patchSize}, rewriter, widthLoc);
FailureOr<Value> paddedOutputRow = createPaddedConvOutputTile(
paddedPatchRow, args.weights.front(), numKSlices, xbarDim, rewriter, widthLoc);
paddedPatchRow, tileWeights, numKSlices, xbarDim, rewriter, widthLoc);
if (failed(paddedOutputRow))
return failure();
Value outputRow = *paddedOutputRow;
if (tileChannels != xbarDim) {
SmallVector<OpFoldResult> rowOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> rowSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels)};
outputRow = tensor::ExtractSliceOp::create(rewriter,
widthLoc,
tileRowType,
outputRow,
rowOffsets,
rowSizes,
getUnitStrides(rewriter, 2));
}
if (state.hasBias)
paddedOutputRow = spatial::SpatVAddOp::create(
rewriter, widthLoc, paddedRowType, *paddedOutputRow, *biasTile).getResult();
Value outputPixel = tensor::ExpandShapeOp::create(
rewriter, widthLoc, tilePixelType, outputRow, SmallVector<ReassociationIndices> {{0}, {1, 2, 3}});
rewriter, widthLoc, tilePixelType, *paddedOutputRow, SmallVector<ReassociationIndices> {{0}, {1, 2, 3}});
SmallVector<OpFoldResult> rowOffsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), widthIndex};
SmallVector<OpFoldResult> rowSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tileChannels),
rewriter.getIndexAttr(xbarDim),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1)};
Value nextFragment = tensor::InsertSliceOp::create(rewriter,
@@ -3024,58 +3014,9 @@ static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLow
rewriter, loc, widthLoop->results.front(), args.outputs.front(), args.lane);
return success();
});
if (failed(tileBatch))
return failure();
outputTiles.push_back(tileBatch->getResult(0));
}
auto fragmentType = getRowStripFragmentType(state.outType);
auto outputStorageType = getRowStripStorageType(state.outType);
auto assemblyBatch = createSpatComputeBatch(rewriter,
loc,
TypeRange {outputStorageType},
state.outHeight,
{},
ValueRange(outputTiles),
[&](detail::SpatComputeBatchBodyArgs args) {
Value fragment = tensor::EmptyOp::create(
rewriter, loc, fragmentType.getShape(), elementType);
for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) {
const int64_t channelOffset = outputTile * xbarDim;
const int64_t tileChannels =
std::min(xbarDim, state.numChannelsOut - channelOffset);
auto tileFragmentType = RankedTensorType::get(
{1, tileChannels, 1, state.outWidth}, elementType);
FailureOr<Value> tileFragment = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs[outputTile], args.lane, tileFragmentType);
if (failed(tileFragment))
return failure();
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(channelOffset),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tileChannels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.outWidth)};
fragment = tensor::InsertSliceOp::create(rewriter,
loc,
*tileFragment,
fragment,
offsets,
sizes,
getUnitStrides(rewriter, 4));
}
insertRowStripFragment(
fragment, args.outputs.front(), state.outType, args.lane, rewriter, loc);
return success();
});
if (failed(assemblyBatch))
if (failed(tileBatch))
return failure();
Value output = assemblyBatch->getResult(0);
if (state.hasBias)
return applyRowStripBiasAdd(output, state.outType, state.b, rewriter, loc);
return output;
return tileBatch->getResult(0);
}
static FailureOr<Value>
@@ -3104,7 +3045,7 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe
weightDenseAttr, state, paddedK, xbarDim, rewriter)
: standard::createPaddedOutputChannelTiledWeightConstant(
weightDenseAttr, state, paddedK, xbarDim, rewriter);
if (!rowStripOutputFitsOneCore(geometry))
if (state.numChannelsOut > xbarDim)
return createOutputChannelTiledRowStripConvOutput(
state, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc);
@@ -3279,11 +3220,106 @@ static FailureOr<Value> createConvOutputFromNchwRowStripFragments(Value rowStrip
return batchOp->getResult(0);
}
static FailureOr<Value> createPointwiseOutputFromRowStripFragments(Value rowStripStorage,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
FailureOr<RowStripPhysicalValue> input = describeRowStripPhysicalValue(rowStripStorage, state.xType);
if (failed(input)) return failure();
ConvGeometry geometry = buildConvGeometry(state);
const int64_t xbarDim = geometry.xbarSize;
const int64_t inputFragmentChannels = input->fragmentType.getDimSize(1);
if (inputFragmentChannels % xbarDim != 0 || state.numChannelsIn % xbarDim != 0)
return failure();
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!weightDenseAttr) return failure();
const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim);
const int64_t numKSlices = state.numChannelsIn / xbarDim;
auto elementType = state.outType.getElementType();
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto inputRowType = RankedTensorType::get({1, inputFragmentChannels}, elementType);
auto weightTileType = RankedTensorType::get({state.numChannelsIn, xbarDim}, state.wType.getElementType());
auto weightSliceType = RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType());
auto outputFragmentType = RankedTensorType::get({1, xbarDim, 1, 1}, elementType);
auto outputStorageType = spatial::getGraphBatchPhysicalResultType(outputTileCount, outputFragmentType);
Value paddedWeights = standard::createPaddedOutputChannelTiledWeightConstant(
weightDenseAttr, state, state.numChannelsIn, xbarDim, rewriter);
FailureOr<Value> paddedBias = failure();
if (state.hasBias) paddedBias = createPaddedBiasTileConstant(state, xbarDim, rewriter);
if (state.hasBias && failed(paddedBias)) return failure();
auto batch = createSpatComputeBatch(rewriter, loc, TypeRange {outputStorageType}, outputTileCount,
ValueRange {paddedWeights},
state.hasBias ? ValueRange {rowStripStorage, *paddedBias} : ValueRange {rowStripStorage},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices);
SmallVector<OpFoldResult> weightOffsets {args.lane, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> weightSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.numChannelsIn), rewriter.getIndexAttr(xbarDim)};
Value weightTile = tensor::ExtractSliceOp::create(
rewriter, loc, weightTileType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3));
auto createPiece = [&](Value kSlice, Location pieceLoc) -> FailureOr<Value> {
Value channelOffset = affineMulConst(rewriter, pieceLoc, kSlice, xbarDim, anchorOp);
Value sourceSlot = affineFloorDivConst(
rewriter, pieceLoc, channelOffset, inputFragmentChannels, anchorOp);
Value sourceOffset = affineModConst(
rewriter, pieceLoc, channelOffset, inputFragmentChannels, anchorOp);
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
rewriter, pieceLoc, args.inputs.front(), sourceSlot, input->fragmentType);
if (failed(fragment)) return failure();
Value inputRow = tensor::CollapseShapeOp::create(rewriter, pieceLoc, inputRowType, *fragment,
SmallVector<ReassociationIndices> {{0}, {1, 2, 3}});
Value inputSlice = tensor::ExtractSliceOp::create(rewriter, pieceLoc, paddedRowType, inputRow,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), sourceOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 2));
Value weightSlice = tensor::ExtractSliceOp::create(rewriter, pieceLoc, weightSliceType, weightTile,
SmallVector<OpFoldResult> {channelOffset, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 2));
return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, weightSlice, inputSlice).getResult();
};
FailureOr<Value> result = createPiece(c0, loc);
if (failed(result)) return failure();
if (numKSlices > 1) {
auto reduction = buildNormalizedScfFor(rewriter, loc, c1, cNumKSlices, c1, ValueRange {*result},
[&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
FailureOr<Value> piece = createPiece(kSlice, reduceLoc);
if (failed(piece)) return failure();
yielded.push_back(spatial::SpatVAddOp::create(
rewriter, reduceLoc, paddedRowType, iterArgs.front(), *piece).getResult());
return success();
});
if (failed(reduction)) return failure();
result = reduction->results.front();
}
if (state.hasBias) {
FailureOr<Value> bias = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs[1], args.lane, paddedRowType);
if (failed(bias)) return failure();
result = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, *result, *bias).getResult();
}
Value fragment = tensor::ExpandShapeOp::create(rewriter, loc, outputFragmentType, *result,
SmallVector<ReassociationIndices> {{0}, {1, 2, 3}});
publishGraphBatchPhysicalFragment(rewriter, loc, fragment, args.outputs.front(), args.lane);
return success();
});
if (failed(batch)) return failure();
return batch->getResult(0);
}
static FailureOr<Value> createConvOutputFromRowStripInput(const ConvLoweringState& state,
[[maybe_unused]] const ConvLoweringDecision& decision,
Value rowStripInput,
PatternRewriter& rewriter,
Location loc) {
if (state.xHeight == 1 && state.xWidth == 1 && state.wHeight == 1 && state.wWidth == 1)
return createPointwiseOutputFromRowStripFragments(rowStripInput, state, rewriter, loc);
return createConvOutputFromNchwRowStripFragments(rowStripInput, state, rewriter, loc);
}
@@ -4187,22 +4223,9 @@ lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp,
}
if (failed(canLowerConvPlanToRowStrip(planOp)))
return planOp.emitOpError("selected row-strip layout is not supported for this Conv plan"), failure();
ConvLoweringState rowState = *state;
const bool applyBiasAfterStorage = rowState.hasBias;
Value originalBias = rowState.b;
if (applyBiasAfterStorage) {
rowState.b = Value();
rowState.hasBias = false;
}
FailureOr<Value> rowStripStorage = createRowStripConvOutputFromDenseInput(rowState, rewriter, planOp.getLoc());
FailureOr<Value> rowStripStorage = createRowStripConvOutputFromDenseInput(*state, rewriter, planOp.getLoc());
if (failed(rowStripStorage))
return planOp.emitOpError("failed to build row-strip fragment storage for the selected Conv plan"), failure();
if (applyBiasAfterStorage) {
rowStripStorage = applyRowStripBiasAdd(*rowStripStorage, state->outType, originalBias, rewriter, planOp.getLoc());
if (failed(rowStripStorage))
return planOp.emitOpError("failed to apply row-strip Conv bias per fragment"), failure();
}
return *rowStripStorage;
}