From 942a9faa4f753006da0538f6453231a3217b2bfa Mon Sep 17 00:00:00 2001 From: NiccoloN Date: Mon, 3 Aug 2026 11:07:28 +0200 Subject: [PATCH] second temp commit: i will soft-reset and recommit after next changes --- README.md | 26 +- src/PIM/Compiler/PimCompilerUtils.cpp | 48 +- .../Conversion/ONNXToSpatial/CMakeLists.txt | 5 +- .../Common/ComputeRegionBuilder.cpp | 8 +- .../Common/ContractionMaterialization.cpp | 39 + .../Common/ContractionMaterialization.hpp | 23 + .../Common/ContractionPlanning.cpp | 68 + .../Common/ContractionPlanning.hpp | 41 + .../Common/ContractionProblem.hpp | 35 + .../Common/MatrixProductLowering.cpp | 55 + .../Common/MatrixProductLowering.hpp | 8 + .../Common/RowStripLayoutUtils.cpp | 47 +- .../Common/RowStripLayoutUtils.hpp | 23 + .../ONNXToSpatial/Common/ShapeTilingUtils.cpp | 11 +- .../ONNXToSpatial/Common/ShapeTilingUtils.hpp | 6 +- .../ONNXToSpatial/LowerSpatialPlansPass.cpp | 1196 +++++++++-------- .../ONNXToSpatial/ONNXToSpatialPass.cpp | 21 +- .../ONNXToSpatial/ONNXToSpatialVerifier.cpp | 3 +- src/PIM/Conversion/ONNXToSpatial/Patterns.cpp | 12 +- src/PIM/Conversion/ONNXToSpatial/Patterns.hpp | 28 +- .../ONNXToSpatial/Patterns/Math/Conv.cpp | 447 +++--- .../Patterns/Math/ConvGeometry.cpp | 296 +++- .../Patterns/Math/ConvGeometry.hpp | 80 +- .../Patterns/Math/Elementwise.cpp | 8 +- .../ONNXToSpatial/Patterns/Math/Gemm.cpp | 267 ++-- .../ONNXToSpatial/Patterns/Math/Gemm.hpp | 27 + .../ONNXToSpatial/Patterns/Math/MatMul.cpp | 214 +-- .../Patterns/Math/ReduceMean.cpp | 6 +- .../ONNXToSpatial/Patterns/NN/Pool.cpp | 252 +++- .../ONNXToSpatial/Patterns/NN/Relu.cpp | 2 +- .../ONNXToSpatial/Patterns/Tensor/Concat.cpp | 3 +- .../ONNXToSpatial/Patterns/Tensor/Flatten.cpp | 14 +- .../ONNXToSpatial/Patterns/Tensor/Resize.cpp | 6 +- .../Patterns/Tensor/Transpose.cpp | 4 +- .../Conversion/ONNXToSpatial/PlanLowering.hpp | 34 +- .../SpatialLayoutCapabilities.cpp | 133 ++ .../SpatialLayoutPlanningPass.cpp | 512 +++---- .../BatchCoreLoweringPatterns.cpp | 9 +- .../SpatialToPim/CoreLoweringPatterns.cpp | 9 +- src/PIM/Conversion/SpatialToPim/Patterns.cpp | 3 +- .../SpatialToPim/ReturnPathNormalization.cpp | 6 +- .../Bufferization/PimBufferizationPass.cpp | 360 +++-- src/PIM/Dialect/Spatial/CMakeLists.txt | 14 +- src/PIM/Dialect/Spatial/Spatial.td | 105 +- .../Dialect/Spatial/SpatialLayoutInterface.td | 24 + src/PIM/Dialect/Spatial/SpatialOps.cpp | 36 +- src/PIM/Dialect/Spatial/SpatialOps.hpp | 64 + src/PIM/Dialect/Spatial/SpatialOpsAsm.cpp | 21 +- src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp | 40 +- src/PIM/Dialect/Spatial/SpatialTargetInfo.hpp | 37 + .../DeferredCommunicationPlanning.cpp | 2 +- .../DeferredTransferPlanning.cpp | 2 +- .../MergeComputeNodesPass.cpp | 128 -- .../ScheduledComputePlanning.cpp | 2 +- .../ScheduledSpatialPasses.cpp | 279 ++++ .../ScheduledSpatialPasses.hpp | 16 + .../Scheduling/PeftScheduler.cpp | 2 +- .../TrivialGraphComputeMergePass.cpp | 1 - src/PIM/Pass/PIMPasses.h | 28 +- src/PIM/PimAccelerator.cpp | 16 +- .../googlenet/googlenet-12.onnx | Bin 28021836 -> 28021885 bytes validation/operations/validation_results.csv | 336 ++--- 62 files changed, 3657 insertions(+), 1891 deletions(-) create mode 100644 src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.cpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.cpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp create mode 100644 src/PIM/Conversion/ONNXToSpatial/SpatialLayoutCapabilities.cpp create mode 100644 src/PIM/Dialect/Spatial/SpatialLayoutInterface.td create mode 100644 src/PIM/Dialect/Spatial/SpatialTargetInfo.hpp delete mode 100644 src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MergeComputeNodesPass.cpp create mode 100644 src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.cpp create mode 100644 src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.hpp diff --git a/README.md b/README.md index f243e35..83d9ebe 100644 --- a/README.md +++ b/README.md @@ -52,12 +52,22 @@ ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> PIM artifacts `Patterns/{Math,NN,Tensor}` and currently cover Conv, Gemm, MatMul, elementwise Add/Mul/Div, ReduceMean, pooling, Relu, Sigmoid, Softmax, Concat, Gather, Reshape, Resize, and Split. + The compiler-layer target adapter supplies the target-neutral + `SpatialTargetInfo`. Layout-aware plan ops advertise typed alternatives + through the Spatial layout interface; the layout planner records the + selected layout and explicit materialization edges. `LowerSpatialPlans` + then pattern-lowers those selected plans. Contraction and Conv lowering + keep semantic problems, target-dependent plans, and IR materializers in + separate layers. -2. **Merge compute nodes** +2. **Merge, schedule, and realize Spatial communication** (`src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes`). - Builds a compute graph, schedules it with the PEFT scheduler, and materializes - the merge schedule into Spatial IR. Supporting scheduling code lives under - `MergeComputeNodes/Scheduling`. + `TrivialGraphComputeMerge` performs local graph merging, then + `ScheduleSpatialGraph` materializes scheduled computes and explicit deferred + communication. `VerifyScheduledSpatial` checks that intermediate contract; + `RealizeSpatialCommunication` resolves transfers and forwarding; and + `VerifyRealizedSpatial` checks the final scheduled graph. Supporting + scheduling code lives under `MergeComputeNodes/Scheduling`. 3. **Spatial -> Pim** (`src/PIM/Conversion/SpatialToPim`). Lowers Spatial operations to the `pim` dialect (`src/PIM/Dialect/Pim`), @@ -65,8 +75,12 @@ ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> PIM artifacts tensor materialization, and return-path normalization. 4. **Bufferization** (`src/PIM/Dialect/Pim/Transforms/Bufferization`). - Converts tensor-semantics PIM IR into memref-semantics PIM IR using MLIR's - bufferization interfaces. + `PimBufferizationPreparation` establishes writable destinations without + duplicating the one-shot copy analysis, `PimOneShotBufferization` runs + MLIR's one-shot analysis, + `PimMemoryNormalization` forwards/removes redundant copies and normalizes + addressable accesses, and `PimBufferizationVerification` checks tensor + absence, contiguity, and copy address spaces. 5. **PIM local-memory planning** (`src/PIM/Dialect/Pim/Transforms/LocalMemoryPlanning`). diff --git a/src/PIM/Compiler/PimCompilerUtils.cpp b/src/PIM/Compiler/PimCompilerUtils.cpp index 6e64791..69355aa 100644 --- a/src/PIM/Compiler/PimCompilerUtils.cpp +++ b/src/PIM/Compiler/PimCompilerUtils.cpp @@ -15,6 +15,8 @@ #include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp" #include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetInfo.hpp" +#include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/SchedulingTarget.hpp" #include "src/Accelerators/PIM/Pass/PIMPasses.h" #include "src/Compiler/CompilerPasses.hpp" @@ -80,6 +82,34 @@ spatial::SchedulingTarget getDefaultPimSchedulingTarget() { return target; } +spatial::ConvLoweringStrategy getSpatialConvLoweringStrategy(PimConvLoweringType strategy) { + switch (strategy) { + case PimConvLoweringAuto: return spatial::ConvLoweringStrategy::Auto; + case PimConvLoweringLegacy: return spatial::ConvLoweringStrategy::Legacy; + case PimConvLoweringDepthwise: return spatial::ConvLoweringStrategy::Depthwise; + case PimConvLoweringPackedIm2Col: return spatial::ConvLoweringStrategy::PackedIm2Col; + case PimConvLoweringStreamedPatch: return spatial::ConvLoweringStrategy::StreamedPatch; + case PimConvLoweringStreamedPacked: return spatial::ConvLoweringStrategy::StreamedPacked; + case PimConvLoweringOutputChannelTiled: return spatial::ConvLoweringStrategy::OutputChannelTiled; + case PimConvLoweringInputKTiled: return spatial::ConvLoweringStrategy::InputKTiled; + case PimConvLoweringTiled2D: return spatial::ConvLoweringStrategy::Tiled2D; + } + llvm_unreachable("unknown PIM Conv lowering strategy"); +} + +spatial::SpatialTargetInfo getPimSpatialTargetInfo(const spatial::SchedulingTarget& target) { + spatial::SpatialTargetInfo info; + info.matrixShape = {target.matrixRows, target.matrixColumns}; + info.matrixUnitsPerProcessor = target.residentWeightCapacity; + info.processorCount = target.processorCount; + info.vectorWidth = target.vectorWidth; + info.convIm2colMaxElements = pimConvIm2colMaxElements.getValue(); + info.convStreamChunkPositions = pimConvStreamChunkPositions.getValue(); + info.convLoweringStrategy = getSpatialConvLoweringStrategy(pimConvLowering.getValue()); + info.useExperimentalConvImplementation = useExperimentalConvImpl.getValue(); + return info; +} + const llvm::json::Object& requireObject(const llvm::json::Object& object, llvm::StringRef key, llvm::StringRef path) { @@ -293,12 +323,17 @@ void addPassesPim(OwningOpRef& module, if (pimEmissionTarget >= EmitSpatial) { spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget(); - pm.addPass(createONNXToSpatialPass()); - pm.addPass(createSpatialLayoutPlanningPass()); - pm.addPass(createLowerSpatialPlansPass()); + spatial::SpatialTargetInfo targetInfo = getPimSpatialTargetInfo(schedulingTarget); + pm.addPass(createONNXToSpatialPass(targetInfo)); + pm.addPass(createSpatialLayoutPlanningPass(targetInfo)); + pm.addPass(createLowerSpatialPlansPass(targetInfo)); pm.addPass(createTrivialGraphComputeMergePass( schedulingTarget.residentWeightCapacity)); - pm.addPass(createMergeComputeNodesPass(schedulingTarget)); + auto scheduledState = std::make_shared(); + pm.addPass(spatial::createScheduleSpatialGraphPass(schedulingTarget, scheduledState)); + pm.addPass(spatial::createVerifyScheduledSpatialPass(scheduledState)); + pm.addPass(spatial::createRealizeSpatialCommunicationPass(schedulingTarget, scheduledState)); + pm.addPass(spatial::createVerifyRealizedSpatialPass(scheduledState)); pm.addPass(createMessagePass("Onnx lowered to Spatial")); } @@ -308,7 +343,10 @@ void addPassesPim(OwningOpRef& module, } if (pimEmissionTarget >= EmitPimBufferized) { - pm.addPass(createPimBufferizationPass()); + pm.addPass(createPimBufferizationPreparationPass()); + pm.addPass(createPimOneShotBufferizationPass()); + pm.addPass(createPimMemoryNormalizationPass()); + pm.addPass(createPimBufferizationVerificationPass()); pm.addPass(createMessagePass("Pim bufferized")); } diff --git a/src/PIM/Conversion/ONNXToSpatial/CMakeLists.txt b/src/PIM/Conversion/ONNXToSpatial/CMakeLists.txt index 6d9395c..ef3bf7c 100644 --- a/src/PIM/Conversion/ONNXToSpatial/CMakeLists.txt +++ b/src/PIM/Conversion/ONNXToSpatial/CMakeLists.txt @@ -27,11 +27,14 @@ add_pim_library(OMONNXToSpatial Patterns/Tensor/Split.cpp Patterns/Tensor/Transpose.cpp ONNXToSpatialPass.cpp + SpatialLayoutCapabilities.cpp SpatialLayoutPlanningPass.cpp LowerSpatialPlansPass.cpp Common/AttributeUtils.cpp Common/BiasAddUtils.cpp Common/ComputeRegionBuilder.cpp + Common/ContractionMaterialization.cpp + Common/ContractionPlanning.cpp Common/MatrixProductLowering.cpp Common/RowStripLayoutUtils.cpp Common/ShapeTilingUtils.cpp @@ -46,8 +49,6 @@ add_pim_library(OMONNXToSpatial MLIRLinalgDialect MLIRSCFDialect MLIRTosaDialect - OMCompilerOptions - OMPimCompilerOptions OMONNXOps SpatialOps OMPimCommon diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.cpp index 93db327..5320441 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.cpp @@ -25,6 +25,9 @@ FailureOr createFragmentAssemblyBlueprint(Value physicalBatch, const int64_t laneCount = physicalType.getDimSize(0); if (laneCount <= 0) return emitError(loc, "fragment assembly requires at least one physical source slot"), failure(); + auto physicalLayoutValue = spatial::symbolizePhysicalLayout(physicalLayout); + if (!physicalLayoutValue) + return emitError(loc, "unknown physical layout for fragment assembly"), failure(); const int64_t fragmentElements = physicalType.getNumElements() / laneCount; SmallVector operandIndices(entries.size(), 0), sourceSlots, sourceOffsets, offsets, sizes, strides(entries.size() * rank, 1); @@ -48,9 +51,10 @@ FailureOr createFragmentAssemblyBlueprint(Value physicalBatch, llvm::append_range(sizes, entry.sizes); } auto blueprint = spatial::SpatBlueprintOp::create(rewriter, loc, logicalType, physicalBatch, ValueRange {}, - rewriter.getStringAttr("nchw"), rewriter.getStringAttr(physicalLayout), + spatial::getNCHWLayout(rewriter.getContext()), + spatial::PhysicalLayoutAttr::get(rewriter.getContext(), *physicalLayoutValue), rewriter.getDenseI64ArrayAttr(offsets), rewriter.getDenseI64ArrayAttr(sizes), - rewriter.getStringAttr(indexMap), rewriter.getStringAttr("fragment_assembly"), + rewriter.getStringAttr(indexMap), spatial::getFragmentAssemblyMode(rewriter.getContext()), rewriter.getDenseI64ArrayAttr(operandIndices), rewriter.getDenseI64ArrayAttr(sourceSlots), rewriter.getDenseI64ArrayAttr(sourceOffsets), rewriter.getDenseI64ArrayAttr(strides), rewriter.getStringAttr("disjoint"), rewriter.getStringAttr("complete")); diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.cpp new file mode 100644 index 0000000..dc11798 --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.cpp @@ -0,0 +1,39 @@ +#include "ContractionMaterialization.hpp" + +#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" +#include "MatrixProductLowering.hpp" + +namespace onnx_mlir { + +mlir::Value materializePaddedContractionInput( + mlir::Value input, + mlir::RankedTensorType paddedType, + mlir::PatternRewriter& rewriter, + mlir::Location loc) { + return createPaddedInputCompute(input, paddedType, rewriter, loc); +} + +mlir::FailureOr materializeTransposedContractionConstant( + mlir::Value input, + mlir::RankedTensorType resultType, + llvm::ArrayRef permutation, + mlir::PatternRewriter& rewriter, + mlir::Location loc) { + auto denseAttr = getHostConstDenseElementsAttr(input); + auto inputType = denseAttr ? mlir::dyn_cast(denseAttr.getType()) : nullptr; + if (!inputType || !inputType.hasStaticShape() || !resultType || !resultType.hasStaticShape() + || inputType.getRank() != resultType.getRank()) + return mlir::failure(); + + auto transposedAttr = transposeDenseElementsAttr(denseAttr, permutation); + if (mlir::failed(transposedAttr) || transposedAttr->getType() != resultType) + return mlir::failure(); + + return getOrCreateConstant(rewriter, + rewriter.getInsertionBlock()->getParentOp(), + *transposedAttr, + resultType); +} + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp new file mode 100644 index 0000000..5f92d3d --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp @@ -0,0 +1,23 @@ +#pragma once + +#include "llvm/ADT/ArrayRef.h" + +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/PatternMatch.h" + +namespace onnx_mlir { + +mlir::Value materializePaddedContractionInput( + mlir::Value input, + mlir::RankedTensorType paddedType, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + +mlir::FailureOr materializeTransposedContractionConstant( + mlir::Value input, + mlir::RankedTensorType resultType, + llvm::ArrayRef permutation, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.cpp new file mode 100644 index 0000000..e8bacf5 --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.cpp @@ -0,0 +1,68 @@ +#include "ContractionPlanning.hpp" + +#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" + +#include + +namespace onnx_mlir { + +namespace { + +static int64_t ceilDivide(int64_t value, int64_t divisor) { + return divisor == 0 ? 0 : (value + divisor - 1) / divisor; +} + +static llvm::SmallVector buildBatchMap( + llvm::ArrayRef sourceShape, + llvm::ArrayRef outputShape) { + llvm::SmallVector map(outputShape.size(), -1); + const int64_t offset = outputShape.size() - sourceShape.size(); + for (int64_t source = 0; source < static_cast(sourceShape.size()); ++source) { + const int64_t output = source + offset; + if (sourceShape[source] != 1) + map[output] = source; + } + return map; +} + +} // namespace + +ContractionPlan makeContractionPlan( + const ContractionProblem& problem, + const spatial::SpatialTargetInfo& target, + ContractionPlanKind kind, + int64_t laneCount, + int64_t fragmentRows) { + ContractionPlan plan; + plan.problem = problem; + plan.kind = kind; + plan.tileM = std::max(1, target.matrixShape.rows); + plan.tileK = std::max(1, target.matrixShape.rows); + plan.tileN = std::max(1, target.matrixShape.columns); + plan.reductionSlices = std::max(1, ceilDivide(problem.k, plan.tileK)); + plan.outputTiles = std::max(1, ceilDivide(problem.n, plan.tileN)); + plan.rowTiles = std::max(1, ceilDivide(problem.m, plan.tileM)); + plan.fragmentRows = std::max( + 1, fragmentRows != 0 ? fragmentRows : plan.tileM); + plan.lhsBatchMap = buildBatchMap(problem.lhsBatchShape, problem.outputBatchShape); + plan.rhsBatchMap = buildBatchMap(problem.rhsBatchShape, problem.outputBatchShape); + + if (laneCount != 0) + plan.laneCount = laneCount; + else if (kind == ContractionPlanKind::StaticTiled) + plan.laneCount = problem.batch * problem.m * plan.reductionSlices * plan.outputTiles; + else if (kind == ContractionPlanKind::GroupedRowDynamicVVD) + plan.laneCount = problem.batch * ceilDivide(problem.m, plan.fragmentRows); + else + plan.laneCount = problem.batch * problem.m * problem.n; + + plan.expectedMvmCount = kind == ContractionPlanKind::StaticTiled ? plan.laneCount : 0; + plan.expectedVvdCount = kind == ContractionPlanKind::StaticTiled ? 0 : plan.laneCount; + plan.expectedVectorCount = plan.laneCount * plan.reductionSlices; + if (problem.resultElementType && problem.n > 0) + plan.physicalFragmentType = mlir::RankedTensorType::get( + {plan.fragmentRows, problem.n}, problem.resultElementType); + return plan; +} + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp new file mode 100644 index 0000000..ef47833 --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp @@ -0,0 +1,41 @@ +#pragma once + +#include "ContractionProblem.hpp" + +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetInfo.hpp" + +namespace onnx_mlir { + +enum class ContractionPlanKind { + StaticTiled, + BatchedDynamicVVD, + GroupedRowDynamicVVD, +}; + +struct ContractionPlan { + ContractionProblem problem; + ContractionPlanKind kind = ContractionPlanKind::StaticTiled; + int64_t tileM = 1; + int64_t tileK = 1; + int64_t tileN = 1; + int64_t fragmentRows = 1; + int64_t reductionSlices = 1; + int64_t outputTiles = 1; + int64_t rowTiles = 1; + int64_t laneCount = 0; + int64_t expectedMvmCount = 0; + int64_t expectedVvdCount = 0; + int64_t expectedVectorCount = 0; + llvm::SmallVector lhsBatchMap; + llvm::SmallVector rhsBatchMap; + mlir::RankedTensorType physicalFragmentType; +}; + +ContractionPlan makeContractionPlan( + const ContractionProblem& problem, + const spatial::SpatialTargetInfo& target, + ContractionPlanKind kind, + int64_t laneCount = 0, + int64_t fragmentRows = 0); + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp new file mode 100644 index 0000000..dd1b628 --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp @@ -0,0 +1,35 @@ +#pragma once + +#include "mlir/IR/BuiltinTypes.h" + +#include "llvm/ADT/SmallVector.h" + +#include + +namespace onnx_mlir { + +enum class ContractionOrigin { Gemm, MatMul }; + +struct ContractionProblem { + llvm::SmallVector lhsBatchShape; + llvm::SmallVector rhsBatchShape; + llvm::SmallVector outputBatchShape; + int64_t lhsBatch = 1; + int64_t rhsBatch = 1; + int64_t batch = 1; + int64_t m = 0; + int64_t k = 0; + int64_t n = 0; + ContractionOrigin origin = ContractionOrigin::MatMul; + mlir::Type lhsElementType; + mlir::Type rhsElementType; + mlir::Type resultElementType; + bool lhsTransposed = false; + bool rhsTransposed = false; + bool lhsWasVector = false; + bool rhsWasVector = false; + float alpha = 1.0f; + float beta = 1.0f; +}; + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.cpp index 4b0b00d..668bf7c 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.cpp @@ -1,15 +1,70 @@ #include "MatrixProductLowering.hpp" #include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/Linalg/IR/Linalg.h" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.hpp" +#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" using namespace mlir; namespace onnx_mlir { +static bool isInsideSpatialCompute(Operation* op) { + for (Operation* parent = op; parent; parent = parent->getParentOp()) + if (spatial::isAnySpatialComputeLike(parent)) + return true; + return false; +} + +static Value buildLinalgTranspose(Value value, + RankedTensorType resultType, + ArrayRef permutation, + PatternRewriter& rewriter, + Location loc) { + Value init = tensor::EmptyOp::create( + rewriter, loc, resultType.getShape(), resultType.getElementType()); + return linalg::TransposeOp::create( + rewriter, loc, value, init, permutation).getResult()[0]; +} + +static Value materializeConstantTranspose(Value value, + RankedTensorType resultType, + ArrayRef permutation, + PatternRewriter& rewriter) { + auto denseAttr = getHostConstDenseElementsAttr(value); + if (!denseAttr) + return {}; + auto transposedAttr = transposeDenseElementsAttr(denseAttr, permutation); + if (failed(transposedAttr) || transposedAttr->getType() != resultType) + return {}; + return getOrCreateConstant( + rewriter, rewriter.getInsertionBlock()->getParentOp(), *transposedAttr, resultType); +} + +Value createLinalgTranspose(Value value, + RankedTensorType resultType, + ArrayRef permutation, + PatternRewriter& rewriter, + Location loc) { + if (Value constant = materializeConstantTranspose(value, resultType, permutation, rewriter)) + return constant; + + if (isInsideSpatialCompute(rewriter.getInsertionBlock()->getParentOp())) + return buildLinalgTranspose(value, resultType, permutation, rewriter, loc); + + auto compute = createSpatCompute<1>( + rewriter, loc, TypeRange {resultType}, {}, ValueRange {value}, + [&](Value input) { + spatial::SpatYieldOp::create( + rewriter, loc, buildLinalgTranspose(input, resultType, permutation, rewriter, loc)); + }); + return compute.getResult(0); +} + Value createZeroPaddedTensor(Value value, RankedTensorType resultType, PatternRewriter& rewriter, Location loc) { auto sourceType = cast(value.getType()); SmallVector lowPads(sourceType.getRank(), rewriter.getIndexAttr(0)); diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp index 45eaff2..1125437 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp @@ -5,8 +5,16 @@ #include "mlir/IR/Value.h" #include "mlir/Transforms/DialectConversion.h" +#include "llvm/ADT/ArrayRef.h" + namespace onnx_mlir { +mlir::Value createLinalgTranspose(mlir::Value value, + mlir::RankedTensorType resultType, + llvm::ArrayRef permutation, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + mlir::Value createZeroPaddedTensor(mlir::Value value, mlir::RankedTensorType resultType, mlir::PatternRewriter& rewriter, diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp index a39de5d..180f972 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp @@ -5,9 +5,9 @@ #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/ComputeRegionBuilder.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" -#include "src/Dialect/ONNX/ONNXOps.hpp" #include @@ -33,6 +33,16 @@ FailureOr describeRowStripPhysicalValue(Value storage, Ra tilesPerRow}; } +FailureOr getRowStripPhysicalValue(Value value) { + auto blueprint = value.getDefiningOp(); + auto logicalType = dyn_cast(value.getType()); + if (!blueprint || !logicalType || blueprint.getOutput() != value + || blueprint.getPhysicalLayout() != spatial::PhysicalLayout::NHWCRowStrip + || !spatial::isPhysicalView(blueprint.getMode())) + return failure(); + return describeRowStripPhysicalValue(blueprint.getInput(), logicalType); +} + RankedTensorType getRowStripFragmentType(RankedTensorType logicalType) { return RankedTensorType::get({logicalType.getDimSize(0), 1, logicalType.getDimSize(3), logicalType.getDimSize(1)}, @@ -144,6 +154,35 @@ FailureOr createRowStripStorageFromRows(Value rows, return batchOp->getResult(0); } +FailureOr createRowStripStorageBlueprint(Value storage, + RankedTensorType logicalType, + PatternRewriter& rewriter, + Location loc) { + FailureOr value = describeRowStripPhysicalValue(storage, logicalType); + if (failed(value)) + return failure(); + + auto blueprint = spatial::SpatBlueprintOp::create( + rewriter, + loc, + logicalType, + storage, + ValueRange {}, + spatial::getNCHWLayout(rewriter.getContext()), + spatial::getNHWCRowStripLayout(rewriter.getContext()), + rewriter.getDenseI64ArrayAttr({}), + rewriter.getDenseI64ArrayAttr({}), + rewriter.getStringAttr(kRowStripIndexMap), + spatial::getPhysicalViewMode(rewriter.getContext()), + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr); + return blueprint.getOutput(); +} + FailureOr createRowStripAssemblyBlueprint(const RowStripPhysicalValue& value, PatternRewriter& rewriter, Location loc) { @@ -160,8 +199,8 @@ FailureOr createRowStripAssemblyBlueprint(const RowStripPhysicalValue& va rewriter, loc, args.inputs.front(), args.lane, value.fragmentType); if (failed(fragment)) return failure(); - Value nchw = ONNXTransposeOp::create( - rewriter, loc, nchwFragmentType, *fragment, rewriter.getI64ArrayAttr({0, 3, 1, 2})); + Value nchw = createLinalgTranspose( + *fragment, nchwFragmentType, {0, 3, 1, 2}, rewriter, loc); publishGraphBatchPhysicalFragment(rewriter, loc, nchw, args.outputs.front(), args.lane); return success(); }); @@ -176,7 +215,7 @@ FailureOr createRowStripAssemblyBlueprint(const RowStripPhysicalValue& va {1, std::min(tileChannels, value.logicalType.getDimSize(1) - channelOffset), 1, value.logicalType.getDimSize(3)}}); } - return createFragmentAssemblyBlueprint(transposed->getResult(0), value.logicalType, entries, "nhwc_row_strip", + return createFragmentAssemblyBlueprint(transposed->getResult(0), value.logicalType, entries, "dense_nchw", kRowStripIndexMap, rewriter, loc); } diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp index a7f75e9..853471e 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp @@ -6,6 +6,12 @@ namespace onnx_mlir { +namespace spatial { +class SpatBlueprintOp; +class SpatGraphCompute; +struct SpatialTargetInfo; +} // namespace spatial + inline constexpr llvm::StringLiteral kRowStripIndexMap = "nhwc_row_strip_fragments"; struct RowStripPhysicalValue { @@ -18,6 +24,8 @@ struct RowStripPhysicalValue { mlir::FailureOr describeRowStripPhysicalValue(mlir::Value storage, mlir::RankedTensorType logicalType); +mlir::FailureOr getRowStripPhysicalValue(mlir::Value value); + std::pair, llvm::SmallVector> buildRowStripMetadata(mlir::RankedTensorType type); @@ -53,6 +61,11 @@ mlir::FailureOr createRowStripStorageFromRows(mlir::Value rows, mlir::PatternRewriter& rewriter, mlir::Location loc); +mlir::FailureOr createRowStripStorageBlueprint(mlir::Value storage, + mlir::RankedTensorType logicalType, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + mlir::FailureOr createRowStripAssemblyBlueprint(const RowStripPhysicalValue& value, mlir::PatternRewriter& rewriter, mlir::Location loc); @@ -80,4 +93,14 @@ mlir::FailureOr applyRowStripConcat(llvm::ArrayRef> -sliceVectorPerCrossbarPerCore(const Value& vectorToSlice, PatternRewriter& rewriter, Location loc) { - SmallVector slices = sliceVector(vectorToSlice, crossbarSize, rewriter, loc); +sliceVectorPerCrossbarPerCore(const Value& vectorToSlice, + PatternRewriter& rewriter, + Location loc, + const spatial::SpatialTargetInfo& target) { + SmallVector slices = sliceVector( + vectorToSlice, static_cast(target.matrixShape.rows), rewriter, loc); DenseMap> slicesPerCore; for (size_t sliceId = 0; sliceId < slices.size(); sliceId++) { - size_t coreId = sliceId / crossbarCountInCore; + size_t coreId = sliceId / target.matrixUnitsPerProcessor; slicesPerCore[coreId].push_back(slices[sliceId]); } return slicesPerCore; diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/ShapeTilingUtils.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/ShapeTilingUtils.hpp index 4fb9021..714ad50 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/ShapeTilingUtils.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/ShapeTilingUtils.hpp @@ -7,6 +7,7 @@ #include "llvm/ADT/SmallVector.h" #include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetInfo.hpp" namespace onnx_mlir { @@ -26,6 +27,9 @@ llvm::SmallVector sliceVector(const mlir::Value& vectorToSlice, /// Partitions one logical vector into per-core crossbar-sized slices using the /// current PIM target geometry. llvm::DenseMap> sliceVectorPerCrossbarPerCore( - const mlir::Value& vectorToSlice, mlir::PatternRewriter& rewriter, mlir::Location loc); + const mlir::Value& vectorToSlice, + mlir::PatternRewriter& rewriter, + mlir::Location loc, + const spatial::SpatialTargetInfo& target); } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp index 74bb30b..25a35b2 100644 --- a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp @@ -5,52 +5,104 @@ #include "mlir/Dialect/SCF/IR/SCF.h" #include "mlir/Dialect/Tensor/IR/Tensor.h" #include "mlir/Pass/Pass.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" #include "mlir/Transforms/DialectConversion.h" -#include "llvm/ADT/DenseMap.h" -#include "llvm/ADT/SmallPtrSet.h" - #include "Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp" #include "mlir/Transforms/Passes.h" #include "src/Accelerators/PIM/Common/PimCommon.hpp" #include "src/Accelerators/PIM/Common/Support/DebugDump.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.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/SpatialDataflowCsvExporter.hpp" #include "src/Accelerators/PIM/Pass/PIMPasses.h" -#include "src/Dialect/ONNX/ONNXOps.hpp" using namespace mlir; namespace onnx_mlir { namespace { -static constexpr StringLiteral kDenseLayout = "dense_nchw"; -static constexpr StringLiteral kRowStripLayout = "nhwc_row_strip"; - -static FailureOr getRowStripValue(llvm::DenseMap& rowStripValues, - Value value) { - auto it = rowStripValues.find(value); - if (it == rowStripValues.end()) - return failure(); - return it->second; +static FailureOr getRowStripValue(Value value) { + return getRowStripPhysicalValue(value); } -static FailureOr buildRowStripValue(spatial::SpatBlueprintOp blueprint, - Value storage) { - auto logicalType = dyn_cast(blueprint.getOutput().getType()); +static FailureOr publishRowStripValue(Operation* planOp, + Value storage, + PatternRewriter& rewriter) { + auto logicalType = dyn_cast(planOp->getResult(0).getType()); if (!logicalType) - return blueprint.emitOpError("requires ranked logical output type"), failure(); - if (blueprint.getIndexMap() != kRowStripIndexMap) - return blueprint.emitOpError("requires the canonical row-strip index map"), failure(); + return planOp->emitOpError("requires ranked logical output type"), failure(); FailureOr value = describeRowStripPhysicalValue(storage, logicalType); if (failed(value)) - return blueprint.emitOpError("requires physical row-strip fragment storage"), failure(); - return *value; + return planOp->emitOpError("lowering produced invalid row-strip physical storage"), failure(); + FailureOr blueprint = createRowStripStorageBlueprint( + storage, logicalType, rewriter, planOp->getLoc()); + if (failed(blueprint)) + return planOp->emitOpError("failed to create row-strip storage Blueprint"), failure(); + rewriter.replaceOp(planOp, *blueprint); + return *blueprint; +} + +static bool isRowStripSelected(Operation* op) { + auto selected = spatial::getSelectedPhysicalLayout(op); + return selected && *selected == spatial::PhysicalLayout::NHWCRowStrip; +} + +static bool isDenseSelected(Operation* op) { + auto selected = spatial::getSelectedPhysicalLayout(op); + return selected && *selected == spatial::PhysicalLayout::DenseNCHW; +} + +static spatial::PhysicalLayout getKnownPhysicalLayout(Value value) { + if (auto materialize = value.getDefiningOp()) + return materialize.getTargetPhysicalLayout(); + if (auto blueprint = value.getDefiningOp()) + return blueprint.getPhysicalLayout(); + if (Operation* producer = value.getDefiningOp()) { + if (auto selected = spatial::getSelectedPhysicalLayout(producer)) + return *selected; + } + return spatial::PhysicalLayout::DenseNCHW; +} + +static LogicalResult verifySelectedLayouts( + func::FuncOp funcOp, const spatial::SpatialTargetInfo& target) { + LogicalResult result = success(); + funcOp.walk([&](Operation* op) { + auto capability = dyn_cast(op); + if (!capability) + return; + auto selected = spatial::getSelectedPhysicalLayout(op); + if (!selected) { + op->emitOpError("requires a selected physical layout from SpatialLayoutPlanning"); + result = failure(); + return; + } + if (*selected != spatial::PhysicalLayout::DenseNCHW + && *selected != spatial::PhysicalLayout::NHWCRowStrip) { + op->emitOpError("has an unsupported selected physical layout"); + result = failure(); + return; + } + SmallVector operandLayouts; + operandLayouts.reserve(op->getNumOperands()); + for (Value operand : op->getOperands()) + operandLayouts.push_back(getKnownPhysicalLayout(operand)); + auto alternatives = capability.getLayoutAlternatives(target, operandLayouts); + if (llvm::none_of(alternatives, [&](const spatial::LayoutAlternative& alternative) { + return alternative.resultLayout == *selected + && alternative.operandLayouts == operandLayouts; + })) { + op->emitOpError("selected physical layout is not lowerable for its explicit operand layouts"); + result = failure(); + } + }); + return result; } static FailureOr @@ -92,6 +144,33 @@ materializeRowStripToDense(const RowStripPhysicalValue& rowStripValue, Location return createRowStripAssemblyBlueprint(rowStripValue, rewriter, loc); } +static FailureOr materializeDenseToRowStrip( + Value input, RankedTensorType logicalType, Location loc, PatternRewriter& rewriter) { + if (!logicalType || !logicalType.hasStaticShape() || logicalType.getRank() != 4 + || logicalType.getDimSize(0) != 1) + return failure(); + auto nhwcType = RankedTensorType::get( + {1, logicalType.getDimSize(2), logicalType.getDimSize(3), logicalType.getDimSize(1)}, + logicalType.getElementType(), logicalType.getEncoding()); + auto rowsType = RankedTensorType::get( + {logicalType.getDimSize(2) * logicalType.getDimSize(3), logicalType.getDimSize(1)}, + logicalType.getElementType(), logicalType.getEncoding()); + auto rowsCompute = createSpatCompute<1>( + rewriter, loc, rowsType, {}, input, [&](Value denseInput) { + Value nhwc = createLinalgTranspose( + denseInput, nhwcType, {0, 2, 3, 1}, rewriter, loc); + Value rows = tensor::CollapseShapeOp::create( + rewriter, loc, rowsType, nhwc, + SmallVector {{0, 1, 2}, {3}}); + spatial::SpatYieldOp::create(rewriter, loc, rows); + }); + Value rows = rowsCompute->getResult(0); + FailureOr storage = createRowStripStorageFromRows(rows, logicalType, rewriter, loc); + if (failed(storage)) + return failure(); + return createRowStripStorageBlueprint(*storage, logicalType, rewriter, loc); +} + static FailureOr lowerDenseBatchBiasAdd(Value input, Value bias, RankedTensorType resultType, PatternRewriter& rewriter, Location loc) { auto producer = input.getDefiningOp(); @@ -143,98 +222,520 @@ static FailureOr lowerDenseBatchBiasAdd(Value input, Value bias, RankedTe return batch->getResult(0); } -static LogicalResult lowerAddPlan(spatial::SpatAddPlanOp planOp, - llvm::DenseMap& rowStripValues, - llvm::SmallPtrSetImpl& eraseAfterLowering, - PatternRewriter& rewriter) { - FailureOr lhs = getRowStripValue(rowStripValues, planOp.getLhs()); - FailureOr rhs = getRowStripValue(rowStripValues, planOp.getRhs()); - if (succeeded(lhs) && succeeded(rhs)) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) - return planOp.emitOpError("row-strip add plan requires a row-strip blueprint result"); +struct LowerDenseReluPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(spatial::SpatReluPlanOp planOp, + PatternRewriter& rewriter) const override { + auto selected = spatial::getSelectedPhysicalLayout(planOp.getOperation()); + if (!selected || *selected != spatial::PhysicalLayout::DenseNCHW) + return failure(); + + auto computeOp = createSpatCompute<1>( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) { + auto relu = spatial::SpatReluOp::create(rewriter, planOp.getLoc(), planOp.getOutput().getType(), x); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), relu.getResult()); + }); + rewriter.replaceOp(planOp, computeOp.getResults()); + return success(); + } +}; + +struct LowerDenseSiluPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatSiluPlanOp planOp, + PatternRewriter& rewriter) const override { + auto selected = spatial::getSelectedPhysicalLayout(planOp.getOperation()); + if (!selected || *selected != spatial::PhysicalLayout::DenseNCHW) + return failure(); + + auto computeOp = createSpatCompute<1>( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) { + Value sigmoid = spatial::SpatSigmoidOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), x).getResult(); + Value silu = spatial::SpatVMulOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, sigmoid).getResult(); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), silu); + }); + rewriter.replaceOp(planOp, computeOp.getResults()); + return success(); + } +}; + +struct LowerDenseResizePlan final : OpRewritePattern { + explicit LowerDenseResizePlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatResizeNearestPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + FailureOr lowered = lowerSelectedResizeNearestPlan(planOp, std::nullopt, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected dense nearest Resize plan"); + rewriter.replaceOp(planOp, *lowered); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerDenseBiasAddPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatBiasAddPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + auto resultType = dyn_cast(planOp.getOutput().getType()); + if (!resultType) + return planOp.emitOpError("requires ranked output type"); + + FailureOr denseBias = materializeDenseBiasAddTensor( + planOp.getBias(), resultType, rewriter, planOp.getLoc()); + if (failed(denseBias)) + return planOp.emitOpError("failed to materialize dense Conv-style bias"); + if (planOp.getInput().getDefiningOp()) { + FailureOr lowered = lowerDenseBatchBiasAdd( + planOp.getInput(), *denseBias, resultType, rewriter, planOp.getLoc()); + if (succeeded(lowered)) { + rewriter.replaceOp(planOp, *lowered); + return success(); + } + } + auto computeOp = createSpatCompute<2>( + rewriter, + planOp.getLoc(), + planOp.getOutput().getType(), + {}, + ValueRange {planOp.getInput(), *denseBias}, + [&](Value x, Value y) { + auto added = spatial::SpatVAddOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, y); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added.getResult()); + }); + rewriter.replaceOp(planOp, computeOp.getResults()); + return success(); + } +}; + +struct LowerDenseAddPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatAddPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + auto compute = createSpatCompute<2>( + rewriter, + planOp.getLoc(), + planOp.getOutput().getType(), + {}, + ValueRange {planOp.getLhs(), planOp.getRhs()}, + [&](Value lhsValue, Value rhsValue) { + Value added = spatial::SpatVAddOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), lhsValue, rhsValue); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added); + }); + rewriter.replaceOp(planOp, compute.getResults()); + return success(); + } +}; + +struct LowerDenseConcatPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatConcatPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + auto compute = createSpatCompute( + rewriter, + planOp.getLoc(), + TypeRange {planOp.getOutput().getType()}, + {}, + planOp.getInputs(), + [&](ValueRange values) { + Value concatenated = spatial::SpatConcatOp::create( + rewriter, + planOp.getLoc(), + planOp.getOutput().getType(), + rewriter.getI64IntegerAttr(planOp.getAxis()), + values); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), concatenated); + }); + rewriter.replaceOp(planOp, compute.getResults()); + return success(); + } +}; + +static LogicalResult lowerAddPlan(spatial::SpatAddPlanOp planOp, + PatternRewriter& rewriter) { + FailureOr lhs = getRowStripValue(planOp.getLhs()); + FailureOr rhs = getRowStripValue(planOp.getRhs()); + if (isRowStripSelected(planOp.getOperation()) && failed(lhs)) { + if (getKnownPhysicalLayout(planOp.getLhs()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip Add plan requires row-strip inputs"); + } + if (isRowStripSelected(planOp.getOperation()) && failed(rhs)) { + if (getKnownPhysicalLayout(planOp.getRhs()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip Add plan requires row-strip inputs"); + } + if (isRowStripSelected(planOp.getOperation())) { rewriter.setInsertionPoint(planOp); FailureOr lowered = lowerRowStripAdd(*lhs, *rhs, planOp, rewriter); if (failed(lowered)) return planOp.emitOpError("failed to lower selected row-strip Spatial add plan"); - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) return failure(); - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); return success(); } - - rewriter.setInsertionPoint(planOp); - auto compute = createSpatCompute<2>(rewriter, - planOp.getLoc(), - planOp.getOutput().getType(), - {}, - ValueRange {planOp.getLhs(), planOp.getRhs()}, - [&](Value lhsValue, Value rhsValue) { - Value added = spatial::SpatVAddOp::create( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), lhsValue, rhsValue); - spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added); - }); - rewriter.replaceOp(planOp, compute.getResults()); - return success(); + return planOp.emitOpError("dense Add plan was not lowered by the selected-plan patterns"); } static LogicalResult lowerConcatPlan(spatial::SpatConcatPlanOp planOp, - llvm::DenseMap& rowStripValues, - llvm::SmallPtrSetImpl& eraseAfterLowering, PatternRewriter& rewriter) { SmallVector inputs; for (Value input : planOp.getInputs()) { - FailureOr physical = getRowStripValue(rowStripValues, input); + FailureOr physical = getRowStripValue(input); if (failed(physical)) { inputs.clear(); break; } inputs.push_back(*physical); } - if (inputs.size() == planOp.getInputs().size()) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) - return planOp.emitOpError("row-strip concat plan requires a row-strip blueprint result"); + if (isRowStripSelected(planOp.getOperation()) && inputs.size() != planOp.getInputs().size()) { + if (llvm::any_of(planOp.getInputs(), [](Value input) { + return getKnownPhysicalLayout(input) == spatial::PhysicalLayout::NHWCRowStrip; + })) + return failure(); + return planOp.emitOpError("selected row-strip Concat plan requires row-strip inputs"); + } + if (isRowStripSelected(planOp.getOperation())) { rewriter.setInsertionPoint(planOp); FailureOr lowered = lowerRowStripConcat(inputs, planOp, rewriter); if (failed(lowered)) return planOp.emitOpError("failed to lower selected row-strip Spatial concat plan"); - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } + return planOp.emitOpError("dense Concat plan was not lowered by the selected-plan patterns"); +} + +struct LowerSelectedConvPlan final : OpRewritePattern { + explicit LowerSelectedConvPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatConv2DPlanOp planOp, + PatternRewriter& rewriter) const override { + if (isDenseSelected(planOp.getOperation())) { + FailureOr lowered = lowerSelectedConv2DPlan( + planOp, std::nullopt, /*emitRowStripLayout=*/false, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected dense Spatial Conv plan"); + rewriter.replaceOp(planOp, *lowered); + return success(); + } + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + + FailureOr rowStripInput = getRowStripValue(planOp.getInput()); + if (failed(rowStripInput) + && getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + std::optional physicalInput; + if (succeeded(rowStripInput)) + physicalInput = rowStripInput->storage; + FailureOr lowered = lowerSelectedConv2DPlan( + planOp, physicalInput, /*emitRowStripLayout=*/true, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial Conv plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) return failure(); - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); return success(); } - rewriter.setInsertionPoint(planOp); - auto compute = createSpatCompute( - rewriter, - planOp.getLoc(), - TypeRange {planOp.getOutput().getType()}, - {}, - planOp.getInputs(), - [&](ValueRange values) { - Value concatenated = spatial::SpatConcatOp::create( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), rewriter.getI64IntegerAttr(planOp.getAxis()), values); - spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), concatenated); - }); - rewriter.replaceOp(planOp, compute.getResults()); - return success(); -} + const spatial::SpatialTargetInfo& target; +}; + +struct LowerRowStripReluPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatReluPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input)) { + if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip ReLU plan requires a row-strip input"); + } + FailureOr lowered = lowerRowStripRelu(*input, planOp, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial ReLU plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } +}; + +struct LowerRowStripSiluPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatSiluPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input)) { + if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip SiLU plan requires a row-strip input"); + } + FailureOr lowered = lowerRowStripSilu(*input, planOp, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial SiLU plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } +}; + +struct LowerRowStripResizePlan final : OpRewritePattern { + explicit LowerRowStripResizePlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatResizeNearestPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input)) { + if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip Resize plan requires a row-strip input"); + } + FailureOr lowered = lowerSelectedResizeNearestPlan(planOp, input->storage, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Resize plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerDenseMaxPoolPlan final : OpRewritePattern { + explicit LowerDenseMaxPoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatMaxPool2DPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + FailureOr lowered = lowerDenseMaxPool2DPlan(planOp, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected dense Spatial MaxPool plan"); + rewriter.replaceOp(planOp, *lowered); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerRowStripMaxPoolPlan final : OpRewritePattern { + explicit LowerRowStripMaxPoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatMaxPool2DPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input) + && getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + std::optional physicalInput; + if (succeeded(input)) + physicalInput = input->storage; + FailureOr lowered = lowerSelectedMaxPool2DPlan(planOp, physicalInput, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial MaxPool plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerRowStripGlobalAveragePoolPlan + final : OpRewritePattern { + explicit LowerRowStripGlobalAveragePoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatGlobalAveragePoolPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input) + && getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + std::optional physicalInput; + if (succeeded(input)) + physicalInput = input->storage; + FailureOr lowered = lowerSelectedGlobalAveragePoolPlan(planOp, physicalInput, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial global AveragePool plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerDenseGlobalAveragePoolPlan + final : OpRewritePattern { + explicit LowerDenseGlobalAveragePoolPlan(MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatGlobalAveragePoolPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isDenseSelected(planOp.getOperation())) + return failure(); + FailureOr lowered = lowerDenseGlobalAveragePoolPlan(planOp, target, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected dense Spatial global AveragePool plan"); + rewriter.replaceOp(planOp, *lowered); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; + +struct LowerRowStripBiasAddPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatBiasAddPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + FailureOr input = getRowStripValue(planOp.getInput()); + if (failed(input)) { + if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip) + return failure(); + return planOp.emitOpError("selected row-strip bias_add plan requires a row-strip input"); + } + FailureOr lowered = lowerRowStripBiasAdd(*input, planOp, rewriter); + if (failed(lowered)) + return planOp.emitOpError("failed to lower selected row-strip Spatial bias_add plan"); + if (failed(publishRowStripValue(planOp, *lowered, rewriter))) + return failure(); + return success(); + } +}; + +struct LowerRowStripAddPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatAddPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + return lowerAddPlan(planOp, rewriter); + } +}; + +struct LowerRowStripConcatPlan final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatConcatPlanOp planOp, + PatternRewriter& rewriter) const override { + if (!isRowStripSelected(planOp.getOperation())) + return failure(); + return lowerConcatPlan(planOp, rewriter); + } +}; + +struct LowerMaterializeLayout final + : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(spatial::SpatMaterializeLayoutOp materializeOp, + PatternRewriter& rewriter) const override { + auto source = materializeOp.getSourcePhysicalLayout(); + auto target = materializeOp.getTargetPhysicalLayout(); + if (source == spatial::PhysicalLayout::DenseNCHW + && target == spatial::PhysicalLayout::DenseNCHW) { + rewriter.replaceOp(materializeOp, materializeOp.getInput()); + return success(); + } + if (source == spatial::PhysicalLayout::DenseNCHW + && target == spatial::PhysicalLayout::NHWCRowStrip) { + auto logicalType = dyn_cast(materializeOp.getInput().getType()); + if (!logicalType) + return materializeOp.emitOpError("requires a ranked dense input"), failure(); + FailureOr rowStrip = materializeDenseToRowStrip( + materializeOp.getInput(), logicalType, materializeOp.getLoc(), rewriter); + if (failed(rowStrip)) + return materializeOp.emitOpError( + "failed to materialize dense NCHW storage to row-strip layout"), failure(); + rewriter.replaceOp(materializeOp, *rowStrip); + return success(); + } + if (source != spatial::PhysicalLayout::NHWCRowStrip + || target != spatial::PhysicalLayout::DenseNCHW) + return materializeOp.emitOpError( + "unsupported Spatial layout materialization direction"), failure(); + auto inputType = dyn_cast(materializeOp.getInput().getType()); + if (!inputType) + return materializeOp.emitOpError("requires a ranked row-strip input"), failure(); + FailureOr rowStripValue = + getRowStripValue(materializeOp.getInput()); + if (failed(rowStripValue)) + return materializeOp.emitOpError( + "requires an explicitly defining row-strip physical value"), failure(); + FailureOr dense = materializeRowStripToDense( + *rowStripValue, materializeOp.getLoc(), rewriter); + if (failed(dense)) + return materializeOp.emitOpError( + "failed to materialize row-strip storage to dense NCHW"), failure(); + rewriter.replaceOp(materializeOp, *dense); + return success(); + } +}; + +struct LowerRowStripFlatten final + : OpRewritePattern { + explicit LowerRowStripFlatten(MLIRContext* context, + const spatial::SpatialTargetInfo& target) + : OpRewritePattern(context), target(target) {} + + LogicalResult matchAndRewrite(spatial::SpatGraphCompute flattenOp, + PatternRewriter& rewriter) const override { + if (flattenOp.getInputs().size() != 1) + return failure(); + FailureOr input = + getRowStripValue(flattenOp.getInputs().front()); + if (failed(input) || failed(canLowerFlattenFromRowStrip(flattenOp, target))) + return failure(); + if (failed(lowerFlattenFromRowStrip(*input, flattenOp, target, rewriter))) + return flattenOp.emitOpError( + "failed to preserve row-strip layout through Flatten"), failure(); + return success(); + } + + const spatial::SpatialTargetInfo& target; +}; struct LowerSpatialPlansPass final : PassWrapper> { MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerSpatialPlansPass) @@ -242,8 +743,17 @@ struct LowerSpatialPlansPass final : PassWrapper rowStripValues; - llvm::SmallPtrSet eraseAfterLowering; auto verifyLogicalPhase = [&](StringRef stage) -> bool { if (succeeded(verifyLogicalSpatialGraphInvariants(*entryFunc))) return true; @@ -265,469 +773,80 @@ struct LowerSpatialPlansPass final : PassWrapper(&op)) { - FailureOr rowStripInput = getRowStripValue(rowStripValues, planOp.getInput()); - auto rowStripBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (rowStripBlueprint != planOp.getResult().getUsers().end()) { - rewriter.setInsertionPoint(planOp); - std::optional physicalInput; - if (succeeded(rowStripInput)) - physicalInput = rowStripInput->storage; - FailureOr lowered = lowerSelectedConv2DPlan( - planOp, - physicalInput, - /*emitRowStripLayout=*/true, - rewriter); - if (failed(lowered)) { - auto diagnostic = planOp.emitOpError("failed to lower selected row-strip Spatial Conv plan with input "); - diagnostic << planOp.getInput().getType() << " and output " << planOp.getResult().getType(); - if (physicalInput) - diagnostic << " from physical storage " << physicalInput->getType(); - signalPassFailure(); - return; - } - auto blueprint = cast(*rowStripBlueprint); - FailureOr rowStripValue = buildRowStripValue(blueprint, *lowered); - if (failed(rowStripValue)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *rowStripValue; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - rewriter.setInsertionPoint(planOp); - FailureOr lowered = - lowerSelectedConv2DPlan(planOp, std::nullopt, /*emitRowStripLayout=*/false, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected Spatial Conv plan"); - signalPassFailure(); - return; - } - rewriter.replaceOp(planOp, *lowered); - continue; - } - - if (auto planOp = dyn_cast(&op)) { - if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("row-strip Relu plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - - FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - FailureOr lowered = lowerRowStripRelu(*input, planOp, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected row-strip Spatial Relu plan"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - - rewriter.setInsertionPoint(planOp); - auto computeOp = createSpatCompute<1>( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) { - auto relu = spatial::SpatReluOp::create(rewriter, planOp.getLoc(), planOp.getOutput().getType(), x); - spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), relu.getResult()); - }); - rewriter.replaceOp(planOp, computeOp.getResults()); - continue; - } - - if (auto planOp = dyn_cast(&op)) { - if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("row-strip SiLU plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - - FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - FailureOr lowered = lowerRowStripSilu(*input, planOp, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected row-strip Spatial SiLU plan"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - - rewriter.setInsertionPoint(planOp); - auto computeOp = createSpatCompute<1>( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) { - Value sigmoid = spatial::SpatSigmoidOp::create( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), x).getResult(); - Value silu = spatial::SpatVMulOp::create( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, sigmoid).getResult(); - spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), silu); - }); - rewriter.replaceOp(planOp, computeOp.getResults()); - continue; - } - if (auto planOp = dyn_cast(&op)) { - FailureOr input = - getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - auto lowered = lowerSelectedResizeNearestPlan( - planOp, succeeded(input) ? std::optional(input->storage) : std::nullopt, - rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected nearest Resize plan"); - signalPassFailure(); - return; - } - if (failed(input)) { - rewriter.replaceOp(planOp, *lowered); - continue; - } - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("row-strip Resize plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - auto output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - if (auto planOp = dyn_cast(&op)) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("selected MaxPool plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - - FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - std::optional physicalInput; - if (succeeded(input)) - physicalInput = input->storage; - FailureOr lowered = lowerSelectedMaxPool2DPlan( - planOp, physicalInput, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected row-strip Spatial MaxPool plan"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - if (auto planOp = dyn_cast(&op)) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("selected global AveragePool plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - - FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - std::optional physicalInput; - if (succeeded(input)) - physicalInput = input->storage; - FailureOr lowered = - lowerSelectedGlobalAveragePoolPlan(planOp, physicalInput, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected row-strip Spatial global AveragePool plan"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - if (auto planOp = dyn_cast(&op)) { - if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) { - auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { - auto blueprint = dyn_cast(user); - return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; - }); - if (outputBlueprint == planOp.getResult().getUsers().end()) { - planOp.emitOpError("row-strip bias_add plan requires a row-strip blueprint result"); - signalPassFailure(); - return; - } - - FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); - rewriter.setInsertionPoint(planOp); - FailureOr lowered = lowerRowStripBiasAdd(*input, planOp, rewriter); - if (failed(lowered)) { - planOp.emitOpError("failed to lower selected row-strip Spatial bias_add plan"); - signalPassFailure(); - return; - } - auto blueprint = cast(*outputBlueprint); - FailureOr output = buildRowStripValue(blueprint, *lowered); - if (failed(output)) { - signalPassFailure(); - return; - } - rowStripValues[blueprint.getResult()] = *output; - eraseAfterLowering.insert(planOp); - eraseAfterLowering.insert(blueprint); - continue; - } - - auto resultType = dyn_cast(planOp.getOutput().getType()); - if (!resultType) { - planOp.emitOpError("requires ranked output type"); - signalPassFailure(); - return; - } - rewriter.setInsertionPoint(planOp); - FailureOr denseBias = materializeDenseBiasAddTensor(planOp.getBias(), resultType, rewriter, planOp.getLoc()); - if (failed(denseBias)) { - planOp.emitOpError("failed to materialize dense Conv-style bias"); - signalPassFailure(); - return; - } - if (planOp.getInput().getDefiningOp()) { - FailureOr lowered = lowerDenseBatchBiasAdd(planOp.getInput(), *denseBias, resultType, rewriter, planOp.getLoc()); - if (succeeded(lowered)) { - rewriter.replaceOp(planOp, *lowered); - continue; - } - } - auto computeOp = createSpatCompute<2>(rewriter, - planOp.getLoc(), - planOp.getOutput().getType(), - {}, - ValueRange {planOp.getInput(), *denseBias}, - [&](Value x, Value y) { - auto added = spatial::SpatVAddOp::create( - rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, y); - spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added.getResult()); - }); - rewriter.replaceOp(planOp, computeOp.getResults()); - continue; - } - if (auto planOp = dyn_cast(&op)) { - if (failed(lowerAddPlan(planOp, rowStripValues, eraseAfterLowering, rewriter))) { - signalPassFailure(); - return; - } - continue; - } - if (auto planOp = dyn_cast(&op)) { - if (failed(lowerConcatPlan(planOp, rowStripValues, eraseAfterLowering, rewriter))) { - signalPassFailure(); - return; - } - continue; - } - if (auto flattenOp = dyn_cast(&op)) { - if (flattenOp.getInputs().size() == 1) { - FailureOr input = - getRowStripValue(rowStripValues, flattenOp.getInputs().front()); - if (succeeded(input) && succeeded(canLowerFlattenFromRowStrip(flattenOp))) { - rewriter.setInsertionPoint(flattenOp); - if (failed(lowerFlattenFromRowStrip(*input, flattenOp, rewriter))) { - flattenOp.emitOpError("failed to preserve row-strip layout through Flatten"); - signalPassFailure(); - return; - } - continue; - } - } - } - if (auto materializeOp = dyn_cast(&op)) { - if (materializeOp.getSourcePhysicalLayout() == kDenseLayout - && materializeOp.getTargetPhysicalLayout() == kDenseLayout) { - rewriter.replaceOp(materializeOp, materializeOp.getInput()); - continue; - } - if (materializeOp.getSourcePhysicalLayout() != kRowStripLayout - || materializeOp.getTargetPhysicalLayout() != kDenseLayout) { - materializeOp.emitOpError("non-dense materialize_layout lowering is not supported yet"); - signalPassFailure(); - return; - } - FailureOr rowStripValue = getRowStripValue(rowStripValues, materializeOp.getInput()); - if (failed(rowStripValue)) { - materializeOp.emitOpError("expected a row-strip blueprint input during row-strip materialization"); - signalPassFailure(); - return; - } - rewriter.setInsertionPoint(materializeOp); - FailureOr dense = materializeRowStripToDense(*rowStripValue, materializeOp.getLoc(), rewriter); - if (failed(dense)) { - materializeOp.emitOpError("failed to materialize selected row-strip layout back to dense NCHW"); - signalPassFailure(); - return; - } - rewriter.replaceOp(materializeOp, *dense); - continue; - } - if (auto blueprintOp = dyn_cast(&op)) { - if (std::optional mode = blueprintOp.getMode(); mode && *mode == "fragment_assembly") - continue; - if (blueprintOp.getPhysicalLayout() == kDenseLayout) { - rewriter.replaceOp(blueprintOp, blueprintOp.getInput()); - continue; - } - if (blueprintOp.getPhysicalLayout() != kRowStripLayout) { - blueprintOp.emitOpError("non-dense blueprint lowering is not supported yet"); - signalPassFailure(); - return; - } - if (!eraseAfterLowering.contains(blueprintOp)) { - blueprintOp.emitOpError("unhandled row-strip blueprint remained during LowerSpatialPlans"); - signalPassFailure(); - return; - } - } - } - bool erasedAny = true; - while (erasedAny) { - erasedAny = false; - for (Operation& op : llvm::make_early_inc_range(funcOp.getBody().front())) { - if (!eraseAfterLowering.contains(&op)) - continue; - if (!op.use_empty()) - continue; - eraseAfterLowering.erase(&op); - rewriter.eraseOp(&op); - erasedAny = true; - } - } - if (!eraseAfterLowering.empty()) { - for (Operation& op : funcOp.getBody().front()) - if (eraseAfterLowering.contains(&op)) - op.emitOpError("selected row-strip planning op could not be fully eliminated during LowerSpatialPlans"); + if (failed(verifySelectedLayouts(funcOp, target))) { + moduleOp.emitError("selected Spatial layout verification failed"); signalPassFailure(); return; } - ConversionTarget helperTarget(*ctx); - helperTarget.addLegalDialect(ctx); + selectedPlanPatterns.add(ctx, target); + if (failed(applyPatternsGreedily(funcOp, std::move(selectedPlanPatterns)))) { + moduleOp.emitError("failed to lower selected Spatial plans"); + signalPassFailure(); + return; + } + + RewritePatternSet layoutPatterns(ctx); + layoutPatterns.add(ctx); + layoutPatterns.add(ctx, target); + ConversionTarget layoutTarget(*ctx); + layoutTarget.addLegalDialect(); - helperTarget.addLegalOp(); - helperTarget.addIllegalOp(); - helperTarget.markOpRecursivelyLegal(); - - RewritePatternSet helperPatterns(ctx); - populateGemmPatterns(helperPatterns, ctx); - populateTransposePatterns(helperPatterns, ctx); - FrozenRewritePatternSet frozenHelperPatterns( - std::move(helperPatterns)); - SmallVector topLevelHelperOps; - funcOp.walk([&](Operation* op) { - if (isa(op)) - return WalkResult::skip(); - if (isa(op)) - topLevelHelperOps.push_back(op); - return WalkResult::advance(); - }); - for (Operation *helper : topLevelHelperOps) { - if (failed(applyPartialConversion( - helper, helperTarget, frozenHelperPatterns))) { - moduleOp.emitError("failed to lower helper ONNX ops emitted by selected Spatial plan lowering"); - signalPassFailure(); - return; - } - } - ConversionTarget nestedHelperTarget(*ctx); - nestedHelperTarget.addLegalDialect(); - nestedHelperTarget.addIllegalOp(); - SmallVector computeLikeOps; - funcOp.walk([&](Operation* op) { - if (isa(op)) - computeLikeOps.push_back(op); - }); - for (Operation* op : computeLikeOps) { - if (failed(applyFullConversion( - op, nestedHelperTarget, frozenHelperPatterns))) { - op->emitOpError("failed to lower nested helper ONNX ops emitted by selected Spatial plan lowering"); - signalPassFailure(); - return; - } - } - if (!verifyLogicalPhase("after nested helper conversions")) + layoutTarget.addIllegalDialect(); + layoutTarget.addIllegalOp(); + layoutTarget.addDynamicallyLegalOp( + [&](spatial::SpatGraphCompute computeOp) { + if (computeOp.getInputs().size() != 1) + return true; + FailureOr input = + getRowStripValue(computeOp.getInputs().front()); + return failed(input) || failed(canLowerFlattenFromRowStrip(computeOp, target)); + }); + FrozenRewritePatternSet frozenLayoutPatterns(std::move(layoutPatterns)); + if (failed(applyFullConversion(funcOp, layoutTarget, + frozenLayoutPatterns))) { + moduleOp.emitError("failed to lower explicit Spatial layout materialization"); + signalPassFailure(); return; + } + + if (!verifyLogicalPhase("after selected-plan conversion")) + return; + SmallVector deadPhysicalViews; + funcOp.walk([&](spatial::SpatBlueprintOp blueprint) { + if (spatial::isPhysicalView(blueprint.getMode()) && blueprint.use_empty()) + deadPhysicalViews.push_back(blueprint); + }); + for (spatial::SpatBlueprintOp blueprint : deadPhysicalViews) + rewriter.eraseOp(blueprint); bool hasIllegalOps = false; moduleOp.walk([&](Operation* op) { if (isa(op)) return; if (auto blueprint = dyn_cast(op)) { - if (std::optional mode = blueprint.getMode(); mode && *mode == "fragment_assembly") + if (spatial::isFragmentAssembly(blueprint.getMode())) return; op->emitOpError("planning blueprint must not remain after LowerSpatialPlans"); hasIllegalOps = true; @@ -767,10 +886,17 @@ struct LowerSpatialPlansPass final : PassWrapper createLowerSpatialPlansPass() { return std::make_unique(); } +std::unique_ptr createLowerSpatialPlansPass(const spatial::SpatialTargetInfo& target) { + return std::make_unique(target); +} + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp index a7b4d72..8bff808 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp @@ -34,9 +34,15 @@ struct ONNXToSpatialPass : PassWrapper(); RewritePatternSet conversionPatterns(ctx); - populateConversionPatterns(conversionPatterns, ctx); + populateConversionPatterns(conversionPatterns, ctx, this->target); if (failed(applyPartialConversion(moduleOp, target, std::move(conversionPatterns)))) { moduleOp.emitError("failed to convert required ONNX ops to Spatial ops"); signalPassFailure(); @@ -258,4 +269,8 @@ void ONNXToSpatialPass::runOnOperation() { std::unique_ptr createONNXToSpatialPass() { return std::make_unique(); } +std::unique_ptr createONNXToSpatialPass(const spatial::SpatialTargetInfo& target) { + return std::make_unique(target); +} + } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp index 258fc89..ab7b38d 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp @@ -130,8 +130,7 @@ template void verifyNoNestedFragmentAssemblyBlueprints(ComputeOpTy compute, pim::CappedDiagnosticReporter& diagnostics) { compute.getBody().walk([&](spatial::SpatBlueprintOp blueprint) { - std::optional mode = blueprint.getMode(); - if (!mode || *mode != "fragment_assembly") + if (!spatial::isFragmentAssembly(blueprint.getMode())) return; diagnostics.report(blueprint.getOperation(), [&](Operation* illegalOp) { illegalOp->emitOpError("fragment assembly blueprint must be host-level after merge materialization"); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns.cpp index b2106fd..e9736ed 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns.cpp @@ -7,12 +7,14 @@ namespace onnx_mlir { void populatePrePatterns(RewritePatternSet& patterns, MLIRContext* ctx) { populateGeneratedPrePatterns(patterns, ctx); } -void populateConversionPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { +void populateConversionPatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { populateElementwisePatterns(patterns, ctx); - populateMatMulRewritePatterns(patterns, ctx); - populateGemmPatterns(patterns, ctx); - populateConvPatterns(patterns, ctx); - populatePoolPatterns(patterns, ctx); + populateMatMulRewritePatterns(patterns, ctx, target); + populateGemmPatterns(patterns, ctx, target); + populateConvPatterns(patterns, ctx, target); + populatePoolPatterns(patterns, ctx, target); populateReduceMeanPatterns(patterns, ctx); populateReluPatterns(patterns, ctx); populateSigmoidPatterns(patterns, ctx); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp b/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp index bb9d069..eea469b 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp @@ -8,20 +8,36 @@ namespace onnx_mlir { +namespace spatial { +struct SpatialTargetInfo; +} + void populatePrePatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populateConversionPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); +void populateConversionPatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); void populatePostPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); void populateGeneratedPrePatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); void populateWeightPromotionPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populateConvPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); +void populateConvPatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); void populateElementwisePatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); void populateElementwiseFusionPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populateGemmPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populateMatMulRewritePatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populateMatMulFusionPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); -void populatePoolPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); +void populateGemmPatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); +void populateMatMulRewritePatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); +void populateMatMulFusionPatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); +void populatePoolPatterns(mlir::RewritePatternSet& patterns, + mlir::MLIRContext* ctx, + const spatial::SpatialTargetInfo& target); void populateReduceMeanPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); void populateReluPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); void populateSigmoidPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp index 4eb5450..292ecdf 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp @@ -14,10 +14,11 @@ #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/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" @@ -30,11 +31,14 @@ namespace onnx_mlir { namespace { struct ConvToGemm : OpConversionPattern { - using OpConversionPattern::OpConversionPattern; + explicit ConvToGemm(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpConversionPattern(ctx), target(target) {} LogicalResult matchAndRewrite(ONNXConvOp convOp, ONNXConvOpAdaptor convOpAdaptor, ConversionPatternRewriter& rewriter) const override; + + const spatial::SpatialTargetInfo& target; }; struct PreparedConvInput { @@ -43,45 +47,21 @@ struct PreparedConvInput { }; static Value createZeroGemmBias(RankedTensorType gemmResultType, PatternRewriter& rewriter); -static StringRef stringifyConvLoweringStrategy(PimConvLoweringType strategy) { +static StringRef stringifyConvLoweringStrategy(spatial::ConvLoweringStrategy strategy) { switch (strategy) { - case PimConvLoweringAuto: return "auto"; - case PimConvLoweringLegacy: return "legacy"; - case PimConvLoweringDepthwise: return "depthwise"; - case PimConvLoweringPackedIm2Col: return "packed-im2col"; - case PimConvLoweringStreamedPatch: return "streamed-patch"; - case PimConvLoweringStreamedPacked: return "streamed-packed"; - case PimConvLoweringOutputChannelTiled: return "output-channel-tiled"; - case PimConvLoweringInputKTiled: return "input-k-tiled"; - case PimConvLoweringTiled2D: return "tiled-2d"; + 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"); } -static PimConvLoweringType chooseConvLoweringStrategy(const ConvGeometry& geo, - PimConvLoweringType requested) { - if (requested != PimConvLoweringAuto) - return requested; - - // Transform-based convolution is intentionally not selected for this ISA: - // it would require explicit transform sequences and staging traffic on top of - // the same crossbar MVM primitive, which is not attractive here. - if (geo.isDepthwise) - return PimConvLoweringDepthwise; - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize && geo.pack >= 2 && geo.im2colElements <= pimConvIm2colMaxElements) - return PimConvLoweringPackedIm2Col; - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize && geo.pack >= 2 && geo.im2colElements > pimConvIm2colMaxElements) - return PimConvLoweringStreamedPacked; - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize) - return PimConvLoweringStreamedPatch; - if (geo.k <= geo.xbarSize && geo.c > geo.xbarSize) - return PimConvLoweringOutputChannelTiled; - if (geo.k > geo.xbarSize && geo.c <= geo.xbarSize) - return PimConvLoweringLegacy; - return PimConvLoweringTiled2D; -} - - static Value expandBiasIfNeeded(Value bias, PatternRewriter& rewriter, Location loc) { auto biasType = cast(bias.getType()); if (biasType.getRank() != 1) @@ -194,7 +174,11 @@ static Value createCollectedConvOutput(ValueRange gemmRows, int64_t packFactor, PatternRewriter& rewriter, Location loc); -static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, Value x, Value w, Value b); +static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, + Value x, + Value w, + Value b, + const spatial::SpatialTargetInfo& target); namespace depthwise { @@ -215,10 +199,10 @@ static std::optional computeTiling(int64_t batchSize, int64_t wHeight, int64_t wWidth, int64_t outHeight, - int64_t outWidth) { + int64_t outWidth, + int64_t xbarDim) { const int64_t kernelElements = wHeight * wWidth; const int64_t outputMultiplier = numChannelsOut / numChannelsIn; - const int64_t xbarDim = static_cast(crossbarSize.getValue()); if (kernelElements <= 0 || outputMultiplier <= 0 || kernelElements > xbarDim || outputMultiplier > xbarDim) return std::nullopt; @@ -249,8 +233,9 @@ static Value buildPackedWeights(DenseElementsAttr wDenseAttr, const Tiling& tiling, PatternRewriter& rewriter, Location loc, + int64_t xbarDim, int64_t paddedInputRows = -1) { - const int64_t paddedOutputChannels = static_cast(crossbarSize.getValue()); + const int64_t paddedOutputChannels = xbarDim; const int64_t packedInputRows = paddedInputRows > 0 ? paddedInputRows : tiling.tileInputRows; auto packedWeightType = RankedTensorType::get( {tiling.numChannelTiles, packedInputRows, paddedOutputChannels}, wType.getElementType()); @@ -396,7 +381,7 @@ static Value createWeightTile(Value packedWeights, PatternRewriter& rewriter, Location loc) { SmallVector offsets {channelTileIndex, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; - const int64_t paddedOutputChannels = static_cast(crossbarSize.getValue()); + const int64_t paddedOutputChannels = packedWeightType.getDimSize(2); SmallVector sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(tiling.tileInputRows), rewriter.getIndexAttr(paddedOutputChannels)}; @@ -528,7 +513,8 @@ static bool canUseStructuredRewrite(const ConvLoweringState& state) { state.wHeight, state.wWidth, state.outHeight, - state.outWidth); + state.outWidth, + state.targetInfo().matrixShape.rows); if (!tiling) return false; @@ -559,7 +545,8 @@ rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& state.wType.getDimSize(2), state.wType.getDimSize(3), state.outType.getDimSize(2), - state.outType.getDimSize(3)); + 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(); @@ -579,9 +566,10 @@ rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& paddedInputType.getDimSize(3), paddedInputType.getDimSize(1)}, paddedInputType.getElementType()); - Value channelLastInput = ONNXTransposeOp::create( - rewriter, loc, channelLastInputType, paddedInput, rewriter.getI64ArrayAttr({0, 2, 3, 1})); - Value packedWeights = buildPackedWeights(wDenseAttr, state.wType, *tiling, rewriter, loc); + 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 batchInputs {channelLastInput}; @@ -600,7 +588,7 @@ rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& 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(crossbarSize.getValue())}, state.outType.getElementType()); + {1, static_cast(state.targetInfo().matrixShape.rows)}, state.outType.getElementType()); auto piecesType = spatial::getGraphBatchPhysicalResultType( tiling->totalPatches * tiling->numChannelTiles, rowTileType); auto inputTileType = @@ -826,8 +814,7 @@ static Value createWeightMatrix( }); if (!transpose) return flattened; - return ONNXTransposeOp::create(rewriter, loc, plan.wTransType, flattened, rewriter.getI64ArrayAttr({1, 0})) - .getResult(); + return createLinalgTranspose(flattened, plan.wTransType, {1, 0}, rewriter, loc); }; if (isCompileTimeComputable(weights)) @@ -963,17 +950,17 @@ static FailureOr rewriteInputKTiledConv(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc); - ConvGeometry geo = buildConvGeometry(state); + 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(1, - static_cast(crossbarCountInCore.getValue()) + static_cast(state.targetInfo().matrixUnitsPerProcessor) / static_cast(std::max(1, numKSlices * 4))); const uint64_t rowChunkWidth = std::max( 1, - std::min({chooseStreamChunkPositions(geo, /*packFactor=*/1), + std::min({chooseStreamChunkPositions(geo, /*packFactor=*/1, state.targetInfo()), maxLanesPerBatch, static_cast(state.outWidth)})); const auto elementType = state.outType.getElementType(); @@ -1227,7 +1214,7 @@ buildConvGemmPlan(const ConvLoweringState& state, const int64_t wMaxDim = std::max(plan.patchSize, state.numChannelsOut); plan.maxParallelPixels = forcedPackFactor ? *forcedPackFactor - : std::max(1, static_cast(crossbarSize.getValue()) / wMaxDim); + : std::max(1, static_cast(state.targetInfo().matrixShape.rows) / wMaxDim); plan.effectiveMaxParallelPixels = (canPackWeightsAsConstants && canPackBiasAsConstants) ? plan.maxParallelPixels : 1; plan.packedNumRows = ceilIntegerDivide(plan.chunkNumPatches, plan.effectiveMaxParallelPixels); @@ -1251,7 +1238,8 @@ static Value createIm2colRows(const ConvLoweringState& state, const ConvGemmPlan& plan, PatternRewriter& rewriter, Location loc) { - if (plan.gemmInputRowsType.getDimSize(1) > crossbarSize.getValue()) { + if (plan.gemmInputRowsType.getDimSize(1) + > static_cast(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) { @@ -1419,15 +1407,15 @@ static Value maybeUnpackChunkRows(Value gemmRows, return unpackCompute.getResult(0); } -static Value createStreamedConvRows(const ConvLoweringState& state, - const PreparedConvInput& preparedInput, - Value weightMatrix, - Value biasMatrix, - DenseElementsAttr wDenseAttr, - DenseElementsAttr biasDenseAttr, - int64_t forcedPackFactor, - PatternRewriter& rewriter, - Location loc) { +static FailureOr 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(wDenseAttr), !state.hasBias || static_cast(biasDenseAttr), 0, totalPatches, forcedPackFactor); @@ -1435,14 +1423,18 @@ static Value createStreamedConvRows(const ConvLoweringState& state, 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); - Value gemmRows = ONNXGemmOp::create(rewriter, loc, plan.gemmOutputRowsType, inputRows, - packedWeights, packedBias, APFloat(1.0f), APFloat(1.0f), 0, !wDenseAttr).getY(); - return maybeUnpackChunkRows(gemmRows, plan, rewriter, loc); + FailureOr 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 Value rewritePackedIm2ColConv(const ConvLoweringState& state, - PatternRewriter& rewriter, - Location loc) { +static FailureOr rewritePackedIm2ColConv(const ConvLoweringState& state, + PatternRewriter& rewriter, + Location loc) { auto wDenseAttr = getHostConstDenseElementsAttr(state.w); PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc); Value biasMatrix; @@ -1466,19 +1458,14 @@ static Value rewritePackedIm2ColConv(const ConvLoweringState& state, gemmBias = state.b; Value gemmC = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, plan, rewriter, loc); - Value gemmRows = ONNXGemmOp::create(rewriter, - loc, - plan.gemmOutputRowsType, - gemmInputRows, - gemmB, - gemmC, - APFloat(1.0f), - APFloat(1.0f), - /*transA=*/0, - /*transB=*/!wDenseAttr) - .getY(); + FailureOr 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}, + return createCollectedConvOutput(ValueRange {*gemmRows}, state.outType, plan.gemmOutType, plan.nhwcType, @@ -1490,10 +1477,10 @@ static Value rewritePackedIm2ColConv(const ConvLoweringState& state, loc); } -static Value rewriteStreamedConv(const ConvLoweringState& state, - PatternRewriter& rewriter, - Location loc, - int64_t forcedPackFactor) { +static FailureOr rewriteStreamedConv(const ConvLoweringState& state, + PatternRewriter& rewriter, + Location loc, + int64_t forcedPackFactor) { auto wDenseAttr = getHostConstDenseElementsAttr(state.w); PreparedConvInput preparedInput = prepareInputForIm2Col(state, rewriter, loc); Value biasMatrix; @@ -1506,20 +1493,22 @@ static Value rewriteStreamedConv(const ConvLoweringState& state, ConvGemmPlan seedPlan = buildConvGemmPlan( state, static_cast(wDenseAttr), !state.hasBias || static_cast(biasDenseAttr), 0, 1, forcedPackFactor); Value weightMatrix = createWeightMatrix(state.w, seedPlan, static_cast(wDenseAttr), rewriter, loc); - Value collectedRows = createStreamedConvRows(state, - preparedInput, - weightMatrix, - biasMatrix, - wDenseAttr, - biasDenseAttr, - forcedPackFactor, - rewriter, - loc); - auto gemmOutType = cast(collectedRows.getType()); + FailureOr collectedRows = createStreamedConvRows(state, + preparedInput, + weightMatrix, + biasMatrix, + wDenseAttr, + biasDenseAttr, + forcedPackFactor, + rewriter, + loc); + if (failed(collectedRows)) + return failure(); + auto gemmOutType = cast(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), + ValueRange {*collectedRows}, state.outType, gemmOutType, nhwcType, state.outType, gemmOutType.getDimSize(0), state.numChannelsOut, /*packFactor=*/1, rewriter, loc); } @@ -1534,12 +1523,12 @@ static Value createZeroGemmBias(RankedTensorType gemmResultType, PatternRewriter static bool rowStripOutputTileFitsOneCore(const ConvGeometry& geometry) { return ceilIntegerDivide(geometry.k, geometry.xbarSize) * ceilIntegerDivide(geometry.c, geometry.xbarSize) - <= static_cast(crossbarCountInCore.getValue()); + <= geometry.matrixUnitsPerProcessor; } static bool rowStripOutputChannelTileFitsOneCore(const ConvGeometry& geometry) { return ceilIntegerDivide(geometry.k, geometry.xbarSize) - <= static_cast(crossbarCountInCore.getValue()); + <= geometry.matrixUnitsPerProcessor; } static int64_t chooseRowStripPixelPackFactor(const ConvLoweringState& state, int64_t xbarDim) { @@ -1581,7 +1570,7 @@ static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state failureReason = "non_constant_weight"; return false; } - if (!rowStripOutputChannelTileFitsOneCore(buildConvGeometry(state))) { + if (!rowStripOutputChannelTileFitsOneCore(buildConvGeometry(state, state.targetInfo()))) { failureReason = "output_channel_tile_does_not_fit_one_core"; return false; } @@ -1790,8 +1779,7 @@ static Value extractDenseConvWindowRow(Value denseInput, rewriter.getIndexAttr(state.xWidth)}; Value nchw = tensor::ExtractSliceOp::create( rewriter, loc, nchwType, denseInput, offsets, sizes, getUnitStrides(rewriter, 4)); - return ONNXTransposeOp::create( - rewriter, loc, fragmentType, nchw, rewriter.getI64ArrayAttr({0, 2, 3, 1})); + return createLinalgTranspose(nchw, fragmentType, {0, 2, 3, 1}, rewriter, loc); } static Value createRowStripWindowMaskTable(const ConvLoweringState& state, PatternRewriter& rewriter) { @@ -2429,7 +2417,7 @@ static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLow static FailureOr createRowStripConvOutputFromDenseInput(const ConvLoweringState& state, PatternRewriter& rewriter, Location loc) { - ConvGeometry geometry = buildConvGeometry(state); + ConvGeometry geometry = buildConvGeometry(state, state.targetInfo()); if (state.group != 1 || state.batchSize != 1 || !rowStripOutputChannelTileFitsOneCore(geometry)) return failure(); @@ -2481,7 +2469,7 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro if (!canConsumePixelMajorRowStripFragments(state, failureReason)) return failure(); - ConvGeometry geometry = buildConvGeometry(state); + 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); @@ -2521,7 +2509,7 @@ static FailureOr createPointwiseOutputFromRowStripFragments(Value rowStri Location loc) { FailureOr input = describeRowStripPhysicalValue(rowStripStorage, state.xType); if (failed(input)) return failure(); - ConvGeometry geometry = buildConvGeometry(state); + 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) @@ -2621,8 +2609,10 @@ static bool canConsumeDepthwiseRowStrip(const ConvLoweringState& state) { state.wHeight, state.wWidth, state.outHeight, - state.outWidth); - return tiling && tiling->numChannelTiles <= static_cast(crossbarCountInCore.getValue()); + state.outWidth, + state.targetInfo().matrixShape.rows); + return tiling && tiling->numChannelTiles + <= static_cast(state.targetInfo().matrixUnitsPerProcessor); } static Value insertDepthwiseInputSegment(Value inputWindow, @@ -2726,16 +2716,23 @@ static FailureOr createDepthwiseOutputFromRowStripFragments(Value rowStri state.wHeight, state.wWidth, state.outHeight, - state.outWidth); + 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(crossbarSize.getValue())); + weight, + state.wType, + *tiling, + rewriter, + loc, + static_cast(state.targetInfo().matrixShape.rows), + static_cast(state.targetInfo().matrixShape.rows)); Value bias = state.hasBias ? expandBiasIfNeeded(state.b, rewriter, loc) : Value(); auto paddedOutputType = RankedTensorType::get( - {1, static_cast(crossbarSize.getValue())}, state.outType.getElementType()); + {1, static_cast(state.targetInfo().matrixShape.rows)}, state.outType.getElementType()); auto outputTileType = RankedTensorType::get( {1, tiling->tileOutputChannels}, state.outType.getElementType()); auto outputPixelType = RankedTensorType::get( @@ -2759,7 +2756,7 @@ static FailureOr createDepthwiseOutputFromRowStripFragments(Value rowStri 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(crossbarSize.getValue()); + const int64_t xbarDim = static_cast(state.targetInfo().matrixShape.rows); auto paddedInputScratchType = RankedTensorType::get( {tiling->numChannelTiles, 1, 1, xbarDim}, state.xType.getElementType(), state.xType.getEncoding()); auto tileScratchType = RankedTensorType::get( @@ -2773,7 +2770,8 @@ static FailureOr createDepthwiseOutputFromRowStripFragments(Value rowStri SmallVector biasTiles; SmallVector tileIndices; SmallVector weightTileSizes { - rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarDim), rewriter.getIndexAttr(xbarDim)}; + 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); @@ -2861,10 +2859,10 @@ static FailureOr createDepthwiseOutputFromRowStripFragments(Value rowStri static FailureOr createConvOutputFromRowStripInput(const ConvLoweringState& state, Value rowStripInput, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter, Location loc) { - if (strategy == PimConvLoweringDepthwise) + 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); @@ -2905,17 +2903,23 @@ static Value createCollectedConvOutput(ValueRange gemmRows, {0, 1, 2}, {3} }); - Value nchwOut = ONNXTransposeOp::create(rewriter, loc, outType, nhwcOut, rewriter.getI64ArrayAttr({0, 3, 1, 2})); + Value nchwOut = createLinalgTranspose(nhwcOut, outType, {0, 3, 1, 2}, rewriter, loc); spatial::SpatYieldOp::create(rewriter, loc, nchwOut); }); return collectComputeOp.getResult(0); } -static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, Value x, Value w, Value b) { +static FailureOr 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 = ⌖ state.xType = cast(state.x.getType()); state.wType = cast(state.w.getType()); state.outType = cast(convOp.getY().getType()); @@ -3019,6 +3023,7 @@ static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, state.padWidthBegin = getI64Attr(*padsAttr, 1); state.padHeightEnd = getI64Attr(*padsAttr, 2); state.padWidthEnd = getI64Attr(*padsAttr, 3); + classifyConvProblem(state); return state; } @@ -3043,6 +3048,7 @@ static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, state.padWidthEnd = totalPadW / 2; state.padWidthBegin = totalPadW - state.padWidthEnd; } + classifyConvProblem(state); return state; } @@ -3051,18 +3057,25 @@ static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, return failure(); } + classifyConvProblem(state); return state; } -static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, ONNXConvOpAdaptor convOpAdaptor) { - return analyzeConvLoweringState(convOp, convOpAdaptor.getX(), convOpAdaptor.getW(), convOpAdaptor.getB()); +static FailureOr analyzeConvLoweringState(ONNXConvOp convOp, + ONNXConvOpAdaptor convOpAdaptor, + const spatial::SpatialTargetInfo& target) { + return analyzeConvLoweringState( + convOp, convOpAdaptor.getX(), convOpAdaptor.getW(), convOpAdaptor.getB(), target); } -static FailureOr analyzeConvLoweringState(spatial::SpatConv2DPlanOp planOp) { +static FailureOr 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 = ⌖ state.xType = dyn_cast(state.x.getType()); state.wType = dyn_cast(state.w.getType()); state.outType = dyn_cast(planOp.getOutput().getType()); @@ -3111,79 +3124,56 @@ static FailureOr analyzeConvLoweringState(spatial::SpatConv2D state.strideWidth = strides[1]; state.dilationHeight = dilations[0]; state.dilationWidth = dilations[1]; + classifyConvProblem(state); return state; } -static FailureOr resolveRequestedConvLoweringStrategy(Operation* op) { - if (!useExperimentalConvImpl) - return pimConvLowering.getValue(); +static FailureOr +resolveRequestedConvLoweringStrategy(Operation* op, const spatial::SpatialTargetInfo& target) { + if (!target.useExperimentalConvImplementation) + return target.convLoweringStrategy; - if (pimConvLowering != PimConvLoweringAuto && pimConvLowering != PimConvLoweringPackedIm2Col) { + if (target.convLoweringStrategy != spatial::ConvLoweringStrategy::Auto + && target.convLoweringStrategy != spatial::ConvLoweringStrategy::PackedIm2Col) { op->emitOpError() << "--use-experimental-conv-impl conflicts with --pim-conv-lowering=" - << stringifyConvLoweringStrategy(pimConvLowering); + << stringifyConvLoweringStrategy(target.convLoweringStrategy); return failure(); } - return PimConvLoweringPackedIm2Col; + return spatial::ConvLoweringStrategy::PackedIm2Col; } -static LogicalResult verifyForcedConvLoweringStrategy(Operation* op, - const ConvGeometry& geo, - PimConvLoweringType strategy) { - switch (strategy) { - case PimConvLoweringAuto: - case PimConvLoweringLegacy: - return success(); - case PimConvLoweringDepthwise: - if (geo.isDepthwise) - return success(); - return op->emitOpError("forced depthwise Conv lowering requires a depthwise convolution"); - case PimConvLoweringPackedIm2Col: - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize && geo.pack >= 2 && geo.im2colElements <= pimConvIm2colMaxElements) - return success(); - return op->emitOpError("forced packed-im2col Conv lowering requires K/C to fit, pack >= 2, and im2col within budget"); - case PimConvLoweringStreamedPatch: - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize) - return success(); - return op->emitOpError("forced streamed-patch Conv lowering requires K and C to each fit one crossbar"); - case PimConvLoweringStreamedPacked: - if (geo.k <= geo.xbarSize && geo.c <= geo.xbarSize && geo.pack >= 2) - return success(); - return op->emitOpError("forced streamed-packed Conv lowering requires K/C to fit and pack >= 2"); - case PimConvLoweringOutputChannelTiled: - if (geo.k <= geo.xbarSize && geo.c > geo.xbarSize) - return success(); - return op->emitOpError("forced output-channel-tiled Conv lowering requires K <= X and C > X"); - case PimConvLoweringInputKTiled: - if (geo.k > geo.xbarSize && geo.c <= geo.xbarSize) - return success(); - return op->emitOpError("forced input-k-tiled Conv lowering requires K > X and C <= X"); - case PimConvLoweringTiled2D: - if (geo.k > geo.xbarSize && geo.c > geo.xbarSize) - return success(); - return op->emitOpError("forced tiled-2d Conv lowering requires K > X and C > X"); - } - llvm_unreachable("unknown conv lowering strategy"); -} - -static FailureOr selectConvLoweringStrategy(Operation* op, - const ConvLoweringState& state) { - FailureOr requested = resolveRequestedConvLoweringStrategy(op); +static FailureOr selectConvLoweringPlan( + Operation* op, const ConvLoweringState& state) { + FailureOr requested = + resolveRequestedConvLoweringStrategy(op, state.targetInfo()); if (failed(requested)) return failure(); - ConvGeometry geometry = buildConvGeometry(state); - PimConvLoweringType strategy = chooseConvLoweringStrategy(geometry, *requested); - if (strategy == PimConvLoweringDepthwise && !depthwise::canUseStructuredRewrite(state) - && *requested == PimConvLoweringAuto) - strategy = PimConvLoweringLegacy; - if (failed(verifyForcedConvLoweringStrategy(op, geometry, strategy))) + if (*requested == spatial::ConvLoweringStrategy::Auto) { + for (const ConvPlan& candidate : buildConvPlanCandidates(state, state.targetInfo())) { + if (candidate.strategy == spatial::ConvLoweringStrategy::Depthwise + && !depthwise::canUseStructuredRewrite(state)) { + continue; + } + return candidate; + } + op->emitOpError("has no applicable Conv lowering candidate for the injected Spatial target"); return failure(); - return strategy; + } + + FailureOr 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(); + } + return *candidate; } static FailureOr lowerDenseSelectedConvPlan(Operation* op, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter, Location loc); @@ -3196,18 +3186,18 @@ static ConvLoweringState makeGroupedConvLoweringState(const ConvLoweringState& p static FailureOr buildConvValueForStrategy(Operation* op, Location loc, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter); static FailureOr buildGroupedConvValue(Operation* op, Location loc, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter); static FailureOr lowerGroupedSelectedConvPlan(Operation* op, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter, Location loc) { return buildGroupedConvValue(op, loc, state, strategy, rewriter); @@ -3215,7 +3205,7 @@ static FailureOr lowerGroupedSelectedConvPlan(Operation* op, static FailureOr lowerDenseSelectedConvPlan(Operation* op, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter, Location loc) { return buildConvValueForStrategy(op, loc, state, strategy, rewriter); @@ -3224,29 +3214,29 @@ static FailureOr lowerDenseSelectedConvPlan(Operation* op, static FailureOr buildConvValueForStrategy(Operation* op, Location loc, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter) { - const ConvGeometry geo = buildConvGeometry(state); + const ConvGeometry geo = buildConvGeometry(state, state.targetInfo()); switch (strategy) { - case PimConvLoweringDepthwise: { + case spatial::ConvLoweringStrategy::Depthwise: { return depthwise::rewriteConv(op, state, rewriter, loc); } - case PimConvLoweringLegacy: - case PimConvLoweringPackedIm2Col: { + case spatial::ConvLoweringStrategy::Legacy: + case spatial::ConvLoweringStrategy::PackedIm2Col: { return standard::rewritePackedIm2ColConv(state, rewriter, loc); } - case PimConvLoweringStreamedPatch: - case PimConvLoweringOutputChannelTiled: - case PimConvLoweringTiled2D: { + case spatial::ConvLoweringStrategy::StreamedPatch: + case spatial::ConvLoweringStrategy::OutputChannelTiled: + case spatial::ConvLoweringStrategy::Tiled2D: { return standard::rewriteStreamedConv(state, rewriter, loc, /*forcedPackFactor=*/1); } - case PimConvLoweringInputKTiled: { + case spatial::ConvLoweringStrategy::InputKTiled: { return standard::rewriteInputKTiledConv(state, rewriter, loc); } - case PimConvLoweringStreamedPacked: { + case spatial::ConvLoweringStrategy::StreamedPacked: { return standard::rewriteStreamedConv(state, rewriter, loc, geo.pack); } - case PimConvLoweringAuto: + case spatial::ConvLoweringStrategy::Auto: break; } op->emitOpError("unexpected auto strategy at Conv lowering dispatch"); @@ -3282,13 +3272,14 @@ static ConvLoweringState makeGroupedConvLoweringState( state.numChannelsInPerGroup = state.numChannelsIn; state.numChannelsOutPerGroup = state.numChannelsOut; state.hasBias = static_cast(groupB); + classifyConvProblem(state); return state; } static FailureOr buildGroupedConvValue(Operation* op, Location loc, const ConvLoweringState& state, - PimConvLoweringType strategy, + spatial::ConvLoweringStrategy strategy, PatternRewriter& rewriter) { SmallVector xSlices = sliceTensor(state.x, /*axis=*/1, state.numChannelsInPerGroup, rewriter, loc); SmallVector wSlices = sliceTensor(state.w, /*axis=*/0, state.numChannelsOutPerGroup, rewriter, loc); @@ -3344,7 +3335,7 @@ static FailureOr buildGroupedConvValue(Operation* op, LogicalResult ConvToGemm::matchAndRewrite(ONNXConvOp convOp, ONNXConvOpAdaptor convOpAdaptor, ConversionPatternRewriter& rewriter) const { - FailureOr state = analyzeConvLoweringState(convOp, convOpAdaptor); + FailureOr state = analyzeConvLoweringState(convOp, convOpAdaptor, target); if (failed(state)) return failure(); SmallVector pads { @@ -3362,15 +3353,20 @@ LogicalResult ConvToGemm::matchAndRewrite(ONNXConvOp convOp, rewriter.getDenseI64ArrayAttr(strides), rewriter.getDenseI64ArrayAttr(dilations), rewriter.getI64IntegerAttr(state->group), - rewriter.getStringAttr("nchw")); + spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(convOp, convPlan.getResult()); return success(); } -void populateConvPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.insert(ctx); } +void populateConvPatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { + patterns.insert(ctx, target); +} -LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) { - FailureOr state = analyzeConvLoweringState(planOp); +LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp, + const spatial::SpatialTargetInfo& target) { + FailureOr state = analyzeConvLoweringState(planOp, target); if (failed(state)) return failure(); @@ -3383,38 +3379,41 @@ LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp) { if (state->hasBias && !isSupportedBiasAddValue(state->b, state->outType)) return failure(); - ConvGeometry geometry = buildConvGeometry(*state); + ConvGeometry geometry = buildConvGeometry(*state, state->targetInfo()); if (!rowStripOutputChannelTileFitsOneCore(geometry)) return failure(); - FailureOr strategy = selectConvLoweringStrategy(planOp.getOperation(), *state); - if (failed(strategy)) + FailureOr plan = + selectConvLoweringPlan(planOp.getOperation(), *state); + if (failed(plan)) return failure(); - switch (*strategy) { - case PimConvLoweringLegacy: - case PimConvLoweringDepthwise: - case PimConvLoweringPackedIm2Col: - case PimConvLoweringStreamedPatch: - case PimConvLoweringOutputChannelTiled: - case PimConvLoweringTiled2D: - case PimConvLoweringStreamedPacked: + 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 PimConvLoweringAuto: - case PimConvLoweringInputKTiled: + case spatial::ConvLoweringStrategy::Auto: + case spatial::ConvLoweringStrategy::InputKTiled: return failure(); } llvm_unreachable("unknown conv lowering strategy"); } -LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp) { - FailureOr state = analyzeConvLoweringState(planOp); +LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp, + const spatial::SpatialTargetInfo& target) { + FailureOr state = analyzeConvLoweringState(planOp, target); if (failed(state)) return failure(); - FailureOr strategy = selectConvLoweringStrategy(planOp.getOperation(), *state); - if (failed(strategy)) + FailureOr plan = + selectConvLoweringPlan(planOp.getOperation(), *state); + if (failed(plan)) return failure(); - if (*strategy == PimConvLoweringDepthwise) + if (plan->strategy == spatial::ConvLoweringStrategy::Depthwise) return canConsumeDepthwiseRowStrip(*state) ? success() : failure(); StringRef failureReason; return canConsumePixelMajorRowStripFragments(*state, failureReason) ? success() : failure(); @@ -3424,23 +3423,25 @@ FailureOr lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp, std::optional rowStripInput, bool emitRowStripLayout, + const spatial::SpatialTargetInfo& target, PatternRewriter& rewriter) { - FailureOr state = analyzeConvLoweringState(planOp); + FailureOr state = analyzeConvLoweringState(planOp, target); if (failed(state)) return failure(); - FailureOr strategy = selectConvLoweringStrategy(planOp.getOperation(), *state); - if (failed(strategy)) + FailureOr plan = + selectConvLoweringPlan(planOp.getOperation(), *state); + if (failed(plan)) return failure(); if (emitRowStripLayout) { if (rowStripInput) { - if (failed(canConsumeAndProduceRowStrip(planOp))) + if (failed(canConsumeAndProduceRowStrip(planOp, target))) return planOp.emitOpError("selected row-strip input/output layout is not supported for this Conv plan"), failure(); return createConvOutputFromRowStripInput( - *state, *rowStripInput, *strategy, rewriter, planOp.getLoc()); + *state, *rowStripInput, plan->strategy, rewriter, planOp.getLoc()); } - if (failed(canLowerConvPlanToRowStrip(planOp))) + if (failed(canLowerConvPlanToRowStrip(planOp, target))) return planOp.emitOpError("selected row-strip layout is not supported for this Conv plan"), failure(); FailureOr rowStripStorage = createRowStripConvOutputFromDenseInput(*state, rewriter, planOp.getLoc()); if (failed(rowStripStorage)) @@ -3448,11 +3449,11 @@ lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp, return *rowStripStorage; } - if (*strategy == PimConvLoweringDepthwise) - return lowerDenseSelectedConvPlan(planOp.getOperation(), *state, *strategy, rewriter, planOp.getLoc()); + 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, *strategy, rewriter, planOp.getLoc()); - return lowerDenseSelectedConvPlan(planOp.getOperation(), *state, *strategy, rewriter, planOp.getLoc()); + return lowerGroupedSelectedConvPlan(planOp.getOperation(), *state, plan->strategy, rewriter, planOp.getLoc()); + return lowerDenseSelectedConvPlan(planOp.getOperation(), *state, plan->strategy, rewriter, planOp.getLoc()); } } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.cpp index 4b1fa71..e15889b 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.cpp @@ -3,47 +3,277 @@ #include #include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" -#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" namespace onnx_mlir { +namespace { + +static int64_t ceilDivide(int64_t value, int64_t divisor) { + return divisor == 0 ? 0 : (value + divisor - 1) / divisor; +} + +} // namespace + bool isDepthwiseConv(int64_t group, int64_t numChannelsIn, int64_t numChannelsOut, int64_t numChannelsInPerGroup) { return group == numChannelsIn && numChannelsInPerGroup == 1 && numChannelsOut % group == 0; } -ConvGeometry buildConvGeometry(const ConvLoweringState& state) { +void classifyConvProblem(ConvProblem& problem) { + problem.isDepthwise = isDepthwiseConv( + problem.group, problem.numChannelsIn, problem.numChannelsOut, + problem.numChannelsInPerGroup); + problem.isGrouped = problem.group > 1; + problem.isPointwise = problem.wHeight == 1 && problem.wWidth == 1 + && problem.strideHeight == 1 && problem.strideWidth == 1 + && problem.dilationHeight == 1 && problem.dilationWidth == 1 + && problem.padHeightBegin == 0 && problem.padHeightEnd == 0 + && problem.padWidthBegin == 0 && problem.padWidthEnd == 0; +} + +ConvGeometry buildConvGeometry(const ConvProblem& problem, + const spatial::SpatialTargetInfo& target) { ConvGeometry geo { - state.batchSize, - state.numChannelsIn, - state.xHeight, - state.xWidth, - state.numChannelsOut, - state.wHeight, - state.wWidth, - state.outHeight, - state.outWidth, - state.group, - state.numChannelsInPerGroup, - state.numChannelsOutPerGroup, - state.numChannelsInPerGroup * state.wHeight * state.wWidth, - state.numChannelsOutPerGroup, - state.batchSize * state.outHeight * state.outWidth, - static_cast(crossbarSize.getValue()), + problem.batchSize, + problem.numChannelsIn, + problem.xHeight, + problem.xWidth, + problem.numChannelsOut, + problem.wHeight, + problem.wWidth, + problem.outHeight, + problem.outWidth, + problem.group, + problem.numChannelsInPerGroup, + problem.numChannelsOutPerGroup, + problem.numChannelsInPerGroup * problem.wHeight * problem.wWidth, + problem.numChannelsOutPerGroup, + problem.batchSize * problem.outHeight * problem.outWidth, + static_cast(target.matrixShape.rows), + static_cast(target.matrixUnitsPerProcessor), 1, 0, - state.hasBias, - isDepthwiseConv(state.group, state.numChannelsIn, state.numChannelsOut, state.numChannelsInPerGroup), + problem.hasBias, + isDepthwiseConv(problem.group, + problem.numChannelsIn, + problem.numChannelsOut, + problem.numChannelsInPerGroup), }; geo.pack = std::max(1, geo.xbarSize / std::max(geo.k, geo.c)); geo.im2colElements = static_cast(std::max(0, geo.p)) * static_cast(std::max(0, geo.k)); return geo; } -uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, int64_t packFactor) { +static ConvMaterializationKind getMaterializationKind( + const ConvProblem& problem, spatial::ConvLoweringStrategy strategy) { + if (strategy == spatial::ConvLoweringStrategy::Depthwise) + return ConvMaterializationKind::StructuredDepthwise; + if (problem.isPointwise) + return ConvMaterializationKind::PointwiseContraction; + switch (strategy) { + case spatial::ConvLoweringStrategy::Depthwise: + return ConvMaterializationKind::StructuredDepthwise; + case spatial::ConvLoweringStrategy::PackedIm2Col: + case spatial::ConvLoweringStrategy::Legacy: + return ConvMaterializationKind::PackedIm2Col; + case spatial::ConvLoweringStrategy::StreamedPatch: + return ConvMaterializationKind::StreamedPatch; + case spatial::ConvLoweringStrategy::StreamedPacked: + return ConvMaterializationKind::StreamedPacked; + case spatial::ConvLoweringStrategy::OutputChannelTiled: + return ConvMaterializationKind::OutputChannelTiled; + case spatial::ConvLoweringStrategy::InputKTiled: + return ConvMaterializationKind::InputKTiled; + case spatial::ConvLoweringStrategy::Tiled2D: + return ConvMaterializationKind::Tiled2D; + case spatial::ConvLoweringStrategy::Auto: + break; + } + llvm_unreachable("auto is not a Conv materialization kind"); +} + +static ConvPlan makeCandidatePlan(const ConvProblem& problem, + spatial::ConvLoweringStrategy strategy, + const spatial::SpatialTargetInfo& target) { + ConvPlan plan; + plan.geometry = buildConvGeometry(problem, target); + plan.strategy = strategy; + plan.materializationKind = getMaterializationKind(problem, strategy); + plan.laneCount = plan.geometry.p; + plan.reductionCount = std::max( + 1, (plan.geometry.k + plan.geometry.xbarSize - 1) / plan.geometry.xbarSize); + plan.mvmCount = plan.laneCount * plan.reductionCount; + plan.vectorCount = plan.mvmCount; + plan.weightElements = static_cast(std::max(0, problem.numChannelsOut)) + * static_cast(std::max(0, plan.geometry.k)); + plan.scratchElements = plan.geometry.im2colElements; + plan.materializationElements = strategy == spatial::ConvLoweringStrategy::Depthwise + ? 0 + : std::min(plan.geometry.im2colElements, target.convIm2colMaxElements); + plan.requiresInputMaterialization = strategy != spatial::ConvLoweringStrategy::Depthwise; + plan.producesRowStrip = strategy != spatial::ConvLoweringStrategy::InputKTiled + && ceilDivide(plan.geometry.k, plan.geometry.xbarSize) <= plan.geometry.matrixUnitsPerProcessor; + plan.consumesRowStrip = plan.producesRowStrip; + // Conv materializers emit local compute and leave inter-core communication + // to Spatial scheduling; zero is an explicit ownership statement here. + plan.communicationElements = 0; + plan.usesContraction = problem.isPointwise || strategy != spatial::ConvLoweringStrategy::Depthwise; + if (problem.isPointwise) { + ContractionProblem contraction; + contraction.origin = ContractionOrigin::Gemm; + contraction.batch = 1; + contraction.m = plan.geometry.p; + contraction.k = plan.geometry.c; + contraction.n = problem.numChannelsOutPerGroup; + contraction.lhsElementType = problem.xType.getElementType(); + contraction.rhsElementType = problem.wType.getElementType(); + contraction.resultElementType = problem.outType.getElementType(); + plan.contraction = makeContractionPlan( + contraction, target, ContractionPlanKind::StaticTiled); + plan.hasContractionPlan = true; + plan.laneCount = plan.contraction.laneCount; + plan.mvmCount = plan.contraction.expectedMvmCount; + plan.vectorCount = plan.contraction.expectedVectorCount; + plan.reductionCount = plan.contraction.reductionSlices; + } + return plan; +} + +static bool fitsSingleCrossbar(const ConvGeometry& geo) { + return geo.k <= geo.xbarSize && geo.c <= geo.xbarSize; +} + +static bool fitsPackedIm2Col(const ConvGeometry& geo, + const spatial::SpatialTargetInfo& target) { + return fitsSingleCrossbar(geo) && geo.pack >= 2 + && geo.im2colElements <= target.convIm2colMaxElements; +} + +static mlir::FailureOr buildDepthwiseCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + if (!problem.isDepthwise) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::Depthwise, target); +} + +static mlir::FailureOr buildPackedIm2ColCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + if (!fitsPackedIm2Col(geo, target)) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::PackedIm2Col, target); +} + +static mlir::FailureOr buildStreamedPatchCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + if (!fitsSingleCrossbar(buildConvGeometry(problem, target))) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::StreamedPatch, target); +} + +static mlir::FailureOr buildStreamedPackedCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + if (!fitsSingleCrossbar(geo) || geo.pack < 2) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::StreamedPacked, target); +} + +static mlir::FailureOr buildOutputChannelTiledCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + if (geo.k > geo.xbarSize || geo.c <= geo.xbarSize) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::OutputChannelTiled, target); +} + +static mlir::FailureOr buildInputKTiledCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + if (geo.k <= geo.xbarSize || geo.c > geo.xbarSize) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::InputKTiled, target); +} + +static mlir::FailureOr buildTiled2DCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + if (geo.k <= geo.xbarSize || geo.c <= geo.xbarSize) + return mlir::failure(); + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::Tiled2D, target); +} + +static mlir::FailureOr buildLegacyCandidate( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + // Legacy is retained as the explicit compatibility/debug materializer and + // as the safe fallback when structured depthwise lowering is unavailable. + return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::Legacy, target); +} + +mlir::FailureOr makeConvPlan(const ConvProblem& problem, + spatial::ConvLoweringStrategy strategy, + const spatial::SpatialTargetInfo& target) { + switch (strategy) { + case spatial::ConvLoweringStrategy::Auto: + return mlir::failure(); + case spatial::ConvLoweringStrategy::Legacy: + return buildLegacyCandidate(problem, target); + case spatial::ConvLoweringStrategy::Depthwise: + return buildDepthwiseCandidate(problem, target); + case spatial::ConvLoweringStrategy::PackedIm2Col: + return buildPackedIm2ColCandidate(problem, target); + case spatial::ConvLoweringStrategy::StreamedPatch: + return buildStreamedPatchCandidate(problem, target); + case spatial::ConvLoweringStrategy::StreamedPacked: + return buildStreamedPackedCandidate(problem, target); + case spatial::ConvLoweringStrategy::OutputChannelTiled: + return buildOutputChannelTiledCandidate(problem, target); + case spatial::ConvLoweringStrategy::InputKTiled: + return buildInputKTiledCandidate(problem, target); + case spatial::ConvLoweringStrategy::Tiled2D: + return buildTiled2DCandidate(problem, target); + } + llvm_unreachable("unknown Conv lowering strategy"); +} + +llvm::SmallVector buildConvPlanCandidates( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target) { + ConvGeometry geo = buildConvGeometry(problem, target); + llvm::SmallVector candidates; + auto append = [&](spatial::ConvLoweringStrategy strategy) { + mlir::FailureOr candidate = makeConvPlan(problem, strategy, target); + if (succeeded(candidate)) + candidates.push_back(*candidate); + }; + + if (problem.isDepthwise) { + append(spatial::ConvLoweringStrategy::Depthwise); + append(spatial::ConvLoweringStrategy::Legacy); + return candidates; + } + if (fitsPackedIm2Col(geo, target)) + append(spatial::ConvLoweringStrategy::PackedIm2Col); + if (fitsSingleCrossbar(geo) && geo.pack >= 2) + append(spatial::ConvLoweringStrategy::StreamedPacked); + if (fitsSingleCrossbar(geo)) + append(spatial::ConvLoweringStrategy::StreamedPatch); + if (geo.k <= geo.xbarSize && geo.c > geo.xbarSize) + append(spatial::ConvLoweringStrategy::OutputChannelTiled); + if (geo.k > geo.xbarSize && geo.c <= geo.xbarSize) + append(spatial::ConvLoweringStrategy::Legacy); + if (geo.k > geo.xbarSize && geo.c <= geo.xbarSize) + append(spatial::ConvLoweringStrategy::InputKTiled); + if (geo.k > geo.xbarSize && geo.c > geo.xbarSize) + append(spatial::ConvLoweringStrategy::Tiled2D); + return candidates; +} + +uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, + int64_t packFactor, + const spatial::SpatialTargetInfo& target) { const uint64_t patchElements = static_cast(std::max(1, geo.k)); - uint64_t chunkPositions = std::max(1, pimConvIm2colMaxElements / patchElements); + uint64_t chunkPositions = std::max(1, target.convIm2colMaxElements / patchElements); chunkPositions = std::min(chunkPositions, static_cast(std::max(1, geo.p))); - chunkPositions = std::min(chunkPositions, std::max(1, pimConvStreamChunkPositions)); + chunkPositions = std::min(chunkPositions, std::max(1, target.convStreamChunkPositions)); if (packFactor > 1 && chunkPositions > static_cast(packFactor)) { chunkPositions -= chunkPositions % static_cast(packFactor); @@ -52,24 +282,26 @@ uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, int64_t packFactor) return std::max(1, chunkPositions); } -RowInterval computeConvInputRowsForOutputRows(RowInterval outputRows, const ConvLoweringState& state) { - const int64_t rawBegin = outputRows.begin * state.strideHeight - state.padHeightBegin; +RowInterval computeConvInputRowsForOutputRows(RowInterval outputRows, const ConvProblem& problem) { + const int64_t rawBegin = outputRows.begin * problem.strideHeight - problem.padHeightBegin; const int64_t rawEnd = - (outputRows.end - 1) * state.strideHeight - state.padHeightBegin + state.dilationHeight * (state.wHeight - 1) + 1; - return {std::max(0, rawBegin), std::min(state.xHeight, rawEnd)}; + (outputRows.end - 1) * problem.strideHeight - problem.padHeightBegin + + problem.dilationHeight * (problem.wHeight - 1) + 1; + return {std::max(0, rawBegin), std::min(problem.xHeight, rawEnd)}; } -ConvRowDemand buildConvRowDemand(RowInterval outputRows, const ConvLoweringState& state) { +ConvRowDemand buildConvRowDemand(RowInterval outputRows, const ConvProblem& problem) { ConvRowDemand demand; demand.outputRows = outputRows; - demand.neededInputRows = computeConvInputRowsForOutputRows(outputRows, state); + demand.neededInputRows = computeConvInputRowsForOutputRows(outputRows, problem); demand.acquiredInputRows = demand.neededInputRows; - const int64_t rawBegin = outputRows.begin * state.strideHeight - state.padHeightBegin; + const int64_t rawBegin = outputRows.begin * problem.strideHeight - problem.padHeightBegin; const int64_t rawEnd = - (outputRows.end - 1) * state.strideHeight - state.padHeightBegin + state.dilationHeight * (state.wHeight - 1) + 1; + (outputRows.end - 1) * problem.strideHeight - problem.padHeightBegin + + problem.dilationHeight * (problem.wHeight - 1) + 1; demand.topHaloRows = std::max(0, -rawBegin); - demand.bottomHaloRows = std::max(0, rawEnd - state.xHeight); + demand.bottomHaloRows = std::max(0, rawEnd - problem.xHeight); demand.acquiredInputRows = demand.neededInputRows; return demand; } diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.hpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.hpp index 60564c6..6eab24a 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ConvGeometry.hpp @@ -3,14 +3,19 @@ #include "mlir/IR/BuiltinTypes.h" #include "mlir/IR/Value.h" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetInfo.hpp" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" + #include +namespace mlir { +class Operation; +} // namespace mlir + namespace onnx_mlir { -struct ConvLoweringState { - mlir::Value x; - mlir::Value w; - mlir::Value b; +struct ConvProblem { mlir::RankedTensorType xType; mlir::RankedTensorType wType; mlir::RankedTensorType outType; @@ -35,6 +40,19 @@ struct ConvLoweringState { int64_t dilationHeight; int64_t dilationWidth; bool hasBias; + bool isDepthwise = false; + bool isGrouped = false; + bool isPointwise = false; +}; + +struct ConvLoweringState : ConvProblem { + mlir::Operation* diagnosticAnchor = nullptr; + mlir::Value x; + mlir::Value w; + mlir::Value b; + const spatial::SpatialTargetInfo* target = nullptr; + + const spatial::SpatialTargetInfo& targetInfo() const { return *target; } }; struct ConvGeometry { @@ -54,6 +72,7 @@ struct ConvGeometry { int64_t c; int64_t p; int64_t xbarSize; + int64_t matrixUnitsPerProcessor; int64_t pack; uint64_t im2colElements; bool hasBias; @@ -73,14 +92,59 @@ struct ConvRowDemand { int64_t bottomHaloRows = 0; }; +enum class ConvMaterializationKind : uint8_t { + StructuredDepthwise, + PointwiseContraction, + PackedIm2Col, + StreamedPatch, + StreamedPacked, + OutputChannelTiled, + InputKTiled, + Tiled2D, +}; + +struct ConvPlan { + ConvGeometry geometry; + spatial::ConvLoweringStrategy strategy = spatial::ConvLoweringStrategy::Auto; + ConvMaterializationKind materializationKind = ConvMaterializationKind::PackedIm2Col; + int64_t laneCount = 0; + int64_t mvmCount = 0; + int64_t vectorCount = 0; + int64_t reductionCount = 0; + uint64_t weightElements = 0; + uint64_t scratchElements = 0; + uint64_t materializationElements = 0; + uint64_t communicationElements = 0; + spatial::PhysicalLayout resultLayout = spatial::PhysicalLayout::DenseNCHW; + bool consumesRowStrip = false; + bool producesRowStrip = false; + bool requiresInputMaterialization = false; + bool requiresOutputMaterialization = false; + bool usesContraction = false; + bool hasContractionPlan = false; + ContractionPlan contraction; +}; + bool isDepthwiseConv(int64_t group, int64_t numChannelsIn, int64_t numChannelsOut, int64_t numChannelsInPerGroup); -ConvGeometry buildConvGeometry(const ConvLoweringState& state); +void classifyConvProblem(ConvProblem& problem); -uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, int64_t packFactor); +ConvGeometry buildConvGeometry(const ConvProblem& problem, + const spatial::SpatialTargetInfo& target); -RowInterval computeConvInputRowsForOutputRows(RowInterval outputRows, const ConvLoweringState& state); +mlir::FailureOr makeConvPlan(const ConvProblem& problem, + spatial::ConvLoweringStrategy strategy, + const spatial::SpatialTargetInfo& target); -ConvRowDemand buildConvRowDemand(RowInterval outputRows, const ConvLoweringState& state); +llvm::SmallVector buildConvPlanCandidates( + const ConvProblem& problem, const spatial::SpatialTargetInfo& target); + +uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, + int64_t packFactor, + const spatial::SpatialTargetInfo& target); + +RowInterval computeConvInputRowsForOutputRows(RowInterval outputRows, const ConvProblem& problem); + +ConvRowDemand buildConvRowDemand(RowInterval outputRows, const ConvProblem& problem); } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp index 871640d..53a767a 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp @@ -31,7 +31,7 @@ struct SiluToSpatialPlan : OpRewritePattern { return failure(); auto plan = spatial::SpatSiluPlanOp::create( - rewriter, mulOp.getLoc(), mulOp.getResult().getType(), input, rewriter.getStringAttr("nchw")); + rewriter, mulOp.getLoc(), mulOp.getResult().getType(), input, spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(mulOp, plan.getResult()); rewriter.eraseOp(sigmoidOp); return success(); @@ -260,14 +260,16 @@ struct AddToSpatialCompute : OpConversionPattern { classifyBiasAddPlanCandidate(adaptor.getA(), adaptor.getB(), resultType); if (succeeded(candidate)) { auto plan = spatial::SpatBiasAddPlanOp::create( - rewriter, op.getLoc(), resultType, candidate->data, candidate->bias, rewriter.getStringAttr("nchw")); + rewriter, op.getLoc(), resultType, candidate->data, candidate->bias, + spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(op, plan.getResult()); return success(); } if (resultType.getRank() == 4 && adaptor.getA().getType() == resultType && adaptor.getB().getType() == resultType) { auto plan = spatial::SpatAddPlanOp::create( - rewriter, op.getLoc(), resultType, adaptor.getA(), adaptor.getB(), rewriter.getStringAttr("nchw")); + rewriter, op.getLoc(), resultType, adaptor.getA(), adaptor.getB(), + spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(op, plan.getResult()); return success(); } diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp index 082ef04..5fe8789 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp @@ -1,5 +1,6 @@ #include "mlir/Dialect/Affine/IR/AffineOps.h" #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/BuiltinTypes.h" @@ -21,6 +22,10 @@ #include "src/Accelerators/PIM/Common/PimCommon.hpp" #include "src/Accelerators/PIM/Common/Support/Diagnostics.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Dialect/ONNX/ONNXOps.hpp" @@ -31,7 +36,7 @@ namespace onnx_mlir { namespace { static FailureOr -materializeScaledConstantTensor(Value value, float factor, ConversionPatternRewriter& rewriter, Location loc) { +materializeScaledConstantTensor(Value value, float factor, PatternRewriter& rewriter, Location loc) { if (factor == 1.0f) return value; @@ -57,7 +62,12 @@ materializeScaledConstantTensor(Value value, float factor, ConversionPatternRewr } static Value createGemmBatchKOffset( - Value lane, int64_t numOutRows, int64_t numKSlices, ConversionPatternRewriter& rewriter, Location loc) { + Value lane, + int64_t numOutRows, + int64_t numKSlices, + int64_t xbarSize, + PatternRewriter& rewriter, + Location loc) { if (numKSlices == 1) return getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0); @@ -65,7 +75,7 @@ static Value createGemmBatchKOffset( AffineExpr d0 = getAffineDimExpr(0, context); return createOrFoldAffineApply(rewriter, loc, - (d0.floorDiv(numOutRows) % numKSlices) * crossbarSize.getValue(), + (d0.floorDiv(numOutRows) % numKSlices) * xbarSize, ValueRange {lane}, rewriter.getInsertionBlock()->getParentOp()); } @@ -74,7 +84,8 @@ static Value createGemmBatchHOffset(Value lane, int64_t numOutRows, int64_t numKSlices, int64_t numOutHSlices, - ConversionPatternRewriter& rewriter, + int64_t xbarSize, + PatternRewriter& rewriter, Location loc) { if (numOutHSlices == 1) return getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0); @@ -83,14 +94,14 @@ static Value createGemmBatchHOffset(Value lane, AffineExpr d0 = getAffineDimExpr(0, context); return createOrFoldAffineApply(rewriter, loc, - d0.floorDiv(numOutRows * numKSlices) * crossbarSize.getValue(), + d0.floorDiv(numOutRows * numKSlices) * xbarSize, ValueRange {lane}, rewriter.getInsertionBlock()->getParentOp()); } static FailureOr materializePaddedConstantMatrix(Value value, RankedTensorType resultType, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { auto sourceType = cast(value.getType()); if (sourceType == resultType) @@ -121,7 +132,7 @@ static FailureOr materializePaddedConstantMatrix(Value value, static FailureOr materializePaddedBroadcastedConstantTensor(Value value, RankedTensorType resultType, int64_t unpaddedColumns, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { auto denseAttr = getHostConstDenseElementsAttr(value); if (!denseAttr) @@ -187,7 +198,7 @@ static FailureOr materializePaddedBroadcastedConstantTensor(Value value, static FailureOr prepareBias(Value c, RankedTensorType outType, RankedTensorType paddedOutType, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { auto cType = cast(c.getType()); if (!cType.hasStaticShape()) @@ -203,9 +214,15 @@ static FailureOr prepareBias(Value c, } static Value extractATile( - Value a, Value row, Value kOffset, RankedTensorType aTileType, ConversionPatternRewriter& rewriter, Location loc) { + Value a, + Value row, + Value kOffset, + RankedTensorType aTileType, + int64_t xbarSize, + PatternRewriter& rewriter, + Location loc) { SmallVector offsets {row, kOffset}; - SmallVector sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())}; + SmallVector sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)}; SmallVector strides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; return tensor::ExtractSliceOp::create(rewriter, loc, aTileType, a, offsets, sizes, strides).getResult(); @@ -219,7 +236,8 @@ static FailureOr createVmmBatch(Value a, int64_t numOutRows, int64_t numKSlices, int64_t numOutHSlices, - ConversionPatternRewriter& rewriter, + int64_t xbarSize, + PatternRewriter& rewriter, Location loc) { const int64_t laneCount = partialPiecesType.getDimSize(0); auto batchOp = createSpatComputeBatch( @@ -232,21 +250,21 @@ static FailureOr createVmmBatch(Value a, [&](detail::SpatComputeBatchBodyArgs args) { Value row = onnx_mlir::affineModConst(rewriter, loc, args.lane, numOutRows, rewriter.getInsertionBlock()->getParentOp()); - Value kOffset = createGemmBatchKOffset(args.lane, numOutRows, numKSlices, rewriter, loc); - Value hOffset = createGemmBatchHOffset(args.lane, numOutRows, numKSlices, numOutHSlices, rewriter, loc); + Value kOffset = createGemmBatchKOffset(args.lane, numOutRows, numKSlices, xbarSize, rewriter, loc); + Value hOffset = createGemmBatchHOffset( + args.lane, numOutRows, numKSlices, numOutHSlices, xbarSize, rewriter, loc); auto aTileType = - RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, aType.getElementType()); + RankedTensorType::get({1, xbarSize}, aType.getElementType()); auto bTileType = RankedTensorType::get( - {static_cast(crossbarSize.getValue()), static_cast(crossbarSize.getValue())}, + {xbarSize, xbarSize}, paddedBType.getElementType()); auto pieceType = - RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, partialPiecesType.getElementType()); - Value aTile = extractATile(args.inputs.front(), row, kOffset, aTileType, rewriter, loc); + RankedTensorType::get({1, xbarSize}, partialPiecesType.getElementType()); + Value aTile = extractATile(args.inputs.front(), row, kOffset, aTileType, xbarSize, rewriter, loc); SmallVector bOffsets {kOffset, hOffset}; - SmallVector bSizes {rewriter.getIndexAttr(crossbarSize.getValue()), - rewriter.getIndexAttr(crossbarSize.getValue())}; + SmallVector bSizes {rewriter.getIndexAttr(xbarSize), rewriter.getIndexAttr(xbarSize)}; SmallVector unitStrides = getUnitStrides(rewriter, 2); Value bTile = extractStaticSliceOrIdentity( rewriter, loc, args.weights.front(), bTileType, bOffsets, bSizes, unitStrides); @@ -260,7 +278,7 @@ static FailureOr createVmmBatch(Value a, } static Value extractDynamicGemmBColumn( - Value matrix, Value column, RankedTensorType vectorType, ConversionPatternRewriter& rewriter, Location loc) { + Value matrix, Value column, RankedTensorType vectorType, PatternRewriter& rewriter, Location loc) { SmallVector offsets {rewriter.getIndexAttr(0), column}; SmallVector sizes {rewriter.getIndexAttr(vectorType.getDimSize(1)), rewriter.getIndexAttr(1)}; SmallVector strides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; @@ -280,7 +298,7 @@ static Value extractDynamicGemmBColumn( } static Value extractDynamicGemmRowVector( - Value matrix, Value row, RankedTensorType vectorType, ConversionPatternRewriter& rewriter, Location loc) { + Value matrix, Value row, RankedTensorType vectorType, PatternRewriter& rewriter, Location loc) { SmallVector offsets {row, rewriter.getIndexAttr(0)}; SmallVector sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(vectorType.getDimSize(1))}; SmallVector strides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; @@ -317,7 +335,7 @@ static bool hasGemmBias(Value c) { static Value createScalarTensorConstant(RankedTensorType scalarType, float value, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { auto elementType = scalarType.getElementType(); auto scalarAttr = rewriter.getFloatAttr(elementType, value); @@ -330,7 +348,7 @@ static Value createBroadcastedBiasScalar(Value bias, Value row, Value column, RankedTensorType scalarType, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { SmallVector unitStrides(biasType.getRank(), rewriter.getIndexAttr(1)); if (biasType.getRank() == 1) { @@ -365,7 +383,7 @@ static FailureOr createVvdmulBatch(Value a, RankedTensorType columnPiecesType, RankedTensorType outType, bool transposeB, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { const int64_t numOutRows = outType.getDimSize(0); const int64_t numOutCols = outType.getDimSize(1); @@ -425,7 +443,7 @@ static FailureOr createDynamicGemmOutputCompute(Value scal RankedTensorType outType, float alpha, float beta, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { const int64_t numOutRows = outType.getDimSize(0); const int64_t numOutCols = outType.getDimSize(1); @@ -510,7 +528,7 @@ static Value createPartialGroupOffset(Value hSlice, int64_t kSlice, int64_t numKSlices, int64_t numOutRows, - ConversionPatternRewriter& rewriter, + PatternRewriter& rewriter, Location loc) { MLIRContext* context = rewriter.getContext(); AffineExpr d0 = getAffineDimExpr(0, context); @@ -527,10 +545,12 @@ static Value extractReductionPiece(Value partialPiecesArg, RankedTensorType pieceType, int64_t numKSlices, int64_t numOutRows, - ConversionPatternRewriter& rewriter, + int64_t xbarSize, + PatternRewriter& rewriter, Location loc) { SmallVector unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; - SmallVector pieceSizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())}; + SmallVector pieceSizes { + rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)}; SmallVector pieceOffsets { createPartialGroupOffset(hSlice, kSlice, numKSlices, numOutRows, rewriter, loc), rewriter.getIndexAttr(0), @@ -545,13 +565,15 @@ static Value reducePartialPiecesForHSlice(Value partialPiecesArg, RankedTensorType pieceType, int64_t numKSlices, int64_t numOutRows, - ConversionPatternRewriter& rewriter, + int64_t xbarSize, + PatternRewriter& rewriter, Location loc) { SmallVector activePieces; activePieces.reserve(numKSlices); for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) activePieces.push_back( - extractReductionPiece(partialPiecesArg, hSlice, kSlice, pieceType, numKSlices, numOutRows, rewriter, loc)); + extractReductionPiece( + partialPiecesArg, hSlice, kSlice, pieceType, numKSlices, numOutRows, xbarSize, rewriter, loc)); while (activePieces.size() > 1) { SmallVector nextPieces; @@ -574,11 +596,12 @@ static FailureOr createReductionOutput(Value partialPieces, RankedTensorType outType, RankedTensorType paddedOutType, int64_t numKSlices, - ConversionPatternRewriter& rewriter, + int64_t xbarSize, + PatternRewriter& rewriter, Location loc) { const int64_t numOutRows = outType.getDimSize(0); - const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(1), crossbarSize.getValue()); - auto pieceType = RankedTensorType::get({numOutRows, static_cast(crossbarSize.getValue())}, + const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(1), xbarSize); + auto pieceType = RankedTensorType::get({numOutRows, xbarSize}, partialPiecesType.getElementType()); if (bias && cast(bias.getType()) != paddedOutType) @@ -590,20 +613,20 @@ static FailureOr createReductionOutput(Value partialPieces, SmallVector outputSlices; outputSlices.reserve(numOutHSlices); for (int64_t hSlice = 0; hSlice < numOutHSlices; ++hSlice) { - const int64_t columnOffset = hSlice * crossbarSize.getValue(); + const int64_t columnOffset = hSlice * xbarSize; const int64_t columns = - std::min(static_cast(crossbarSize.getValue()), outType.getDimSize(1) - columnOffset); + std::min(xbarSize, outType.getDimSize(1) - columnOffset); auto outputSliceType = RankedTensorType::get({numOutRows, columns}, outType.getElementType()); auto computeOp = createSpatCompute( rewriter, loc, TypeRange {outputSliceType}, {}, inputs, [&](ValueRange blockArgs) -> LogicalResult { Value hSliceValue = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), hSlice); Value reduced = reducePartialPiecesForHSlice( - blockArgs[0], hSliceValue, pieceType, numKSlices, numOutRows, rewriter, loc); + blockArgs[0], hSliceValue, pieceType, numKSlices, numOutRows, xbarSize, rewriter, loc); if (bias) { SmallVector biasOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(columnOffset)}; SmallVector pieceSizes {rewriter.getIndexAttr(numOutRows), - rewriter.getIndexAttr(crossbarSize.getValue())}; + rewriter.getIndexAttr(xbarSize)}; SmallVector unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; Value biasSlice = tensor::ExtractSliceOp::create(rewriter, loc, pieceType, blockArgs[1], biasOffsets, pieceSizes, unitStrides) @@ -637,79 +660,101 @@ static FailureOr createReductionOutput(Value partialPieces, } struct GemmToSpatialComputes : OpConversionPattern { - using OpConversionPattern::OpConversionPattern; + explicit GemmToSpatialComputes(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpConversionPattern(ctx), target(target) {} LogicalResult matchAndRewrite(ONNXGemmOp gemmOp, ONNXGemmOpAdaptor gemmOpAdaptor, ConversionPatternRewriter& rewriter) const override; + + const spatial::SpatialTargetInfo& target; }; } // namespace -LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, - ONNXGemmOpAdaptor gemmOpAdaptor, - ConversionPatternRewriter& rewriter) const { - Location loc = gemmOp.getLoc(); - Value a = gemmOpAdaptor.getA(); - Value b = gemmOpAdaptor.getB(); - Value c = gemmOpAdaptor.getC(); - +FailureOr lowerGemmToSpatial( + Operation* diagnosticAnchor, + Value a, + Value b, + Value c, + RankedTensorType outType, + bool transA, + bool transB, + float alpha, + float beta, + const spatial::SpatialTargetInfo& target, + PatternRewriter& rewriter, + Location loc) { auto aType = dyn_cast(a.getType()); auto bType = dyn_cast(b.getType()); - auto outType = dyn_cast(gemmOp.getY().getType()); - if (!aType || !bType || !outType) + if (!diagnosticAnchor || !aType || !bType || !outType) return failure(); if (!aType.hasStaticShape()) { - pim::emitUnsupportedStaticShapeDiagnostic(gemmOp, "Gemm input A"); + pim::emitUnsupportedStaticShapeDiagnostic(diagnosticAnchor, "Gemm input A"); return failure(); } if (!bType.hasStaticShape()) { - pim::emitUnsupportedStaticShapeDiagnostic(gemmOp, "Gemm input B"); + pim::emitUnsupportedStaticShapeDiagnostic(diagnosticAnchor, "Gemm input B"); return failure(); } if (!outType.hasStaticShape()) { - pim::emitUnsupportedStaticShapeDiagnostic(gemmOp, "Gemm result"); + pim::emitUnsupportedStaticShapeDiagnostic(diagnosticAnchor, "Gemm result"); return failure(); } if (aType.getRank() != 2) { - pim::emitUnsupportedRankDiagnostic(gemmOp, "Gemm input A", aType.getRank(), {2}); + pim::emitUnsupportedRankDiagnostic(diagnosticAnchor, "Gemm input A", aType.getRank(), {2}); return failure(); } if (bType.getRank() != 2) { - pim::emitUnsupportedRankDiagnostic(gemmOp, "Gemm input B", bType.getRank(), {2}); + pim::emitUnsupportedRankDiagnostic(diagnosticAnchor, "Gemm input B", bType.getRank(), {2}); return failure(); } if (outType.getRank() != 2) { - pim::emitUnsupportedRankDiagnostic(gemmOp, "Gemm result", outType.getRank(), {2}); + pim::emitUnsupportedRankDiagnostic(diagnosticAnchor, "Gemm result", outType.getRank(), {2}); return failure(); } - if (gemmOpAdaptor.getTransA()) { + if (transA) { auto aShape = aType.getShape(); auto transposedType = RankedTensorType::get({aShape[1], aShape[0]}, aType.getElementType(), aType.getEncoding()); - a = ONNXTransposeOp::create(rewriter, loc, transposedType, a, rewriter.getI64ArrayAttr({1, 0})).getResult(); + a = createLinalgTranspose(a, transposedType, {1, 0}, rewriter, loc); aType = transposedType; } - const int64_t numOutRows = outType.getDimSize(0); - const int64_t numOutCols = outType.getDimSize(1); - const int64_t reductionSize = aType.getDimSize(1); - const bool transposeB = gemmOpAdaptor.getTransB(); + ContractionProblem problem; + problem.lhsBatchShape = {}; + problem.rhsBatchShape = {}; + problem.outputBatchShape = {}; + problem.lhsBatch = 1; + problem.rhsBatch = 1; + problem.batch = 1; + problem.m = outType.getDimSize(0); + problem.k = aType.getDimSize(1); + problem.n = outType.getDimSize(1); + problem.origin = ContractionOrigin::Gemm; + problem.lhsElementType = aType.getElementType(); + problem.rhsElementType = bType.getElementType(); + problem.resultElementType = outType.getElementType(); + problem.lhsTransposed = transA; + problem.rhsTransposed = transB; + problem.alpha = alpha; + problem.beta = beta; + const bool transposeB = transB; if (!isCompileTimeComputable(b)) { + ContractionPlan plan = makeContractionPlan( + problem, target, ContractionPlanKind::BatchedDynamicVVD); bool hasC = hasGemmBias(c); - float alpha = gemmOpAdaptor.getAlpha().convertToFloat(); - float beta = gemmOpAdaptor.getBeta().convertToFloat(); RankedTensorType biasType; if (hasC) { auto cType = dyn_cast(c.getType()); if (!cType || !cType.hasStaticShape()) { - pim::emitUnsupportedStaticShapeDiagnostic(gemmOp, "Gemm bias"); + pim::emitUnsupportedStaticShapeDiagnostic(diagnosticAnchor, "Gemm bias"); return failure(); } auto verifiedBiasType = verifyDynamicGemmBiasType(cType, outType); if (failed(verifiedBiasType)) { - gemmOp.emitOpError("requires Gemm bias C to be broadcastable to the output shape"); + diagnosticAnchor->emitOpError("requires Gemm bias C to be broadcastable to the output shape"); return failure(); } biasType = *verifiedBiasType; @@ -717,19 +762,19 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, const int64_t bReductionSize = bType.getDimSize(transposeB ? 1 : 0); const int64_t bOutputColumns = bType.getDimSize(transposeB ? 0 : 1); - if (aType.getDimSize(0) != numOutRows || bReductionSize != reductionSize || bOutputColumns != numOutCols) { - gemmOp.emitOpError("has inconsistent A, B, and output shapes"); + if (aType.getDimSize(0) != problem.m || bReductionSize != problem.k || bOutputColumns != problem.n) { + diagnosticAnchor->emitOpError("has inconsistent A, B, and output shapes"); return failure(); } - const int64_t laneCount64 = numOutRows * numOutCols; + const int64_t laneCount64 = plan.laneCount; if (laneCount64 > std::numeric_limits::max()) { - gemmOp.emitOpError("requires Gemm dynamic batch lane count to fit in i32"); + diagnosticAnchor->emitOpError("requires Gemm dynamic batch lane count to fit in i32"); return failure(); } - auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType()); - auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType(numOutCols, columnType); + auto columnType = RankedTensorType::get({problem.m, 1}, outType.getElementType()); + auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType(problem.n, columnType); auto batchOp = createVvdmulBatch(a, b, aType, bType, scalarPiecesType, outType, transposeB, rewriter, loc); if (failed(batchOp)) return failure(); @@ -737,94 +782,122 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, batchOp->getResult(0), hasC ? c : Value(), scalarPiecesType, biasType, outType, alpha, beta, rewriter, loc); if (failed(outputCompute)) return failure(); - rewriter.replaceOp(gemmOp, outputCompute->getResults()); - return success(); + return outputCompute->getResult(0); } if (transposeB) { auto bShape = bType.getShape(); auto transposedType = RankedTensorType::get({bShape[1], bShape[0]}, bType.getElementType(), bType.getEncoding()); - b = ONNXTransposeOp::create(rewriter, loc, transposedType, b, rewriter.getI64ArrayAttr({1, 0})).getResult(); + if (isCompileTimeComputable(b)) { + auto transposedConstant = materializeTransposedContractionConstant( + b, transposedType, {1, 0}, rewriter, loc); + if (failed(transposedConstant)) { + diagnosticAnchor->emitOpError("requires Gemm input B transpose to remain statically materializable"); + return failure(); + } + b = *transposedConstant; + } else { + b = createLinalgTranspose(b, transposedType, {1, 0}, rewriter, loc); + } bType = transposedType; } - auto scaledB = materializeScaledConstantTensor(b, gemmOpAdaptor.getAlpha().convertToFloat(), rewriter, loc); + auto scaledB = materializeScaledConstantTensor(b, alpha, rewriter, loc); if (failed(scaledB)) { - gemmOp.emitOpError("requires constant Gemm input B when alpha is not 1.0"); + diagnosticAnchor->emitOpError("requires constant Gemm input B when alpha is not 1.0"); return failure(); } b = *scaledB; bType = cast(b.getType()); - if (aType.getDimSize(0) != numOutRows || bType.getDimSize(0) != reductionSize || bType.getDimSize(1) != numOutCols) { - gemmOp.emitOpError("has inconsistent A, B, and output shapes after transpose handling"); + if (aType.getDimSize(0) != problem.m || bType.getDimSize(0) != problem.k || bType.getDimSize(1) != problem.n) { + diagnosticAnchor->emitOpError("has inconsistent A, B, and output shapes after transpose handling"); return failure(); } - const int64_t numKSlices = ceilIntegerDivide(reductionSize, crossbarSize.getValue()); - const int64_t numOutHSlices = ceilIntegerDivide(numOutCols, crossbarSize.getValue()); - const int64_t paddedReductionSize = numKSlices * static_cast(crossbarSize.getValue()); - const int64_t paddedOutCols = numOutHSlices * static_cast(crossbarSize.getValue()); + ContractionPlan plan = makeContractionPlan( + problem, target, ContractionPlanKind::StaticTiled); + const int64_t xbarSize = plan.tileK; + const int64_t numKSlices = plan.reductionSlices; + const int64_t numOutHSlices = plan.outputTiles; + const int64_t paddedReductionSize = numKSlices * plan.tileK; + const int64_t paddedOutCols = numOutHSlices * plan.tileN; auto paddedBType = RankedTensorType::get({paddedReductionSize, paddedOutCols}, bType.getElementType()); auto paddedB = materializePaddedConstantMatrix(b, paddedBType, rewriter, loc); if (failed(paddedB)) { - gemmOp.emitOpError("requires constant Gemm input B so tiled weights can be padded statically"); + diagnosticAnchor->emitOpError("requires constant Gemm input B so tiled weights can be padded statically"); return failure(); } b = *paddedB; - auto paddedAType = RankedTensorType::get({numOutRows, paddedReductionSize}, aType.getElementType()); - a = createPaddedInputCompute(a, paddedAType, rewriter, loc); + auto paddedAType = RankedTensorType::get({problem.m, paddedReductionSize}, aType.getElementType()); + a = materializePaddedContractionInput(a, paddedAType, rewriter, loc); aType = paddedAType; Value bias; bool hasC = hasGemmBias(c); - auto paddedOutType = RankedTensorType::get({numOutRows, paddedOutCols}, outType.getElementType()); + auto paddedOutType = RankedTensorType::get({problem.m, paddedOutCols}, outType.getElementType()); if (hasC) { auto cType = dyn_cast(c.getType()); if (!cType || !cType.hasStaticShape()) { - pim::emitUnsupportedStaticShapeDiagnostic(gemmOp, "Gemm bias"); + pim::emitUnsupportedStaticShapeDiagnostic(diagnosticAnchor, "Gemm bias"); return failure(); } - auto scaledC = materializeScaledConstantTensor(c, gemmOpAdaptor.getBeta().convertToFloat(), rewriter, loc); + auto scaledC = materializeScaledConstantTensor(c, beta, rewriter, loc); if (failed(scaledC)) { - gemmOp.emitOpError("requires constant Gemm bias C when beta is not 1.0"); + diagnosticAnchor->emitOpError("requires constant Gemm bias C when beta is not 1.0"); return failure(); } c = *scaledC; auto preparedBias = prepareBias(c, outType, paddedOutType, rewriter, loc); if (failed(preparedBias)) { - gemmOp.emitOpError("requires Gemm bias C to be broadcastable to the output shape"); + diagnosticAnchor->emitOpError("requires Gemm bias C to be broadcastable to the output shape"); return failure(); } bias = *preparedBias; } - const int64_t laneCount64 = numOutHSlices * numKSlices * numOutRows; + const int64_t laneCount64 = plan.laneCount; if (laneCount64 > std::numeric_limits::max()) { - gemmOp.emitOpError("requires Gemm tiled batch lane count to fit in i32"); + diagnosticAnchor->emitOpError("requires Gemm tiled batch lane count to fit in i32"); return failure(); } auto partialPiecesType = spatial::getGraphBatchPhysicalResultType( - laneCount64, RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, outType.getElementType())); + laneCount64, RankedTensorType::get({1, xbarSize}, outType.getElementType())); auto batchOp = - createVmmBatch(a, b, aType, paddedBType, partialPiecesType, numOutRows, numKSlices, numOutHSlices, rewriter, loc); + createVmmBatch( + a, b, aType, paddedBType, partialPiecesType, problem.m, numKSlices, numOutHSlices, xbarSize, rewriter, loc); if (failed(batchOp)) return failure(); auto reductionOutput = createReductionOutput( - batchOp->getResult(0), bias, partialPiecesType, outType, paddedOutType, numKSlices, rewriter, loc); + batchOp->getResult(0), bias, partialPiecesType, outType, paddedOutType, numKSlices, xbarSize, rewriter, loc); if (failed(reductionOutput)) return failure(); - rewriter.replaceOp(gemmOp, *reductionOutput); + return *reductionOutput; +} + +LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, + ONNXGemmOpAdaptor gemmOpAdaptor, + ConversionPatternRewriter& rewriter) const { + FailureOr result = lowerGemmToSpatial( + gemmOp.getOperation(), gemmOpAdaptor.getA(), gemmOpAdaptor.getB(), gemmOpAdaptor.getC(), + cast(gemmOp.getY().getType()), gemmOpAdaptor.getTransA(), + gemmOpAdaptor.getTransB(), gemmOpAdaptor.getAlpha().convertToFloat(), + gemmOpAdaptor.getBeta().convertToFloat(), target, rewriter, gemmOp.getLoc()); + if (failed(result)) + return failure(); + rewriter.replaceOp(gemmOp, *result); return success(); } -void populateGemmPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { - patterns.insert(ctx); +void populateGemmPatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { + patterns.insert(ctx, target); } } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp new file mode 100644 index 0000000..97bcd1b --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.hpp @@ -0,0 +1,27 @@ +#pragma once + +#include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/Location.h" +#include "mlir/IR/Value.h" +#include "mlir/IR/PatternMatch.h" + +namespace onnx_mlir { +namespace spatial { +struct SpatialTargetInfo; +} + +mlir::FailureOr lowerGemmToSpatial( + mlir::Operation* diagnosticAnchor, + mlir::Value a, + mlir::Value b, + mlir::Value c, + mlir::RankedTensorType outputType, + bool transA, + bool transB, + float alpha, + float beta, + const spatial::SpatialTargetInfo& target, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + +} // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp index 3332073..d4e7e5b 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp @@ -11,6 +11,9 @@ #include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp" #include "src/Accelerators/PIM/Common/IR/TensorSliceUtils.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionProblem.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionMaterialization.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ContractionPlanning.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" @@ -118,7 +121,7 @@ static FailureOr collapseFragmentAssemblyBatchDims(Value value, auto fragmentStrides = blueprint ? blueprint.getFragmentStrides() : std::nullopt; if (!blueprint || !inputType || !storageType || !inputType.hasStaticShape() || !storageType.hasStaticShape() || inputType.getRank() <= 3 || resultType.getRank() != 3 || !blueprint.getFragments().empty() - || blueprint.getMode() != "fragment_assembly" || !operandIndices || !sourceOffsets || !fragmentStrides + || !spatial::isFragmentAssembly(blueprint.getMode()) || !operandIndices || !sourceOffsets || !fragmentStrides || storageType.getRank() != inputType.getRank() + 1) return failure(); if (blueprint.getIndexMap() == spatial::kContiguousRowMajorFragments @@ -176,7 +179,7 @@ static FailureOr collapseFragmentAssemblyBatchDims(Value value, collapsedStorage, ValueRange {}, blueprint.getLogicalLayoutAttr(), - rewriter.getStringAttr("fragmented"), + spatial::getFragmentedLayout(rewriter.getContext()), rewriter.getDenseI64ArrayAttr(offsets), rewriter.getDenseI64ArrayAttr(sizes), rewriter.getStringAttr("collapsed_fragments"), @@ -462,6 +465,7 @@ static FailureOr createBatchedVmmBatch(Value a, int64_t numOutRows, int64_t numKSlices, int64_t numOutHSlices, + int64_t xbarSize, PatternRewriter& rewriter, Location loc) { const int64_t laneCount = partialPiecesType.getDimSize(0); @@ -480,16 +484,16 @@ static FailureOr createBatchedVmmBatch(Value a, Value sliceLane = affineModConst(rewriter, loc, outerLane, numKSlices * numOutHSlices, anchorOp); Value kSlice = affineModConst(rewriter, loc, sliceLane, numKSlices, anchorOp); Value hSlice = affineFloorDivConst(rewriter, loc, sliceLane, numKSlices, anchorOp); - Value kOffset = affineMulConst(rewriter, loc, kSlice, crossbarSize.getValue(), anchorOp); - Value hOffset = affineMulConst(rewriter, loc, hSlice, crossbarSize.getValue(), anchorOp); + Value kOffset = affineMulConst(rewriter, loc, kSlice, xbarSize, anchorOp); + Value hOffset = affineMulConst(rewriter, loc, hSlice, xbarSize, anchorOp); auto aTileType = - RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, aType.getElementType()); + RankedTensorType::get({1, xbarSize}, aType.getElementType()); auto bTileType = RankedTensorType::get( - {static_cast(crossbarSize.getValue()), static_cast(crossbarSize.getValue())}, + {xbarSize, xbarSize}, bType.getElementType()); auto pieceType = - RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, partialPiecesType.getElementType()); + RankedTensorType::get({1, xbarSize}, partialPiecesType.getElementType()); Value aTile = extractBatchedATile( args.inputs.front(), aBatchShape, outputBatchShape, batch, row, kOffset, aTileType, rewriter, loc); @@ -522,8 +526,11 @@ static Value extractDynamicBatchedRowVector(Value matrix, {offsets, sizes, getUnitStrides(rewriter, 3)}); } -static int64_t chooseDynamicMatMulRowsPerLane(int64_t rows, int64_t reductionSize, int64_t columns) { - const int64_t crossbarElements = static_cast(crossbarSize.getValue() * crossbarSize.getValue()); +static int64_t chooseDynamicMatMulRowsPerLane(int64_t rows, + int64_t reductionSize, + int64_t columns, + int64_t xbarSize) { + const int64_t crossbarElements = xbarSize * xbarSize; const int64_t target = std::min(rows, ceilIntegerDivide(reductionSize * columns, crossbarElements)); int64_t rowsPerLane = 1; for (int64_t candidate = 2; candidate <= target; ++candidate) @@ -678,6 +685,7 @@ static Value extractBatchedReductionPiece(Value partialPiecesArg, int64_t numKSlices, int64_t numOutHSlices, int64_t numOutRows, + int64_t xbarSize, PatternRewriter& rewriter, Location loc) { Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); @@ -687,7 +695,8 @@ static Value extractBatchedReductionPiece(Value partialPiecesArg, Value batchAndHSlice = arith::AddIOp::create(rewriter, loc, batchOffset, hOffset); Value pieceOffset = arith::AddIOp::create(rewriter, loc, batchAndHSlice, kOffset); SmallVector offsets {pieceOffset, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}; - SmallVector sizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())}; + SmallVector sizes { + rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)}; return extractMixedSliceOrIdentity( rewriter, loc, partialPiecesArg, pieceType, {offsets, sizes, getUnitStrides(rewriter, 3)}); @@ -700,13 +709,24 @@ static Value reduceBatchedPartialPiecesForHSlice(Value partialPiecesArg, int64_t numKSlices, int64_t numOutHSlices, int64_t numOutRows, + int64_t xbarSize, PatternRewriter& rewriter, Location loc) { SmallVector activePieces; activePieces.reserve(numKSlices); for (int64_t kSlice = 0; kSlice < numKSlices; ++kSlice) activePieces.push_back(extractBatchedReductionPiece( - partialPiecesArg, batch, hSlice, kSlice, pieceType, numKSlices, numOutHSlices, numOutRows, rewriter, loc)); + partialPiecesArg, + batch, + hSlice, + kSlice, + pieceType, + numKSlices, + numOutHSlices, + numOutRows, + xbarSize, + rewriter, + loc)); while (activePieces.size() > 1) { SmallVector nextPieces; @@ -729,13 +749,14 @@ static FailureOr createBatchedReductionCompute(Value partialPieces, RankedTensorType paddedOutType, int64_t numBatches, int64_t numKSlices, + int64_t xbarSize, PatternRewriter& rewriter, Location loc) { auto computeOp = createSpatCompute<1>( rewriter, loc, TypeRange {outType}, {}, ValueRange {partialPieces}, [&](Value partialPiecesArg) -> LogicalResult { const int64_t numOutRows = outType.getDimSize(1); - const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(2), crossbarSize.getValue()); - auto pieceType = RankedTensorType::get({numOutRows, static_cast(crossbarSize.getValue())}, + const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(2), xbarSize); + auto pieceType = RankedTensorType::get({numOutRows, xbarSize}, partialPiecesType.getElementType()); Value outputInit = @@ -765,13 +786,22 @@ static FailureOr createBatchedReductionCompute(Value partialPieces, [&](OpBuilder&, Location hLoc, Value hSlice, ValueRange hIterArgs, SmallVectorImpl& hYielded) { Value outputAcc = hIterArgs.front(); Value reduced = reduceBatchedPartialPiecesForHSlice( - partialPiecesArg, batch, hSlice, pieceType, numKSlices, numOutHSlices, numOutRows, rewriter, hLoc); + partialPiecesArg, + batch, + hSlice, + pieceType, + numKSlices, + numOutHSlices, + numOutRows, + xbarSize, + rewriter, + hLoc); Value hOffset = affineMulConst( - rewriter, hLoc, hSlice, crossbarSize.getValue(), rewriter.getInsertionBlock()->getParentOp()); + rewriter, hLoc, hSlice, xbarSize, rewriter.getInsertionBlock()->getParentOp()); SmallVector outputOffsets {batch, rewriter.getIndexAttr(0), hOffset}; SmallVector outputSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(numOutRows), - rewriter.getIndexAttr(crossbarSize.getValue())}; + rewriter.getIndexAttr(xbarSize)}; Value next = tensor::InsertSliceOp::create( rewriter, hLoc, reduced, outputAcc, outputOffsets, outputSizes, getUnitStrides(rewriter, 3)) @@ -805,39 +835,45 @@ static FailureOr createBatchedReductionCompute(Value partialPieces, return computeOp->getResult(0); } -struct NormalizedMatMulInfo { +struct NormalizedMatMulInfo : ContractionProblem { + NormalizedMatMulInfo(RankedTensorType lhsType, + RankedTensorType rhsType, + RankedTensorType outType, + RankedTensorType normalizedLhsType, + RankedTensorType normalizedRhsType, + ContractionProblem problem, + bool lhsWasVector, + bool rhsWasVector) + : ContractionProblem(std::move(problem)), + lhsType(lhsType), + rhsType(rhsType), + outType(outType), + normalizedLhsType(normalizedLhsType), + normalizedRhsType(normalizedRhsType), + lhsWasVector(lhsWasVector), + rhsWasVector(rhsWasVector) {} + RankedTensorType lhsType; RankedTensorType rhsType; RankedTensorType outType; RankedTensorType normalizedLhsType; RankedTensorType normalizedRhsType; - SmallVector lhsBatchShape; - SmallVector rhsBatchShape; - SmallVector outputBatchShape; bool lhsWasVector; bool rhsWasVector; - int64_t lhsBatch; - int64_t rhsBatch; - int64_t batch; - int64_t m; - int64_t k; - int64_t n; }; -struct MatMulLoweringPlan { +struct MatMulLoweringPlan : ContractionProblem { + MatMulLoweringPlan(Value lhs, Value rhs, const NormalizedMatMulInfo& info) + : ContractionProblem(info), + lhs(lhs), + rhs(rhs), + lhsType(cast(lhs.getType())), + rhsType(cast(rhs.getType())) {} + Value lhs; Value rhs; RankedTensorType lhsType; RankedTensorType rhsType; - SmallVector lhsBatchShape; - SmallVector rhsBatchShape; - SmallVector outputBatchShape; - int64_t lhsBatch; - int64_t rhsBatch; - int64_t batch; - int64_t m; - int64_t k; - int64_t n; bool transposedResult; }; @@ -901,22 +937,31 @@ static FailureOr analyzeMatMulShape(ONNXMatMulOp matmulOp) return failure(); } - return NormalizedMatMulInfo {lhsType, - rhsType, - outType, - normalizedLhsType, - normalizedRhsType, - lhsBatchShape, - rhsBatchShape, - *outputBatchShape, - lhsWasVector, - rhsWasVector, - lhsBatch, - rhsBatch, - batch, - m, - k, - n}; + return NormalizedMatMulInfo( + lhsType, + rhsType, + outType, + normalizedLhsType, + normalizedRhsType, + ContractionProblem {lhsBatchShape, + rhsBatchShape, + *outputBatchShape, + lhsBatch, + rhsBatch, + batch, + m, + k, + n, + ContractionOrigin::MatMul, + lhsType.getElementType(), + rhsType.getElementType(), + outType.getElementType(), + false, + false, + lhsWasVector, + rhsWasVector}, + lhsWasVector, + rhsWasVector); } static MatMulLoweringPlan buildLoweringPlan(Value normalizedLhs, @@ -925,20 +970,8 @@ static MatMulLoweringPlan buildLoweringPlan(Value normalizedLhs, bool useTransposedForm, PatternRewriter& rewriter, Location loc) { - MatMulLoweringPlan plan {normalizedLhs, - normalizedRhs, - cast(normalizedLhs.getType()), - cast(normalizedRhs.getType()), - info.lhsBatchShape, - info.rhsBatchShape, - info.outputBatchShape, - info.lhsBatch, - info.rhsBatch, - info.batch, - info.m, - info.k, - info.n, - false}; + MatMulLoweringPlan plan(normalizedLhs, normalizedRhs, info); + plan.transposedResult = false; if (!useTransposedForm) return plan; @@ -1057,7 +1090,9 @@ struct MatMulToGemm : OpRewritePattern { }; struct MatMulBatchedToSpatialComputes : OpRewritePattern { - using OpRewritePattern::OpRewritePattern; + explicit MatMulBatchedToSpatialComputes(MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) + : OpRewritePattern(ctx), target(target) {} LogicalResult matchAndRewrite(ONNXMatMulOp matmulOp, PatternRewriter& rewriter) const override { auto shapeInfo = analyzeMatMulShape(matmulOp); @@ -1067,6 +1102,7 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { return failure(); Location loc = matmulOp.getLoc(); + const int64_t xbarSize = static_cast(target.matrixShape.rows); bool useTransposedForm = !shapeInfo->lhsWasVector && !shapeInfo->rhsWasVector && isCompileTimeComputable(matmulOp.getA()) && !isCompileTimeComputable(matmulOp.getB()); Value rhsRows = getLastTwoTransposeInput(matmulOp.getB()); @@ -1088,7 +1124,8 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { rhsStoredAsRows ? shapeInfo->k : shapeInfo->n, rewriter, loc); - MatMulLoweringPlan plan = buildLoweringPlan(lhs, rhs, *shapeInfo, useTransposedForm, rewriter, loc); + MatMulLoweringPlan plan = buildLoweringPlan( + lhs, rhs, *shapeInfo, useTransposedForm, rewriter, loc); plan.lhs = ensureBatchedTensor(plan.lhs, plan.lhsBatch, plan.m, plan.k, rewriter, loc); plan.rhs = ensureBatchedTensor(plan.rhs, @@ -1103,10 +1140,12 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { {plan.batch, plan.m, plan.n}, shapeInfo->outType.getElementType(), shapeInfo->outType.getEncoding()); if (isCompileTimeComputable(plan.rhs)) { - const int64_t numKSlices = ceilIntegerDivide(plan.k, crossbarSize.getValue()); - const int64_t numOutHSlices = ceilIntegerDivide(plan.n, crossbarSize.getValue()); - const int64_t paddedReductionSize = numKSlices * static_cast(crossbarSize.getValue()); - const int64_t paddedOutCols = numOutHSlices * static_cast(crossbarSize.getValue()); + ContractionPlan contractionPlan = makeContractionPlan( + plan, target, ContractionPlanKind::StaticTiled); + const int64_t numKSlices = contractionPlan.reductionSlices; + const int64_t numOutHSlices = contractionPlan.outputTiles; + const int64_t paddedReductionSize = numKSlices * xbarSize; + const int64_t paddedOutCols = numOutHSlices * xbarSize; auto paddedLhsType = RankedTensorType::get( {plan.lhsBatch, plan.m, paddedReductionSize}, plan.lhsType.getElementType(), plan.lhsType.getEncoding()); auto paddedRhsType = RankedTensorType::get( @@ -1117,10 +1156,11 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { auto paddedRhs = materializePaddedBatchedWeight(plan.rhs, plan.rhsBatchShape, plan.outputBatchShape, paddedRhsType, rewriter); if (succeeded(paddedRhs)) { - Value paddedLhs = createPaddedInputCompute(plan.lhs, paddedLhsType, rewriter, loc); - const int64_t laneCount = plan.batch * plan.m * numKSlices * numOutHSlices; + Value paddedLhs = materializePaddedContractionInput( + plan.lhs, paddedLhsType, rewriter, loc); + const int64_t laneCount = contractionPlan.laneCount; auto partialPiecesType = spatial::getGraphBatchPhysicalResultType( - laneCount, RankedTensorType::get({1, static_cast(crossbarSize.getValue())}, shapeInfo->outType.getElementType())); + laneCount, RankedTensorType::get({1, xbarSize}, shapeInfo->outType.getElementType())); auto batchOp = createBatchedVmmBatch(paddedLhs, *paddedRhs, paddedLhsType, @@ -1132,6 +1172,7 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { plan.m, numKSlices, numOutHSlices, + xbarSize, rewriter, loc); if (failed(batchOp)) @@ -1142,6 +1183,7 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { paddedOutType, plan.batch, numKSlices, + xbarSize, rewriter, loc); if (failed(result)) @@ -1164,8 +1206,11 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { ? shapeInfo->outType : directOutType; SmallVector blueprintBatchShape = !shapeInfo->lhsWasVector && !shapeInfo->rhsWasVector ? shapeInfo->outputBatchShape : SmallVector {plan.batch}; - const int64_t rowsPerLane = chooseDynamicMatMulRowsPerLane(plan.m, plan.k, plan.n); - const int64_t laneCount = plan.batch * plan.m / rowsPerLane; + const int64_t rowsPerLane = chooseDynamicMatMulRowsPerLane(plan.m, plan.k, plan.n, xbarSize); + ContractionPlan contractionPlan = makeContractionPlan( + plan, target, ContractionPlanKind::GroupedRowDynamicVVD, + /*laneCount=*/plan.batch * plan.m / rowsPerLane, rowsPerLane); + const int64_t laneCount = contractionPlan.laneCount; SmallVector fragmentShape(blueprintType.getRank(), 1); fragmentShape[fragmentShape.size() - 2] = rowsPerLane; fragmentShape.back() = plan.n; @@ -1210,6 +1255,8 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { rewriter.eraseOp(foldedMultiply); return success(); } + + const spatial::SpatialTargetInfo& target; }; struct TransposedRhsMatMulToSpatial : MatMulBatchedToSpatialComputes { @@ -1224,12 +1271,17 @@ struct TransposedRhsMatMulToSpatial : MatMulBatchedToSpatialComputes { } // namespace -void populateMatMulFusionPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { - patterns.add(ctx); +void populateMatMulFusionPatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { + patterns.add(ctx, target); } -void populateMatMulRewritePatterns(RewritePatternSet& patterns, MLIRContext* ctx) { - patterns.insert(ctx); +void populateMatMulRewritePatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { + patterns.insert(ctx); + patterns.insert(ctx, target); } } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ReduceMean.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ReduceMean.cpp index 8904feb..359c1d3 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ReduceMean.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/ReduceMean.cpp @@ -280,12 +280,12 @@ static FailureOr buildReduceMeanKeepdimsBlueprint( SmallVector fragmentStrides(fragmentOffsets.size(), 1); return spatial::SpatBlueprintOp::create( rewriter, loc, keepdimsType, batchValue, ValueRange {}, - rewriter.getStringAttr("nchw"), - rewriter.getStringAttr("fragmented"), + spatial::getNCHWLayout(rewriter.getContext()), + spatial::getFragmentedLayout(rewriter.getContext()), rewriter.getDenseI64ArrayAttr(fragmentOffsets), rewriter.getDenseI64ArrayAttr(fragmentSizes), rewriter.getStringAttr("reduce_mean_keepdims_fragments"), - rewriter.getStringAttr("fragment_assembly"), + spatial::getFragmentAssemblyMode(rewriter.getContext()), rewriter.getDenseI64ArrayAttr(operandIndices), rewriter.getDenseI64ArrayAttr(sourceSlots), rewriter.getDenseI64ArrayAttr(sourceOffsets), diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp index 26692de..0fce40e 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Pool.cpp @@ -14,8 +14,8 @@ #include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp" #include "src/Accelerators/PIM/Common/PimCommon.hpp" -#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" @@ -32,8 +32,10 @@ static Value materializeTileTensor(PatternRewriter& rewriter, Location loc, Valu return insertStaticSlice(rewriter, loc, tile, empty, getZeroOffsets(rewriter, tileType.getRank())); } -static Value -createPoolFillElement(ConversionPatternRewriter& rewriter, Location loc, Type elementType, bool useMinimumValue) { +static Value createPoolFillElement(OpBuilder& rewriter, + Location loc, + Type elementType, + bool useMinimumValue) { Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); if (!useMinimumValue) return getOrCreateConstant(rewriter, anchorOp, rewriter.getZeroAttr(elementType), elementType); @@ -51,7 +53,7 @@ createPoolFillElement(ConversionPatternRewriter& rewriter, Location loc, Type el llvm_unreachable("unsupported pool element type"); } -static Value createPoolFillTensor(ConversionPatternRewriter& rewriter, +static Value createPoolFillTensor(OpBuilder& rewriter, Location loc, RankedTensorType tensorType, bool useMinimumValue) { @@ -59,16 +61,15 @@ static Value createPoolFillTensor(ConversionPatternRewriter& rewriter, return tensor::SplatOp::create(rewriter, loc, tensorType, fillElement); } -template -static Value createPaddedPoolInput(ConversionPatternRewriter& rewriter, +static Value createPaddedPoolInput(OpBuilder& rewriter, Location loc, - PoolOp poolOp, Value input, RankedTensorType inputType, int64_t padTop, int64_t padLeft, int64_t padBottom, - int64_t padRight) { + int64_t padRight, + bool useMinimumValue) { if (padTop == 0 && padLeft == 0 && padBottom == 0 && padRight == 0) return input; @@ -90,8 +91,8 @@ static Value createPaddedPoolInput(ConversionPatternRewriter& rewriter, padBlock->addArgument(rewriter.getIndexType(), loc); padOp.getRegion().push_back(padBlock); rewriter.setInsertionPointToStart(padBlock); - Value padValue = - createPoolFillElement(rewriter, loc, inputType.getElementType(), std::is_same_v); + Value padValue = createPoolFillElement( + rewriter, loc, inputType.getElementType(), useMinimumValue); tensor::YieldOp::create(rewriter, loc, padValue); rewriter.setInsertionPointAfter(padOp); return padOp.getResult(); @@ -160,7 +161,10 @@ struct PoolToSpatialCompute; template struct PoolToSpatialComputeBase : public OpConversionPattern { - using OpConversionPattern::OpConversionPattern; + PoolToSpatialComputeBase(MLIRContext* ctx, const spatial::SpatialTargetInfo& target) + : OpConversionPattern(ctx), target(target) {} + + const spatial::SpatialTargetInfo& target; LogicalResult matchAndRewrite(PoolOp poolOp, PoolOpAdaptor adaptor, ConversionPatternRewriter& rewriter) const final { Location loc = poolOp.getLoc(); @@ -241,7 +245,7 @@ struct PoolToSpatialComputeBase : public OpConversionPattern { rewriter.getDenseI64ArrayAttr({padTop, padLeft, padBottom, padRight}), rewriter.getDenseI64ArrayAttr({strideHeight, strideWidth}), rewriter.getDenseI64ArrayAttr({dilationHeight, dilationWidth}), - rewriter.getStringAttr("nchw")); + spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(poolOp, plan.getResult()); return success(); } @@ -251,12 +255,12 @@ struct PoolToSpatialComputeBase : public OpConversionPattern { && dilationHeight == 1 && dilationWidth == 1 && padTop == 0 && padLeft == 0 && padBottom == 0 && padRight == 0) { auto plan = spatial::SpatGlobalAveragePoolPlanOp::create( - rewriter, loc, outType, x, rewriter.getStringAttr("nchw")); + rewriter, loc, outType, x, spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(poolOp, plan.getResult()); return success(); } - const int64_t xbarSize = static_cast(crossbarSize.getValue()); + const int64_t xbarSize = static_cast(target.matrixShape.rows); const int64_t channelTileCount = (channels + xbarSize - 1) / xbarSize; const int64_t outputPatchCount = batchSize * outputHeight * outputWidth; const bool countIncludePad = [&]() { @@ -292,7 +296,9 @@ struct PoolToSpatialComputeBase : public OpConversionPattern { auto computeOp = createSpatCompute(rewriter, loc, outType, {}, ValueRange {x}, [&](Value xArg) -> LogicalResult { Value paddedInput = - createPaddedPoolInput(rewriter, loc, poolOp, xArg, xType, padTop, padLeft, padBottom, padRight); + createPaddedPoolInput(rewriter, loc, xArg, xType, padTop, padLeft, + padBottom, padRight, + std::is_same_v); Value pooledOutputInit = tensor::EmptyOp::create(rewriter, loc, outType.getShape(), outType.getElementType()); Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); @@ -424,7 +430,8 @@ struct PoolToSpatialCompute } // namespace -LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) { +LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp, + const spatial::SpatialTargetInfo&) { auto inputType = dyn_cast(planOp.getInput().getType()); auto outputType = dyn_cast(planOp.getOutput().getType()); if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape()) @@ -439,6 +446,118 @@ LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) return success(); } +FailureOr lowerDenseMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, + const spatial::SpatialTargetInfo& target, + PatternRewriter& rewriter) { + auto inputType = dyn_cast(planOp.getInput().getType()); + auto outputType = dyn_cast(planOp.getOutput().getType()); + if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape() + || inputType.getRank() != 4 || outputType.getRank() != 4) + return planOp.emitOpError("dense MaxPool lowering requires static rank-4 tensors"), failure(); + + auto kernel = planOp.getKernelShape(); + auto pads = planOp.getPads(); + auto strides = planOp.getStrides(); + auto dilations = planOp.getDilations(); + if (kernel.size() != 2 || pads.size() != 4 || strides.size() != 2 || dilations.size() != 2 + || llvm::any_of(kernel, [](int64_t value) { return value <= 0; }) + || llvm::any_of(strides, [](int64_t value) { return value <= 0; }) + || llvm::any_of(dilations, [](int64_t value) { return value <= 0; }) + || llvm::any_of(pads, [](int64_t value) { return value < 0; })) + return planOp.emitOpError("dense MaxPool lowering requires valid kernel, padding, stride, and dilation attributes"), + failure(); + + const int64_t batchSize = inputType.getDimSize(0); + const int64_t channels = inputType.getDimSize(1); + const int64_t outputHeight = outputType.getDimSize(2); + const int64_t outputWidth = outputType.getDimSize(3); + const int64_t tileWidth = std::max(1, target.matrixShape.rows); + const int64_t channelTileCount = (channels + tileWidth - 1) / tileWidth; + const int64_t outputPatchCount = batchSize * outputHeight * outputWidth; + + auto compute = createSpatCompute<1>( + rewriter, planOp.getLoc(), outputType, {}, planOp.getInput(), + [&](Value input) -> LogicalResult { + Value paddedInput = createPaddedPoolInput( + rewriter, planOp.getLoc(), input, inputType, + pads[0], pads[1], pads[2], pads[3], /*useMinimumValue=*/true); + Value outputInit = tensor::EmptyOp::create( + rewriter, planOp.getLoc(), outputType.getShape(), outputType.getElementType()); + Operation* anchor = rewriter.getInsertionBlock()->getParentOp(); + Value zero = getOrCreateIndexConstant(rewriter, anchor, 0); + Value one = getOrCreateIndexConstant(rewriter, anchor, 1); + Value patchCount = getOrCreateIndexConstant(rewriter, anchor, outputPatchCount); + Value pixelsPerBatch = getOrCreateIndexConstant( + rewriter, anchor, outputHeight * outputWidth); + Value outputWidthValue = getOrCreateIndexConstant(rewriter, anchor, outputWidth); + Value strideHeight = getOrCreateIndexConstant(rewriter, anchor, strides[0]); + Value strideWidth = getOrCreateIndexConstant(rewriter, anchor, strides[1]); + + auto loop = buildNormalizedScfFor( + rewriter, planOp.getLoc(), zero, patchCount, one, ValueRange {outputInit}, + [&](OpBuilder&, Location loc, Value patch, ValueRange iterArgs, + SmallVectorImpl& yielded) { + Value batch = arith::DivUIOp::create(rewriter, loc, patch, pixelsPerBatch); + Value batchPatch = arith::RemUIOp::create(rewriter, loc, patch, pixelsPerBatch); + Value outputRow = arith::DivUIOp::create(rewriter, loc, batchPatch, outputWidthValue); + Value outputColumn = arith::RemUIOp::create(rewriter, loc, batchPatch, outputWidthValue); + Value windowRow = arith::MulIOp::create(rewriter, loc, outputRow, strideHeight); + Value windowColumn = arith::MulIOp::create(rewriter, loc, outputColumn, strideWidth); + Value updated = iterArgs.front(); + + for (int64_t tile = 0; tile < channelTileCount; ++tile) { + const int64_t tileChannels = std::min(tileWidth, channels - tile * tileWidth); + auto tileType = RankedTensorType::get( + {1, tileChannels, 1, 1}, outputType.getElementType()); + Value reduced = createPoolFillTensor( + rewriter, loc, tileType, /*useMinimumValue=*/true); + for (int64_t kernelRow = 0; kernelRow < kernel[0]; ++kernelRow) { + Value sourceRow = windowRow; + if (kernelRow * dilations[0] != 0) + sourceRow = arith::AddIOp::create( + rewriter, loc, sourceRow, + getOrCreateIndexConstant(rewriter, anchor, kernelRow * dilations[0])); + for (int64_t kernelColumn = 0; kernelColumn < kernel[1]; ++kernelColumn) { + Value sourceColumn = windowColumn; + if (kernelColumn * dilations[1] != 0) + sourceColumn = arith::AddIOp::create( + rewriter, loc, sourceColumn, + getOrCreateIndexConstant(rewriter, anchor, kernelColumn * dilations[1])); + Value point = tensor::ExtractSliceOp::create( + rewriter, loc, tileType, paddedInput, + SmallVector { + batch, rewriter.getIndexAttr(tile * tileWidth), sourceRow, sourceColumn}, + SmallVector { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + point = materializeTileTensor(rewriter, loc, point); + reduced = spatial::SpatVMaxOp::create( + rewriter, loc, tileType, reduced, point); + } + } + updated = tensor::InsertSliceOp::create( + rewriter, loc, reduced, updated, + SmallVector { + batch, rewriter.getIndexAttr(tile * tileWidth), outputRow, outputColumn}, + SmallVector { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + } + yielded.push_back(updated); + return success(); + }); + if (failed(loop)) + return failure(); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), loop->results.front()); + return success(); + }); + if (failed(compute)) + return failure(); + return compute->getResult(0); +} + static Value createClampedPoolIndexTable(PatternRewriter& rewriter, Operation* anchorOp, int64_t outputSize, @@ -497,8 +616,9 @@ static Value extractPoolIndex(PatternRewriter& rewriter, FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo& target, PatternRewriter& rewriter) { - if (failed(canLowerMaxPoolPlanToRowStrip(planOp))) + if (failed(canLowerMaxPoolPlanToRowStrip(planOp, target))) return failure(); Location loc = planOp.getLoc(); @@ -590,8 +710,8 @@ FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, rewriter.getIndexAttr(1), rewriter.getIndexAttr(inputWidth)}, getUnitStrides(rewriter, 4)); - inputRows.push_back(ONNXTransposeOp::create( - rewriter, loc, inputFragmentType, nchw, rewriter.getI64ArrayAttr({0, 2, 3, 1}))); + inputRows.push_back(createLinalgTranspose( + nchw, inputFragmentType, {0, 2, 3, 1}, rewriter, loc)); } } @@ -685,7 +805,8 @@ FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, return batch->getResult(0); } -LogicalResult canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp) { +LogicalResult canLowerGlobalAveragePoolPlanToRowStrip( + spatial::SpatGlobalAveragePoolPlanOp planOp, const spatial::SpatialTargetInfo&) { auto inputType = dyn_cast(planOp.getInput().getType()); auto outputType = dyn_cast(planOp.getOutput().getType()); if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape()) @@ -697,10 +818,87 @@ LogicalResult canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAverage return success(); } +FailureOr lowerDenseGlobalAveragePoolPlan( + spatial::SpatGlobalAveragePoolPlanOp planOp, + const spatial::SpatialTargetInfo& target, + PatternRewriter& rewriter) { + auto inputType = dyn_cast(planOp.getInput().getType()); + auto outputType = dyn_cast(planOp.getOutput().getType()); + if (!inputType || !outputType || !inputType.hasStaticShape() + || !outputType.hasStaticShape() || inputType.getRank() != 4 + || outputType.getRank() != 4 || inputType.getDimSize(0) != 1 + || outputType.getDimSize(0) != 1 || inputType.getDimSize(1) != outputType.getDimSize(1) + || outputType.getDimSize(2) != 1 || outputType.getDimSize(3) != 1) + return planOp.emitOpError("dense global AveragePool lowering requires static rank-4 floating-point tensors"), + failure(); + auto elementType = dyn_cast(inputType.getElementType()); + if (!elementType) + return planOp.emitOpError("dense global AveragePool lowering requires floating-point tensors"), + failure(); + + const int64_t channels = inputType.getDimSize(1); + const int64_t height = inputType.getDimSize(2); + const int64_t width = inputType.getDimSize(3); + const int64_t tileWidth = std::max(1, target.matrixShape.rows); + const int64_t channelTileCount = (channels + tileWidth - 1) / tileWidth; + const double scaleValue = 1.0 / static_cast(height * width); + + auto compute = createSpatCompute<1>( + rewriter, planOp.getLoc(), outputType, {}, planOp.getInput(), + [&](Value input) -> LogicalResult { + Value output = tensor::EmptyOp::create( + rewriter, planOp.getLoc(), outputType.getShape(), outputType.getElementType()); + Operation* anchor = rewriter.getInsertionBlock()->getParentOp(); + for (int64_t tile = 0; tile < channelTileCount; ++tile) { + const int64_t tileChannels = std::min(tileWidth, channels - tile * tileWidth); + auto tileType = RankedTensorType::get( + {1, tileChannels, 1, 1}, outputType.getElementType()); + Value reduced = createPoolFillTensor( + rewriter, planOp.getLoc(), tileType, /*useMinimumValue=*/false); + for (int64_t row = 0; row < height; ++row) { + for (int64_t column = 0; column < width; ++column) { + Value point = tensor::ExtractSliceOp::create( + rewriter, planOp.getLoc(), tileType, input, + SmallVector { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(tile * tileWidth), + rewriter.getIndexAttr(row), rewriter.getIndexAttr(column)}, + SmallVector { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + point = materializeTileTensor(rewriter, planOp.getLoc(), point); + reduced = spatial::SpatVAddOp::create( + rewriter, planOp.getLoc(), tileType, reduced, point); + } + } + auto scaleAttr = DenseElementsAttr::get( + tileType, rewriter.getFloatAttr(elementType, scaleValue)); + Value scale = getOrCreateConstant(rewriter, anchor, scaleAttr, tileType); + reduced = spatial::SpatVMulOp::create( + rewriter, planOp.getLoc(), tileType, reduced, scale); + output = tensor::InsertSliceOp::create( + rewriter, planOp.getLoc(), reduced, output, + SmallVector { + rewriter.getIndexAttr(0), rewriter.getIndexAttr(tile * tileWidth), + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, + SmallVector { + rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels), + rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 4)); + } + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), output); + return success(); + }); + if (failed(compute)) + return failure(); + return compute->getResult(0); +} + FailureOr lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo& target, PatternRewriter& rewriter) { - if (failed(canLowerGlobalAveragePoolPlanToRowStrip(planOp))) + if (failed(canLowerGlobalAveragePoolPlanToRowStrip(planOp, target))) return failure(); Location loc = planOp.getLoc(); @@ -777,8 +975,8 @@ FailureOr lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePo rewriter.getIndexAttr(1), rewriter.getIndexAttr(width)}, getUnitStrides(rewriter, 4)); - fragment = ONNXTransposeOp::create( - rewriter, loc, inputFragmentType, nchw, rewriter.getI64ArrayAttr({0, 2, 3, 1})); + fragment = createLinalgTranspose( + nchw, inputFragmentType, {0, 2, 3, 1}, rewriter, loc); } for (int64_t column = 0; column < width; ++column) { Value point = tensor::ExtractSliceOp::create( @@ -811,9 +1009,11 @@ FailureOr lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePo return batch->getResult(0); } -void populatePoolPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { - patterns.insert>(ctx); - patterns.insert>(ctx); +void populatePoolPatterns(RewritePatternSet& patterns, + MLIRContext* ctx, + const spatial::SpatialTargetInfo& target) { + patterns.insert>(ctx, target); + patterns.insert>(ctx, target); } } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Relu.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Relu.cpp index 008865e..6866b7e 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Relu.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/NN/Relu.cpp @@ -17,7 +17,7 @@ struct ReluToSpatialCompute : OpConversionPattern { Location loc = reluOp.getLoc(); Type resultType = reluOp.getResult().getType(); auto reluPlan = spatial::SpatReluPlanOp::create( - rewriter, loc, resultType, adaptor.getX(), rewriter.getStringAttr("nchw")); + rewriter, loc, resultType, adaptor.getX(), spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(reluOp, reluPlan.getResult()); return success(); } diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Concat.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Concat.cpp index a23076e..07a1a2d 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Concat.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Concat.cpp @@ -32,7 +32,8 @@ struct Concat : public OpConversionPattern { return type && type.hasStaticShape() && type.getRank() == 4; })) { rewriter.replaceOpWithNewOp( - maxpoolOp, resultType, inputs, rewriter.getI64IntegerAttr(axis), rewriter.getStringAttr("nchw")); + maxpoolOp, resultType, inputs, rewriter.getI64IntegerAttr(axis), + spatial::getNCHWLayout(rewriter.getContext())); return success(); } diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp index 9fdb01b..3a86733 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Flatten.cpp @@ -4,7 +4,6 @@ #include "llvm/ADT/SmallVector.h" #include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp" -#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp" @@ -119,7 +118,8 @@ struct RowStripFlattenAnalysis { DenseElementsAttr weight; }; -static FailureOr analyzeRowStripFlatten(spatial::SpatGraphCompute flattenOp) { +static FailureOr analyzeRowStripFlatten( + spatial::SpatGraphCompute flattenOp, const spatial::SpatialTargetInfo& target) { if (flattenOp.getWeights().size() != 0 || flattenOp.getInputs().size() != 1 || flattenOp.getOutputs().size() != 1) return failure(); @@ -130,7 +130,7 @@ static FailureOr analyzeRowStripFlatten(spatial::SpatGr || resultType.getDimSize(0) != 1 || resultType.getDimSize(1) != sourceType.getNumElements()) return failure(); const int64_t channels = sourceType.getDimSize(1); - const int64_t xbarDim = static_cast(crossbarSize.getValue()); + const int64_t xbarDim = static_cast(target.matrixShape.rows); if (channels > xbarDim && channels % xbarDim != 0) return failure(); @@ -162,14 +162,16 @@ static FailureOr analyzeRowStripFlatten(spatial::SpatGr void populateFlattenPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.add(ctx); } -LogicalResult canLowerFlattenFromRowStrip(spatial::SpatGraphCompute flattenOp) { - return succeeded(analyzeRowStripFlatten(flattenOp)) ? success() : failure(); +LogicalResult canLowerFlattenFromRowStrip(spatial::SpatGraphCompute flattenOp, + const spatial::SpatialTargetInfo& target) { + return succeeded(analyzeRowStripFlatten(flattenOp, target)) ? success() : failure(); } LogicalResult lowerFlattenFromRowStrip(const RowStripPhysicalValue& input, spatial::SpatGraphCompute flattenOp, + const spatial::SpatialTargetInfo& target, PatternRewriter& rewriter) { - FailureOr analysis = analyzeRowStripFlatten(flattenOp); + FailureOr analysis = analyzeRowStripFlatten(flattenOp, target); if (failed(analysis)) return failure(); auto storageType = dyn_cast(input.storage.getType()); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Resize.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Resize.cpp index 764ea71..9401dc6 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Resize.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Resize.cpp @@ -183,7 +183,7 @@ struct Resize : OpConversionPattern { return rewriter.notifyMatchFailure(resizeOp, "resize lowering requires positive static dimensions."); auto plan = spatial::SpatResizeNearestPlanOp::create( - rewriter, resizeOp.getLoc(), resultType, adaptor.getX(), rewriter.getStringAttr("nchw")); + rewriter, resizeOp.getLoc(), resultType, adaptor.getX(), spatial::getNCHWLayout(rewriter.getContext())); rewriter.replaceOp(resizeOp, plan.getResult()); return success(); } @@ -192,7 +192,8 @@ struct Resize : OpConversionPattern { } // namespace LogicalResult canLowerResizeNearestPlanToRowStrip( - spatial::SpatResizeNearestPlanOp planOp) { + spatial::SpatResizeNearestPlanOp planOp, + const spatial::SpatialTargetInfo&) { auto inputType = dyn_cast(planOp.getInput().getType()); auto outputType = dyn_cast(planOp.getOutput().getType()); return success(inputType && outputType && inputType.hasStaticShape() @@ -204,6 +205,7 @@ LogicalResult canLowerResizeNearestPlanToRowStrip( FailureOr lowerSelectedResizeNearestPlan( spatial::SpatResizeNearestPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo&, PatternRewriter& rewriter) { auto inputType = cast(planOp.getInput().getType()); auto outputType = cast(planOp.getOutput().getType()); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Transpose.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Transpose.cpp index 2088717..b77a9dc 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Transpose.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Tensor/Transpose.cpp @@ -70,7 +70,7 @@ static FailureOr transposeFragmentAssemblyBlueprint(spatial::SpatBlueprin auto sourceOffsets = blueprint.getFragmentSourceOffsets(); auto fragmentStrides = blueprint.getFragmentStrides(); if (!storageType || !storageType.hasStaticShape() || !resultType.hasStaticShape() - || !blueprint.getFragments().empty() || blueprint.getMode() != "fragment_assembly" + || !blueprint.getFragments().empty() || !spatial::isFragmentAssembly(blueprint.getMode()) || !blueprint.getFragmentOperandIndices() || !sourceOffsets || !fragmentStrides || llvm::any_of(*sourceOffsets, [](int64_t offset) { return offset != 0; }) || storageType.getRank() != resultType.getRank() + 1) @@ -113,7 +113,7 @@ static FailureOr transposeFragmentAssemblyBlueprint(spatial::SpatBlueprin *mapped, ValueRange {}, blueprint.getLogicalLayoutAttr(), - rewriter.getStringAttr("fragmented"), + spatial::getFragmentedLayout(rewriter.getContext()), rewriter.getDenseI64ArrayAttr(offsets), rewriter.getDenseI64ArrayAttr(sizes), rewriter.getStringAttr("permuted_fragments"), diff --git a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp index 7a952a5..2060508 100644 --- a/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp @@ -15,38 +15,50 @@ mlir::FailureOr lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp, std::optional rowStripInput, bool emitRowStripLayout, + const spatial::SpatialTargetInfo& target, mlir::PatternRewriter& rewriter); -mlir::LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp); -mlir::LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp); +mlir::LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp, + const spatial::SpatialTargetInfo& target); +mlir::LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp, + const spatial::SpatialTargetInfo& target); mlir::LogicalResult canLowerResizeNearestPlanToRowStrip( - spatial::SpatResizeNearestPlanOp planOp); + spatial::SpatResizeNearestPlanOp planOp, const spatial::SpatialTargetInfo& target); mlir::FailureOr lowerSelectedResizeNearestPlan( spatial::SpatResizeNearestPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo& target, mlir::PatternRewriter& rewriter); -mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp); +mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp, + const spatial::SpatialTargetInfo& target); + +mlir::FailureOr +lowerDenseMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, + const spatial::SpatialTargetInfo& target, + mlir::PatternRewriter& rewriter); mlir::FailureOr lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo& target, mlir::PatternRewriter& rewriter); mlir::LogicalResult -canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp); +canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp, + const spatial::SpatialTargetInfo& target); + +mlir::FailureOr +lowerDenseGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp, + const spatial::SpatialTargetInfo& target, + mlir::PatternRewriter& rewriter); mlir::FailureOr lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp, std::optional rowStripInput, + const spatial::SpatialTargetInfo& target, mlir::PatternRewriter& rewriter); -mlir::LogicalResult canLowerFlattenFromRowStrip(spatial::SpatGraphCompute flattenOp); - -mlir::LogicalResult lowerFlattenFromRowStrip(const RowStripPhysicalValue& input, - spatial::SpatGraphCompute flattenOp, - mlir::PatternRewriter& rewriter); - } // namespace onnx_mlir diff --git a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutCapabilities.cpp b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutCapabilities.cpp new file mode 100644 index 0000000..e96a8a5 --- /dev/null +++ b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutCapabilities.cpp @@ -0,0 +1,133 @@ +#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 diff --git a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp index 0b39303..56eb1e6 100644 --- a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp @@ -6,358 +6,260 @@ #include "Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp" #include "src/Accelerators/PIM/Common/PimCommon.hpp" -#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp" -#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Accelerators/PIM/Pass/PIMPasses.h" +#include + using namespace mlir; namespace onnx_mlir { namespace { -static constexpr StringLiteral kLogicalLayout = "nchw"; -static constexpr StringLiteral kDenseLayout = "dense_nchw"; -static constexpr StringLiteral kRowStripLayout = "nhwc_row_strip"; +using LayoutMap = llvm::DenseMap; -enum class SelectedLayout { - DenseNchw, - PixelMajorRowStrip, -}; - -static SelectedLayout getSelectedLayout(llvm::DenseMap& layouts, Value value) { - auto it = layouts.find(value); - return it == layouts.end() ? SelectedLayout::DenseNchw : it->second; +static spatial::PhysicalLayout getSelectedLayout(const LayoutMap& layouts, Value value) { + if (auto it = layouts.find(value); it != layouts.end()) + return it->second; + if (auto materialize = value.getDefiningOp()) + return materialize.getTargetPhysicalLayout(); + if (auto blueprint = value.getDefiningOp()) + return blueprint.getPhysicalLayout(); + return spatial::PhysicalLayout::DenseNCHW; } -static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap& layouts) { - if (auto reluPlan = dyn_cast(user)) - return getSelectedLayout(layouts, reluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto siluPlan = dyn_cast(user)) - return getSelectedLayout(layouts, siluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto resizePlan = dyn_cast(user)) - return getSelectedLayout(layouts, resizePlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto biasAddPlan = dyn_cast(user)) - return getSelectedLayout(layouts, biasAddPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto addPlan = dyn_cast(user)) - return getSelectedLayout(layouts, addPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto concatPlan = dyn_cast(user)) - return getSelectedLayout(layouts, concatPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto convPlan = dyn_cast(user)) - return getSelectedLayout(layouts, convPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto maxPoolPlan = dyn_cast(user)) - return getSelectedLayout(layouts, maxPoolPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto averagePoolPlan = dyn_cast(user)) - return getSelectedLayout(layouts, averagePoolPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; - if (auto flattenCompute = dyn_cast(user)) - return succeeded(canLowerFlattenFromRowStrip(flattenCompute)); - return false; +static SmallVector getOperandLayouts( + Operation* op, const LayoutMap& layouts) { + SmallVector operandLayouts; + operandLayouts.reserve(op->getNumOperands()); + for (Value operand : op->getOperands()) + operandLayouts.push_back(getSelectedLayout(layouts, operand)); + return operandLayouts; } -static bool allUsersCanHandleRowStrip(Value value, llvm::DenseMap& layouts) { - for (Operation* user : value.getUsers()) { - if (usesSelectedRowStrip(user, layouts)) +static FailureOr> getAlternatives( + Operation* op, const LayoutMap& layouts, const spatial::SpatialTargetInfo& target) { + auto capability = dyn_cast(op); + if (!capability) + return failure(); + SmallVector alternatives = + capability.getLayoutAlternatives(target, getOperandLayouts(op, layouts)); + if (alternatives.empty()) + return op->emitOpError("does not advertise a legal Spatial layout alternative"), failure(); + for (const spatial::LayoutAlternative& alternative : alternatives) + if (alternative.operandLayouts.size() != op->getNumOperands()) + return op->emitOpError("advertises a layout alternative with the wrong operand count"), failure(); + return alternatives; +} + +static unsigned findCurrentAlternative( + Operation* op, ArrayRef alternatives, + spatial::PhysicalLayout selectedResult) { + for (auto [index, alternative] : llvm::enumerate(alternatives)) + if (alternative.resultLayout == selectedResult) + return index; + return 0; +} + +static int64_t alternativeCost(Operation* op, + const spatial::LayoutAlternative& alternative, + const LayoutMap& layouts, + const LayoutMap& selectedResults, + const spatial::SpatialTargetInfo& target) { + int64_t cost = alternative.intrinsicCost; + SmallVector operandLayouts = getOperandLayouts(op, layouts); + for (auto [actual, required] : llvm::zip(operandLayouts, alternative.operandLayouts)) + cost += actual != required; + + Value result = op->getResult(0); + for (OpOperand& use : result.getUses()) { + auto user = dyn_cast(use.getOwner()); + if (!user) { + if (alternative.resultLayout != spatial::PhysicalLayout::DenseNCHW) { + auto flatten = dyn_cast(use.getOwner()); + if (!flatten || failed(canLowerFlattenFromRowStrip(flatten, target))) + ++cost; + } continue; - // Dense-only users must be materialized explicitly. - continue; - } - return true; -} - -static bool canConsumeRowStripAsUser(Operation* user) { - if (isa(user)) - return true; - if (auto resizePlan = dyn_cast(user)) - return succeeded(canLowerResizeNearestPlanToRowStrip(resizePlan)); - if (auto biasAddPlan = dyn_cast(user)) { - auto resultType = dyn_cast(biasAddPlan.getOutput().getType()); - return resultType && isSupportedBiasAddValue(biasAddPlan.getBias(), resultType); - } - if (isa(user)) - return true; - if (isa(user)) - return true; - if (auto convPlan = dyn_cast(user)) - return succeeded(canConsumeAndProduceRowStrip(convPlan)); - if (auto maxPoolPlan = dyn_cast(user)) - return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan)); - if (auto averagePoolPlan = dyn_cast(user)) - return succeeded(canLowerGlobalAveragePoolPlanToRowStrip(averagePoolPlan)); - return false; -} - -static bool hasRowStripConsumer(Value value) { - for (Operation* user : value.getUsers()) - if (canConsumeRowStripAsUser(user)) - return true; - return false; -} - -static bool canSelectConvRowStrip(spatial::SpatConv2DPlanOp convPlan, - llvm::DenseMap& layouts) { - SelectedLayout inputLayout = getSelectedLayout(layouts, convPlan.getInput()); - if (inputLayout == SelectedLayout::PixelMajorRowStrip) - return succeeded(canConsumeAndProduceRowStrip(convPlan)); - return succeeded(canLowerConvPlanToRowStrip(convPlan)); -} - -static SelectedLayout chooseConvLayout(spatial::SpatConv2DPlanOp convPlan, - llvm::DenseMap& layouts) { - if (!canSelectConvRowStrip(convPlan, layouts)) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(convPlan.getResult(), layouts)) - return SelectedLayout::DenseNchw; - return SelectedLayout::PixelMajorRowStrip; -} - -static SelectedLayout chooseActivationLayout(Value input, - Value result, - llvm::DenseMap& layouts) { - if (getSelectedLayout(layouts, input) != SelectedLayout::PixelMajorRowStrip) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(result, layouts)) - return SelectedLayout::DenseNchw; - return SelectedLayout::PixelMajorRowStrip; -} - -static SelectedLayout chooseResizeLayout( - spatial::SpatResizeNearestPlanOp resizePlan, - llvm::DenseMap& layouts) { - return getSelectedLayout(layouts, resizePlan.getInput()) == SelectedLayout::PixelMajorRowStrip - && succeeded(canLowerResizeNearestPlanToRowStrip(resizePlan)) - ? SelectedLayout::PixelMajorRowStrip : SelectedLayout::DenseNchw; -} - -static SelectedLayout chooseBiasAddLayout(spatial::SpatBiasAddPlanOp biasAddPlan, - llvm::DenseMap& layouts) { - if (getSelectedLayout(layouts, biasAddPlan.getInput()) != SelectedLayout::PixelMajorRowStrip) - return SelectedLayout::DenseNchw; - auto resultType = dyn_cast(biasAddPlan.getOutput().getType()); - if (!resultType || !isSupportedBiasAddValue(biasAddPlan.getBias(), resultType)) - return SelectedLayout::DenseNchw; - if (!hasRowStripConsumer(biasAddPlan.getResult())) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(biasAddPlan.getResult(), layouts)) - return SelectedLayout::DenseNchw; - return SelectedLayout::PixelMajorRowStrip; -} - -static SelectedLayout chooseAddLayout(spatial::SpatAddPlanOp addPlan, llvm::DenseMap& layouts) { - if (getSelectedLayout(layouts, addPlan.getLhs()) != SelectedLayout::PixelMajorRowStrip - || getSelectedLayout(layouts, addPlan.getRhs()) != SelectedLayout::PixelMajorRowStrip) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(addPlan.getResult(), layouts)) - return SelectedLayout::DenseNchw; - return SelectedLayout::PixelMajorRowStrip; -} - -static SelectedLayout chooseConcatLayout(spatial::SpatConcatPlanOp concatPlan, - llvm::DenseMap& layouts) { - if (llvm::any_of(concatPlan.getInputs(), [&](Value input) { - return getSelectedLayout(layouts, input) != SelectedLayout::PixelMajorRowStrip; - })) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(concatPlan.getResult(), layouts)) - return SelectedLayout::DenseNchw; - return SelectedLayout::PixelMajorRowStrip; -} - -static SelectedLayout chooseMaxPoolLayout(spatial::SpatMaxPool2DPlanOp maxPoolPlan) { - return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan)) ? SelectedLayout::PixelMajorRowStrip - : SelectedLayout::DenseNchw; -} - -static SelectedLayout chooseGlobalAveragePoolLayout( - spatial::SpatGlobalAveragePoolPlanOp averagePoolPlan) { - return succeeded(canLowerGlobalAveragePoolPlanToRowStrip(averagePoolPlan)) - ? SelectedLayout::PixelMajorRowStrip - : SelectedLayout::DenseNchw; -} - -static spatial::SpatBlueprintOp insertRowStripBlueprint(IRRewriter& rewriter, Value value) { - auto outputType = cast(value.getType()); - auto [offsets, sizes] = buildRowStripMetadata(outputType); - return spatial::SpatBlueprintOp::create(rewriter, - value.getLoc(), - outputType, - value, - ValueRange {}, - rewriter.getStringAttr(kLogicalLayout), - rewriter.getStringAttr(kRowStripLayout), - rewriter.getDenseI64ArrayAttr(offsets), - rewriter.getDenseI64ArrayAttr(sizes), - rewriter.getStringAttr(kRowStripIndexMap), - nullptr, - nullptr, - nullptr, - nullptr, - nullptr, - nullptr, - nullptr); -} - -static void materializeDenseUses(IRRewriter& rewriter, - Value layoutValue, - llvm::DenseMap& layouts) { - SmallVector denseUses; - for (OpOperand& use : layoutValue.getUses()) { - if (usesSelectedRowStrip(use.getOwner(), layouts)) + } + auto userAlternatives = getAlternatives(use.getOwner(), selectedResults, target); + if (failed(userAlternatives)) continue; - denseUses.push_back(&use); + spatial::PhysicalLayout userResult = + selectedResults.lookup(use.getOwner()->getResult(0)); + unsigned userIndex = findCurrentAlternative(use.getOwner(), *userAlternatives, userResult); + if (use.getOperandNumber() < (*userAlternatives)[userIndex].operandLayouts.size() + && (*userAlternatives)[userIndex].operandLayouts[use.getOperandNumber()] + != alternative.resultLayout) + ++cost; + } + return cost; +} + +static LogicalResult materializeMismatchedUses( + IRRewriter& rewriter, Value value, const LayoutMap& layouts, + const spatial::SpatialTargetInfo& target) { + spatial::PhysicalLayout sourceLayout = getSelectedLayout(layouts, value); + SmallVector> mismatches; + for (OpOperand& use : value.getUses()) { + Operation* userOp = use.getOwner(); + spatial::PhysicalLayout required = spatial::PhysicalLayout::DenseNCHW; + if (auto capability = dyn_cast(userOp)) { + auto alternatives = getAlternatives(userOp, layouts, target); + if (failed(alternatives)) + return failure(); + spatial::PhysicalLayout selected = + getSelectedLayout(layouts, userOp->getResult(0)); + unsigned selectedIndex = findCurrentAlternative(userOp, *alternatives, selected); + required = (*alternatives)[selectedIndex].operandLayouts[use.getOperandNumber()]; + } + else if (auto flatten = dyn_cast(userOp); + flatten && sourceLayout == spatial::PhysicalLayout::NHWCRowStrip + && succeeded(canLowerFlattenFromRowStrip(flatten, target))) { + continue; + } + if (required != sourceLayout) + mismatches.push_back({&use, required}); } - for (OpOperand* use : denseUses) { - Operation* owner = use->getOwner(); - rewriter.setInsertionPoint(owner); - auto materialized = spatial::SpatMaterializeLayoutOp::create(rewriter, - owner->getLoc(), - use->get().getType(), - use->get(), - rewriter.getStringAttr(kLogicalLayout), - rewriter.getStringAttr(kRowStripLayout), - rewriter.getStringAttr(kDenseLayout)); + for (auto [use, required] : mismatches) { + Operation* userOp = use->getOwner(); + rewriter.setInsertionPoint(userOp); + auto materialized = spatial::SpatMaterializeLayoutOp::create( + rewriter, userOp->getLoc(), use->get().getType(), use->get(), + spatial::LogicalLayoutAttr::get( + rewriter.getContext(), spatial::LogicalLayout::NCHW), + spatial::PhysicalLayoutAttr::get(rewriter.getContext(), sourceLayout), + spatial::PhysicalLayoutAttr::get(rewriter.getContext(), + required)); use->set(materialized.getResult()); } + return success(); } -struct SpatialLayoutPlanningPass final : PassWrapper> { +static LogicalResult verifySelectedLayouts( + ArrayRef planOps, const LayoutMap& layouts, + const spatial::SpatialTargetInfo& target) { + for (Operation* op : planOps) { + auto selected = spatial::getSelectedPhysicalLayout(op); + if (!selected) + return op->emitOpError("requires a selected physical layout"), failure(); + auto alternatives = getAlternatives(op, layouts, target); + if (failed(alternatives)) + return failure(); + if (llvm::none_of(*alternatives, [&](const spatial::LayoutAlternative& alternative) { + return alternative.resultLayout == *selected; + })) + return op->emitOpError("selected physical layout is not advertised by its layout contract"), failure(); + } + return success(); +} + +struct SpatialLayoutPlanningPass final + : PassWrapper> { MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SpatialLayoutPlanningPass) StringRef getArgument() const override { return "spatial-layout-planning"; } - StringRef getDescription() const override { return "Select conservative Spatial layouts and insert reconciliation barriers."; } + StringRef getDescription() const override { + return "Select Spatial layout alternatives and insert explicit reconciliation barriers."; + } + + SpatialLayoutPlanningPass() = default; + explicit SpatialLayoutPlanningPass(const spatial::SpatialTargetInfo& target) + : target(target), hasTarget(true) {} void runOnOperation() override { - auto entryFunc = getPimEntryFunc(getOperation()); + ModuleOp moduleOp = getOperation(); + if (!hasTarget) { + moduleOp.emitError("Spatial layout planning requires an injected SpatialTargetInfo"); + signalPassFailure(); + return; + } + auto entryFunc = getPimEntryFunc(moduleOp); if (failed(entryFunc)) { - getOperation().emitError("failed to locate the PIM entry function during Spatial layout planning"); + moduleOp.emitError("failed to locate the PIM entry function during Spatial layout planning"); signalPassFailure(); return; } func::FuncOp funcOp = *entryFunc; - IRRewriter rewriter(&getContext()); - llvm::DenseMap layouts; + SmallVector planOps; + for (Operation& op : funcOp.getBody().front()) + if (isa(&op)) + planOps.push_back(&op); - bool changed = true; - while (changed) { - changed = false; - for (Operation& op : llvm::make_early_inc_range(funcOp.getBody().front())) { - if (auto convPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseConvLayout(convPlan, layouts); - if (layouts[convPlan.getResult()] != selected) { - layouts[convPlan.getResult()] = selected; - changed = true; - } - continue; + LayoutMap layouts; + for (Operation* op : planOps) + layouts[op->getResult(0)] = spatial::PhysicalLayout::DenseNCHW; + + const size_t maxRounds = 2 * planOps.size() + 1; + bool converged = false; + for (size_t round = 0; round < maxRounds && !converged; ++round) { + converged = true; + SmallVector order(planOps); + if (round % 2) + std::reverse(order.begin(), order.end()); + for (Operation* op : order) { + auto alternatives = getAlternatives(op, layouts, target); + if (failed(alternatives)) { + signalPassFailure(); + return; } - if (auto reluPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseActivationLayout(reluPlan.getInput(), reluPlan.getResult(), layouts); - if (layouts[reluPlan.getResult()] != selected) { - layouts[reluPlan.getResult()] = selected; - changed = true; + spatial::PhysicalLayout current = layouts.lookup(op->getResult(0)); + unsigned currentIndex = findCurrentAlternative(op, *alternatives, current); + int64_t bestCost = alternativeCost( + op, (*alternatives)[currentIndex], layouts, layouts, target); + unsigned bestIndex = currentIndex; + for (auto [index, alternative] : llvm::enumerate(*alternatives)) { + int64_t cost = alternativeCost(op, alternative, layouts, layouts, target); + if (cost < bestCost) { + bestCost = cost; + bestIndex = index; } - continue; } - if (auto siluPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseActivationLayout(siluPlan.getInput(), siluPlan.getResult(), layouts); - if (layouts[siluPlan.getResult()] != selected) { - layouts[siluPlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto resizePlan = dyn_cast(&op)) { - SelectedLayout selected = chooseResizeLayout(resizePlan, layouts); - if (layouts[resizePlan.getResult()] != selected) { - layouts[resizePlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto biasAddPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseBiasAddLayout(biasAddPlan, layouts); - if (layouts[biasAddPlan.getResult()] != selected) { - layouts[biasAddPlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto addPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseAddLayout(addPlan, layouts); - if (layouts[addPlan.getResult()] != selected) { - layouts[addPlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto concatPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseConcatLayout(concatPlan, layouts); - if (layouts[concatPlan.getResult()] != selected) { - layouts[concatPlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto maxPoolPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseMaxPoolLayout(maxPoolPlan); - if (layouts[maxPoolPlan.getResult()] != selected) { - layouts[maxPoolPlan.getResult()] = selected; - changed = true; - } - continue; - } - if (auto averagePoolPlan = dyn_cast(&op)) { - SelectedLayout selected = chooseGlobalAveragePoolLayout(averagePoolPlan); - if (layouts[averagePoolPlan.getResult()] != selected) { - layouts[averagePoolPlan.getResult()] = selected; - changed = true; - } - continue; + spatial::PhysicalLayout selected = (*alternatives)[bestIndex].resultLayout; + if (selected != current) { + layouts[op->getResult(0)] = selected; + converged = false; } } } - - for (Operation& op : llvm::make_early_inc_range(funcOp.getBody().front())) { - Value producedValue; - if (auto convPlan = dyn_cast(&op)) - producedValue = convPlan.getResult(); - else if (auto biasAddPlan = dyn_cast(&op)) - producedValue = biasAddPlan.getResult(); - else if (auto addPlan = dyn_cast(&op)) - producedValue = addPlan.getResult(); - else if (auto concatPlan = dyn_cast(&op)) - producedValue = concatPlan.getResult(); - else if (auto reluPlan = dyn_cast(&op)) - producedValue = reluPlan.getResult(); - else if (auto siluPlan = dyn_cast(&op)) - producedValue = siluPlan.getResult(); - else if (auto resizePlan = dyn_cast(&op)) - producedValue = resizePlan.getResult(); - else if (auto maxPoolPlan = dyn_cast(&op)) - producedValue = maxPoolPlan.getResult(); - else if (auto averagePoolPlan = dyn_cast(&op)) - producedValue = averagePoolPlan.getResult(); - else - continue; - - if (getSelectedLayout(layouts, producedValue) != SelectedLayout::PixelMajorRowStrip) - continue; - - rewriter.setInsertionPointAfter(&op); - auto blueprint = insertRowStripBlueprint(rewriter, producedValue); - rewriter.replaceAllUsesExcept(producedValue, blueprint.getResult(), blueprint); - materializeDenseUses(rewriter, blueprint.getResult(), layouts); + if (!converged) { + moduleOp.emitError("Spatial layout selection did not converge within its bounded iteration budget"); + signalPassFailure(); + return; } - if (failed(verifyLogicalSpatialGraphInvariants(*entryFunc))) { - getOperation().emitError("logical Spatial graph verification failed after SpatialLayoutPlanning"); + IRRewriter rewriter(&getContext()); + for (Operation* op : planOps) { + op->setAttr(spatial::kSelectedLayoutAttrName, + spatial::PhysicalLayoutAttr::get( + rewriter.getContext(), layouts.lookup(op->getResult(0)))); + if (failed(materializeMismatchedUses(rewriter, op->getResult(0), layouts, target))) { + signalPassFailure(); + return; + } + } + if (failed(verifySelectedLayouts(planOps, layouts, target)) + || failed(verifyLogicalSpatialGraphInvariants(*entryFunc))) { + moduleOp.emitError("Spatial layout planning verification failed"); signalPassFailure(); } } + + spatial::SpatialTargetInfo target; + bool hasTarget = false; }; } // namespace -std::unique_ptr createSpatialLayoutPlanningPass() { return std::make_unique(); } +std::unique_ptr createSpatialLayoutPlanningPass() { + return std::make_unique(); +} + +std::unique_ptr createSpatialLayoutPlanningPass( + const spatial::SpatialTargetInfo& target) { + return std::make_unique(target); +} } // namespace onnx_mlir diff --git a/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp b/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp index 7a58074..83e9384 100644 --- a/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp +++ b/src/PIM/Conversion/SpatialToPim/BatchCoreLoweringPatterns.cpp @@ -149,11 +149,10 @@ collectTopLevelFragmentAssemblyCopies(OpResult result, RankedTensorType packedRe auto blueprint = dyn_cast(use.getOwner()); if (!blueprint || blueprint->getParentOp() != blueprint->getParentOfType()) return failure(); - std::optional mode = blueprint.getMode(); std::optional> operandIndicesAttr = blueprint.getFragmentOperandIndices(); std::optional> sourceOffsetsAttr = blueprint.getFragmentSourceOffsets(); std::optional> sourceSlotsAttr = blueprint.getFragmentSourceSlots(); - if (!mode || *mode != "fragment_assembly" || !operandIndicesAttr || !sourceOffsetsAttr || !sourceSlotsAttr) + if (!spatial::isFragmentAssembly(blueprint.getMode()) || !operandIndicesAttr || !sourceOffsetsAttr || !sourceSlotsAttr) return failure(); if (!blueprint.getOutput().hasOneUse() || !isa(*blueprint.getOutput().getUsers().begin())) return failure(); @@ -418,8 +417,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul rewriter.setInsertionPointToEnd(newBlock); if (auto blueprint = dyn_cast(op)) { - std::optional modeAttr = blueprint.getMode(); - if (modeAttr && *modeAttr == "fragment_assembly") { + if (spatial::isFragmentAssembly(blueprint.getMode())) { for (Operation* user : blueprint.getOutput().getUsers()) { if (!isa(user)) return blueprint.emitOpError( @@ -483,8 +481,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul auto hostTargetType = cast(hostTarget.getType()); if (auto blueprint = insertSlice.getSource().getDefiningOp()) { - std::optional modeAttr = blueprint.getMode(); - if (modeAttr && *modeAttr == "fragment_assembly") { + if (spatial::isFragmentAssembly(blueprint.getMode())) { FailureOr> fragmentAssemblyCopies = collectFragmentAssemblyCopiesFromBlueprint(blueprint, mapper, /*lane=*/0, /*hostTargetIndex=*/0); if (failed(fragmentAssemblyCopies)) diff --git a/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp b/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp index f2cb16d..318a1e2 100644 --- a/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp +++ b/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp @@ -42,12 +42,11 @@ static FailureOr lowerFragmentAssemblyBlueprint(IRRewriter& rewriter, if (!resultType || !resultType.hasStaticShape()) return blueprint.emitOpError("fragment assembly lowering requires a static ranked tensor result"); - std::optional modeAttr = blueprint.getMode(); std::optional> operandIndicesAttr = blueprint.getFragmentOperandIndices(); std::optional> sourceSlotsAttr = blueprint.getFragmentSourceSlots(); std::optional> sourceOffsetsAttr = blueprint.getFragmentSourceOffsets(); std::optional> fragmentStridesAttr = blueprint.getFragmentStrides(); - if (!modeAttr || *modeAttr != "fragment_assembly" || !operandIndicesAttr || !sourceSlotsAttr + if (!spatial::isFragmentAssembly(blueprint.getMode()) || !operandIndicesAttr || !sourceSlotsAttr || !sourceOffsetsAttr || !fragmentStridesAttr) return blueprint.emitOpError("fragment assembly lowering requires explicit fragment metadata"); @@ -203,8 +202,7 @@ static bool isHostMaterializableHelperOp(Operation* op) { if (isa(op) || op->hasTrait()) return true; if (auto blueprint = dyn_cast(op)) { - std::optional mode = blueprint.getMode(); - return mode && *mode == "fragment_assembly"; + return spatial::isFragmentAssembly(blueprint.getMode()); } return isShapingOnlyOp(op) || isPureIndexComputationOp(op); } @@ -291,8 +289,7 @@ static bool inlineInputlessHelperComputeForWeightLikeUsers(spatial::SpatSchedule } for (Operation& op : block.without_terminator()) { if (auto blueprint = dyn_cast(op)) { - std::optional modeAttr = blueprint.getMode(); - if (modeAttr && *modeAttr == "fragment_assembly") { + if (spatial::isFragmentAssembly(blueprint.getMode())) { auto lowered = lowerFragmentAssemblyBlueprint(rewriter, blueprint, mapping); if (failed(lowered)) return false; diff --git a/src/PIM/Conversion/SpatialToPim/Patterns.cpp b/src/PIM/Conversion/SpatialToPim/Patterns.cpp index c450e0b..ebb9f34 100644 --- a/src/PIM/Conversion/SpatialToPim/Patterns.cpp +++ b/src/PIM/Conversion/SpatialToPim/Patterns.cpp @@ -22,8 +22,7 @@ struct LowerFragmentAssemblyBlueprintPattern LogicalResult matchAndRewrite(spatial::SpatBlueprintOp op, OpAdaptor adaptor, ConversionPatternRewriter& rewriter) const override { - std::optional modeAttr = op.getMode(); - if (!modeAttr || *modeAttr != "fragment_assembly") + if (!spatial::isFragmentAssembly(op.getMode())) return failure(); auto resultType = dyn_cast(op.getOutput().getType()); diff --git a/src/PIM/Conversion/SpatialToPim/ReturnPathNormalization.cpp b/src/PIM/Conversion/SpatialToPim/ReturnPathNormalization.cpp index 0c8e9a6..b33c6d0 100644 --- a/src/PIM/Conversion/SpatialToPim/ReturnPathNormalization.cpp +++ b/src/PIM/Conversion/SpatialToPim/ReturnPathNormalization.cpp @@ -158,8 +158,7 @@ analyzeTopLevelFragmentAssemblyUses(Value value) { auto blueprint = dyn_cast(use.getOwner()); if (!blueprint || blueprint->getParentOp() != blueprint->getParentOfType()) return failure(); - std::optional mode = blueprint.getMode(); - if (!mode || *mode != "fragment_assembly") + if (!spatial::isFragmentAssembly(blueprint.getMode())) return failure(); if (!blueprint.getOutput().hasOneUse() || !isa(*blueprint.getOutput().getUsers().begin())) return failure(); @@ -819,8 +818,7 @@ void raptor::SpatialToPimPass::replaceReturnWithOutputBuffers(func::ReturnOp ret } if (auto blueprint = dyn_cast(op)) { - std::optional mode = blueprint.getMode(); - if (mode && *mode == "fragment_assembly") { + if (spatial::isFragmentAssembly(blueprint.getMode())) { markOpToRemove(blueprint.getOperation()); for (Value operand : blueprint->getOperands()) markOwnedReturnChain(operand.getDefiningOp(), markOwnedReturnChain); diff --git a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp index 038d428..a77bac2 100644 --- a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp @@ -33,6 +33,9 @@ using namespace pim; namespace onnx_mlir { +static void annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncOp funcOp); +static FailureOr requirePimEntryFunc(ModuleOp moduleOp, StringRef phase); + namespace { struct MemRefCopyWorkItem { @@ -333,22 +336,6 @@ static LogicalResult verifyPimCopyEndpoints(Operation* copy, return success(valid); } -struct PimBufferizationPass : PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PimBufferizationPass) - StringRef getArgument() const override { return "bufferize-pim"; } - StringRef getDescription() const override { return "Bufferize PIM and Spatial ops."; } - - PimBufferizationPass() = default; - PimBufferizationPass(const PimBufferizationPass& pass) {} - - void runOnOperation() final; - -private: - void annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncOp funcOp) const; - LogicalResult verifyContiguousRuntimeOperands(ModuleOp moduleOp) const; - LogicalResult verifyPimCopyAddressSpaces(ModuleOp moduleOp) const; -}; - static void materializeWritableConstantDestinations(func::FuncOp funcOp) { SmallVector constantBackedRoots; llvm::SmallPtrSet seenRoots; @@ -387,65 +374,23 @@ static void materializeWritableConstantDestinations(func::FuncOp funcOp) { } } -static LogicalResult verifyPimCoresNeedNoTensorCopies( - ModuleOp module, const bufferization::OneShotBufferizationOptions& baseOptions) { - static constexpr StringLiteral kExistingAlloc = "raptor.existing_core_alloc"; - OwningOpRef clone = module.clone(); - clone->walk([&](bufferization::AllocTensorOp alloc) { - if (alloc->getParentOfType() - || alloc->getParentOfType()) - alloc->setAttr(kExistingAlloc, UnitAttr::get(module.getContext())); - }); - - auto options = baseOptions; - options.bufferizeFunctionBoundaries = false; - options.opFilter.allowOperation([](Operation* op) { - return isa(op) - || op->getParentOfType() - || op->getParentOfType(); - }); - - bufferization::BufferizationState state; - if (failed(bufferization::insertTensorCopies(*clone, options, state))) { - module.emitError("official one-shot analysis failed while verifying PIM core copy freedom"); - return failure(); - } - - CappedDiagnosticReporter diagnostics; - clone->walk([&](bufferization::AllocTensorOp alloc) { - if (alloc->hasAttr(kExistingAlloc) - || (!alloc->getParentOfType() - && !alloc->getParentOfType())) - return; - Operation* requiredBy = alloc->getUsers().empty() - ? alloc.getOperation() : *alloc->getUsers().begin(); - diagnostics.report(requiredBy, [](Operation* op) { - op->emitOpError("official one-shot bufferization requires a tensor copy inside a PIM core"); - }); - }); - diagnostics.emitSuppressedSummary(module, "required PIM core tensor copies"); - return success(!diagnostics.hasFailure()); -} - -} // namespace - -void PimBufferizationPass::runOnOperation() { - auto moduleOp = getOperation(); - auto funcOp = *getPimEntryFunc(moduleOp); - +static bufferization::OneShotBufferizationOptions makePimBufferizationOptions() { bufferization::OneShotBufferizationOptions options; options.allowUnknownOps = true; options.bufferizeFunctionBoundaries = true; options.setFunctionBoundaryTypeConversion(bufferization::LayoutMapOption::IdentityLayoutMap); + return options; +} +static LogicalResult preparePimBufferization(func::FuncOp funcOp) { materializeWritableConstantDestinations(funcOp); - if (failed(verifyPimCoresNeedNoTensorCopies(moduleOp, options))) { - signalPassFailure(); - return; - } + return success(); +} +static LogicalResult runOneShotPimBufferization( + ModuleOp moduleOp, const bufferization::OneShotBufferizationOptions& options) { auto hostOptions = options; - hostOptions.opFilter.denyOperation([](Operation *op) { + hostOptions.opFilter.denyOperation([](Operation* op) { return op->getParentOfType() || op->getParentOfType(); }); @@ -453,84 +398,14 @@ void PimBufferizationPass::runOnOperation() { if (failed(bufferization::insertTensorCopies(moduleOp, hostOptions, state)) || failed(bufferization::bufferizeModuleOp(moduleOp, options, state))) { moduleOp.emitError("Failed to bufferize PIM and Spatial ops"); - signalPassFailure(); - return; + return failure(); } - - forwardSingleConsumerReceiveCopies(funcOp); - forwardSingleConsumerContiguousInputCopies(funcOp); - forwardSingleConsumerPimOutputCopies(funcOp); - - MLIRContext* ctx = moduleOp.getContext(); - PatternRewriter rewriter(ctx); - - SmallVector copyWorklist; - llvm::SmallPtrSet seenCopyOps; - auto addCopyOp = [&](memref::CopyOp copyOp, const StaticValueKnowledge& knowledge) { - if (seenCopyOps.insert(copyOp.getOperation()).second) - copyWorklist.push_back({copyOp, knowledge}); - }; - - moduleOp.walk([&](pim::PimCoreOp coreOp) { - StaticValueKnowledge knowledge = seedCoreKnowledge(coreOp); - (void) walkPimCoreBlockStructurally( - coreOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { - if (auto copyOp = dyn_cast(&op)) - addCopyOp(copyOp, opKnowledge); - return success(); - }); - }); - moduleOp.walk([&](pim::PimCoreBatchOp coreBatchOp) { - for (unsigned lane = 0; lane < coreBatchOp.getLaneCount(); ++lane) { - StaticValueKnowledge knowledge = seedCoreBatchKnowledge(coreBatchOp, lane); - (void) walkPimCoreBlockStructurally( - coreBatchOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { - if (auto copyOp = dyn_cast(&op)) - addCopyOp(copyOp, opKnowledge); - return success(); - }); - } - }); - - bool hasFailed = false; - Value zeroOffset = getOrCreateIndexConstant(rewriter, funcOp, 0); - for (const MemRefCopyWorkItem& workItem : copyWorklist) { - memref::CopyOp copyOp = workItem.copyOp; - rewriter.setInsertionPoint(copyOp); - if (failed(lowerMemRefCopyToPimCopy(copyOp, zeroOffset, rewriter, workItem.knowledge))) - hasFailed = true; - } - if (hasFailed) { - signalPassFailure(); - return; - } - - RewritePatternSet contiguityPatterns(ctx); - populatePimContiguityNormalizationPatterns(contiguityPatterns); - - GreedyRewriteConfig contiguityConfig; - contiguityConfig.enableFolding(false); - if (failed(applyPatternsGreedily(moduleOp, std::move(contiguityPatterns), contiguityConfig))) { - moduleOp.emitError("failed to normalize PIM copy contiguity during bufferization"); - signalPassFailure(); - return; - } - if (failed(verifyContiguousRuntimeOperands(moduleOp))) { - signalPassFailure(); - return; - } - if (failed(verifyPimCopyAddressSpaces(moduleOp))) { - signalPassFailure(); - return; - } - - annotateWeightsMemrefs(moduleOp, funcOp); - - // Dump to file for debug - dumpModule(moduleOp, "pim1_buff"); + return success(); } -void PimBufferizationPass::annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncOp funcOp) const { +} // namespace + +static void annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncOp funcOp) { auto markWeights = [&](Operation* op) { walkPimMvmVmmWeightUses(op, [&](OpOperand& weightUse) { Value weight = weightUse.get(); @@ -548,7 +423,7 @@ void PimBufferizationPass::annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncO funcOp.walk([&](PimCoreBatchOp coreBatchOp) { markWeights(coreBatchOp); }); } -LogicalResult PimBufferizationPass::verifyContiguousRuntimeOperands(ModuleOp moduleOp) const { +static LogicalResult verifyContiguousRuntimeOperands(ModuleOp moduleOp) { bool hasFailure = false; auto verifyWithKnowledge = [&](auto coreLikeOp, const StaticValueKnowledge& initialKnowledge) { @@ -640,7 +515,7 @@ LogicalResult PimBufferizationPass::verifyContiguousRuntimeOperands(ModuleOp mod return success(); } -LogicalResult PimBufferizationPass::verifyPimCopyAddressSpaces(ModuleOp moduleOp) const { +static LogicalResult verifyPimCopyAddressSpaces(ModuleOp moduleOp) { size_t failureCount = 0; auto verifyWithKnowledge = [&](auto coreLikeOp, const StaticValueKnowledge& initialKnowledge) { (void) walkPimCoreBlockStructurally( @@ -675,6 +550,201 @@ LogicalResult PimBufferizationPass::verifyPimCopyAddressSpaces(ModuleOp moduleOp return success(failureCount == 0); } -std::unique_ptr createPimBufferizationPass() { return std::make_unique(); } +static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp) { + forwardSingleConsumerReceiveCopies(funcOp); + forwardSingleConsumerContiguousInputCopies(funcOp); + forwardSingleConsumerPimOutputCopies(funcOp); + + MLIRContext* ctx = moduleOp.getContext(); + PatternRewriter rewriter(ctx); + + SmallVector copyWorklist; + llvm::SmallPtrSet seenCopyOps; + auto addCopyOp = [&](memref::CopyOp copyOp, const StaticValueKnowledge& knowledge) { + if (seenCopyOps.insert(copyOp.getOperation()).second) + copyWorklist.push_back({copyOp, knowledge}); + }; + + moduleOp.walk([&](pim::PimCoreOp coreOp) { + StaticValueKnowledge knowledge = seedCoreKnowledge(coreOp); + (void) walkPimCoreBlockStructurally( + coreOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { + if (auto copyOp = dyn_cast(&op)) + addCopyOp(copyOp, opKnowledge); + return success(); + }); + }); + moduleOp.walk([&](pim::PimCoreBatchOp coreBatchOp) { + for (unsigned lane = 0; lane < coreBatchOp.getLaneCount(); ++lane) { + StaticValueKnowledge knowledge = seedCoreBatchKnowledge(coreBatchOp, lane); + (void) walkPimCoreBlockStructurally( + coreBatchOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { + if (auto copyOp = dyn_cast(&op)) + addCopyOp(copyOp, opKnowledge); + return success(); + }); + } + }); + + bool hasFailed = false; + Value zeroOffset = getOrCreateIndexConstant(rewriter, funcOp, 0); + for (const MemRefCopyWorkItem& workItem : copyWorklist) { + memref::CopyOp copyOp = workItem.copyOp; + rewriter.setInsertionPoint(copyOp); + if (failed(lowerMemRefCopyToPimCopy(copyOp, zeroOffset, rewriter, workItem.knowledge))) + hasFailed = true; + } + if (hasFailed) + return failure(); + + RewritePatternSet contiguityPatterns(ctx); + populatePimContiguityNormalizationPatterns(contiguityPatterns); + + GreedyRewriteConfig contiguityConfig; + contiguityConfig.enableFolding(false); + if (failed(applyPatternsGreedily(moduleOp, std::move(contiguityPatterns), contiguityConfig))) { + moduleOp.emitError("failed to normalize PIM copy contiguity during bufferization"); + return failure(); + } + annotateWeightsMemrefs(moduleOp, funcOp); + dumpModule(moduleOp, "pim1_buff"); + return success(); +} + +static FailureOr requirePimEntryFunc(ModuleOp moduleOp, StringRef phase) { + auto entryFunc = getPimEntryFunc(moduleOp); + if (failed(entryFunc)) { + moduleOp.emitError("failed to locate the PIM entry function during ") << phase; + return failure(); + } + return *entryFunc; +} + +namespace { + +struct PimBufferizationPreparationPass + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PimBufferizationPreparationPass) + + StringRef getArgument() const override { return "pim-bufferization-preparation"; } + StringRef getDescription() const override { + return "Prepare writable tensor destinations for PIM one-shot bufferization."; + } + + void runOnOperation() final { + ModuleOp moduleOp = getOperation(); + auto funcOp = requirePimEntryFunc(moduleOp, "PIM bufferization preparation"); + if (failed(funcOp)) { + signalPassFailure(); + return; + } + if (failed(preparePimBufferization(*funcOp))) + signalPassFailure(); + } +}; + +struct PimOneShotBufferizationPass + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PimOneShotBufferizationPass) + + StringRef getArgument() const override { return "pim-one-shot-bufferization"; } + StringRef getDescription() const override { + return "Run one-shot bufferization for PIM and Spatial tensors."; + } + + void runOnOperation() final { + if (failed(runOneShotPimBufferization(getOperation(), makePimBufferizationOptions()))) + signalPassFailure(); + } +}; + +struct PimMemoryNormalizationPass + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PimMemoryNormalizationPass) + + StringRef getArgument() const override { return "pim-memory-normalization"; } + StringRef getDescription() const override { + return "Normalize PIM memory copies and verify addressable operands."; + } + + void runOnOperation() final { + ModuleOp moduleOp = getOperation(); + auto funcOp = requirePimEntryFunc(moduleOp, "PIM memory normalization"); + if (failed(funcOp)) { + signalPassFailure(); + return; + } + if (failed(normalizePimMemory(moduleOp, *funcOp))) + signalPassFailure(); + } +}; + +static LogicalResult verifyNoTensorValues(ModuleOp moduleOp) { + size_t failureCount = 0; + moduleOp.walk([&](Operation* op) { + if (failureCount >= 8) + return; + if (op->getDialect()->getNamespace() == "tensor") { + op->emitOpError("tensor operation remains after PIM bufferization"); + ++failureCount; + return; + } + for (Value value : op->getOperands()) { + if (isa(value.getType())) { + op->emitOpError("tensor operand remains after PIM bufferization"); + ++failureCount; + return; + } + } + for (Value value : op->getResults()) { + if (isa(value.getType())) { + op->emitOpError("tensor result remains after PIM bufferization"); + ++failureCount; + return; + } + } + }); + if (failureCount != 0) + moduleOp.emitError() << "found " << failureCount + << " tensor value(s) after PIM bufferization" + << (failureCount == 8 ? " (first 8 reported)" : ""); + return success(failureCount == 0); +} + +struct PimBufferizationVerificationPass + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PimBufferizationVerificationPass) + + StringRef getArgument() const override { return "pim-bufferization-verification"; } + StringRef getDescription() const override { + return "Verify tensor elimination, contiguity, and PIM copy address spaces."; + } + + void runOnOperation() final { + ModuleOp moduleOp = getOperation(); + if (failed(verifyNoTensorValues(moduleOp)) + || failed(verifyContiguousRuntimeOperands(moduleOp)) + || failed(verifyPimCopyAddressSpaces(moduleOp))) + signalPassFailure(); + } +}; + +} // namespace + +std::unique_ptr createPimBufferizationPreparationPass() { + return std::make_unique(); +} + +std::unique_ptr createPimOneShotBufferizationPass() { + return std::make_unique(); +} + +std::unique_ptr createPimMemoryNormalizationPass() { + return std::make_unique(); +} + +std::unique_ptr createPimBufferizationVerificationPass() { + return std::make_unique(); +} } // namespace onnx_mlir diff --git a/src/PIM/Dialect/Spatial/CMakeLists.txt b/src/PIM/Dialect/Spatial/CMakeLists.txt index 983cfa9..7ee433c 100644 --- a/src/PIM/Dialect/Spatial/CMakeLists.txt +++ b/src/PIM/Dialect/Spatial/CMakeLists.txt @@ -1,6 +1,16 @@ add_onnx_mlir_dialect(Spatial spat) add_onnx_mlir_dialect_doc(spat Spatial.td) +set(LLVM_TARGET_DEFINITIONS Spatial.td) +mlir_tablegen(SpatialEnums.hpp.inc -gen-enum-decls "-I${ONNX_MLIR_SRC_ROOT}") +mlir_tablegen(SpatialEnums.cpp.inc -gen-enum-defs "-I${ONNX_MLIR_SRC_ROOT}") +add_public_tablegen_target(OMSpatialEnumsIncGen) + +set(LLVM_TARGET_DEFINITIONS SpatialLayoutInterface.td) +mlir_tablegen(SpatialLayoutInterface.hpp.inc -gen-op-interface-decls "-I${ONNX_MLIR_SRC_ROOT}") +mlir_tablegen(SpatialLayoutInterface.cpp.inc -gen-op-interface-defs "-I${ONNX_MLIR_SRC_ROOT}") +add_public_tablegen_target(OMSpatialLayoutInterfaceIncGen) + add_pim_library(SpatialOps SpatialOps.cpp SpatialOpsAsm.cpp @@ -18,7 +28,7 @@ add_pim_library(SpatialOps Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp Transforms/MergeComputeNodes/DeferredResultRealization.cpp Transforms/MergeComputeNodes/DeferredCommunicationRealization.cpp - Transforms/MergeComputeNodes/MergeComputeNodesPass.cpp + Transforms/MergeComputeNodes/ScheduledSpatialPasses.cpp Transforms/MergeComputeNodes/ScheduledComputeMaterialization.cpp Transforms/MergeComputeNodes/ScheduledComputePlanning.cpp Transforms/MergeComputeNodes/ScheduledComputeReport.cpp @@ -33,6 +43,8 @@ add_pim_library(SpatialOps DEPENDS OMONNXIncGen OMSpatialIncGen + OMSpatialEnumsIncGen + OMSpatialLayoutInterfaceIncGen LINK_LIBS PUBLIC MLIRIR diff --git a/src/PIM/Dialect/Spatial/Spatial.td b/src/PIM/Dialect/Spatial/Spatial.td index 22f099f..7fdd8ee 100644 --- a/src/PIM/Dialect/Spatial/Spatial.td +++ b/src/PIM/Dialect/Spatial/Spatial.td @@ -5,20 +5,77 @@ include "mlir/IR/OpBase.td" include "mlir/IR/OpAsmInterface.td" include "mlir/IR/BuiltinTypes.td" include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/EnumAttr.td" include "mlir/IR/RegionKindInterface.td" include "mlir/Interfaces/ControlFlowInterfaces.td" include "mlir/Interfaces/ParallelCombiningOpInterface.td" include "mlir/Interfaces/SideEffectInterfaces.td" +include "src/Accelerators/PIM/Dialect/Spatial/SpatialLayoutInterface.td" def SpatialDialect : Dialect { let name = "spat"; let summary = "Dialect designed for deep learning computation in a spatial architecture"; let cppNamespace = "::onnx_mlir::spatial"; + let useDefaultAttributePrinterParser = 0; + let extraClassDeclaration = [{ + ::mlir::Attribute parseAttribute(::mlir::DialectAsmParser &parser, + ::mlir::Type type) const override; + void printAttribute(::mlir::Attribute attr, + ::mlir::DialectAsmPrinter &printer) const override; + }]; +} + +def SpatLogicalLayoutNCHW : I32EnumAttrCase<"NCHW", 0, "nchw">; +def SpatLogicalLayout : I32EnumAttr<"LogicalLayout", "Logical tensor layout", [ + SpatLogicalLayoutNCHW +]> { + let genSpecializedAttr = 0; + let cppNamespace = "::onnx_mlir::spatial"; +} + +def SpatLogicalLayoutAttr : EnumAttr { + let assemblyFormat = "$value"; +} + +def SpatPhysicalLayoutDenseNCHW : I32EnumAttrCase<"DenseNCHW", 0, "dense_nchw">; +def SpatPhysicalLayoutNCHWRowStrip : I32EnumAttrCase<"NCHWRowStrip", 1, "nchw_row_strip">; +def SpatPhysicalLayoutNHWCRowStrip : I32EnumAttrCase<"NHWCRowStrip", 2, "nhwc_row_strip">; +def SpatPhysicalLayoutFragmented : I32EnumAttrCase<"Fragmented", 3, "fragmented">; +def SpatPhysicalLayout : I32EnumAttr<"PhysicalLayout", "Physical tensor layout", [ + SpatPhysicalLayoutDenseNCHW, + SpatPhysicalLayoutNCHWRowStrip, + SpatPhysicalLayoutNHWCRowStrip, + SpatPhysicalLayoutFragmented +]> { + let genSpecializedAttr = 0; + let cppNamespace = "::onnx_mlir::spatial"; +} + +def SpatPhysicalLayoutAttr : EnumAttr { + let assemblyFormat = "$value"; +} + +def SpatBlueprintModePhysicalView : I32EnumAttrCase<"PhysicalView", 0, "physical_view">; +def SpatBlueprintModeFragmentAssembly : I32EnumAttrCase<"FragmentAssembly", 1, "fragment_assembly">; +def SpatBlueprintMode : I32EnumAttr<"BlueprintMode", "Blueprint reconstruction mode", [ + SpatBlueprintModePhysicalView, + SpatBlueprintModeFragmentAssembly +]> { + let genSpecializedAttr = 0; + let cppNamespace = "::onnx_mlir::spatial"; +} + +def SpatBlueprintModeAttr : EnumAttr { + let assemblyFormat = "$value"; } class SpatOp traits = []> : Op; +class SpatLayoutPlanOp : SpatOp]>; + // TODO maybe remove and use AnyRankedTensor directly def SpatTensor : AnyTypeOf<[AnyMemRef, AnyRankedTensor], "", "::mlir::ShapedType">; @@ -252,7 +309,7 @@ def SpatConcatOp : SpatOp<"concat", []> { // Planning //===----------------------------------------------------------------------===// -def SpatConv2DPlanOp : SpatOp<"conv2d_plan", []> { +def SpatConv2DPlanOp : SpatLayoutPlanOp<"conv2d_plan"> { let summary = "Structured Conv2D planning op that preserves logical ONNX geometry"; let arguments = (ins @@ -263,7 +320,7 @@ def SpatConv2DPlanOp : SpatOp<"conv2d_plan", []> { DenseI64ArrayAttr:$strides, DenseI64ArrayAttr:$dilations, I64Attr:$group, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -273,12 +330,12 @@ def SpatConv2DPlanOp : SpatOp<"conv2d_plan", []> { let hasVerifier = 1; } -def SpatReluPlanOp : SpatOp<"relu_plan", []> { +def SpatReluPlanOp : SpatLayoutPlanOp<"relu_plan"> { let summary = "Layout-aware ReLU planning op"; let arguments = (ins SpatTensor:$input, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -288,12 +345,12 @@ def SpatReluPlanOp : SpatOp<"relu_plan", []> { let hasVerifier = 1; } -def SpatSiluPlanOp : SpatOp<"silu_plan", []> { +def SpatSiluPlanOp : SpatLayoutPlanOp<"silu_plan"> { let summary = "Layout-aware SiLU planning op"; let arguments = (ins SpatTensor:$input, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -303,12 +360,12 @@ def SpatSiluPlanOp : SpatOp<"silu_plan", []> { let hasVerifier = 1; } -def SpatResizeNearestPlanOp : SpatOp<"resize_nearest_plan", []> { +def SpatResizeNearestPlanOp : SpatLayoutPlanOp<"resize_nearest_plan"> { let summary = "Layout-aware nearest asymmetric Resize planning op"; let arguments = (ins SpatTensor:$input, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -318,7 +375,7 @@ def SpatResizeNearestPlanOp : SpatOp<"resize_nearest_plan", []> { let hasVerifier = 1; } -def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> { +def SpatMaxPool2DPlanOp : SpatLayoutPlanOp<"max_pool2d_plan"> { let summary = "Layout-aware 2D NCHW MaxPool planning op"; let arguments = (ins @@ -327,7 +384,7 @@ def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> { DenseI64ArrayAttr:$pads, DenseI64ArrayAttr:$strides, DenseI64ArrayAttr:$dilations, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -337,12 +394,12 @@ def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> { let hasVerifier = 1; } -def SpatGlobalAveragePoolPlanOp : SpatOp<"global_average_pool_plan", []> { +def SpatGlobalAveragePoolPlanOp : SpatLayoutPlanOp<"global_average_pool_plan"> { let summary = "Layout-aware NCHW global average-pool planning op"; let arguments = (ins SpatTensor:$input, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -352,13 +409,13 @@ def SpatGlobalAveragePoolPlanOp : SpatOp<"global_average_pool_plan", []> { let hasVerifier = 1; } -def SpatBiasAddPlanOp : SpatOp<"bias_add_plan", []> { +def SpatBiasAddPlanOp : SpatLayoutPlanOp<"bias_add_plan"> { let summary = "Layout-aware Conv-style bias add planning op"; let arguments = (ins SpatTensor:$input, SpatTensor:$bias, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -368,13 +425,13 @@ def SpatBiasAddPlanOp : SpatOp<"bias_add_plan", []> { let hasVerifier = 1; } -def SpatAddPlanOp : SpatOp<"add_plan", []> { +def SpatAddPlanOp : SpatLayoutPlanOp<"add_plan"> { let summary = "Layout-aware elementwise add planning op"; let arguments = (ins SpatTensor:$lhs, SpatTensor:$rhs, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -384,13 +441,13 @@ def SpatAddPlanOp : SpatOp<"add_plan", []> { let hasVerifier = 1; } -def SpatConcatPlanOp : SpatOp<"concat_plan", []> { +def SpatConcatPlanOp : SpatLayoutPlanOp<"concat_plan"> { let summary = "Layout-aware tensor concatenation planning op"; let arguments = (ins Variadic:$inputs, I64Attr:$axis, - StrAttr:$logicalLayout + SpatLogicalLayoutAttr:$logicalLayout ); let results = (outs @@ -406,12 +463,12 @@ def SpatBlueprintOp : SpatOp<"blueprint", []> { let arguments = (ins SpatTensor:$input, Variadic:$fragments, - StrAttr:$logicalLayout, - StrAttr:$physicalLayout, + SpatLogicalLayoutAttr:$logicalLayout, + SpatPhysicalLayoutAttr:$physicalLayout, DenseI64ArrayAttr:$fragmentOffsets, DenseI64ArrayAttr:$fragmentSizes, StrAttr:$indexMap, - OptionalAttr:$mode, + OptionalAttr:$mode, OptionalAttr:$fragmentOperandIndices, OptionalAttr:$fragmentSourceSlots, OptionalAttr:$fragmentSourceOffsets, @@ -433,9 +490,9 @@ def SpatMaterializeLayoutOp : SpatOp<"materialize_layout", []> { let arguments = (ins SpatTensor:$input, - StrAttr:$logicalLayout, - StrAttr:$sourcePhysicalLayout, - StrAttr:$targetPhysicalLayout + SpatLogicalLayoutAttr:$logicalLayout, + SpatPhysicalLayoutAttr:$sourcePhysicalLayout, + SpatPhysicalLayoutAttr:$targetPhysicalLayout ); let results = (outs diff --git a/src/PIM/Dialect/Spatial/SpatialLayoutInterface.td b/src/PIM/Dialect/Spatial/SpatialLayoutInterface.td new file mode 100644 index 0000000..2cbb68e --- /dev/null +++ b/src/PIM/Dialect/Spatial/SpatialLayoutInterface.td @@ -0,0 +1,24 @@ +#ifndef SPATIAL_LAYOUT_INTERFACE_TD +#define SPATIAL_LAYOUT_INTERFACE_TD + +include "mlir/IR/OpBase.td" + +def SpatialLayoutCapabilityInterface : OpInterface<"SpatialLayoutCapabilityInterface"> { + let description = [{ + Contract implemented by logical Spatial planning operations that expose + their legal physical layout alternatives to the Spatial planner. + }]; + + let methods = [ + InterfaceMethod< + "Return legal physical layout alternatives for this operation and its current operand layouts.", + "::llvm::SmallVector<::onnx_mlir::spatial::LayoutAlternative>", + "getLayoutAlternatives", + (ins "const ::onnx_mlir::spatial::SpatialTargetInfo &":$target, + "::llvm::ArrayRef<::onnx_mlir::spatial::PhysicalLayout>":$operandLayouts)> + ]; + + let cppNamespace = "::onnx_mlir::spatial"; +} + +#endif diff --git a/src/PIM/Dialect/Spatial/SpatialOps.cpp b/src/PIM/Dialect/Spatial/SpatialOps.cpp index 166b8a8..ff14951 100644 --- a/src/PIM/Dialect/Spatial/SpatialOps.cpp +++ b/src/PIM/Dialect/Spatial/SpatialOps.cpp @@ -4,6 +4,8 @@ #include +#include "mlir/IR/DialectImplementation.h" + #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" using namespace mlir; @@ -47,7 +49,7 @@ bool isCanonicalContiguousRowMajorFragmentAssembly(SpatBlueprintOp blueprint) { auto fragmentStrides = blueprint.getFragmentStrides(); if (!logicalType || !physicalType || !logicalType.hasStaticShape() || !physicalType.hasStaticShape() || logicalType.getRank() < 2 || !blueprint.getFragments().empty() - || blueprint.getMode() != "fragment_assembly" || !operandIndices || !sourceSlots || !sourceOffsets + || !isFragmentAssembly(blueprint.getMode()) || !operandIndices || !sourceSlots || !sourceOffsets || !fragmentStrides) return false; @@ -438,6 +440,11 @@ OpResult SpatInParallelOp::getParentResult(int64_t idx) { llvm::iterator_range SpatInParallelOp::getYieldingOps() { return getRegion().front().getOperations(); } void SpatialDialect::initialize() { + addAttributes< +#define GET_ATTRDEF_LIST +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialAttributes.cpp.inc" + + >(); addTypes< #define GET_TYPEDEF_LIST #include "src/Accelerators/PIM/Dialect/Spatial/SpatialTypes.cpp.inc" @@ -459,6 +466,33 @@ void SpatialDialect::initialize() { #define GET_OP_CLASSES #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.cpp.inc" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialEnums.cpp.inc" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialLayoutInterface.cpp.inc" + +#define GET_ATTRDEF_CLASSES +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialAttributes.cpp.inc" + +namespace onnx_mlir { +namespace spatial { + +Attribute SpatialDialect::parseAttribute(DialectAsmParser& parser, Type type) const { + StringRef attrTag; + if (Attribute attr; generatedAttributeParser(parser, &attrTag, type, attr).has_value()) + return attr; + parser.emitError(parser.getCurrentLocation()) << "unknown attribute `" << attrTag + << "` in dialect `spat`"; + return {}; +} + +void SpatialDialect::printAttribute(Attribute attr, DialectAsmPrinter& printer) const { + if (succeeded(generatedAttributePrinter(attr, printer))) + return; + llvm_unreachable("unknown attribute in Spatial dialect"); +} + +} // namespace spatial +} // namespace onnx_mlir + #define GET_TYPEDEF_CLASSES #include "src/Accelerators/PIM/Dialect/Spatial/SpatialDialect.cpp.inc" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialTypes.cpp.inc" diff --git a/src/PIM/Dialect/Spatial/SpatialOps.hpp b/src/PIM/Dialect/Spatial/SpatialOps.hpp index 0953af3..5724659 100644 --- a/src/PIM/Dialect/Spatial/SpatialOps.hpp +++ b/src/PIM/Dialect/Spatial/SpatialOps.hpp @@ -11,15 +11,37 @@ #include "mlir/Interfaces/ParallelCombiningOpInterface.h" #include "llvm/ADT/DenseSet.h" +#include "llvm/ADT/SmallVector.h" #include "llvm/ADT/SetVector.h" +#include "llvm/ADT/TypeSwitch.h" #include #include #include #include +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetInfo.hpp" + /// Include the auto-generated header files containing the declarations #include "src/Accelerators/PIM/Dialect/Spatial/SpatialDialect.hpp.inc" +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialEnums.hpp.inc" + +namespace onnx_mlir { +namespace spatial { + +struct LayoutAlternative { + llvm::SmallVector operandLayouts; + PhysicalLayout resultLayout = PhysicalLayout::DenseNCHW; + int64_t intrinsicCost = 0; +}; + +} // namespace spatial +} // namespace onnx_mlir + +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialLayoutInterface.hpp.inc" + +#define GET_ATTRDEF_CLASSES +#include "src/Accelerators/PIM/Dialect/Spatial/SpatialAttributes.hpp.inc" #define GET_TYPEDEF_CLASSES #include "src/Accelerators/PIM/Dialect/Spatial/SpatialTypes.hpp.inc" @@ -31,6 +53,48 @@ namespace onnx_mlir { namespace spatial { inline constexpr llvm::StringLiteral kContiguousRowMajorFragments = "contiguous_row_major_fragments"; +inline constexpr llvm::StringLiteral kSelectedLayoutAttrName = "spat.selected_layout"; + +inline LogicalLayoutAttr getNCHWLayout(mlir::MLIRContext* context) { + return LogicalLayoutAttr::get(context, LogicalLayout::NCHW); +} + +inline PhysicalLayoutAttr getDenseNCHWLayout(mlir::MLIRContext* context) { + return PhysicalLayoutAttr::get(context, PhysicalLayout::DenseNCHW); +} + +inline PhysicalLayoutAttr getNCHWRowStripLayout(mlir::MLIRContext* context) { + return PhysicalLayoutAttr::get(context, PhysicalLayout::NCHWRowStrip); +} + +inline PhysicalLayoutAttr getNHWCRowStripLayout(mlir::MLIRContext* context) { + return PhysicalLayoutAttr::get(context, PhysicalLayout::NHWCRowStrip); +} + +inline PhysicalLayoutAttr getFragmentedLayout(mlir::MLIRContext* context) { + return PhysicalLayoutAttr::get(context, PhysicalLayout::Fragmented); +} + +inline BlueprintModeAttr getFragmentAssemblyMode(mlir::MLIRContext* context) { + return BlueprintModeAttr::get(context, BlueprintMode::FragmentAssembly); +} + +inline BlueprintModeAttr getPhysicalViewMode(mlir::MLIRContext* context) { + return BlueprintModeAttr::get(context, BlueprintMode::PhysicalView); +} + +inline bool isPhysicalView(std::optional mode) { + return mode && *mode == BlueprintMode::PhysicalView; +} + +inline bool isFragmentAssembly(std::optional mode) { + return mode && *mode == BlueprintMode::FragmentAssembly; +} + +inline std::optional getSelectedPhysicalLayout(mlir::Operation* op) { + auto attr = op->getAttrOfType(kSelectedLayoutAttrName); + return attr ? std::optional(attr.getValue()) : std::nullopt; +} bool hasCanonicalContiguousRowMajorFragments(mlir::RankedTensorType logicalType, llvm::ArrayRef offsets, diff --git a/src/PIM/Dialect/Spatial/SpatialOpsAsm.cpp b/src/PIM/Dialect/Spatial/SpatialOpsAsm.cpp index 1f1f90e..608b002 100644 --- a/src/PIM/Dialect/Spatial/SpatialOpsAsm.cpp +++ b/src/PIM/Dialect/Spatial/SpatialOpsAsm.cpp @@ -616,7 +616,7 @@ void SpatBlueprintOp::print(OpAsmPrinter& printer) { printer << " sizes "; printCompressedIntegerList(printer, getFragmentSizes()); printer << " map " << getIndexMap(); - if (std::optional mode = getMode()) + if (auto mode = getMode()) printer << " mode " << *mode; if (std::optional> operandIndices = getFragmentOperandIndices()) { printer << " operandIndices "; @@ -712,14 +712,25 @@ ParseResult SpatBlueprintOp::parse(OpAsmParser& parser, OperationState& result) if (operands.size() != operandTypes.size()) return parser.emitError(parser.getCurrentLocation(), "number of fragment operands and types must match"); + auto logicalLayoutValue = symbolizeLogicalLayout(logicalLayout.getValue()); + auto physicalLayoutValue = symbolizePhysicalLayout(physicalLayout.getValue()); + if (!logicalLayoutValue || !physicalLayoutValue) + return parser.emitError(parser.getCurrentLocation(), "unknown Blueprint layout"); + std::optional modeValue; + if (mode) { + modeValue = symbolizeBlueprintMode(mode.getValue()); + if (!modeValue) + return parser.emitError(parser.getCurrentLocation(), "unknown Blueprint mode"); + } + auto& builder = parser.getBuilder(); - result.addAttribute("logicalLayout", logicalLayout); - result.addAttribute("physicalLayout", physicalLayout); + result.addAttribute("logicalLayout", LogicalLayoutAttr::get(builder.getContext(), *logicalLayoutValue)); + result.addAttribute("physicalLayout", PhysicalLayoutAttr::get(builder.getContext(), *physicalLayoutValue)); result.addAttribute("fragmentOffsets", builder.getDenseI64ArrayAttr(fragmentOffsets)); result.addAttribute("fragmentSizes", builder.getDenseI64ArrayAttr(fragmentSizes)); result.addAttribute("indexMap", indexMap); - if (mode) - result.addAttribute("mode", mode); + if (modeValue) + result.addAttribute("mode", BlueprintModeAttr::get(builder.getContext(), *modeValue)); if (!fragmentOperandIndices.empty()) result.addAttribute("fragmentOperandIndices", builder.getDenseI64ArrayAttr(fragmentOperandIndices)); if (!fragmentSourceSlots.empty()) diff --git a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp index 7001d62..f753832 100644 --- a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp +++ b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp @@ -17,7 +17,6 @@ #include "src/Accelerators/PIM/Common/IR/AffineUtils.hpp" #include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp" #include "src/Accelerators/PIM/Common/PimCommon.hpp" -#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" #include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" @@ -405,7 +404,7 @@ static LogicalResult verifyConcatTypes(Operation* op, ValueRange inputs, Value o LogicalResult SpatConcatOp::verify() { return verifyConcatTypes(getOperation(), getInputs(), getOutput(), getAxis()); } LogicalResult SpatConcatPlanOp::verify() { - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logicalLayout = \"nchw\""); auto outputType = dyn_cast(getOutput().getType()); if (!outputType || !outputType.hasStaticShape() || outputType.getRank() != 4) @@ -415,11 +414,11 @@ LogicalResult SpatConcatPlanOp::verify() { return verifyConcatTypes(getOperation(), getInputs(), getOutput(), getAxis()); } -static bool isKnownLogicalLayout(StringRef layout) { return layout == "nchw"; } +static bool isKnownLogicalLayout(LogicalLayout layout) { return layout == LogicalLayout::NCHW; } -static bool isKnownPhysicalLayout(StringRef layout) { - return layout == "dense_nchw" || layout == "nchw_row_strip" || layout == "nhwc_row_strip" - || layout == "fragmented"; +static bool isKnownPhysicalLayout(PhysicalLayout layout) { + return layout == PhysicalLayout::DenseNCHW || layout == PhysicalLayout::NCHWRowStrip + || layout == PhysicalLayout::NHWCRowStrip || layout == PhysicalLayout::Fragmented; } static LogicalResult verifyPlanTensorTypes(Operation* op, Value input, Value output, StringRef kind) { @@ -489,7 +488,7 @@ LogicalResult SpatResizeNearestPlanOp::verify() { if (!inputType.hasStaticShape() || !outputType.hasStaticShape() || inputType.getRank() != 4 || outputType.getRank() != 4) return emitError("requires static rank-4 input and output tensors"); - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logical layout \"nchw\""); if (llvm::any_of(inputType.getShape(), [](int64_t dim) { return dim <= 0; }) || llvm::any_of(outputType.getShape(), [](int64_t dim) { return dim <= 0; })) @@ -505,7 +504,7 @@ LogicalResult SpatMaxPool2DPlanOp::verify() { if (!inputType.hasStaticShape() || !outputType.hasStaticShape() || inputType.getRank() != 4 || outputType.getRank() != 4) return emitError("requires static rank-4 input and output tensors"); - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logical layout \"nchw\""); if (getKernelShape().size() != 2 || getStrides().size() != 2 || getDilations().size() != 2) return emitError("requires two kernel, stride, and dilation values"); @@ -526,7 +525,7 @@ LogicalResult SpatGlobalAveragePoolPlanOp::verify() { if (!inputType.hasStaticShape() || !outputType.hasStaticShape() || inputType.getRank() != 4 || outputType.getRank() != 4) return emitError("requires static rank-4 input and output tensors"); - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logical layout \"nchw\""); if (inputType.getDimSize(0) != 1 || outputType.getDimSize(0) != 1 || inputType.getDimSize(1) != outputType.getDimSize(1) @@ -552,7 +551,7 @@ LogicalResult SpatBiasAddPlanOp::verify() { return emitError("requires matching input and output tensor types"); if (outputType.getRank() != 4) return emitError("requires rank-4 input/output tensors"); - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logical layout \"nchw\""); if (biasType.getElementType() != outputType.getElementType()) return emitError("requires bias element type to match the output element type"); @@ -580,14 +579,31 @@ LogicalResult SpatAddPlanOp::verify() { return emitError("requires matching operand and output tensor types"); if (outputType.getRank() != 4) return emitError("requires rank-4 operands and output"); - if (getLogicalLayout() != "nchw") + if (getLogicalLayout() != LogicalLayout::NCHW) return emitError("requires logical layout \"nchw\""); return success(); } LogicalResult SpatBlueprintOp::verify() { auto modeAttr = getModeAttr(); - bool isFragmentAssembly = modeAttr && modeAttr.getValue() == "fragment_assembly"; + bool isPhysicalView = modeAttr && modeAttr.getValue() == BlueprintMode::PhysicalView; + bool isFragmentAssembly = modeAttr && modeAttr.getValue() == BlueprintMode::FragmentAssembly; + if (isPhysicalView) { + auto inputType = dyn_cast(getInput().getType()); + auto outputType = dyn_cast(getOutput().getType()); + if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape()) + return emitError("physical view requires static ranked tensor input and output"); + if (inputType.getRank() != outputType.getRank() + 1) + return emitError("physical view requires one leading physical slot dimension"); + if (!getFragments().empty() || !getFragmentOffsets().empty() || !getFragmentSizes().empty() + || getFragmentOperandIndicesAttr() || getFragmentSourceSlotsAttr() + || getFragmentSourceOffsetsAttr() || getFragmentStridesAttr() + || getConflictPolicyAttr() || getCoveragePolicyAttr()) + return emitError("physical view does not accept fragment assembly metadata"); + if (!isKnownLogicalLayout(getLogicalLayout()) || !isKnownPhysicalLayout(getPhysicalLayout())) + return emitError("physical view requires known logical and physical layouts"); + return success(); + } if (!isFragmentAssembly && failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.blueprint"))) return failure(); if (!isKnownLogicalLayout(getLogicalLayout())) diff --git a/src/PIM/Dialect/Spatial/SpatialTargetInfo.hpp b/src/PIM/Dialect/Spatial/SpatialTargetInfo.hpp new file mode 100644 index 0000000..b7d02c5 --- /dev/null +++ b/src/PIM/Dialect/Spatial/SpatialTargetInfo.hpp @@ -0,0 +1,37 @@ +#pragma once + +#include +#include + +namespace onnx_mlir::spatial { + +struct MatrixUnitShape { + size_t rows = 128; + size_t columns = 128; +}; + +enum class ConvLoweringStrategy : uint8_t { + Auto, + Legacy, + Depthwise, + PackedIm2Col, + StreamedPatch, + StreamedPacked, + OutputChannelTiled, + InputKTiled, + Tiled2D, +}; + +struct SpatialTargetInfo { + MatrixUnitShape matrixShape; + size_t matrixUnitsPerProcessor = 64; + size_t processorCount = 1; + size_t vectorWidth = 16; + + uint64_t convIm2colMaxElements = 1ull << 20; + uint64_t convStreamChunkPositions = 1024; + ConvLoweringStrategy convLoweringStrategy = ConvLoweringStrategy::Auto; + bool useExperimentalConvImplementation = false; +}; + +} // namespace onnx_mlir::spatial diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp index 0fc3985..712110c 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp @@ -387,7 +387,7 @@ static void collectClosure(Value value, Block &body, const DeferredInputPlan &pl bool isDeferredFragmentAssemblyInput(Value input, size_t processorCount) { auto blueprint = input.getDefiningOp(); - if (!blueprint || blueprint.getMode() != "fragment_assembly") + if (!blueprint || !isFragmentAssembly(blueprint.getMode())) return false; return llvm::all_of(getBlueprintFragments(blueprint), [&](Value fragment) { return getProducerValueRef(fragment, nullptr, processorCount).has_value(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp index fcb6729..9cd9444 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp @@ -395,7 +395,7 @@ static LogicalResult buildExchanges(func::FuncOp funcOp, DeferredTransferPlan& p static LogicalResult retargetBlueprint(DeferredTransferPlan& plan, SpatBlueprintOp blueprint, GraphBatchPublicationCache& publicationCache) { - if (blueprint.getMode() != "fragment_assembly") + if (!isFragmentAssembly(blueprint.getMode())) return success(); bool escapesScheduledGraph = llvm::any_of( blueprint.getOutput().getUses(), [](OpOperand &use) { diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MergeComputeNodesPass.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MergeComputeNodesPass.cpp deleted file mode 100644 index d33a6dc..0000000 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MergeComputeNodesPass.cpp +++ /dev/null @@ -1,128 +0,0 @@ -#include "mlir/Pass/Pass.h" - -#include "DeferredCommunicationRealization.hpp" -#include "ScheduledComputeMaterialization.hpp" -#include "ScheduledComputeReport.hpp" -#include "ScheduledComputeVerification.hpp" -#include "Scheduling/MergeSchedulingAnalysis.hpp" -#include "SpatialDataflowCsvExporter.hpp" -#include "src/Accelerators/PIM/Common/PimCommon.hpp" -#include "src/Accelerators/PIM/Common/Support/DebugDump.hpp" -#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp" -#include "src/Accelerators/PIM/Pass/PIMPasses.h" - -using namespace mlir; - -namespace onnx_mlir { -namespace spatial { -namespace { - -struct MergeComputeNodesPass final : PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(MergeComputeNodesPass) - - MergeComputeNodesPass() = default; - explicit MergeComputeNodesPass(const SchedulingTarget& schedulingTarget) - : target(schedulingTarget), hasTarget(true) {} - - StringRef getArgument() const override { return "pim-merge-compute-nodes"; } - StringRef getDescription() const override { - return "Materialize scheduled Spatial compute with deferred communication placeholders."; - } - - void runOnOperation() override { - ModuleOp moduleOp = getOperation(); - if (!hasTarget || target.processorCount == 0 || target.residentWeightCapacity == 0 || target.transferWidthBytes == 0 - || target.interProcessorLatencyNs.size() != target.processorCount * target.processorCount - || (target.processorCount > 1 && target.averageInterProcessorLatencyNs == 0)) { - moduleOp.emitError("MergeComputeNodes requires an explicit valid Spatial scheduling target"); - signalPassFailure(); - return; - } - auto entryFunc = getPimEntryFunc(moduleOp); - if (failed(entryFunc)) { - moduleOp.emitError("failed to locate the PIM entry function during MergeComputeNodes"); - signalPassFailure(); - return; - } - - func::FuncOp funcOp = *entryFunc; - MergeScheduleResult logicalSchedule = MergeSchedulingAnalysis(funcOp, target).getResult(); - PatternRewriter rewriter(moduleOp.getContext()); - FailureOr materialization = - materializeScheduledCompute(funcOp, logicalSchedule, rewriter); - if (failed(materialization)) { - signalPassFailure(); - return; - } - // Phase 1 is intentionally dumped before its verifier: malformed deferred - // payloads must be diagnosed from the producer-owned body. - dumpModule(moduleOp, "spatial3_scheduled_no_comm", /*assumeVerified=*/true); - if (failed(verifyMaterializedScheduleMapping(funcOp, - logicalSchedule, - materialization->peftClassPlans, - materialization->graphComputeToBlockMap, - materialization->materializedSchedules))) { - moduleOp.emitError("scheduled Spatial materialization mapping verification failed"); - signalPassFailure(); - return; - } - if (failed(verifyDeferredTransferPhase1Invariants(funcOp))) { - moduleOp.emitError("scheduled Spatial deferred communication verification failed"); - signalPassFailure(); - return; - } - if (failed(verifyScheduledMaterializationRecords(materialization->materializedSchedules))) { - moduleOp.emitError("scheduled Spatial materialization record verification failed"); - signalPassFailure(); - return; - } - if (failed(verifyScheduledSpatialInvariants(funcOp))) { - moduleOp.emitError("scheduled Spatial phase 1 verification failed"); - signalPassFailure(); - return; - } - - SpatialDataflowExportStage exportMode = getSpatialDataflowExportStage(); - if (shouldExportSpatialDataflowStage(exportMode, SpatialDataflowExportStage::Spatial3) - && failed(exportSpatialDataflowCsvScheduled( - funcOp, materialization->materializedSchedules, "spatial3_scheduled_no_comm", "spatial3"))) { - signalPassFailure(); - return; - } - - dumpScheduledComputeReport( - moduleOp, funcOp, logicalSchedule, materialization->peftClassPlans, materialization->materializedSchedules); - if (failed(realizeDeferredCommunication(funcOp, *materialization, target))) { - moduleOp.emitError("MergeComputeNodes phase 2 communication realization failed"); - signalPassFailure(); - return; - } - dumpModule(moduleOp, "spatial4_scheduled", /*assumeVerified=*/true); - if (failed(verifyScheduledResultsLive(materialization->materializedSchedules)) - || failed(verifyScheduledSpatialInvariants(funcOp))) { - moduleOp.emitError("scheduled Spatial phase 2 verification failed"); - signalPassFailure(); - return; - } - if (shouldExportSpatialDataflowStage(exportMode, SpatialDataflowExportStage::Spatial4) - && failed(exportSpatialDataflowCsvScheduled( - funcOp, materialization->materializedSchedules, "spatial4_scheduled", "spatial4"))) { - signalPassFailure(); - } - } - -private: - SchedulingTarget target; - bool hasTarget = false; -}; - -} // namespace -} // namespace spatial - -std::unique_ptr createMergeComputeNodesPass() { return std::make_unique(); } - -std::unique_ptr createMergeComputeNodesPass(const spatial::SchedulingTarget& target) { - return std::make_unique(target); -} - -} // namespace onnx_mlir diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledComputePlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledComputePlanning.cpp index c98f3ae..8cdf867 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledComputePlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledComputePlanning.cpp @@ -24,7 +24,7 @@ bool requiresScheduledPublication(Value value, DenseSet &visited) { SpatDeferredCommunicationOp>(user)) return false; auto blueprint = dyn_cast(user); - return !blueprint || blueprint.getMode() != "fragment_assembly" + return !blueprint || !isFragmentAssembly(blueprint.getMode()) || requiresScheduledPublication(blueprint.getOutput(), visited); }); } diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.cpp new file mode 100644 index 0000000..4c54ac6 --- /dev/null +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.cpp @@ -0,0 +1,279 @@ +#include "ScheduledSpatialPasses.hpp" + +#include "mlir/Pass/Pass.h" + +#include "DeferredCommunicationRealization.hpp" +#include "ScheduledComputeReport.hpp" +#include "ScheduledComputeVerification.hpp" +#include "SpatialDataflowCsvExporter.hpp" +#include "src/Accelerators/PIM/Common/PimCommon.hpp" +#include "src/Accelerators/PIM/Common/Support/DebugDump.hpp" +#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp" +#include "src/Accelerators/PIM/Pass/PIMPasses.h" + +using namespace mlir; + +namespace onnx_mlir { +namespace spatial { +namespace { + +static bool hasValidTarget(const SchedulingTarget& target) { + return target.processorCount != 0 && target.residentWeightCapacity != 0 + && target.transferWidthBytes != 0 + && target.interProcessorLatencyNs.size() == target.processorCount * target.processorCount + && (target.processorCount == 1 || target.averageInterProcessorLatencyNs != 0); +} + +static FailureOr requireEntry(ModuleOp moduleOp, StringRef passName) { + auto entry = getPimEntryFunc(moduleOp); + if (failed(entry)) { + moduleOp.emitError("failed to locate the PIM entry function during ") << passName; + return failure(); + } + return *entry; +} + +static LogicalResult requireState(ModuleOp moduleOp, + const std::shared_ptr& state, + StringRef passName) { + if (state && state->logicalSchedule && state->materialization) + return success(); + moduleOp.emitError() << passName << " requires scheduling state from ScheduleSpatialGraph"; + return failure(); +} + +struct ScheduleSpatialGraphPass final + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ScheduleSpatialGraphPass) + + ScheduleSpatialGraphPass() = default; + ScheduleSpatialGraphPass(const SchedulingTarget& target, + std::shared_ptr state) + : target(target), state(std::move(state)), hasTarget(true) {} + + StringRef getArgument() const override { return "schedule-spatial-graph"; } + StringRef getDescription() const override { + return "Schedule Spatial graph computes and materialize deferred communication boundaries."; + } + + void runOnOperation() override { + ModuleOp moduleOp = getOperation(); + if (!hasTarget || !hasValidTarget(target) || !state) { + moduleOp.emitError("ScheduleSpatialGraph requires an explicit target and shared pass state"); + signalPassFailure(); + return; + } + auto entry = requireEntry(moduleOp, "ScheduleSpatialGraph"); + if (failed(entry)) { + signalPassFailure(); + return; + } + + MergeSchedulingAnalysis analysis(*entry, target); + MergeScheduleResult schedule = std::move(analysis.getResult()); + PatternRewriter rewriter(moduleOp.getContext()); + FailureOr materialization = + materializeScheduledCompute(*entry, schedule, rewriter); + if (failed(materialization)) { + signalPassFailure(); + return; + } + + state->logicalSchedule = std::move(schedule); + state->materialization = std::move(*materialization); + dumpModule(moduleOp, "spatial3_scheduled_no_comm", /*assumeVerified=*/true); + } + +private: + SchedulingTarget target; + std::shared_ptr state; + bool hasTarget = false; +}; + +struct VerifyScheduledSpatialPass final + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(VerifyScheduledSpatialPass) + + explicit VerifyScheduledSpatialPass(std::shared_ptr state = {}) + : state(std::move(state)) {} + + StringRef getArgument() const override { return "verify-scheduled-spatial"; } + StringRef getDescription() const override { + return "Verify scheduled Spatial compute, deferred communication, and materialization records."; + } + + void runOnOperation() override { + ModuleOp moduleOp = getOperation(); + if (failed(requireState(moduleOp, state, "VerifyScheduledSpatial"))) { + signalPassFailure(); + return; + } + auto entry = requireEntry(moduleOp, "VerifyScheduledSpatial"); + if (failed(entry)) { + signalPassFailure(); + return; + } + const auto& schedule = *state->logicalSchedule; + const auto& materialization = *state->materialization; + if (failed(verifyMaterializedScheduleMapping( + *entry, schedule, materialization.peftClassPlans, + materialization.graphComputeToBlockMap, + materialization.materializedSchedules)) + || failed(verifyDeferredTransferPhase1Invariants(*entry)) + || failed(verifyScheduledMaterializationRecords(materialization.materializedSchedules)) + || failed(verifyScheduledSpatialInvariants(*entry))) { + moduleOp.emitError("scheduled Spatial phase verification failed"); + signalPassFailure(); + return; + } + SpatialDataflowExportStage exportMode = getSpatialDataflowExportStage(); + if (shouldExportSpatialDataflowStage(exportMode, SpatialDataflowExportStage::Spatial3) + && failed(exportSpatialDataflowCsvScheduled( + *entry, materialization.materializedSchedules, + "spatial3_scheduled_no_comm", "spatial3"))) { + signalPassFailure(); + return; + } + dumpScheduledComputeReport( + moduleOp, *entry, schedule, materialization.peftClassPlans, + materialization.materializedSchedules); + } + +private: + std::shared_ptr state; +}; + +struct RealizeSpatialCommunicationPass final + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(RealizeSpatialCommunicationPass) + + RealizeSpatialCommunicationPass() = default; + RealizeSpatialCommunicationPass(const SchedulingTarget& target, + std::shared_ptr state) + : target(target), state(std::move(state)), hasTarget(true) {} + + StringRef getArgument() const override { return "realize-spatial-communication"; } + StringRef getDescription() const override { + return "Realize deferred Spatial communication after scheduled graph verification."; + } + + void runOnOperation() override { + ModuleOp moduleOp = getOperation(); + if (!hasTarget || !hasValidTarget(target) + || failed(requireState(moduleOp, state, "RealizeSpatialCommunication"))) { + signalPassFailure(); + return; + } + auto entry = requireEntry(moduleOp, "RealizeSpatialCommunication"); + if (failed(entry)) { + signalPassFailure(); + return; + } + if (failed(realizeDeferredCommunication(*entry, *state->materialization, target))) { + moduleOp.emitError("Spatial communication realization failed"); + signalPassFailure(); + return; + } + dumpModule(moduleOp, "spatial4_scheduled", /*assumeVerified=*/true); + } + +private: + SchedulingTarget target; + std::shared_ptr state; + bool hasTarget = false; +}; + +struct VerifyRealizedSpatialPass final + : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(VerifyRealizedSpatialPass) + + explicit VerifyRealizedSpatialPass(std::shared_ptr state = {}) + : state(std::move(state)) {} + + StringRef getArgument() const override { return "verify-realized-spatial"; } + StringRef getDescription() const override { + return "Verify realized Spatial communication and scheduled result liveness."; + } + + void runOnOperation() override { + ModuleOp moduleOp = getOperation(); + if (failed(requireState(moduleOp, state, "VerifyRealizedSpatial"))) { + signalPassFailure(); + return; + } + auto entry = requireEntry(moduleOp, "VerifyRealizedSpatial"); + if (failed(entry)) { + signalPassFailure(); + return; + } + const auto& records = state->materialization->materializedSchedules; + bool deferredRemains = false; + (*entry).walk([&](SpatDeferredCommunicationOp deferred) { + if (deferredRemains) + return; + deferred.emitOpError("realized Spatial graph still contains deferred communication"); + deferredRemains = true; + }); + if (deferredRemains + || failed(verifyScheduledResultsLive(records)) + || failed(verifyScheduledSpatialInvariants(*entry))) { + moduleOp.emitError("realized Spatial communication verification failed"); + signalPassFailure(); + return; + } + SpatialDataflowExportStage exportMode = getSpatialDataflowExportStage(); + if (shouldExportSpatialDataflowStage(exportMode, SpatialDataflowExportStage::Spatial4) + && failed(exportSpatialDataflowCsvScheduled( + *entry, records, "spatial4_scheduled", "spatial4"))) + signalPassFailure(); + } + +private: + std::shared_ptr state; +}; + +} // namespace + +std::unique_ptr createScheduleSpatialGraphPass() { + return std::make_unique(); +} + +std::unique_ptr createScheduleSpatialGraphPass(const SchedulingTarget& target) { + return std::make_unique( + target, std::make_shared()); +} + +std::unique_ptr createScheduleSpatialGraphPass( + const SchedulingTarget& target, std::shared_ptr state) { + return std::make_unique(target, std::move(state)); +} + +std::unique_ptr createVerifyScheduledSpatialPass() { + return std::make_unique(); +} + +std::unique_ptr createVerifyScheduledSpatialPass( + std::shared_ptr state) { + return std::make_unique(std::move(state)); +} + +std::unique_ptr createRealizeSpatialCommunicationPass() { + return std::make_unique(); +} + +std::unique_ptr createRealizeSpatialCommunicationPass( + const SchedulingTarget& target, std::shared_ptr state) { + return std::make_unique(target, std::move(state)); +} + +std::unique_ptr createVerifyRealizedSpatialPass() { + return std::make_unique(); +} + +std::unique_ptr createVerifyRealizedSpatialPass( + std::shared_ptr state) { + return std::make_unique(std::move(state)); +} + +} // namespace spatial +} // namespace onnx_mlir diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.hpp new file mode 100644 index 0000000..4c98411 --- /dev/null +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/ScheduledSpatialPasses.hpp @@ -0,0 +1,16 @@ +#pragma once + +#include "ScheduledComputeMaterialization.hpp" +#include "Scheduling/MergeSchedulingAnalysis.hpp" + +#include +#include + +namespace onnx_mlir::spatial { + +struct ScheduledSpatialState { + std::optional logicalSchedule; + std::optional materialization; +}; + +} // namespace onnx_mlir::spatial diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp index 0f0a5dd..a02a0f5 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/PeftScheduler.cpp @@ -181,7 +181,7 @@ FailureOr buildLanePublicationSignatures(SpatComputeB for (auto [useIndex, use] : llvm::enumerate(result.getUses())) { auto blueprint = dyn_cast(use.getOwner()); - if (!blueprint || blueprint.getMode() != "fragment_assembly") + if (!blueprint || !isFragmentAssembly(blueprint.getMode())) continue; auto operandIndices = blueprint.getFragmentOperandIndices(); auto sourceSlots = blueprint.getFragmentSourceSlots(); diff --git a/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp b/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp index 3426795..a178d7a 100644 --- a/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp @@ -14,7 +14,6 @@ #include "src/Accelerators/PIM/Common/IR/TensorSliceUtils.hpp" #include "src/Accelerators/PIM/Common/PimCommon.hpp" #include "src/Accelerators/PIM/Common/Support/DebugDump.hpp" -#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/ComputeRegionBuilder.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/ComputeGraph.hpp" diff --git a/src/PIM/Pass/PIMPasses.h b/src/PIM/Pass/PIMPasses.h index cc5d513..ebdc169 100644 --- a/src/PIM/Pass/PIMPasses.h +++ b/src/PIM/Pass/PIMPasses.h @@ -9,18 +9,40 @@ namespace onnx_mlir { namespace spatial { struct SchedulingTarget; +struct ScheduledSpatialState; +struct SpatialTargetInfo; + +std::unique_ptr createScheduleSpatialGraphPass(); +std::unique_ptr createScheduleSpatialGraphPass(const SchedulingTarget& target); +std::unique_ptr createScheduleSpatialGraphPass( + const SchedulingTarget& target, + std::shared_ptr state); +std::unique_ptr createVerifyScheduledSpatialPass(); +std::unique_ptr createVerifyScheduledSpatialPass( + std::shared_ptr state); +std::unique_ptr createRealizeSpatialCommunicationPass(); +std::unique_ptr createRealizeSpatialCommunicationPass( + const SchedulingTarget& target, + std::shared_ptr state); +std::unique_ptr createVerifyRealizedSpatialPass(); +std::unique_ptr createVerifyRealizedSpatialPass( + std::shared_ptr state); } std::unique_ptr createONNXToSpatialPass(); +std::unique_ptr createONNXToSpatialPass(const spatial::SpatialTargetInfo& target); std::unique_ptr createSpatialLayoutPlanningPass(); +std::unique_ptr createSpatialLayoutPlanningPass(const spatial::SpatialTargetInfo& target); std::unique_ptr createLowerSpatialPlansPass(); +std::unique_ptr createLowerSpatialPlansPass(const spatial::SpatialTargetInfo& target); std::unique_ptr createSpatialToPimPass(); -std::unique_ptr createPimBufferizationPass(); +std::unique_ptr createPimBufferizationPreparationPass(); +std::unique_ptr createPimOneShotBufferizationPass(); +std::unique_ptr createPimMemoryNormalizationPass(); +std::unique_ptr createPimBufferizationVerificationPass(); -std::unique_ptr createMergeComputeNodesPass(); -std::unique_ptr createMergeComputeNodesPass(const spatial::SchedulingTarget& target); std::unique_ptr createTrivialGraphComputeMergePass(); std::unique_ptr createTrivialGraphComputeMergePass( diff --git a/src/PIM/PimAccelerator.cpp b/src/PIM/PimAccelerator.cpp index 2c2ae5e..0516344 100644 --- a/src/PIM/PimAccelerator.cpp +++ b/src/PIM/PimAccelerator.cpp @@ -71,13 +71,19 @@ void PimAccelerator::registerDialects(mlir::DialectRegistry& registry) const { void PimAccelerator::registerPasses(int optLevel) const { LLVM_DEBUG(llvm::dbgs() << "Registering passes for PIM accelerator\n"); - registerPass(createONNXToSpatialPass); - registerPass(createSpatialLayoutPlanningPass); - registerPass(createLowerSpatialPlansPass); + mlir::registerPass([] { return createONNXToSpatialPass(); }); + mlir::registerPass([] { return createSpatialLayoutPlanningPass(); }); + mlir::registerPass([] { return createLowerSpatialPlansPass(); }); registerPass(createSpatialToPimPass); - registerPass(createPimBufferizationPass); + registerPass(createPimBufferizationPreparationPass); + registerPass(createPimOneShotBufferizationPass); + registerPass(createPimMemoryNormalizationPass); + registerPass(createPimBufferizationVerificationPass); mlir::registerPass([] { return createTrivialGraphComputeMergePass(); }); - mlir::registerPass([] { return createMergeComputeNodesPass(); }); + mlir::registerPass([] { return spatial::createScheduleSpatialGraphPass(); }); + mlir::registerPass([] { return spatial::createVerifyScheduledSpatialPass(); }); + mlir::registerPass([] { return spatial::createRealizeSpatialCommunicationPass(); }); + mlir::registerPass([] { return spatial::createVerifyRealizedSpatialPass(); }); registerPass(createPimHostConstantFoldingPass); registerPass(createPimInstructionSelectionPass); registerPass(createPimLocalMemoryPlanningPass); diff --git a/validation/networks/pimcomp_models/googlenet/googlenet-12.onnx b/validation/networks/pimcomp_models/googlenet/googlenet-12.onnx index 1865acc0a19390cd1f562b32f4f87751b873fdae..97a4db3480c969c153e3549f588b2e1860852633 100644 GIT binary patch delta 1526 zcmWm6XMhj_07r30D0EvOGb?v*`e!4vylSXQFf9FQc)^NWhs;@QY2NS znpBq>Qd4S4lG;*7>PkJSFAb!jG?JZV7ila_q^UHM=F&oTmEB}_X(_E_57|>%%U;q( z_LjD?kF=A0rM+~Jjss~ji?$-&Z1x=Rn~DZQk(^pQj4P&rHvmm}mz zIZBR}zS2+n%Q14S43L2`NCwMsa=Z+Y6XZlWNluobGE9o)6e*EYWw?xxkupj~%NRLL z#>zMuFB7Cx%4DLPE@w!&oGFuJvP_Ywa+XY!>2kK5BQs>C%#ztMN6wYGqRf-?F1cIok$YvU+$Zl#k?N z`9waI&*XFYLcWx*=-*mg{T;nqH+{Ql_-j;Q7x)Rji?#5;{QkOs1tRgUeu2U(J>&aq20jwaDGnnm+y z5xd52v3saD9 zjdig;E{`kX%D5`7jty~5Y>Z8@xnL_>3OX%sI<{R=-QtO5W#c=w>oT=Vae4dV2_@y@ cird$2P*kt5%DM-aZ7HZzsaJ)f9coqk4=vWv8UO$Q delta 1489 zcmWm6XTS&q00!YcgzU4Eor;9)k+MP=84+ z=jYwuJFhRjFDPEDe9383r_O9NX4LrcaP(ni`!J2^_)%h7U-bdZkHNjggxIaaz#H|Z|N$??)ddP*

g%2FAvCr@{l|%kI19)m^>~| z$dmGv?31VE8F^Noljr3Hc~M@Hm*o|CRbG?VzzLKxy8~IkglkepR`B8q7pXC?%ReqD-5onOq7js@&BWIREUaEDJn;ms2T@GwKyoMM~$c%wW4;^iMnxc z)Qdyn&^RpWM}spaXPC>lqTI5L_>vuGYIqGhy-*3l-~M!PsF+Q-pxOmv8j(J4Ad zmpC@MMz`o5$Hno{BYH-!=p83SpXeL?qJNwi17cvD6oX=L42hFtXbg+tF(O9BDKRQW z$Cwx!<6?YFh>0;NPL0ViB~FW}F)dDy=`kbDh%@7?I6KaXnQ?BM7w5+Xaba8(7ssra z9dlxC%nM_FEQp1%C@zV`u_TtpvRED~;?lS*E{~OQMO+!HVs)&EwQ*Ifi}kS~HpZsd j99!b**c#hnd+dmvaZT(h*vsyMmMe>PE>dE@a;5$Qj6=l; diff --git a/validation/operations/validation_results.csv b/validation/operations/validation_results.csv index c14c1c3..76cf774 100644 --- a/validation/operations/validation_results.csv +++ b/validation/operations/validation_results.csv @@ -1,169 +1,169 @@ Operation,Result,Compile,Host mem,Cores mem,Cores,Xbars,Latency,Power,Energy -add/after_gemm,PASS,0.072 s,0.01 MiB,0.01 MiB,5,4,0.007784 ms,104.703618 mW,815012.960000 pJ -add/basic,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -add/broadcast_row,PASS,0.048 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -add/channel_broadcast_1024,PASS,0.051 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ -add/leading_dimension_broadcast,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -concat/channel_axis,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,0.000457 ms,78.157549 mW,35718.000000 pJ -concat/negative_axis,PASS,0.061 s,0.00 MiB,0.00 MiB,1,0,0.001043 ms,78.092042 mW,81450.000000 pJ -concat/three_inputs_channel_axis,PASS,0.054 s,0.00 MiB,0.00 MiB,1,0,0.000644 ms,78.149068 mW,50328.000000 pJ -conv/batch_2,PASS,0.059 s,0.00 MiB,0.00 MiB,2,2,0.013694 ms,82.623885 mW,1131451.480000 pJ -conv/batch_4_pointwise,PASS,0.056 s,0.00 MiB,0.01 MiB,5,4,0.003932 ms,116.078576 mW,456420.960000 pJ -conv/depthwise_1024_channels,PASS,0.081 s,0.19 MiB,0.38 MiB,129,128,0.220751 ms,178.454307 mW,39393966.720000 pJ -conv/depthwise_grouped,PASS,0.059 s,0.01 MiB,0.00 MiB,5,4,0.006024 ms,108.326521 mW,652558.960000 pJ -conv/dilated_3x3,PASS,0.059 s,0.00 MiB,0.00 MiB,3,3,0.004045 ms,110.234541 mW,445898.720000 pJ -conv/dynamic,PASS,0.059 s,0.00 MiB,0.00 MiB,5,0,0.001835 ms,92.281199 mW,169336.000000 pJ -conv/explicit_padding,PASS,0.058 s,0.00 MiB,0.00 MiB,4,4,0.004327 ms,115.794768 mW,501043.960000 pJ -conv/grouped_many_groups,PASS,0.505 s,0.05 MiB,0.09 MiB,65,64,0.181845 ms,142.210104 mW,25860196.360000 pJ -conv/grouped_two_groups,PASS,0.061 s,0.00 MiB,0.00 MiB,3,2,0.005360 ms,101.459418 mW,543822.480000 pJ -conv/huge_pointwise_1024,PASS,0.643 s,0.01 MiB,0.01 MiB,1,64,0.028261 ms,133.488743 mW,3772525.360000 pJ -conv/huge_pointwise_1024_dynamic,PASS,0.081 s,8.04 MiB,12.61 MiB,168,0,2.627964 ms,169.518697 mW,445489032.000000 pJ -conv/kernel_3x3,PASS,0.056 s,0.00 MiB,0.00 MiB,3,3,0.003091 ms,115.862414 mW,358130.720000 pJ -conv/kernel_equals_input_spatial,PASS,0.060 s,0.00 MiB,0.00 MiB,1,2,0.008443 ms,83.863849 mW,708062.480000 pJ -conv/large_input_channels_1x1,PASS,0.088 s,0.01 MiB,0.01 MiB,1,8,0.017167 ms,89.416900 mW,1535019.920000 pJ -conv/large_output_channels_1x1,PASS,0.087 s,0.00 MiB,0.01 MiB,1,8,0.004964 ms,117.628106 mW,583905.920000 pJ -conv/large_spatial,PASS,0.055 s,0.00 MiB,0.01 MiB,6,6,0.004096 ms,129.015000 mW,528445.440000 pJ -conv/multi_channel,PASS,0.056 s,0.00 MiB,0.00 MiB,3,3,0.005148 ms,106.453520 mW,548022.720000 pJ -conv/non_square_kernel_1x3,PASS,0.056 s,0.00 MiB,0.00 MiB,5,5,0.004029 ms,123.600943 mW,497988.200000 pJ -conv/non_square_kernel_3x1,PASS,0.058 s,0.00 MiB,0.00 MiB,3,3,0.005526 ms,105.464843 mW,582798.720000 pJ -conv/non_uniform_stride,PASS,0.062 s,0.00 MiB,0.00 MiB,4,4,0.005808 ms,110.081433 mW,639352.960000 pJ -conv/pointwise_1x1,PASS,0.057 s,0.00 MiB,0.00 MiB,4,4,0.004539 ms,114.835858 mW,521239.960000 pJ -conv/pointwise_tiled_chain,PASS,0.771 s,0.01 MiB,0.02 MiB,2,80,0.084437 ms,102.289307 mW,8637002.200000 pJ -conv/real_asymmetric_padding,PASS,0.055 s,0.00 MiB,0.00 MiB,4,4,0.005232 ms,111.870214 mW,585304.960000 pJ -conv/relu_conv_store,PASS,0.076 s,0.02 MiB,0.10 MiB,32,32,0.064390 ms,243.916397 mW,15705776.800000 pJ -conv/same_lower_3x3,PASS,0.057 s,0.00 MiB,0.00 MiB,5,5,0.004700 ms,119.232170 mW,560391.200000 pJ -conv/same_padding_3x3,PASS,0.060 s,0.00 MiB,0.00 MiB,5,5,0.004700 ms,119.232170 mW,560391.200000 pJ -conv/simple,PASS,0.055 s,0.00 MiB,0.00 MiB,2,2,0.003148 ms,94.665972 mW,298008.480000 pJ -conv/stride_2,PASS,0.053 s,0.00 MiB,0.00 MiB,2,2,0.002827 ms,96.393873 mW,272505.480000 pJ -conv/with_bias_3x3,PASS,0.055 s,0.00 MiB,0.00 MiB,3,3,0.004898 ms,107.176546 mW,524950.720000 pJ -conv/with_constant,PASS,0.073 s,0.00 MiB,0.00 MiB,3,3,0.004273 ms,109.362677 mW,467306.720000 pJ -conv/without_kernel_shape_attr,PASS,0.057 s,0.00 MiB,0.00 MiB,3,3,0.003091 ms,115.862414 mW,358130.720000 pJ -conv/yolo11n_depthwise_head,PASS,0.509 s,4.82 MiB,15.92 MiB,160,720,4.105780 ms,492.826315 mW,2023436428.000010 pJ -conv/yolo11n_heavy,PASS,0.488 s,4.82 MiB,20.66 MiB,160,800,6.275385 ms,418.028199 mW,2623287892.000010 pJ -conv/yolo11n_stem,PASS,1.730 s,12.86 MiB,31.38 MiB,168,488,9.799204 ms,361.075380 mW,3538251304.000010 pJ -div/after_gemm,PASS,0.059 s,0.01 MiB,0.01 MiB,5,4,0.007784 ms,104.703618 mW,815012.960000 pJ -div/basic,PASS,0.048 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -div/channel_broadcast_1024,PASS,0.051 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ -div/leading_dimension_broadcast,PASS,0.063 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -div/runtime_scalar_rhs,PASS,0.057 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ -div/scalar_constant,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -gather/3d_input_axis1,PASS,0.050 s,0.00 MiB,0.00 MiB,1,0,0.000589 ms,78.081494 mW,45990.000000 pJ -gather/axis0_matrix_indices,PASS,0.050 s,0.00 MiB,0.00 MiB,1,0,0.000697 ms,78.068867 mW,54414.000000 pJ -gather/axis1,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,0.000801 ms,78.059925 mW,62526.000000 pJ -gather/negative_axis,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,0.001437 ms,78.033403 mW,112134.000000 pJ -gather/negative_indices,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,0.000376 ms,78.127660 mW,29376.000000 pJ -gemm/alpha_beta,PASS,0.056 s,0.01 MiB,0.01 MiB,5,4,0.007456 ms,105.272125 mW,784908.960000 pJ -gemm/bias_rank2_broadcast,PASS,0.056 s,0.00 MiB,0.01 MiB,5,4,0.007072 ms,105.979208 mW,749484.960000 pJ -gemm/dynamic,PASS,0.054 s,0.00 MiB,0.00 MiB,5,0,0.002421 ms,91.480793 mW,221475.000000 pJ -gemm/dynamic_alpha,PASS,0.056 s,0.00 MiB,0.00 MiB,5,0,0.003262 ms,91.415696 mW,298198.000000 pJ -gemm/dynamic_beta,PASS,0.059 s,0.00 MiB,0.00 MiB,5,0,0.004365 ms,91.316151 mW,398595.000000 pJ -gemm/dynamic_bias,PASS,0.056 s,0.00 MiB,0.00 MiB,5,0,0.002665 ms,91.445779 mW,243703.000000 pJ -gemm/dynamic_bias_alpha_beta,PASS,0.053 s,0.00 MiB,0.00 MiB,5,0,0.005629 ms,91.279268 mW,513811.000000 pJ -gemm/dynamic_transB,PASS,0.054 s,0.00 MiB,0.00 MiB,5,0,0.001301 ms,91.378171 mW,118883.000000 pJ -gemm/huge_1024,PASS,0.157 s,0.01 MiB,0.10 MiB,73,64,0.017522 ms,215.037402 mW,3767885.360000 pJ -gemm/large,PASS,0.067 s,0.02 MiB,0.03 MiB,17,16,0.011229 ms,140.152181 mW,1573768.840000 pJ -gemm/large_k_small_n,PASS,0.097 s,0.01 MiB,0.01 MiB,9,8,0.004748 ms,133.481449 mW,633769.920000 pJ -gemm/non_square,PASS,0.060 s,0.00 MiB,0.01 MiB,5,4,0.003527 ms,118.958310 mW,419565.960000 pJ -gemm/scalar_bias,PASS,0.054 s,0.00 MiB,0.01 MiB,5,4,0.007072 ms,105.979208 mW,749484.960000 pJ -gemm/simple,PASS,0.071 s,0.03 MiB,0.08 MiB,42,40,0.021640 ms,151.774196 mW,3284393.600000 pJ -gemm/small,PASS,0.052 s,0.00 MiB,0.00 MiB,2,2,0.004420 ms,90.144000 mW,398436.480000 pJ -gemm/small_k_large_n,PASS,0.089 s,0.01 MiB,0.02 MiB,17,8,0.007962 ms,131.005014 mW,1043061.920000 pJ -gemm/transA,PASS,0.057 s,0.00 MiB,0.01 MiB,5,4,0.005762 ms,109.140743 mW,628868.960000 pJ -gemm/transA_transB,PASS,0.053 s,0.00 MiB,0.01 MiB,5,4,0.005762 ms,109.140743 mW,628868.960000 pJ -gemm/transB,PASS,0.062 s,0.00 MiB,0.01 MiB,5,4,0.003527 ms,118.958310 mW,419565.960000 pJ -gemm/transB_with_bias,PASS,0.057 s,0.01 MiB,0.01 MiB,5,4,0.005046 ms,110.546762 mW,557818.960000 pJ -gemm/with_bias,PASS,0.056 s,0.01 MiB,0.01 MiB,5,4,0.005562 ms,108.767882 mW,604966.960000 pJ -gemv/constant,PASS,0.051 s,0.00 MiB,0.00 MiB,0,0,0.000000 ms,2.000000 mW,0.000000 pJ -gemv/simple,PASS,0.066 s,0.00 MiB,0.01 MiB,6,4,0.005160 ms,111.150380 mW,573535.960000 pJ -gemv/with_heterogeneous_constant,PASS,0.066 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ -gemv/with_homogeneous_constant,PASS,0.063 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ -gemv/with_scalar_constant,PASS,0.064 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ -matmul/basic,PASS,0.064 s,0.00 MiB,0.00 MiB,2,2,0.004420 ms,90.144000 mW,398436.480000 pJ -matmul/batched_3d,PASS,0.058 s,0.00 MiB,0.01 MiB,5,4,0.005958 ms,108.588949 mW,646972.960000 pJ -matmul/batched_3d_dynamic,PASS,0.053 s,0.00 MiB,0.00 MiB,4,0,0.001822 ms,92.192645 mW,167975.000000 pJ -matmul/batched_left_constant,PASS,0.058 s,0.00 MiB,0.02 MiB,9,8,0.008822 ms,114.385164 mW,1009105.920000 pJ -matmul/batched_lhs_broadcast,PASS,0.056 s,0.00 MiB,0.01 MiB,5,4,0.005681 ms,109.389361 mW,621440.960000 pJ -matmul/batched_rhs_broadcast,PASS,0.059 s,0.00 MiB,0.01 MiB,5,4,0.005958 ms,108.588949 mW,646972.960000 pJ -matmul/dynamic,PASS,0.054 s,0.00 MiB,0.00 MiB,5,0,0.001621 ms,91.421962 mW,148195.000000 pJ -matmul/huge_1024,PASS,0.142 s,0.01 MiB,0.10 MiB,73,64,0.017522 ms,215.037402 mW,3767885.360000 pJ -matmul/left_constant,PASS,0.064 s,0.00 MiB,0.01 MiB,5,4,0.005853 ms,108.861944 mW,637168.960000 pJ -matmul/matrix_vector,PASS,0.097 s,0.52 MiB,0.78 MiB,168,173,0.384660 ms,202.131271 mW,77751814.880000 pJ -matmul/vector_matrix,PASS,0.086 s,0.01 MiB,0.01 MiB,9,8,0.007409 ms,118.680243 mW,879301.920000 pJ -matmul/yolo_attention,PASS,0.417 s,1.02 MiB,43.44 MiB,168,0,8.151445 ms,170.003707 mW,1385775865.000000 pJ -mul/after_conv,PASS,0.060 s,0.00 MiB,0.00 MiB,4,3,0.005453 ms,107.639046 mW,586955.720000 pJ -mul/basic,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -mul/channel_broadcast_1024,PASS,0.060 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ -mul/leading_dimension_broadcast,PASS,0.058 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -mul/scalar_constant,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -pool/avg_basic,PASS,0.057 s,0.00 MiB,0.00 MiB,1,0,0.011939 ms,78.022112 mW,931506.000000 pJ -pool/avg_ceil_mode,PASS,0.054 s,0.00 MiB,0.00 MiB,1,0,0.004359 ms,78.033035 mW,340146.000000 pJ -pool/avg_explicit_padding,PASS,0.062 s,0.00 MiB,0.00 MiB,1,0,0.008822 ms,78.027205 mW,688356.000000 pJ -pool/avg_include_pad,PASS,0.055 s,0.00 MiB,0.00 MiB,1,0,0.008506 ms,78.016929 mW,663612.000000 pJ -pool/avg_large_channels,PASS,0.067 s,0.04 MiB,0.02 MiB,1,0,0.178249 ms,78.280327 mW,13953390.000000 pJ -pool/avg_non_uniform_stride,PASS,0.060 s,0.00 MiB,0.00 MiB,1,0,0.014513 ms,78.016537 mW,1132254.000000 pJ -pool/avg_real_asymmetric_padding,PASS,0.071 s,0.00 MiB,0.00 MiB,1,0,0.025206 ms,78.024756 mW,1966692.000000 pJ -pool/max_after_conv,PASS,0.066 s,0.00 MiB,0.00 MiB,6,4,0.006452 ms,96.374606 mW,621808.960000 pJ -pool/max_basic,PASS,0.063 s,0.00 MiB,0.00 MiB,3,0,0.001634 ms,92.132191 mW,150544.000000 pJ -pool/max_ceil_mode,PASS,0.078 s,0.00 MiB,0.00 MiB,2,0,0.001297 ms,79.111025 mW,102607.000000 pJ -pool/max_global_style_kernel_equals_input,PASS,0.067 s,0.00 MiB,0.00 MiB,1,0,0.004366 ms,78.010994 mW,340596.000000 pJ -pool/max_non_square_kernel,PASS,0.073 s,0.00 MiB,0.00 MiB,4,0,0.003409 ms,93.253447 mW,317901.000000 pJ -pool/max_real_asymmetric_padding,PASS,0.064 s,0.00 MiB,0.00 MiB,4,0,0.003078 ms,93.124756 mW,286638.000000 pJ -pool/max_same_upper,PASS,0.063 s,0.00 MiB,0.00 MiB,3,0,0.003024 ms,92.095238 mW,278496.000000 pJ -pool/max_stride2_multichannel,PASS,0.066 s,0.00 MiB,0.00 MiB,3,0,0.004012 ms,92.269192 mW,370184.000000 pJ -reduce_mean/4d_spatial,PASS,0.078 s,0.00 MiB,0.00 MiB,3,0,0.000321 ms,92.448598 mW,29676.000000 pJ -reduce_mean/4d_spatial_keepdims_0,PASS,0.066 s,0.00 MiB,0.00 MiB,4,0,0.000655 ms,94.352672 mW,61801.000000 pJ -reduce_mean/after_conv,PASS,0.068 s,0.00 MiB,0.00 MiB,5,3,0.005342 ms,106.951089 mW,571332.720000 pJ -reduce_mean/all_axes_keepdims_0,PASS,0.067 s,0.00 MiB,0.00 MiB,2,0,0.000391 ms,79.237852 mW,30982.000000 pJ -reduce_mean/all_axes_keepdims_1,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ -reduce_mean/basic,PASS,0.052 s,0.00 MiB,0.00 MiB,4,0,0.000373 ms,93.514745 mW,34881.000000 pJ -reduce_mean/channel_axis_nchw,PASS,0.054 s,0.03 MiB,0.02 MiB,4,0,0.164926 ms,93.596631 mW,15436518.000000 pJ -reduce_mean/keepdims_0,PASS,0.054 s,0.00 MiB,0.00 MiB,5,0,0.000748 ms,91.401070 mW,68368.000000 pJ -reduce_mean/large_dimension_1024,PASS,0.056 s,0.01 MiB,0.00 MiB,1,0,0.002785 ms,78.017235 mW,217278.000000 pJ -reduce_mean/legacy_axes_1_2_keepdims_1,PASS,0.064 s,0.00 MiB,0.00 MiB,2,0,0.000271 ms,79.354244 mW,21505.000000 pJ -reduce_mean/legacy_axis1_keepdims_0,PASS,0.071 s,0.00 MiB,0.00 MiB,9,0,0.001986 ms,92.501511 mW,183708.000000 pJ -reduce_mean/legacy_axis1_keepdims_1,PASS,0.065 s,0.00 MiB,0.00 MiB,8,0,0.001373 ms,94.559359 mW,129830.000000 pJ -reduce_mean/legacy_empty_axes_noop,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ -reduce_mean/legacy_nchw_spatial,PASS,0.050 s,0.00 MiB,0.00 MiB,3,0,0.000321 ms,92.448598 mW,29676.000000 pJ -reduce_mean/legacy_negative_axis,PASS,0.064 s,0.00 MiB,0.00 MiB,6,0,0.000553 ms,93.520796 mW,51717.000000 pJ -reduce_mean/legacy_reduce_all_keepdims_1,PASS,0.054 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ -reduce_mean/negative_axis,PASS,0.059 s,0.00 MiB,0.00 MiB,6,0,0.000553 ms,93.520796 mW,51717.000000 pJ -relu/4d,PASS,0.068 s,0.00 MiB,0.00 MiB,1,0,0.000521 ms,78.184261 mW,40734.000000 pJ -relu/after_conv,PASS,0.057 s,0.00 MiB,0.00 MiB,3,3,0.004956 ms,106.998935 mW,530286.720000 pJ -relu/after_gemm,PASS,0.064 s,0.01 MiB,0.01 MiB,5,4,0.007513 ms,105.158653 mW,790056.960000 pJ -relu/basic,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ -reshape/4d_to_2d_flatten,PASS,0.068 s,0.00 MiB,0.00 MiB,1,0,0.000258 ms,78.279070 mW,20196.000000 pJ -reshape/infer_dim_minus_one,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,0.000162 ms,78.296296 mW,12684.000000 pJ -reshape/same_rank,PASS,0.065 s,0.00 MiB,0.00 MiB,1,0,0.000162 ms,78.296296 mW,12684.000000 pJ -reshape/zero_copies_input_dim,PASS,0.057 s,0.00 MiB,0.00 MiB,1,0,0.000162 ms,78.296296 mW,12684.000000 pJ -resize/height_only,PASS,0.063 s,0.00 MiB,0.00 MiB,4,0,0.000693 ms,93.554113 mW,64833.000000 pJ -resize/nearest_2x,PASS,0.064 s,0.00 MiB,0.00 MiB,4,0,0.001173 ms,93.572890 mW,109761.000000 pJ -resize/nearest_downsample,PASS,0.062 s,0.00 MiB,0.00 MiB,2,0,0.000427 ms,79.449649 mW,33925.000000 pJ -resize/non_uniform,PASS,0.068 s,0.00 MiB,0.00 MiB,6,0,0.001753 ms,93.575014 mW,164037.000000 pJ -resize/width_only,PASS,0.070 s,0.00 MiB,0.00 MiB,2,0,0.000667 ms,79.503748 mW,53029.000000 pJ -resize/with_sizes,PASS,0.052 s,0.00 MiB,0.00 MiB,3,0,0.000797 ms,92.542033 mW,73756.000000 pJ -sigmoid/4d,PASS,0.067 s,0.00 MiB,0.00 MiB,1,0,0.000521 ms,78.184261 mW,40734.000000 pJ -sigmoid/after_gemm,PASS,0.070 s,0.01 MiB,0.01 MiB,5,4,0.007513 ms,105.158653 mW,790056.960000 pJ -sigmoid/basic,PASS,0.062 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ -slice/2d_basic,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,0.000242 ms,78.297521 mW,18948.000000 pJ -slice/after_conv,PASS,0.062 s,0.00 MiB,0.01 MiB,7,6,0.011296 ms,118.190765 mW,1335082.880000 pJ -slice/default_axes,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,0.000242 ms,78.297521 mW,18948.000000 pJ -slice/large_channel_1024,PASS,0.049 s,0.01 MiB,0.00 MiB,1,0,0.002832 ms,78.144068 mW,221304.000000 pJ -slice/nchw_spatial_crop,PASS,0.063 s,0.00 MiB,0.00 MiB,1,0,0.001302 ms,78.239631 mW,101868.000000 pJ -slice/negative_axis,PASS,0.058 s,0.00 MiB,0.00 MiB,1,0,0.000562 ms,78.298932 mW,44004.000000 pJ -slice/negative_indices,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,0.000322 ms,78.298137 mW,25212.000000 pJ -slice/step2,PASS,0.054 s,0.00 MiB,0.00 MiB,1,0,0.002042 ms,78.293830 mW,159876.000000 pJ -softmax/3d_last_axis,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED -softmax/basic,PASS,0.069 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED -softmax/channel_axis,PASS,0.061 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED -softmax/large_dimension_1024,PASS,0.045 s,0.01 MiB,0.01 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED -softmax/negative_axis,PASS,0.064 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED -split/basic,PASS,0.066 s,0.00 MiB,0.00 MiB,1,0,0.000403 ms,78.297767 mW,31554.000000 pJ -split/equal_three_way,PASS,0.065 s,0.00 MiB,0.00 MiB,1,0,0.000564 ms,78.297872 mW,44160.000000 pJ -split/negative_axis,PASS,0.058 s,0.00 MiB,0.00 MiB,1,0,0.001083 ms,78.288089 mW,84786.000000 pJ -split/uneven_channel_axis_4d,PASS,0.062 s,0.00 MiB,0.00 MiB,1,0,0.000242 ms,78.297521 mW,18948.000000 pJ -sub/after_gemm,PASS,0.068 s,0.01 MiB,0.01 MiB,5,4,0.007784 ms,104.703618 mW,815012.960000 pJ -sub/basic,PASS,0.060 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -sub/broadcast_row,PASS,0.062 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ -sub/channel_broadcast_1024,PASS,0.057 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ -sub/constant_lhs_broadcast,PASS,0.057 s,0.00 MiB,0.00 MiB,1,0,0.000322 ms,78.223602 mW,25188.000000 pJ -sub/leading_dimension_broadcast,PASS,0.058 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ +add/after_gemm,PASS,0.062 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +add/basic,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +add/broadcast_row,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +add/channel_broadcast_1024,PASS,0.061 s,0.02 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +add/leading_dimension_broadcast,PASS,0.057 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +concat/channel_axis,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +concat/negative_axis,PASS,0.061 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +concat/three_inputs_channel_axis,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +conv/batch_2,PASS,0.056 s,0.00 MiB,0.00 MiB,2,2,SKIP,SKIP,SKIP +conv/batch_4_pointwise,PASS,0.058 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +conv/depthwise_1024_channels,PASS,0.085 s,0.19 MiB,0.38 MiB,129,128,SKIP,SKIP,SKIP +conv/depthwise_grouped,PASS,0.074 s,0.01 MiB,0.00 MiB,5,4,SKIP,SKIP,SKIP +conv/dilated_3x3,PASS,0.067 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/dynamic,PASS,0.059 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +conv/explicit_padding,PASS,0.065 s,0.00 MiB,0.00 MiB,4,4,SKIP,SKIP,SKIP +conv/grouped_many_groups,PASS,0.500 s,0.05 MiB,0.09 MiB,65,64,SKIP,SKIP,SKIP +conv/grouped_two_groups,PASS,0.070 s,0.00 MiB,0.00 MiB,3,2,SKIP,SKIP,SKIP +conv/huge_pointwise_1024,PASS,0.652 s,0.01 MiB,0.01 MiB,1,64,SKIP,SKIP,SKIP +conv/huge_pointwise_1024_dynamic,PASS,0.083 s,8.04 MiB,12.61 MiB,168,0,SKIP,SKIP,SKIP +conv/kernel_3x3,PASS,0.067 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/kernel_equals_input_spatial,PASS,0.072 s,0.00 MiB,0.00 MiB,1,2,SKIP,SKIP,SKIP +conv/large_input_channels_1x1,PASS,0.098 s,0.01 MiB,0.01 MiB,1,8,SKIP,SKIP,SKIP +conv/large_output_channels_1x1,PASS,0.089 s,0.00 MiB,0.01 MiB,1,8,SKIP,SKIP,SKIP +conv/large_spatial,PASS,0.063 s,0.00 MiB,0.01 MiB,6,6,SKIP,SKIP,SKIP +conv/multi_channel,PASS,0.068 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/non_square_kernel_1x3,PASS,0.080 s,0.00 MiB,0.00 MiB,5,5,SKIP,SKIP,SKIP +conv/non_square_kernel_3x1,PASS,0.068 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/non_uniform_stride,PASS,0.068 s,0.00 MiB,0.00 MiB,4,4,SKIP,SKIP,SKIP +conv/pointwise_1x1,PASS,0.071 s,0.00 MiB,0.00 MiB,4,4,SKIP,SKIP,SKIP +conv/pointwise_tiled_chain,PASS,0.911 s,0.01 MiB,0.02 MiB,2,80,SKIP,SKIP,SKIP +conv/real_asymmetric_padding,PASS,0.057 s,0.00 MiB,0.00 MiB,4,4,SKIP,SKIP,SKIP +conv/relu_conv_store,PASS,0.070 s,0.02 MiB,0.10 MiB,32,32,SKIP,SKIP,SKIP +conv/same_lower_3x3,PASS,0.061 s,0.00 MiB,0.00 MiB,5,5,SKIP,SKIP,SKIP +conv/same_padding_3x3,PASS,0.075 s,0.00 MiB,0.00 MiB,5,5,SKIP,SKIP,SKIP +conv/simple,PASS,0.055 s,0.00 MiB,0.00 MiB,2,2,SKIP,SKIP,SKIP +conv/stride_2,PASS,0.057 s,0.00 MiB,0.00 MiB,2,2,SKIP,SKIP,SKIP +conv/with_bias_3x3,PASS,0.067 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/with_constant,PASS,0.063 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/without_kernel_shape_attr,PASS,0.059 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +conv/yolo11n_depthwise_head,PASS,0.624 s,4.82 MiB,15.92 MiB,160,720,SKIP,SKIP,SKIP +conv/yolo11n_heavy,PASS,0.496 s,4.82 MiB,20.66 MiB,160,800,SKIP,SKIP,SKIP +conv/yolo11n_stem,PASS,0.935 s,12.86 MiB,31.38 MiB,168,488,SKIP,SKIP,SKIP +div/after_gemm,PASS,0.060 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +div/basic,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +div/channel_broadcast_1024,PASS,0.049 s,0.02 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +div/leading_dimension_broadcast,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +div/runtime_scalar_rhs,PASS,0.053 s,0.02 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +div/scalar_constant,PASS,0.078 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gather/3d_input_axis1,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gather/axis0_matrix_indices,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gather/axis1,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gather/negative_axis,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gather/negative_indices,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +gemm/alpha_beta,PASS,0.078 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/bias_rank2_broadcast,PASS,0.055 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/dynamic,PASS,0.058 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/dynamic_alpha,PASS,0.056 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/dynamic_beta,PASS,0.053 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/dynamic_bias,PASS,0.062 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/dynamic_bias_alpha_beta,PASS,0.114 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/dynamic_transB,PASS,0.055 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +gemm/huge_1024,PASS,0.150 s,0.01 MiB,0.10 MiB,73,64,SKIP,SKIP,SKIP +gemm/large,PASS,0.059 s,0.02 MiB,0.03 MiB,17,16,SKIP,SKIP,SKIP +gemm/large_k_small_n,PASS,0.087 s,0.01 MiB,0.01 MiB,9,8,SKIP,SKIP,SKIP +gemm/non_square,PASS,0.058 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/scalar_bias,PASS,0.058 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/simple,PASS,0.075 s,0.03 MiB,0.08 MiB,42,40,SKIP,SKIP,SKIP +gemm/small,PASS,0.056 s,0.00 MiB,0.00 MiB,2,2,SKIP,SKIP,SKIP +gemm/small_k_large_n,PASS,0.091 s,0.01 MiB,0.02 MiB,17,8,SKIP,SKIP,SKIP +gemm/transA,PASS,0.057 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/transA_transB,PASS,0.079 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/transB,PASS,0.065 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/transB_with_bias,PASS,0.058 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemm/with_bias,PASS,0.054 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +gemv/constant,PASS,0.049 s,0.00 MiB,0.00 MiB,0,0,SKIP,SKIP,SKIP +gemv/simple,PASS,0.066 s,0.00 MiB,0.01 MiB,6,4,SKIP,SKIP,SKIP +gemv/with_heterogeneous_constant,PASS,0.069 s,0.00 MiB,0.01 MiB,6,4,SKIP,SKIP,SKIP +gemv/with_homogeneous_constant,PASS,0.065 s,0.00 MiB,0.01 MiB,6,4,SKIP,SKIP,SKIP +gemv/with_scalar_constant,PASS,0.100 s,0.00 MiB,0.01 MiB,6,4,SKIP,SKIP,SKIP +matmul/basic,PASS,0.056 s,0.00 MiB,0.00 MiB,2,2,SKIP,SKIP,SKIP +matmul/batched_3d,PASS,0.057 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +matmul/batched_3d_dynamic,PASS,0.053 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +matmul/batched_left_constant,PASS,0.062 s,0.00 MiB,0.02 MiB,9,8,SKIP,SKIP,SKIP +matmul/batched_lhs_broadcast,PASS,0.059 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +matmul/batched_rhs_broadcast,PASS,0.087 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +matmul/dynamic,PASS,0.058 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +matmul/huge_1024,PASS,0.164 s,0.01 MiB,0.10 MiB,73,64,SKIP,SKIP,SKIP +matmul/left_constant,PASS,0.058 s,0.00 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +matmul/matrix_vector,PASS,0.100 s,0.52 MiB,0.78 MiB,168,173,SKIP,SKIP,SKIP +matmul/vector_matrix,PASS,0.148 s,0.01 MiB,0.01 MiB,9,8,SKIP,SKIP,SKIP +matmul/yolo_attention,PASS,0.466 s,1.02 MiB,43.44 MiB,168,0,SKIP,SKIP,SKIP +mul/after_conv,PASS,0.059 s,0.00 MiB,0.00 MiB,4,3,SKIP,SKIP,SKIP +mul/basic,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +mul/channel_broadcast_1024,PASS,0.051 s,0.02 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +mul/leading_dimension_broadcast,PASS,0.071 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +mul/scalar_constant,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_basic,PASS,0.055 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_ceil_mode,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_explicit_padding,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_include_pad,PASS,0.062 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_large_channels,PASS,0.078 s,0.04 MiB,0.02 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_non_uniform_stride,PASS,0.057 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/avg_real_asymmetric_padding,PASS,0.061 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/max_after_conv,PASS,0.061 s,0.00 MiB,0.00 MiB,6,4,SKIP,SKIP,SKIP +pool/max_basic,PASS,0.060 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +pool/max_ceil_mode,PASS,0.055 s,0.00 MiB,0.00 MiB,2,0,SKIP,SKIP,SKIP +pool/max_global_style_kernel_equals_input,PASS,0.100 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +pool/max_non_square_kernel,PASS,0.062 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +pool/max_real_asymmetric_padding,PASS,0.062 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +pool/max_same_upper,PASS,0.059 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +pool/max_stride2_multichannel,PASS,0.056 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +reduce_mean/4d_spatial,PASS,0.053 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +reduce_mean/4d_spatial_keepdims_0,PASS,0.070 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +reduce_mean/after_conv,PASS,0.063 s,0.00 MiB,0.00 MiB,5,3,SKIP,SKIP,SKIP +reduce_mean/all_axes_keepdims_0,PASS,0.053 s,0.00 MiB,0.00 MiB,2,0,SKIP,SKIP,SKIP +reduce_mean/all_axes_keepdims_1,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reduce_mean/basic,PASS,0.052 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +reduce_mean/channel_axis_nchw,PASS,0.053 s,0.03 MiB,0.02 MiB,4,0,SKIP,SKIP,SKIP +reduce_mean/keepdims_0,PASS,0.053 s,0.00 MiB,0.00 MiB,5,0,SKIP,SKIP,SKIP +reduce_mean/large_dimension_1024,PASS,0.061 s,0.01 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reduce_mean/legacy_axes_1_2_keepdims_1,PASS,0.052 s,0.00 MiB,0.00 MiB,2,0,SKIP,SKIP,SKIP +reduce_mean/legacy_axis1_keepdims_0,PASS,0.052 s,0.00 MiB,0.00 MiB,9,0,SKIP,SKIP,SKIP +reduce_mean/legacy_axis1_keepdims_1,PASS,0.058 s,0.00 MiB,0.00 MiB,8,0,SKIP,SKIP,SKIP +reduce_mean/legacy_empty_axes_noop,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reduce_mean/legacy_nchw_spatial,PASS,0.058 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +reduce_mean/legacy_negative_axis,PASS,0.059 s,0.00 MiB,0.00 MiB,6,0,SKIP,SKIP,SKIP +reduce_mean/legacy_reduce_all_keepdims_1,PASS,0.068 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reduce_mean/negative_axis,PASS,0.065 s,0.00 MiB,0.00 MiB,6,0,SKIP,SKIP,SKIP +relu/4d,PASS,0.066 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +relu/after_conv,PASS,0.075 s,0.00 MiB,0.00 MiB,3,3,SKIP,SKIP,SKIP +relu/after_gemm,PASS,0.079 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +relu/basic,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reshape/4d_to_2d_flatten,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reshape/infer_dim_minus_one,PASS,0.061 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reshape/same_rank,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +reshape/zero_copies_input_dim,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +resize/height_only,PASS,0.079 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +resize/nearest_2x,PASS,0.053 s,0.00 MiB,0.00 MiB,4,0,SKIP,SKIP,SKIP +resize/nearest_downsample,PASS,0.055 s,0.00 MiB,0.00 MiB,2,0,SKIP,SKIP,SKIP +resize/non_uniform,PASS,0.053 s,0.00 MiB,0.00 MiB,6,0,SKIP,SKIP,SKIP +resize/width_only,PASS,0.052 s,0.00 MiB,0.00 MiB,2,0,SKIP,SKIP,SKIP +resize/with_sizes,PASS,0.053 s,0.00 MiB,0.00 MiB,3,0,SKIP,SKIP,SKIP +sigmoid/4d,PASS,0.077 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +sigmoid/after_gemm,PASS,0.058 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +sigmoid/basic,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/2d_basic,PASS,0.055 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/after_conv,PASS,0.061 s,0.00 MiB,0.01 MiB,7,6,SKIP,SKIP,SKIP +slice/default_axes,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/large_channel_1024,PASS,0.068 s,0.01 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/nchw_spatial_crop,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/negative_axis,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/negative_indices,PASS,0.048 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +slice/step2,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +softmax/3d_last_axis,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +softmax/basic,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +softmax/channel_axis,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +softmax/large_dimension_1024,PASS,0.048 s,0.01 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +softmax/negative_axis,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +split/basic,PASS,0.050 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +split/equal_three_way,PASS,0.049 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +split/negative_axis,PASS,0.067 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +split/uneven_channel_axis_4d,PASS,0.053 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +sub/after_gemm,PASS,0.059 s,0.01 MiB,0.01 MiB,5,4,SKIP,SKIP,SKIP +sub/basic,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +sub/broadcast_row,PASS,0.051 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +sub/channel_broadcast_1024,PASS,0.056 s,0.02 MiB,0.01 MiB,1,0,SKIP,SKIP,SKIP +sub/constant_lhs_broadcast,PASS,0.052 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP +sub/leading_dimension_broadcast,PASS,0.066 s,0.00 MiB,0.00 MiB,1,0,SKIP,SKIP,SKIP