Resnet is fast
This commit is contained in:
@@ -80,8 +80,17 @@ Value extractMixedSliceOrIdentity(RewriterBase &rewriter,
|
||||
|
||||
Value insertMixedSlice(OpBuilder &builder, Location loc, Value source,
|
||||
Value dest, const MixedSliceGeometry &geometry) {
|
||||
SmallVector<OpFoldResult> sizes(geometry.sizes);
|
||||
auto sourceType = dyn_cast<RankedTensorType>(source.getType());
|
||||
auto destType = dyn_cast<RankedTensorType>(dest.getType());
|
||||
if (sourceType && destType && sourceType.hasStaticShape()
|
||||
&& sourceType.getRank() == destType.getRank()) {
|
||||
sizes.clear();
|
||||
for (int64_t dimension : sourceType.getShape())
|
||||
sizes.push_back(builder.getIndexAttr(dimension));
|
||||
}
|
||||
return tensor::InsertSliceOp::create(builder, loc, source, dest,
|
||||
geometry.offsets, geometry.sizes,
|
||||
geometry.offsets, sizes,
|
||||
geometry.strides);
|
||||
}
|
||||
|
||||
|
||||
@@ -1590,7 +1590,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
|
||||
}
|
||||
|
||||
for (auto [slot, fileName] : llvm::enumerate(weightFiles)) {
|
||||
xbarsPerGroup.push_back(static_cast<int64_t>(slot));
|
||||
xbarsPerGroup.push_back(1);
|
||||
std::string sourcePath = outputDirPath + "/weights/" + fileName;
|
||||
std::string targetPath = coreWeightsDirPath + "/crossbar_" + std::to_string(slot) + ".bin";
|
||||
sys::fs::remove(targetPath);
|
||||
|
||||
@@ -191,6 +191,37 @@ struct LowerSpatialPlansPass final : PassWrapper<LowerSpatialPlansPass, Operatio
|
||||
rewriter.replaceOp(planOp, computeOp.getResults());
|
||||
continue;
|
||||
}
|
||||
if (auto planOp = dyn_cast<spatial::SpatMaxPool2DPlanOp>(&op)) {
|
||||
auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) {
|
||||
auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(user);
|
||||
return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout;
|
||||
});
|
||||
if (outputBlueprint == planOp.getResult().getUsers().end()) {
|
||||
planOp.emitOpError("selected MaxPool plan requires a row-strip blueprint result");
|
||||
signalPassFailure();
|
||||
return;
|
||||
}
|
||||
|
||||
FailureOr<RowStripPhysicalValue> input = getRowStripValue(rowStripValues, planOp.getInput());
|
||||
rewriter.setInsertionPoint(planOp);
|
||||
FailureOr<Value> lowered = lowerSelectedMaxPool2DPlan(
|
||||
planOp, succeeded(input) ? std::optional<Value> {input->storage} : std::nullopt, rewriter);
|
||||
if (failed(lowered)) {
|
||||
planOp.emitOpError("failed to lower selected row-strip Spatial MaxPool plan");
|
||||
signalPassFailure();
|
||||
return;
|
||||
}
|
||||
auto blueprint = cast<spatial::SpatBlueprintOp>(*outputBlueprint);
|
||||
FailureOr<RowStripPhysicalValue> output = buildRowStripValue(blueprint, *lowered);
|
||||
if (failed(output)) {
|
||||
signalPassFailure();
|
||||
return;
|
||||
}
|
||||
rowStripValues[blueprint.getResult()] = *output;
|
||||
eraseAfterLowering.insert(planOp);
|
||||
eraseAfterLowering.insert(blueprint);
|
||||
continue;
|
||||
}
|
||||
if (auto planOp = dyn_cast<spatial::SpatBiasAddPlanOp>(&op)) {
|
||||
if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) {
|
||||
auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) {
|
||||
@@ -386,6 +417,7 @@ struct LowerSpatialPlansPass final : PassWrapper<LowerSpatialPlansPass, Operatio
|
||||
} else if (isa<spatial::SpatConv2DPlanOp,
|
||||
spatial::SpatBiasAddPlanOp,
|
||||
spatial::SpatReluPlanOp,
|
||||
spatial::SpatMaxPool2DPlanOp,
|
||||
spatial::SpatMaterializeLayoutOp>(op)
|
||||
|| op->getDialect()->getNamespace() == "onnx") {
|
||||
op->emitOpError("operation must not remain after LowerSpatialPlans");
|
||||
|
||||
@@ -48,10 +48,11 @@ static void populateEmptyFunction(func::FuncOp funcOp) {
|
||||
SmallVector<spatial::SpatConv2DPlanOp> convPlans(funcOp.getOps<spatial::SpatConv2DPlanOp>());
|
||||
SmallVector<spatial::SpatBiasAddPlanOp> biasAddPlans(funcOp.getOps<spatial::SpatBiasAddPlanOp>());
|
||||
SmallVector<spatial::SpatReluPlanOp> reluPlans(funcOp.getOps<spatial::SpatReluPlanOp>());
|
||||
SmallVector<spatial::SpatMaxPool2DPlanOp> maxPoolPlans(funcOp.getOps<spatial::SpatMaxPool2DPlanOp>());
|
||||
SmallVector<spatial::SpatBlueprintOp> blueprints(funcOp.getOps<spatial::SpatBlueprintOp>());
|
||||
SmallVector<spatial::SpatMaterializeLayoutOp> materializers(funcOp.getOps<spatial::SpatMaterializeLayoutOp>());
|
||||
if (!computes.empty() || !computeBatches.empty() || !convPlans.empty() || !biasAddPlans.empty() || !reluPlans.empty()
|
||||
|| !blueprints.empty() || !materializers.empty()) {
|
||||
|| !maxPoolPlans.empty() || !blueprints.empty() || !materializers.empty()) {
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
@@ -147,6 +147,7 @@ void verifyLogicalTopLevelOps(func::FuncOp funcOp, pim::CappedDiagnosticReporter
|
||||
spatial::SpatConv2DPlanOp,
|
||||
spatial::SpatBiasAddPlanOp,
|
||||
spatial::SpatReluPlanOp,
|
||||
spatial::SpatMaxPool2DPlanOp,
|
||||
spatial::SpatBlueprintOp,
|
||||
spatial::SpatMaterializeLayoutOp>(&op)) {
|
||||
continue;
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -16,6 +16,8 @@
|
||||
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
|
||||
#include "src/Dialect/ONNX/ONNXOps.hpp"
|
||||
|
||||
@@ -24,7 +26,7 @@ using namespace mlir;
|
||||
namespace onnx_mlir {
|
||||
namespace {
|
||||
|
||||
static Value materializeTileTensor(ConversionPatternRewriter& rewriter, Location loc, Value tile) {
|
||||
static Value materializeTileTensor(PatternRewriter& rewriter, Location loc, Value tile) {
|
||||
auto tileType = cast<RankedTensorType>(tile.getType());
|
||||
Value empty = tensor::EmptyOp::create(rewriter, loc, tileType.getShape(), tileType.getElementType());
|
||||
return insertStaticSlice(rewriter, loc, tile, empty, getZeroOffsets(rewriter, tileType.getRank()));
|
||||
@@ -228,6 +230,23 @@ struct PoolToSpatialComputeBase : public OpConversionPattern<PoolOp> {
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (std::is_same_v<PoolOp, ONNXMaxPoolSingleOutOp>) {
|
||||
if (batchSize == 1) {
|
||||
auto plan = spatial::SpatMaxPool2DPlanOp::create(
|
||||
rewriter,
|
||||
loc,
|
||||
outType,
|
||||
x,
|
||||
rewriter.getDenseI64ArrayAttr({kernelHeight, kernelWidth}),
|
||||
rewriter.getDenseI64ArrayAttr({padTop, padLeft, padBottom, padRight}),
|
||||
rewriter.getDenseI64ArrayAttr({strideHeight, strideWidth}),
|
||||
rewriter.getDenseI64ArrayAttr({dilationHeight, dilationWidth}),
|
||||
rewriter.getStringAttr("nchw"));
|
||||
rewriter.replaceOp(poolOp, plan.getResult());
|
||||
return success();
|
||||
}
|
||||
}
|
||||
|
||||
const int64_t xbarSize = static_cast<int64_t>(crossbarSize.getValue());
|
||||
const int64_t channelTileCount = (channels + xbarSize - 1) / xbarSize;
|
||||
const int64_t outputPatchCount = batchSize * outputHeight * outputWidth;
|
||||
@@ -396,6 +415,220 @@ struct PoolToSpatialCompute<ONNXAveragePoolOp>
|
||||
|
||||
} // namespace
|
||||
|
||||
LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) {
|
||||
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
|
||||
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
|
||||
if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape())
|
||||
return failure();
|
||||
if (inputType.getRank() != 4 || outputType.getRank() != 4 || inputType.getDimSize(0) != 1
|
||||
|| outputType.getDimSize(0) != 1 || inputType.getDimSize(1) != outputType.getDimSize(1))
|
||||
return failure();
|
||||
if (llvm::any_of(planOp.getKernelShape(), [](int64_t value) { return value <= 0; })
|
||||
|| llvm::any_of(planOp.getStrides(), [](int64_t value) { return value <= 0; })
|
||||
|| llvm::any_of(planOp.getDilations(), [](int64_t value) { return value <= 0; }))
|
||||
return failure();
|
||||
return success();
|
||||
}
|
||||
|
||||
static Value createClampedPoolIndexTable(PatternRewriter& rewriter,
|
||||
Operation* anchorOp,
|
||||
int64_t outputSize,
|
||||
int64_t kernelSize,
|
||||
int64_t stride,
|
||||
int64_t dilation,
|
||||
int64_t padBegin,
|
||||
int64_t inputSize) {
|
||||
auto tableType = RankedTensorType::get({outputSize * kernelSize}, rewriter.getIndexType());
|
||||
SmallVector<Attribute> values;
|
||||
values.reserve(tableType.getNumElements());
|
||||
for (int64_t output = 0; output < outputSize; ++output)
|
||||
for (int64_t kernel = 0; kernel < kernelSize; ++kernel)
|
||||
values.push_back(rewriter.getIndexAttr(
|
||||
std::clamp(output * stride + kernel * dilation - padBegin, int64_t {0}, inputSize - 1)));
|
||||
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
|
||||
}
|
||||
|
||||
static Value extractPoolIndex(PatternRewriter& rewriter,
|
||||
Location loc,
|
||||
Operation* anchorOp,
|
||||
Value table,
|
||||
Value outputIndex,
|
||||
int64_t kernelIndex,
|
||||
int64_t kernelSize) {
|
||||
Value tableIndex = arith::MulIOp::create(
|
||||
rewriter, loc, outputIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelSize));
|
||||
if (kernelIndex != 0)
|
||||
tableIndex = arith::AddIOp::create(
|
||||
rewriter, loc, tableIndex, getOrCreateIndexConstant(rewriter, anchorOp, kernelIndex));
|
||||
return tensor::ExtractOp::create(rewriter, loc, table, tableIndex);
|
||||
}
|
||||
|
||||
FailureOr<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
std::optional<Value> rowStripInput,
|
||||
PatternRewriter& rewriter) {
|
||||
if (failed(canLowerMaxPoolPlanToRowStrip(planOp)))
|
||||
return failure();
|
||||
|
||||
Location loc = planOp.getLoc();
|
||||
auto inputType = cast<RankedTensorType>(planOp.getInput().getType());
|
||||
auto outputType = cast<RankedTensorType>(planOp.getOutput().getType());
|
||||
const int64_t channels = inputType.getDimSize(1);
|
||||
const int64_t inputHeight = inputType.getDimSize(2);
|
||||
const int64_t inputWidth = inputType.getDimSize(3);
|
||||
const int64_t outputHeight = outputType.getDimSize(2);
|
||||
const int64_t outputWidth = outputType.getDimSize(3);
|
||||
const int64_t kernelHeight = planOp.getKernelShape()[0];
|
||||
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);
|
||||
if (!physicalInput && actualInputType != inputType)
|
||||
return failure();
|
||||
|
||||
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value rowTable = createClampedPoolIndexTable(rewriter,
|
||||
anchorOp,
|
||||
outputHeight,
|
||||
kernelHeight,
|
||||
planOp.getStrides()[0],
|
||||
planOp.getDilations()[0],
|
||||
planOp.getPads()[0],
|
||||
inputHeight);
|
||||
Value columnTable = createClampedPoolIndexTable(rewriter,
|
||||
anchorOp,
|
||||
outputWidth,
|
||||
kernelWidth,
|
||||
planOp.getStrides()[1],
|
||||
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 batch = createSpatComputeBatch(
|
||||
rewriter,
|
||||
loc,
|
||||
TypeRange {outputStorageType},
|
||||
outputHeight,
|
||||
{},
|
||||
ValueRange {input},
|
||||
[&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult {
|
||||
SmallVector<Value> inputRows;
|
||||
inputRows.reserve(kernelHeight);
|
||||
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
|
||||
Value sourceRow =
|
||||
extractPoolIndex(rewriter, loc, anchorOp, rowTable, args.lane, kernelRow, kernelHeight);
|
||||
if (physicalInput) {
|
||||
inputRows.push_back(
|
||||
extractRowStripFragment(args.inputs.front(), inputType, sourceRow, rewriter, loc));
|
||||
}
|
||||
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)));
|
||||
}
|
||||
}
|
||||
|
||||
auto windowType = RankedTensorType::get(
|
||||
{1, channels, kernelHeight, inputWidth}, 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)};
|
||||
window = tensor::InsertSliceOp::create(rewriter,
|
||||
loc,
|
||||
inputRows[kernelRow],
|
||||
window,
|
||||
offsets,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(channels),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(inputWidth)},
|
||||
getUnitStrides(rewriter, 4));
|
||||
}
|
||||
|
||||
Value outputInit = tensor::EmptyOp::create(
|
||||
rewriter, loc, outputFragmentType.getShape(), outputFragmentType.getElementType());
|
||||
Operation* bodyAnchor = rewriter.getInsertionBlock()->getParentOp();
|
||||
Value c0 = getOrCreateIndexConstant(rewriter, bodyAnchor, 0);
|
||||
Value c1 = getOrCreateIndexConstant(rewriter, bodyAnchor, 1);
|
||||
Value cOutputWidth = getOrCreateIndexConstant(rewriter, bodyAnchor, outputWidth);
|
||||
auto outputLoop = buildNormalizedScfFor(
|
||||
rewriter,
|
||||
loc,
|
||||
c0,
|
||||
cOutputWidth,
|
||||
c1,
|
||||
ValueRange {outputInit},
|
||||
[&](OpBuilder&, Location nestedLoc, Value outputColumn, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
|
||||
Value reduced;
|
||||
for (int64_t kernelRow = 0; kernelRow < kernelHeight; ++kernelRow) {
|
||||
for (int64_t kernelColumn = 0; kernelColumn < kernelWidth; ++kernelColumn) {
|
||||
Value sourceColumn = extractPoolIndex(rewriter,
|
||||
nestedLoc,
|
||||
bodyAnchor,
|
||||
columnTable,
|
||||
outputColumn,
|
||||
kernelColumn,
|
||||
kernelWidth);
|
||||
SmallVector<OpFoldResult> offsets {
|
||||
rewriter.getIndexAttr(0),
|
||||
rewriter.getIndexAttr(0),
|
||||
rewriter.getIndexAttr(kernelRow),
|
||||
sourceColumn};
|
||||
Value point = tensor::ExtractSliceOp::create(rewriter,
|
||||
nestedLoc,
|
||||
tileType,
|
||||
window,
|
||||
offsets,
|
||||
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(channels),
|
||||
rewriter.getIndexAttr(1),
|
||||
rewriter.getIndexAttr(1)},
|
||||
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};
|
||||
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)},
|
||||
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);
|
||||
return success();
|
||||
});
|
||||
if (failed(batch))
|
||||
return failure();
|
||||
return batch->getResult(0);
|
||||
}
|
||||
|
||||
void populatePoolPatterns(RewritePatternSet& patterns, MLIRContext* ctx) {
|
||||
patterns.insert<PoolToSpatialCompute<ONNXMaxPoolSingleOutOp>>(ctx);
|
||||
patterns.insert<PoolToSpatialCompute<ONNXAveragePoolOp>>(ctx);
|
||||
|
||||
@@ -18,4 +18,11 @@ lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp,
|
||||
mlir::LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp);
|
||||
mlir::LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp);
|
||||
|
||||
mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp);
|
||||
|
||||
mlir::FailureOr<mlir::Value>
|
||||
lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
|
||||
std::optional<mlir::Value> rowStripInput,
|
||||
mlir::PatternRewriter& rewriter);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -38,6 +38,8 @@ static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap<Value, Selected
|
||||
return getSelectedLayout(layouts, biasAddPlan.getResult()) == SelectedLayout::NchwRowStrip;
|
||||
if (auto convPlan = dyn_cast<spatial::SpatConv2DPlanOp>(user))
|
||||
return getSelectedLayout(layouts, convPlan.getResult()) == SelectedLayout::NchwRowStrip;
|
||||
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(user))
|
||||
return getSelectedLayout(layouts, maxPoolPlan.getResult()) == SelectedLayout::NchwRowStrip;
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -60,6 +62,8 @@ static bool canConsumeRowStripAsUser(Operation* user) {
|
||||
}
|
||||
if (auto convPlan = dyn_cast<spatial::SpatConv2DPlanOp>(user))
|
||||
return succeeded(canConsumeAndProduceRowStrip(convPlan));
|
||||
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(user))
|
||||
return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan));
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -70,7 +74,6 @@ static bool hasRowStripConsumer(Value value) {
|
||||
return false;
|
||||
}
|
||||
|
||||
|
||||
static bool canSelectConvRowStrip(spatial::SpatConv2DPlanOp convPlan,
|
||||
llvm::DenseMap<Value, SelectedLayout>& layouts) {
|
||||
SelectedLayout inputLayout = getSelectedLayout(layouts, convPlan.getInput());
|
||||
@@ -83,9 +86,6 @@ static SelectedLayout chooseConvLayout(spatial::SpatConv2DPlanOp convPlan,
|
||||
llvm::DenseMap<Value, SelectedLayout>& layouts) {
|
||||
if (!canSelectConvRowStrip(convPlan, layouts))
|
||||
return SelectedLayout::DenseNchw;
|
||||
if (getSelectedLayout(layouts, convPlan.getInput()) != SelectedLayout::NchwRowStrip
|
||||
&& !hasRowStripConsumer(convPlan.getResult()))
|
||||
return SelectedLayout::DenseNchw;
|
||||
if (!allUsersCanHandleRowStrip(convPlan.getResult(), layouts))
|
||||
return SelectedLayout::DenseNchw;
|
||||
return SelectedLayout::NchwRowStrip;
|
||||
@@ -116,6 +116,11 @@ static SelectedLayout chooseBiasAddLayout(spatial::SpatBiasAddPlanOp biasAddPlan
|
||||
return SelectedLayout::NchwRowStrip;
|
||||
}
|
||||
|
||||
static SelectedLayout chooseMaxPoolLayout(spatial::SpatMaxPool2DPlanOp maxPoolPlan) {
|
||||
return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan)) ? SelectedLayout::NchwRowStrip
|
||||
: SelectedLayout::DenseNchw;
|
||||
}
|
||||
|
||||
static spatial::SpatBlueprintOp insertRowStripBlueprint(IRRewriter& rewriter, Value value) {
|
||||
auto outputType = cast<RankedTensorType>(value.getType());
|
||||
auto [offsets, sizes] = buildRowStripMetadata(outputType);
|
||||
@@ -208,6 +213,14 @@ struct SpatialLayoutPlanningPass final : PassWrapper<SpatialLayoutPlanningPass,
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(&op)) {
|
||||
SelectedLayout selected = chooseMaxPoolLayout(maxPoolPlan);
|
||||
if (layouts[maxPoolPlan.getResult()] != selected) {
|
||||
layouts[maxPoolPlan.getResult()] = selected;
|
||||
changed = true;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -219,6 +232,8 @@ struct SpatialLayoutPlanningPass final : PassWrapper<SpatialLayoutPlanningPass,
|
||||
producedValue = biasAddPlan.getResult();
|
||||
else if (auto reluPlan = dyn_cast<spatial::SpatReluPlanOp>(&op))
|
||||
producedValue = reluPlan.getResult();
|
||||
else if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(&op))
|
||||
producedValue = maxPoolPlan.getResult();
|
||||
else
|
||||
continue;
|
||||
|
||||
|
||||
@@ -29,6 +29,11 @@ static bool isUsedOnlyAsExplicitHostOperand(Value value) {
|
||||
});
|
||||
}
|
||||
|
||||
static bool isUsedOnlyByExtractSlices(Value value) {
|
||||
return !value.use_empty()
|
||||
&& llvm::all_of(value.getUsers(), [](Operation* user) { return isa<tensor::ExtractSliceOp>(user); });
|
||||
}
|
||||
|
||||
static FailureOr<unsigned> getDirectReturnOperandIndex(OpResult result) {
|
||||
if (!result.hasOneUse())
|
||||
return failure();
|
||||
@@ -357,6 +362,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
|
||||
rewriter.createBlock(&coreBatchOp.getBody(), coreBatchOp.getBody().end(), TypeRange(blockArgTypes), blockArgLocs);
|
||||
|
||||
IRMapping mapper;
|
||||
SmallPtrSet<Value, 4> hostResidentTensors;
|
||||
rewriter.setInsertionPointToStart(newBlock);
|
||||
auto oldLaneArg = computeBatchOp.getLaneArgument();
|
||||
if (!oldLaneArg)
|
||||
@@ -523,8 +529,10 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
|
||||
if (isa_and_present<memref::GetGlobalOp>(toTensorOp.getBuffer().getDefiningOp())) {
|
||||
Operation* cloned = rewriter.clone(op, mapper);
|
||||
auto clonedTensor = cloned->getResult(0);
|
||||
if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult())) {
|
||||
if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult())
|
||||
|| isUsedOnlyByExtractSlices(toTensorOp.getResult())) {
|
||||
mapper.map(toTensorOp.getResult(), clonedTensor);
|
||||
hostResidentTensors.insert(toTensorOp.getResult());
|
||||
continue;
|
||||
}
|
||||
auto clonedType = cast<ShapedType>(clonedTensor.getType());
|
||||
@@ -542,6 +550,28 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
|
||||
}
|
||||
}
|
||||
|
||||
if (auto extractSlice = dyn_cast<tensor::ExtractSliceOp>(op);
|
||||
extractSlice && hostResidentTensors.contains(extractSlice.getSource())) {
|
||||
Operation* cloned = rewriter.clone(op, mapper);
|
||||
Value hostSlice = cloned->getResult(0);
|
||||
auto outputBuffer = createEmptyTensorFromShaped(rewriter, loc, cast<ShapedType>(hostSlice.getType()));
|
||||
Value zeroOffset = getOrCreateIndexConstant(rewriter, coreBatchOp.getOperation(), 0);
|
||||
auto sizeAttr = getTensorSizeInBytesAttr(rewriter, coreBatchOp.getOperation(), hostSlice);
|
||||
if (failed(sizeAttr))
|
||||
return failure();
|
||||
auto copied = pim::PimMemCopyHostToDevOp::create(rewriter,
|
||||
loc,
|
||||
outputBuffer.getType(),
|
||||
zeroOffset,
|
||||
zeroOffset,
|
||||
outputBuffer,
|
||||
hostSlice,
|
||||
*sizeAttr)
|
||||
.getOutput();
|
||||
mapper.map(extractSlice.getResult(), copied);
|
||||
continue;
|
||||
}
|
||||
|
||||
for (auto [operandIndex, operand] : llvm::enumerate(op.getOperands())) {
|
||||
if (!isa<TensorType>(operand.getType()) || mapper.contains(operand))
|
||||
continue;
|
||||
|
||||
@@ -270,6 +270,25 @@ def SpatReluPlanOp : SpatOp<"relu_plan", []> {
|
||||
let hasVerifier = 1;
|
||||
}
|
||||
|
||||
def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> {
|
||||
let summary = "Layout-aware 2D NCHW MaxPool planning op";
|
||||
|
||||
let arguments = (ins
|
||||
SpatTensor:$input,
|
||||
DenseI64ArrayAttr:$kernelShape,
|
||||
DenseI64ArrayAttr:$pads,
|
||||
DenseI64ArrayAttr:$strides,
|
||||
DenseI64ArrayAttr:$dilations,
|
||||
StrAttr:$logicalLayout
|
||||
);
|
||||
|
||||
let results = (outs
|
||||
SpatTensor:$output
|
||||
);
|
||||
|
||||
let hasVerifier = 1;
|
||||
}
|
||||
|
||||
def SpatBiasAddPlanOp : SpatOp<"bias_add_plan", []> {
|
||||
let summary = "Layout-aware Conv-style bias add planning op";
|
||||
|
||||
|
||||
@@ -486,6 +486,26 @@ LogicalResult SpatReluPlanOp::verify() {
|
||||
return success();
|
||||
}
|
||||
|
||||
LogicalResult SpatMaxPool2DPlanOp::verify() {
|
||||
if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.max_pool2d_plan")))
|
||||
return failure();
|
||||
auto inputType = dyn_cast<RankedTensorType>(getInput().getType());
|
||||
auto outputType = dyn_cast<RankedTensorType>(getOutput().getType());
|
||||
if (!inputType.hasStaticShape() || !outputType.hasStaticShape() || inputType.getRank() != 4
|
||||
|| outputType.getRank() != 4)
|
||||
return emitError("requires static rank-4 input and output tensors");
|
||||
if (getLogicalLayout() != "nchw")
|
||||
return emitError("requires logical layout \"nchw\"");
|
||||
if (getKernelShape().size() != 2 || getStrides().size() != 2 || getDilations().size() != 2)
|
||||
return emitError("requires two kernel, stride, and dilation values");
|
||||
if (getPads().size() != 4)
|
||||
return emitError("requires four pad values");
|
||||
if (inputType.getDimSize(0) != outputType.getDimSize(0)
|
||||
|| inputType.getDimSize(1) != outputType.getDimSize(1))
|
||||
return emitError("requires matching input/output batch and channel dimensions");
|
||||
return success();
|
||||
}
|
||||
|
||||
LogicalResult SpatBiasAddPlanOp::verify() {
|
||||
if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.bias_add_plan")))
|
||||
return failure();
|
||||
|
||||
@@ -780,7 +780,7 @@ ComputeGraph buildComputeGraph(Operation* entryOp) {
|
||||
if (auto batch = dyn_cast<SpatComputeBatch>(&op)) {
|
||||
if (isUsedAsWeightOnly(batch.getOperation()))
|
||||
continue;
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch.getLaneCount());
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch);
|
||||
for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) {
|
||||
ComputeInstance instance = getBatchChunkForIndex(batch, chunkIndex);
|
||||
size_t index = graph.nodes.size();
|
||||
|
||||
+51
-13
@@ -1,9 +1,11 @@
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
|
||||
#include "ComputeGraph.hpp"
|
||||
#include "ComputeInstanceUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
|
||||
@@ -12,20 +14,16 @@ using namespace mlir;
|
||||
namespace onnx_mlir {
|
||||
namespace spatial {
|
||||
|
||||
static constexpr llvm::StringLiteral kMergeChunkCountAttr = "spat.merge_chunk_count";
|
||||
|
||||
size_t getSchedulingCpuBudget() {
|
||||
if (coresCount.getValue() > 0)
|
||||
return static_cast<size_t>(coresCount.getValue());
|
||||
return std::numeric_limits<size_t>::max();
|
||||
}
|
||||
|
||||
size_t getBatchChunkTargetCount(int32_t laneCount) {
|
||||
static BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkCount, size_t chunkIndex) {
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
return std::min(static_cast<size_t>(laneCount), getSchedulingCpuBudget());
|
||||
}
|
||||
|
||||
BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) {
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
size_t chunkCount = getBatchChunkTargetCount(laneCount);
|
||||
assert(chunkIndex < chunkCount && "chunkIndex out of range");
|
||||
|
||||
size_t laneCountSize = static_cast<size_t>(laneCount);
|
||||
@@ -38,11 +36,51 @@ BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) {
|
||||
return {static_cast<uint32_t>(start), static_cast<uint32_t>(count)};
|
||||
}
|
||||
|
||||
size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) {
|
||||
static bool batchChunksFit(SpatComputeBatch batch, size_t chunkCount, size_t crossbarCapacity) {
|
||||
for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) {
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkCount, chunkIndex);
|
||||
ComputeInstance instance {batch.getOperation(), chunk.laneStart, chunk.laneCount};
|
||||
if (getComputeInstanceCrossbarUsage(instance).size() > crossbarCapacity)
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t getBatchChunkTargetCount(SpatComputeBatch batch) {
|
||||
if (auto chunkCount = batch->getAttrOfType<IntegerAttr>(kMergeChunkCountAttr))
|
||||
return static_cast<size_t>(chunkCount.getInt());
|
||||
|
||||
int32_t laneCount = batch.getLaneCount();
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
size_t maxChunkCount = std::min(static_cast<size_t>(laneCount), getSchedulingCpuBudget());
|
||||
size_t crossbarCapacity = crossbarCountInCore.getValue();
|
||||
CrossbarUsage fullUsage = collectDistinctCrossbarWeights(batch.getOperation());
|
||||
if (fullUsage.empty() || crossbarCapacity == 0) {
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount));
|
||||
return maxChunkCount;
|
||||
}
|
||||
|
||||
size_t chunkCount = std::max<size_t>(1, (fullUsage.size() + crossbarCapacity - 1) / crossbarCapacity);
|
||||
for (; chunkCount <= maxChunkCount; ++chunkCount) {
|
||||
if (batchChunksFit(batch, chunkCount, crossbarCapacity)) {
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), chunkCount));
|
||||
return chunkCount;
|
||||
}
|
||||
}
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount));
|
||||
return maxChunkCount;
|
||||
}
|
||||
|
||||
BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex) {
|
||||
return getBatchChunkRange(batch.getLaneCount(), getBatchChunkTargetCount(batch), chunkIndex);
|
||||
}
|
||||
|
||||
size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane) {
|
||||
int32_t laneCount = batch.getLaneCount();
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
assert(lane < static_cast<uint32_t>(laneCount) && "lane out of range");
|
||||
|
||||
size_t chunkCount = getBatchChunkTargetCount(laneCount);
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch);
|
||||
size_t laneCountSize = static_cast<size_t>(laneCount);
|
||||
size_t baseChunkSize = laneCountSize / chunkCount;
|
||||
size_t remainder = laneCountSize % chunkCount;
|
||||
@@ -56,12 +94,12 @@ size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) {
|
||||
}
|
||||
|
||||
ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex) {
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkIndex);
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch, chunkIndex);
|
||||
return {batch.getOperation(), chunk.laneStart, chunk.laneCount};
|
||||
}
|
||||
|
||||
ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane) {
|
||||
return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch.getLaneCount(), lane));
|
||||
return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch, lane));
|
||||
}
|
||||
|
||||
llvm::SmallVector<ComputeInstance, 4>
|
||||
@@ -74,8 +112,8 @@ getBatchChunksForRange(SpatComputeBatch batch, uint32_t laneStart, uint32_t lane
|
||||
assert(laneEnd >= laneStart && "lane range overflow");
|
||||
assert(laneEnd <= static_cast<uint32_t>(batch.getLaneCount()) && "lane range out of bounds");
|
||||
|
||||
size_t firstChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneStart);
|
||||
size_t lastChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneEnd - 1);
|
||||
size_t firstChunk = getBatchChunkIndexForLane(batch, laneStart);
|
||||
size_t lastChunk = getBatchChunkIndexForLane(batch, laneEnd - 1);
|
||||
chunks.reserve(lastChunk - firstChunk + 1);
|
||||
for (size_t chunkIndex = firstChunk; chunkIndex <= lastChunk; ++chunkIndex)
|
||||
chunks.push_back(getBatchChunkForIndex(batch, chunkIndex));
|
||||
|
||||
+3
-3
@@ -27,9 +27,9 @@ struct BatchChunkRange {
|
||||
};
|
||||
|
||||
size_t getSchedulingCpuBudget();
|
||||
size_t getBatchChunkTargetCount(int32_t laneCount);
|
||||
BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex);
|
||||
size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane);
|
||||
size_t getBatchChunkTargetCount(SpatComputeBatch batch);
|
||||
BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex);
|
||||
size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane);
|
||||
ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex);
|
||||
ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane);
|
||||
llvm::SmallVector<ComputeInstance, 4>
|
||||
|
||||
@@ -6,7 +6,9 @@
|
||||
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <queue>
|
||||
#include <tuple>
|
||||
#include <vector>
|
||||
|
||||
#include "PeftScheduler.hpp"
|
||||
@@ -133,6 +135,55 @@ void verifyOctTableSize(size_t nodeCount, size_t processorCount) {
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<CrossbarUsage> planCrossbarResidency(const ComputeGraph& graph,
|
||||
size_t processorCount,
|
||||
size_t crossbarCapacity,
|
||||
const MeshModel& mesh) {
|
||||
std::vector<size_t> weightedTasks;
|
||||
for (size_t task = 0; task < graph.nodes.size(); ++task)
|
||||
if (!graph.nodes[task].crossbarUsage.empty())
|
||||
weightedTasks.push_back(task);
|
||||
llvm::sort(weightedTasks, [&](size_t lhs, size_t rhs) {
|
||||
if (graph.nodes[lhs].crossbarUsage.size() != graph.nodes[rhs].crossbarUsage.size())
|
||||
return graph.nodes[lhs].crossbarUsage.size() > graph.nodes[rhs].crossbarUsage.size();
|
||||
return graph.nodes[lhs].originalOrder < graph.nodes[rhs].originalOrder;
|
||||
});
|
||||
|
||||
std::vector<CrossbarUsage> residency(processorCount);
|
||||
for (size_t task : weightedTasks) {
|
||||
size_t bestProcessor = std::numeric_limits<size_t>::max();
|
||||
using ResidencyScore = std::tuple<size_t, size_t, size_t, size_t>;
|
||||
std::optional<ResidencyScore> bestScore;
|
||||
for (size_t processor = 0; processor < processorCount; ++processor) {
|
||||
size_t crossbarUnion = getCrossbarUnionSize(residency[processor], graph.nodes[task].crossbarUsage);
|
||||
if (crossbarUnion > crossbarCapacity)
|
||||
continue;
|
||||
size_t addedCrossbars = crossbarUnion - residency[processor].size();
|
||||
ResidencyScore score {addedCrossbars,
|
||||
crossbarCapacity - crossbarUnion,
|
||||
mesh.getCenterDistance(processor),
|
||||
processor};
|
||||
if (!bestScore || score < *bestScore) {
|
||||
bestProcessor = processor;
|
||||
bestScore = score;
|
||||
}
|
||||
}
|
||||
if (bestProcessor == std::numeric_limits<size_t>::max()) {
|
||||
std::string message =
|
||||
llvm::formatv("PEFT residency planner: cannot place task {0} with {1} distinct weights in {2} "
|
||||
"processors of capacity {3}",
|
||||
graph.nodes[task].originalOrder,
|
||||
graph.nodes[task].crossbarUsage.size(),
|
||||
processorCount,
|
||||
crossbarCapacity)
|
||||
.str();
|
||||
llvm::report_fatal_error(llvm::StringRef(message));
|
||||
}
|
||||
insertCrossbarWeights(residency[bestProcessor], graph.nodes[task].crossbarUsage);
|
||||
}
|
||||
return residency;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
Time getPeftTransferTime(Time transferCost, size_t sourceProcessor, size_t targetProcessor, size_t processorCount) {
|
||||
@@ -145,6 +196,8 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
if (processorCount == 0)
|
||||
llvm::report_fatal_error("PEFT scheduler: processor count must be positive");
|
||||
MeshModel mesh = MeshModel::infer(processorCount);
|
||||
std::vector<CrossbarUsage> plannedResidency =
|
||||
planCrossbarResidency(graph, processorCount, options.crossbarCapacity, mesh);
|
||||
|
||||
verifyOctTableSize(nodeCount, processorCount);
|
||||
std::vector<std::vector<size_t>> reverseLevels = buildReverseLevels(graph);
|
||||
@@ -237,20 +290,23 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
size_t bestProcessor = std::numeric_limits<size_t>::max();
|
||||
Time bestEst = 0;
|
||||
Time bestEft = 0;
|
||||
Time bestOeft = std::numeric_limits<Time>::max();
|
||||
unsigned int bestOverlapCount = 0;
|
||||
size_t bestCenterDistance = std::numeric_limits<size_t>::max();
|
||||
using CandidateScore = std::tuple<size_t, size_t, Time, Time, Time, size_t, unsigned int>;
|
||||
std::optional<CandidateScore> bestScore;
|
||||
size_t smallestCrossbarUnion = std::numeric_limits<size_t>::max();
|
||||
bool crossbarRejected = false;
|
||||
|
||||
for (size_t processor = 0; processor < processorCount; ++processor) {
|
||||
unsigned int overlapCount = countCrossbarOverlap(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
if (!graph.nodes[task].crossbarUsage.empty()
|
||||
&& getCrossbarUnionSize(processorCrossbars[processor], graph.nodes[task].crossbarUsage)
|
||||
> options.crossbarCapacity) {
|
||||
&& countCrossbarOverlap(plannedResidency[processor], graph.nodes[task].crossbarUsage)
|
||||
!= graph.nodes[task].crossbarUsage.size())
|
||||
continue;
|
||||
unsigned int overlapCount = countCrossbarOverlap(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
size_t crossbarUnion = getCrossbarUnionSize(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
smallestCrossbarUnion = std::min(smallestCrossbarUnion, crossbarUnion);
|
||||
if (!graph.nodes[task].crossbarUsage.empty() && crossbarUnion > options.crossbarCapacity) {
|
||||
crossbarRejected = true;
|
||||
continue;
|
||||
}
|
||||
|
||||
Time dataReady = 0;
|
||||
for (const auto& [pred, comm] : graph.predecessors[task]) {
|
||||
const ScheduledTask& predSchedule = schedules[pred];
|
||||
@@ -282,41 +338,32 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
Time eft = addOrMax(est, computeCost);
|
||||
Time oeft = addOrMax(eft, oct[task * processorCount + processor]);
|
||||
size_t centerDistance = mesh.getCenterDistance(processor);
|
||||
|
||||
if (oeft < bestOeft || (oeft == bestOeft && eft < bestEft)
|
||||
|| (oeft == bestOeft && eft == bestEft && est < bestEst)) {
|
||||
CandidateScore score {0,
|
||||
0,
|
||||
oeft,
|
||||
eft,
|
||||
est,
|
||||
centerDistance,
|
||||
overlapCount};
|
||||
if (!bestScore || score < *bestScore) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
}
|
||||
else if (oeft == bestOeft && eft == bestEft && est == bestEst
|
||||
&& centerDistance < bestCenterDistance) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
}
|
||||
else if (oeft == bestOeft && eft == bestEft && est == bestEst
|
||||
&& centerDistance == bestCenterDistance && overlapCount < bestOverlapCount) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
bestScore = score;
|
||||
}
|
||||
}
|
||||
|
||||
if (bestProcessor == std::numeric_limits<size_t>::max()) {
|
||||
if (crossbarRejected) {
|
||||
const ComputeInstance& instance = graph.nodes[task].instance;
|
||||
std::string message =
|
||||
llvm::formatv("PEFT scheduler: no valid processor for task {0}; crossbar capacity {1} is exhausted",
|
||||
llvm::formatv("PEFT scheduler: no valid processor for task {0} (lanes {1}..{2}, {3} distinct weights); "
|
||||
"smallest processor union is {4}, exceeding crossbar capacity {5}",
|
||||
graph.nodes[task].originalOrder,
|
||||
instance.laneStart,
|
||||
instance.laneStart + instance.laneCount,
|
||||
graph.nodes[task].crossbarUsage.size(),
|
||||
smallestCrossbarUnion,
|
||||
options.crossbarCapacity)
|
||||
.str();
|
||||
llvm::report_fatal_error(llvm::StringRef(message));
|
||||
|
||||
Reference in New Issue
Block a user