Files
Raptor/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp
T
ilgeco 0aa3840a72
Validate Operations / validate-operations (push) Has been cancelled
Conv lowering report
2026-08-03 12:05:03 +02:00

3654 lines
184 KiB
C++

#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/IR/BuiltinAttributes.h"
#include "mlir/IR/BuiltinTypes.h"
#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/SmallPtrSet.h"
#include "llvm/ADT/SmallVector.h"
#include "llvm/Support/raw_ostream.h"
#include <algorithm>
#include <mutex>
#include <optional>
#include <string>
#include "src/Accelerators/PIM/Common/IR/AffineUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/TensorSliceUtils.hpp"
#include "src/Accelerators/PIM/Common/Support/Diagnostics.hpp"
#include "src/Accelerators/PIM/Common/Support/ReportUtils.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Dialect/ONNX/ONNXOps.hpp"
using namespace mlir;
namespace onnx_mlir {
namespace {
struct ConvToGemm : OpConversionPattern<ONNXConvOp> {
explicit ConvToGemm(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpConversionPattern<ONNXConvOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(ONNXConvOp convOp,
ONNXConvOpAdaptor convOpAdaptor,
ConversionPatternRewriter& rewriter) const override;
const spatial::SpatialTargetInfo& target;
};
struct PreparedConvInput {
Value value;
RankedTensorType type;
};
static Value createZeroGemmBias(RankedTensorType gemmResultType, PatternRewriter& rewriter);
static StringRef stringifyConvLoweringStrategy(spatial::ConvLoweringStrategy strategy) {
switch (strategy) {
case spatial::ConvLoweringStrategy::Auto: return "auto";
case spatial::ConvLoweringStrategy::Legacy: return "legacy";
case spatial::ConvLoweringStrategy::Depthwise: return "depthwise";
case spatial::ConvLoweringStrategy::PackedIm2Col: return "packed-im2col";
case spatial::ConvLoweringStrategy::StreamedPatch: return "streamed-patch";
case spatial::ConvLoweringStrategy::StreamedPacked: return "streamed-packed";
case spatial::ConvLoweringStrategy::OutputChannelTiled: return "output-channel-tiled";
case spatial::ConvLoweringStrategy::InputKTiled: return "input-k-tiled";
case spatial::ConvLoweringStrategy::Tiled2D: return "tiled-2d";
}
llvm_unreachable("unknown conv lowering strategy");
}
enum class ConvLoweringReportPhase { Planning, Realization };
struct ConvLoweringReportEntry {
size_t convId;
std::string phase;
std::string location;
std::string strategy;
std::string implementation;
};
struct ConvLoweringReportState {
std::mutex mutex;
llvm::SmallPtrSet<Operation*, 32> planned;
llvm::SmallPtrSet<Operation*, 32> realized;
llvm::DenseMap<Operation*, size_t> convIds;
llvm::SmallVector<ConvLoweringReportEntry, 256> entries;
size_t nextConvId = 1;
};
static StringRef stringifyConvLoweringReportPhase(ConvLoweringReportPhase phase) {
return phase == ConvLoweringReportPhase::Planning ? "planning" : "realization";
}
static std::string convReportLocation(Operation* op) {
std::string location;
llvm::raw_string_ostream stream(location);
op->getLoc().print(stream);
if (location.size() > 120)
location.replace(117, std::string::npos, "...");
return location;
}
static StringRef convLoweringImplementation(spatial::ConvLoweringStrategy strategy) {
switch (strategy) {
case spatial::ConvLoweringStrategy::Depthwise:
return "DW";
case spatial::ConvLoweringStrategy::Legacy:
case spatial::ConvLoweringStrategy::PackedIm2Col:
return "PIC";
case spatial::ConvLoweringStrategy::StreamedPatch:
case spatial::ConvLoweringStrategy::OutputChannelTiled:
case spatial::ConvLoweringStrategy::Tiled2D:
return "STR";
case spatial::ConvLoweringStrategy::InputKTiled:
return "IKT";
case spatial::ConvLoweringStrategy::StreamedPacked:
return "STP";
case spatial::ConvLoweringStrategy::Auto:
return "AUTO";
}
llvm_unreachable("unknown conv lowering implementation");
}
static StringRef convRowStripInputImplementation(const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy) {
if (strategy == spatial::ConvLoweringStrategy::Depthwise)
return "RSDW";
if (state.xHeight == 1 && state.xWidth == 1 && state.wHeight == 1 && state.wWidth == 1)
return "RSP";
return "RSM";
}
static constexpr size_t kConvReportIdWidth = 4;
static constexpr size_t kConvReportLocationWidth = 24;
static constexpr size_t kConvReportStrategyWidth = 20;
static constexpr size_t kConvReportCodeWidth = 8;
static std::string convReportCell(StringRef value, size_t width) {
std::string cell = value.str();
if (cell.size() > width) {
cell = width <= 3 ? std::string(width, '.') : cell.substr(0, width - 3) + "...";
}
cell.append(width - cell.size(), ' ');
return cell;
}
static void writeConvReportTableHeader(std::fstream& reportFile, StringRef fourthColumn) {
reportFile << "+------+--------------------------+----------------------+----------+\n";
reportFile << "| " << convReportCell("Conv", kConvReportIdWidth) << " | "
<< convReportCell("Location", kConvReportLocationWidth) << " | "
<< convReportCell("Strategy", kConvReportStrategyWidth) << " | "
<< convReportCell(fourthColumn, kConvReportCodeWidth) << " |\n";
reportFile << "+------+--------------------------+----------------------+----------+\n";
}
static void writeConvReportRow(std::fstream& reportFile,
const ConvLoweringReportEntry& entry) {
reportFile << "| " << convReportCell(std::to_string(entry.convId), kConvReportIdWidth) << " | "
<< convReportCell(entry.location, kConvReportLocationWidth) << " | "
<< convReportCell(entry.strategy, kConvReportStrategyWidth) << " | "
<< convReportCell(entry.implementation, kConvReportCodeWidth) << " |\n";
}
static void writeConvReportLegend(std::fstream& reportFile) {
reportFile << "Legend: Conv is shared by both sections; codes expand to:\n";
reportFile << " SEL selectConvLoweringPlan\n";
reportFile << " DW depthwise::rewriteConv\n";
reportFile << " PIC standard::rewritePackedIm2ColConv\n";
reportFile << " STR standard::rewriteStreamedConv(pack=1)\n";
reportFile << " IKT standard::rewriteInputKTiledConv\n";
reportFile << " STP standard::rewriteStreamedConv(pack=geo.pack)\n";
reportFile << " AUTO unresolved strategy\n";
reportFile << " RSD createRowStripConvOutputFromDenseInput -> createRowStripConvOutput\n";
reportFile << " RSDW createConvOutputFromRowStripInput -> createDepthwiseOutputFromRowStripFragments\n";
reportFile << " RSP createConvOutputFromRowStripInput -> createPointwiseOutputFromRowStripFragments\n";
reportFile << " RSM createConvOutputFromRowStripInput -> createConvOutputFromPixelMajorRowStripFragments\n\n";
}
static bool writeConvLoweringReport(const ConvLoweringReportEntry& entry,
ConvLoweringReportState& state) {
state.entries.push_back(entry);
std::fstream reportFile = openReportFile("conv_lowering_report");
if (!reportFile.is_open()) {
state.entries.pop_back();
return false;
}
reportFile << "# PIM Conv Lowering Report (bounded to 512 rows)\n\n";
reportFile << "## Plan selection\n";
writeConvReportTableHeader(reportFile, "Selector");
bool realizationSectionStarted = false;
for (const ConvLoweringReportEntry& reportEntry : state.entries) {
if (reportEntry.phase == "realization" && !realizationSectionStarted) {
reportFile << "\n## Realization\n";
writeConvReportTableHeader(reportFile, "Code");
realizationSectionStarted = true;
}
writeConvReportRow(reportFile, reportEntry);
}
reportFile << "\n";
writeConvReportLegend(reportFile);
if (!reportFile.good()) {
state.entries.pop_back();
return false;
}
return true;
}
static void recordConvLoweringReport(Operation* op,
ConvLoweringReportPhase phase,
spatial::ConvLoweringStrategy strategy,
StringRef implementation) {
if (!pimReportConvLowering)
return;
static ConvLoweringReportState state;
std::lock_guard<std::mutex> lock(state.mutex);
if (state.entries.size() >= 512)
return;
if (phase == ConvLoweringReportPhase::Planning) {
if (state.planned.contains(op))
return;
}
else {
if (state.realized.contains(op))
return;
}
size_t convId = state.convIds.lookup(op);
if (!convId) {
convId = state.nextConvId++;
state.convIds[op] = convId;
}
ConvLoweringReportEntry entry {convId,
stringifyConvLoweringReportPhase(phase).str(),
convReportLocation(op),
stringifyConvLoweringStrategy(strategy).str(),
implementation.str()};
if (!writeConvLoweringReport(entry, state))
return;
if (phase == ConvLoweringReportPhase::Planning)
state.planned.insert(op);
else
state.realized.insert(op);
}
static Value expandBiasIfNeeded(Value bias, PatternRewriter& rewriter, Location loc) {
auto biasType = cast<RankedTensorType>(bias.getType());
if (biasType.getRank() != 1)
return bias;
auto expandedBiasType = RankedTensorType::get({1, biasType.getDimSize(0)}, biasType.getElementType());
return tensor::ExpandShapeOp::create(rewriter,
loc,
expandedBiasType,
bias,
SmallVector<ReassociationIndices> {
{0, 1}
});
}
static int64_t findLargestDivisorAtMost(int64_t value, int64_t limit) {
assert(value > 0 && "expected positive value");
limit = std::min(value, limit);
for (int64_t candidate = limit; candidate >= 1; --candidate)
if (value % candidate == 0)
return candidate;
return 1;
}
static Value createZeroPaddedTensor(Value value,
RankedTensorType resultType,
ArrayRef<int64_t> lowPadValues,
ArrayRef<int64_t> highPadValues,
PatternRewriter& rewriter,
Location loc) {
auto valueType = cast<RankedTensorType>(value.getType());
if (valueType == resultType)
return value;
SmallVector<OpFoldResult> lowPads;
SmallVector<OpFoldResult> highPads;
lowPads.reserve(lowPadValues.size());
highPads.reserve(highPadValues.size());
for (auto lowPad : lowPadValues)
lowPads.push_back(rewriter.getIndexAttr(lowPad));
for (auto highPad : highPadValues)
highPads.push_back(rewriter.getIndexAttr(highPad));
auto padOp = tensor::PadOp::create(rewriter, loc, resultType, value, lowPads, highPads);
auto* padBlock = new Block();
for (int64_t dim = 0, rank = resultType.getRank(); dim < rank; ++dim)
padBlock->addArgument(rewriter.getIndexType(), loc);
padOp.getRegion().push_back(padBlock);
rewriter.setInsertionPointToStart(padBlock);
auto zero = getOrCreateConstant(
rewriter, padOp.getOperation(), rewriter.getZeroAttr(resultType.getElementType()), resultType.getElementType());
tensor::YieldOp::create(rewriter, loc, zero);
rewriter.setInsertionPointAfter(padOp);
return padOp.getResult();
}
static Value createConvInputPatch(Value input,
RankedTensorType patchType,
Value batchIndex,
Value channelOffset,
Value inputHeightOffset,
Value inputWidthOffset,
int64_t dilationHeight,
int64_t dilationWidth,
PatternRewriter& rewriter,
Location loc) {
const int64_t patchChannels = patchType.getDimSize(1);
const int64_t kernelHeight = patchType.getDimSize(2);
const int64_t kernelWidth = patchType.getDimSize(3);
if (dilationHeight == 1 && dilationWidth == 1) {
SmallVector<OpFoldResult> offsets {batchIndex, channelOffset, inputHeightOffset, inputWidthOffset};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(patchChannels),
rewriter.getIndexAttr(kernelHeight),
rewriter.getIndexAttr(kernelWidth)};
return tensor::ExtractSliceOp::create(rewriter, loc, patchType, input, offsets, sizes, getUnitStrides(rewriter, 4));
}
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto elementType = patchType.getElementType();
auto pixelType = RankedTensorType::get({1, patchChannels, 1, 1}, elementType, patchType.getEncoding());
Value patch = tensor::EmptyOp::create(rewriter, loc, patchType.getShape(), elementType);
for (int64_t kernelH = 0; kernelH < kernelHeight; ++kernelH) {
Value sourceHeightOffset = affineAddConst(rewriter, loc, inputHeightOffset, kernelH * dilationHeight, anchorOp);
for (int64_t kernelW = 0; kernelW < kernelWidth; ++kernelW) {
Value sourceWidthOffset = affineAddConst(rewriter, loc, inputWidthOffset, kernelW * dilationWidth, anchorOp);
SmallVector<OpFoldResult> sourceOffsets {batchIndex, channelOffset, sourceHeightOffset, sourceWidthOffset};
SmallVector<OpFoldResult> sourceSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(patchChannels),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1)};
Value sourcePixel = tensor::ExtractSliceOp::create(
rewriter, loc, pixelType, input, sourceOffsets, sourceSizes, getUnitStrides(rewriter, 4));
SmallVector<OpFoldResult> targetOffsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(kernelH), rewriter.getIndexAttr(kernelW)};
patch = tensor::InsertSliceOp::create(
rewriter, loc, sourcePixel, patch, targetOffsets, sourceSizes, getUnitStrides(rewriter, 4));
}
}
return patch;
}
static Value createCollectedConvOutput(ValueRange gemmRows,
Type convType,
RankedTensorType gemmOutType,
RankedTensorType nhwcType,
RankedTensorType outType,
int64_t numPatches,
int64_t numChannelsOut,
int64_t packFactor,
PatternRewriter& rewriter,
Location loc);
static FailureOr<ConvLoweringState> analyzeConvLoweringState(ONNXConvOp convOp,
Value x,
Value w,
Value b,
const spatial::SpatialTargetInfo& target);
namespace depthwise {
struct Tiling {
int64_t outputMultiplier;
int64_t kernelElements;
int64_t channelsPerTile;
int64_t tileInputRows;
int64_t tileOutputChannels;
int64_t numChannelTiles;
int64_t spatialPatchesPerBatch;
int64_t totalPatches;
};
static std::optional<Tiling> computeTiling(int64_t batchSize,
int64_t numChannelsIn,
int64_t numChannelsOut,
int64_t wHeight,
int64_t wWidth,
int64_t outHeight,
int64_t outWidth,
int64_t xbarDim) {
const int64_t kernelElements = wHeight * wWidth;
const int64_t outputMultiplier = numChannelsOut / numChannelsIn;
if (kernelElements <= 0 || outputMultiplier <= 0 || kernelElements > xbarDim || outputMultiplier > xbarDim)
return std::nullopt;
const int64_t maxChannelsPerTile = std::min(xbarDim / kernelElements, xbarDim / outputMultiplier);
if (maxChannelsPerTile <= 0)
return std::nullopt;
const int64_t channelsPerTile = findLargestDivisorAtMost(numChannelsIn, maxChannelsPerTile);
const int64_t tileInputRows = channelsPerTile * kernelElements;
const int64_t tileOutputChannels = channelsPerTile * outputMultiplier;
if (tileInputRows > xbarDim || tileOutputChannels > xbarDim)
return std::nullopt;
return Tiling {
outputMultiplier,
kernelElements,
channelsPerTile,
tileInputRows,
tileOutputChannels,
numChannelsIn / channelsPerTile,
outHeight * outWidth,
batchSize * outHeight * outWidth,
};
}
static Value buildPackedWeights(DenseElementsAttr wDenseAttr,
RankedTensorType wType,
const Tiling& tiling,
PatternRewriter& rewriter,
Location loc,
int64_t xbarDim,
int64_t paddedInputRows = -1) {
const int64_t paddedOutputChannels = xbarDim;
const int64_t packedInputRows = paddedInputRows > 0 ? paddedInputRows : tiling.tileInputRows;
auto packedWeightType = RankedTensorType::get(
{tiling.numChannelTiles, packedInputRows, paddedOutputChannels}, wType.getElementType());
SmallVector<Attribute> packedValues(packedWeightType.getNumElements(),
cast<Attribute>(rewriter.getZeroAttr(wType.getElementType())));
SmallVector<Attribute> sourceValues(wDenseAttr.getValues<Attribute>());
for (int64_t tileIndex = 0; tileIndex < tiling.numChannelTiles; ++tileIndex) {
const int64_t channelBase = tileIndex * tiling.channelsPerTile;
for (int64_t localChannel = 0; localChannel < tiling.channelsPerTile; ++localChannel) {
const int64_t globalChannel = channelBase + localChannel;
for (int64_t kernelIndex = 0; kernelIndex < tiling.kernelElements; ++kernelIndex) {
const int64_t kernelH = kernelIndex / wType.getDimSize(3);
const int64_t kernelW = kernelIndex % wType.getDimSize(3);
const int64_t targetRow = kernelIndex * tiling.channelsPerTile + localChannel;
for (int64_t multiplierIndex = 0; multiplierIndex < tiling.outputMultiplier; ++multiplierIndex) {
const int64_t globalOutChannel = globalChannel * tiling.outputMultiplier + multiplierIndex;
const int64_t sourceFlatIndex =
((globalOutChannel * wType.getDimSize(1) * wType.getDimSize(2)) + kernelH) * wType.getDimSize(3) + kernelW;
const int64_t targetCol = localChannel * tiling.outputMultiplier + multiplierIndex;
const int64_t targetFlatIndex =
((tileIndex * packedInputRows) + targetRow) * paddedOutputChannels + targetCol;
packedValues[targetFlatIndex] = sourceValues[sourceFlatIndex];
}
}
}
}
auto packedAttr = DenseElementsAttr::get(packedWeightType, packedValues);
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), packedAttr, packedWeightType);
}
static Value createPaddedInput(Value input,
RankedTensorType inputType,
int64_t padHeightBegin,
int64_t padHeightEnd,
int64_t padWidthBegin,
int64_t padWidthEnd,
PatternRewriter& rewriter,
Location loc) {
if (padHeightBegin == 0 && padHeightEnd == 0 && padWidthBegin == 0 && padWidthEnd == 0)
return input;
auto paddedInputType = RankedTensorType::get({inputType.getDimSize(0),
inputType.getDimSize(1),
inputType.getDimSize(2) + padHeightBegin + padHeightEnd,
inputType.getDimSize(3) + padWidthBegin + padWidthEnd},
inputType.getElementType());
auto computeOp = createSpatCompute<1>(rewriter, loc, TypeRange {paddedInputType}, {}, input, [&](Value computeInput) {
Value padded = createZeroPaddedTensor(computeInput,
paddedInputType,
{0, 0, padHeightBegin, padWidthBegin},
{0, 0, padHeightEnd, padWidthEnd},
rewriter,
loc);
spatial::SpatYieldOp::create(rewriter, loc, padded);
});
return computeOp.getResult(0);
}
static Value createInputTile(Value input,
Value patchIndex,
Value channelTileIndex,
RankedTensorType inputTileType,
const Tiling& tiling,
int64_t strideHeight,
int64_t strideWidth,
int64_t dilationHeight,
int64_t dilationWidth,
int64_t outWidth,
PatternRewriter& rewriter,
Location loc) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value batchIndex = affineFloorDivConst(rewriter, loc, patchIndex, tiling.spatialPatchesPerBatch, anchorOp);
Value batchPatchIndex = affineModConst(rewriter, loc, patchIndex, tiling.spatialPatchesPerBatch, anchorOp);
Value outHeightIndex = affineFloorDivConst(rewriter, loc, batchPatchIndex, outWidth, anchorOp);
Value outWidthIndex = affineModConst(rewriter, loc, batchPatchIndex, outWidth, anchorOp);
Value inputHeightOffset =
strideHeight == 1 ? outHeightIndex : affineMulConst(rewriter, loc, outHeightIndex, strideHeight, anchorOp);
Value inputWidthOffset =
strideWidth == 1 ? outWidthIndex : affineMulConst(rewriter, loc, outWidthIndex, strideWidth, anchorOp);
Value channelOffset = tiling.channelsPerTile == 1
? channelTileIndex
: affineMulConst(rewriter, loc, channelTileIndex, tiling.channelsPerTile, anchorOp);
Value tile4D;
if (dilationHeight == 1 && dilationWidth == 1) {
SmallVector<OpFoldResult> offsets {batchIndex, inputHeightOffset, inputWidthOffset, channelOffset};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(inputTileType.getDimSize(1)),
rewriter.getIndexAttr(inputTileType.getDimSize(2)),
rewriter.getIndexAttr(tiling.channelsPerTile)};
tile4D = tensor::ExtractSliceOp::create(
rewriter, loc, inputTileType, input, offsets, sizes, getUnitStrides(rewriter, 4));
}
else {
auto pixelType = RankedTensorType::get(
{1, 1, 1, tiling.channelsPerTile}, inputTileType.getElementType(), inputTileType.getEncoding());
tile4D = tensor::EmptyOp::create(rewriter, loc, inputTileType.getShape(), inputTileType.getElementType());
for (int64_t kernelH = 0; kernelH < inputTileType.getDimSize(1); ++kernelH)
for (int64_t kernelW = 0; kernelW < inputTileType.getDimSize(2); ++kernelW) {
Value sourceHeight = affineAddConst(rewriter, loc, inputHeightOffset, kernelH * dilationHeight, anchorOp);
Value sourceWidth = affineAddConst(rewriter, loc, inputWidthOffset, kernelW * dilationWidth, anchorOp);
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tiling.channelsPerTile)};
Value pixel = tensor::ExtractSliceOp::create(
rewriter,
loc,
pixelType,
input,
SmallVector<OpFoldResult> {batchIndex, sourceHeight, sourceWidth, channelOffset},
sizes,
getUnitStrides(rewriter, 4));
tile4D = tensor::InsertSliceOp::create(
rewriter,
loc,
pixel,
tile4D,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelH),
rewriter.getIndexAttr(kernelW),
rewriter.getIndexAttr(0)},
sizes,
getUnitStrides(rewriter, 4));
}
}
auto collapsedType = RankedTensorType::get({1, tiling.tileInputRows}, inputTileType.getElementType());
return tensor::CollapseShapeOp::create(rewriter,
loc,
collapsedType,
tile4D,
SmallVector<ReassociationIndices> {
{0},
{1, 2, 3}
});
}
static Value createWeightTile(Value packedWeights,
Value channelTileIndex,
RankedTensorType packedWeightType,
const Tiling& tiling,
PatternRewriter& rewriter,
Location loc) {
SmallVector<OpFoldResult> offsets {channelTileIndex, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
const int64_t paddedOutputChannels = packedWeightType.getDimSize(2);
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tiling.tileInputRows),
rewriter.getIndexAttr(paddedOutputChannels)};
auto collapsedType =
RankedTensorType::get({tiling.tileInputRows, paddedOutputChannels}, packedWeightType.getElementType());
return extractMixedSliceOrIdentity(
rewriter, loc, packedWeights, collapsedType,
{offsets, sizes, getUnitStrides(rewriter, 3)});
}
static Value createBiasTile(
Value bias, Value channelTileIndex, const Tiling& tiling, PatternRewriter& rewriter, Location loc) {
auto biasType = cast<RankedTensorType>(bias.getType());
auto biasTileType = RankedTensorType::get({1, tiling.tileOutputChannels}, biasType.getElementType());
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value channelOffset = tiling.tileOutputChannels == 1
? channelTileIndex
: affineMulConst(rewriter, loc, channelTileIndex, tiling.tileOutputChannels, anchorOp);
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0), channelOffset};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(tiling.tileOutputChannels)};
return tensor::ExtractSliceOp::create(rewriter, loc, biasTileType, bias, offsets, sizes, getUnitStrides(rewriter, 2));
}
static Value insertOutputTile(Value rowTile,
Value rowAccumulator,
Value channelTileIndex,
const Tiling& tiling,
PatternRewriter& rewriter,
Location loc) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value channelOffset = tiling.tileOutputChannels == 1
? channelTileIndex
: affineMulConst(rewriter, loc, channelTileIndex, tiling.tileOutputChannels, anchorOp);
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0), channelOffset};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(tiling.tileOutputChannels)};
return tensor::InsertSliceOp::create(
rewriter, loc, rowTile, rowAccumulator, offsets, sizes, getUnitStrides(rewriter, 2));
}
static FailureOr<Value> reconstructDepthwiseGemmRows(Value pieces,
RankedTensorType piecesType,
RankedTensorType gemmOutType,
const Tiling& tiling,
PatternRewriter& rewriter,
Location loc) {
auto collectedOp = createSpatCompute<1>(rewriter, loc, TypeRange {gemmOutType}, {}, pieces, [&](Value piecesArg) {
auto rowType = RankedTensorType::get({1, gemmOutType.getDimSize(1)}, gemmOutType.getElementType());
Value outputInit = tensor::EmptyOp::create(rewriter, loc, gemmOutType.getShape(), gemmOutType.getElementType());
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cNumPatches = getOrCreateIndexConstant(rewriter, anchorOp, tiling.totalPatches);
Value cNumChannelTiles = getOrCreateIndexConstant(rewriter, anchorOp, tiling.numChannelTiles);
auto patchLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cNumPatches,
c1,
ValueRange {outputInit},
[&](OpBuilder&,
Location nestedLoc,
Value patchIndex,
ValueRange patchIterArgs,
SmallVectorImpl<Value>& patchYielded) {
Value outputAcc = patchIterArgs.front();
Value rowInit = tensor::EmptyOp::create(rewriter, nestedLoc, rowType.getShape(), rowType.getElementType());
auto tileLoop = buildNormalizedScfFor(
rewriter,
nestedLoc,
c0,
cNumChannelTiles,
c1,
ValueRange {rowInit},
[&](OpBuilder&,
Location tileLoc,
Value channelTileIndex,
ValueRange tileIterArgs,
SmallVectorImpl<Value>& tileYielded) {
Value rowAcc = tileIterArgs.front();
MLIRContext* context = rewriter.getContext();
AffineExpr d0 = getAffineDimExpr(0, context);
AffineExpr d1 = getAffineDimExpr(1, context);
Value laneIndex = createOrFoldAffineApply(
rewriter, tileLoc, (d0 * tiling.totalPatches) + d1, ValueRange {channelTileIndex, patchIndex}, anchorOp);
auto rowTileType = RankedTensorType::get({1, tiling.tileOutputChannels}, piecesType.getElementType());
FailureOr<Value> rowTile = extractGraphBatchPhysicalFragment(rewriter, tileLoc, piecesArg, laneIndex, rowTileType);
if (failed(rowTile))
return failure();
Value rowNext = insertOutputTile(*rowTile, rowAcc, channelTileIndex, tiling, rewriter, tileLoc);
tileYielded.push_back(rowNext);
return success();
});
if (failed(tileLoop))
return failure();
SmallVector<OpFoldResult> rowOffsets {patchIndex, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> rowSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(gemmOutType.getDimSize(1))};
Value outputNext = tensor::InsertSliceOp::create(rewriter,
nestedLoc,
tileLoop->results.front(),
outputAcc,
rowOffsets,
rowSizes,
getUnitStrides(rewriter, 2))
.getResult();
patchYielded.push_back(outputNext);
return success();
});
if (failed(patchLoop))
return failure();
spatial::SpatYieldOp::create(rewriter, loc, patchLoop->results.front());
return success();
});
if (failed(collectedOp))
return failure();
return collectedOp->getResult(0);
}
static bool canUseStructuredRewrite(const ConvLoweringState& state) {
if (!getHostConstDenseElementsAttr(state.w))
return false;
auto tiling = computeTiling(state.batchSize,
state.numChannelsIn,
state.numChannelsOut,
state.wHeight,
state.wWidth,
state.outHeight,
state.outWidth,
state.targetInfo().matrixShape.rows);
if (!tiling)
return false;
if (!state.hasBias)
return true;
auto biasType = dyn_cast<RankedTensorType>(state.b.getType());
if (!biasType)
return false;
if (biasType.getRank() == 1)
return biasType.getDimSize(0) == state.numChannelsOut;
if (biasType.getRank() != 2)
return false;
return biasType.getDimSize(0) == 1 && biasType.getDimSize(1) == state.numChannelsOut;
}
static FailureOr<Value>
rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) {
auto wDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!wDenseAttr) {
convOp->emitOpError("requires constant-derived weights for structured depthwise Spatial lowering");
return failure();
}
auto tiling = computeTiling(state.xType.getDimSize(0),
state.xType.getDimSize(1),
state.outType.getDimSize(1),
state.wType.getDimSize(2),
state.wType.getDimSize(3),
state.outType.getDimSize(2),
state.outType.getDimSize(3),
state.targetInfo().matrixShape.rows);
if (!tiling) {
convOp->emitOpError("failed to derive a structured depthwise tiling that fits Spatial weighted VMM lowering");
return failure();
}
Value paddedInput = createPaddedInput(state.x,
state.xType,
state.padHeightBegin,
state.padHeightEnd,
state.padWidthBegin,
state.padWidthEnd,
rewriter,
loc);
auto paddedInputType = cast<RankedTensorType>(paddedInput.getType());
auto channelLastInputType = RankedTensorType::get({paddedInputType.getDimSize(0),
paddedInputType.getDimSize(2),
paddedInputType.getDimSize(3),
paddedInputType.getDimSize(1)},
paddedInputType.getElementType());
Value channelLastInput = createLinalgTranspose(
paddedInput, channelLastInputType, {0, 2, 3, 1}, rewriter, loc);
Value packedWeights = buildPackedWeights(
wDenseAttr, state.wType, *tiling, rewriter, loc, state.targetInfo().matrixShape.rows);
Value expandedBias;
SmallVector<Value> batchInputs {channelLastInput};
if (state.hasBias) {
expandedBias = expandBiasIfNeeded(state.b, rewriter, loc);
auto biasType = dyn_cast<RankedTensorType>(expandedBias.getType());
if (!biasType || biasType.getRank() != 2 || biasType.getDimSize(0) != 1
|| biasType.getDimSize(1) != state.outType.getDimSize(1)) {
convOp->emitOpError("requires bias sliceable as tensor<1xCout> for structured depthwise Spatial lowering");
return failure();
}
batchInputs.push_back(expandedBias);
}
auto gemmOutType =
RankedTensorType::get({tiling->totalPatches, state.outType.getDimSize(1)}, state.outType.getElementType());
auto rowTileType = RankedTensorType::get({1, tiling->tileOutputChannels}, state.outType.getElementType());
auto paddedRowTileType = RankedTensorType::get(
{1, static_cast<int64_t>(state.targetInfo().matrixShape.rows)}, state.outType.getElementType());
auto piecesType = spatial::getGraphBatchPhysicalResultType(
tiling->totalPatches * tiling->numChannelTiles, rowTileType);
auto inputTileType =
RankedTensorType::get({1, state.wType.getDimSize(2), state.wType.getDimSize(3), tiling->channelsPerTile},
paddedInputType.getElementType());
SmallVector<Value> batchWeights;
if (tiling->numChannelTiles == 1) {
Value c0 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0);
batchWeights.push_back(createWeightTile(packedWeights,
c0,
cast<RankedTensorType>(packedWeights.getType()),
*tiling,
rewriter,
loc));
}
else {
batchWeights.push_back(packedWeights);
}
auto batchOp = createSpatComputeBatch(
rewriter,
loc,
TypeRange {piecesType},
tiling->totalPatches * tiling->numChannelTiles,
batchWeights,
batchInputs,
[&](detail::SpatComputeBatchBodyArgs args) {
auto pickInputByRank = [&](int64_t rank) -> Value {
for (Value input : args.inputs) {
auto inputType = dyn_cast<RankedTensorType>(input.getType());
if (inputType && inputType.getRank() == rank)
return input;
}
return Value();
};
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value patchIndex = tiling->numChannelTiles == 1
? args.lane
: affineModConst(rewriter, loc, args.lane, tiling->totalPatches, anchorOp);
Value channelTileIndex = tiling->numChannelTiles == 1
? getOrCreateIndexConstant(rewriter, anchorOp, 0)
: affineFloorDivConst(rewriter, loc, args.lane, tiling->totalPatches, anchorOp);
Value paddedInputArg = pickInputByRank(/*rank=*/4);
if (!paddedInputArg) {
convOp->emitOpError("structured depthwise batch body requires a rank-4 padded input block argument");
return failure();
}
Value inputTile = createInputTile(paddedInputArg,
patchIndex,
channelTileIndex,
inputTileType,
*tiling,
state.strideHeight,
state.strideWidth,
state.dilationHeight,
state.dilationWidth,
state.outType.getDimSize(3),
rewriter,
loc);
Value weightTile = tiling->numChannelTiles == 1
? args.weights.front()
: createWeightTile(args.weights.front(),
channelTileIndex,
cast<RankedTensorType>(args.weights.front().getType()),
*tiling,
rewriter,
loc);
Value paddedRowTile =
spatial::SpatVMMOp::create(rewriter, loc, paddedRowTileType, weightTile, inputTile).getResult();
Value rowTile = tensor::ExtractSliceOp::create(
rewriter,
loc,
rowTileType,
paddedRowTile,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tiling->tileOutputChannels)},
getUnitStrides(rewriter, 2));
if (args.inputs.size() > 1) {
Value biasArg = pickInputByRank(/*rank=*/2);
if (!biasArg) {
convOp->emitOpError("structured depthwise batch body requires a rank-2 bias block argument when bias is present");
return failure();
}
Value biasTile = tiling->numChannelTiles == 1 ? biasArg : createBiasTile(biasArg, channelTileIndex, *tiling, rewriter, loc);
rowTile = spatial::SpatVAddOp::create(rewriter, loc, rowTileType, rowTile, biasTile).getResult();
}
publishGraphBatchPhysicalFragment(rewriter, loc, rowTile, args.outputs.front(), args.lane);
return success();
});
if (failed(batchOp))
return failure();
auto nhwcType = RankedTensorType::get(
{state.xType.getDimSize(0), state.outType.getDimSize(2), state.outType.getDimSize(3), state.outType.getDimSize(1)},
state.outType.getElementType());
auto reconstructedRows = reconstructDepthwiseGemmRows(batchOp->getResult(0), piecesType, gemmOutType, *tiling, rewriter, loc);
if (failed(reconstructedRows))
return failure();
Value collectedRows = *reconstructedRows;
return createCollectedConvOutput(ValueRange {collectedRows},
state.outType,
gemmOutType,
nhwcType,
state.outType,
tiling->totalPatches,
state.outType.getDimSize(1),
/*packFactor=*/1,
rewriter,
loc);
}
} // namespace depthwise
namespace standard {
struct ConvGemmPlan {
int64_t patchSize;
int64_t numPatchesPerBatch;
int64_t globalNumPatches;
int64_t chunkStart;
int64_t chunkNumPatches;
int64_t maxParallelPixels;
int64_t effectiveMaxParallelPixels;
int64_t packedNumRows;
RankedTensorType gemmInputRowsType;
RankedTensorType wFlatType;
RankedTensorType wTransType;
RankedTensorType gemmOutType;
RankedTensorType gemmOutputRowsType;
RankedTensorType nhwcType;
};
static ConvGemmPlan
buildConvGemmPlan(const ConvLoweringState& state,
bool canPackWeightsAsConstants,
bool canPackBiasAsConstants,
int64_t chunkStart,
int64_t chunkNumPatches,
std::optional<int64_t> forcedPackFactor = std::nullopt);
static PreparedConvInput prepareInputForIm2Col(const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
if (state.padHeightBegin == 0 && state.padHeightEnd == 0 && state.padWidthBegin == 0 && state.padWidthEnd == 0)
return {state.x, state.xType};
auto paddedType = RankedTensorType::get({state.batchSize,
state.numChannelsIn,
state.xHeight + state.padHeightBegin + state.padHeightEnd,
state.xWidth + state.padWidthBegin + state.padWidthEnd},
state.xType.getElementType());
auto paddedInputOp =
createSpatCompute<1>(rewriter, loc, TypeRange {paddedType}, {}, state.x, [&](Value inputArg) {
Value paddedInput = createZeroPaddedTensor(inputArg,
paddedType,
{0, 0, state.padHeightBegin, state.padWidthBegin},
{0, 0, state.padHeightEnd, state.padWidthEnd},
rewriter,
loc);
spatial::SpatYieldOp::create(rewriter, loc, paddedInput);
});
return {paddedInputOp.getResult(0), paddedType};
}
static Value unpackRowsFromParallelGemm(Value packedRows,
RankedTensorType packedRowsType,
int64_t unpackedRows,
int64_t rowWidth,
int64_t packFactor,
PatternRewriter& rewriter,
Location loc) {
if (packFactor == 1)
return packedRows;
const int64_t packedNumRows = packedRowsType.getDimSize(0);
const int64_t paddedNumRows = packedNumRows * packFactor;
auto expandedType = RankedTensorType::get(
{packedNumRows, packFactor, rowWidth}, packedRowsType.getElementType(), packedRowsType.getEncoding());
auto paddedType =
RankedTensorType::get({paddedNumRows, rowWidth}, packedRowsType.getElementType(), packedRowsType.getEncoding());
auto unpackedType =
RankedTensorType::get({unpackedRows, rowWidth}, packedRowsType.getElementType(), packedRowsType.getEncoding());
Value expanded = tensor::ExpandShapeOp::create(rewriter,
loc,
expandedType,
packedRows,
SmallVector<ReassociationIndices> {
{0},
{1, 2}
});
Value padded = tensor::CollapseShapeOp::create(rewriter,
loc,
paddedType,
expanded,
SmallVector<ReassociationIndices> {
{0, 1},
{2}
});
if (paddedNumRows == unpackedRows)
return padded;
SmallVector<OpFoldResult> offsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(unpackedRows), rewriter.getIndexAttr(rowWidth)};
return tensor::ExtractSliceOp::create(rewriter, loc, unpackedType, padded, offsets, sizes, getUnitStrides(rewriter, 2));
}
static Value createWeightMatrix(
Value weights, const ConvGemmPlan& plan, bool transpose, PatternRewriter& rewriter, Location loc) {
auto buildWeightMatrix = [&](Value weight) -> Value {
Value flattened = tensor::CollapseShapeOp::create(rewriter,
loc,
plan.wFlatType,
weight,
SmallVector<ReassociationIndices> {
{0},
{1, 2, 3}
});
if (!transpose)
return flattened;
return createLinalgTranspose(flattened, plan.wTransType, {1, 0}, rewriter, loc);
};
if (isCompileTimeComputable(weights))
return buildWeightMatrix(weights);
RankedTensorType resultType = transpose ? plan.wTransType : plan.wFlatType;
auto computeOp =
createSpatCompute<1>(rewriter, loc, TypeRange {resultType}, {}, ValueRange {weights}, [&](Value weight) {
spatial::SpatYieldOp::create(rewriter, loc, buildWeightMatrix(weight));
});
return computeOp.getResult(0);
}
static Value createPaddedConvMatrix(Value matrix,
RankedTensorType sourceType,
RankedTensorType paddedType,
PatternRewriter& rewriter,
Location loc) {
if (sourceType == paddedType)
return matrix;
return createZeroPaddedTensor(matrix,
paddedType,
{0, 0},
{paddedType.getDimSize(0) - sourceType.getDimSize(0),
paddedType.getDimSize(1) - sourceType.getDimSize(1)},
rewriter,
loc);
}
static Value createPaddedConstantMatrix(DenseElementsAttr sourceAttr,
RankedTensorType sourceType,
RankedTensorType paddedType,
PatternRewriter& rewriter) {
SmallVector<Attribute> paddedValues(
paddedType.getNumElements(), cast<Attribute>(rewriter.getZeroAttr(paddedType.getElementType())));
SmallVector<Attribute> sourceValues(sourceAttr.getValues<Attribute>());
const int64_t sourceRows = sourceType.getDimSize(0);
const int64_t sourceCols = sourceType.getDimSize(1);
const int64_t paddedCols = paddedType.getDimSize(1);
for (int64_t row = 0; row < sourceRows; ++row)
for (int64_t col = 0; col < sourceCols; ++col)
paddedValues[row * paddedCols + col] = sourceValues[row * sourceCols + col];
auto paddedAttr = DenseElementsAttr::get(paddedType, paddedValues);
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType);
}
static Value createPaddedInputKTiledWeightConstant(DenseElementsAttr sourceAttr,
const ConvLoweringState& state,
int64_t paddedK,
int64_t paddedC,
PatternRewriter& rewriter) {
auto paddedType = RankedTensorType::get({paddedK, paddedC}, 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) {
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;
paddedValues[patchIndex * paddedC + outChannel] = sourceValues[sourceFlatIndex];
}
}
}
}
auto paddedAttr = DenseElementsAttr::get(paddedType, paddedValues);
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), paddedAttr, paddedType);
}
static Value createPaddedPixelMajorWeightConstant(DenseElementsAttr sourceAttr,
const ConvLoweringState& state,
int64_t paddedK,
int64_t paddedC,
int64_t packFactor,
PatternRewriter& rewriter) {
auto paddedType = RankedTensorType::get({paddedK, paddedC}, state.wType.getElementType());
SmallVector<Attribute> sourceValues(sourceAttr.getValues<Attribute>());
SmallVector<Attribute> paddedValues(
paddedType.getNumElements(), cast<Attribute>(rewriter.getZeroAttr(paddedType.getElementType())));
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
for (int64_t copy = 0; copy < packFactor; ++copy)
for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel)
for (int64_t kernelH = 0; kernelH < state.wHeight; ++kernelH)
for (int64_t kernelW = 0; kernelW < state.wWidth; ++kernelW)
for (int64_t inChannel = 0; inChannel < state.numChannelsIn; ++inChannel) {
const int64_t sourceFlatIndex =
(((outChannel * state.numChannelsIn) + inChannel) * state.wHeight + kernelH) * state.wWidth + kernelW;
const int64_t patchIndex =
((kernelH * state.wWidth) + kernelW) * state.numChannelsIn + inChannel;
const int64_t packedRow = copy * patchSize + patchIndex;
const int64_t packedColumn = copy * state.numChannelsOut + outChannel;
paddedValues[packedRow * paddedC + packedColumn] = sourceValues[sourceFlatIndex];
}
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(),
DenseElementsAttr::get(paddedType, paddedValues), 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 =
((kernelH * state.wWidth) + kernelW) * state.numChannelsIn + inChannel;
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,
PatternRewriter& rewriter,
Location loc) {
PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc);
ConvGeometry geo = buildConvGeometry(state, state.targetInfo());
const int64_t xbarDim = geo.xbarSize;
const int64_t numKSlices = ceilIntegerDivide(geo.k, xbarDim);
const int64_t paddedK = numKSlices * xbarDim;
const uint64_t maxLanesPerBatch =
std::max<uint64_t>(1,
static_cast<uint64_t>(state.targetInfo().matrixUnitsPerProcessor)
/ static_cast<uint64_t>(std::max<int64_t>(1, numKSlices * 4)));
const uint64_t rowChunkWidth = std::max<uint64_t>(
1,
std::min<uint64_t>({chooseStreamChunkPositions(geo, /*packFactor=*/1, state.targetInfo()),
maxLanesPerBatch,
static_cast<uint64_t>(state.outWidth)}));
const auto elementType = state.outType.getElementType();
auto wDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!wDenseAttr)
return failure();
Value paddedWeight = createPaddedInputKTiledWeightConstant(wDenseAttr, state, paddedK, xbarDim, rewriter);
Value paddedBias;
RankedTensorType paddedBiasType;
if (state.hasBias) {
Value biasMatrix = expandBiasIfNeeded(state.b, rewriter, loc);
auto biasMatrixType = cast<RankedTensorType>(biasMatrix.getType());
paddedBiasType = RankedTensorType::get({1, xbarDim}, elementType);
if (auto biasDenseAttr = getHostConstDenseElementsAttr(state.b))
paddedBias = createPaddedConstantMatrix(biasDenseAttr, biasMatrixType, paddedBiasType, rewriter);
else
paddedBias = materializeOrComputeUnary(
biasMatrix, paddedBiasType, rewriter, loc, [&](Value biasValue) {
return createPaddedConvMatrix(biasValue, biasMatrixType, paddedBiasType, rewriter, loc);
});
}
SmallVector<Value> chunkRows;
const int64_t totalPatches = state.batchSize * state.outHeight * state.outWidth;
chunkRows.reserve(
state.batchSize * state.outHeight * ceilIntegerDivide(state.outWidth, static_cast<int64_t>(rowChunkWidth)));
for (int64_t batchIndex = 0; batchIndex < state.batchSize; ++batchIndex) {
for (int64_t outHeightIndex = 0; outHeightIndex < state.outHeight; ++outHeightIndex) {
for (int64_t outWidthChunkStart = 0; outWidthChunkStart < state.outWidth;
outWidthChunkStart += static_cast<int64_t>(rowChunkWidth)) {
const int64_t chunkNumPatches =
std::min<int64_t>(static_cast<int64_t>(rowChunkWidth), state.outWidth - outWidthChunkStart);
auto chunkRowsType = RankedTensorType::get({chunkNumPatches, state.numChannelsOut}, elementType);
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto paddedChunkRowType = RankedTensorType::get({1, paddedK}, elementType);
auto patchType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.wWidth}, elementType);
auto collapsedPatchType = RankedTensorType::get({1, geo.k}, elementType);
auto weightTileType = RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType());
auto rowType = RankedTensorType::get({1, state.numChannelsOut}, elementType);
SmallVector<Value> inputsStorage {preparedInput.value};
if (state.hasBias)
inputsStorage.push_back(paddedBias);
ValueRange inputs(inputsStorage);
auto chunkCompute = spatial::SpatCompute::create(rewriter, loc, TypeRange {chunkRowsType}, ValueRange {paddedWeight}, inputs);
auto* block = new Block();
block->addArgument(paddedWeight.getType(), loc);
for (Value input : inputs)
block->addArgument(input.getType(), loc);
chunkCompute.getBody().push_back(block);
rewriter.setInsertionPointToStart(block);
auto buildChunk = [&]() -> LogicalResult {
Value weightArg = block->getArgument(0);
Value inputArg = block->getArgument(1);
Value biasArg = state.hasBias ? block->getArgument(2) : Value();
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value cBatchIndex = getOrCreateIndexConstant(rewriter, anchorOp, batchIndex);
Value cZero = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value cKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices);
Value cOne = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cXbar = getOrCreateIndexConstant(rewriter, anchorOp, xbarDim);
Value cInputHeightOffset =
getOrCreateIndexConstant(rewriter, anchorOp, outHeightIndex * state.strideHeight);
Value chunkRowsValue = tensor::EmptyOp::create(rewriter, loc, chunkRowsType.getShape(), elementType);
auto widthLoop = buildNormalizedScfFor(
rewriter,
loc,
cZero,
getOrCreateIndexConstant(rewriter, anchorOp, chunkNumPatches),
cOne,
ValueRange {chunkRowsValue},
[&](OpBuilder&, Location nestedLoc, Value widthIndex, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value laneWithChunkOffset = affineAddConst(rewriter, nestedLoc, widthIndex, outWidthChunkStart, anchorOp);
Value inputWidthOffset = createOrFoldAffineApply(rewriter,
nestedLoc,
getAffineDimExpr(0, rewriter.getContext()) * state.strideWidth,
ValueRange {laneWithChunkOffset},
anchorOp);
Value patch = createConvInputPatch(inputArg,
patchType,
cBatchIndex,
cZero,
cInputHeightOffset,
inputWidthOffset,
state.dilationHeight,
state.dilationWidth,
rewriter,
nestedLoc);
Value patchRow = tensor::CollapseShapeOp::create(rewriter,
nestedLoc,
collapsedPatchType,
patch,
SmallVector<ReassociationIndices> {
{0},
{1, 2, 3}
});
Value paddedPatchRow = createZeroPaddedTensor(
patchRow, paddedChunkRowType, {0, 0}, {0, paddedK - geo.k}, rewriter, nestedLoc);
auto zeroAttr = DenseElementsAttr::get(paddedRowType, rewriter.getZeroAttr(elementType));
Value zeroRow = getOrCreateConstant(rewriter, anchorOp, zeroAttr, paddedRowType);
auto kLoop = buildNormalizedScfFor(
rewriter,
nestedLoc,
cZero,
cKSlices,
cOne,
ValueRange {zeroRow},
[&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange reduceIterArgs, SmallVectorImpl<Value>& reduceYielded) {
Value acc = reduceIterArgs.front();
Value kOffset = arith::MulIOp::create(rewriter, reduceLoc, kSlice, cXbar);
SmallVector<OpFoldResult> aOffsets {rewriter.getIndexAttr(0), kOffset};
SmallVector<OpFoldResult> aSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)};
SmallVector<OpFoldResult> unitStrides = getUnitStrides(rewriter, 2);
Value aTile = tensor::ExtractSliceOp::create(
rewriter, reduceLoc, paddedRowType, paddedPatchRow, aOffsets, aSizes, unitStrides);
SmallVector<OpFoldResult> bOffsets {kOffset, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)};
Value bTile = extractStaticSliceOrIdentity(
rewriter, reduceLoc, weightArg, weightTileType, bOffsets, bSizes, unitStrides);
Value piece = spatial::SpatVMMOp::create(rewriter, reduceLoc, paddedRowType, bTile, aTile).getResult();
reduceYielded.push_back(
spatial::SpatVAddOp::create(rewriter, reduceLoc, paddedRowType, acc, piece).getResult());
return success();
});
if (failed(kLoop))
return failure();
Value reduced = kLoop->results.front();
if (state.hasBias)
reduced = spatial::SpatVAddOp::create(rewriter, nestedLoc, paddedRowType, reduced, biasArg).getResult();
Value row = reduced;
if (state.numChannelsOut != xbarDim) {
SmallVector<OpFoldResult> rowOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> rowSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)};
row = tensor::ExtractSliceOp::create(
rewriter, nestedLoc, rowType, reduced, rowOffsets, rowSizes, getUnitStrides(rewriter, 2));
}
SmallVector<OpFoldResult> outputOffsets {widthIndex, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> outputSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(state.numChannelsOut)};
Value updatedRows = tensor::InsertSliceOp::create(
rewriter, nestedLoc, row, iterArgs.front(), outputOffsets, outputSizes, getUnitStrides(rewriter, 2));
yielded.push_back(updatedRows);
return success();
});
if (failed(widthLoop))
return failure();
spatial::SpatYieldOp::create(rewriter, loc, widthLoop->results.front());
return success();
};
if (failed(buildChunk())) {
rewriter.setInsertionPointAfter(chunkCompute);
rewriter.eraseOp(chunkCompute);
return failure();
}
rewriter.setInsertionPointAfter(chunkCompute);
chunkRows.push_back(chunkCompute.getResult(0));
}
}
}
auto nhwcType = RankedTensorType::get({state.batchSize, state.outHeight, state.outWidth, state.numChannelsOut},
elementType);
return createCollectedConvOutput(
chunkRows, state.outType, cast<RankedTensorType>(chunkRows.front().getType()), nhwcType, state.outType, totalPatches,
state.numChannelsOut, /*packFactor=*/1, rewriter, loc);
}
static Value buildPackedWeights(DenseElementsAttr wDenseAttr,
Value wTrans,
const ConvLoweringState& state,
const ConvGemmPlan& plan,
PatternRewriter& rewriter,
Location loc) {
if (plan.effectiveMaxParallelPixels == 1)
return wTrans;
auto packedWeightType = RankedTensorType::get(
{plan.effectiveMaxParallelPixels * plan.patchSize, plan.effectiveMaxParallelPixels * state.numChannelsOut},
state.wType.getElementType());
SmallVector<Attribute> sourceValues(wDenseAttr.getValues<Attribute>());
SmallVector<Attribute> packedValues(packedWeightType.getNumElements(),
cast<Attribute>(rewriter.getZeroAttr(state.wType.getElementType())));
for (int64_t copyId = 0; copyId < plan.effectiveMaxParallelPixels; ++copyId) {
for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel) {
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 targetRow = copyId * plan.patchSize + patchIndex;
const int64_t targetCol = copyId * state.numChannelsOut + outChannel;
packedValues[targetRow * packedWeightType.getDimSize(1) + targetCol] = sourceValues[sourceFlatIndex];
}
}
}
}
}
auto packedAttr = DenseElementsAttr::get(packedWeightType, packedValues);
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), packedAttr, packedWeightType);
}
static Value buildPackedBias(Value gemmBias,
Value biasMatrix,
DenseElementsAttr biasDenseAttr,
const ConvLoweringState& state,
const ConvGemmPlan& plan,
PatternRewriter& rewriter,
Location loc) {
if (!state.hasBias)
return gemmBias;
if (plan.effectiveMaxParallelPixels == 1)
return biasMatrix;
SmallVector<Attribute> sourceValues(biasDenseAttr.getValues<Attribute>());
SmallVector<Attribute> packedValues;
packedValues.reserve(plan.effectiveMaxParallelPixels * state.numChannelsOut);
for (int64_t copyId = 0; copyId < plan.effectiveMaxParallelPixels; ++copyId)
packedValues.append(sourceValues.begin(), sourceValues.end());
auto packedBiasType =
RankedTensorType::get({1, plan.effectiveMaxParallelPixels * state.numChannelsOut}, state.outType.getElementType());
auto packedBiasAttr = DenseElementsAttr::get(packedBiasType, packedValues);
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), packedBiasAttr, packedBiasType);
}
static ConvGemmPlan
buildConvGemmPlan(const ConvLoweringState& state,
bool canPackWeightsAsConstants,
bool canPackBiasAsConstants,
int64_t chunkStart,
int64_t chunkNumPatches,
std::optional<int64_t> forcedPackFactor) {
ConvGemmPlan plan;
plan.patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
plan.numPatchesPerBatch = state.outHeight * state.outWidth;
plan.globalNumPatches = state.batchSize * plan.numPatchesPerBatch;
plan.chunkStart = chunkStart;
plan.chunkNumPatches = chunkNumPatches;
const int64_t wMaxDim = std::max(plan.patchSize, state.numChannelsOut);
plan.maxParallelPixels = forcedPackFactor
? *forcedPackFactor
: std::max<int64_t>(1, static_cast<int64_t>(state.targetInfo().matrixShape.rows) / wMaxDim);
plan.effectiveMaxParallelPixels =
(canPackWeightsAsConstants && canPackBiasAsConstants) ? plan.maxParallelPixels : 1;
plan.packedNumRows = ceilIntegerDivide(plan.chunkNumPatches, plan.effectiveMaxParallelPixels);
auto elemType = state.xType.getElementType();
auto outElemType = state.outType.getElementType();
plan.gemmInputRowsType =
RankedTensorType::get({plan.packedNumRows, plan.effectiveMaxParallelPixels * plan.patchSize}, elemType);
plan.wFlatType = RankedTensorType::get({state.numChannelsOut, plan.patchSize}, state.wType.getElementType());
plan.wTransType = RankedTensorType::get({plan.patchSize, state.numChannelsOut}, state.wType.getElementType());
plan.gemmOutType = RankedTensorType::get({plan.chunkNumPatches, state.numChannelsOut}, outElemType);
plan.gemmOutputRowsType =
RankedTensorType::get({plan.packedNumRows, plan.effectiveMaxParallelPixels * state.numChannelsOut}, outElemType);
plan.nhwcType =
RankedTensorType::get({state.batchSize, state.outHeight, state.outWidth, state.numChannelsOut}, outElemType);
return plan;
}
static Value createIm2colRows(const ConvLoweringState& state,
const PreparedConvInput& preparedInput,
const ConvGemmPlan& plan,
PatternRewriter& rewriter,
Location loc) {
if (plan.gemmInputRowsType.getDimSize(1)
> static_cast<int64_t>(state.targetInfo().matrixShape.rows)) {
assert(plan.effectiveMaxParallelPixels == 1 && "multi-crossbar im2col rows cannot pack pixels");
auto compute = createSpatCompute<1>(
rewriter, loc, TypeRange {plan.gemmInputRowsType}, {}, preparedInput.value, [&](Value input) {
auto elemType = preparedInput.type.getElementType();
Value empty = tensor::EmptyOp::create(rewriter, loc, plan.gemmInputRowsType.getShape(), elemType);
Operation *anchor = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchor, 1);
Value upper = getOrCreateIndexConstant(rewriter, anchor, plan.chunkNumPatches);
auto patchType = RankedTensorType::get(
{1, state.numChannelsIn, state.wHeight, state.wWidth}, elemType);
auto rowType = RankedTensorType::get({plan.patchSize}, elemType);
auto loop = buildNormalizedScfFor(
rewriter, loc, c0, upper, c1, ValueRange {empty},
[&](OpBuilder &, Location nestedLoc, Value patchIndex, ValueRange iterArgs,
SmallVectorImpl<Value> &yielded) {
Value batchIndex = affineAddFloorDivConst(
rewriter, nestedLoc, patchIndex, plan.chunkStart, plan.numPatchesPerBatch, anchor);
Value batchPatchIndex = affineAddModConst(
rewriter, nestedLoc, patchIndex, plan.chunkStart, plan.numPatchesPerBatch, anchor);
Value outHeight = affineFloorDivConst(
rewriter, nestedLoc, batchPatchIndex, state.outWidth, anchor);
Value outWidth = affineModConst(
rewriter, nestedLoc, batchPatchIndex, state.outWidth, anchor);
Value patch = createConvInputPatch(
input, patchType, batchIndex, c0,
affineMulConst(rewriter, nestedLoc, outHeight, state.strideHeight, anchor),
affineMulConst(rewriter, nestedLoc, outWidth, state.strideWidth, anchor),
state.dilationHeight, state.dilationWidth, rewriter, nestedLoc);
Value row = tensor::CollapseShapeOp::create(
rewriter, nestedLoc, rowType, patch,
SmallVector<ReassociationIndices> {{0, 1, 2, 3}});
Value next = tensor::InsertSliceOp::create(
rewriter, nestedLoc, row, iterArgs.front(),
SmallVector<OpFoldResult> {patchIndex, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(plan.patchSize)},
getUnitStrides(rewriter, 2));
yielded.push_back(next);
return success();
});
if (failed(loop))
return failure();
spatial::SpatYieldOp::create(rewriter, loc, loop->results.front());
return success();
});
assert(succeeded(compute) && "Conv im2col compute construction must succeed");
return compute->getResult(0);
}
auto elemType = preparedInput.type.getElementType();
auto packedRowType = RankedTensorType::get(
{plan.effectiveMaxParallelPixels * plan.patchSize}, elemType, plan.gemmInputRowsType.getEncoding());
auto patchType = RankedTensorType::get({1, state.numChannelsIn, state.wHeight, state.wWidth}, elemType);
auto patchRowType = RankedTensorType::get({plan.patchSize}, elemType);
bool hasPartialLane = plan.chunkNumPatches % plan.effectiveMaxParallelPixels != 0;
SmallVector<Value> im2colInputs {preparedInput.value};
auto im2colComputeOp = createSpatComputeBatch(
rewriter,
loc,
TypeRange {plan.gemmInputRowsType},
plan.packedNumRows,
{},
im2colInputs,
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cPack = getOrCreateIndexConstant(rewriter, anchorOp, plan.effectiveMaxParallelPixels);
Value laneStart = affineMulConst(rewriter, loc, args.lane, plan.effectiveMaxParallelPixels, anchorOp);
Value lanePatches = cPack;
if (hasPartialLane) {
Value cNumPatches = getOrCreateIndexConstant(rewriter, anchorOp, plan.chunkNumPatches);
Value remaining = arith::SubIOp::create(rewriter, loc, cNumPatches, laneStart);
Value isPartial = arith::CmpIOp::create(
rewriter, loc, arith::CmpIPredicate::ult, remaining, cPack);
lanePatches = arith::SelectOp::create(rewriter, loc, isPartial, remaining, cPack);
}
Value rowInit = tensor::EmptyOp::create(rewriter, loc, packedRowType.getShape(), elemType);
if (hasPartialLane) {
auto zeroAttr = cast<TypedAttr>(rewriter.getZeroAttr(elemType));
rowInit = linalg::MapOp::create(
rewriter, loc, ValueRange {}, rowInit,
[&](OpBuilder& builder, Location nestedLoc, ValueRange) {
Value zero = arith::ConstantOp::create(builder, nestedLoc, zeroAttr);
linalg::YieldOp::create(builder, nestedLoc, zero);
}).getResult().front();
}
auto rowLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
lanePatches,
c1,
ValueRange {rowInit},
[&](OpBuilder&, Location nestedLoc, Value copyIndex, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value patchIndex = arith::AddIOp::create(rewriter, nestedLoc, laneStart, copyIndex);
Value batchIndex = state.batchSize == 1
? c0
: affineAddFloorDivConst(
rewriter, nestedLoc, patchIndex, plan.chunkStart,
plan.numPatchesPerBatch, anchorOp);
Value batchPatchIndex =
affineAddModConst(rewriter, nestedLoc, patchIndex, plan.chunkStart, plan.numPatchesPerBatch, anchorOp);
Value outHeightIndex = affineFloorDivConst(rewriter, nestedLoc, batchPatchIndex, state.outWidth, anchorOp);
Value outWidthIndex = affineModConst(rewriter, nestedLoc, batchPatchIndex, state.outWidth, anchorOp);
Value inputHeightOffset = affineMulConst(rewriter, nestedLoc, outHeightIndex, state.strideHeight, anchorOp);
Value inputWidthOffset = affineMulConst(rewriter, nestedLoc, outWidthIndex, state.strideWidth, anchorOp);
Value patch = createConvInputPatch(args.inputs.front(),
patchType,
batchIndex,
c0,
inputHeightOffset,
inputWidthOffset,
state.dilationHeight,
state.dilationWidth,
rewriter,
nestedLoc);
Value patchRow = tensor::CollapseShapeOp::create(rewriter,
nestedLoc,
patchRowType,
patch,
SmallVector<ReassociationIndices> {
{0, 1, 2, 3}
});
Value rowOffset = affineMulConst(rewriter, nestedLoc, copyIndex, plan.patchSize, anchorOp);
Value next = tensor::InsertSliceOp::create(rewriter,
nestedLoc,
patchRow,
iterArgs.front(),
SmallVector<OpFoldResult> {rowOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(plan.patchSize)},
getUnitStrides(rewriter, 1));
yielded.push_back(next);
return success();
});
if (failed(rowLoop))
return failure();
Value row = rowLoop->results.front();
publishGraphBatchPhysicalFragment(rewriter, loc, row, args.outputs.front(), args.lane);
return success();
});
assert(succeeded(im2colComputeOp) && "Conv im2col compute construction must succeed");
return im2colComputeOp->getResult(0);
}
static Value maybeUnpackChunkRows(Value gemmRows,
const ConvGemmPlan& plan,
PatternRewriter& rewriter,
Location loc) {
if (plan.effectiveMaxParallelPixels == 1)
return gemmRows;
auto unpackedType = RankedTensorType::get(
{plan.chunkNumPatches, plan.gemmOutType.getDimSize(1)}, plan.gemmOutType.getElementType(), plan.gemmOutType.getEncoding());
auto unpackCompute = createSpatCompute<1>(rewriter, loc, TypeRange {unpackedType}, {}, gemmRows, [&](Value rowsArg) {
Value unpacked = unpackRowsFromParallelGemm(rowsArg,
cast<RankedTensorType>(rowsArg.getType()),
plan.chunkNumPatches,
plan.gemmOutType.getDimSize(1),
plan.effectiveMaxParallelPixels,
rewriter,
loc);
spatial::SpatYieldOp::create(rewriter, loc, unpacked);
});
return unpackCompute.getResult(0);
}
static FailureOr<Value> createStreamedConvRows(const ConvLoweringState& state,
const PreparedConvInput& preparedInput,
Value weightMatrix,
Value biasMatrix,
DenseElementsAttr wDenseAttr,
DenseElementsAttr biasDenseAttr,
int64_t forcedPackFactor,
PatternRewriter& rewriter,
Location loc) {
const int64_t totalPatches = state.batchSize * state.outHeight * state.outWidth;
ConvGemmPlan plan = buildConvGemmPlan(state, static_cast<bool>(wDenseAttr),
!state.hasBias || static_cast<bool>(biasDenseAttr), 0, totalPatches, forcedPackFactor);
Value inputRows = createIm2colRows(state, preparedInput, plan, rewriter, loc);
Value packedWeights = buildPackedWeights(wDenseAttr, weightMatrix, state, plan, rewriter, loc);
Value gemmBias = state.hasBias ? state.b : createZeroGemmBias(plan.gemmOutputRowsType, rewriter);
Value packedBias = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, plan, rewriter, loc);
FailureOr<Value> gemmRows = lowerGemmToSpatial(
state.diagnosticAnchor, inputRows, packedWeights, packedBias,
plan.gemmOutputRowsType, /*transA=*/false, /*transB=*/!wDenseAttr,
/*alpha=*/1.0f, /*beta=*/1.0f, state.targetInfo(), rewriter, loc);
if (failed(gemmRows))
return failure();
return maybeUnpackChunkRows(*gemmRows, plan, rewriter, loc);
}
static FailureOr<Value> rewritePackedIm2ColConv(const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
auto wDenseAttr = getHostConstDenseElementsAttr(state.w);
PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc);
Value biasMatrix;
DenseElementsAttr biasDenseAttr;
if (state.hasBias) {
biasDenseAttr = getHostConstDenseElementsAttr(state.b);
biasMatrix = expandBiasIfNeeded(state.b, rewriter, loc);
}
ConvGemmPlan plan =
buildConvGemmPlan(state, static_cast<bool>(wDenseAttr), !state.hasBias || static_cast<bool>(biasDenseAttr), 0,
state.batchSize * state.outHeight * state.outWidth);
// Static weights use the crossbar [patchSize, Cout] layout. Runtime weights
// stay in ONNX's contiguous [Cout, patchSize] layout and Gemm consumes them
// through transB without materializing a transpose.
Value weightMatrix = createWeightMatrix(state.w, plan, static_cast<bool>(wDenseAttr), rewriter, loc);
Value gemmInputRows = createIm2colRows(state, preparedInput, plan, rewriter, loc);
Value gemmB = buildPackedWeights(wDenseAttr, weightMatrix, state, plan, rewriter, loc);
Value gemmBias = createZeroGemmBias(plan.gemmOutputRowsType, rewriter);
if (state.hasBias)
gemmBias = state.b;
Value gemmC = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, plan, rewriter, loc);
FailureOr<Value> gemmRows = lowerGemmToSpatial(
state.diagnosticAnchor, gemmInputRows, gemmB, gemmC,
plan.gemmOutputRowsType, /*transA=*/false, /*transB=*/!wDenseAttr,
/*alpha=*/1.0f, /*beta=*/1.0f, state.targetInfo(), rewriter, loc);
if (failed(gemmRows))
return failure();
return createCollectedConvOutput(ValueRange {*gemmRows},
state.outType,
plan.gemmOutType,
plan.nhwcType,
state.outType,
plan.chunkNumPatches,
state.numChannelsOut,
plan.effectiveMaxParallelPixels,
rewriter,
loc);
}
static FailureOr<Value> rewriteStreamedConv(const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc,
int64_t forcedPackFactor) {
auto wDenseAttr = getHostConstDenseElementsAttr(state.w);
PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc);
Value biasMatrix;
DenseElementsAttr biasDenseAttr;
if (state.hasBias) {
biasDenseAttr = getHostConstDenseElementsAttr(state.b);
biasMatrix = expandBiasIfNeeded(state.b, rewriter, loc);
}
ConvGemmPlan seedPlan = buildConvGemmPlan(
state, static_cast<bool>(wDenseAttr), !state.hasBias || static_cast<bool>(biasDenseAttr), 0, 1, forcedPackFactor);
Value weightMatrix = createWeightMatrix(state.w, seedPlan, static_cast<bool>(wDenseAttr), rewriter, loc);
FailureOr<Value> collectedRows = createStreamedConvRows(state,
preparedInput,
weightMatrix,
biasMatrix,
wDenseAttr,
biasDenseAttr,
forcedPackFactor,
rewriter,
loc);
if (failed(collectedRows))
return failure();
auto gemmOutType = cast<RankedTensorType>(collectedRows->getType());
auto nhwcType = RankedTensorType::get({state.batchSize, state.outHeight, state.outWidth, state.numChannelsOut},
state.outType.getElementType());
return createCollectedConvOutput(
ValueRange {*collectedRows}, state.outType, gemmOutType, nhwcType, state.outType, gemmOutType.getDimSize(0),
state.numChannelsOut, /*packFactor=*/1, rewriter, loc);
}
} // namespace standard
static Value createZeroGemmBias(RankedTensorType gemmResultType, PatternRewriter& rewriter) {
auto zeroAttr = DenseElementsAttr::get(gemmResultType, rewriter.getZeroAttr(gemmResultType.getElementType()));
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), zeroAttr, gemmResultType);
}
static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) {
return ceilIntegerDivide(geometry.k, geometry.xbarSize)
* ceilIntegerDivide(geometry.c, geometry.xbarSize)
<= geometry.matrixUnitsPerProcessor;
}
static bool rowStripOutputChannelTileFitsOneCore(const ConvGeometry& geometry) {
return ceilIntegerDivide(geometry.k, geometry.xbarSize)
<= geometry.matrixUnitsPerProcessor;
}
static int64_t chooseRowStripPixelPackFactor(const ConvLoweringState& state, int64_t xbarDim) {
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
const int64_t baseWeightGroups = ceilIntegerDivide(patchSize, xbarDim)
* ceilIntegerDivide(state.numChannelsOut, xbarDim);
int64_t factor = std::min(state.outWidth, xbarDim / state.numChannelsOut);
while (factor > 1
&& (state.outWidth % factor != 0
|| ceilIntegerDivide(factor * patchSize, xbarDim)
* ceilIntegerDivide(factor * state.numChannelsOut, xbarDim)
> baseWeightGroups))
--factor;
return std::max<int64_t>(factor, 1);
}
static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state, StringRef& failureReason) {
if (state.batchSize != 1) {
failureReason = "batch_not_one";
return false;
}
if (state.group != 1) {
failureReason = "grouped_conv";
return false;
}
if (!state.xType.hasStaticShape() || !state.wType.hasStaticShape() || !state.outType.hasStaticShape()) {
failureReason = "dynamic_shape";
return false;
}
if (!isa<FloatType>(state.xType.getElementType())) {
failureReason = "non_float_input";
return false;
}
if (state.dilationHeight != 1 || state.dilationWidth != 1) {
failureReason = "dilation_not_one";
return false;
}
if (!getHostConstDenseElementsAttr(state.w)) {
failureReason = "non_constant_weight";
return false;
}
if (!rowStripOutputChannelTileFitsOneCore(buildConvGeometry(state, state.targetInfo()))) {
failureReason = "output_channel_tile_does_not_fit_one_core";
return false;
}
if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)) {
failureReason = "unsupported_bias";
return false;
}
return true;
}
static Value createZeroTensorConstant(RankedTensorType type, PatternRewriter& rewriter) {
auto zeroAttr = DenseElementsAttr::get(type, rewriter.getZeroAttr(type.getElementType()));
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), zeroAttr, type);
}
static FailureOr<Value> createBiasRowConstant(const ConvLoweringState& state,
int64_t packFactor,
PatternRewriter& rewriter) {
DenseElementsAttr denseAttr;
if (!isSupportedBiasAddValue(state.b, state.outType, &denseAttr))
return failure();
FailureOr<SmallVector<Attribute>> channelValues = getBiasChannelValues(denseAttr, state.outType);
if (failed(channelValues))
return failure();
SmallVector<Attribute> packedValues;
packedValues.reserve(packFactor * state.numChannelsOut);
for (int64_t copy = 0; copy < packFactor; ++copy)
packedValues.append(channelValues->begin(), channelValues->end());
auto biasType = RankedTensorType::get({1, packFactor * state.numChannelsOut}, state.outType.getElementType());
return getOrCreateConstant(rewriter,
rewriter.getInsertionBlock()->getParentOp(),
DenseElementsAttr::get(biasType, packedValues),
biasType);
}
static FailureOr<Value> createPaddedBiasTileConstant(const ConvLoweringState& state,
int64_t tileChannels,
PatternRewriter& rewriter) {
DenseElementsAttr denseAttr;
if (!isSupportedBiasAddValue(state.b, state.outType, &denseAttr))
return failure();
FailureOr<SmallVector<Attribute>> channelValues = getBiasChannelValues(denseAttr, state.outType);
if (failed(channelValues))
return failure();
const int64_t tileCount = ceilIntegerDivide(state.numChannelsOut, tileChannels);
auto tileType = RankedTensorType::get({tileCount, 1, tileChannels}, state.outType.getElementType());
SmallVector<Attribute> values(
tileType.getNumElements(), cast<Attribute>(rewriter.getZeroAttr(tileType.getElementType())));
for (int64_t channel = 0; channel < state.numChannelsOut; ++channel)
values[channel] = (*channelValues)[channel];
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(),
DenseElementsAttr::get(tileType, values), tileType);
}
static Value createHorizontallyPaddedRowStripFragment(Value fragment,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
auto paddedType = RankedTensorType::get(
{1, 1, state.xWidth + state.padWidthBegin + state.padWidthEnd, state.numChannelsIn},
state.xType.getElementType(),
state.xType.getEncoding());
return createZeroPaddedTensor(fragment,
paddedType,
{0, 0, state.padWidthBegin, 0},
{0, 0, state.padWidthEnd, 0},
rewriter,
loc);
}
static Value createRowStripWindowSourceRowTable(const ConvLoweringState& state, PatternRewriter& rewriter) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto tableType = RankedTensorType::get({state.outHeight * state.wHeight}, rewriter.getIndexType());
SmallVector<Attribute> values;
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 * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin;
sourceRow = std::clamp(sourceRow, int64_t {0}, state.xHeight - 1);
values.push_back(rewriter.getIndexAttr(sourceRow));
}
}
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value createRowStripWindowSourceSlotTable(const ConvLoweringState& state,
int64_t tilesPerRow,
PatternRewriter& rewriter) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto tableType =
RankedTensorType::get({state.outHeight * state.wHeight * tilesPerRow}, rewriter.getIndexType());
SmallVector<Attribute> values;
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 * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin;
sourceRow = std::clamp(sourceRow, int64_t {0}, state.xHeight - 1);
for (int64_t tile = 0; tile < tilesPerRow; ++tile)
values.push_back(rewriter.getIndexAttr(sourceRow * tilesPerRow + tile));
}
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value createRowStripWindowTableIndex(Value outputHeight,
Value kernelRow,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
MLIRContext* ctx = rewriter.getContext();
AffineExpr outputRowExpr = getAffineDimExpr(0, ctx);
AffineExpr kernelRowExpr = getAffineDimExpr(1, ctx);
return createOrFoldAffineApply(
rewriter, loc, outputRowExpr * state.wHeight + kernelRowExpr, ValueRange {outputHeight, kernelRow}, anchorOp);
}
static FailureOr<Value> extractProjectedRowStripWindowRow(Value rowStripStorage,
Value sourceSlotTable,
const ConvLoweringState& state,
Value outputHeight,
Value kernelRow,
PatternRewriter& rewriter,
Location loc) {
FailureOr<RowStripPhysicalValue> physical = describeRowStripPhysicalValue(rowStripStorage, state.xType);
if (failed(physical))
return failure();
Value tableIndex = createRowStripWindowTableIndex(outputHeight, kernelRow, state, rewriter, loc);
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value tileTableIndex =
affineMulConst(rewriter, loc, tableIndex, physical->tilesPerRow, anchorOp);
if (physical->tilesPerRow == 1) {
Value sourceSlot =
tensor::ExtractOp::create(rewriter, loc, sourceSlotTable, ValueRange {tileTableIndex}).getResult();
return extractGraphBatchPhysicalFragment(
rewriter, loc, rowStripStorage, sourceSlot, physical->fragmentType);
}
auto fullFragmentType = getRowStripFragmentType(state.xType);
Value fullFragment = tensor::EmptyOp::create(
rewriter, loc, fullFragmentType.getShape(), fullFragmentType.getElementType());
const int64_t tileChannels = physical->fragmentType.getDimSize(3);
for (int64_t tile = 0; tile < physical->tilesPerRow; ++tile) {
Value slotTableIndex = affineAddConst(rewriter, loc, tileTableIndex, tile, anchorOp);
Value tileSlot =
tensor::ExtractOp::create(rewriter, loc, sourceSlotTable, ValueRange {slotTableIndex}).getResult();
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
rewriter, loc, rowStripStorage, tileSlot, physical->fragmentType);
if (failed(fragment))
return failure();
const int64_t channelOffset = tile * tileChannels;
const int64_t validChannels = std::min(tileChannels, state.numChannelsIn - channelOffset);
auto validType = RankedTensorType::get(
{1, 1, state.xWidth, validChannels}, state.xType.getElementType(), state.xType.getEncoding());
Value validFragment = *fragment;
if (validChannels != tileChannels)
validFragment = tensor::ExtractSliceOp::create(
rewriter,
loc,
validType,
*fragment,
SmallVector<OpFoldResult>(4, rewriter.getIndexAttr(0)),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.xWidth),
rewriter.getIndexAttr(validChannels)},
getUnitStrides(rewriter, 4));
fullFragment = tensor::InsertSliceOp::create(
rewriter,
loc,
validFragment,
fullFragment,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
rewriter.getIndexAttr(channelOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.xWidth),
rewriter.getIndexAttr(validChannels)},
getUnitStrides(rewriter, 4));
}
return fullFragment;
}
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 nchwType = RankedTensorType::get(
{1, state.numChannelsIn, 1, state.xWidth}, state.xType.getElementType(), state.xType.getEncoding());
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)};
Value nchw = tensor::ExtractSliceOp::create(
rewriter, loc, nchwType, denseInput, offsets, sizes, getUnitStrides(rewriter, 4));
return createLinalgTranspose(nchw, fragmentType, {0, 2, 3, 1}, rewriter, loc);
}
static Value createRowStripWindowMaskTable(const ConvLoweringState& state, PatternRewriter& rewriter) {
auto elementType = cast<FloatType>(state.xType.getElementType());
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto tableType = RankedTensorType::get(
{2, 1, state.xWidth, state.numChannelsIn}, elementType, state.xType.getEncoding());
SmallVector<Attribute> values(tableType.getNumElements(), rewriter.getZeroAttr(elementType));
std::fill(values.begin() + tableType.getNumElements() / 2, values.end(), rewriter.getFloatAttr(elementType, 1.0));
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value createRowStripWindowMaskIndexTable(const ConvLoweringState& state, PatternRewriter& rewriter) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto tableType = RankedTensorType::get({state.outHeight * state.wHeight}, rewriter.getIndexType());
SmallVector<Attribute> values;
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 * state.strideHeight + kernelRow * state.dilationHeight - state.padHeightBegin;
values.push_back(rewriter.getIndexAttr(sourceRow >= 0 && sourceRow < state.xHeight));
}
return getOrCreateConstant(rewriter, anchorOp, DenseElementsAttr::get(tableType, values), tableType);
}
static Value extractProjectedRowStripWindowMask(Value maskTable,
Value maskIndexTable,
const ConvLoweringState& state,
Value outputHeight,
Value kernelRow,
PatternRewriter& rewriter,
Location loc) {
Value tableIndex = createRowStripWindowTableIndex(outputHeight, kernelRow, state, rewriter, loc);
Value maskIndex = tensor::ExtractOp::create(rewriter, loc, maskIndexTable, ValueRange {tableIndex}).getResult();
auto fragmentType = getRowStripFragmentType(state.xType);
return tensor::ExtractSliceOp::create(
rewriter,
loc,
fragmentType,
maskTable,
SmallVector<OpFoldResult> {
maskIndex, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.xWidth), rewriter.getIndexAttr(state.numChannelsIn)},
getUnitStrides(rewriter, 4));
}
static FailureOr<Value> createConvInputWindow(Value input,
const ConvLoweringState& state,
Value outputHeight,
PatternRewriter& rewriter,
Location loc) {
auto fragmentType = getRowStripFragmentType(state.xType);
auto inputType = dyn_cast<RankedTensorType>(input.getType());
const bool denseInput = inputType == state.xType;
FailureOr<RowStripPhysicalValue> physicalInput = describeRowStripPhysicalValue(input, state.xType);
if (!denseInput && failed(physicalInput))
return failure();
if (!denseInput && physicalInput->tilesPerRow == 1 && state.wHeight == 1 && state.wWidth == 1
&& state.strideHeight == 1 && state.strideWidth == 1 && state.padHeightBegin == 0
&& state.padHeightEnd == 0 && state.padWidthBegin == 0 && state.padWidthEnd == 0)
return extractGraphBatchPhysicalFragment(
rewriter, loc, input, outputHeight, physicalInput->fragmentType);
auto paddedWindowType = RankedTensorType::get(
{1, state.wHeight, state.xWidth + state.padWidthBegin + state.padWidthEnd, state.numChannelsIn},
state.xType.getElementType(),
state.xType.getEncoding());
Value sourceIndexTable =
denseInput ? createRowStripWindowSourceRowTable(state, rewriter)
: createRowStripWindowSourceSlotTable(state, physicalInput->tilesPerRow, rewriter);
Value maskTable = createRowStripWindowMaskTable(state, rewriter);
Value maskIndexTable = createRowStripWindowMaskIndexTable(state, rewriter);
Value initWindow = createZeroTensorConstant(paddedWindowType, rewriter);
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cKernelRows = getOrCreateIndexConstant(rewriter, anchorOp, state.wHeight);
auto loop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cKernelRows,
c1,
ValueRange {initWindow},
[&](OpBuilder&, Location rowLoc, Value kernelRow, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) -> LogicalResult {
FailureOr<Value> sourceRow =
denseInput
? FailureOr<Value>(
extractDenseConvWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, rowLoc))
: extractProjectedRowStripWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, rowLoc);
if (failed(sourceRow))
return failure();
Value semanticRow = *sourceRow;
if (state.padHeightBegin != 0 || state.padHeightEnd != 0) {
Value mask = extractProjectedRowStripWindowMask(
maskTable, maskIndexTable, state, outputHeight, kernelRow, rewriter, rowLoc);
semanticRow = spatial::SpatVMulOp::create(rewriter, rowLoc, fragmentType, semanticRow, mask).getResult();
}
Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, rowLoc);
yielded.push_back(tensor::InsertSliceOp::create(
rewriter,
rowLoc,
paddedRow,
iterArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), kernelRow,
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.xWidth + state.padWidthBegin
+ state.padWidthEnd),
rewriter.getIndexAttr(state.numChannelsIn)},
getUnitStrides(rewriter, 4)));
return success();
});
return failed(loop) ? FailureOr<Value>(failure())
: FailureOr<Value>(loop->results.front());
}
static FailureOr<Value> createPixelMajorConvPatchRow(Value paddedWindow,
const ConvLoweringState& state,
Value outputWidth,
PatternRewriter& rewriter,
Location loc) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
auto patchType = RankedTensorType::get({1, state.wHeight, state.wWidth, state.numChannelsIn},
state.xType.getElementType(),
state.xType.getEncoding());
auto rowType = RankedTensorType::get({1, patchSize}, state.xType.getElementType(), state.xType.getEncoding());
Value inputWidthOffset = affineMulConst(rewriter, loc, outputWidth, state.strideWidth, anchorOp);
SmallVector<OpFoldResult> offsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), inputWidthOffset, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.wHeight),
rewriter.getIndexAttr(state.wWidth),
rewriter.getIndexAttr(state.numChannelsIn)};
Value patch;
if (state.dilationWidth == 1)
patch = tensor::ExtractSliceOp::create(
rewriter, loc, patchType, paddedWindow, offsets, sizes, getUnitStrides(rewriter, 4));
else {
auto columnType = RankedTensorType::get({1, state.wHeight, 1, state.numChannelsIn},
state.xType.getElementType(), state.xType.getEncoding());
patch = tensor::EmptyOp::create(rewriter, loc, patchType.getShape(), patchType.getElementType());
for (int64_t kernelColumn = 0; kernelColumn < state.wWidth; ++kernelColumn) {
Value sourceWidth =
affineAddConst(rewriter, loc, inputWidthOffset, kernelColumn * state.dilationWidth, anchorOp);
SmallVector<OpFoldResult> columnSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.wHeight),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.numChannelsIn)};
Value column = tensor::ExtractSliceOp::create(
rewriter, loc, columnType, paddedWindow,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceWidth,
rewriter.getIndexAttr(0)},
columnSizes, getUnitStrides(rewriter, 4));
patch = tensor::InsertSliceOp::create(
rewriter, loc, column, patch,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelColumn), rewriter.getIndexAttr(0)},
columnSizes, getUnitStrides(rewriter, 4));
}
}
return tensor::CollapseShapeOp::create(
rewriter, loc, rowType, patch, SmallVector<ReassociationIndices> {{0}, {1, 2, 3}})
.getResult();
}
static FailureOr<Value> createPackedPixelMajorConvPatchRow(Value paddedWindow,
const ConvLoweringState& state,
Value outputGroup,
int64_t packFactor,
PatternRewriter& rewriter,
Location loc) {
if (packFactor == 1)
return createPixelMajorConvPatchRow(paddedWindow, state, outputGroup, rewriter, loc);
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth;
auto packedType = RankedTensorType::get(
{1, packFactor * patchSize}, state.xType.getElementType(), state.xType.getEncoding());
Value packed = tensor::EmptyOp::create(rewriter, loc, packedType.getShape(), packedType.getElementType());
Value outputStart = affineMulConst(rewriter, loc, outputGroup, packFactor, anchorOp);
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cPackFactor = getOrCreateIndexConstant(rewriter, anchorOp, packFactor);
auto loop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cPackFactor,
c1,
ValueRange {packed},
[&](OpBuilder&, Location copyLoc, Value copy, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) -> LogicalResult {
Value outputWidth = createOrFoldAffineApply(
rewriter, copyLoc, rewriter.getAffineDimExpr(0) + rewriter.getAffineDimExpr(1),
ValueRange {outputStart, copy}, anchorOp);
FailureOr<Value> patch =
createPixelMajorConvPatchRow(paddedWindow, state, outputWidth, rewriter, copyLoc);
if (failed(patch))
return failure();
Value packedOffset = affineMulConst(rewriter, copyLoc, copy, patchSize, anchorOp);
yielded.push_back(tensor::InsertSliceOp::create(
rewriter,
copyLoc,
*patch,
iterArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), packedOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(patchSize)},
getUnitStrides(rewriter, 2)));
return success();
});
if (failed(loop))
return failure();
return loop->results.front();
}
static FailureOr<SmallVector<Value>> createConvInputTiles(Value paddedWindow,
const ConvLoweringState& state,
Value outputWidth,
int64_t packFactor,
Value& partialInputScratch,
int64_t patchSize,
int64_t numKSlices,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
auto elementType = state.xType.getElementType();
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
SmallVector<Value> inputTiles;
inputTiles.reserve(numKSlices);
if (packFactor == 1 && state.numChannelsIn % xbarDim == 0) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
auto inputTileType = RankedTensorType::get(
{1, 1, 1, xbarDim}, elementType, state.xType.getEncoding());
for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) {
const int64_t linearOffset = kSlice * xbarDim;
const int64_t kernelPixel = linearOffset / state.numChannelsIn;
const int64_t kernelRow = kernelPixel / state.wWidth;
const int64_t kernelColumn = kernelPixel % state.wWidth;
const int64_t channelOffset = linearOffset % state.numChannelsIn;
Value inputWidthOffset =
affineMulConst(rewriter, loc, outputWidth, state.strideWidth, anchorOp);
inputWidthOffset = affineAddConst(
rewriter, loc, inputWidthOffset, kernelColumn * state.dilationWidth, anchorOp);
Value inputTile = tensor::ExtractSliceOp::create(
rewriter,
loc,
inputTileType,
paddedWindow,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(kernelRow),
inputWidthOffset,
rewriter.getIndexAttr(channelOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 4));
inputTiles.push_back(tensor::CollapseShapeOp::create(
rewriter, loc, paddedRowType, inputTile, SmallVector<ReassociationIndices> {{0, 1, 2}, {3}})
.getResult());
}
return inputTiles;
}
FailureOr<Value> patchRow = createPackedPixelMajorConvPatchRow(
paddedWindow, state, outputWidth, packFactor, rewriter, loc);
if (failed(patchRow))
return failure();
for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) {
const int64_t kOffset = kSlice * xbarDim;
const int64_t sliceSize = std::min(xbarDim, patchSize - kOffset);
Value inputTile;
if (sliceSize == xbarDim) {
inputTile = extractStaticSliceOrIdentity(
rewriter,
loc,
*patchRow,
paddedRowType,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 2));
}
else {
if (!partialInputScratch)
return failure();
auto partialType = RankedTensorType::get({1, sliceSize}, elementType);
Value partial = extractStaticSliceOrIdentity(
rewriter,
loc,
*patchRow,
partialType,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(kOffset)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)},
getUnitStrides(rewriter, 2));
partialInputScratch = tensor::InsertSliceOp::create(
rewriter,
loc,
partial,
partialInputScratch,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(sliceSize)},
getUnitStrides(rewriter, 2));
inputTile = partialInputScratch;
}
inputTiles.push_back(inputTile);
}
return inputTiles;
}
static FailureOr<Value> createConvOutputTile(ValueRange inputTiles,
Value tileWeights,
int64_t outputChannels,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
auto elementType = cast<RankedTensorType>(inputTiles.front().getType()).getElementType();
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto resultType = RankedTensorType::get({1, outputChannels}, elementType);
auto weightElementType = cast<RankedTensorType>(tileWeights.getType()).getElementType();
auto paddedWeightTileType = RankedTensorType::get({xbarDim, xbarDim}, weightElementType);
Value tileResult;
for (auto [kSlice, inputTile] : llvm::enumerate(inputTiles)) {
const int64_t kOffset = static_cast<int64_t>(kSlice) * xbarDim;
SmallVector<OpFoldResult> bOffsets {
rewriter.getIndexAttr(kOffset), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)};
Value bTile = extractStaticSliceOrIdentity(
rewriter, loc, tileWeights, paddedWeightTileType, bOffsets, bSizes, getUnitStrides(rewriter, 2));
Value piece = spatial::SpatVMMOp::create(
rewriter, loc, paddedRowType, bTile, inputTile).getResult();
if (outputChannels != xbarDim)
piece = tensor::ExtractSliceOp::create(
rewriter,
loc,
resultType,
piece,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(outputChannels)},
getUnitStrides(rewriter, 2));
tileResult = tileResult
? spatial::SpatVAddOp::create(
rewriter, loc, resultType, tileResult, piece).getResult()
: piece;
}
return tileResult;
}
static FailureOr<Value> createConvOutputRow(ValueRange inputTiles,
int64_t paddedK,
int64_t outputChannels,
Value paddedWeights,
Value bias,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
auto elementType = cast<RankedTensorType>(inputTiles.front().getType()).getElementType();
auto rowType = RankedTensorType::get({1, outputChannels}, elementType);
const int64_t outputTileCount = ceilIntegerDivide(outputChannels, xbarDim);
const int64_t paddedOutputChannels = outputTileCount * xbarDim;
auto paddedOutputType = RankedTensorType::get({1, paddedOutputChannels}, elementType);
auto weightSliceType = RankedTensorType::get(
{xbarDim, paddedOutputChannels},
cast<RankedTensorType>(paddedWeights.getType()).getElementType());
Value paddedOutput;
for (auto [kSlice, inputTile] : llvm::enumerate(inputTiles)) {
const int64_t kOffset = static_cast<int64_t>(kSlice) * xbarDim;
Value weightSlice = extractStaticSliceOrIdentity(
rewriter,
loc,
paddedWeights,
weightSliceType,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(kOffset), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(xbarDim),
rewriter.getIndexAttr(paddedOutputChannels)},
getUnitStrides(rewriter, 2));
Value piece =
spatial::SpatVMMOp::create(rewriter, loc, paddedOutputType, weightSlice, inputTile).getResult();
paddedOutput = paddedOutput
? spatial::SpatVAddOp::create(
rewriter, loc, paddedOutputType, paddedOutput, piece).getResult()
: piece;
}
Value validRow = outputChannels == paddedOutputChannels
? paddedOutput
: tensor::ExtractSliceOp::create(
rewriter,
loc,
rowType,
paddedOutput,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(outputChannels)},
getUnitStrides(rewriter, 2))
.getResult();
if (bias)
validRow = spatial::SpatVAddOp::create(rewriter, loc, rowType, validRow, bias).getResult();
return validRow;
}
static FailureOr<Value> createRowStripConvOutput(const ConvLoweringState& state,
Value input,
Value paddedWeights,
Value bias,
int64_t packFactor,
int64_t paddedK,
int64_t numKSlices,
int64_t xbarDim,
PatternRewriter& rewriter,
Location loc) {
const int64_t laneCount = state.outHeight;
const int64_t patchSize = packFactor * state.numChannelsIn * state.wHeight * state.wWidth;
const bool hasPartialInputTile = patchSize % xbarDim != 0;
auto elementType = state.outType.getElementType();
auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType);
auto outputPixelType = RankedTensorType::get({1, 1, packFactor, state.numChannelsOut}, elementType);
auto fragmentType = getRowStripFragmentType(state.outType);
auto storageType = getRowStripStorageType(state.outType);
auto batch = createSpatComputeBatch(
rewriter,
loc,
TypeRange {storageType},
laneCount,
ValueRange {paddedWeights},
bias ? ValueRange {input, bias} : ValueRange {input},
[&](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 / packFactor);
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);
SmallVector<Value> loopInit {fragmentInit};
if (hasPartialInputTile)
loopInit.push_back(createZeroTensorConstant(partialInputScratchType, rewriter));
auto loop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cOutWidth,
c1,
loopInit,
[&](OpBuilder&,
Location pixelLoc,
Value localColumn,
ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
Value partialInputScratch = hasPartialInputTile ? iterArgs[1] : Value();
FailureOr<SmallVector<Value>> inputTiles = createConvInputTiles(*inputWindow,
state,
localColumn,
packFactor,
partialInputScratch,
patchSize,
numKSlices,
xbarDim,
rewriter,
pixelLoc);
if (failed(inputTiles))
return failure();
FailureOr<Value> output = createConvOutputRow(*inputTiles,
paddedK,
packFactor * state.numChannelsOut,
args.weights.front(),
bias ? args.inputs[1] : Value(),
xbarDim,
rewriter,
pixelLoc);
if (failed(output))
return failure();
Value outputPixel = tensor::ExpandShapeOp::create(
rewriter, pixelLoc, outputPixelType, *output, SmallVector<ReassociationIndices> {{0, 1}, {2, 3}});
Value next = tensor::InsertSliceOp::create(
rewriter,
pixelLoc,
outputPixel,
iterArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
affineMulConst(rewriter, pixelLoc, localColumn, packFactor, anchorOp),
rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(packFactor),
rewriter.getIndexAttr(state.numChannelsOut)},
getUnitStrides(rewriter, 4));
yielded.push_back(next);
if (hasPartialInputTile)
yielded.push_back(partialInputScratch);
return success();
});
if (failed(loop))
return failure();
publishGraphBatchPhysicalFragment(
rewriter, loc, loop->results.front(), args.outputs.front(), args.lane);
return success();
});
if (failed(batch))
return failure();
return batch->getResult(0);
}
static FailureOr<Value> createOutputChannelTiledRowStripConvOutput(const ConvLoweringState& state,
Value input,
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;
const bool hasPartialInputTile = patchSize % xbarDim != 0;
auto elementType = state.outType.getElementType();
auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType);
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto tilePixelType = RankedTensorType::get({1, 1, 1, xbarDim}, elementType);
auto tileFragmentType = RankedTensorType::get({1, 1, state.outWidth, xbarDim}, elementType);
auto tileWeightsType = RankedTensorType::get({paddedK, xbarDim}, state.wType.getElementType());
const int64_t laneCount = state.outHeight * outputTileCount;
auto tileStorageType = spatial::getGraphBatchPhysicalResultType(laneCount, tileFragmentType);
FailureOr<Value> paddedBias = failure();
if (state.hasBias)
paddedBias = createPaddedBiasTileConstant(state, xbarDim, rewriter);
if (state.hasBias && failed(paddedBias))
return failure();
auto tileBatch = createSpatComputeBatch(
rewriter,
loc,
TypeRange {tileStorageType},
laneCount,
ValueRange {paddedWeights},
state.hasBias ? ValueRange {input, *paddedBias} : ValueRange {input},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth);
Value outputRow = affineFloorDivConst(rewriter, loc, args.lane, outputTileCount, anchorOp);
Value outputTile = affineModConst(rewriter, loc, args.lane, outputTileCount, anchorOp);
SmallVector<OpFoldResult> weightOffsets {
outputTile, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> weightSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(paddedK), rewriter.getIndexAttr(xbarDim)};
Value tileWeights = tensor::ExtractSliceOp::create(
rewriter, loc, tileWeightsType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3));
FailureOr<Value> biasTile = failure();
if (state.hasBias)
biasTile = extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[1], outputTile, paddedRowType);
if (state.hasBias && failed(biasTile))
return failure();
FailureOr<Value> inputWindow =
createConvInputWindow(args.inputs.front(), state, outputRow, rewriter, loc);
if (failed(inputWindow))
return failure();
Value fragmentInit = tensor::EmptyOp::create(rewriter, loc, tileFragmentType.getShape(), elementType);
SmallVector<Value> widthLoopInit {fragmentInit};
if (hasPartialInputTile)
widthLoopInit.push_back(createZeroTensorConstant(partialInputScratchType, rewriter));
auto widthLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cOutWidth,
c1,
widthLoopInit,
[&](OpBuilder&,
Location widthLoc,
Value widthIndex,
ValueRange widthIterArgs,
SmallVectorImpl<Value>& widthYielded) {
Value partialInputScratch = hasPartialInputTile ? widthIterArgs[1] : Value();
FailureOr<SmallVector<Value>> inputTiles = createConvInputTiles(*inputWindow,
state,
widthIndex,
/*packFactor=*/1,
partialInputScratch,
patchSize,
numKSlices,
xbarDim,
rewriter,
widthLoc);
if (failed(inputTiles))
return failure();
FailureOr<Value> paddedOutputRow =
createConvOutputTile(*inputTiles, tileWeights, xbarDim, xbarDim, rewriter, widthLoc);
if (failed(paddedOutputRow))
return failure();
if (state.hasBias)
paddedOutputRow =
spatial::SpatVAddOp::create(rewriter, widthLoc, paddedRowType, *paddedOutputRow, *biasTile).getResult();
Value outputPixel = tensor::ExpandShapeOp::create(
rewriter, widthLoc, tilePixelType, *paddedOutputRow, SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
SmallVector<OpFoldResult> rowOffsets {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), widthIndex, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> rowSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(xbarDim)};
Value nextFragment = tensor::InsertSliceOp::create(rewriter,
widthLoc,
outputPixel,
widthIterArgs.front(),
rowOffsets,
rowSizes,
getUnitStrides(rewriter, 4));
widthYielded.push_back(nextFragment);
if (hasPartialInputTile)
widthYielded.push_back(partialInputScratch);
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();
return tileBatch->getResult(0);
}
static FailureOr<Value>
createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) {
ConvGeometry geometry = buildConvGeometry(state, state.targetInfo());
if (state.group != 1 || state.batchSize != 1 || !rowStripOutputChannelTileFitsOneCore(geometry))
return failure();
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!weightDenseAttr)
return failure();
if (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType))
return failure();
const int64_t xbarDim = geometry.xbarSize;
const int64_t basePatchSize = state.numChannelsIn * state.wHeight * state.wWidth;
const int64_t baseNumKSlices = ceilIntegerDivide(basePatchSize, xbarDim);
const int64_t basePaddedK = baseNumKSlices * xbarDim;
if (!rowStripOutputTileFitsOneCore(geometry)) {
Value tiledWeights =
standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, basePaddedK, xbarDim, rewriter);
return createOutputChannelTiledRowStripConvOutput(
state, state.x, tiledWeights, basePaddedK, baseNumKSlices, xbarDim, rewriter, loc);
}
const int64_t packFactor = chooseRowStripPixelPackFactor(state, xbarDim);
const int64_t packedPatchSize = packFactor * basePatchSize;
const int64_t numKSlices = ceilIntegerDivide(packedPatchSize, xbarDim);
const int64_t paddedK = numKSlices * xbarDim;
const int64_t packedOutputChannels = packFactor * state.numChannelsOut;
const int64_t paddedOutputChannels = ceilIntegerDivide(packedOutputChannels, xbarDim) * xbarDim;
Value paddedWeights = standard::createPaddedPixelMajorWeightConstant(
weightDenseAttr, state, paddedK, paddedOutputChannels, packFactor, rewriter);
FailureOr<Value> bias = failure();
if (state.hasBias)
bias = createBiasRowConstant(state, packFactor, rewriter);
if (state.hasBias && failed(bias))
return failure();
return createRowStripConvOutput(
state, state.x, paddedWeights, state.hasBias ? *bias : Value(),
packFactor, paddedK, numKSlices, xbarDim, rewriter, loc);
}
static FailureOr<Value> createConvOutputFromPixelMajorRowStripFragments(Value rowStripStorage,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
if (failed(describeRowStripPhysicalValue(rowStripStorage, state.xType)))
return failure();
StringRef failureReason;
if (!canConsumePixelMajorRowStripFragments(state, failureReason))
return failure();
ConvGeometry geometry = buildConvGeometry(state, state.targetInfo());
const int64_t xbarDim = geometry.xbarSize;
const int64_t basePatchSize = state.numChannelsIn * state.wHeight * state.wWidth;
const int64_t baseNumKSlices = ceilIntegerDivide(basePatchSize, xbarDim);
const int64_t basePaddedK = baseNumKSlices * xbarDim;
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!weightDenseAttr)
return failure();
if (!rowStripOutputTileFitsOneCore(geometry)) {
Value tiledWeights =
standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, basePaddedK, xbarDim, rewriter);
return createOutputChannelTiledRowStripConvOutput(
state, rowStripStorage, tiledWeights, basePaddedK, baseNumKSlices, xbarDim, rewriter, loc);
}
const int64_t packFactor = chooseRowStripPixelPackFactor(state, xbarDim);
const int64_t packedPatchSize = packFactor * basePatchSize;
const int64_t numKSlices = ceilIntegerDivide(packedPatchSize, xbarDim);
const int64_t paddedK = numKSlices * xbarDim;
const int64_t packedOutputChannels = packFactor * state.numChannelsOut;
const int64_t paddedOutputChannels = ceilIntegerDivide(packedOutputChannels, xbarDim) * xbarDim;
Value paddedWeights = standard::createPaddedPixelMajorWeightConstant(
weightDenseAttr, state, paddedK, paddedOutputChannels, packFactor, rewriter);
FailureOr<Value> bias = failure();
if (state.hasBias)
bias = createBiasRowConstant(state, packFactor, rewriter);
if (state.hasBias && failed(bias))
return failure();
return createRowStripConvOutput(
state, rowStripStorage, paddedWeights, state.hasBias ? *bias : Value(),
packFactor, paddedK, numKSlices, xbarDim, rewriter, loc);
}
static FailureOr<Value> createPointwiseOutputFromRowStripFragments(Value rowStripStorage,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
FailureOr<RowStripPhysicalValue> input = describeRowStripPhysicalValue(rowStripStorage, state.xType);
if (failed(input)) return failure();
ConvGeometry geometry = buildConvGeometry(state, state.targetInfo());
const int64_t xbarDim = geometry.xbarSize;
const int64_t inputFragmentChannels = input->fragmentType.getDimSize(3);
if (inputFragmentChannels % xbarDim != 0 || state.numChannelsIn % xbarDim != 0)
return failure();
auto weightDenseAttr = getHostConstDenseElementsAttr(state.w);
if (!weightDenseAttr) return failure();
const int64_t outputTileCount = ceilIntegerDivide(state.numChannelsOut, xbarDim);
const int64_t numKSlices = state.numChannelsIn / xbarDim;
auto elementType = state.outType.getElementType();
auto paddedRowType = RankedTensorType::get({1, xbarDim}, elementType);
auto inputRowType = RankedTensorType::get({1, inputFragmentChannels}, elementType);
auto weightTileType = RankedTensorType::get({state.numChannelsIn, xbarDim}, state.wType.getElementType());
auto weightSliceType = RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType());
auto outputFragmentType = RankedTensorType::get({1, 1, 1, xbarDim}, elementType);
auto outputStorageType = spatial::getGraphBatchPhysicalResultType(outputTileCount, outputFragmentType);
Value paddedWeights = standard::createPaddedOutputChannelTiledWeightConstant(
weightDenseAttr, state, state.numChannelsIn, xbarDim, rewriter);
FailureOr<Value> paddedBias = failure();
if (state.hasBias) paddedBias = createPaddedBiasTileConstant(state, xbarDim, rewriter);
if (state.hasBias && failed(paddedBias)) return failure();
auto batch = createSpatComputeBatch(rewriter, loc, TypeRange {outputStorageType}, outputTileCount,
ValueRange {paddedWeights},
state.hasBias ? ValueRange {rowStripStorage, *paddedBias} : ValueRange {rowStripStorage},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cNumKSlices = getOrCreateIndexConstant(rewriter, anchorOp, numKSlices);
SmallVector<OpFoldResult> weightOffsets {args.lane, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> weightSizes {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(state.numChannelsIn), rewriter.getIndexAttr(xbarDim)};
Value weightTile = tensor::ExtractSliceOp::create(
rewriter, loc, weightTileType, args.weights.front(), weightOffsets, weightSizes, getUnitStrides(rewriter, 3));
auto createPiece = [&](Value kSlice, Location pieceLoc) -> FailureOr<Value> {
Value channelOffset = affineMulConst(rewriter, pieceLoc, kSlice, xbarDim, anchorOp);
Value sourceSlot = affineFloorDivConst(
rewriter, pieceLoc, channelOffset, inputFragmentChannels, anchorOp);
Value sourceOffset = affineModConst(
rewriter, pieceLoc, channelOffset, inputFragmentChannels, anchorOp);
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
rewriter, pieceLoc, args.inputs.front(), sourceSlot, input->fragmentType);
if (failed(fragment)) return failure();
Value inputRow = tensor::CollapseShapeOp::create(rewriter, pieceLoc, inputRowType, *fragment,
SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
Value inputSlice = tensor::ExtractSliceOp::create(rewriter, pieceLoc, paddedRowType, inputRow,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), sourceOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 2));
Value weightSlice = tensor::ExtractSliceOp::create(rewriter, pieceLoc, weightSliceType, weightTile,
SmallVector<OpFoldResult> {channelOffset, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 2));
return spatial::SpatVMMOp::create(rewriter, pieceLoc, paddedRowType, weightSlice, inputSlice).getResult();
};
FailureOr<Value> result = createPiece(c0, loc);
if (failed(result)) return failure();
if (numKSlices > 1) {
auto reduction = buildNormalizedScfFor(rewriter, loc, c1, cNumKSlices, c1, ValueRange {*result},
[&](OpBuilder&, Location reduceLoc, Value kSlice, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
FailureOr<Value> piece = createPiece(kSlice, reduceLoc);
if (failed(piece)) return failure();
yielded.push_back(spatial::SpatVAddOp::create(
rewriter, reduceLoc, paddedRowType, iterArgs.front(), *piece).getResult());
return success();
});
if (failed(reduction)) return failure();
result = reduction->results.front();
}
if (state.hasBias) {
FailureOr<Value> bias = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs[1], args.lane, paddedRowType);
if (failed(bias)) return failure();
result = spatial::SpatVAddOp::create(rewriter, loc, paddedRowType, *result, *bias).getResult();
}
Value fragment = tensor::ExpandShapeOp::create(rewriter, loc, outputFragmentType, *result,
SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
publishGraphBatchPhysicalFragment(rewriter, loc, fragment, args.outputs.front(), args.lane);
return success();
});
if (failed(batch)) return failure();
return batch->getResult(0);
}
static bool canConsumeDepthwiseRowStrip(const ConvLoweringState& state) {
if (state.batchSize != 1 || state.group != state.numChannelsIn
|| state.dilationHeight != 1 || state.dilationWidth != 1
|| !isa<FloatType>(state.xType.getElementType())
|| !getHostConstDenseElementsAttr(state.w)
|| (state.hasBias && !isSupportedBiasAddValue(state.b, state.outType)))
return false;
auto tiling = depthwise::computeTiling(state.batchSize,
state.numChannelsIn,
state.numChannelsOut,
state.wHeight,
state.wWidth,
state.outHeight,
state.outWidth,
state.targetInfo().matrixShape.rows);
return tiling && tiling->numChannelTiles
<= static_cast<int64_t>(state.targetInfo().matrixUnitsPerProcessor);
}
static Value insertDepthwiseInputSegment(Value inputWindow,
Value scratch,
Value tileIndex,
Value kernelRow,
Value sourceWidth,
Value scratchOffset,
int64_t inputChannel,
const depthwise::Tiling& tiling,
PatternRewriter& rewriter,
Location loc) {
auto inputWindowType = cast<RankedTensorType>(inputWindow.getType());
auto inputPixelType = RankedTensorType::get(
{1, 1, 1, tiling.channelsPerTile}, inputWindowType.getElementType(), inputWindowType.getEncoding());
Value inputPixel = tensor::ExtractSliceOp::create(
rewriter,
loc,
inputPixelType,
inputWindow,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), kernelRow, sourceWidth,
rewriter.getIndexAttr(inputChannel)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1), rewriter.getIndexAttr(tiling.channelsPerTile)},
getUnitStrides(rewriter, 4));
return tensor::InsertSliceOp::create(
rewriter,
loc,
inputPixel,
scratch,
SmallVector<OpFoldResult> {tileIndex, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), scratchOffset},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tiling.channelsPerTile)},
getUnitStrides(rewriter, 4));
}
static FailureOr<Value> assembleDepthwiseInputScratch(Value inputWindow,
Value scratch,
Value inputWidth,
const ConvLoweringState& state,
const depthwise::Tiling& tiling,
PatternRewriter& rewriter,
Location loc) {
Operation* anchor = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchor, 1);
Value cKernelElements = getOrCreateIndexConstant(rewriter, anchor, tiling.kernelElements);
SmallVector<Value> tileIndices;
tileIndices.reserve(tiling.numChannelTiles);
for (int64_t tile = 0; tile < tiling.numChannelTiles; ++tile)
tileIndices.push_back(getOrCreateIndexConstant(rewriter, anchor, tile));
auto kernelLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cKernelElements,
c1,
ValueRange {scratch},
[&](OpBuilder&, Location kernelLoc, Value kernelIndex, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
Value kernelRow = affineFloorDivConst(rewriter, kernelLoc, kernelIndex, state.wWidth, anchor);
Value kernelColumn = affineModConst(rewriter, kernelLoc, kernelIndex, state.wWidth, anchor);
Value sourceWidth = createOrFoldAffineApply(
rewriter,
kernelLoc,
getAffineDimExpr(0, rewriter.getContext()) + getAffineDimExpr(1, rewriter.getContext()),
ValueRange {inputWidth, kernelColumn},
anchor);
Value scratchOffset = affineMulConst(
rewriter, kernelLoc, kernelIndex, tiling.channelsPerTile, anchor);
Value nextScratch = iterArgs.front();
for (int64_t tile = 0; tile < tiling.numChannelTiles; ++tile)
nextScratch = insertDepthwiseInputSegment(inputWindow,
nextScratch,
tileIndices[tile],
kernelRow,
sourceWidth,
scratchOffset,
tile * tiling.channelsPerTile,
tiling,
rewriter,
kernelLoc);
yielded.push_back(nextScratch);
return success();
});
if (failed(kernelLoop))
return failure();
return kernelLoop->results.front();
}
static FailureOr<Value> createDepthwiseOutputFromRowStripFragments(Value rowStripStorage,
const ConvLoweringState& state,
PatternRewriter& rewriter,
Location loc) {
if (!canConsumeDepthwiseRowStrip(state)
|| failed(describeRowStripPhysicalValue(rowStripStorage, state.xType)))
return failure();
auto tiling = depthwise::computeTiling(state.batchSize,
state.numChannelsIn,
state.numChannelsOut,
state.wHeight,
state.wWidth,
state.outHeight,
state.outWidth,
state.targetInfo().matrixShape.rows);
auto weight = getHostConstDenseElementsAttr(state.w);
if (!tiling || !weight)
return failure();
Value packedWeights = depthwise::buildPackedWeights(
weight,
state.wType,
*tiling,
rewriter,
loc,
static_cast<int64_t>(state.targetInfo().matrixShape.rows),
static_cast<int64_t>(state.targetInfo().matrixShape.rows));
Value bias = state.hasBias ? expandBiasIfNeeded(state.b, rewriter, loc) : Value();
auto paddedOutputType = RankedTensorType::get(
{1, static_cast<int64_t>(state.targetInfo().matrixShape.rows)}, state.outType.getElementType());
auto outputTileType = RankedTensorType::get(
{1, tiling->tileOutputChannels}, state.outType.getElementType());
auto outputPixelType = RankedTensorType::get(
{1, 1, 1, tiling->tileOutputChannels}, state.outType.getElementType());
auto fragmentType = getRowStripFragmentType(state.outType);
auto storageType = getRowStripStorageType(state.outType);
auto batch = createSpatComputeBatch(
rewriter,
loc,
TypeRange {storageType},
state.outHeight,
ValueRange {packedWeights},
state.hasBias ? ValueRange {rowStripStorage, bias} : ValueRange {rowStripStorage},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchor = rewriter.getInsertionBlock()->getParentOp();
FailureOr<Value> inputWindow =
createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc);
if (failed(inputWindow))
return failure();
Value c0 = getOrCreateIndexConstant(rewriter, anchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchor, 1);
Value cOutWidth = getOrCreateIndexConstant(rewriter, anchor, state.outWidth);
const int64_t xbarDim = static_cast<int64_t>(state.targetInfo().matrixShape.rows);
auto paddedInputScratchType = RankedTensorType::get(
{tiling->numChannelTiles, 1, 1, xbarDim}, state.xType.getElementType(), state.xType.getEncoding());
auto tileScratchType = RankedTensorType::get(
{1, 1, 1, xbarDim}, state.xType.getElementType(), state.xType.getEncoding());
auto vmmInputType = RankedTensorType::get(
{1, xbarDim}, state.xType.getElementType(), state.xType.getEncoding());
Value zeroScratch = createZeroTensorConstant(paddedInputScratchType, rewriter);
Value fragment = tensor::EmptyOp::create(
rewriter, loc, fragmentType.getShape(), fragmentType.getElementType());
SmallVector<Value> weightTiles;
SmallVector<Value> biasTiles;
SmallVector<Value> tileIndices;
SmallVector<OpFoldResult> weightTileSizes {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim),
rewriter.getIndexAttr(xbarDim)};
for (int64_t tile = 0; tile < tiling->numChannelTiles; ++tile) {
Value tileIndex = getOrCreateIndexConstant(rewriter, anchor, tile);
tileIndices.push_back(tileIndex);
weightTiles.push_back(extractMixedSliceOrIdentity(
rewriter,
loc,
args.weights.front(),
RankedTensorType::get({xbarDim, xbarDim}, state.wType.getElementType()),
{SmallVector<OpFoldResult> {tileIndex, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
weightTileSizes,
getUnitStrides(rewriter, 3)}));
if (state.hasBias)
biasTiles.push_back(depthwise::createBiasTile(args.inputs[1], tileIndex, *tiling, rewriter, loc));
}
auto widthLoop = buildNormalizedScfFor(
rewriter, loc, c0, cOutWidth, c1, ValueRange {fragment, zeroScratch},
[&](OpBuilder&, Location widthLoc, Value width, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
Value next = iterArgs.front();
Value scratch = iterArgs[1];
// Valid prefix entries are overwritten per tile; the padded tail stays zero.
Value inputWidth = affineMulConst(
rewriter, widthLoc, width, state.strideWidth, anchor);
FailureOr<Value> nextScratch = assembleDepthwiseInputScratch(
*inputWindow, scratch, inputWidth, state, *tiling, rewriter, widthLoc);
if (failed(nextScratch))
return failure();
scratch = *nextScratch;
for (int64_t tile = 0; tile < tiling->numChannelTiles; ++tile) {
Value tileScratch = tensor::ExtractSliceOp::create(
rewriter,
widthLoc,
tileScratchType,
scratch,
SmallVector<OpFoldResult> {tileIndices[tile], rewriter.getIndexAttr(0), rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim)},
getUnitStrides(rewriter, 4));
Value vmmInput = tensor::CollapseShapeOp::create(
rewriter,
widthLoc,
vmmInputType,
tileScratch,
SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
Value output = spatial::SpatVMMOp::create(
rewriter, widthLoc, paddedOutputType, weightTiles[tile], vmmInput);
Value validOutput = tensor::ExtractSliceOp::create(
rewriter,
widthLoc,
outputTileType,
output,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(tiling->tileOutputChannels)},
getUnitStrides(rewriter, 2));
if (state.hasBias)
validOutput = spatial::SpatVAddOp::create(
rewriter, widthLoc, outputTileType, validOutput, biasTiles[tile]);
Value pixel = tensor::ExpandShapeOp::create(
rewriter, widthLoc, outputPixelType, validOutput,
SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
next = tensor::InsertSliceOp::create(
rewriter,
widthLoc,
pixel,
next,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), width,
rewriter.getIndexAttr(tile * tiling->tileOutputChannels)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1), rewriter.getIndexAttr(tiling->tileOutputChannels)},
getUnitStrides(rewriter, 4));
}
yielded.push_back(next);
yielded.push_back(scratch);
return success();
});
if (failed(widthLoop))
return failure();
publishGraphBatchPhysicalFragment(rewriter, loc, widthLoop->results.front(), args.outputs.front(), args.lane);
return success();
});
return failed(batch) ? FailureOr<Value>(failure()) : FailureOr<Value>(batch->getResult(0));
}
static FailureOr<Value> createConvOutputFromRowStripInput(const ConvLoweringState& state,
Value rowStripInput,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter,
Location loc) {
if (strategy == spatial::ConvLoweringStrategy::Depthwise)
return createDepthwiseOutputFromRowStripFragments(rowStripInput, state, rewriter, loc);
if (state.xHeight == 1 && state.xWidth == 1 && state.wHeight == 1 && state.wWidth == 1)
return createPointwiseOutputFromRowStripFragments(rowStripInput, state, rewriter, loc);
return createConvOutputFromPixelMajorRowStripFragments(rowStripInput, state, rewriter, loc);
}
static Value createCollectedConvOutput(ValueRange gemmRows,
Type convType,
RankedTensorType gemmOutType,
RankedTensorType nhwcType,
RankedTensorType outType,
int64_t numPatches,
int64_t numChannelsOut,
int64_t packFactor,
PatternRewriter& rewriter,
Location loc) {
auto collectComputeOp = createSpatCompute(rewriter, loc, convType, {}, gemmRows, [&](ValueRange gemmRowArgs) {
Value gemmOut;
if (packFactor == 1) {
gemmOut = createSpatConcat(rewriter, loc, /*axis=*/0, gemmRowArgs);
}
else {
Value packedOutput = createSpatConcat(rewriter, loc, /*axis=*/0, gemmRowArgs);
gemmOut = standard::unpackRowsFromParallelGemm(
packedOutput, cast<RankedTensorType>(packedOutput.getType()), numPatches, numChannelsOut, packFactor, rewriter, loc);
}
// Restore output layout:
// [numPatches, numChannelsOut]
// -> [N, Hout, Wout, Cout]
// -> [N, Cout, Hout, Wout]
Value nhwcOut = tensor::ExpandShapeOp::create(rewriter,
loc,
nhwcType,
gemmOut,
SmallVector<ReassociationIndices> {
{0, 1, 2},
{3}
});
Value nchwOut = createLinalgTranspose(nhwcOut, outType, {0, 3, 1, 2}, rewriter, loc);
spatial::SpatYieldOp::create(rewriter, loc, nchwOut);
});
return collectComputeOp.getResult(0);
}
static FailureOr<ConvLoweringState> analyzeConvLoweringState(ONNXConvOp convOp,
Value x,
Value w,
Value b,
const spatial::SpatialTargetInfo& target) {
ConvLoweringState state;
state.diagnosticAnchor = convOp.getOperation();
state.x = x;
state.w = w;
state.b = b;
state.target = &target;
state.xType = cast<RankedTensorType>(state.x.getType());
state.wType = cast<RankedTensorType>(state.w.getType());
state.outType = cast<RankedTensorType>(convOp.getY().getType());
if (!state.xType.hasStaticShape()) {
pim::emitUnsupportedStaticShapeDiagnostic(convOp, "conv input");
return failure();
}
if (!state.wType.hasStaticShape()) {
pim::emitUnsupportedStaticShapeDiagnostic(convOp, "conv weight");
return failure();
}
if (!state.outType.hasStaticShape()) {
pim::emitUnsupportedStaticShapeDiagnostic(convOp, "conv result");
return failure();
}
if (state.xType.getRank() != 4) {
pim::emitUnsupportedRankDiagnostic(convOp, "conv input", state.xType.getRank(), {4});
return failure();
}
if (state.wType.getRank() != 4) {
pim::emitUnsupportedRankDiagnostic(convOp, "conv weight", state.wType.getRank(), {4});
return failure();
}
if (state.outType.getRank() != 4) {
pim::emitUnsupportedRankDiagnostic(convOp, "conv result", state.outType.getRank(), {4});
return failure();
}
state.group = convOp.getGroup();
if (state.group < 1) {
convOp.emitOpError("requires group >= 1 for Spatial lowering");
return failure();
}
state.batchSize = state.xType.getDimSize(0);
state.numChannelsIn = state.xType.getDimSize(1);
state.xHeight = state.xType.getDimSize(2);
state.xWidth = state.xType.getDimSize(3);
state.numChannelsOut = state.wType.getDimSize(0);
state.wHeight = state.wType.getDimSize(2);
state.wWidth = state.wType.getDimSize(3);
state.outHeight = state.outType.getDimSize(2);
state.outWidth = state.outType.getDimSize(3);
state.hasBias =
state.b && !isa<ONNXNoneOp>(state.b.getDefiningOp()) && !isZeroSplatHostConstant(state.b);
if (state.numChannelsIn % state.group != 0) {
convOp.emitOpError() << "requires input channels " << state.numChannelsIn << " to be divisible by group "
<< state.group << " for Spatial lowering";
return failure();
}
if (state.numChannelsOut % state.group != 0) {
convOp.emitOpError() << "requires output channels " << state.numChannelsOut << " to be divisible by group "
<< state.group << " for Spatial lowering";
return failure();
}
state.numChannelsInPerGroup = state.numChannelsIn / state.group;
state.numChannelsOutPerGroup = state.numChannelsOut / state.group;
if (state.wType.getDimSize(1) != state.numChannelsInPerGroup) {
convOp.emitOpError() << "requires grouped conv weight input channels " << state.wType.getDimSize(1)
<< " to match input channels per group " << state.numChannelsInPerGroup
<< " for Spatial lowering";
return failure();
}
if (state.wType.getDimSize(0) != state.numChannelsOut) {
convOp.emitOpError() << "requires weight output channels " << state.wType.getDimSize(0)
<< " to match result channels " << state.numChannelsOut << " for Spatial lowering";
return failure();
}
const auto stridesAttr = convOp.getStrides();
const auto dilationsAttr = convOp.getDilations();
const auto padsAttr = convOp.getPads();
if (stridesAttr && stridesAttr->size() != 2) {
convOp.emitOpError("requires exactly two stride values for Spatial lowering");
return failure();
}
if (dilationsAttr && dilationsAttr->size() != 2) {
convOp.emitOpError("requires exactly two dilation values for Spatial lowering");
return failure();
}
if (padsAttr && padsAttr->size() != 4) {
convOp.emitOpError("requires exactly four pad values for 2D Spatial lowering");
return failure();
}
state.strideHeight = getOptionalI64Attr(stridesAttr, 0, 1);
state.strideWidth = getOptionalI64Attr(stridesAttr, 1, 1);
state.dilationHeight = getOptionalI64Attr(dilationsAttr, 0, 1);
state.dilationWidth = getOptionalI64Attr(dilationsAttr, 1, 1);
state.padHeightBegin = 0;
state.padHeightEnd = 0;
state.padWidthBegin = 0;
state.padWidthEnd = 0;
if (padsAttr) {
state.padHeightBegin = getI64Attr(*padsAttr, 0);
state.padWidthBegin = getI64Attr(*padsAttr, 1);
state.padHeightEnd = getI64Attr(*padsAttr, 2);
state.padWidthEnd = getI64Attr(*padsAttr, 3);
classifyConvProblem(state);
return state;
}
const auto autoPad = convOp.getAutoPad();
if (autoPad == "SAME_UPPER" || autoPad == "SAME_LOWER") {
const int64_t effectiveKernelH = (state.wHeight - 1) * state.dilationHeight + 1;
const int64_t effectiveKernelW = (state.wWidth - 1) * state.dilationWidth + 1;
const int64_t totalPadH =
std::max(static_cast<int64_t>(0), (state.outHeight - 1) * state.strideHeight + effectiveKernelH - state.xHeight);
const int64_t totalPadW =
std::max(static_cast<int64_t>(0), (state.outWidth - 1) * state.strideWidth + effectiveKernelW - state.xWidth);
if (autoPad == "SAME_UPPER") {
state.padHeightBegin = totalPadH / 2;
state.padHeightEnd = totalPadH - state.padHeightBegin;
state.padWidthBegin = totalPadW / 2;
state.padWidthEnd = totalPadW - state.padWidthBegin;
}
else {
state.padHeightEnd = totalPadH / 2;
state.padHeightBegin = totalPadH - state.padHeightEnd;
state.padWidthEnd = totalPadW / 2;
state.padWidthBegin = totalPadW - state.padWidthEnd;
}
classifyConvProblem(state);
return state;
}
if (autoPad != "NOTSET" && autoPad != "VALID") {
convOp.emitOpError() << "unsupported auto_pad value `" << autoPad << "` for Spatial lowering";
return failure();
}
classifyConvProblem(state);
return state;
}
static FailureOr<ConvLoweringState> analyzeConvLoweringState(ONNXConvOp convOp,
ONNXConvOpAdaptor convOpAdaptor,
const spatial::SpatialTargetInfo& target) {
return analyzeConvLoweringState(
convOp, convOpAdaptor.getX(), convOpAdaptor.getW(), convOpAdaptor.getB(), target);
}
static FailureOr<ConvLoweringState> analyzeConvLoweringState(
spatial::SpatConv2DPlanOp planOp, const spatial::SpatialTargetInfo& target) {
ConvLoweringState state;
state.diagnosticAnchor = planOp.getOperation();
state.x = planOp.getInput();
state.w = planOp.getWeight();
state.b = planOp.getBias() ? planOp.getBias() : Value();
state.target = &target;
state.xType = dyn_cast<RankedTensorType>(state.x.getType());
state.wType = dyn_cast<RankedTensorType>(state.w.getType());
state.outType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!state.xType || !state.wType || !state.outType)
return planOp.emitOpError("requires ranked tensor input, weight, and output"), failure();
if (!state.xType.hasStaticShape() || !state.wType.hasStaticShape() || !state.outType.hasStaticShape())
return planOp.emitOpError("requires static input, weight, and output shapes"), failure();
if (state.xType.getRank() != 4 || state.wType.getRank() != 4 || state.outType.getRank() != 4)
return planOp.emitOpError("requires rank-4 input, weight, and output tensors"), failure();
state.group = planOp.getGroup();
if (state.group < 1)
return planOp.emitOpError("requires group >= 1"), failure();
state.batchSize = state.xType.getDimSize(0);
state.numChannelsIn = state.xType.getDimSize(1);
state.xHeight = state.xType.getDimSize(2);
state.xWidth = state.xType.getDimSize(3);
state.numChannelsOut = state.wType.getDimSize(0);
state.wHeight = state.wType.getDimSize(2);
state.wWidth = state.wType.getDimSize(3);
state.outHeight = state.outType.getDimSize(2);
state.outWidth = state.outType.getDimSize(3);
state.hasBias = planOp.getBias() && !isZeroSplatHostConstant(planOp.getBias());
if (state.numChannelsIn % state.group != 0 || state.numChannelsOut % state.group != 0)
return planOp.emitOpError("requires input and output channels divisible by group"), failure();
state.numChannelsInPerGroup = state.numChannelsIn / state.group;
state.numChannelsOutPerGroup = state.numChannelsOut / state.group;
if (state.wType.getDimSize(1) != state.numChannelsInPerGroup)
return planOp.emitOpError("requires grouped conv weight channels to match input channels per group"), failure();
auto pads = planOp.getPads();
auto strides = planOp.getStrides();
auto dilations = planOp.getDilations();
if (pads.size() != 4 || strides.size() != 2 || dilations.size() != 2)
return planOp.emitOpError("requires 4 pads, 2 strides, and 2 dilations"), failure();
state.padHeightBegin = pads[0];
state.padWidthBegin = pads[1];
state.padHeightEnd = pads[2];
state.padWidthEnd = pads[3];
state.strideHeight = strides[0];
state.strideWidth = strides[1];
state.dilationHeight = dilations[0];
state.dilationWidth = dilations[1];
classifyConvProblem(state);
return state;
}
static FailureOr<spatial::ConvLoweringStrategy>
resolveRequestedConvLoweringStrategy(Operation* op, const spatial::SpatialTargetInfo& target) {
if (!target.useExperimentalConvImplementation)
return target.convLoweringStrategy;
if (target.convLoweringStrategy != spatial::ConvLoweringStrategy::Auto
&& target.convLoweringStrategy != spatial::ConvLoweringStrategy::PackedIm2Col) {
op->emitOpError() << "--use-experimental-conv-impl conflicts with --pim-conv-lowering="
<< stringifyConvLoweringStrategy(target.convLoweringStrategy);
return failure();
}
return spatial::ConvLoweringStrategy::PackedIm2Col;
}
static FailureOr<ConvPlan> selectConvLoweringPlan(
Operation* op, const ConvLoweringState& state, bool reportPlanning) {
FailureOr<spatial::ConvLoweringStrategy> requested =
resolveRequestedConvLoweringStrategy(op, state.targetInfo());
if (failed(requested))
return failure();
if (*requested == spatial::ConvLoweringStrategy::Auto) {
for (const ConvPlan& candidate : buildConvPlanCandidates(state, state.targetInfo())) {
if (candidate.strategy == spatial::ConvLoweringStrategy::Depthwise
&& !depthwise::canUseStructuredRewrite(state)) {
continue;
}
if (reportPlanning)
recordConvLoweringReport(op, ConvLoweringReportPhase::Planning, candidate.strategy, "SEL");
return candidate;
}
op->emitOpError("has no applicable Conv lowering candidate for the injected Spatial target");
return failure();
}
FailureOr<ConvPlan> candidate = makeConvPlan(state, *requested, state.targetInfo());
if (failed(candidate)) {
op->emitOpError() << "forced Conv lowering `"
<< stringifyConvLoweringStrategy(*requested)
<< "` is not applicable to this Conv problem";
return failure();
}
if (reportPlanning)
recordConvLoweringReport(op, ConvLoweringReportPhase::Planning, candidate->strategy, "SEL");
return *candidate;
}
static FailureOr<Value> lowerDenseSelectedConvPlan(Operation* op,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter,
Location loc);
static ConvLoweringState makeGroupedConvLoweringState(const ConvLoweringState& parent,
Value groupX,
Value groupW,
Value groupB,
RankedTensorType groupOutType);
static FailureOr<Value> buildConvValueForStrategy(Operation* op,
Location loc,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter);
static FailureOr<Value> buildGroupedConvValue(Operation* op,
Location loc,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter);
static FailureOr<Value> lowerGroupedSelectedConvPlan(Operation* op,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter,
Location loc) {
return buildGroupedConvValue(op, loc, state, strategy, rewriter);
}
static FailureOr<Value> lowerDenseSelectedConvPlan(Operation* op,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter,
Location loc) {
return buildConvValueForStrategy(op, loc, state, strategy, rewriter);
}
static FailureOr<Value> buildConvValueForStrategy(Operation* op,
Location loc,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter) {
const ConvGeometry geo = buildConvGeometry(state, state.targetInfo());
switch (strategy) {
case spatial::ConvLoweringStrategy::Depthwise: {
return depthwise::rewriteConv(op, state, rewriter, loc);
}
case spatial::ConvLoweringStrategy::Legacy:
case spatial::ConvLoweringStrategy::PackedIm2Col: {
return standard::rewritePackedIm2ColConv(state, rewriter, loc);
}
case spatial::ConvLoweringStrategy::StreamedPatch:
case spatial::ConvLoweringStrategy::OutputChannelTiled:
case spatial::ConvLoweringStrategy::Tiled2D: {
return standard::rewriteStreamedConv(state, rewriter, loc, /*forcedPackFactor=*/1);
}
case spatial::ConvLoweringStrategy::InputKTiled: {
return standard::rewriteInputKTiledConv(state, rewriter, loc);
}
case spatial::ConvLoweringStrategy::StreamedPacked: {
return standard::rewriteStreamedConv(state, rewriter, loc, geo.pack);
}
case spatial::ConvLoweringStrategy::Auto:
break;
}
op->emitOpError("unexpected auto strategy at Conv lowering dispatch");
return failure();
}
static ConvLoweringState makeGroupedConvLoweringState(const ConvLoweringState& parent,
Value groupX,
Value groupW,
Value groupB,
RankedTensorType groupOutType);
static ConvLoweringState makeGroupedConvLoweringState(
const ConvLoweringState& parent, Value groupX, Value groupW, Value groupB, RankedTensorType groupOutType) {
ConvLoweringState state = parent;
state.x = groupX;
state.w = groupW;
state.b = groupB;
state.xType = cast<RankedTensorType>(groupX.getType());
state.wType = cast<RankedTensorType>(groupW.getType());
state.outType = groupOutType;
state.batchSize = state.xType.getDimSize(0);
state.numChannelsIn = state.xType.getDimSize(1);
state.xHeight = state.xType.getDimSize(2);
state.xWidth = state.xType.getDimSize(3);
state.numChannelsOut = state.wType.getDimSize(0);
state.wHeight = state.wType.getDimSize(2);
state.wWidth = state.wType.getDimSize(3);
state.outHeight = state.outType.getDimSize(2);
state.outWidth = state.outType.getDimSize(3);
state.group = 1;
state.numChannelsInPerGroup = state.numChannelsIn;
state.numChannelsOutPerGroup = state.numChannelsOut;
state.hasBias = static_cast<bool>(groupB);
classifyConvProblem(state);
return state;
}
static FailureOr<Value> buildGroupedConvValue(Operation* op,
Location loc,
const ConvLoweringState& state,
spatial::ConvLoweringStrategy strategy,
PatternRewriter& rewriter) {
SmallVector<Value> xSlices = sliceTensor(state.x, /*axis=*/1, state.numChannelsInPerGroup, rewriter, loc);
SmallVector<Value> wSlices = sliceTensor(state.w, /*axis=*/0, state.numChannelsOutPerGroup, rewriter, loc);
SmallVector<Value> bSlices;
if (state.hasBias) {
auto biasType = cast<RankedTensorType>(state.b.getType());
int64_t biasAxis = -1;
if (biasType.getRank() == 1)
biasAxis = 0;
else if (biasType.getRank() == 2)
biasAxis = biasType.getDimSize(0) != 1 ? 0 : 1;
else {
op->emitOpError() << "requires rank-1 or rank-2 bias for grouped convolution Spatial lowering, but got rank "
<< biasType.getRank();
return failure();
}
bSlices = sliceTensor(state.b, biasAxis, state.numChannelsOutPerGroup, rewriter, loc);
}
if (xSlices.size() != static_cast<size_t>(state.group) || wSlices.size() != static_cast<size_t>(state.group)
|| (state.hasBias && bSlices.size() != static_cast<size_t>(state.group))) {
op->emitOpError("failed to partition grouped convolution operands for Spatial lowering");
return failure();
}
SmallVector<Value> groupResults;
groupResults.reserve(state.group);
auto groupOutType = RankedTensorType::get(
{state.batchSize, state.numChannelsOutPerGroup, state.outHeight, state.outWidth}, state.outType.getElementType());
for (int64_t groupId = 0; groupId < state.group; groupId++) {
Value groupX = xSlices[groupId];
Value groupW = wSlices[groupId];
Value groupB = state.hasBias ? bSlices[groupId] : Value();
ConvLoweringState groupState = makeGroupedConvLoweringState(state, groupX, groupW, groupB, groupOutType);
FailureOr<Value> groupResult = buildConvValueForStrategy(op, loc, groupState, strategy, rewriter);
if (failed(groupResult))
return failure();
groupResults.push_back(*groupResult);
}
if (llvm::all_of(groupResults, isCompileTimeComputable))
return createSpatConcat(rewriter, loc, /*axis=*/1, groupResults);
auto concatCompute = createSpatCompute(rewriter, loc, TypeRange {state.outType}, {}, groupResults, [&](ValueRange args) {
spatial::SpatYieldOp::create(rewriter, loc, createSpatConcat(rewriter, loc, /*axis=*/1, args));
});
return concatCompute.getResult(0);
}
} // namespace
LogicalResult ConvToGemm::matchAndRewrite(ONNXConvOp convOp,
ONNXConvOpAdaptor convOpAdaptor,
ConversionPatternRewriter& rewriter) const {
FailureOr<ConvLoweringState> state = analyzeConvLoweringState(convOp, convOpAdaptor, target);
if (failed(state))
return failure();
SmallVector<int64_t> pads {
state->padHeightBegin, state->padWidthBegin, state->padHeightEnd, state->padWidthEnd};
SmallVector<int64_t> strides {state->strideHeight, state->strideWidth};
SmallVector<int64_t> dilations {state->dilationHeight, state->dilationWidth};
Value bias = state->hasBias ? convOpAdaptor.getB() : Value();
auto convPlan = spatial::SpatConv2DPlanOp::create(rewriter,
convOp.getLoc(),
convOp.getY().getType(),
convOpAdaptor.getX(),
convOpAdaptor.getW(),
bias,
rewriter.getDenseI64ArrayAttr(pads),
rewriter.getDenseI64ArrayAttr(strides),
rewriter.getDenseI64ArrayAttr(dilations),
rewriter.getI64IntegerAttr(state->group),
spatial::getNCHWLayout(rewriter.getContext()));
rewriter.replaceOp(convOp, convPlan.getResult());
return success();
}
void populateConvPatterns(RewritePatternSet& patterns,
MLIRContext* ctx,
const spatial::SpatialTargetInfo& target) {
patterns.insert<ConvToGemm>(ctx, target);
}
LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp,
const spatial::SpatialTargetInfo& target) {
FailureOr<ConvLoweringState> state = analyzeConvLoweringState(planOp, target);
if (failed(state))
return failure();
if (state->group != 1 || state->batchSize != 1)
return failure();
if (state->outType.getRank() != 4 || !state->outType.hasStaticShape())
return failure();
if (!getHostConstDenseElementsAttr(state->w))
return failure();
if (state->hasBias && !isSupportedBiasAddValue(state->b, state->outType))
return failure();
ConvGeometry geometry = buildConvGeometry(*state, state->targetInfo());
if (!rowStripOutputChannelTileFitsOneCore(geometry))
return failure();
FailureOr<ConvPlan> plan =
selectConvLoweringPlan(planOp.getOperation(), *state, /*reportPlanning=*/true);
if (failed(plan))
return failure();
switch (plan->strategy) {
case spatial::ConvLoweringStrategy::Legacy:
case spatial::ConvLoweringStrategy::Depthwise:
case spatial::ConvLoweringStrategy::PackedIm2Col:
case spatial::ConvLoweringStrategy::StreamedPatch:
case spatial::ConvLoweringStrategy::OutputChannelTiled:
case spatial::ConvLoweringStrategy::Tiled2D:
case spatial::ConvLoweringStrategy::StreamedPacked:
return success();
case spatial::ConvLoweringStrategy::Auto:
case spatial::ConvLoweringStrategy::InputKTiled:
return failure();
}
llvm_unreachable("unknown conv lowering strategy");
}
LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp,
const spatial::SpatialTargetInfo& target) {
FailureOr<ConvLoweringState> state = analyzeConvLoweringState(planOp, target);
if (failed(state))
return failure();
FailureOr<ConvPlan> plan =
selectConvLoweringPlan(planOp.getOperation(), *state, /*reportPlanning=*/true);
if (failed(plan))
return failure();
if (plan->strategy == spatial::ConvLoweringStrategy::Depthwise)
return canConsumeDepthwiseRowStrip(*state) ? success() : failure();
StringRef failureReason;
return canConsumePixelMajorRowStripFragments(*state, failureReason) ? success() : failure();
}
FailureOr<Value>
lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp,
std::optional<Value> rowStripInput,
bool emitRowStripLayout,
const spatial::SpatialTargetInfo& target,
PatternRewriter& rewriter) {
FailureOr<ConvLoweringState> state = analyzeConvLoweringState(planOp, target);
if (failed(state))
return failure();
FailureOr<ConvPlan> plan =
selectConvLoweringPlan(planOp.getOperation(), *state, /*reportPlanning=*/false);
if (failed(plan))
return failure();
auto reportRealization = [&](StringRef implementation) {
recordConvLoweringReport(
planOp.getOperation(), ConvLoweringReportPhase::Realization, plan->strategy, implementation);
};
if (emitRowStripLayout) {
if (rowStripInput) {
if (failed(canConsumeAndProduceRowStrip(planOp, target)))
return planOp.emitOpError("selected row-strip input/output layout is not supported for this Conv plan"), failure();
reportRealization(convRowStripInputImplementation(*state, plan->strategy));
return createConvOutputFromRowStripInput(
*state, *rowStripInput, plan->strategy, rewriter, planOp.getLoc());
}
if (failed(canLowerConvPlanToRowStrip(planOp, target)))
return planOp.emitOpError("selected row-strip layout is not supported for this Conv plan"), failure();
reportRealization("RSD");
FailureOr<Value> rowStripStorage = createRowStripConvOutputFromDenseInput(*state, rewriter, planOp.getLoc());
if (failed(rowStripStorage))
return planOp.emitOpError("failed to build row-strip fragment storage for the selected Conv plan"), failure();
return *rowStripStorage;
}
reportRealization(convLoweringImplementation(plan->strategy));
if (plan->strategy == spatial::ConvLoweringStrategy::Depthwise)
return lowerDenseSelectedConvPlan(planOp.getOperation(), *state, plan->strategy, rewriter, planOp.getLoc());
if (state->group != 1)
return lowerGroupedSelectedConvPlan(planOp.getOperation(), *state, plan->strategy, rewriter, planOp.getLoc());
return lowerDenseSelectedConvPlan(planOp.getOperation(), *state, plan->strategy, rewriter, planOp.getLoc());
}
} // namespace onnx_mlir