#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 operandLayouts) { LayoutAlternative alternative; alternative.operandLayouts.assign(operandLayouts.begin(), operandLayouts.end()); alternative.resultLayout = PhysicalLayout::NHWCRowStrip; alternative.intrinsicCost = -2; return alternative; } static bool hasRowStripInput(ArrayRef operandLayouts, unsigned index) { return index < operandLayouts.size() && operandLayouts[index] == PhysicalLayout::NHWCRowStrip; } SmallVector SpatConv2DPlanOp::getLayoutAlternatives( const SpatialTargetInfo& target, ArrayRef operandLayouts) { SmallVector 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 SpatReluPlanOp::getLayoutAlternatives( const SpatialTargetInfo&, ArrayRef operandLayouts) { SmallVector alternatives {denseAlternative(getOperation())}; if (hasRowStripInput(operandLayouts, 0)) alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts)); return alternatives; } SmallVector SpatSiluPlanOp::getLayoutAlternatives( const SpatialTargetInfo&, ArrayRef operandLayouts) { SmallVector alternatives {denseAlternative(getOperation())}; if (hasRowStripInput(operandLayouts, 0)) { LayoutAlternative alternative = rowStripAlternative(getOperation(), operandLayouts); alternative.intrinsicCost = -3; alternatives.push_back(std::move(alternative)); } return alternatives; } SmallVector SpatResizeNearestPlanOp::getLayoutAlternatives( const SpatialTargetInfo& target, ArrayRef operandLayouts) { SmallVector alternatives {denseAlternative(getOperation())}; if (hasRowStripInput(operandLayouts, 0) && succeeded(canLowerResizeNearestPlanToRowStrip(*this, target))) alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts)); return alternatives; } SmallVector SpatMaxPool2DPlanOp::getLayoutAlternatives( const SpatialTargetInfo& target, ArrayRef operandLayouts) { SmallVector 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 SpatGlobalAveragePoolPlanOp::getLayoutAlternatives( const SpatialTargetInfo& target, ArrayRef operandLayouts) { SmallVector 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 SpatBiasAddPlanOp::getLayoutAlternatives( const SpatialTargetInfo&, ArrayRef operandLayouts) { SmallVector alternatives {denseAlternative(getOperation())}; auto resultType = dyn_cast(getOutput().getType()); if (resultType && hasRowStripInput(operandLayouts, 0) && isSupportedBiasAddValue(getBias(), resultType)) alternatives.push_back(rowStripAlternative(getOperation(), {PhysicalLayout::NHWCRowStrip, PhysicalLayout::DenseNCHW})); return alternatives; } SmallVector SpatAddPlanOp::getLayoutAlternatives( const SpatialTargetInfo&, ArrayRef operandLayouts) { SmallVector alternatives {denseAlternative(getOperation())}; if (operandLayouts.size() >= 2 && hasRowStripInput(operandLayouts, 0) && hasRowStripInput(operandLayouts, 1)) alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts)); return alternatives; } SmallVector SpatConcatPlanOp::getLayoutAlternatives( const SpatialTargetInfo&, ArrayRef operandLayouts) { SmallVector 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