134 lines
6.0 KiB
C++
134 lines
6.0 KiB
C++
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp"
|
|
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
|
|
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
|
|
|
|
using namespace mlir;
|
|
|
|
namespace onnx_mlir::spatial {
|
|
|
|
static LayoutAlternative denseAlternative(Operation *op) {
|
|
LayoutAlternative alternative;
|
|
alternative.operandLayouts.assign(op->getNumOperands(), PhysicalLayout::DenseNCHW);
|
|
alternative.resultLayout = PhysicalLayout::DenseNCHW;
|
|
return alternative;
|
|
}
|
|
|
|
static LayoutAlternative rowStripAlternative(Operation *op,
|
|
ArrayRef<PhysicalLayout> operandLayouts) {
|
|
LayoutAlternative alternative;
|
|
alternative.operandLayouts.assign(operandLayouts.begin(), operandLayouts.end());
|
|
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
|
|
alternative.intrinsicCost = -2;
|
|
return alternative;
|
|
}
|
|
|
|
static bool hasRowStripInput(ArrayRef<PhysicalLayout> operandLayouts, unsigned index) {
|
|
return index < operandLayouts.size()
|
|
&& operandLayouts[index] == PhysicalLayout::NHWCRowStrip;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatConv2DPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (hasRowStripInput(operandLayouts, 0)) {
|
|
if (succeeded(canConsumeAndProduceRowStrip(*this, target)))
|
|
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
|
|
}
|
|
else if (succeeded(canLowerConvPlanToRowStrip(*this, target))) {
|
|
LayoutAlternative alternative = denseAlternative(getOperation());
|
|
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
|
|
alternative.intrinsicCost = -2;
|
|
alternatives.push_back(std::move(alternative));
|
|
}
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatReluPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (hasRowStripInput(operandLayouts, 0))
|
|
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatSiluPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (hasRowStripInput(operandLayouts, 0)) {
|
|
LayoutAlternative alternative = rowStripAlternative(getOperation(), operandLayouts);
|
|
alternative.intrinsicCost = -3;
|
|
alternatives.push_back(std::move(alternative));
|
|
}
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatResizeNearestPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (hasRowStripInput(operandLayouts, 0)
|
|
&& succeeded(canLowerResizeNearestPlanToRowStrip(*this, target)))
|
|
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatMaxPool2DPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (succeeded(canLowerMaxPoolPlanToRowStrip(*this, target))) {
|
|
LayoutAlternative alternative = denseAlternative(getOperation());
|
|
if (hasRowStripInput(operandLayouts, 0))
|
|
alternative = rowStripAlternative(getOperation(), operandLayouts);
|
|
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
|
|
alternative.intrinsicCost = -2;
|
|
alternatives.push_back(std::move(alternative));
|
|
}
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatGlobalAveragePoolPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (succeeded(canLowerGlobalAveragePoolPlanToRowStrip(*this, target))) {
|
|
LayoutAlternative alternative = denseAlternative(getOperation());
|
|
if (hasRowStripInput(operandLayouts, 0))
|
|
alternative = rowStripAlternative(getOperation(), operandLayouts);
|
|
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
|
|
alternative.intrinsicCost = -2;
|
|
alternatives.push_back(std::move(alternative));
|
|
}
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatBiasAddPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
auto resultType = dyn_cast<RankedTensorType>(getOutput().getType());
|
|
if (resultType && hasRowStripInput(operandLayouts, 0)
|
|
&& isSupportedBiasAddValue(getBias(), resultType))
|
|
alternatives.push_back(rowStripAlternative(getOperation(),
|
|
{PhysicalLayout::NHWCRowStrip,
|
|
PhysicalLayout::DenseNCHW}));
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatAddPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (operandLayouts.size() >= 2 && hasRowStripInput(operandLayouts, 0)
|
|
&& hasRowStripInput(operandLayouts, 1))
|
|
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
|
|
return alternatives;
|
|
}
|
|
|
|
SmallVector<LayoutAlternative> SpatConcatPlanOp::getLayoutAlternatives(
|
|
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
|
|
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
|
|
if (!operandLayouts.empty() && llvm::all_of(operandLayouts, [](PhysicalLayout layout) {
|
|
return layout == PhysicalLayout::NHWCRowStrip;
|
|
}))
|
|
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
|
|
return alternatives;
|
|
}
|
|
|
|
} // namespace onnx_mlir::spatial
|