Resnet is fast
This commit is contained in:
@@ -1890,6 +1890,37 @@ static Value createPaddedInputKTiledWeightConstant(DenseElementsAttr sourceAttr,
|
||||
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType);
|
||||
}
|
||||
|
||||
static Value createPaddedOutputChannelTiledWeightConstant(DenseElementsAttr sourceAttr,
|
||||
const ConvLoweringState& state,
|
||||
int64_t paddedK,
|
||||
int64_t xbarDim,
|
||||
PatternRewriter& rewriter) {
|
||||
const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim);
|
||||
auto paddedType =
|
||||
RankedTensorType::get({outputTileCount, paddedK, xbarDim}, state.wType.getElementType());
|
||||
SmallVector<Attribute> sourceValues(sourceAttr.getValues<Attribute>());
|
||||
SmallVector<Attribute> paddedValues(
|
||||
paddedType.getNumElements(), cast<Attribute>(rewriter.getZeroAttr(paddedType.getElementType())));
|
||||
for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel) {
|
||||
const int64_t outputTile = outChannel / xbarDim;
|
||||
const int64_t tileChannel = outChannel % xbarDim;
|
||||
for (int64_t inChannel = 0; inChannel < state.numChannelsIn; ++inChannel) {
|
||||
for (int64_t kernelH = 0; kernelH < state.wHeight; ++kernelH) {
|
||||
for (int64_t kernelW = 0; kernelW < state.wWidth; ++kernelW) {
|
||||
const int64_t sourceFlatIndex =
|
||||
(((outChannel * state.numChannelsIn) + inChannel) * state.wHeight + kernelH) * state.wWidth + kernelW;
|
||||
const int64_t patchIndex = ((inChannel * state.wHeight) + kernelH) * state.wWidth + kernelW;
|
||||
const int64_t destinationFlatIndex =
|
||||
((outputTile * paddedK) + patchIndex) * xbarDim + tileChannel;
|
||||
paddedValues[destinationFlatIndex] = sourceValues[sourceFlatIndex];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
auto paddedAttr = DenseElementsAttr::get(paddedType, paddedValues);
|
||||
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType);
|
||||
}
|
||||
|
||||
static FailureOr<Value> rewriteInputKTiledConv(const ConvLoweringState& state,
|
||||
ArrayRef<DistributedTensorStep> distributedConsumers,
|
||||
PatternRewriter& rewriter,
|
||||
@@ -2541,10 +2572,16 @@ static Value createHorizontallyPaddedRowStripFragment(Value fragment,
|
||||
const ConvLoweringState& state,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
auto paddedType = RankedTensorType::get({1, state.numChannelsIn, 1, state.xWidth + 2},
|
||||
state.xType.getElementType(),
|
||||
state.xType.getEncoding());
|
||||
return createZeroPaddedTensor(fragment, paddedType, {0, 0, 0, 1}, {0, 0, 0, 1}, rewriter, loc);
|
||||
auto paddedType = RankedTensorType::get(
|
||||
{1, state.numChannelsIn, 1, state.xWidth + state.padWidthBegin + state.padWidthEnd},
|
||||
state.xType.getElementType(),
|
||||
state.xType.getEncoding());
|
||||
return createZeroPaddedTensor(fragment,
|
||||
paddedType,
|
||||
{0, 0, 0, state.padWidthBegin},
|
||||
{0, 0, 0, state.padWidthEnd},
|
||||
rewriter,
|
||||
loc);
|
||||
}
|
||||
|
||||
static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state, PatternRewriter& rewriter) {
|
||||
@@ -2554,7 +2591,8 @@ static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state,
|
||||
values.reserve(tableType.getNumElements());
|
||||
for (int64_t outputRow = 0; outputRow < state.outHeight; ++outputRow) {
|
||||
for (int64_t kernelRow = 0; kernelRow < state.wHeight; ++kernelRow) {
|
||||
int64_t sourceRow = outputRow + kernelRow - state.padHeightBegin;
|
||||
int64_t sourceRow =
|
||||
outputRow * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin;
|
||||
sourceRow = std::clamp(sourceRow, int64_t {0}, state.xHeight - 1);
|
||||
values.push_back(rewriter.getIndexAttr(sourceRow));
|
||||
}
|
||||
@@ -2588,6 +2626,26 @@ static Value extractProjectedRowStripWindowRow(Value rowStripStorage,
|
||||
return extractRowStripFragment(rowStripStorage, state.xType, sourceRow, rewriter, loc);
|
||||
}
|
||||
|
||||
static Value extractDenseConvWindowRow(Value denseInput,
|
||||
Value sourceRowTable,
|
||||
const ConvLoweringState& state,
|
||||
Value outputHeight,
|
||||
Value kernelRow,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
Value tableIndex = createRowStripWindowTableIndex(outputHeight, kernelRow, state, rewriter, loc);
|
||||
Value sourceRow = tensor::ExtractOp::create(rewriter, loc, sourceRowTable, ValueRange {tableIndex}).getResult();
|
||||
auto fragmentType = getRowStripFragmentType(state.xType);
|
||||
SmallVector<OpFoldResult> offsets {
|
||||
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceRow, rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(state.numChannelsIn),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(state.xWidth)};
|
||||
return tensor::ExtractSliceOp::create(
|
||||
rewriter, loc, fragmentType, denseInput, offsets, sizes, getUnitStrides(rewriter, 4));
|
||||
}
|
||||
|
||||
static FailureOr<Value> createRowStripWindowMaskTable(const ConvLoweringState& state, PatternRewriter& rewriter) {
|
||||
auto elementType = state.xType.getElementType();
|
||||
auto floatType = dyn_cast<FloatType>(elementType);
|
||||
@@ -2604,7 +2662,8 @@ static FailureOr<Value> createRowStripWindowMaskTable(const ConvLoweringState& s
|
||||
values.reserve(tableType.getNumElements());
|
||||
for (int64_t outputRow = 0; outputRow < state.outHeight; ++outputRow) {
|
||||
for (int64_t kernelRow = 0; kernelRow < state.wHeight; ++kernelRow) {
|
||||
int64_t sourceRow = outputRow + kernelRow - state.padHeightBegin;
|
||||
int64_t sourceRow =
|
||||
outputRow * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin;
|
||||
Attribute value = (sourceRow < 0 || sourceRow >= state.xHeight) ? zero : one;
|
||||
for (int64_t channel = 0; channel < state.numChannelsIn; ++channel)
|
||||
for (int64_t width = 0; width < state.xWidth; ++width)
|
||||
@@ -2638,15 +2697,20 @@ static Value extractProjectedRowStripWindowMask(Value maskTable,
|
||||
getUnitStrides(rewriter, 4));
|
||||
}
|
||||
|
||||
static FailureOr<Value> createNchwRowStripConvWindow(Value rowStripStorage,
|
||||
const ConvLoweringState& state,
|
||||
Value outputHeight,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
static FailureOr<Value> createConvInputWindow(Value input,
|
||||
const ConvLoweringState& state,
|
||||
Value outputHeight,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
auto fragmentType = getRowStripFragmentType(state.xType);
|
||||
auto paddedWindowType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.xWidth + 2},
|
||||
state.xType.getElementType(),
|
||||
state.xType.getEncoding());
|
||||
auto inputType = dyn_cast<RankedTensorType>(input.getType());
|
||||
const bool denseInput = inputType == state.xType;
|
||||
if (!denseInput && inputType != getRowStripStorageType(state.xType))
|
||||
return failure();
|
||||
auto paddedWindowType = RankedTensorType::get(
|
||||
{1, state.numChannelsIn, state.wHeight, state.xWidth + state.padWidthBegin + state.padWidthEnd},
|
||||
state.xType.getElementType(),
|
||||
state.xType.getEncoding());
|
||||
Value sourceRowTable = createRowStripWindowSourceRowTable(state, rewriter);
|
||||
FailureOr<Value> maskTable = createRowStripWindowMaskTable(state, rewriter);
|
||||
if (failed(maskTable))
|
||||
@@ -2657,8 +2721,10 @@ static FailureOr<Value> createNchwRowStripConvWindow(Value rowStripStorage,
|
||||
Value window = initWindow;
|
||||
for (int64_t kernelRowIndex = 0; kernelRowIndex < state.wHeight; ++kernelRowIndex) {
|
||||
Value kernelRow = getOrCreateIndexConstant(rewriter, anchorOp, kernelRowIndex);
|
||||
Value sourceRow =
|
||||
extractProjectedRowStripWindowRow(rowStripStorage, sourceRowTable, state, outputHeight, kernelRow, rewriter, loc);
|
||||
Value sourceRow = denseInput
|
||||
? 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 paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, loc);
|
||||
@@ -2673,7 +2739,9 @@ static FailureOr<Value> createNchwRowStripConvWindow(Value rowStripStorage,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(state.numChannelsIn),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(state.xWidth + 2)},
|
||||
rewriter.getIndexAttr(
|
||||
state.xWidth + state.padWidthBegin
|
||||
+ state.padWidthEnd)},
|
||||
getUnitStrides(rewriter, 4));
|
||||
}
|
||||
return window;
|
||||
@@ -2691,12 +2759,13 @@ static FailureOr<Value> createNchwRowStripConvPatchRow(Value paddedWindow,
|
||||
state.xType.getEncoding());
|
||||
auto rowType = RankedTensorType::get({1, patchSize}, state.xType.getElementType(), state.xType.getEncoding());
|
||||
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
|
||||
Value inputWidthOffset = affineMulConst(rewriter, loc, outputWidth, state.strideWidth, anchorOp);
|
||||
Value patch = createConvInputPatch(paddedWindow,
|
||||
patchType,
|
||||
c0,
|
||||
c0,
|
||||
c0,
|
||||
outputWidth,
|
||||
inputWidthOffset,
|
||||
state.dilationHeight,
|
||||
state.dilationWidth,
|
||||
rewriter,
|
||||
@@ -2706,6 +2775,57 @@ static FailureOr<Value> createNchwRowStripConvPatchRow(Value paddedWindow,
|
||||
.getResult();
|
||||
}
|
||||
|
||||
static FailureOr<Value> createPaddedConvOutputTile(Value paddedPatchRow,
|
||||
Value tileWeights,
|
||||
int64_t numKSlices,
|
||||
int64_t xbarDim,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
auto elementType = cast<RankedTensorType>(paddedPatchRow.getType()).getElementType();
|
||||
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
|
||||
auto weightElementType = cast<RankedTensorType>(tileWeights.getType()).getElementType();
|
||||
auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType);
|
||||
|
||||
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
|
||||
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
|
||||
Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices);
|
||||
Value cXbar = getOrCreateIndexConstant(rewriter, anchorOp, xbarDim);
|
||||
auto createPiece = [&](Value kSlice, Location pieceLoc) -> Value {
|
||||
Value kOffset = arith::MulIOp::create(rewriter, pieceLoc, kSlice, cXbar);
|
||||
SmallVector<OpFoldResult> aOffsets {rewriter.getIndexAttr(0), kOffset};
|
||||
SmallVector<OpFoldResult> aSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)};
|
||||
Value aTile = extractStaticSliceOrIdentity(
|
||||
rewriter, pieceLoc, paddedPatchRow, paddedRowType, aOffsets, aSizes, getUnitStrides(rewriter, 2));
|
||||
SmallVector<OpFoldResult> bOffsets {kOffset, rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)};
|
||||
Value bTile = extractStaticSliceOrIdentity(
|
||||
rewriter, pieceLoc, tileWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2));
|
||||
return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, bTile, aTile).getResult();
|
||||
};
|
||||
|
||||
Value tileResult = createPiece(c0, loc);
|
||||
if (numKSlices == 1)
|
||||
return tileResult;
|
||||
|
||||
auto kLoop = buildNormalizedScfFor(
|
||||
rewriter,
|
||||
loc,
|
||||
c1,
|
||||
cNumKSlices,
|
||||
c1,
|
||||
ValueRange {tileResult},
|
||||
[&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange reduceIterArgs, SmallVectorImpl<Value>& reduceYielded) {
|
||||
Value piece = createPiece(kSlice, reduceLoc);
|
||||
reduceYielded.push_back(
|
||||
spatial::SpatVAddOp::create(rewriter, reduceLoc, paddedRowType, reduceIterArgs.front(), piece).getResult());
|
||||
return success();
|
||||
});
|
||||
if (failed(kLoop))
|
||||
return failure();
|
||||
return kLoop->results.front();
|
||||
}
|
||||
|
||||
static FailureOr<Value> createPaddedConvOutputRow(Value patchRow,
|
||||
const ConvLoweringState& state,
|
||||
Value paddedWeights,
|
||||
@@ -2720,67 +2840,241 @@ static FailureOr<Value> createPaddedConvOutputRow(Value patchRow,
|
||||
auto rowType = RankedTensorType::get({1, state.numChannelsOut}, elementType);
|
||||
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
|
||||
auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, elementType);
|
||||
auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType());
|
||||
auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, state.wType.getElementType());
|
||||
const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim);
|
||||
|
||||
Value paddedPatchRow = patchRow;
|
||||
if (patchSize != paddedK)
|
||||
paddedPatchRow = createZeroPaddedTensor(
|
||||
paddedPatchRow, paddedPatchRowType, {0, 0}, {0, paddedK - patchSize}, rewriter, loc);
|
||||
|
||||
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
|
||||
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
|
||||
Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices);
|
||||
Value cXbar = getOrCreateIndexConstant(rewriter, anchorOp, xbarDim);
|
||||
auto createPiece = [&](Value kSlice, Location pieceLoc) -> Value {
|
||||
Value kOffset = arith::MulIOp::create(rewriter, pieceLoc, kSlice, cXbar);
|
||||
SmallVector<OpFoldResult> aOffsets {rewriter.getIndexAttr(0), kOffset};
|
||||
SmallVector<OpFoldResult> aSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)};
|
||||
Value aTile = extractStaticSliceOrIdentity(
|
||||
rewriter, pieceLoc, paddedPatchRow, paddedRowType, aOffsets, aSizes, getUnitStrides(rewriter, 2));
|
||||
SmallVector<OpFoldResult> bOffsets {kOffset, rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)};
|
||||
Value bTile = extractStaticSliceOrIdentity(
|
||||
rewriter, pieceLoc, paddedWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2));
|
||||
return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, bTile, aTile).getResult();
|
||||
auto getTileWeights = [&](int64_t outputTile) {
|
||||
if (outputTileCount == 1)
|
||||
return paddedWeights;
|
||||
SmallVector<OpFoldResult> offsets {
|
||||
rewriter.getIndexAttr(outputTile), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> sizes {
|
||||
rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)};
|
||||
return extractStaticSliceOrIdentity(
|
||||
rewriter, loc, paddedWeights, tileWeightsType, offsets, sizes, getUnitStrides(rewriter, 3));
|
||||
};
|
||||
|
||||
Value rowResult = createPiece(c0, loc);
|
||||
if (numKSlices > 1) {
|
||||
auto kLoop = buildNormalizedScfFor(
|
||||
rewriter,
|
||||
loc,
|
||||
c1,
|
||||
cNumKSlices,
|
||||
c1,
|
||||
ValueRange {rowResult},
|
||||
[&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange reduceIterArgs, SmallVectorImpl<Value>& reduceYielded) {
|
||||
Value piece = createPiece(kSlice, reduceLoc);
|
||||
reduceYielded.push_back(
|
||||
spatial::SpatVAddOp::create(rewriter, reduceLoc, paddedRowType, reduceIterArgs.front(), piece).getResult());
|
||||
return success();
|
||||
});
|
||||
if (failed(kLoop))
|
||||
if (outputTileCount == 1) {
|
||||
FailureOr<Value> rowResult = createPaddedConvOutputTile(
|
||||
paddedPatchRow, getTileWeights(0), numKSlices, xbarDim, rewriter, loc);
|
||||
if (failed(rowResult))
|
||||
return failure();
|
||||
rowResult = kLoop->results.front();
|
||||
if (paddedBias)
|
||||
rowResult = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, *rowResult, paddedBias).getResult();
|
||||
if (state.numChannelsOut == xbarDim)
|
||||
return *rowResult;
|
||||
|
||||
SmallVector<OpFoldResult> outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> outputSizes {
|
||||
rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)};
|
||||
return tensor::ExtractSliceOp::create(
|
||||
rewriter, loc, rowType, *rowResult, outputOffsets, outputSizes, getUnitStrides(rewriter, 2))
|
||||
.getResult();
|
||||
}
|
||||
|
||||
if (paddedBias)
|
||||
rowResult = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, rowResult, paddedBias).getResult();
|
||||
if (state.numChannelsOut == xbarDim)
|
||||
return rowResult;
|
||||
const int64_t paddedOutputChannels = outputTileCount * xbarDim;
|
||||
auto paddedOutputType = RankedTensorType::get({1, paddedOutputChannels}, elementType);
|
||||
Value paddedOutput = tensor::EmptyOp::create(rewriter, loc, paddedOutputType.getShape(), elementType);
|
||||
for (int64_t outputTile = 0; outputTile < outputTileCount; ++outputTile) {
|
||||
FailureOr<Value> tileResult = createPaddedConvOutputTile(
|
||||
paddedPatchRow, getTileWeights(outputTile), numKSlices, xbarDim, rewriter, loc);
|
||||
if (failed(tileResult))
|
||||
return failure();
|
||||
SmallVector<OpFoldResult> tileOffsets {
|
||||
rewriter.getIndexAttr(0), rewriter.getIndexAttr(outputTile * xbarDim)};
|
||||
SmallVector<OpFoldResult> tileSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)};
|
||||
paddedOutput = tensor::InsertSliceOp::create(
|
||||
rewriter, loc, *tileResult, paddedOutput, tileOffsets, tileSizes, getUnitStrides(rewriter, 2));
|
||||
}
|
||||
|
||||
SmallVector<OpFoldResult> outputOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
|
||||
SmallVector<OpFoldResult> outputSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)};
|
||||
return tensor::ExtractSliceOp::create(
|
||||
rewriter, loc, rowType, rowResult, outputOffsets, outputSizes, getUnitStrides(rewriter, 2))
|
||||
rewriter, loc, rowType, paddedOutput, outputOffsets, outputSizes, getUnitStrides(rewriter, 2))
|
||||
.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());
|
||||
}
|
||||
|
||||
static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLoweringState& state,
|
||||
Value paddedWeights,
|
||||
int64_t paddedK,
|
||||
int64_t numKSlices,
|
||||
int64_t xbarDim,
|
||||
PatternRewriter& rewriter,
|
||||
Location loc) {
|
||||
const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim);
|
||||
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
|
||||
auto elementType = state.outType.getElementType();
|
||||
auto paddedPatchRowType = RankedTensorType::get({1, paddedK}, 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) {
|
||||
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
|
||||
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
|
||||
Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth);
|
||||
FailureOr<Value> inputWindow =
|
||||
createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc);
|
||||
if (failed(inputWindow))
|
||||
return failure();
|
||||
Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType);
|
||||
auto widthLoop = buildNormalizedScfFor(
|
||||
rewriter,
|
||||
loc,
|
||||
c0,
|
||||
cOutWidth,
|
||||
c1,
|
||||
ValueRange {fragmentInit},
|
||||
[&](OpBuilder&,
|
||||
Location widthLoc,
|
||||
Value widthIndex,
|
||||
ValueRange widthIterArgs,
|
||||
SmallVectorImpl<Value>& widthYielded) {
|
||||
FailureOr<Value> patchRow =
|
||||
createNchwRowStripConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc);
|
||||
if (failed(patchRow))
|
||||
return failure();
|
||||
Value paddedPatchRow = *patchRow;
|
||||
if (patchSize != paddedK)
|
||||
paddedPatchRow = createZeroPaddedTensor(
|
||||
paddedPatchRow, paddedPatchRowType, {0, 0}, {0, paddedK - patchSize}, rewriter, widthLoc);
|
||||
FailureOr<Value> paddedOutputRow = createPaddedConvOutputTile(
|
||||
paddedPatchRow, args.weights.front(), 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));
|
||||
}
|
||||
Value outputPixel = tensor::ExpandShapeOp::create(
|
||||
rewriter, widthLoc, tilePixelType, outputRow, 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(1),
|
||||
rewriter.getIndexAttr(1)};
|
||||
Value nextFragment = tensor::InsertSliceOp::create(rewriter,
|
||||
widthLoc,
|
||||
outputPixel,
|
||||
widthIterArgs.front(),
|
||||
rowOffsets,
|
||||
rowSizes,
|
||||
getUnitStrides(rewriter, 4));
|
||||
widthYielded.push_back(nextFragment);
|
||||
return success();
|
||||
});
|
||||
if (failed(widthLoop))
|
||||
return failure();
|
||||
publishGraphBatchPhysicalFragment(
|
||||
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))
|
||||
return failure();
|
||||
Value output = assemblyBatch->getResult(0);
|
||||
if (state.hasBias)
|
||||
return applyRowStripBiasAdd(output, state.outType, state.b, rewriter, loc);
|
||||
return output;
|
||||
}
|
||||
|
||||
static FailureOr<Value>
|
||||
createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) {
|
||||
ConvGeometry geometry = buildConvGeometry(state);
|
||||
if (state.group != 1 || state.batchSize != 1 || geometry.c > geometry.xbarSize)
|
||||
if (state.group != 1 || state.batchSize != 1 || !rowStripOutputTileFitsOneCore(geometry))
|
||||
return failure();
|
||||
|
||||
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
|
||||
@@ -2796,12 +3090,17 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe
|
||||
auto elementType = state.outType.getElementType();
|
||||
auto fragmentType = getRowStripFragmentType(state.outType);
|
||||
auto outputPixelType = RankedTensorType::get({1, state.numChannelsOut, 1, 1}, elementType);
|
||||
auto patchType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.wWidth}, state.xType.getElementType());
|
||||
auto patchRowType = RankedTensorType::get({1, patchSize}, state.xType.getElementType());
|
||||
auto outputStorageType = getRowStripStorageType(state.outType);
|
||||
|
||||
PreparedConvInput preparedInput = standard::prepareInputForIm2Col(state, rewriter, loc);
|
||||
Value paddedWeights = standard::createPaddedInputKTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter);
|
||||
Value paddedWeights = state.numChannelsOut <= xbarDim
|
||||
? standard::createPaddedInputKTiledWeightConstant(
|
||||
weightDenseAttr, state, paddedK, xbarDim, rewriter)
|
||||
: standard::createPaddedOutputChannelTiledWeightConstant(
|
||||
weightDenseAttr, state, paddedK, xbarDim, rewriter);
|
||||
if (!rowStripOutputFitsOneCore(geometry))
|
||||
return createOutputChannelTiledRowStripConvOutput(
|
||||
state, paddedWeights, paddedK, numKSlices, xbarDim, rewriter, loc);
|
||||
|
||||
FailureOr<Value> paddedBias = failure();
|
||||
if (state.hasBias)
|
||||
paddedBias = createPaddedBiasRowConstant(state, xbarDim, rewriter);
|
||||
@@ -2814,13 +3113,16 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe
|
||||
TypeRange {outputStorageType},
|
||||
state.outHeight,
|
||||
ValueRange {paddedWeights},
|
||||
state.hasBias ? ValueRange {preparedInput.value, *paddedBias} : ValueRange {preparedInput.value},
|
||||
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 inputHeightOffset = affineMulConst(rewriter, loc, args.lane, state.strideHeight, anchorOp);
|
||||
FailureOr<Value> inputWindow =
|
||||
createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc);
|
||||
if (failed(inputWindow))
|
||||
return failure();
|
||||
Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType);
|
||||
auto widthLoop = buildNormalizedScfFor(
|
||||
rewriter,
|
||||
@@ -2830,20 +3132,11 @@ createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRe
|
||||
c1,
|
||||
ValueRange {fragmentInit},
|
||||
[&](OpBuilder&, Location widthLoc, Value widthIndex, ValueRange widthIterArgs, SmallVectorImpl<Value>& widthYielded) {
|
||||
Value inputWidthOffset = affineMulConst(rewriter, widthLoc, widthIndex, state.strideWidth, anchorOp);
|
||||
Value patch = createConvInputPatch(args.inputs.front(),
|
||||
patchType,
|
||||
c0,
|
||||
c0,
|
||||
inputHeightOffset,
|
||||
inputWidthOffset,
|
||||
state.dilationHeight,
|
||||
state.dilationWidth,
|
||||
rewriter,
|
||||
widthLoc);
|
||||
Value patchRow = tensor::CollapseShapeOp::create(
|
||||
rewriter, widthLoc, patchRowType, patch, SmallVector<ReassociationIndices> {{0}, {1, 2, 3}});
|
||||
FailureOr<Value> outputRow = createPaddedConvOutputRow(patchRow,
|
||||
FailureOr<Value> patchRow =
|
||||
createNchwRowStripConvPatchRow(*inputWindow, state, widthIndex, rewriter, widthLoc);
|
||||
if (failed(patchRow))
|
||||
return failure();
|
||||
FailureOr<Value> outputRow = createPaddedConvOutputRow(*patchRow,
|
||||
state,
|
||||
args.weights.front(),
|
||||
state.hasBias ? args.inputs[1] : Value(),
|
||||
@@ -2924,7 +3217,7 @@ static FailureOr<Value> createConvOutputFromNchwRowStripFragments(Value rowStrip
|
||||
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
|
||||
Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth);
|
||||
auto fragmentType = getRowStripFragmentType(state.outType);
|
||||
FailureOr<Value> inputWindow = createNchwRowStripConvWindow(args.inputs.front(), state, args.lane, rewriter, loc);
|
||||
FailureOr<Value> inputWindow = createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc);
|
||||
if (failed(inputWindow))
|
||||
return failure();
|
||||
Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, fragmentType.getShape(), elementType);
|
||||
@@ -3811,7 +4104,7 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) {
|
||||
analysis.barrierKind = DistributedConvBarrierKind::UnsupportedConsumer;
|
||||
analysis.barrierDetail = "selected row-strip layout";
|
||||
ConvGeometry geometry = buildConvGeometry(*state);
|
||||
if (geometry.c > geometry.xbarSize)
|
||||
if (!rowStripOutputTileFitsOneCore(geometry))
|
||||
return failure();
|
||||
ConvLoweringDecision decision = chooseConvLoweringStrategy(geometry, *requestedStrategy, analysis);
|
||||
if (decision.strategy == PimConvLoweringDepthwise && !depthwise::canUseStructuredRewrite(*state)
|
||||
|
||||
Reference in New Issue
Block a user