Compare commits

2 Commits

Author SHA1 Message Date
NiccoloN 942a9faa4f second temp commit: i will soft-reset and recommit after next changes
Validate Operations / validate-operations (push) Has been cancelled
2026-08-03 11:07:28 +02:00
NiccoloN 893e90feac temp commit: i will soft-reset and recommit after next changes 2026-08-02 11:37:22 +02:00
81 changed files with 5523 additions and 3646 deletions
+20 -8
View File
@@ -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`).
@@ -118,8 +132,6 @@ options; `onnx-mlir --help` lists the inherited ONNX-MLIR options.
elements per convolution before streaming. Default is `1048576`.
- `--pim-conv-stream-chunk-positions=<N>` - maximum output positions per
streamed convolution chunk. Default is `1024`.
- `--pim-report-conv-lowering=<true|false>` - emit the bounded convolution
lowering report. Default is `true`.
- `--use-experimental-conv-impl` - use the alternate convolution lowering.
- `--pim-detect-communication-deadlock` - statically simulate expanded
send/receive ordering and reject blocking deadlocks. Default is off.
+62 -14
View File
@@ -1,10 +1,13 @@
#include "mlir/Dialect/Affine/IR/AffineOps.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Bufferization/IR/Bufferization.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/IR/BuiltinAttributes.h"
#include "mlir/Interfaces/DestinationStyleOpInterface.h"
#include "llvm/ADT/SmallPtrSet.h"
#include <limits>
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
@@ -36,6 +39,10 @@ mlir::Value resolveAlias(mlir::Value value, const StaticValueKnowledge* knowledg
llvm::FailureOr<CompiledIndexExpr> compileIndexValueImpl(mlir::Value value);
llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExprImpl(mlir::Value value);
using AliasResolutionSet = llvm::SmallPtrSet<mlir::Value, 8>;
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value,
const StaticValueKnowledge* knowledge,
AliasResolutionSet& visited);
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge);
template <typename... Args>
@@ -45,18 +52,23 @@ CompiledIndexExpr makeCompiledIndexExpr(Args&&... args) {
static mlir::Value resolveForYieldedAliasToInit(mlir::scf::ForOp forOp,
mlir::Value yieldedValue,
const StaticValueKnowledge* knowledge) {
yieldedValue = resolveLoopCarriedAliasImpl(yieldedValue, knowledge);
const StaticValueKnowledge* knowledge,
AliasResolutionSet& visited) {
yieldedValue = resolveLoopCarriedAliasImpl(yieldedValue, knowledge, visited);
if (auto blockArgument = mlir::dyn_cast<mlir::BlockArgument>(yieldedValue)) {
if (blockArgument.getOwner() == forOp.getBody() && blockArgument.getArgNumber() > 0
&& static_cast<unsigned>(blockArgument.getArgNumber() - 1) < forOp.getInitArgs().size())
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[blockArgument.getArgNumber() - 1], knowledge);
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[blockArgument.getArgNumber() - 1], knowledge, visited);
}
return yieldedValue;
}
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge) {
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value,
const StaticValueKnowledge* knowledge,
AliasResolutionSet& visited) {
value = resolveAlias(value, knowledge);
if (!value || !visited.insert(value).second)
return value;
if (auto blockArgument = mlir::dyn_cast<mlir::BlockArgument>(value)) {
auto forOp = mlir::dyn_cast_or_null<mlir::scf::ForOp>(blockArgument.getOwner()->getParentOp());
@@ -64,9 +76,12 @@ mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnow
const unsigned iterArgIndex = blockArgument.getArgNumber() - 1;
auto yieldOp = mlir::dyn_cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
if (iterArgIndex < forOp.getInitArgs().size() && yieldOp
&& iterArgIndex < yieldOp.getNumOperands()
&& resolveAlias(yieldOp.getOperand(iterArgIndex), knowledge) == blockArgument)
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[iterArgIndex], knowledge);
&& iterArgIndex < yieldOp.getNumOperands()) {
mlir::Value yieldedValue = resolveAlias(yieldOp.getOperand(iterArgIndex), knowledge);
if (yieldedValue == blockArgument
|| (yieldedValue && resolveLoopCarriedAliasImpl(yieldedValue, knowledge, visited) == blockArgument))
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[iterArgIndex], knowledge, visited);
}
}
return value;
}
@@ -75,10 +90,15 @@ mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnow
if (!definingOp)
return value;
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp))
return resolveLoopCarriedAliasImpl(toBufferOp.getTensor(), knowledge, visited);
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp))
return resolveLoopCarriedAliasImpl(toTensorOp.getBuffer(), knowledge, visited);
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
if (auto result = mlir::dyn_cast<mlir::OpResult>(value))
if (mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(result))
return resolveLoopCarriedAliasImpl(tiedOperand->get(), knowledge);
return resolveLoopCarriedAliasImpl(tiedOperand->get(), knowledge, visited);
}
if (auto forOp = mlir::dyn_cast<mlir::scf::ForOp>(definingOp)) {
@@ -86,20 +106,26 @@ mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnow
if (result) {
auto yieldOp = mlir::dyn_cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
if (yieldOp && result.getResultNumber() < yieldOp.getNumOperands())
return resolveForYieldedAliasToInit(forOp, yieldOp.getOperand(result.getResultNumber()), knowledge);
return resolveForYieldedAliasToInit(
forOp, yieldOp.getOperand(result.getResultNumber()), knowledge, visited);
}
}
if (auto castOp = mlir::dyn_cast<mlir::memref::CastOp>(definingOp))
return resolveLoopCarriedAliasImpl(castOp.getSource(), knowledge);
return resolveLoopCarriedAliasImpl(castOp.getSource(), knowledge, visited);
if (auto collapseOp = mlir::dyn_cast<mlir::memref::CollapseShapeOp>(definingOp))
return resolveLoopCarriedAliasImpl(collapseOp.getSrc(), knowledge);
return resolveLoopCarriedAliasImpl(collapseOp.getSrc(), knowledge, visited);
if (auto expandOp = mlir::dyn_cast<mlir::memref::ExpandShapeOp>(definingOp))
return resolveLoopCarriedAliasImpl(expandOp.getSrc(), knowledge);
return resolveLoopCarriedAliasImpl(expandOp.getSrc(), knowledge, visited);
return value;
}
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge) {
AliasResolutionSet visited;
return resolveLoopCarriedAliasImpl(value, knowledge, visited);
}
llvm::FailureOr<int64_t> resolveOpFoldResult(mlir::OpFoldResult ofr, const StaticValueKnowledge* knowledge);
llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticValueKnowledge* knowledge);
@@ -524,6 +550,15 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
if (!definingOp)
return mlir::failure();
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp)) {
value = resolveAlias(toBufferOp.getTensor(), knowledge);
continue;
}
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp)) {
value = resolveAlias(toTensorOp.getBuffer(), knowledge);
continue;
}
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(mlir::dyn_cast<mlir::OpResult>(value));
if (!tiedOperand)
@@ -538,7 +573,9 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
return mlir::failure();
auto yieldOp = mlir::cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
value = resolveForYieldedAliasToInit(forOp, yieldOp.getOperand(result.getResultNumber()), knowledge);
AliasResolutionSet visited;
value = resolveForYieldedAliasToInit(
forOp, yieldOp.getOperand(result.getResultNumber()), knowledge, visited);
continue;
}
@@ -643,6 +680,15 @@ llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExprImpl(mlir::Valu
if (!definingOp)
return mlir::failure();
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp)) {
value = toBufferOp.getTensor();
continue;
}
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp)) {
value = toTensorOp.getBuffer();
continue;
}
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(mlir::dyn_cast<mlir::OpResult>(value));
if (!tiedOperand)
@@ -657,7 +703,9 @@ llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExprImpl(mlir::Valu
return mlir::failure();
auto yieldOp = mlir::cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
value = resolveForYieldedAliasToInit(forOp, yieldOp.getOperand(result.getResultNumber()), nullptr);
AliasResolutionSet visited;
value = resolveForYieldedAliasToInit(
forOp, yieldOp.getOperand(result.getResultNumber()), nullptr, visited);
continue;
}
-5
View File
@@ -87,11 +87,6 @@ llvm::cl::opt<uint64_t> pimConvStreamChunkPositions(
llvm::cl::init(1024),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimReportConvLowering("pim-report-conv-lowering",
llvm::cl::desc("Emit a bounded Conv lowering report"),
llvm::cl::init(true),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimEmitJson("pim-emit-json",
llvm::cl::desc("Also emit per-core JSON instruction files alongside binary .pim files"),
llvm::cl::init(false),
-1
View File
@@ -57,7 +57,6 @@ extern llvm::cl::opt<PimSpatialDataflowExportType> pimExportSpatialDataflow;
extern llvm::cl::opt<bool> pimOnlyCodegen;
extern llvm::cl::opt<bool> useExperimentalConvImpl;
extern llvm::cl::opt<bool> pimEmitJson;
extern llvm::cl::opt<bool> pimReportConvLowering;
extern llvm::cl::opt<bool> pimDetectCommunicationDeadlock;
extern llvm::cl::opt<bool> pimMaterializeScalarFanoutGlobalOrder;
extern llvm::cl::opt<bool> pimTraceCommunicationMaterialization;
+43 -5
View File
@@ -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<ModuleOp>& 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<spatial::ScheduledSpatialState>();
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<ModuleOp>& module,
}
if (pimEmissionTarget >= EmitPimBufferized) {
pm.addPass(createPimBufferizationPass());
pm.addPass(createPimBufferizationPreparationPass());
pm.addPass(createPimOneShotBufferizationPass());
pm.addPass(createPimMemoryNormalizationPass());
pm.addPass(createPimBufferizationVerificationPass());
pm.addPass(createMessagePass("Pim bufferized"));
}
@@ -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
@@ -25,6 +25,9 @@ FailureOr<Value> 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<int64_t> operandIndices(entries.size(), 0), sourceSlots, sourceOffsets, offsets, sizes,
strides(entries.size() * rank, 1);
@@ -47,13 +50,18 @@ FailureOr<Value> createFragmentAssemblyBlueprint(Value physicalBatch,
llvm::append_range(offsets, entry.destinationOffsets);
llvm::append_range(sizes, entry.sizes);
}
return spatial::SpatBlueprintOp::create(rewriter, loc, logicalType, physicalBatch, ValueRange {},
rewriter.getStringAttr("nchw"), rewriter.getStringAttr(physicalLayout),
auto blueprint = spatial::SpatBlueprintOp::create(rewriter, loc, logicalType, physicalBatch, ValueRange {},
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")).getOutput();
rewriter.getStringAttr("disjoint"), rewriter.getStringAttr("complete"));
if (indexMap == spatial::kContiguousRowMajorFragments
&& !spatial::isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
blueprint.setIndexMapAttr(rewriter.getStringAttr("fragment_assembly"));
return blueprint.getOutput();
}
Value sumTensors(ArrayRef<Value> tensors, PatternRewriter& rewriter) {
@@ -394,6 +394,39 @@ extractGraphBatchPhysicalFragment(mlir::PatternRewriter& rewriter,
rewriter, loc, physicalBatch, fragmentType, {offsets, sizes, strides});
}
template <typename BodyFn>
mlir::FailureOr<mlir::Value> mapGraphBatchFragments(mlir::Value input,
mlir::RankedTensorType outputType,
mlir::PatternRewriter& rewriter,
mlir::Location loc,
BodyFn&& build) {
auto inputType = mlir::dyn_cast<mlir::RankedTensorType>(input.getType());
if (!inputType || !inputType.hasStaticShape() || !outputType.hasStaticShape()
|| inputType.getRank() != outputType.getRank() || inputType.getRank() < 2
|| inputType.getDimSize(0) != outputType.getDimSize(0))
return mlir::failure();
auto inputFragmentType = mlir::RankedTensorType::get(
inputType.getShape().drop_front(), inputType.getElementType(), inputType.getEncoding());
auto outputFragmentType = mlir::RankedTensorType::get(
outputType.getShape().drop_front(), outputType.getElementType(), outputType.getEncoding());
auto batch = createSpatComputeBatch(
rewriter, loc, mlir::TypeRange {outputType}, inputType.getDimSize(0), {}, mlir::ValueRange {input},
[&](detail::SpatComputeBatchBodyArgs args) -> mlir::LogicalResult {
auto fragment = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs.front(), args.lane, inputFragmentType);
if (mlir::failed(fragment))
return mlir::failure();
mlir::FailureOr<mlir::Value> result = build(*fragment, outputFragmentType);
if (mlir::failed(result) || result->getType() != outputFragmentType)
return mlir::failure();
publishGraphBatchPhysicalFragment(rewriter, loc, *result, args.outputs.front(), args.lane);
return mlir::success();
});
if (mlir::failed(batch))
return mlir::failure();
return batch->getResult(0);
}
template <typename BodyFn>
mlir::Value materializeOrComputeUnary(mlir::Value input,
mlir::RankedTensorType resultType,
@@ -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<mlir::Value> materializeTransposedContractionConstant(
mlir::Value input,
mlir::RankedTensorType resultType,
llvm::ArrayRef<int64_t> permutation,
mlir::PatternRewriter& rewriter,
mlir::Location loc) {
auto denseAttr = getHostConstDenseElementsAttr(input);
auto inputType = denseAttr ? mlir::dyn_cast<mlir::RankedTensorType>(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
@@ -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<mlir::Value> materializeTransposedContractionConstant(
mlir::Value input,
mlir::RankedTensorType resultType,
llvm::ArrayRef<int64_t> permutation,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
} // namespace onnx_mlir
@@ -0,0 +1,68 @@
#include "ContractionPlanning.hpp"
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
#include <algorithm>
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<int64_t> buildBatchMap(
llvm::ArrayRef<int64_t> sourceShape,
llvm::ArrayRef<int64_t> outputShape) {
llvm::SmallVector<int64_t> map(outputShape.size(), -1);
const int64_t offset = outputShape.size() - sourceShape.size();
for (int64_t source = 0; source < static_cast<int64_t>(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<int64_t>(1, target.matrixShape.rows);
plan.tileK = std::max<int64_t>(1, target.matrixShape.rows);
plan.tileN = std::max<int64_t>(1, target.matrixShape.columns);
plan.reductionSlices = std::max<int64_t>(1, ceilDivide(problem.k, plan.tileK));
plan.outputTiles = std::max<int64_t>(1, ceilDivide(problem.n, plan.tileN));
plan.rowTiles = std::max<int64_t>(1, ceilDivide(problem.m, plan.tileM));
plan.fragmentRows = std::max<int64_t>(
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
@@ -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<int64_t> lhsBatchMap;
llvm::SmallVector<int64_t> 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
@@ -0,0 +1,35 @@
#pragma once
#include "mlir/IR/BuiltinTypes.h"
#include "llvm/ADT/SmallVector.h"
#include <cstdint>
namespace onnx_mlir {
enum class ContractionOrigin { Gemm, MatMul };
struct ContractionProblem {
llvm::SmallVector<int64_t> lhsBatchShape;
llvm::SmallVector<int64_t> rhsBatchShape;
llvm::SmallVector<int64_t> 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
@@ -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<int64_t> 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<int64_t> 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<int64_t> 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<RankedTensorType>(value.getType());
SmallVector<OpFoldResult> lowPads(sourceType.getRank(), rewriter.getIndexAttr(0));
@@ -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<int64_t> permutation,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
mlir::Value createZeroPaddedTensor(mlir::Value value,
mlir::RankedTensorType resultType,
mlir::PatternRewriter& rewriter,
@@ -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 <numeric>
@@ -33,6 +33,16 @@ FailureOr<RowStripPhysicalValue> describeRowStripPhysicalValue(Value storage, Ra
tilesPerRow};
}
FailureOr<RowStripPhysicalValue> getRowStripPhysicalValue(Value value) {
auto blueprint = value.getDefiningOp<spatial::SpatBlueprintOp>();
auto logicalType = dyn_cast<RankedTensorType>(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<Value> createRowStripStorageFromRows(Value rows,
return batchOp->getResult(0);
}
FailureOr<Value> createRowStripStorageBlueprint(Value storage,
RankedTensorType logicalType,
PatternRewriter& rewriter,
Location loc) {
FailureOr<RowStripPhysicalValue> 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<Value> createRowStripAssemblyBlueprint(const RowStripPhysicalValue& value,
PatternRewriter& rewriter,
Location loc) {
@@ -160,8 +199,8 @@ FailureOr<Value> 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<Value> 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);
}
@@ -186,25 +225,9 @@ static FailureOr<Value> applyRowStripActivation(const RowStripPhysicalValue& val
Location loc,
BuildActivation buildActivation) {
auto storageType = cast<RankedTensorType>(value.storage.getType());
const int64_t laneCount = storageType.getDimSize(0);
auto batchOp = createSpatComputeBatch(rewriter,
loc,
TypeRange {storageType},
laneCount,
{},
ValueRange {value.storage},
[&](detail::SpatComputeBatchBodyArgs args) {
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs.front(), args.lane, value.fragmentType);
if (failed(fragment)) return failure();
Value result = buildActivation(*fragment);
publishGraphBatchPhysicalFragment(
rewriter, loc, result, args.outputs.front(), args.lane);
return success();
});
if (failed(batchOp))
return failure();
return batchOp->getResult(0);
return mapGraphBatchFragments(value.storage, storageType, rewriter, loc, [&](Value fragment, RankedTensorType) {
return FailureOr<Value>(buildActivation(fragment));
});
}
FailureOr<Value> applyRowStripRelu(const RowStripPhysicalValue& value, PatternRewriter& rewriter, Location loc) {
@@ -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<RowStripPhysicalValue> describeRowStripPhysicalValue(mlir::Value storage,
mlir::RankedTensorType logicalType);
mlir::FailureOr<RowStripPhysicalValue> getRowStripPhysicalValue(mlir::Value value);
std::pair<llvm::SmallVector<int64_t>, llvm::SmallVector<int64_t>>
buildRowStripMetadata(mlir::RankedTensorType type);
@@ -53,6 +61,11 @@ mlir::FailureOr<mlir::Value> createRowStripStorageFromRows(mlir::Value rows,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
mlir::FailureOr<mlir::Value> createRowStripStorageBlueprint(mlir::Value storage,
mlir::RankedTensorType logicalType,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
mlir::FailureOr<mlir::Value> createRowStripAssemblyBlueprint(const RowStripPhysicalValue& value,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
@@ -80,4 +93,14 @@ mlir::FailureOr<mlir::Value> applyRowStripConcat(llvm::ArrayRef<RowStripPhysical
mlir::PatternRewriter& rewriter,
mlir::Location loc);
mlir::LogicalResult canLowerFlattenFromRowStrip(
spatial::SpatGraphCompute flattenOp,
const spatial::SpatialTargetInfo& target);
mlir::LogicalResult lowerFlattenFromRowStrip(
const RowStripPhysicalValue& input,
spatial::SpatGraphCompute flattenOp,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
} // namespace onnx_mlir
@@ -5,7 +5,6 @@
#include "ShapeTilingUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
using namespace mlir;
@@ -67,11 +66,15 @@ sliceVector(const Value& vectorToSlice, int64_t sliceSize, PatternRewriter& rewr
}
DenseMap<CoreId, SmallVector<Value>>
sliceVectorPerCrossbarPerCore(const Value& vectorToSlice, PatternRewriter& rewriter, Location loc) {
SmallVector<Value> slices = sliceVector(vectorToSlice, crossbarSize, rewriter, loc);
sliceVectorPerCrossbarPerCore(const Value& vectorToSlice,
PatternRewriter& rewriter,
Location loc,
const spatial::SpatialTargetInfo& target) {
SmallVector<Value> slices = sliceVector(
vectorToSlice, static_cast<int64_t>(target.matrixShape.rows), rewriter, loc);
DenseMap<CoreId, SmallVector<Value>> 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;
@@ -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<mlir::Value> sliceVector(const mlir::Value& vectorToSlice,
/// Partitions one logical vector into per-core crossbar-sized slices using the
/// current PIM target geometry.
llvm::DenseMap<CoreId, llvm::SmallVector<mlir::Value>> 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
File diff suppressed because it is too large Load Diff
@@ -34,9 +34,15 @@ struct ONNXToSpatialPass : PassWrapper<ONNXToSpatialPass, OperationPass<ModuleOp
StringRef getDescription() const override { return "Lower ONNX ops to Spatial ops."; }
ONNXToSpatialPass() = default;
ONNXToSpatialPass(const ONNXToSpatialPass& pass) {}
explicit ONNXToSpatialPass(const spatial::SpatialTargetInfo& target)
: target(target), hasTarget(true) {}
ONNXToSpatialPass(const ONNXToSpatialPass& pass)
: target(pass.target), hasTarget(pass.hasTarget) {}
void runOnOperation() override;
spatial::SpatialTargetInfo target;
bool hasTarget = false;
};
} // namespace
@@ -52,13 +58,16 @@ static void populateEmptyFunction(func::FuncOp funcOp) {
SmallVector<spatial::SpatConcatPlanOp> concatPlans(funcOp.getOps<spatial::SpatConcatPlanOp>());
SmallVector<spatial::SpatReluPlanOp> reluPlans(funcOp.getOps<spatial::SpatReluPlanOp>());
SmallVector<spatial::SpatSiluPlanOp> siluPlans(funcOp.getOps<spatial::SpatSiluPlanOp>());
SmallVector<spatial::SpatResizeNearestPlanOp> resizePlans(
funcOp.getOps<spatial::SpatResizeNearestPlanOp>());
SmallVector<spatial::SpatMaxPool2DPlanOp> maxPoolPlans(funcOp.getOps<spatial::SpatMaxPool2DPlanOp>());
SmallVector<spatial::SpatGlobalAveragePoolPlanOp> globalAveragePoolPlans(
funcOp.getOps<spatial::SpatGlobalAveragePoolPlanOp>());
SmallVector<spatial::SpatBlueprintOp> blueprints(funcOp.getOps<spatial::SpatBlueprintOp>());
SmallVector<spatial::SpatMaterializeLayoutOp> materializers(funcOp.getOps<spatial::SpatMaterializeLayoutOp>());
if (!computes.empty() || !computeBatches.empty() || !convPlans.empty() || !biasAddPlans.empty() || !addPlans.empty()
|| !concatPlans.empty() || !reluPlans.empty() || !siluPlans.empty() || !maxPoolPlans.empty() || !blueprints.empty()
|| !concatPlans.empty() || !reluPlans.empty() || !siluPlans.empty() || !resizePlans.empty()
|| !maxPoolPlans.empty() || !blueprints.empty()
|| !globalAveragePoolPlans.empty() || !materializers.empty()) {
return;
}
@@ -103,6 +112,11 @@ static void populateEmptyFunction(func::FuncOp funcOp) {
void ONNXToSpatialPass::runOnOperation() {
ModuleOp moduleOp = getOperation();
if (!hasTarget) {
moduleOp.emitError("ONNX-to-Spatial lowering requires an injected SpatialTargetInfo");
signalPassFailure();
return;
}
MLIRContext* ctx = &getContext();
ConversionTarget preTarget(*ctx);
@@ -123,6 +137,14 @@ void ONNXToSpatialPass::runOnOperation() {
return;
}
RewritePatternSet matmulPatterns(ctx);
populateMatMulFusionPatterns(matmulPatterns, ctx, target);
if (failed(applyPatternsGreedily(moduleOp, std::move(matmulPatterns)))) {
moduleOp.emitError("failed to lower MatMul before producer conversion");
signalPassFailure();
return;
}
RewritePatternSet fusionPatterns(ctx);
populateElementwiseFusionPatterns(fusionPatterns, ctx);
if (failed(applyPatternsGreedily(moduleOp, std::move(fusionPatterns)))) {
@@ -171,7 +193,7 @@ void ONNXToSpatialPass::runOnOperation() {
target.addIllegalOp<ONNXSplitOp>();
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();
@@ -247,4 +269,8 @@ void ONNXToSpatialPass::runOnOperation() {
std::unique_ptr<Pass> createONNXToSpatialPass() { return std::make_unique<ONNXToSpatialPass>(); }
std::unique_ptr<Pass> createONNXToSpatialPass(const spatial::SpatialTargetInfo& target) {
return std::make_unique<ONNXToSpatialPass>(target);
}
} // namespace onnx_mlir
@@ -130,8 +130,7 @@ template <typename ComputeOpTy>
void verifyNoNestedFragmentAssemblyBlueprints(ComputeOpTy compute,
pim::CappedDiagnosticReporter& diagnostics) {
compute.getBody().walk([&](spatial::SpatBlueprintOp blueprint) {
std::optional<StringRef> 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");
@@ -150,6 +149,7 @@ void verifyLogicalTopLevelOps(func::FuncOp funcOp, pim::CappedDiagnosticReporter
spatial::SpatConcatPlanOp,
spatial::SpatReluPlanOp,
spatial::SpatSiluPlanOp,
spatial::SpatResizeNearestPlanOp,
spatial::SpatMaxPool2DPlanOp,
spatial::SpatGlobalAveragePoolPlanOp,
spatial::SpatBlueprintOp,
@@ -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);
+22 -5
View File
@@ -8,19 +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 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);
File diff suppressed because it is too large Load Diff
@@ -3,47 +3,277 @@
#include <algorithm>
#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<int64_t>(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<int64_t>(target.matrixShape.rows),
static_cast<int64_t>(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<int64_t>(1, geo.xbarSize / std::max<int64_t>(geo.k, geo.c));
geo.im2colElements = static_cast<uint64_t>(std::max<int64_t>(0, geo.p)) * static_cast<uint64_t>(std::max<int64_t>(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<int64_t>(
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<uint64_t>(std::max<int64_t>(0, problem.numChannelsOut))
* static_cast<uint64_t>(std::max<int64_t>(0, plan.geometry.k));
plan.scratchElements = plan.geometry.im2colElements;
plan.materializationElements = strategy == spatial::ConvLoweringStrategy::Depthwise
? 0
: std::min<uint64_t>(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<ConvPlan> buildDepthwiseCandidate(
const ConvProblem& problem, const spatial::SpatialTargetInfo& target) {
if (!problem.isDepthwise)
return mlir::failure();
return makeCandidatePlan(problem, spatial::ConvLoweringStrategy::Depthwise, target);
}
static mlir::FailureOr<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan> 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<ConvPlan, 8> buildConvPlanCandidates(
const ConvProblem& problem, const spatial::SpatialTargetInfo& target) {
ConvGeometry geo = buildConvGeometry(problem, target);
llvm::SmallVector<ConvPlan, 8> candidates;
auto append = [&](spatial::ConvLoweringStrategy strategy) {
mlir::FailureOr<ConvPlan> 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<uint64_t>(std::max<int64_t>(1, geo.k));
uint64_t chunkPositions = std::max<uint64_t>(1, pimConvIm2colMaxElements / patchElements);
uint64_t chunkPositions = std::max<uint64_t>(1, target.convIm2colMaxElements / patchElements);
chunkPositions = std::min<uint64_t>(chunkPositions, static_cast<uint64_t>(std::max<int64_t>(1, geo.p)));
chunkPositions = std::min<uint64_t>(chunkPositions, std::max<uint64_t>(1, pimConvStreamChunkPositions));
chunkPositions = std::min<uint64_t>(chunkPositions, std::max<uint64_t>(1, target.convStreamChunkPositions));
if (packFactor > 1 && chunkPositions > static_cast<uint64_t>(packFactor)) {
chunkPositions -= chunkPositions % static_cast<uint64_t>(packFactor);
@@ -52,24 +282,26 @@ uint64_t chooseStreamChunkPositions(const ConvGeometry& geo, int64_t packFactor)
return std::max<uint64_t>(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<int64_t>(0, rawBegin), std::min<int64_t>(state.xHeight, rawEnd)};
(outputRows.end - 1) * problem.strideHeight - problem.padHeightBegin
+ problem.dilationHeight * (problem.wHeight - 1) + 1;
return {std::max<int64_t>(0, rawBegin), std::min<int64_t>(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<int64_t>(0, -rawBegin);
demand.bottomHaloRows = std::max<int64_t>(0, rawEnd - state.xHeight);
demand.bottomHaloRows = std::max<int64_t>(0, rawEnd - problem.xHeight);
demand.acquiredInputRows = demand.neededInputRows;
return demand;
}
@@ -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 <cstdint>
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<ConvPlan> makeConvPlan(const ConvProblem& problem,
spatial::ConvLoweringStrategy strategy,
const spatial::SpatialTargetInfo& target);
ConvRowDemand buildConvRowDemand(RowInterval outputRows, const ConvLoweringState& state);
llvm::SmallVector<ConvPlan, 8> 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
@@ -31,7 +31,7 @@ struct SiluToSpatialPlan : OpRewritePattern<ONNXMulOp> {
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();
@@ -48,6 +48,56 @@ static DenseElementsAttr getDenseConstantAttr(Value value) {
return nullptr;
}
struct BlueprintSplatMulToSpatial : OpConversionPattern<ONNXMulOp> {
explicit BlueprintSplatMulToSpatial(MLIRContext* ctx) : OpConversionPattern(ctx, 2) {}
LogicalResult
matchAndRewrite(ONNXMulOp op, ONNXMulOpAdaptor adaptor, ConversionPatternRewriter& rewriter) const override {
auto blueprint = adaptor.getA().getDefiningOp<spatial::SpatBlueprintOp>();
Value scalar = adaptor.getB();
if (!blueprint) {
blueprint = adaptor.getB().getDefiningOp<spatial::SpatBlueprintOp>();
scalar = adaptor.getA();
}
auto scalarAttr = getDenseConstantAttr(scalar);
auto resultType = dyn_cast<RankedTensorType>(op.getResult().getType());
auto storageType = blueprint ? dyn_cast<RankedTensorType>(blueprint.getInput().getType()) : RankedTensorType();
if (!blueprint || !blueprint.getFragments().empty() || !scalarAttr || !scalarAttr.isSplat() || !resultType
|| resultType != blueprint.getOutput().getType() || !storageType)
return failure();
auto mapped = mapGraphBatchFragments(
blueprint.getInput(), storageType, rewriter, op.getLoc(), [&](Value fragment, RankedTensorType fragmentType) {
auto splat = DenseElementsAttr::get(fragmentType, scalarAttr.getSplatValue<Attribute>());
Value constant = arith::ConstantOp::create(rewriter, op.getLoc(), fragmentType, splat);
return FailureOr<Value>(
spatial::SpatVMulOp::create(rewriter, op.getLoc(), fragmentType, fragment, constant).getResult());
});
if (failed(mapped))
return failure();
auto result = spatial::SpatBlueprintOp::create(rewriter,
op.getLoc(),
resultType,
*mapped,
ValueRange {},
blueprint.getLogicalLayoutAttr(),
blueprint.getPhysicalLayoutAttr(),
blueprint.getFragmentOffsetsAttr(),
blueprint.getFragmentSizesAttr(),
blueprint.getIndexMapAttr(),
blueprint.getModeAttr(),
blueprint.getFragmentOperandIndicesAttr(),
blueprint.getFragmentSourceSlotsAttr(),
blueprint.getFragmentSourceOffsetsAttr(),
blueprint.getFragmentStridesAttr(),
blueprint.getConflictPolicyAttr(),
blueprint.getCoveragePolicyAttr());
rewriter.replaceOp(op, result.getOutput());
return success();
}
};
static FailureOr<Value> materializeBroadcastedConstantTensor(Value value,
RankedTensorType resultType,
ConversionPatternRewriter& rewriter,
@@ -210,14 +260,16 @@ struct AddToSpatialCompute : OpConversionPattern<ONNXAddOp> {
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();
}
@@ -246,6 +298,7 @@ void populateElementwiseFusionPatterns(RewritePatternSet& patterns, MLIRContext*
}
void populateElementwisePatterns(RewritePatternSet& patterns, MLIRContext* ctx) {
patterns.add<BlueprintSplatMulToSpatial>(ctx);
patterns.add<AddToSpatialCompute>(ctx);
patterns.add<BinaryElementwiseToSpatialCompute<ONNXSubOp, spatial::SpatVSubOp>>(ctx);
patterns.add<BinaryElementwiseToSpatialCompute<ONNXMulOp, spatial::SpatVMulOp>>(ctx);
@@ -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<Value>
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<Value> materializePaddedConstantMatrix(Value value,
RankedTensorType resultType,
ConversionPatternRewriter& rewriter,
PatternRewriter& rewriter,
Location loc) {
auto sourceType = cast<RankedTensorType>(value.getType());
if (sourceType == resultType)
@@ -121,7 +132,7 @@ static FailureOr<Value> materializePaddedConstantMatrix(Value value,
static FailureOr<Value> 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<Value> materializePaddedBroadcastedConstantTensor(Value value,
static FailureOr<Value> prepareBias(Value c,
RankedTensorType outType,
RankedTensorType paddedOutType,
ConversionPatternRewriter& rewriter,
PatternRewriter& rewriter,
Location loc) {
auto cType = cast<RankedTensorType>(c.getType());
if (!cType.hasStaticShape())
@@ -203,9 +214,15 @@ static FailureOr<Value> 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<OpFoldResult> offsets {row, kOffset};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)};
SmallVector<OpFoldResult> 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<spatial::SpatComputeBatch> 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<spatial::SpatComputeBatch> 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<int64_t>(crossbarSize.getValue())}, aType.getElementType());
RankedTensorType::get({1, xbarSize}, aType.getElementType());
auto bTileType = RankedTensorType::get(
{static_cast<int64_t>(crossbarSize.getValue()), static_cast<int64_t>(crossbarSize.getValue())},
{xbarSize, xbarSize},
paddedBType.getElementType());
auto pieceType =
RankedTensorType::get({1, static_cast<int64_t>(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<OpFoldResult> bOffsets {kOffset, hOffset};
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(crossbarSize.getValue()),
rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> bSizes {rewriter.getIndexAttr(xbarSize), rewriter.getIndexAttr(xbarSize)};
SmallVector<OpFoldResult> unitStrides = getUnitStrides(rewriter, 2);
Value bTile = extractStaticSliceOrIdentity(
rewriter, loc, args.weights.front(), bTileType, bOffsets, bSizes, unitStrides);
@@ -260,7 +278,7 @@ static FailureOr<spatial::SpatComputeBatch> 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<OpFoldResult> offsets {rewriter.getIndexAttr(0), column};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(vectorType.getDimSize(1)), rewriter.getIndexAttr(1)};
SmallVector<OpFoldResult> 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<OpFoldResult> offsets {row, rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(vectorType.getDimSize(1))};
SmallVector<OpFoldResult> 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<OpFoldResult> unitStrides(biasType.getRank(), rewriter.getIndexAttr(1));
if (biasType.getRank() == 1) {
@@ -365,7 +383,7 @@ static FailureOr<spatial::SpatComputeBatch> 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<spatial::SpatCompute> 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<OpFoldResult> unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)};
SmallVector<OpFoldResult> pieceSizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> pieceSizes {
rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)};
SmallVector<OpFoldResult> 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<Value> 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<Value> nextPieces;
@@ -574,11 +596,12 @@ static FailureOr<Value> 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<int64_t>(crossbarSize.getValue())},
const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(1), xbarSize);
auto pieceType = RankedTensorType::get({numOutRows, xbarSize},
partialPiecesType.getElementType());
if (bias && cast<RankedTensorType>(bias.getType()) != paddedOutType)
@@ -590,20 +613,20 @@ static FailureOr<Value> createReductionOutput(Value partialPieces,
SmallVector<Value> 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<int64_t>(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<OpFoldResult> biasOffsets {rewriter.getIndexAttr(0), rewriter.getIndexAttr(columnOffset)};
SmallVector<OpFoldResult> pieceSizes {rewriter.getIndexAttr(numOutRows),
rewriter.getIndexAttr(crossbarSize.getValue())};
rewriter.getIndexAttr(xbarSize)};
SmallVector<OpFoldResult> 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<Value> createReductionOutput(Value partialPieces,
}
struct GemmToSpatialComputes : OpConversionPattern<ONNXGemmOp> {
using OpConversionPattern::OpConversionPattern;
explicit GemmToSpatialComputes(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpConversionPattern<ONNXGemmOp>(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<Value> 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<RankedTensorType>(a.getType());
auto bType = dyn_cast<RankedTensorType>(b.getType());
auto outType = dyn_cast<RankedTensorType>(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<RankedTensorType>(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<int32_t>::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<RankedTensorType>(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<int64_t>(crossbarSize.getValue());
const int64_t paddedOutCols = numOutHSlices * static_cast<int64_t>(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<RankedTensorType>(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<int32_t>::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<int64_t>(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<Value> result = lowerGemmToSpatial(
gemmOp.getOperation(), gemmOpAdaptor.getA(), gemmOpAdaptor.getB(), gemmOpAdaptor.getC(),
cast<RankedTensorType>(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<GemmToSpatialComputes>(ctx);
void populateGemmPatterns(RewritePatternSet& patterns,
MLIRContext* ctx,
const spatial::SpatialTargetInfo& target) {
patterns.insert<GemmToSpatialComputes>(ctx, target);
}
} // namespace onnx_mlir
@@ -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<mlir::Value> 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
@@ -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"
@@ -106,6 +109,92 @@ static Value mapOutputBatchIndexToSourceBatchIndex(Value outputBatchIndex,
return sourceBatchIndex;
}
static FailureOr<Value> collapseFragmentAssemblyBatchDims(Value value,
RankedTensorType resultType,
PatternRewriter& rewriter,
Location loc) {
auto blueprint = value.getDefiningOp<spatial::SpatBlueprintOp>();
auto inputType = dyn_cast<RankedTensorType>(value.getType());
auto storageType = blueprint ? dyn_cast<RankedTensorType>(blueprint.getInput().getType()) : RankedTensorType();
auto operandIndices = blueprint ? blueprint.getFragmentOperandIndices() : std::nullopt;
auto sourceOffsets = blueprint ? blueprint.getFragmentSourceOffsets() : std::nullopt;
auto fragmentStrides = blueprint ? blueprint.getFragmentStrides() : std::nullopt;
if (!blueprint || !inputType || !storageType || !inputType.hasStaticShape() || !storageType.hasStaticShape()
|| inputType.getRank() <= 3 || resultType.getRank() != 3 || !blueprint.getFragments().empty()
|| !spatial::isFragmentAssembly(blueprint.getMode()) || !operandIndices || !sourceOffsets || !fragmentStrides
|| storageType.getRank() != inputType.getRank() + 1)
return failure();
if (blueprint.getIndexMap() == spatial::kContiguousRowMajorFragments
&& !spatial::isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
return blueprint.emitOpError("contiguous row-major fragment physical source order or storage is not canonical"), failure();
const int64_t batchRank = inputType.getRank() - 2;
SmallVector<ReassociationIndices> reassociation {ReassociationIndices {},
ReassociationIndices {batchRank},
ReassociationIndices {batchRank + 1}};
for (int64_t dim = 0; dim < batchRank; ++dim)
reassociation.front().push_back(dim);
SmallVector<int64_t> outputFragmentShape {1, storageType.getDimSize(batchRank + 1), storageType.getDimSize(batchRank + 2)};
auto outputFragmentType = RankedTensorType::get(outputFragmentShape, storageType.getElementType());
auto outputStorageType = spatial::getGraphBatchPhysicalResultType(storageType.getDimSize(0), outputFragmentType);
SmallVector<ReassociationIndices> storageReassociation {
ReassociationIndices {0}, ReassociationIndices {}, ReassociationIndices {batchRank + 1},
ReassociationIndices {batchRank + 2}};
for (int64_t dim = 0; dim < batchRank; ++dim)
storageReassociation[1].push_back(dim + 1);
Value collapsedStorage = tensor::CollapseShapeOp::create(
rewriter, loc, outputStorageType, blueprint.getInput(), storageReassociation);
const int64_t inputRank = inputType.getRank();
const int64_t fragmentCount = operandIndices->size();
ArrayRef<int64_t> inputOffsets = blueprint.getFragmentOffsets();
ArrayRef<int64_t> inputSizes = blueprint.getFragmentSizes();
SmallVector<int64_t> batchShape(inputType.getShape().drop_back(2));
SmallVector<int64_t> batchStrides = computeRowMajorStrides(batchShape);
SmallVector<int64_t> offsets, sizes, strides;
offsets.reserve(fragmentCount * 3);
sizes.reserve(fragmentCount * 3);
strides.reserve(fragmentCount * 3);
for (int64_t fragment = 0; fragment < fragmentCount; ++fragment) {
int64_t flatBatch = 0;
for (int64_t dim = 0; dim < batchRank; ++dim) {
const int64_t index = fragment * inputRank + dim;
if (inputSizes[index] != 1 || (*fragmentStrides)[index] != 1)
return failure();
flatBatch += inputOffsets[index] * batchStrides[dim];
}
offsets.push_back(flatBatch);
sizes.push_back(1);
strides.push_back(1);
for (int64_t dim = batchRank; dim < inputRank; ++dim) {
const int64_t index = fragment * inputRank + dim;
offsets.push_back(inputOffsets[index]);
sizes.push_back(inputSizes[index]);
strides.push_back((*fragmentStrides)[index]);
}
}
auto collapsedBlueprint = spatial::SpatBlueprintOp::create(rewriter,
loc,
resultType,
collapsedStorage,
ValueRange {},
blueprint.getLogicalLayoutAttr(),
spatial::getFragmentedLayout(rewriter.getContext()),
rewriter.getDenseI64ArrayAttr(offsets),
rewriter.getDenseI64ArrayAttr(sizes),
rewriter.getStringAttr("collapsed_fragments"),
blueprint.getModeAttr(),
blueprint.getFragmentOperandIndicesAttr(),
blueprint.getFragmentSourceSlotsAttr(),
blueprint.getFragmentSourceOffsetsAttr(),
rewriter.getDenseI64ArrayAttr(strides),
blueprint.getConflictPolicyAttr(),
blueprint.getCoveragePolicyAttr());
if (spatial::isCanonicalContiguousRowMajorFragmentAssembly(collapsedBlueprint))
collapsedBlueprint.setIndexMapAttr(rewriter.getStringAttr(spatial::kContiguousRowMajorFragments));
return collapsedBlueprint.getOutput();
}
static Value
collapseBatchDims(Value value, int64_t batchSize, int64_t rows, int64_t cols, PatternRewriter& rewriter, Location loc) {
auto type = cast<RankedTensorType>(value.getType());
@@ -113,6 +202,8 @@ collapseBatchDims(Value value, int64_t batchSize, int64_t rows, int64_t cols, Pa
return value;
auto collapsedType = RankedTensorType::get({batchSize, rows, cols}, type.getElementType(), type.getEncoding());
if (auto collapsed = collapseFragmentAssemblyBatchDims(value, collapsedType, rewriter, loc); succeeded(collapsed))
return *collapsed;
SmallVector<ReassociationIndices> reassociation = {ReassociationIndices {},
ReassociationIndices {static_cast<int64_t>(type.getRank() - 2)},
ReassociationIndices {static_cast<int64_t>(type.getRank() - 1)}};
@@ -241,7 +332,34 @@ static Value extractBatchMatrix(Value value,
return materializeOrComputeUnary(value, matrixType, rewriter, loc, buildMatrix);
}
static Value getLastTwoTransposeInput(Value value) {
auto type = cast<RankedTensorType>(value.getType());
if (auto transpose = value.getDefiningOp<ONNXTransposeOp>()) {
auto permutation = getTransposePermutationChecked(transpose.getPermAttr(), type.getRank());
if (succeeded(permutation) && llvm::all_of(llvm::seq<int64_t>(0, type.getRank()), [&](int64_t dim) {
return (*permutation)[dim] == (dim < type.getRank() - 2 ? dim : 2 * type.getRank() - 3 - dim);
}))
return transpose.getData();
}
return {};
}
static std::pair<Value, Value> splitSplatMultiply(Value value) {
if (!value)
return {};
auto multiply = value.getDefiningOp<ONNXMulOp>();
if (!multiply)
return {};
for (auto [data, scale] : {std::pair {multiply.getA(), multiply.getB()},
std::pair {multiply.getB(), multiply.getA()}})
if (auto constant = getHostConstDenseElementsAttr(scale); constant && constant.isSplat())
return {data, scale};
return {};
}
static Value transposeLastTwoDims(Value value, PatternRewriter& rewriter, Location loc) {
if (Value input = getLastTwoTransposeInput(value))
return input;
auto type = cast<RankedTensorType>(value.getType());
auto shape = type.getShape();
auto createONNXTranspose = [&](RankedTensorType resultType, ArrayRef<int64_t> permutation) {
@@ -347,6 +465,7 @@ static FailureOr<spatial::SpatComputeBatch> 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);
@@ -365,16 +484,16 @@ static FailureOr<spatial::SpatComputeBatch> 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<int64_t>(crossbarSize.getValue())}, aType.getElementType());
RankedTensorType::get({1, xbarSize}, aType.getElementType());
auto bTileType = RankedTensorType::get(
{static_cast<int64_t>(crossbarSize.getValue()), static_cast<int64_t>(crossbarSize.getValue())},
{xbarSize, xbarSize},
bType.getElementType());
auto pieceType =
RankedTensorType::get({1, static_cast<int64_t>(crossbarSize.getValue())}, partialPiecesType.getElementType());
RankedTensorType::get({1, xbarSize}, partialPiecesType.getElementType());
Value aTile = extractBatchedATile(
args.inputs.front(), aBatchShape, outputBatchShape, batch, row, kOffset, aTileType, rewriter, loc);
@@ -407,123 +526,155 @@ static Value extractDynamicBatchedRowVector(Value matrix,
{offsets, sizes, getUnitStrides(rewriter, 3)});
}
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)
if (rows % candidate == 0)
rowsPerLane = candidate;
return rowsPerLane;
}
static FailureOr<spatial::SpatComputeBatch> createBatchedVvdmulBatch(Value a,
ArrayRef<int64_t> aBatchShape,
Value b,
ArrayRef<int64_t> bBatchShape,
ArrayRef<int64_t> outputBatchShape,
RankedTensorType aType,
RankedTensorType bType,
RankedTensorType columnPiecesType,
int64_t reductionSize,
int64_t rowsPerLane,
RankedTensorType rowPiecesType,
RankedTensorType outType,
RankedTensorType publicationFragmentType,
PatternRewriter& rewriter,
Location loc) {
const int64_t numBatches = outType.getDimSize(0);
const int64_t numOutRows = outType.getDimSize(1);
const int64_t numOutCols = outType.getDimSize(2);
const int64_t reductionSize = aType.getDimSize(2);
const int64_t laneCount = numBatches * numOutCols;
auto vectorType = RankedTensorType::get({1, reductionSize}, aType.getElementType());
const int64_t rowGroups = numOutRows / rowsPerLane;
const int64_t laneCount = numBatches * rowGroups;
auto vectorType = RankedTensorType::get({1, reductionSize}, outType.getElementType());
auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType());
auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType());
auto rowType = RankedTensorType::get({1, numOutCols}, outType.getElementType());
auto rowsType = RankedTensorType::get({rowsPerLane, numOutCols}, outType.getElementType());
auto batchOp = createSpatComputeBatch(
rewriter,
loc,
TypeRange {columnPiecesType},
TypeRange {rowPiecesType},
laneCount,
ValueRange {},
ValueRange {a, b},
[&](detail::SpatComputeBatchBodyArgs args) {
[&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult {
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value batch = affineFloorDivConst(rewriter, loc, args.lane, numOutCols, anchorOp);
Value column = affineModConst(rewriter, loc, args.lane, numOutCols, anchorOp);
Value bVector = extractDynamicBatchedRowVector(
args.inputs[1], bBatchShape, outputBatchShape, batch, column, vectorType, rewriter, loc);
Value columnInit = tensor::EmptyOp::create(rewriter, loc, columnType.getShape(), columnType.getElementType());
Value batch = affineFloorDivConst(rewriter, loc, args.lane, rowGroups, anchorOp);
Value rowGroup = affineModConst(rewriter, loc, args.lane, rowGroups, anchorOp);
Value rowBase = affineMulConst(rewriter, loc, rowGroup, rowsPerLane, anchorOp);
Value c0 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0);
Value c1 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 1);
Value cNumOutRows =
getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), numOutRows);
auto loop = buildNormalizedScfFor(
Value cRows = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), rowsPerLane);
Value cNumOutCols =
getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), numOutCols);
Value rowsInit = tensor::EmptyOp::create(rewriter, loc, rowsType.getShape(), rowsType.getElementType());
auto rowsLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cNumOutRows,
cRows,
c1,
ValueRange {columnInit},
[&](OpBuilder&, Location nestedLoc, Value row, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
ValueRange {rowsInit},
[&](OpBuilder&, Location nestedLoc, Value rowOffset, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value row = arith::AddIOp::create(rewriter, nestedLoc, rowBase, rowOffset);
Value aVector = extractDynamicBatchedRowVector(
args.inputs[0], aBatchShape, outputBatchShape, batch, row, vectorType, rewriter, nestedLoc);
Value scalar = spatial::SpatVVDMulOp::create(rewriter, nestedLoc, scalarType, aVector, bVector).getResult();
Value next = tensor::InsertSliceOp::create(rewriter,
nestedLoc,
scalar,
iterArgs.front(),
SmallVector<OpFoldResult> {row, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1)},
getUnitStrides(rewriter, 2));
yielded.push_back(next);
Value rowInit = tensor::EmptyOp::create(rewriter, nestedLoc, rowType.getShape(), rowType.getElementType());
auto columnsLoop = buildNormalizedScfFor(
rewriter,
nestedLoc,
c0,
cNumOutCols,
c1,
ValueRange {rowInit},
[&](OpBuilder&, Location columnLoc, Value column, ValueRange columnArgs, SmallVectorImpl<Value>& rowYielded) {
Value bVector = extractDynamicBatchedRowVector(
args.inputs[1], bBatchShape, outputBatchShape, batch, column, vectorType, rewriter, columnLoc);
Value scalar = spatial::SpatVVDMulOp::create(rewriter, columnLoc, scalarType, aVector, bVector).getResult();
rowYielded.push_back(tensor::InsertSliceOp::create(
rewriter,
columnLoc,
scalar,
columnArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), column},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)},
getUnitStrides(rewriter, 2)));
return success();
});
assert(succeeded(columnsLoop) && "dynamic MatMul column loop construction must succeed");
yielded.push_back(tensor::InsertSliceOp::create(
rewriter,
nestedLoc,
columnsLoop->results.front(),
iterArgs.front(),
SmallVector<OpFoldResult> {rowOffset, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(numOutCols)},
getUnitStrides(rewriter, 2)));
return success();
});
assert(succeeded(loop) && "dynamic MatMul row loop construction must succeed");
publishGraphBatchPhysicalFragment(rewriter, loc, loop->results.front(), args.outputs.front(), args.lane);
assert(succeeded(rowsLoop) && "dynamic MatMul row-group loop construction must succeed");
Value fragment = rowsLoop->results.front();
while (fragment.getType() != publicationFragmentType) {
auto expanded = addLeadingUnitTensorDimension(rewriter, loc, fragment);
if (failed(expanded))
return failure();
fragment = *expanded;
}
publishGraphBatchPhysicalFragment(rewriter, loc, fragment, args.outputs.front(), args.lane);
return success();
});
if (failed(batchOp))
return failure();
return *batchOp;
}
static FailureOr<Value> createBatchedDynamicOutputCompute(Value scalarPieces,
RankedTensorType scalarPiecesType,
RankedTensorType outType,
PatternRewriter& rewriter,
Location loc) {
const int64_t laneCount = scalarPiecesType.getDimSize(0);
const int64_t numOutCols = outType.getDimSize(2);
auto columnType = RankedTensorType::get({outType.getDimSize(1), 1}, outType.getElementType());
auto computeOp = createSpatCompute<1>(
rewriter, loc, TypeRange {outType}, {}, ValueRange {scalarPieces}, [&](Value pieces) -> LogicalResult {
Value outputInit =
tensor::EmptyOp::create(rewriter, loc, outType.getShape(), outType.getElementType()).getResult();
Value c0 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 0);
Value c1 = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), 1);
Value cLaneCount = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), laneCount);
auto loop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cLaneCount,
c1,
ValueRange {outputInit},
[&](OpBuilder&, Location nestedLoc, Value lane, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value outputAcc = iterArgs.front();
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value batch = affineFloorDivConst(rewriter, nestedLoc, lane, numOutCols, anchorOp);
Value column = affineModConst(rewriter, nestedLoc, lane, numOutCols, anchorOp);
FailureOr<Value> columnPiece =
extractGraphBatchPhysicalFragment(rewriter, nestedLoc, pieces, lane, columnType);
if (failed(columnPiece))
return failure();
SmallVector<OpFoldResult> outputOffsets {batch, rewriter.getIndexAttr(0), column};
SmallVector<OpFoldResult> outputSizes = {
rewriter.getIndexAttr(1), rewriter.getIndexAttr(outType.getDimSize(1)), rewriter.getIndexAttr(1)};
Value next =
tensor::InsertSliceOp::create(
rewriter, nestedLoc, *columnPiece, outputAcc, outputOffsets, outputSizes, getUnitStrides(rewriter, 3))
.getResult();
yielded.push_back(next);
return success();
});
if (failed(loop))
return failure();
spatial::SpatYieldOp::create(rewriter, loc, loop->results.front());
return success();
});
if (failed(computeOp))
return failure();
return computeOp->getResult(0);
static FailureOr<Value> createBatchedRowOutputBlueprint(Value rowPieces,
RankedTensorType outType,
ArrayRef<int64_t> batchShape,
int64_t rowsPerFragment,
PatternRewriter& rewriter,
Location loc) {
SmallVector<FragmentAssemblyEntry> entries;
const int64_t rowAxis = outType.getRank() - 2;
const int64_t rows = outType.getDimSize(rowAxis);
const int64_t columns = outType.getDimSize(rowAxis + 1);
const int64_t batches = batchShape.empty() ? 1 : getStaticShapeElementCount(batchShape);
SmallVector<int64_t> batchStrides = computeRowMajorStrides(batchShape);
const int64_t rowGroups = rows / rowsPerFragment;
entries.reserve(batches * rowGroups);
for (int64_t batch = 0; batch < batches; ++batch)
for (int64_t row = 0; row < rows; row += rowsPerFragment) {
SmallVector<int64_t, 4> offsets;
for (auto [dim, size] : llvm::enumerate(batchShape))
offsets.push_back((batch / batchStrides[dim]) % size);
offsets.push_back(row);
offsets.push_back(0);
SmallVector<int64_t, 4> sizes(outType.getRank(), 1);
sizes[rowAxis] = rowsPerFragment;
sizes.back() = columns;
entries.push_back({batch * rowGroups + row / rowsPerFragment,
0,
std::move(offsets),
std::move(sizes)});
}
return createFragmentAssemblyBlueprint(
rowPieces,
outType,
entries,
"dense_nchw",
rowsPerFragment == 1 ? spatial::kContiguousRowMajorFragments : "row_group_fragments",
rewriter,
loc);
}
static Value extractBatchedReductionPiece(Value partialPiecesArg,
@@ -534,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();
@@ -543,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<OpFoldResult> offsets {pieceOffset, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
SmallVector<OpFoldResult> sizes {rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(crossbarSize.getValue())};
SmallVector<OpFoldResult> sizes {
rewriter.getIndexAttr(numOutRows), rewriter.getIndexAttr(1), rewriter.getIndexAttr(xbarSize)};
return extractMixedSliceOrIdentity(
rewriter, loc, partialPiecesArg, pieceType,
{offsets, sizes, getUnitStrides(rewriter, 3)});
@@ -556,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<Value> 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<Value> nextPieces;
@@ -585,13 +749,14 @@ static FailureOr<Value> 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<int64_t>(crossbarSize.getValue())},
const int64_t numOutHSlices = ceilIntegerDivide(outType.getDimSize(2), xbarSize);
auto pieceType = RankedTensorType::get({numOutRows, xbarSize},
partialPiecesType.getElementType());
Value outputInit =
@@ -621,13 +786,22 @@ static FailureOr<Value> createBatchedReductionCompute(Value partialPieces,
[&](OpBuilder&, Location hLoc, Value hSlice, ValueRange hIterArgs, SmallVectorImpl<Value>& 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<OpFoldResult> outputOffsets {batch, rewriter.getIndexAttr(0), hOffset};
SmallVector<OpFoldResult> 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))
@@ -661,39 +835,45 @@ static FailureOr<Value> 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<int64_t> lhsBatchShape;
SmallVector<int64_t> rhsBatchShape;
SmallVector<int64_t> 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<RankedTensorType>(lhs.getType())),
rhsType(cast<RankedTensorType>(rhs.getType())) {}
Value lhs;
Value rhs;
RankedTensorType lhsType;
RankedTensorType rhsType;
SmallVector<int64_t> lhsBatchShape;
SmallVector<int64_t> rhsBatchShape;
SmallVector<int64_t> outputBatchShape;
int64_t lhsBatch;
int64_t rhsBatch;
int64_t batch;
int64_t m;
int64_t k;
int64_t n;
bool transposedResult;
};
@@ -757,22 +937,31 @@ static FailureOr<NormalizedMatMulInfo> 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,
@@ -781,20 +970,8 @@ static MatMulLoweringPlan buildLoweringPlan(Value normalizedLhs,
bool useTransposedForm,
PatternRewriter& rewriter,
Location loc) {
MatMulLoweringPlan plan {normalizedLhs,
normalizedRhs,
cast<RankedTensorType>(normalizedLhs.getType()),
cast<RankedTensorType>(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;
@@ -822,6 +999,8 @@ static Value finalizeNormalizedMatMulResult(Value value,
const NormalizedMatMulInfo& info,
PatternRewriter& rewriter,
Location loc) {
if (value.getType() == info.outType)
return value;
// The direct lowered result is always [flatBatch, normalizedM, normalizedN].
// Restore ONNX MatMul result rank by expanding right-aligned batch dimensions
// and removing the synthetic unit matrix axes introduced for vector operands.
@@ -911,7 +1090,9 @@ struct MatMulToGemm : OpRewritePattern<ONNXMatMulOp> {
};
struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
using OpRewritePattern::OpRewritePattern;
explicit MatMulBatchedToSpatialComputes(MLIRContext* ctx,
const spatial::SpatialTargetInfo& target)
: OpRewritePattern<ONNXMatMulOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(ONNXMatMulOp matmulOp, PatternRewriter& rewriter) const override {
auto shapeInfo = analyzeMatMulShape(matmulOp);
@@ -921,29 +1102,50 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
return failure();
Location loc = matmulOp.getLoc();
const int64_t xbarSize = static_cast<int64_t>(target.matrixShape.rows);
bool useTransposedForm = !shapeInfo->lhsWasVector && !shapeInfo->rhsWasVector
&& isCompileTimeComputable(matmulOp.getA()) && !isCompileTimeComputable(matmulOp.getB());
Value rhsRows = getLastTwoTransposeInput(matmulOp.getB());
ONNXTransposeOp foldedTranspose = matmulOp.getB().getDefiningOp<ONNXTransposeOp>();
ONNXMulOp foldedMultiply = rhsRows ? rhsRows.getDefiningOp<ONNXMulOp>() : ONNXMulOp {};
auto [unscaledRhsRows, outputScale] = splitSplatMultiply(rhsRows);
if (unscaledRhsRows)
rhsRows = unscaledRhsRows;
const bool rhsStoredAsRows = rhsRows && !useTransposedForm;
Value lhs =
normalizeMatMulOperand(matmulOp.getA(), shapeInfo->normalizedLhsType, shapeInfo->lhsWasVector, rewriter, loc);
Value rhs =
normalizeMatMulOperand(matmulOp.getB(), shapeInfo->normalizedRhsType, shapeInfo->rhsWasVector, rewriter, loc);
Value rhs = normalizeMatMulOperand(
rhsStoredAsRows ? rhsRows : matmulOp.getB(), shapeInfo->normalizedRhsType, shapeInfo->rhsWasVector, rewriter, loc);
lhs = collapseBatchDims(lhs, shapeInfo->lhsBatch, shapeInfo->m, shapeInfo->k, rewriter, loc);
rhs = collapseBatchDims(rhs, shapeInfo->rhsBatch, shapeInfo->k, shapeInfo->n, rewriter, loc);
MatMulLoweringPlan plan = buildLoweringPlan(lhs, rhs, *shapeInfo, useTransposedForm, rewriter, loc);
rhs = collapseBatchDims(rhs,
shapeInfo->rhsBatch,
rhsStoredAsRows ? shapeInfo->n : shapeInfo->k,
rhsStoredAsRows ? shapeInfo->k : shapeInfo->n,
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, plan.rhsBatch, plan.k, plan.n, rewriter, loc);
plan.rhs = ensureBatchedTensor(plan.rhs,
plan.rhsBatch,
rhsStoredAsRows ? plan.n : plan.k,
rhsStoredAsRows ? plan.k : plan.n,
rewriter,
loc);
plan.lhsType = cast<RankedTensorType>(plan.lhs.getType());
plan.rhsType = cast<RankedTensorType>(plan.rhs.getType());
auto directOutType = RankedTensorType::get(
{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<int64_t>(crossbarSize.getValue());
const int64_t paddedOutCols = numOutHSlices * static_cast<int64_t>(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(
@@ -954,10 +1156,11 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
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<int64_t>(crossbarSize.getValue())}, shapeInfo->outType.getElementType()));
laneCount, RankedTensorType::get({1, xbarSize}, shapeInfo->outType.getElementType()));
auto batchOp = createBatchedVmmBatch(paddedLhs,
*paddedRhs,
paddedLhsType,
@@ -969,6 +1172,7 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
plan.m,
numKSlices,
numOutHSlices,
xbarSize,
rewriter,
loc);
if (failed(batchOp))
@@ -979,6 +1183,7 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
paddedOutType,
plan.batch,
numKSlices,
xbarSize,
rewriter,
loc);
if (failed(result))
@@ -997,26 +1202,37 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
return success();
}
}
const int64_t laneCount = plan.batch * plan.n;
auto columnType = RankedTensorType::get({plan.m, 1}, shapeInfo->outType.getElementType());
auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType(
laneCount, columnType);
Value transposedRhs = transposeLastTwoDims(plan.rhs, rewriter, loc);
RankedTensorType blueprintType = !shapeInfo->lhsWasVector && !shapeInfo->rhsWasVector
? shapeInfo->outType : directOutType;
SmallVector<int64_t> blueprintBatchShape = !shapeInfo->lhsWasVector && !shapeInfo->rhsWasVector
? shapeInfo->outputBatchShape : SmallVector<int64_t> {plan.batch};
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<int64_t> fragmentShape(blueprintType.getRank(), 1);
fragmentShape[fragmentShape.size() - 2] = rowsPerLane;
fragmentShape.back() = plan.n;
auto fragmentType = RankedTensorType::get(fragmentShape, shapeInfo->outType.getElementType());
auto rowPiecesType = spatial::getGraphBatchPhysicalResultType(laneCount, fragmentType);
Value transposedRhs = rhsStoredAsRows ? plan.rhs : transposeLastTwoDims(plan.rhs, rewriter, loc);
auto batchOp = createBatchedVvdmulBatch(plan.lhs,
plan.lhsBatchShape,
transposedRhs,
plan.rhsBatchShape,
plan.outputBatchShape,
plan.lhsType,
plan.rhsType,
scalarPiecesType,
plan.k,
rowsPerLane,
rowPiecesType,
directOutType,
fragmentType,
rewriter,
loc);
if (failed(batchOp))
return failure();
auto result =
createBatchedDynamicOutputCompute(batchOp->getResult(0), scalarPiecesType, directOutType, rewriter, loc);
auto result = createBatchedRowOutputBlueprint(
batchOp->getResult(0), blueprintType, blueprintBatchShape, rowsPerLane, rewriter, loc);
if (failed(result))
return failure();
Value finalResult = *result;
@@ -1029,15 +1245,43 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern<ONNXMatMulOp> {
.getResult();
}
finalResult = finalizeNormalizedMatMulResult(finalResult, directOutType, *shapeInfo, rewriter, loc);
if (outputScale)
finalResult = ONNXMulOp::create(
rewriter, loc, shapeInfo->outType, finalResult, outputScale).getResult();
rewriter.replaceOp(matmulOp, finalResult);
if (foldedTranspose && foldedTranspose->use_empty())
rewriter.eraseOp(foldedTranspose);
if (foldedMultiply && foldedMultiply->use_empty())
rewriter.eraseOp(foldedMultiply);
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct TransposedRhsMatMulToSpatial : MatMulBatchedToSpatialComputes {
using MatMulBatchedToSpatialComputes::MatMulBatchedToSpatialComputes;
LogicalResult matchAndRewrite(ONNXMatMulOp matmulOp, PatternRewriter& rewriter) const override {
if (!getLastTwoTransposeInput(matmulOp.getB()))
return failure();
return MatMulBatchedToSpatialComputes::matchAndRewrite(matmulOp, rewriter);
}
};
} // namespace
void populateMatMulRewritePatterns(RewritePatternSet& patterns, MLIRContext* ctx) {
patterns.insert<MatMulToGemm, MatMulBatchedToSpatialComputes>(ctx);
void populateMatMulFusionPatterns(RewritePatternSet& patterns,
MLIRContext* ctx,
const spatial::SpatialTargetInfo& target) {
patterns.add<TransposedRhsMatMulToSpatial>(ctx, target);
}
void populateMatMulRewritePatterns(RewritePatternSet& patterns,
MLIRContext* ctx,
const spatial::SpatialTargetInfo& target) {
patterns.insert<MatMulToGemm>(ctx);
patterns.insert<MatMulBatchedToSpatialComputes>(ctx, target);
}
} // namespace onnx_mlir
@@ -280,12 +280,12 @@ static FailureOr<Value> buildReduceMeanKeepdimsBlueprint(
SmallVector<int64_t> 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),
@@ -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 <typename PoolOp>
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<PoolOp, ONNXMaxPoolSingleOutOp>);
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 <typename PoolOp, typename PoolOpAdaptor, typename ReduceOp>
struct PoolToSpatialComputeBase : public OpConversionPattern<PoolOp> {
using OpConversionPattern<PoolOp>::OpConversionPattern;
PoolToSpatialComputeBase(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpConversionPattern<PoolOp>(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<PoolOp> {
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<PoolOp> {
&& 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<int64_t>(crossbarSize.getValue());
const int64_t xbarSize = static_cast<int64_t>(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<PoolOp> {
auto computeOp =
createSpatCompute<numInputs>(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<PoolOp, ONNXMaxPoolSingleOutOp>);
Value pooledOutputInit = tensor::EmptyOp::create(rewriter, loc, outType.getShape(), outType.getElementType());
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
@@ -424,7 +430,8 @@ struct PoolToSpatialCompute<ONNXAveragePoolOp>
} // namespace
LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp) {
LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo&) {
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape())
@@ -439,6 +446,118 @@ LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp)
return success();
}
FailureOr<Value> lowerDenseMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo& target,
PatternRewriter& rewriter) {
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(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<int64_t>(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<Value>& 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<int64_t>(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<OpFoldResult> {
batch, rewriter.getIndexAttr(tile * tileWidth), sourceRow, sourceColumn},
SmallVector<OpFoldResult> {
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<OpFoldResult> {
batch, rewriter.getIndexAttr(tile * tileWidth), outputRow, outputColumn},
SmallVector<OpFoldResult> {
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<Value> lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
std::optional<Value> 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<Value> 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<Value> 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<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!inputType || !outputType || !inputType.hasStaticShape() || !outputType.hasStaticShape())
@@ -697,10 +818,87 @@ LogicalResult canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAverage
return success();
}
FailureOr<Value> lowerDenseGlobalAveragePoolPlan(
spatial::SpatGlobalAveragePoolPlanOp planOp,
const spatial::SpatialTargetInfo& target,
PatternRewriter& rewriter) {
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(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<FloatType>(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<int64_t>(1, target.matrixShape.rows);
const int64_t channelTileCount = (channels + tileWidth - 1) / tileWidth;
const double scaleValue = 1.0 / static_cast<double>(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<int64_t>(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<OpFoldResult> {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(tile * tileWidth),
rewriter.getIndexAttr(row), rewriter.getIndexAttr(column)},
SmallVector<OpFoldResult> {
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<OpFoldResult> {
rewriter.getIndexAttr(0), rewriter.getIndexAttr(tile * tileWidth),
rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {
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<Value> lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp,
std::optional<Value> 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<Value> 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<Value> lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePo
return batch->getResult(0);
}
void populatePoolPatterns(RewritePatternSet& patterns, MLIRContext* ctx) {
patterns.insert<PoolToSpatialCompute<ONNXMaxPoolSingleOutOp>>(ctx);
patterns.insert<PoolToSpatialCompute<ONNXAveragePoolOp>>(ctx);
void populatePoolPatterns(RewritePatternSet& patterns,
MLIRContext* ctx,
const spatial::SpatialTargetInfo& target) {
patterns.insert<PoolToSpatialCompute<ONNXMaxPoolSingleOutOp>>(ctx, target);
patterns.insert<PoolToSpatialCompute<ONNXAveragePoolOp>>(ctx, target);
}
} // namespace onnx_mlir
@@ -17,7 +17,7 @@ struct ReluToSpatialCompute : OpConversionPattern<ONNXReluOp> {
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();
}
@@ -32,7 +32,8 @@ struct Concat : public OpConversionPattern<ONNXConcatOp> {
return type && type.hasStaticShape() && type.getRank() == 4;
})) {
rewriter.replaceOpWithNewOp<spatial::SpatConcatPlanOp>(
maxpoolOp, resultType, inputs, rewriter.getI64IntegerAttr(axis), rewriter.getStringAttr("nchw"));
maxpoolOp, resultType, inputs, rewriter.getI64IntegerAttr(axis),
spatial::getNCHWLayout(rewriter.getContext()));
return success();
}
@@ -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<RowStripFlattenAnalysis> analyzeRowStripFlatten(spatial::SpatGraphCompute flattenOp) {
static FailureOr<RowStripFlattenAnalysis> 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<RowStripFlattenAnalysis> 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<int64_t>(crossbarSize.getValue());
const int64_t xbarDim = static_cast<int64_t>(target.matrixShape.rows);
if (channels > xbarDim && channels % xbarDim != 0)
return failure();
@@ -162,14 +162,16 @@ static FailureOr<RowStripFlattenAnalysis> analyzeRowStripFlatten(spatial::SpatGr
void populateFlattenPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.add<Flatten>(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<RowStripFlattenAnalysis> analysis = analyzeRowStripFlatten(flattenOp);
FailureOr<RowStripFlattenAnalysis> analysis = analyzeRowStripFlatten(flattenOp, target);
if (failed(analysis))
return failure();
auto storageType = dyn_cast<RankedTensorType>(input.storage.getType());
@@ -5,8 +5,11 @@
#include "llvm/ADT/STLExtras.h"
#include "src/Accelerators/PIM/Common/IR/AffineUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Dialect/ONNX/ONNXOps.hpp"
@@ -17,126 +20,144 @@ namespace onnx_mlir {
namespace {
static Value buildNearestAsymmetricIndex(
Value outputIndex, int64_t inputDim, int64_t outputDim, ConversionPatternRewriter& rewriter, Location loc) {
Value outputIndex, int64_t inputDim, int64_t outputDim, PatternRewriter& rewriter, Location loc) {
if (inputDim == outputDim)
return outputIndex;
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
if (outputDim % inputDim == 0)
return affineFloorDivConst(rewriter, loc, outputIndex, outputDim / inputDim, anchorOp);
if (inputDim % outputDim == 0)
return affineMulConst(rewriter, loc, outputIndex, inputDim / outputDim, anchorOp);
Value cInputDim = getOrCreateIndexConstant(rewriter, anchorOp, inputDim);
Value cOutputDim = getOrCreateIndexConstant(rewriter, anchorOp, outputDim);
Value cInputDimLast = getOrCreateIndexConstant(rewriter, anchorOp, inputDim - 1);
Value scaledIndex = arith::MulIOp::create(rewriter, loc, outputIndex, cInputDim);
Value inputIndex = arith::DivUIOp::create(rewriter, loc, scaledIndex, cOutputDim);
return arith::MinUIOp::create(rewriter, loc, inputIndex, cInputDimLast);
return arith::DivUIOp::create(rewriter, loc, scaledIndex, cOutputDim);
}
static FailureOr<Value> buildNearestResizeLoop(Value input,
RankedTensorType inputType,
RankedTensorType resultType,
ConversionPatternRewriter& rewriter,
Location loc) {
auto elemType = resultType.getElementType();
SmallVector<int64_t> unitShape(resultType.getRank(), 1);
auto unitTensorType = RankedTensorType::get(unitShape, elemType);
SmallVector<OpFoldResult> unitSizes(resultType.getRank(), rewriter.getIndexAttr(1));
SmallVector<OpFoldResult> unitStrides(resultType.getRank(), rewriter.getIndexAttr(1));
Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp();
Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1);
Value cOutputN = getOrCreateIndexConstant(rewriter, anchorOp, resultType.getDimSize(0));
Value cOutputC = getOrCreateIndexConstant(rewriter, anchorOp, resultType.getDimSize(1));
Value cOutputH = getOrCreateIndexConstant(rewriter, anchorOp, resultType.getDimSize(2));
Value cOutputW = getOrCreateIndexConstant(rewriter, anchorOp, resultType.getDimSize(3));
Value outputInit = tensor::EmptyOp::create(rewriter, loc, resultType.getShape(), elemType);
auto batchLoop = buildNormalizedScfFor(
rewriter,
loc,
c0,
cOutputN,
c1,
ValueRange {outputInit},
[&](OpBuilder&, Location nestedLoc, Value outputN, ValueRange batchIterArgs, SmallVectorImpl<Value>& batchYielded) {
Value outputBatchAcc = batchIterArgs.front();
Value inputN =
buildNearestAsymmetricIndex(outputN, inputType.getDimSize(0), resultType.getDimSize(0), rewriter, nestedLoc);
auto channelLoop = buildNormalizedScfFor(
rewriter,
nestedLoc,
c0,
cOutputC,
c1,
ValueRange {outputBatchAcc},
[&](OpBuilder&,
Location channelLoc,
Value outputC,
ValueRange channelIterArgs,
SmallVectorImpl<Value>& channelYielded) {
Value outputChannelAcc = channelIterArgs.front();
Value inputC = buildNearestAsymmetricIndex(
outputC, inputType.getDimSize(1), resultType.getDimSize(1), rewriter, channelLoc);
auto heightLoop = buildNormalizedScfFor(
rewriter,
channelLoc,
c0,
cOutputH,
c1,
ValueRange {outputChannelAcc},
[&](OpBuilder&,
Location heightLoc,
Value outputH,
ValueRange heightIterArgs,
SmallVectorImpl<Value>& heightYielded) {
Value outputHeightAcc = heightIterArgs.front();
Value inputH = buildNearestAsymmetricIndex(
outputH, inputType.getDimSize(2), resultType.getDimSize(2), rewriter, heightLoc);
auto widthLoop = buildNormalizedScfFor(
rewriter,
heightLoc,
c0,
cOutputW,
c1,
ValueRange {outputHeightAcc},
[&](OpBuilder&,
Location widthLoc,
Value outputW,
ValueRange widthIterArgs,
SmallVectorImpl<Value>& widthYielded) {
Value outputWidthAcc = widthIterArgs.front();
Value inputW = buildNearestAsymmetricIndex(
outputW, inputType.getDimSize(3), resultType.getDimSize(3), rewriter, widthLoc);
SmallVector<OpFoldResult> inputOffsets = {inputN, inputC, inputH, inputW};
Value inputSlice = tensor::ExtractSliceOp::create(
rewriter, widthLoc, unitTensorType, input, inputOffsets, unitSizes, unitStrides);
SmallVector<OpFoldResult> outputOffsets = {outputN, outputC, outputH, outputW};
Value updatedOutput = tensor::InsertSliceOp::create(
rewriter, widthLoc, inputSlice, outputWidthAcc, outputOffsets, unitSizes, unitStrides);
widthYielded.push_back(updatedOutput);
return success();
});
if (failed(widthLoop))
return failure();
heightYielded.push_back(widthLoop->results.front());
return success();
});
if (failed(heightLoop))
return failure();
channelYielded.push_back(heightLoop->results.front());
static FailureOr<Value> buildDenseNearestResize(Value input,
RankedTensorType inputType,
RankedTensorType resultType,
PatternRewriter& rewriter,
Location loc) {
ArrayRef<int64_t> shape = resultType.getShape();
int64_t rowCount = shape[0] * shape[1] * shape[2];
auto scalarType = RankedTensorType::get({1, 1, 1, 1}, resultType.getElementType());
auto rowType = RankedTensorType::get({1, 1, 1, shape[3]}, resultType.getElementType());
auto rowsType = RankedTensorType::get({rowCount, 1, 1, 1, shape[3]}, resultType.getElementType());
auto batch = createSpatComputeBatch(
rewriter, loc, TypeRange {rowsType}, rowCount, {}, ValueRange {input},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchor = rewriter.getInsertionBlock()->getParentOp();
Value outputN = affineFloorDivConst(rewriter, loc, args.lane, shape[1] * shape[2], anchor);
Value channelRow = affineModConst(rewriter, loc, args.lane, shape[1] * shape[2], anchor);
Value outputC = affineFloorDivConst(rewriter, loc, channelRow, shape[2], anchor);
Value outputH = affineModConst(rewriter, loc, channelRow, shape[2], anchor);
Value inputN = buildNearestAsymmetricIndex(outputN, inputType.getDimSize(0), shape[0], rewriter, loc);
Value inputC = buildNearestAsymmetricIndex(outputC, inputType.getDimSize(1), shape[1], rewriter, loc);
Value inputH = buildNearestAsymmetricIndex(outputH, inputType.getDimSize(2), shape[2], rewriter, loc);
Value row = tensor::EmptyOp::create(rewriter, loc, rowType.getShape(), rowType.getElementType());
Value c0 = getOrCreateIndexConstant(rewriter, anchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchor, 1);
Value width = getOrCreateIndexConstant(rewriter, anchor, shape[3]);
auto loop = buildNormalizedScfFor(
rewriter, loc, c0, width, c1, ValueRange {row},
[&](OpBuilder&, Location nestedLoc, Value outputW, ValueRange iterArgs, SmallVectorImpl<Value>& yielded) {
Value inputW = buildNearestAsymmetricIndex(
outputW, inputType.getDimSize(3), shape[3], rewriter, nestedLoc);
SmallVector<OpFoldResult> unitSizes(4, rewriter.getIndexAttr(1));
SmallVector<OpFoldResult> unitStrides(4, rewriter.getIndexAttr(1));
Value scalar = tensor::ExtractSliceOp::create(
rewriter, nestedLoc, scalarType, args.inputs.front(),
SmallVector<OpFoldResult> {inputN, inputC, inputH, inputW}, unitSizes, unitStrides);
yielded.push_back(tensor::InsertSliceOp::create(
rewriter, nestedLoc, scalar, iterArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0), outputW},
unitSizes, unitStrides));
return success();
});
if (failed(channelLoop))
assert(succeeded(loop) && "nearest Resize row loop construction must succeed");
publishGraphBatchPhysicalFragment(rewriter, loc, loop->results.front(), args.outputs.front(), args.lane);
});
if (failed(batch))
return failure();
SmallVector<FragmentAssemblyEntry> entries;
entries.reserve(rowCount);
for (int64_t n = 0; n < shape[0]; ++n)
for (int64_t c = 0; c < shape[1]; ++c)
for (int64_t h = 0; h < shape[2]; ++h)
entries.push_back({(n * shape[1] + c) * shape[2] + h, 0, {n, c, h, 0}, {1, 1, 1, shape[3]}});
return createFragmentAssemblyBlueprint(
batch->getResult(0), resultType, entries, "dense_nchw", spatial::kContiguousRowMajorFragments, rewriter, loc);
}
static FailureOr<Value> buildRowStripNearestResize(
Value storage, RankedTensorType inputType, RankedTensorType resultType,
PatternRewriter& rewriter, Location loc) {
auto input = describeRowStripPhysicalValue(storage, inputType);
if (failed(input))
return failure();
int64_t tilesPerRow = input->tilesPerRow;
int64_t outputHeight = resultType.getDimSize(2);
int64_t outputWidth = resultType.getDimSize(3);
int64_t tileChannels = input->fragmentType.getDimSize(3);
int64_t laneCount = outputHeight * tilesPerRow;
auto outputFragmentType = RankedTensorType::get(
{1, 1, outputWidth, tileChannels}, resultType.getElementType());
auto outputStorageType = spatial::getGraphBatchPhysicalResultType(
laneCount, outputFragmentType);
auto pixelType = RankedTensorType::get(
{1, 1, 1, tileChannels}, resultType.getElementType());
auto batch = createSpatComputeBatch(
rewriter, loc, TypeRange {outputStorageType}, laneCount, {}, ValueRange {storage},
[&](detail::SpatComputeBatchBodyArgs args) {
Operation* anchor = rewriter.getInsertionBlock()->getParentOp();
Value outputRow = affineFloorDivConst(rewriter, loc, args.lane, tilesPerRow, anchor);
Value tile = affineModConst(rewriter, loc, args.lane, tilesPerRow, anchor);
Value inputRow = buildNearestAsymmetricIndex(
outputRow, inputType.getDimSize(2), outputHeight, rewriter, loc);
Value inputSlot = arith::AddIOp::create(
rewriter, loc, affineMulConst(rewriter, loc, inputRow, tilesPerRow, anchor), tile);
auto source = extractGraphBatchPhysicalFragment(
rewriter, loc, args.inputs.front(), inputSlot, input->fragmentType);
if (failed(source))
return failure();
batchYielded.push_back(channelLoop->results.front());
Value initial = tensor::EmptyOp::create(
rewriter, loc, outputFragmentType.getShape(), resultType.getElementType());
Value c0 = getOrCreateIndexConstant(rewriter, anchor, 0);
Value c1 = getOrCreateIndexConstant(rewriter, anchor, 1);
Value width = getOrCreateIndexConstant(rewriter, anchor, outputWidth);
auto loop = buildNormalizedScfFor(
rewriter, loc, c0, width, c1, ValueRange {initial},
[&](OpBuilder&, Location nestedLoc, Value outputColumn, ValueRange iterArgs,
SmallVectorImpl<Value>& yielded) {
Value inputColumn = buildNearestAsymmetricIndex(
outputColumn, inputType.getDimSize(3), outputWidth, rewriter, nestedLoc);
Value pixel = tensor::ExtractSliceOp::create(
rewriter, nestedLoc, pixelType, *source,
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0),
inputColumn, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels)},
getUnitStrides(rewriter, 4));
yielded.push_back(tensor::InsertSliceOp::create(
rewriter, nestedLoc, pixel, iterArgs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0),
outputColumn, rewriter.getIndexAttr(0)},
SmallVector<OpFoldResult> {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1),
rewriter.getIndexAttr(1), rewriter.getIndexAttr(tileChannels)},
getUnitStrides(rewriter, 4)));
return success();
});
if (failed(loop))
return failure();
publishGraphBatchPhysicalFragment(
rewriter, loc, loop->results.front(), args.outputs.front(), args.lane);
return success();
});
if (failed(batchLoop))
return failure();
return batchLoop->results.front();
return failed(batch) ? FailureOr<Value>(failure())
: FailureOr<Value>(batch->getResult(0));
}
struct Resize : OpConversionPattern<ONNXResizeOp> {
@@ -161,23 +182,40 @@ struct Resize : OpConversionPattern<ONNXResizeOp> {
|| llvm::any_of(resultType.getShape(), [](int64_t dim) { return dim <= 0; }))
return rewriter.notifyMatchFailure(resizeOp, "resize lowering requires positive static dimensions.");
auto computeOp = createSpatCompute<1>(
rewriter, resizeOp.getLoc(), TypeRange {resultType}, {}, adaptor.getX(), [&](Value x) -> LogicalResult {
auto result = buildNearestResizeLoop(x, inputType, resultType, rewriter, resizeOp.getLoc());
if (failed(result))
return failure();
spatial::SpatYieldOp::create(rewriter, resizeOp.getLoc(), *result);
return success();
});
if (failed(computeOp))
return failure();
rewriter.replaceOp(resizeOp, computeOp->getResults());
auto plan = spatial::SpatResizeNearestPlanOp::create(
rewriter, resizeOp.getLoc(), resultType, adaptor.getX(), spatial::getNCHWLayout(rewriter.getContext()));
rewriter.replaceOp(resizeOp, plan.getResult());
return success();
}
};
} // namespace
LogicalResult canLowerResizeNearestPlanToRowStrip(
spatial::SpatResizeNearestPlanOp planOp,
const spatial::SpatialTargetInfo&) {
auto inputType = dyn_cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
return success(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));
}
FailureOr<Value> lowerSelectedResizeNearestPlan(
spatial::SpatResizeNearestPlanOp planOp, std::optional<Value> rowStripInput,
const spatial::SpatialTargetInfo&,
PatternRewriter& rewriter) {
auto inputType = cast<RankedTensorType>(planOp.getInput().getType());
auto outputType = cast<RankedTensorType>(planOp.getOutput().getType());
if (rowStripInput)
return buildRowStripNearestResize(
*rowStripInput, inputType, outputType, rewriter, planOp.getLoc());
return buildDenseNearestResize(
planOp.getInput(), inputType, outputType, rewriter, planOp.getLoc());
}
void populateResizePatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.add<Resize>(ctx); }
} // namespace onnx_mlir
@@ -61,6 +61,74 @@ static FailureOr<Value> materializeTransposedConstant(Value input,
resultType);
}
static FailureOr<Value> transposeFragmentAssemblyBlueprint(spatial::SpatBlueprintOp blueprint,
RankedTensorType resultType,
ArrayRef<int64_t> permutation,
ConversionPatternRewriter& rewriter,
Location loc) {
auto storageType = dyn_cast<RankedTensorType>(blueprint.getInput().getType());
auto sourceOffsets = blueprint.getFragmentSourceOffsets();
auto fragmentStrides = blueprint.getFragmentStrides();
if (!storageType || !storageType.hasStaticShape() || !resultType.hasStaticShape()
|| !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)
return failure();
if (blueprint.getIndexMap() == spatial::kContiguousRowMajorFragments
&& !spatial::isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
return blueprint.emitOpError("contiguous row-major fragment physical source order or storage is not canonical"), failure();
SmallVector<int64_t> outputStorageShape {storageType.getDimSize(0)};
for (int64_t sourceDim : permutation)
outputStorageShape.push_back(storageType.getDimSize(sourceDim + 1));
auto outputStorageType = RankedTensorType::get(outputStorageShape, storageType.getElementType());
auto mapped = mapGraphBatchFragments(
blueprint.getInput(), outputStorageType, rewriter, loc, [&](Value fragment, RankedTensorType fragmentType) {
Value init = createTransposeInit(fragment, fragmentType, permutation, rewriter, loc);
return FailureOr<Value>(
linalg::TransposeOp::create(rewriter, loc, fragment, init, permutation).getResult()[0]);
});
if (failed(mapped))
return failure();
const int64_t rank = resultType.getRank();
const int64_t fragmentCount = blueprint.getFragmentOperandIndices()->size();
SmallVector<int64_t> offsets, sizes, strides;
offsets.reserve(fragmentCount * rank);
sizes.reserve(fragmentCount * rank);
strides.reserve(fragmentCount * rank);
ArrayRef<int64_t> inputOffsets = blueprint.getFragmentOffsets();
ArrayRef<int64_t> inputSizes = blueprint.getFragmentSizes();
for (int64_t fragment = 0; fragment < fragmentCount; ++fragment)
for (int64_t sourceDim : permutation) {
const int64_t index = fragment * rank + sourceDim;
offsets.push_back(inputOffsets[index]);
sizes.push_back(inputSizes[index]);
strides.push_back((*fragmentStrides)[index]);
}
auto transposedBlueprint = spatial::SpatBlueprintOp::create(rewriter,
loc,
resultType,
*mapped,
ValueRange {},
blueprint.getLogicalLayoutAttr(),
spatial::getFragmentedLayout(rewriter.getContext()),
rewriter.getDenseI64ArrayAttr(offsets),
rewriter.getDenseI64ArrayAttr(sizes),
rewriter.getStringAttr("permuted_fragments"),
blueprint.getModeAttr(),
blueprint.getFragmentOperandIndicesAttr(),
blueprint.getFragmentSourceSlotsAttr(),
blueprint.getFragmentSourceOffsetsAttr(),
rewriter.getDenseI64ArrayAttr(strides),
blueprint.getConflictPolicyAttr(),
blueprint.getCoveragePolicyAttr());
if (spatial::isCanonicalContiguousRowMajorFragmentAssembly(transposedBlueprint))
transposedBlueprint.setIndexMapAttr(rewriter.getStringAttr(spatial::kContiguousRowMajorFragments));
return transposedBlueprint.getOutput();
}
struct TransposeToLinalgTranspose : OpConversionPattern<ONNXTransposeOp> {
using OpConversionPattern::OpConversionPattern;
@@ -75,6 +143,14 @@ struct TransposeToLinalgTranspose : OpConversionPattern<ONNXTransposeOp> {
auto permutation = getTransposePermutationChecked(transposeOp.getPermAttr(), inputType.getRank());
if (failed(permutation))
return failure();
if (auto blueprint = adaptor.getData().getDefiningOp<spatial::SpatBlueprintOp>()) {
auto transposed =
transposeFragmentAssemblyBlueprint(blueprint, resultType, *permutation, rewriter, transposeOp.getLoc());
if (succeeded(transposed)) {
rewriter.replaceOp(transposeOp, *transposed);
return success();
}
}
if (isCompileTimeComputable(adaptor.getData())) {
auto constantTranspose =
materializeTransposedConstant(adaptor.getData(), resultType, *permutation, rewriter, transposeOp.getLoc());
@@ -15,30 +15,50 @@ mlir::FailureOr<mlir::Value>
lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp,
std::optional<mlir::Value> 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 canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp);
mlir::LogicalResult canLowerResizeNearestPlanToRowStrip(
spatial::SpatResizeNearestPlanOp planOp, const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value> lowerSelectedResizeNearestPlan(
spatial::SpatResizeNearestPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value>
lowerDenseMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::FailureOr<mlir::Value>
lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::LogicalResult
canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp);
canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value>
lowerDenseGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::FailureOr<mlir::Value>
lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp,
std::optional<mlir::Value> 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
@@ -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<PhysicalLayout> operandLayouts) {
LayoutAlternative alternative;
alternative.operandLayouts.assign(operandLayouts.begin(), operandLayouts.end());
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
return alternative;
}
static bool hasRowStripInput(ArrayRef<PhysicalLayout> operandLayouts, unsigned index) {
return index < operandLayouts.size()
&& operandLayouts[index] == PhysicalLayout::NHWCRowStrip;
}
SmallVector<LayoutAlternative> SpatConv2DPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)) {
if (succeeded(canConsumeAndProduceRowStrip(*this, target)))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
}
else if (succeeded(canLowerConvPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatReluPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatSiluPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)) {
LayoutAlternative alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.intrinsicCost = -3;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatResizeNearestPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)
&& succeeded(canLowerResizeNearestPlanToRowStrip(*this, target)))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatMaxPool2DPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (succeeded(canLowerMaxPoolPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
if (hasRowStripInput(operandLayouts, 0))
alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatGlobalAveragePoolPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (succeeded(canLowerGlobalAveragePoolPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
if (hasRowStripInput(operandLayouts, 0))
alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatBiasAddPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
auto resultType = dyn_cast<RankedTensorType>(getOutput().getType());
if (resultType && hasRowStripInput(operandLayouts, 0)
&& isSupportedBiasAddValue(getBias(), resultType))
alternatives.push_back(rowStripAlternative(getOperation(),
{PhysicalLayout::NHWCRowStrip,
PhysicalLayout::DenseNCHW}));
return alternatives;
}
SmallVector<LayoutAlternative> SpatAddPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (operandLayouts.size() >= 2 && hasRowStripInput(operandLayouts, 0)
&& hasRowStripInput(operandLayouts, 1))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatConcatPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (!operandLayouts.empty() && llvm::all_of(operandLayouts, [](PhysicalLayout layout) {
return layout == PhysicalLayout::NHWCRowStrip;
}))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
} // namespace onnx_mlir::spatial
@@ -6,336 +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 <algorithm>
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<Value, spatial::PhysicalLayout>;
enum class SelectedLayout {
DenseNchw,
PixelMajorRowStrip,
};
static SelectedLayout getSelectedLayout(llvm::DenseMap<Value, SelectedLayout>& 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<spatial::SpatMaterializeLayoutOp>())
return materialize.getTargetPhysicalLayout();
if (auto blueprint = value.getDefiningOp<spatial::SpatBlueprintOp>())
return blueprint.getPhysicalLayout();
return spatial::PhysicalLayout::DenseNCHW;
}
static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap<Value, SelectedLayout>& layouts) {
if (auto reluPlan = dyn_cast<spatial::SpatReluPlanOp>(user))
return getSelectedLayout(layouts, reluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto siluPlan = dyn_cast<spatial::SpatSiluPlanOp>(user))
return getSelectedLayout(layouts, siluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto biasAddPlan = dyn_cast<spatial::SpatBiasAddPlanOp>(user))
return getSelectedLayout(layouts, biasAddPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto addPlan = dyn_cast<spatial::SpatAddPlanOp>(user))
return getSelectedLayout(layouts, addPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto concatPlan = dyn_cast<spatial::SpatConcatPlanOp>(user))
return getSelectedLayout(layouts, concatPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto convPlan = dyn_cast<spatial::SpatConv2DPlanOp>(user))
return getSelectedLayout(layouts, convPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(user))
return getSelectedLayout(layouts, maxPoolPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto averagePoolPlan = dyn_cast<spatial::SpatGlobalAveragePoolPlanOp>(user))
return getSelectedLayout(layouts, averagePoolPlan.getResult()) == SelectedLayout::PixelMajorRowStrip;
if (auto flattenCompute = dyn_cast<spatial::SpatGraphCompute>(user))
return succeeded(canLowerFlattenFromRowStrip(flattenCompute));
return false;
static SmallVector<spatial::PhysicalLayout> getOperandLayouts(
Operation* op, const LayoutMap& layouts) {
SmallVector<spatial::PhysicalLayout> 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<Value, SelectedLayout>& layouts) {
for (Operation* user : value.getUsers()) {
if (usesSelectedRowStrip(user, layouts))
static FailureOr<SmallVector<spatial::LayoutAlternative>> getAlternatives(
Operation* op, const LayoutMap& layouts, const spatial::SpatialTargetInfo& target) {
auto capability = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(op);
if (!capability)
return failure();
SmallVector<spatial::LayoutAlternative> 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<spatial::LayoutAlternative> 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<spatial::PhysicalLayout> 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<spatial::SpatialLayoutCapabilityInterface>(use.getOwner());
if (!user) {
if (alternative.resultLayout != spatial::PhysicalLayout::DenseNCHW) {
auto flatten = dyn_cast<spatial::SpatGraphCompute>(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<spatial::SpatReluPlanOp, spatial::SpatSiluPlanOp>(user))
return true;
if (auto biasAddPlan = dyn_cast<spatial::SpatBiasAddPlanOp>(user)) {
auto resultType = dyn_cast<RankedTensorType>(biasAddPlan.getOutput().getType());
return resultType && isSupportedBiasAddValue(biasAddPlan.getBias(), resultType);
}
if (isa<spatial::SpatAddPlanOp>(user))
return true;
if (isa<spatial::SpatConcatPlanOp>(user))
return true;
if (auto convPlan = dyn_cast<spatial::SpatConv2DPlanOp>(user))
return succeeded(canConsumeAndProduceRowStrip(convPlan));
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(user))
return succeeded(canLowerMaxPoolPlanToRowStrip(maxPoolPlan));
if (auto averagePoolPlan = dyn_cast<spatial::SpatGlobalAveragePoolPlanOp>(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<Value, SelectedLayout>& 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<Value, SelectedLayout>& 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<Value, SelectedLayout>& layouts) {
if (getSelectedLayout(layouts, input) != SelectedLayout::PixelMajorRowStrip)
return SelectedLayout::DenseNchw;
if (!allUsersCanHandleRowStrip(result, layouts))
return SelectedLayout::DenseNchw;
return SelectedLayout::PixelMajorRowStrip;
}
static SelectedLayout chooseBiasAddLayout(spatial::SpatBiasAddPlanOp biasAddPlan,
llvm::DenseMap<Value, SelectedLayout>& layouts) {
if (getSelectedLayout(layouts, biasAddPlan.getInput()) != SelectedLayout::PixelMajorRowStrip)
return SelectedLayout::DenseNchw;
auto resultType = dyn_cast<RankedTensorType>(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<Value, SelectedLayout>& 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<Value, SelectedLayout>& 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<RankedTensorType>(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<Value, SelectedLayout>& layouts) {
SmallVector<OpOperand*> 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<std::pair<OpOperand*, spatial::PhysicalLayout>> mismatches;
for (OpOperand& use : value.getUses()) {
Operation* userOp = use.getOwner();
spatial::PhysicalLayout required = spatial::PhysicalLayout::DenseNCHW;
if (auto capability = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(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<spatial::SpatGraphCompute>(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<SpatialLayoutPlanningPass, OperationPass<ModuleOp>> {
static LogicalResult verifySelectedLayouts(
ArrayRef<Operation*> 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<SpatialLayoutPlanningPass, OperationPass<ModuleOp>> {
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<Value, SelectedLayout> layouts;
SmallVector<Operation*> planOps;
for (Operation& op : funcOp.getBody().front())
if (isa<spatial::SpatialLayoutCapabilityInterface>(&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<spatial::SpatConv2DPlanOp>(&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<Operation*> 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<spatial::SpatReluPlanOp>(&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<spatial::SpatSiluPlanOp>(&op)) {
SelectedLayout selected = chooseActivationLayout(siluPlan.getInput(), siluPlan.getResult(), layouts);
if (layouts[siluPlan.getResult()] != selected) {
layouts[siluPlan.getResult()] = selected;
changed = true;
}
continue;
}
if (auto biasAddPlan = dyn_cast<spatial::SpatBiasAddPlanOp>(&op)) {
SelectedLayout selected = chooseBiasAddLayout(biasAddPlan, layouts);
if (layouts[biasAddPlan.getResult()] != selected) {
layouts[biasAddPlan.getResult()] = selected;
changed = true;
}
continue;
}
if (auto addPlan = dyn_cast<spatial::SpatAddPlanOp>(&op)) {
SelectedLayout selected = chooseAddLayout(addPlan, layouts);
if (layouts[addPlan.getResult()] != selected) {
layouts[addPlan.getResult()] = selected;
changed = true;
}
continue;
}
if (auto concatPlan = dyn_cast<spatial::SpatConcatPlanOp>(&op)) {
SelectedLayout selected = chooseConcatLayout(concatPlan, layouts);
if (layouts[concatPlan.getResult()] != selected) {
layouts[concatPlan.getResult()] = selected;
changed = true;
}
continue;
}
if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(&op)) {
SelectedLayout selected = chooseMaxPoolLayout(maxPoolPlan);
if (layouts[maxPoolPlan.getResult()] != selected) {
layouts[maxPoolPlan.getResult()] = selected;
changed = true;
}
continue;
}
if (auto averagePoolPlan = dyn_cast<spatial::SpatGlobalAveragePoolPlanOp>(&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<spatial::SpatConv2DPlanOp>(&op))
producedValue = convPlan.getResult();
else if (auto biasAddPlan = dyn_cast<spatial::SpatBiasAddPlanOp>(&op))
producedValue = biasAddPlan.getResult();
else if (auto addPlan = dyn_cast<spatial::SpatAddPlanOp>(&op))
producedValue = addPlan.getResult();
else if (auto concatPlan = dyn_cast<spatial::SpatConcatPlanOp>(&op))
producedValue = concatPlan.getResult();
else if (auto reluPlan = dyn_cast<spatial::SpatReluPlanOp>(&op))
producedValue = reluPlan.getResult();
else if (auto siluPlan = dyn_cast<spatial::SpatSiluPlanOp>(&op))
producedValue = siluPlan.getResult();
else if (auto maxPoolPlan = dyn_cast<spatial::SpatMaxPool2DPlanOp>(&op))
producedValue = maxPoolPlan.getResult();
else if (auto averagePoolPlan = dyn_cast<spatial::SpatGlobalAveragePoolPlanOp>(&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<Pass> createSpatialLayoutPlanningPass() { return std::make_unique<SpatialLayoutPlanningPass>(); }
std::unique_ptr<Pass> createSpatialLayoutPlanningPass() {
return std::make_unique<SpatialLayoutPlanningPass>();
}
std::unique_ptr<Pass> createSpatialLayoutPlanningPass(
const spatial::SpatialTargetInfo& target) {
return std::make_unique<SpatialLayoutPlanningPass>(target);
}
} // namespace onnx_mlir
@@ -149,11 +149,10 @@ collectTopLevelFragmentAssemblyCopies(OpResult result, RankedTensorType packedRe
auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(use.getOwner());
if (!blueprint || blueprint->getParentOp() != blueprint->getParentOfType<func::FuncOp>())
return failure();
std::optional<StringRef> mode = blueprint.getMode();
std::optional<ArrayRef<int64_t>> operandIndicesAttr = blueprint.getFragmentOperandIndices();
std::optional<ArrayRef<int64_t>> sourceOffsetsAttr = blueprint.getFragmentSourceOffsets();
std::optional<ArrayRef<int64_t>> 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<func::ReturnOp>(*blueprint.getOutput().getUsers().begin()))
return failure();
@@ -418,8 +417,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
rewriter.setInsertionPointToEnd(newBlock);
if (auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(op)) {
std::optional<StringRef> modeAttr = blueprint.getMode();
if (modeAttr && *modeAttr == "fragment_assembly") {
if (spatial::isFragmentAssembly(blueprint.getMode())) {
for (Operation* user : blueprint.getOutput().getUsers()) {
if (!isa<tensor::ParallelInsertSliceOp>(user))
return blueprint.emitOpError(
@@ -483,8 +481,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
auto hostTargetType = cast<ShapedType>(hostTarget.getType());
if (auto blueprint =
insertSlice.getSource().getDefiningOp<spatial::SpatBlueprintOp>()) {
std::optional<StringRef> modeAttr = blueprint.getMode();
if (modeAttr && *modeAttr == "fragment_assembly") {
if (spatial::isFragmentAssembly(blueprint.getMode())) {
FailureOr<SmallVector<FragmentAssemblyCopy, 8>> fragmentAssemblyCopies =
collectFragmentAssemblyCopiesFromBlueprint(blueprint, mapper, /*lane=*/0, /*hostTargetIndex=*/0);
if (failed(fragmentAssemblyCopies))
@@ -129,6 +129,32 @@ LogicalResult validateFragmentAssemblyMetadata(spatial::SpatBlueprintOp blueprin
return success();
}
FailureOr<mlir::Value> reshapeContiguousRowMajorFragments(RewriterBase& rewriter,
Location loc,
mlir::Value source,
RankedTensorType resultType) {
auto sourceType = dyn_cast<RankedTensorType>(source.getType());
if (!sourceType || !sourceType.hasStaticShape() || !resultType.hasStaticShape() || resultType.getRank() < 2
|| sourceType.getRank() != resultType.getRank() + 1 || sourceType.getElementType() != resultType.getElementType()
|| sourceType.getNumElements() != resultType.getNumElements()
|| sourceType.getDimSize(0) != getStaticShapeElementCount(resultType.getShape().drop_back())
|| sourceType.getDimSize(sourceType.getRank() - 1) != resultType.getDimSize(resultType.getRank() - 1)
|| llvm::any_of(sourceType.getShape().slice(1, sourceType.getRank() - 2), [](int64_t dim) { return dim != 1; }))
return failure();
SmallVector<ReassociationIndices> collapse {{}, {sourceType.getRank() - 1}};
for (int64_t dim = 0; dim < sourceType.getRank() - 1; ++dim)
collapse.front().push_back(dim);
auto flatType = RankedTensorType::get(
{sourceType.getDimSize(0), sourceType.getDimSize(sourceType.getRank() - 1)}, resultType.getElementType());
mlir::Value flat = tensor::CollapseShapeOp::create(rewriter, loc, flatType, source, collapse);
SmallVector<ReassociationIndices> expand {{}, {resultType.getRank() - 1}};
for (int64_t dim = 0; dim < resultType.getRank() - 1; ++dim)
expand.front().push_back(dim);
return tensor::ExpandShapeOp::create(rewriter, loc, resultType, flat, expand).getResult();
}
static SmallVector<int64_t, 4> expandFlatElementIndex(int64_t flatIndex, ArrayRef<int64_t> shape) {
SmallVector<int64_t, 4> indices(shape.size(), 0);
for (int64_t dim = static_cast<int64_t>(shape.size()) - 1; dim >= 0; --dim) {
@@ -51,6 +51,11 @@ mlir::LogicalResult validateFragmentAssemblyMetadata(onnx_mlir::spatial::SpatBlu
llvm::ArrayRef<int64_t> flatSizes,
llvm::ArrayRef<int64_t> flatStrides);
mlir::FailureOr<mlir::Value> reshapeContiguousRowMajorFragments(mlir::RewriterBase& rewriter,
mlir::Location loc,
mlir::Value source,
mlir::RankedTensorType resultType);
mlir::FailureOr<mlir::SmallVector<int64_t, 4>>
getStaticSliceOffsetsForElementOffset(mlir::Operation* anchor,
mlir::ShapedType sourceType,
@@ -42,12 +42,11 @@ static FailureOr<Value> lowerFragmentAssemblyBlueprint(IRRewriter& rewriter,
if (!resultType || !resultType.hasStaticShape())
return blueprint.emitOpError("fragment assembly lowering requires a static ranked tensor result");
std::optional<StringRef> modeAttr = blueprint.getMode();
std::optional<ArrayRef<int64_t>> operandIndicesAttr = blueprint.getFragmentOperandIndices();
std::optional<ArrayRef<int64_t>> sourceSlotsAttr = blueprint.getFragmentSourceSlots();
std::optional<ArrayRef<int64_t>> sourceOffsetsAttr = blueprint.getFragmentSourceOffsets();
std::optional<ArrayRef<int64_t>> 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");
@@ -71,6 +70,16 @@ static FailureOr<Value> lowerFragmentAssemblyBlueprint(IRRewriter& rewriter,
flatStrides)))
return failure();
if (blueprint.getIndexMap() == spatial::kContiguousRowMajorFragments) {
if (!spatial::isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
return blueprint.emitOpError("contiguous row-major fragment physical source order or storage is not canonical"), failure();
Value source = mapping.lookupOrDefault(blueprint.getInput());
auto reshaped = reshapeContiguousRowMajorFragments(
rewriter, blueprint.getLoc(), source, cast<RankedTensorType>(resultType));
if (failed(reshaped))
return blueprint.emitOpError("contiguous row-major fragment storage does not match its logical result"), failure();
return *reshaped;
}
SmallVector<int64_t> hostStrides = computeRowMajorStrides(resultType.getShape());
SmallVector<FragmentAssemblyCopy, 8> copies;
for (int64_t fragmentIndex = 0; fragmentIndex < static_cast<int64_t>(operandIndices.size()); ++fragmentIndex) {
@@ -193,8 +202,7 @@ static bool isHostMaterializableHelperOp(Operation* op) {
if (isa<arith::ConstantOp>(op) || op->hasTrait<OpTrait::ConstantLike>())
return true;
if (auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(op)) {
std::optional<StringRef> mode = blueprint.getMode();
return mode && *mode == "fragment_assembly";
return spatial::isFragmentAssembly(blueprint.getMode());
}
return isShapingOnlyOp(op) || isPureIndexComputationOp(op);
}
@@ -281,8 +289,7 @@ static bool inlineInputlessHelperComputeForWeightLikeUsers(spatial::SpatSchedule
}
for (Operation& op : block.without_terminator()) {
if (auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(op)) {
std::optional<StringRef> modeAttr = blueprint.getMode();
if (modeAttr && *modeAttr == "fragment_assembly") {
if (spatial::isFragmentAssembly(blueprint.getMode())) {
auto lowered = lowerFragmentAssemblyBlueprint(rewriter, blueprint, mapping);
if (failed(lowered))
return false;
+11 -2
View File
@@ -22,8 +22,7 @@ struct LowerFragmentAssemblyBlueprintPattern
LogicalResult matchAndRewrite(spatial::SpatBlueprintOp op,
OpAdaptor adaptor,
ConversionPatternRewriter& rewriter) const override {
std::optional<StringRef> modeAttr = op.getMode();
if (!modeAttr || *modeAttr != "fragment_assembly")
if (!spatial::isFragmentAssembly(op.getMode()))
return failure();
auto resultType = dyn_cast<ShapedType>(op.getOutput().getType());
@@ -49,6 +48,16 @@ struct LowerFragmentAssemblyBlueprintPattern
op, rank, fragmentOperands.size(), operandIndices, sourceOffsets, flatOffsets, flatSizes, flatStrides)))
return failure();
if (op.getIndexMap() == spatial::kContiguousRowMajorFragments) {
if (!spatial::isCanonicalContiguousRowMajorFragmentAssembly(op))
return op.emitOpError("contiguous row-major fragment physical source order or storage is not canonical");
auto reshaped = reshapeContiguousRowMajorFragments(
rewriter, op.getLoc(), adaptor.getInput(), cast<RankedTensorType>(resultType));
if (failed(reshaped))
return op.emitOpError("contiguous row-major fragment storage does not match its logical result");
rewriter.replaceOp(op, *reshaped);
return success();
}
Value currentOutput =
tensor::EmptyOp::create(rewriter, op.getLoc(), resultType.getShape(), resultType.getElementType()).getResult();
for (int64_t fragmentIndex = 0; fragmentIndex < static_cast<int64_t>(operandIndices.size()); ++fragmentIndex) {
@@ -158,8 +158,7 @@ analyzeTopLevelFragmentAssemblyUses(Value value) {
auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(use.getOwner());
if (!blueprint || blueprint->getParentOp() != blueprint->getParentOfType<func::FuncOp>())
return failure();
std::optional<StringRef> mode = blueprint.getMode();
if (!mode || *mode != "fragment_assembly")
if (!spatial::isFragmentAssembly(blueprint.getMode()))
return failure();
if (!blueprint.getOutput().hasOneUse() || !isa<func::ReturnOp>(*blueprint.getOutput().getUsers().begin()))
return failure();
@@ -819,8 +818,7 @@ void raptor::SpatialToPimPass::replaceReturnWithOutputBuffers(func::ReturnOp ret
}
if (auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(op)) {
std::optional<StringRef> 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);
@@ -33,6 +33,9 @@ using namespace pim;
namespace onnx_mlir {
static void annotateWeightsMemrefs(ModuleOp moduleOp, func::FuncOp funcOp);
static FailureOr<func::FuncOp> requirePimEntryFunc(ModuleOp moduleOp, StringRef phase);
namespace {
struct MemRefCopyWorkItem {
@@ -333,22 +336,6 @@ static LogicalResult verifyPimCopyEndpoints(Operation* copy,
return success(valid);
}
struct PimBufferizationPass : PassWrapper<PimBufferizationPass, OperationPass<ModuleOp>> {
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<OpOperand*> constantBackedRoots;
llvm::SmallPtrSet<OpOperand*, 8> 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<ModuleOp> clone = module.clone();
clone->walk([&](bufferization::AllocTensorOp alloc) {
if (alloc->getParentOfType<pim::PimCoreOp>()
|| alloc->getParentOfType<pim::PimCoreBatchOp>())
alloc->setAttr(kExistingAlloc, UnitAttr::get(module.getContext()));
});
auto options = baseOptions;
options.bufferizeFunctionBoundaries = false;
options.opFilter.allowOperation([](Operation* op) {
return isa<pim::PimCoreOp, pim::PimCoreBatchOp>(op)
|| op->getParentOfType<pim::PimCoreOp>()
|| op->getParentOfType<pim::PimCoreBatchOp>();
});
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<pim::PimCoreOp>()
&& !alloc->getParentOfType<pim::PimCoreBatchOp>()))
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<pim::PimCoreOp>()
|| op->getParentOfType<pim::PimCoreBatchOp>();
});
@@ -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<MemRefCopyWorkItem> copyWorklist;
llvm::SmallPtrSet<Operation*, 16> 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<memref::CopyOp>(&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<memref::CopyOp>(&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<Pass> createPimBufferizationPass() { return std::make_unique<PimBufferizationPass>(); }
static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp) {
forwardSingleConsumerReceiveCopies(funcOp);
forwardSingleConsumerContiguousInputCopies(funcOp);
forwardSingleConsumerPimOutputCopies(funcOp);
MLIRContext* ctx = moduleOp.getContext();
PatternRewriter rewriter(ctx);
SmallVector<MemRefCopyWorkItem> copyWorklist;
llvm::SmallPtrSet<Operation*, 16> 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<memref::CopyOp>(&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<memref::CopyOp>(&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<func::FuncOp> 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<PimBufferizationPreparationPass, OperationPass<ModuleOp>> {
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<PimOneShotBufferizationPass, OperationPass<ModuleOp>> {
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<PimMemoryNormalizationPass, OperationPass<ModuleOp>> {
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<TensorType>(value.getType())) {
op->emitOpError("tensor operand remains after PIM bufferization");
++failureCount;
return;
}
}
for (Value value : op->getResults()) {
if (isa<TensorType>(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<PimBufferizationVerificationPass, OperationPass<ModuleOp>> {
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<Pass> createPimBufferizationPreparationPass() {
return std::make_unique<PimBufferizationPreparationPass>();
}
std::unique_ptr<Pass> createPimOneShotBufferizationPass() {
return std::make_unique<PimOneShotBufferizationPass>();
}
std::unique_ptr<Pass> createPimMemoryNormalizationPass() {
return std::make_unique<PimMemoryNormalizationPass>();
}
std::unique_ptr<Pass> createPimBufferizationVerificationPass() {
return std::make_unique<PimBufferizationVerificationPass>();
}
} // namespace onnx_mlir
@@ -532,54 +532,74 @@ struct FoldConstantMemCpPattern final : OpRewritePattern<pim::PimMemCopyOp> {
}
};
static bool isOne(Attribute value) {
if (auto floatValue = dyn_cast<FloatAttr>(value))
return floatValue.getValue().isExactlyValue(1.0);
if (auto integerValue = dyn_cast<IntegerAttr>(value))
return integerValue.getValue() == 1;
return false;
enum class MultiplicationConstant { Other, Zero, One };
static MultiplicationConstant classifyMultiplicationConstant(Attribute value) {
if (auto floatValue = dyn_cast<FloatAttr>(value)) {
const APFloat& number = floatValue.getValue();
if (number.isZero() && !number.isNegative())
return MultiplicationConstant::Zero;
if (number.isExactlyValue(1.0))
return MultiplicationConstant::One;
}
if (auto integerValue = dyn_cast<IntegerAttr>(value)) {
if (integerValue.getValue().isZero())
return MultiplicationConstant::Zero;
if (integerValue.getValue() == 1)
return MultiplicationConstant::One;
}
return MultiplicationConstant::Other;
}
static bool isAllOneHostCopy(pim::PimMemCopyHostToDevOp copyOp, ModuleOp moduleOp, MemRefType copiedType) {
static MultiplicationConstant classifyUniformHostCopy(
pim::PimMemCopyHostToDevOp copyOp, ModuleOp moduleOp, MemRefType copiedType) {
auto targetOffset = resolveIndexValue(copyOp.getDeviceTargetOffset());
auto sourceOffset = resolveIndexValue(copyOp.getHostSourceOffset());
if (failed(targetOffset) || failed(sourceOffset) || *targetOffset != 0)
return false;
return MultiplicationConstant::Other;
Type elementType = copiedType.getElementType();
if (!elementType.isIntOrFloat())
return false;
return MultiplicationConstant::Other;
unsigned bitWidth = elementType.getIntOrFloatBitWidth();
if (bitWidth == 0 || bitWidth % 8 != 0)
return false;
return MultiplicationConstant::Other;
int64_t elementBytes = bitWidth / 8;
int64_t copiedElements = copiedType.getNumElements();
if (*sourceOffset % elementBytes != 0 || copyOp.getSize() != copiedElements * elementBytes)
return false;
return MultiplicationConstant::Other;
auto source = getDenseGlobalValue(moduleOp, copyOp.getHostSource());
if (failed(source) || source->getElementType() != elementType)
return false;
return MultiplicationConstant::Other;
int64_t firstElement = *sourceOffset / elementBytes;
int64_t endElement = firstElement + copiedElements;
if (firstElement < 0 || endElement > source->getNumElements())
return false;
return MultiplicationConstant::Other;
if (source->isSplat())
return isOne(source->getSplatValue<Attribute>());
return classifyMultiplicationConstant(source->getSplatValue<Attribute>());
MultiplicationConstant classification = MultiplicationConstant::Other;
int64_t index = 0;
for (Attribute value : source->getValues<Attribute>()) {
if (index >= firstElement && index < endElement && !isOne(value))
return false;
if (index >= firstElement && index < endElement) {
MultiplicationConstant current = classifyMultiplicationConstant(value);
if (current == MultiplicationConstant::Other)
return current;
if (classification == MultiplicationConstant::Other)
classification = current;
else if (classification != current)
return MultiplicationConstant::Other;
}
if (++index >= endElement)
break;
}
return true;
return classification;
}
struct FoldMultiplyByOnePattern final : OpRewritePattern<pim::PimVVMulOp> {
struct FoldMultiplyByConstantPattern final : OpRewritePattern<pim::PimVVMulOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(pim::PimVVMulOp mulOp, PatternRewriter& rewriter) const override {
@@ -605,14 +625,19 @@ struct FoldMultiplyByOnePattern final : OpRewritePattern<pim::PimVVMulOp> {
copyOp = candidate;
}
auto maskType = dyn_cast<MemRefType>(mask.getType());
if (!copyOp || !copyOp.use_empty() || !maskType || !isAllOneHostCopy(copyOp, moduleOp, maskType))
if (!copyOp || !copyOp.use_empty() || !maskType)
continue;
MultiplicationConstant constant = classifyUniformHostCopy(copyOp, moduleOp, maskType);
if (constant == MultiplicationConstant::Other)
continue;
auto outputAlloc = mulOp.getOutputBuffer().getDefiningOp<memref::AllocOp>();
rewriter.replaceOp(mulOp, input);
rewriter.eraseOp(copyOp);
if (maskAlloc.use_empty())
rewriter.eraseOp(maskAlloc);
rewriter.replaceOp(mulOp, constant == MultiplicationConstant::One ? input : mask);
if (constant == MultiplicationConstant::One) {
rewriter.eraseOp(copyOp);
if (maskAlloc.use_empty())
rewriter.eraseOp(maskAlloc);
}
if (outputAlloc && outputAlloc.use_empty())
rewriter.eraseOp(outputAlloc);
return success();
@@ -629,7 +654,7 @@ void populateConstantFoldingConstantPatterns(RewritePatternSet& patterns) {
FoldConstantCoreMapPattern,
FoldConstantHostCopyPattern,
FoldConstantMemCpPattern,
FoldMultiplyByOnePattern>(patterns.getContext());
FoldMultiplyByConstantPattern>(patterns.getContext());
}
} // namespace onnx_mlir
+13 -1
View File
@@ -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
+94 -22
View File
@@ -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<SpatialDialect, SpatLogicalLayout, "logical_layout"> {
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<SpatialDialect, SpatPhysicalLayout, "physical_layout"> {
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<SpatialDialect, SpatBlueprintMode, "blueprint_mode"> {
let assemblyFormat = "$value";
}
class SpatOp<string mnemonic, list<Trait> traits = []> :
Op<SpatialDialect, mnemonic, traits>;
class SpatLayoutPlanOp<string mnemonic> : SpatOp<mnemonic,
[SpatialLayoutCapabilityInterface,
DeclareOpInterfaceMethods<SpatialLayoutCapabilityInterface>]>;
// 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,7 +360,22 @@ def SpatSiluPlanOp : SpatOp<"silu_plan", []> {
let hasVerifier = 1;
}
def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> {
def SpatResizeNearestPlanOp : SpatLayoutPlanOp<"resize_nearest_plan"> {
let summary = "Layout-aware nearest asymmetric Resize planning op";
let arguments = (ins
SpatTensor:$input,
SpatLogicalLayoutAttr:$logicalLayout
);
let results = (outs
SpatTensor:$output
);
let hasVerifier = 1;
}
def SpatMaxPool2DPlanOp : SpatLayoutPlanOp<"max_pool2d_plan"> {
let summary = "Layout-aware 2D NCHW MaxPool planning op";
let arguments = (ins
@@ -312,7 +384,7 @@ def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> {
DenseI64ArrayAttr:$pads,
DenseI64ArrayAttr:$strides,
DenseI64ArrayAttr:$dilations,
StrAttr:$logicalLayout
SpatLogicalLayoutAttr:$logicalLayout
);
let results = (outs
@@ -322,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
@@ -337,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
@@ -353,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
@@ -369,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<SpatTensor>:$inputs,
I64Attr:$axis,
StrAttr:$logicalLayout
SpatLogicalLayoutAttr:$logicalLayout
);
let results = (outs
@@ -391,12 +463,12 @@ def SpatBlueprintOp : SpatOp<"blueprint", []> {
let arguments = (ins
SpatTensor:$input,
Variadic<SpatTensor>:$fragments,
StrAttr:$logicalLayout,
StrAttr:$physicalLayout,
SpatLogicalLayoutAttr:$logicalLayout,
SpatPhysicalLayoutAttr:$physicalLayout,
DenseI64ArrayAttr:$fragmentOffsets,
DenseI64ArrayAttr:$fragmentSizes,
StrAttr:$indexMap,
OptionalAttr<StrAttr>:$mode,
OptionalAttr<SpatBlueprintModeAttr>:$mode,
OptionalAttr<DenseI64ArrayAttr>:$fragmentOperandIndices,
OptionalAttr<DenseI64ArrayAttr>:$fragmentSourceSlots,
OptionalAttr<DenseI64ArrayAttr>:$fragmentSourceOffsets,
@@ -418,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
@@ -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
+98
View File
@@ -4,6 +4,8 @@
#include <string>
#include "mlir/IR/DialectImplementation.h"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
using namespace mlir;
@@ -11,6 +13,70 @@ using namespace mlir;
namespace onnx_mlir {
namespace spatial {
bool hasCanonicalContiguousRowMajorFragments(RankedTensorType logicalType,
ArrayRef<int64_t> offsets,
ArrayRef<int64_t> sizes,
ArrayRef<int64_t> strides) {
if (!logicalType || !logicalType.hasStaticShape() || logicalType.getRank() <= 0
|| logicalType.getDimSize(logicalType.getRank() - 1) <= 0)
return false;
const int64_t rank = logicalType.getRank();
const int64_t rowCount = logicalType.getNumElements() / logicalType.getDimSize(rank - 1);
if (offsets.size() != static_cast<size_t>(rowCount * rank) || sizes.size() != offsets.size()
|| strides.size() != offsets.size())
return false;
for (int64_t row = 0; row < rowCount; ++row) {
int64_t remaining = row;
for (int64_t dim = rank - 2; dim >= 0; --dim) {
const int64_t index = row * rank + dim;
if (offsets[index] != remaining % logicalType.getDimSize(dim) || sizes[index] != 1 || strides[index] != 1)
return false;
remaining /= logicalType.getDimSize(dim);
}
const int64_t last = row * rank + rank - 1;
if (offsets[last] != 0 || sizes[last] != logicalType.getDimSize(rank - 1) || strides[last] != 1)
return false;
}
return true;
}
bool isCanonicalContiguousRowMajorFragmentAssembly(SpatBlueprintOp blueprint) {
auto logicalType = dyn_cast<RankedTensorType>(blueprint.getOutput().getType());
auto physicalType = dyn_cast<RankedTensorType>(blueprint.getInput().getType());
auto operandIndices = blueprint.getFragmentOperandIndices();
auto sourceSlots = blueprint.getFragmentSourceSlots();
auto sourceOffsets = blueprint.getFragmentSourceOffsets();
auto fragmentStrides = blueprint.getFragmentStrides();
if (!logicalType || !physicalType || !logicalType.hasStaticShape() || !physicalType.hasStaticShape()
|| logicalType.getRank() < 2 || !blueprint.getFragments().empty()
|| !isFragmentAssembly(blueprint.getMode()) || !operandIndices || !sourceSlots || !sourceOffsets
|| !fragmentStrides)
return false;
ArrayRef<int64_t> offsets = blueprint.getFragmentOffsets();
ArrayRef<int64_t> sizes = blueprint.getFragmentSizes();
if (!hasCanonicalContiguousRowMajorFragments(logicalType, offsets, sizes, *fragmentStrides)
|| operandIndices->empty() || operandIndices->size() != sourceSlots->size()
|| operandIndices->size() != sourceOffsets->size()
|| operandIndices->size() * static_cast<size_t>(logicalType.getRank()) != offsets.size()
|| physicalType.getRank() != logicalType.getRank() + 1
|| physicalType.getDimSize(0) != static_cast<int64_t>(operandIndices->size())
|| physicalType.getDimSize(0)
!= logicalType.getNumElements() / logicalType.getDimSize(logicalType.getRank() - 1)
|| physicalType.getElementType() != logicalType.getElementType()
|| physicalType.getNumElements() != logicalType.getNumElements()
|| physicalType.getDimSize(physicalType.getRank() - 1) != logicalType.getDimSize(logicalType.getRank() - 1)
|| llvm::any_of(physicalType.getShape().slice(1, physicalType.getRank() - 2),
[](int64_t dim) { return dim != 1; }))
return false;
for (auto [fragmentIndex, operandIndex] : llvm::enumerate(*operandIndices))
if (operandIndex != 0 || (*sourceSlots)[fragmentIndex] != static_cast<int64_t>(fragmentIndex)
|| (*sourceOffsets)[fragmentIndex] != 0)
return false;
return true;
}
RankedTensorType getGraphBatchPhysicalResultType(int64_t laneCount, RankedTensorType fragmentType) {
SmallVector<int64_t> shape {laneCount};
llvm::append_range(shape, fragmentType.getShape());
@@ -374,6 +440,11 @@ OpResult SpatInParallelOp::getParentResult(int64_t idx) {
llvm::iterator_range<Block::iterator> 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"
@@ -395,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"
+73
View File
@@ -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 <map>
#include <optional>
#include <string>
#include <tuple>
#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<PhysicalLayout> 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"
@@ -30,6 +52,57 @@
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<BlueprintMode> mode) {
return mode && *mode == BlueprintMode::PhysicalView;
}
inline bool isFragmentAssembly(std::optional<BlueprintMode> mode) {
return mode && *mode == BlueprintMode::FragmentAssembly;
}
inline std::optional<PhysicalLayout> getSelectedPhysicalLayout(mlir::Operation* op) {
auto attr = op->getAttrOfType<PhysicalLayoutAttr>(kSelectedLayoutAttrName);
return attr ? std::optional<PhysicalLayout>(attr.getValue()) : std::nullopt;
}
bool hasCanonicalContiguousRowMajorFragments(mlir::RankedTensorType logicalType,
llvm::ArrayRef<int64_t> offsets,
llvm::ArrayRef<int64_t> sizes,
llvm::ArrayRef<int64_t> strides);
bool isCanonicalContiguousRowMajorFragmentAssembly(SpatBlueprintOp blueprint);
mlir::RankedTensorType getGraphBatchPhysicalResultType(int64_t laneCount, mlir::RankedTensorType fragmentType);
mlir::FailureOr<mlir::RankedTensorType>
getGraphBatchFragmentType(mlir::RankedTensorType physicalType, int64_t expectedLaneCount);
+16 -5
View File
@@ -616,7 +616,7 @@ void SpatBlueprintOp::print(OpAsmPrinter& printer) {
printer << " sizes ";
printCompressedIntegerList(printer, getFragmentSizes());
printer << " map " << getIndexMap();
if (std::optional<StringRef> mode = getMode())
if (auto mode = getMode())
printer << " mode " << *mode;
if (std::optional<ArrayRef<int64_t>> 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<BlueprintMode> 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())
+72 -24
View File
@@ -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<RankedTensorType>(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) {
@@ -480,6 +479,23 @@ LogicalResult SpatSiluPlanOp::verify() {
return success();
}
LogicalResult SpatResizeNearestPlanOp::verify() {
if (failed(verifyPlanTensorTypes(
getOperation(), getInput(), getOutput(), "spat.resize_nearest_plan")))
return failure();
auto inputType = dyn_cast<RankedTensorType>(getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(getOutput().getType());
if (!inputType.hasStaticShape() || !outputType.hasStaticShape()
|| inputType.getRank() != 4 || outputType.getRank() != 4)
return emitError("requires static rank-4 input and output tensors");
if (getLogicalLayout() != 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; }))
return emitError("requires positive dimensions");
return success();
}
LogicalResult SpatMaxPool2DPlanOp::verify() {
if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.max_pool2d_plan")))
return failure();
@@ -488,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");
@@ -509,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)
@@ -535,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");
@@ -563,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<RankedTensorType>(getInput().getType());
auto outputType = dyn_cast<RankedTensorType>(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()))
@@ -587,8 +620,12 @@ LogicalResult SpatBlueprintOp::verify() {
if (offsets.size() != sizes.size())
return emitError("fragment offset and size arrays must have the same length");
int64_t rank = logicalType.getRank();
if (offsets.empty())
bool isContiguousRowMajor = getIndexMap() == kContiguousRowMajorFragments;
if (offsets.empty()) {
if (isContiguousRowMajor)
return emitError("contiguous row-major fragment destination geometry is not canonical");
return success();
}
if (rank <= 0 || offsets.size() % rank != 0)
return emitError("fragment metadata must be a whole number of rank-sized fragments");
@@ -611,6 +648,8 @@ LogicalResult SpatBlueprintOp::verify() {
};
if (!isFragmentAssembly) {
if (isContiguousRowMajor)
return emitError("contiguous row-major fragments require fragment assembly metadata");
if (failed(verifyBoundsOnly({})))
return failure();
if (!getFragments().empty())
@@ -662,6 +701,13 @@ LogicalResult SpatBlueprintOp::verify() {
if (failed(verifyBoundsOnly(strides)))
return failure();
if (isContiguousRowMajor) {
if (!hasCanonicalContiguousRowMajorFragments(logicalType, offsets, sizes, strides))
return emitError("contiguous row-major fragment destination geometry is not canonical");
if (!isCanonicalContiguousRowMajorFragmentAssembly(*this))
return emitError("contiguous row-major fragment physical source order or storage is not canonical");
}
SmallVector<std::pair<SmallVector<int64_t, 4>, SmallVector<int64_t, 4>>, 8> slices;
slices.reserve(static_cast<size_t>(fragmentCount));
SmallVector<int64_t, 8> fragmentCountsByOperand(static_cast<size_t>(operandCount), 0);
@@ -713,20 +759,22 @@ LogicalResult SpatBlueprintOp::verify() {
if (sourceSliceOffsets[dim] + fragmentSizes[dim] > fragmentType.getDimSize(dim))
return emitError("fragment assembly source offset must describe a valid unit-stride slice");
for (const auto& [existingOffsets, existingSizes] : slices) {
bool overlaps = true;
for (int64_t dim = 0; dim < rank; ++dim) {
int64_t begin = fragmentOffsets[dim];
int64_t end = begin + fragmentSizes[dim];
int64_t existingBegin = existingOffsets[dim];
int64_t existingEnd = existingBegin + existingSizes[dim];
if (end <= existingBegin || existingEnd <= begin) {
overlaps = false;
break;
if (!isContiguousRowMajor) {
for (const auto& [existingOffsets, existingSizes] : slices) {
bool overlaps = true;
for (int64_t dim = 0; dim < rank; ++dim) {
int64_t begin = fragmentOffsets[dim];
int64_t end = begin + fragmentSizes[dim];
int64_t existingBegin = existingOffsets[dim];
int64_t existingEnd = existingBegin + existingSizes[dim];
if (end <= existingBegin || existingEnd <= begin) {
overlaps = false;
break;
}
}
if (overlaps)
return emitError("fragment assembly blueprint requires disjoint static slices");
}
if (overlaps)
return emitError("fragment assembly blueprint requires disjoint static slices");
}
slices.push_back({std::move(fragmentOffsets), std::move(fragmentSizes)});
}
@@ -0,0 +1,37 @@
#pragma once
#include <cstddef>
#include <cstdint>
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
@@ -9,6 +9,7 @@
#include "llvm/ADT/SmallPtrSet.h"
#include "src/Accelerators/PIM/Common/IR/AffineUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/ShapingUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/StaticIntSequence.hpp"
#include "src/Accelerators/PIM/Common/IR/TensorSliceUtils.hpp"
@@ -37,6 +38,63 @@ static SmallVector<Value> getBlueprintFragments(SpatBlueprintOp blueprint) {
return fragments;
}
static FailureOr<Value> buildContiguousRowMajorReconstruction(
OpBuilder &builder, Location loc, SpatBlueprintOp blueprint,
Value source) {
auto resultType = dyn_cast<RankedTensorType>(blueprint.getOutput().getType());
auto sourceType = dyn_cast<RankedTensorType>(source.getType());
if (!resultType || !sourceType || !resultType.hasStaticShape()
|| !sourceType.hasStaticShape() || resultType.getRank() <= 0
|| sourceType.getRank() != resultType.getRank() + 1)
return failure();
int64_t rank = resultType.getRank();
int64_t width = resultType.getDimSize(rank - 1);
int64_t rowCount = resultType.getNumElements() / width;
if (sourceType.getDimSize(0) != rowCount || sourceType.getDimSize(rank) != width)
return failure();
auto flatType = RankedTensorType::get({rowCount, width}, resultType.getElementType());
auto rowType = RankedTensorType::get({1, width}, resultType.getElementType());
auto physicalRowType = RankedTensorType::get(sourceType.getShape().drop_front(), resultType.getElementType());
Value init = tensor::EmptyOp::create(builder, loc, flatType.getShape(), flatType.getElementType());
Value c0 = arith::ConstantIndexOp::create(builder, loc, 0);
Value c1 = arith::ConstantIndexOp::create(builder, loc, 1);
Value rows = arith::ConstantIndexOp::create(builder, loc, rowCount);
auto loop = buildNormalizedScfFor(
builder, loc, c0, rows, c1, ValueRange {init},
[&](OpBuilder &nested, Location nestedLoc, Value row, ValueRange iterArgs,
SmallVectorImpl<Value> &yielded) {
SmallVector<OpFoldResult> offsets {row};
SmallVector<OpFoldResult> sizes {nested.getIndexAttr(1)};
SmallVector<OpFoldResult> strides {nested.getIndexAttr(1)};
for (int64_t dim : physicalRowType.getShape()) {
offsets.push_back(nested.getIndexAttr(0));
sizes.push_back(nested.getIndexAttr(dim));
strides.push_back(nested.getIndexAttr(1));
}
Value physicalRow = tensor::ExtractSliceOp::create(
nested, nestedLoc, physicalRowType, source, offsets, sizes, strides);
SmallVector<ReassociationIndices> collapse {{}};
for (int64_t dim = 0; dim < rank - 1; ++dim)
collapse.front().push_back(dim);
collapse.push_back({rank - 1});
Value flatRow = tensor::CollapseShapeOp::create(
nested, nestedLoc, rowType, physicalRow, collapse);
yielded.push_back(tensor::InsertSliceOp::create(
nested, nestedLoc, flatRow, iterArgs.front(),
SmallVector<OpFoldResult> {row, nested.getIndexAttr(0)},
SmallVector<OpFoldResult> {nested.getIndexAttr(1), nested.getIndexAttr(width)},
SmallVector<OpFoldResult> {nested.getIndexAttr(1), nested.getIndexAttr(1)}));
return success();
});
if (failed(loop))
return failure();
SmallVector<ReassociationIndices> expand {{}, {rank - 1}};
for (int64_t dim = 0; dim < rank - 1; ++dim)
expand.front().push_back(dim);
return tensor::ExpandShapeOp::create(builder, loc, resultType, loop->results.front(), expand).getResult();
}
static FailureOr<Value> buildBlueprintReconstruction(
OpBuilder &builder, Location loc, SpatBlueprintOp blueprint,
ValueRange sourceBlockArgs) {
@@ -57,6 +115,13 @@ static FailureOr<Value> buildBlueprintReconstruction(
sourceOffsets->size() != operandIndices->size())
return blueprint.emitOpError("phase 1 fragment assembly metadata has inconsistent sizes"), failure();
if (blueprint.getIndexMap() == kContiguousRowMajorFragments) {
if (!isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
return blueprint.emitOpError("contiguous row-major fragment physical source order or storage is not canonical"), failure();
if (sourceBlockArgs.size() != 1)
return blueprint.emitOpError("contiguous row-major fragment reconstruction requires one physical source"), failure();
return buildContiguousRowMajorReconstruction(builder, loc, blueprint, sourceBlockArgs.front());
}
Value result = tensor::EmptyOp::create(builder, loc, resultType.getShape(),
resultType.getElementType());
for (auto [fragmentIndex, operandIndex] : llvm::enumerate(*operandIndices)) {
@@ -178,6 +243,15 @@ static Operation *getTopLevelDeferredOperation(
return op && isTopLevelDeferredOperation(op, body, plan) ? op : nullptr;
}
static bool isDefinedInside(Operation *owner, Value value) {
if (Operation *definition = value.getDefiningOp())
return owner->isProperAncestor(definition);
auto argument = dyn_cast<BlockArgument>(value);
Region *region = argument ? argument.getOwner()->getParent() : nullptr;
return region == &owner->getRegion(0)
|| (region && owner->getRegion(0).isAncestor(region));
}
static bool isEligible(Value value, Block &body, const DeferredInputPlan &plan,
llvm::SmallPtrSetImpl<Operation *> &seen) {
if (value == plan.graphInput || value == plan.graphLane || value == plan.scheduledLane)
@@ -195,18 +269,10 @@ static bool isEligible(Value value, Block &body, const DeferredInputPlan &plan,
loop.getRegion().walk([&](Operation *nested) {
if (isa<scf::ForOp>(nested) && nested != loop)
eligible = false;
for (Value operand : nested->getOperands()) {
Operation *definition = operand.getDefiningOp();
auto argument = dyn_cast<BlockArgument>(operand);
Region *argumentRegion = argument
? argument.getOwner()->getParent() : nullptr;
bool definedInside = definition
? loop->isProperAncestor(definition)
: argumentRegion == &loop.getRegion()
|| (argumentRegion && loop.getRegion().isAncestor(argumentRegion));
if (!definedInside && !isEligible(operand, body, plan, seen))
for (Value operand : nested->getOperands())
if (!isDefinedInside(loop, operand)
&& !isEligible(operand, body, plan, seen))
eligible = false;
}
});
if (!eligible)
return false;
@@ -267,19 +333,9 @@ static FailureOr<Value> clonePayloadRoot(Value root, Block &body, const Deferred
if (auto loop = dyn_cast<scf::ForOp>(op)) {
SmallVector<Value> captures;
loop.getRegion().walk([&](Operation *nested) {
for (Value operand : nested->getOperands()) {
Operation *definition = operand.getDefiningOp();
auto argument = dyn_cast<BlockArgument>(operand);
Region *argumentRegion = argument
? argument.getOwner()->getParent() : nullptr;
bool definedInside = definition
? loop->isProperAncestor(definition)
: argumentRegion == &loop.getRegion()
|| (argumentRegion
&& loop.getRegion().isAncestor(argumentRegion));
if (!definedInside && !mapping.contains(operand))
for (Value operand : nested->getOperands())
if (!isDefinedInside(loop, operand) && !mapping.contains(operand))
captures.push_back(operand);
}
});
for (Value capture : captures)
if (!mapping.contains(capture) && failed(clone(capture)))
@@ -304,7 +360,10 @@ static bool dependsOnGraphLane(Value value, Value graphLane, Block &body,
if (auto loop = dyn_cast<scf::ForOp>(op)) {
bool depends = false;
loop.getRegion().walk([&](Operation *nested) {
depends |= llvm::is_contained(nested->getOperands(), graphLane);
for (Value operand : nested->getOperands())
if (!isDefinedInside(loop, operand)
&& dependsOnGraphLane(operand, graphLane, body, plan, seen))
depends = true;
});
if (depends)
return true;
@@ -328,7 +387,7 @@ static void collectClosure(Value value, Block &body, const DeferredInputPlan &pl
bool isDeferredFragmentAssemblyInput(Value input, size_t processorCount) {
auto blueprint = input.getDefiningOp<SpatBlueprintOp>();
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();
@@ -14,6 +14,26 @@ namespace onnx_mlir::spatial {
using namespace mlir;
namespace {
static LogicalResult verifyNoEscapingRegionValues(Operation* owner, StringRef phase) {
Operation* escapingDefinition = nullptr;
Operation* escapingUser = nullptr;
owner->walk([&](Operation* nested) {
for (Value result : nested->getResults())
for (Operation* user : result.getUsers())
if (!owner->isProperAncestor(user)) {
escapingDefinition = nested;
escapingUser = user;
return WalkResult::interrupt();
}
return WalkResult::advance();
});
if (!escapingDefinition)
return success();
return owner->emitOpError() << phase << " left a value defined by " << escapingDefinition->getName()
<< " at " << escapingDefinition->getLoc() << " captured by "
<< escapingUser->getName() << " at " << escapingUser->getLoc();
}
static LogicalResult placeLogicalProcessorsOnPhysicalCores(DeferredTransferPlan& plan, const SchedulingTarget& target) {
std::vector<Cost> logicalTrafficFlits(target.processorCount * target.processorCount, 0);
for (const std::unique_ptr<DeferredExchangePlan>& exchange : plan.exchanges)
@@ -138,6 +158,8 @@ static LogicalResult eraseOldGraph(func::FuncOp funcOp, IRRewriter& rewriter) {
}
}
}
if (failed(verifyNoEscapingRegionValues(op, "phase 2")))
return failure();
rewriter.eraseOp(op);
}
return success();
@@ -217,6 +239,8 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
op->getResult(0).replaceAllUsesWith(replacement);
if (!op->use_empty())
return op->emitOpError("phase 2 cannot erase deferred communication with live uses");
if (failed(verifyNoEscapingRegionValues(op, "phase 2 deferred communication")))
return failure();
rewriter.eraseOp(op);
}
if (failed(eraseDeferredSourceSelectors(funcOp, rewriter)) || failed(eraseOldGraph(funcOp, rewriter))
@@ -180,13 +180,18 @@ static bool originatesFromDeferredSource(
return originatesFromDeferredSource(value, deferred, visited);
}
static bool isInsideDeferredLoop(Operation *op,
SpatDeferredCommunicationOp deferred) {
static scf::ForOp getEnclosingDeferredLoop(
Operation *op, SpatDeferredCommunicationOp deferred) {
for (Operation *parent = op->getParentOp(); parent && parent != deferred;
parent = parent->getParentOp())
if (isa<scf::ForOp>(parent))
return true;
return false;
if (auto loop = dyn_cast<scf::ForOp>(parent))
return loop;
return {};
}
static bool isInsideDeferredLoop(
Operation *op, SpatDeferredCommunicationOp deferred) {
return static_cast<bool>(getEnclosingDeferredLoop(op, deferred));
}
static FailureOr<unsigned> getLoopIterationCount(
@@ -297,7 +302,7 @@ static LogicalResult validateDeferredProgram(
&& llvm::any_of(op->getOperands(), [&](Value operand) {
return originatesFromDeferredSource(operand, deferred);
})) {
auto loop = op->getParentOfType<scf::ForOp>();
auto loop = getEnclosingDeferredLoop(op, deferred);
scf::ForOp outerLoop;
for (Operation *parent = loop ? loop->getParentOp() : nullptr;
parent && parent != deferred; parent = parent->getParentOp())
@@ -578,7 +583,7 @@ FailureOr<DeferredProgramTemplate> analyzeDeferredProgramTemplate(
SmallVector<OpFoldResult>(
ArrayRef(slice.getMixedStrides()).drop_front())};
leaf.reconstructedType = cast<RankedTensorType>(value.getType());
leaf.enclosingLoop = slice->getParentOfType<scf::ForOp>();
leaf.enclosingLoop = getEnclosingDeferredLoop(slice, deferred);
if (graphProjection
&& slice.getSourceType().getRank()
== leaf.reconstructedType.getRank() + 1
@@ -609,8 +614,11 @@ FailureOr<DeferredProgramTemplate> analyzeDeferredProgramTemplate(
program.leaves.push_back(std::move(leaf));
return success();
}
if (value.getType().isIndex() || isa<IntegerType>(value.getType()))
return success();
if (value.getType().isIndex() || isa<IntegerType>(value.getType())) {
Operation *definition = value.getDefiningOp();
if (!definition || !deferred->isProperAncestor(definition))
return success();
}
if (auto argument = dyn_cast<BlockArgument>(value)) {
auto loop = dyn_cast_or_null<scf::ForOp>(
argument.getOwner()->getParentOp());
@@ -619,7 +627,7 @@ FailureOr<DeferredProgramTemplate> analyzeDeferredProgramTemplate(
}
Operation *op = value.getDefiningOp();
if (!op || (op->getBlock() != &body
&& !op->getParentOfType<scf::ForOp>()))
&& !getEnclosingDeferredLoop(op, deferred)))
return deferred.emitOpError(
"deferred residual escapes its verified body: ") << value;
if (auto loop = dyn_cast<scf::ForOp>(op)) {
@@ -289,7 +289,9 @@ static Value cloneResidual(
mapping.map(oldValue, newValue);
}
for (Operation *op : exchange.program.residualOps) {
if (op->hasTrait<OpTrait::ConstantLike>())
if (op->hasTrait<OpTrait::ConstantLike>()
|| llvm::all_of(op->getResults(),
[&](Value result) { return mapping.contains(result); }))
continue;
if (auto oldLoop = dyn_cast<scf::ForOp>(op)) {
SmallVector<Value> initArgs;
@@ -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) {
@@ -454,6 +454,9 @@ retargetBlueprint(DeferredTransferPlan& plan, SpatBlueprintOp blueprint, GraphBa
OpBuilder builder(blueprint);
blueprint->setAttr("fragmentOperandIndices", builder.getDenseI64ArrayAttr(newOperands));
blueprint->setAttr("fragmentSourceSlots", builder.getDenseI64ArrayAttr(newSlots));
if (blueprint.getIndexMap() == kContiguousRowMajorFragments
&& !isCanonicalContiguousRowMajorFragmentAssembly(blueprint))
blueprint.setIndexMapAttr(builder.getStringAttr("fragment_assembly"));
return success();
}
@@ -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<MergeComputeNodesPass, OperationPass<ModuleOp>> {
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<ScheduledComputeMaterializationResult> 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<Pass> createMergeComputeNodesPass() { return std::make_unique<spatial::MergeComputeNodesPass>(); }
std::unique_ptr<Pass> createMergeComputeNodesPass(const spatial::SchedulingTarget& target) {
return std::make_unique<spatial::MergeComputeNodesPass>(target);
}
} // namespace onnx_mlir
@@ -24,7 +24,7 @@ bool requiresScheduledPublication(Value value, DenseSet<Value> &visited) {
SpatDeferredCommunicationOp>(user))
return false;
auto blueprint = dyn_cast<SpatBlueprintOp>(user);
return !blueprint || blueprint.getMode() != "fragment_assembly"
return !blueprint || !isFragmentAssembly(blueprint.getMode())
|| requiresScheduledPublication(blueprint.getOutput(), visited);
});
}
@@ -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<func::FuncOp> 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<ScheduledSpatialState>& 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<ScheduleSpatialGraphPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ScheduleSpatialGraphPass)
ScheduleSpatialGraphPass() = default;
ScheduleSpatialGraphPass(const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> 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<ScheduledComputeMaterializationResult> 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<ScheduledSpatialState> state;
bool hasTarget = false;
};
struct VerifyScheduledSpatialPass final
: PassWrapper<VerifyScheduledSpatialPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(VerifyScheduledSpatialPass)
explicit VerifyScheduledSpatialPass(std::shared_ptr<ScheduledSpatialState> 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<ScheduledSpatialState> state;
};
struct RealizeSpatialCommunicationPass final
: PassWrapper<RealizeSpatialCommunicationPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(RealizeSpatialCommunicationPass)
RealizeSpatialCommunicationPass() = default;
RealizeSpatialCommunicationPass(const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> 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<ScheduledSpatialState> state;
bool hasTarget = false;
};
struct VerifyRealizedSpatialPass final
: PassWrapper<VerifyRealizedSpatialPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(VerifyRealizedSpatialPass)
explicit VerifyRealizedSpatialPass(std::shared_ptr<ScheduledSpatialState> 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<ScheduledSpatialState> state;
};
} // namespace
std::unique_ptr<Pass> createScheduleSpatialGraphPass() {
return std::make_unique<ScheduleSpatialGraphPass>();
}
std::unique_ptr<Pass> createScheduleSpatialGraphPass(const SchedulingTarget& target) {
return std::make_unique<ScheduleSpatialGraphPass>(
target, std::make_shared<ScheduledSpatialState>());
}
std::unique_ptr<Pass> createScheduleSpatialGraphPass(
const SchedulingTarget& target, std::shared_ptr<ScheduledSpatialState> state) {
return std::make_unique<ScheduleSpatialGraphPass>(target, std::move(state));
}
std::unique_ptr<Pass> createVerifyScheduledSpatialPass() {
return std::make_unique<VerifyScheduledSpatialPass>();
}
std::unique_ptr<Pass> createVerifyScheduledSpatialPass(
std::shared_ptr<ScheduledSpatialState> state) {
return std::make_unique<VerifyScheduledSpatialPass>(std::move(state));
}
std::unique_ptr<Pass> createRealizeSpatialCommunicationPass() {
return std::make_unique<RealizeSpatialCommunicationPass>();
}
std::unique_ptr<Pass> createRealizeSpatialCommunicationPass(
const SchedulingTarget& target, std::shared_ptr<ScheduledSpatialState> state) {
return std::make_unique<RealizeSpatialCommunicationPass>(target, std::move(state));
}
std::unique_ptr<Pass> createVerifyRealizedSpatialPass() {
return std::make_unique<VerifyRealizedSpatialPass>();
}
std::unique_ptr<Pass> createVerifyRealizedSpatialPass(
std::shared_ptr<ScheduledSpatialState> state) {
return std::make_unique<VerifyRealizedSpatialPass>(std::move(state));
}
} // namespace spatial
} // namespace onnx_mlir
@@ -0,0 +1,16 @@
#pragma once
#include "ScheduledComputeMaterialization.hpp"
#include "Scheduling/MergeSchedulingAnalysis.hpp"
#include <memory>
#include <optional>
namespace onnx_mlir::spatial {
struct ScheduledSpatialState {
std::optional<MergeScheduleResult> logicalSchedule;
std::optional<ScheduledComputeMaterializationResult> materialization;
};
} // namespace onnx_mlir::spatial
@@ -181,7 +181,7 @@ FailureOr<LanePublicationSignatures> buildLanePublicationSignatures(SpatComputeB
for (auto [useIndex, use] : llvm::enumerate(result.getUses())) {
auto blueprint = dyn_cast<SpatBlueprintOp>(use.getOwner());
if (!blueprint || blueprint.getMode() != "fragment_assembly")
if (!blueprint || !isFragmentAssembly(blueprint.getMode()))
continue;
auto operandIndices = blueprint.getFragmentOperandIndices();
auto sourceSlots = blueprint.getFragmentSourceSlots();
@@ -400,6 +400,7 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
std::vector<ResidentWeightSet> processorResidentWeights(processorCount);
std::vector<ScheduledTask> schedules(nodeCount);
std::vector<std::vector<size_t>> tasksByProcessor(processorCount);
std::vector<std::vector<size_t>> timelineByProcessor(processorCount);
size_t scheduledCount = 0;
while (!readyQueue.empty()) {
@@ -441,7 +442,7 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
Time currentEnd = 0;
bool foundGap = false;
for (size_t schedTaskIndex : tasksByProcessor[processor]) {
for (size_t schedTaskIndex : timelineByProcessor[processor]) {
const ScheduledTask& schedTask = schedules[schedTaskIndex];
Time gapStart = std::max(currentEnd, dataReady);
@@ -532,10 +533,13 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
insertResidentWeights(capacityReservations[bestProcessor], graph.nodes[task].residentWeights);
insertResidentWeights(processorResidentWeights[bestProcessor], graph.nodes[task].residentWeights);
// 3. CRITICAL FIX: Topological Append
// Because the readyQueue pops in strict topological order, simply pushing to the
// back guarantees the Monoliths will be physically generated cycle-free.
// The hardware will still benefit from the processor assignment chosen by PEFT.
auto& timeline = timelineByProcessor[bestProcessor];
timeline.insert(llvm::upper_bound(timeline, task, [&](size_t lhs, size_t rhs) {
return schedules[lhs].startTime < schedules[rhs].startTime;
}), task);
// Materialization requires topological order; gap placement requires the
// separate chronological timeline above.
tasksByProcessor[bestProcessor].push_back(task);
for (const auto& [child, weight] : graph.successors[task]) {
@@ -1,5 +1,6 @@
#include "mlir/Dialect/Affine/IR/AffineOps.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/IR/AsmState.h"
#include "mlir/IR/BuiltinAttributes.h"
@@ -16,6 +17,7 @@
#include <cstdint>
#include <fstream>
#include <limits>
#include <optional>
#include <string>
#include <utility>
@@ -54,6 +56,12 @@ struct ChannelSendRecord {
std::optional<uint32_t> sourceLane;
};
struct ChannelEvaluationContext {
Value laneArg;
uint32_t lane = 0;
DenseMap<Value, int64_t> bindings;
};
enum class LogicalNodeSelector {
Scalar,
Lane,
@@ -249,10 +257,20 @@ void addBatchNodeRows(std::fstream& nodesFile,
}
}
std::optional<int64_t> evaluateIndexLike(Value value, Value laneArg, uint32_t lane);
std::optional<int64_t> evaluateIndexLike(Value value,
Value laneArg,
uint32_t lane,
const DenseMap<Value, int64_t>* bindings);
std::optional<int64_t> evaluateIndexLike(Value value, Value laneArg, uint32_t lane) {
if (value == laneArg)
std::optional<int64_t> evaluateIndexLike(Value value,
Value laneArg,
uint32_t lane,
const DenseMap<Value, int64_t>* bindings) {
if (bindings)
if (auto it = bindings->find(value); it != bindings->end())
return it->second;
if (laneArg && value == laneArg)
return static_cast<int64_t>(lane);
if (std::optional<int64_t> constant = matchConstantIndexValue(value))
@@ -270,7 +288,8 @@ std::optional<int64_t> evaluateIndexLike(Value value, Value laneArg, uint32_t la
if (!elements || !shapedType || shapedType.getRank() != 1 || extract.getIndices().size() != 1)
return std::nullopt;
std::optional<int64_t> index = evaluateIndexLike(extract.getIndices().front(), laneArg, lane);
std::optional<int64_t> index =
evaluateIndexLike(extract.getIndices().front(), laneArg, lane, bindings);
if (!index || *index < 0 || *index >= static_cast<int64_t>(elements.getNumElements()))
return std::nullopt;
@@ -279,11 +298,104 @@ std::optional<int64_t> evaluateIndexLike(Value value, Value laneArg, uint32_t la
return std::nullopt;
}
if (auto indexCast = value.getDefiningOp<arith::IndexCastOp>())
return evaluateIndexLike(indexCast.getIn(), laneArg, lane, bindings);
if (auto add = value.getDefiningOp<arith::AddIOp>()) {
auto lhs = evaluateIndexLike(add.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(add.getRhs(), laneArg, lane, bindings);
if (lhs && rhs)
return *lhs + *rhs;
return std::nullopt;
}
if (auto sub = value.getDefiningOp<arith::SubIOp>()) {
auto lhs = evaluateIndexLike(sub.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(sub.getRhs(), laneArg, lane, bindings);
if (lhs && rhs)
return *lhs - *rhs;
return std::nullopt;
}
if (auto mul = value.getDefiningOp<arith::MulIOp>()) {
auto lhs = evaluateIndexLike(mul.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(mul.getRhs(), laneArg, lane, bindings);
if (lhs && rhs)
return *lhs * *rhs;
return std::nullopt;
}
if (auto div = value.getDefiningOp<arith::DivSIOp>()) {
auto lhs = evaluateIndexLike(div.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(div.getRhs(), laneArg, lane, bindings);
if (!lhs || !rhs || *rhs == 0
|| (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1))
return std::nullopt;
return *lhs / *rhs;
}
if (auto div = value.getDefiningOp<arith::DivUIOp>()) {
auto lhs = evaluateIndexLike(div.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(div.getRhs(), laneArg, lane, bindings);
if (!lhs || !rhs || *rhs == 0)
return std::nullopt;
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) / static_cast<uint64_t>(*rhs));
}
if (auto rem = value.getDefiningOp<arith::RemSIOp>()) {
auto lhs = evaluateIndexLike(rem.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(rem.getRhs(), laneArg, lane, bindings);
if (!lhs || !rhs || *rhs == 0)
return std::nullopt;
if (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1)
return 0;
return *lhs % *rhs;
}
if (auto rem = value.getDefiningOp<arith::RemUIOp>()) {
auto lhs = evaluateIndexLike(rem.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(rem.getRhs(), laneArg, lane, bindings);
if (!lhs || !rhs || *rhs == 0)
return std::nullopt;
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) % static_cast<uint64_t>(*rhs));
}
if (auto cmp = value.getDefiningOp<arith::CmpIOp>()) {
auto lhs = evaluateIndexLike(cmp.getLhs(), laneArg, lane, bindings);
auto rhs = evaluateIndexLike(cmp.getRhs(), laneArg, lane, bindings);
if (!lhs || !rhs)
return std::nullopt;
bool result = false;
switch (cmp.getPredicate()) {
case arith::CmpIPredicate::eq: result = *lhs == *rhs; break;
case arith::CmpIPredicate::ne: result = *lhs != *rhs; break;
case arith::CmpIPredicate::slt: result = *lhs < *rhs; break;
case arith::CmpIPredicate::sle: result = *lhs <= *rhs; break;
case arith::CmpIPredicate::sgt: result = *lhs > *rhs; break;
case arith::CmpIPredicate::sge: result = *lhs >= *rhs; break;
case arith::CmpIPredicate::ult: result = static_cast<uint64_t>(*lhs) < static_cast<uint64_t>(*rhs); break;
case arith::CmpIPredicate::ule: result = static_cast<uint64_t>(*lhs) <= static_cast<uint64_t>(*rhs); break;
case arith::CmpIPredicate::ugt: result = static_cast<uint64_t>(*lhs) > static_cast<uint64_t>(*rhs); break;
case arith::CmpIPredicate::uge: result = static_cast<uint64_t>(*lhs) >= static_cast<uint64_t>(*rhs); break;
}
return result ? 1 : 0;
}
if (auto select = value.getDefiningOp<arith::SelectOp>()) {
auto condition = evaluateIndexLike(select.getCondition(), laneArg, lane, bindings);
if (!condition)
return std::nullopt;
return evaluateIndexLike(*condition ? select.getTrueValue() : select.getFalseValue(),
laneArg,
lane,
bindings);
}
if (auto affineApply = value.getDefiningOp<affine::AffineApplyOp>())
if (FailureOr<int64_t> folded = evaluateAffineApply(affineApply,
[&](Value operand) -> FailureOr<int64_t> {
if (std::optional<int64_t> resolved =
evaluateIndexLike(operand, laneArg, lane))
evaluateIndexLike(operand, laneArg, lane, bindings))
return *resolved;
return failure();
});
@@ -294,24 +406,70 @@ std::optional<int64_t> evaluateIndexLike(Value value, Value laneArg, uint32_t la
return std::nullopt;
}
SmallVector<int64_t, 8> collectPossibleIntValues(Value value, Value laneArg, uint32_t lane) {
if (std::optional<int64_t> exact = evaluateIndexLike(value, laneArg, lane))
return {*exact};
bool containsChannelOperation(Operation* root) {
bool found = false;
root->walk([&](Operation* op) {
found |= isa<SpatChannelSendOp, SpatChannelReceiveOp>(op);
});
return found;
}
auto extract = value.getDefiningOp<tensor::ExtractOp>();
auto constant = extract ? extract.getTensor().getDefiningOp<arith::ConstantOp>() : nullptr;
auto elements = constant ? dyn_cast<ElementsAttr>(constant.getValue()) : nullptr;
if (!elements)
return {};
template <typename Emit>
LogicalResult walkChannelRegion(Region& region, const ChannelEvaluationContext& context, Emit& emit) {
if (region.empty())
return success();
SmallVector<int64_t, 8> values;
if (auto denseInts = dyn_cast<DenseIntElementsAttr>(elements)) {
values.reserve(elements.getNumElements());
for (APInt element : denseInts.getValues<APInt>())
if (!llvm::is_contained(values, element.getSExtValue()))
values.push_back(element.getSExtValue());
for (Operation& op : region.front()) {
if (auto ifOp = dyn_cast<scf::IfOp>(&op)) {
auto condition = evaluateIndexLike(
ifOp.getCondition(), context.laneArg, context.lane, &context.bindings);
if (!condition)
return ifOp.emitOpError("has an unresolved condition in Spatial dataflow export");
Region& selected = *condition ? ifOp.getThenRegion() : ifOp.getElseRegion();
if (failed(walkChannelRegion(selected, context, emit)))
return failure();
continue;
}
if (auto forOp = dyn_cast<scf::ForOp>(&op)) {
if (!containsChannelOperation(forOp))
continue;
auto lower = evaluateIndexLike(forOp.getLowerBound(), context.laneArg, context.lane, &context.bindings);
auto upper = evaluateIndexLike(forOp.getUpperBound(), context.laneArg, context.lane, &context.bindings);
auto step = evaluateIndexLike(forOp.getStep(), context.laneArg, context.lane, &context.bindings);
if (!lower || !upper || !step || *step == 0)
return forOp.emitOpError("has unresolved or invalid bounds in Spatial dataflow export");
constexpr uint64_t kMaxExportedLoopIterations = 1 << 20;
uint64_t iterationCount = 0;
int64_t induction = *lower;
while ((*step > 0 && induction < *upper) || (*step < 0 && induction > *upper)) {
if (++iterationCount > kMaxExportedLoopIterations)
return forOp.emitOpError("exceeds the bounded iteration limit in Spatial dataflow export");
ChannelEvaluationContext iterationContext = context;
iterationContext.bindings[forOp.getInductionVar()] = induction;
if (failed(walkChannelRegion(forOp.getRegion(), iterationContext, emit)))
return failure();
if ((*step > 0 && induction > std::numeric_limits<int64_t>::max() - *step)
|| (*step < 0 && induction < std::numeric_limits<int64_t>::min() - *step))
return forOp.emitOpError("overflows while enumerating Spatial dataflow export iterations");
induction += *step;
}
continue;
}
if (!isa<SpatChannelSendOp, SpatChannelReceiveOp>(&op))
continue;
auto channel = dyn_cast<SpatChannelSendOp>(&op);
Value channelValue = channel ? channel.getChannelId() : cast<SpatChannelReceiveOp>(&op).getChannelId();
auto channelId = evaluateIndexLike(channelValue, context.laneArg, context.lane, &context.bindings);
if (!channelId)
return op.emitError("has an unresolved channel identity in Spatial dataflow export");
if (failed(emit(op, *channelId, context)))
return failure();
}
return values;
return success();
}
template <typename BatchOpTy>
@@ -604,50 +762,44 @@ LogicalResult emitDataEdges(std::fstream& edgesFile,
}
template <typename BatchOpTy>
void collectChannelSends(DenseMap<int64_t, SmallVector<ChannelSendRecord, 4>>& sendsByChannelId,
const DenseMap<std::pair<Operation*, uint32_t>, ExpandedNodeInfo>& expandedNodes,
BatchOpTy batch) {
LogicalResult collectChannelSends(DenseMap<int64_t, SmallVector<ChannelSendRecord, 4>>& sendsByChannelId,
const DenseMap<std::pair<Operation*, uint32_t>, ExpandedNodeInfo>& expandedNodes,
BatchOpTy batch) {
std::optional<BlockArgument> laneArg = batch.getLaneArgument();
if (!laneArg)
return;
return success();
for (uint32_t lane = 0; lane < static_cast<uint32_t>(batch.getLaneCount()); ++lane) {
std::string sourceId = getExpandedNodeId(expandedNodes, batch.getOperation(), lane);
if (sourceId.empty())
continue;
batch.getBody().walk([&](SpatChannelSendOp send) {
std::optional<int64_t> channelId = evaluateIndexLike(send.getChannelId(), *laneArg, lane);
if (!channelId)
return;
sendsByChannelId[*channelId].push_back({sourceId, lane});
});
ChannelEvaluationContext context;
context.laneArg = *laneArg;
context.lane = lane;
auto emit = [&](Operation& op, int64_t channelId, const ChannelEvaluationContext&) {
if (auto send = dyn_cast<SpatChannelSendOp>(&op))
sendsByChannelId[channelId].push_back({sourceId, lane});
return success();
};
if (failed(walkChannelRegion(batch.getBody(), context, emit)))
return failure();
}
return success();
}
void collectChannelSends(DenseMap<int64_t, SmallVector<ChannelSendRecord, 4>>& sendsByChannelId,
const DenseMap<std::pair<Operation*, uint32_t>, ExpandedNodeInfo>& expandedNodes,
SpatScheduledCompute compute) {
LogicalResult collectChannelSends(DenseMap<int64_t, SmallVector<ChannelSendRecord, 4>>& sendsByChannelId,
const DenseMap<std::pair<Operation*, uint32_t>, ExpandedNodeInfo>& expandedNodes,
SpatScheduledCompute compute) {
std::string sourceId = getExpandedNodeId(expandedNodes, compute.getOperation(), 0);
if (sourceId.empty())
return;
compute.getBody().walk([&](SpatChannelSendOp send) {
std::optional<int64_t> channelId = evaluateIndexLike(send.getChannelId(), Value(), 0);
if (!channelId)
return;
sendsByChannelId[*channelId].push_back({sourceId, std::nullopt});
});
}
DenseMap<int32_t, SmallVector<ChannelSendRecord, 4>>
buildNodesByCore(const DenseMap<std::pair<Operation*, uint32_t>, ExpandedNodeInfo>& expandedNodes) {
DenseMap<int32_t, SmallVector<ChannelSendRecord, 4>> nodesByCore;
for (const auto& entry : expandedNodes) {
const ExpandedNodeInfo& node = entry.second;
if (!node.core)
continue;
nodesByCore[*node.core].push_back({node.id, node.lane});
}
return nodesByCore;
return success();
ChannelEvaluationContext context;
auto emit = [&](Operation& op, int64_t channelId, const ChannelEvaluationContext&) {
if (isa<SpatChannelSendOp>(&op))
sendsByChannelId[channelId].push_back({sourceId, std::nullopt});
return success();
};
return walkChannelRegion(compute.getBody(), context, emit);
}
template <typename ComputeOpTy, typename BatchOpTy, typename ResolveChannelSourcesFn>
@@ -660,14 +812,18 @@ LogicalResult emitExplicitChannelEdges(std::fstream& edgesFile,
const TopLevelOpInfo& info = entry.second;
if (auto compute = dyn_cast<ComputeOpTy>(op)) {
compute.getBody().walk([&](SpatChannelReceiveOp receive) {
SmallVector<ChannelSendRecord, 4> sources = resolveChannelSources(receive, 0);
if (sources.empty())
return;
std::optional<int64_t> channelId = evaluateIndexLike(receive.getChannelId(), Value(), 0);
ChannelEvaluationContext context;
auto emit = [&](Operation& channelOp, int64_t channelId, const ChannelEvaluationContext&) {
auto receive = dyn_cast<SpatChannelReceiveOp>(&channelOp);
if (!receive)
return success();
FailureOr<SmallVector<ChannelSendRecord, 4>> sources =
resolveChannelSources(receive, channelId, 0);
if (failed(sources))
return failure();
std::string targetId = getScalarId(info.isScheduled, info.opId);
std::optional<uint64_t> byteSize = getTypeSizeBytes(receive.getType());
for (const ChannelSendRecord& source : sources)
for (const ChannelSendRecord& source : *sources)
emitEdgeRow(edgesFile,
source.sourceId,
targetId,
@@ -677,7 +833,10 @@ LogicalResult emitExplicitChannelEdges(std::fstream& edgesFile,
source.sourceLane,
std::nullopt,
channelId);
});
return success();
};
if (failed(walkChannelRegion(compute.getBody(), context, emit)))
return failure();
continue;
}
@@ -688,14 +847,20 @@ LogicalResult emitExplicitChannelEdges(std::fstream& edgesFile,
if (!laneArg)
continue;
for (uint32_t lane = 0; lane < static_cast<uint32_t>(batch.getLaneCount()); ++lane) {
std::string targetId = getBatchLaneId(info.isScheduled, info.opId, lane);
batch.getBody().walk([&](SpatChannelReceiveOp receive) {
SmallVector<ChannelSendRecord, 4> sources = resolveChannelSources(receive, lane);
if (sources.empty())
return;
std::optional<int64_t> channelId = evaluateIndexLike(receive.getChannelId(), *laneArg, lane);
ChannelEvaluationContext context;
context.laneArg = *laneArg;
context.lane = lane;
auto emit = [&](Operation& channelOp, int64_t channelId, const ChannelEvaluationContext& eventContext) {
auto receive = dyn_cast<SpatChannelReceiveOp>(&channelOp);
if (!receive)
return success();
FailureOr<SmallVector<ChannelSendRecord, 4>> sources =
resolveChannelSources(receive, channelId, eventContext.lane);
if (failed(sources))
return failure();
std::string targetId = getBatchLaneId(info.isScheduled, info.opId, eventContext.lane);
std::optional<uint64_t> byteSize = getTypeSizeBytes(receive.getType());
for (const ChannelSendRecord& source : sources)
for (const ChannelSendRecord& source : *sources)
emitEdgeRow(edgesFile,
source.sourceId,
targetId,
@@ -703,9 +868,12 @@ LogicalResult emitExplicitChannelEdges(std::fstream& edgesFile,
receive.getType(),
stage,
source.sourceLane,
lane,
eventContext.lane,
channelId);
});
return success();
};
if (failed(walkChannelRegion(batch.getBody(), context, emit)))
return failure();
}
}
@@ -810,33 +978,25 @@ LogicalResult exportScheduled(func::FuncOp func,
DenseMap<int64_t, SmallVector<ChannelSendRecord, 4>> sendsByChannelId;
for (const auto& entry : topLevelInfo) {
Operation* op = entry.first;
LogicalResult collected = success();
if (auto compute = dyn_cast<SpatScheduledCompute>(op))
collectChannelSends(sendsByChannelId, expandedNodes, compute);
collected = collectChannelSends(sendsByChannelId, expandedNodes, compute);
else if (auto batch = dyn_cast<SpatScheduledComputeBatch>(op))
collectChannelSends(sendsByChannelId, expandedNodes, batch);
collected = collectChannelSends(sendsByChannelId, expandedNodes, batch);
if (failed(collected))
return failure();
}
DenseMap<int32_t, SmallVector<ChannelSendRecord, 4>> nodesByCore = buildNodesByCore(expandedNodes);
auto resolveChannelSources = [&](SpatChannelReceiveOp receive, uint32_t lane) {
DenseMap<int64_t, size_t> consumedSendsByChannelId;
auto resolveChannelSources = [&](SpatChannelReceiveOp receive, int64_t channelId, uint32_t) {
SmallVector<ChannelSendRecord, 4> sources;
Value laneArg;
if (auto owner = receive->getParentOfType<SpatScheduledComputeBatch>())
if (auto maybeLaneArg = owner.getLaneArgument())
laneArg = *maybeLaneArg;
if (std::optional<int64_t> channelId = evaluateIndexLike(receive.getChannelId(), laneArg, lane)) {
if (auto it = sendsByChannelId.find(*channelId); it != sendsByChannelId.end())
return it->second;
}
for (int64_t sourceCore : collectPossibleIntValues(receive.getSourceCoreId(), laneArg, lane)) {
auto it = nodesByCore.find(static_cast<int32_t>(sourceCore));
if (it == nodesByCore.end())
continue;
llvm::append_range(sources, it->second);
}
return sources;
auto sends = sendsByChannelId.find(channelId);
size_t& consumed = consumedSendsByChannelId[channelId];
if (sends == sendsByChannelId.end() || consumed >= sends->second.size())
return receive.emitOpError("has no matching realized channel send in Spatial dataflow export"),
FailureOr<SmallVector<ChannelSendRecord, 4>>(failure());
sources.push_back(sends->second[consumed++]);
return FailureOr<SmallVector<ChannelSendRecord, 4>>(std::move(sources));
};
return emitExplicitChannelEdges<SpatScheduledCompute, SpatScheduledComputeBatch>(
@@ -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"
+25 -3
View File
@@ -9,18 +9,40 @@
namespace onnx_mlir {
namespace spatial {
struct SchedulingTarget;
struct ScheduledSpatialState;
struct SpatialTargetInfo;
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass();
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass(const SchedulingTarget& target);
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass(
const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createVerifyScheduledSpatialPass();
std::unique_ptr<mlir::Pass> createVerifyScheduledSpatialPass(
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createRealizeSpatialCommunicationPass();
std::unique_ptr<mlir::Pass> createRealizeSpatialCommunicationPass(
const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createVerifyRealizedSpatialPass();
std::unique_ptr<mlir::Pass> createVerifyRealizedSpatialPass(
std::shared_ptr<ScheduledSpatialState> state);
}
std::unique_ptr<mlir::Pass> createONNXToSpatialPass();
std::unique_ptr<mlir::Pass> createONNXToSpatialPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass();
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass();
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createSpatialToPimPass();
std::unique_ptr<mlir::Pass> createPimBufferizationPass();
std::unique_ptr<mlir::Pass> createPimBufferizationPreparationPass();
std::unique_ptr<mlir::Pass> createPimOneShotBufferizationPass();
std::unique_ptr<mlir::Pass> createPimMemoryNormalizationPass();
std::unique_ptr<mlir::Pass> createPimBufferizationVerificationPass();
std::unique_ptr<mlir::Pass> createMergeComputeNodesPass();
std::unique_ptr<mlir::Pass> createMergeComputeNodesPass(const spatial::SchedulingTarget& target);
std::unique_ptr<mlir::Pass> createTrivialGraphComputeMergePass();
std::unique_ptr<mlir::Pass> createTrivialGraphComputeMergePass(
+11 -5
View File
@@ -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);
+3
View File
@@ -3,6 +3,7 @@ operations/**/outputs
operations/**/raptor
operations/**/runner
operations/**/simulation
operations/**/*.csv
networks/**/inputs
networks/**/outputs
networks/**/raptor
@@ -14,3 +15,5 @@ networks/**/*.png
networks/**/*.jpg
networks/**/*.csv
!networks/pimcomp_models/results.csv
!networks/pimcomp_models/validation_results.csv
!operations/validation_results.csv
@@ -0,0 +1,6 @@
Operation,Result,Compile,Host mem,Cores mem,Cores,Xbars,Latency,Power,Energy
vgg8-mnist-reconstructed,PASS,1.009 s,1.37 MiB,3.14 MiB,141,761,1.465778 ms,325.627854 mW,477298145.040001 pJ
resnet18-v1-7,PASS,11.548 s,9.89 MiB,40.24 MiB,168,7676,28.099952 ms,312.513408 mW,8781611766.119984 pJ
resnet34-v1-7,PASS,28.495 s,9.90 MiB,48.89 MiB,168,15292,45.781486 ms,326.833870 mW,14962940227.679951 pJ
googlenet-12-latency,PASS,6.573 s,10.74 MiB,22.41 MiB,168,7176,13.371204 ms,457.538139 mW,6117835798.919991 pJ
yolo11n-latency,FAIL,58.572 s,82.55 MiB,185.68 MiB,168,6484,885.264931 ms,189.218985 mW,167508931321.001465 pJ
1 Operation Result Compile Host mem Cores mem Cores Xbars Latency Power Energy
2 vgg8-mnist-reconstructed PASS 1.009 s 1.37 MiB 3.14 MiB 141 761 1.465778 ms 325.627854 mW 477298145.040001 pJ
3 resnet18-v1-7 PASS 11.548 s 9.89 MiB 40.24 MiB 168 7676 28.099952 ms 312.513408 mW 8781611766.119984 pJ
4 resnet34-v1-7 PASS 28.495 s 9.90 MiB 48.89 MiB 168 15292 45.781486 ms 326.833870 mW 14962940227.679951 pJ
5 googlenet-12-latency PASS 6.573 s 10.74 MiB 22.41 MiB 168 7176 13.371204 ms 457.538139 mW 6117835798.919991 pJ
6 yolo11n-latency FAIL 58.572 s 82.55 MiB 185.68 MiB 168 6484 885.264931 ms 189.218985 mW 167508931321.001465 pJ
+6 -3
View File
@@ -43,7 +43,7 @@ and writes the same rows to `validation_results.csv`.
## Complete inventory
The suite contains 165 models. Tensor shapes, attributes, and constants are
The suite contains 168 models. Tensor shapes, attributes, and constants are
defined in `gen_tests.py` and in the checked-in ONNX models.
### Add (5)
@@ -64,7 +64,7 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `negative_axis` | Concatenates tensors using a negative axis. |
| `three_inputs_channel_axis` | Concatenates three runtime NCHW tensors along the channel axis. |
### Conv (32)
### Conv (34)
| Case | Description |
|---|---|
@@ -99,6 +99,8 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `with_bias_3x3` | Multi-channel 3x3 Conv with bias. |
| `with_constant` | Hand-authored SAME_UPPER Conv with constant weight and bias. |
| `without_kernel_shape_attr` | Conv whose kernel shape is inferred from its weight tensor. |
| `yolo11n_depthwise_head` | YOLO11n pointwise-to-depthwise head boundary at `80x80`, preserving row fragments. |
| `yolo11n_heavy` | Two largest standard YOLO11n Conv-SiLU blocks by MAC count at `64x80x80`. |
| `yolo11n_stem` | First two YOLO11n `Conv-SiLU` blocks at `640x640`, including the distributed activation boundary. |
### Div (6)
@@ -158,7 +160,7 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `with_homogeneous_constant` | Adds a constant bias matching the output shape. |
| `with_scalar_constant` | Adds a scalar broadcast bias. |
### MatMul (11)
### MatMul (12)
| Case | Description |
|---|---|
@@ -173,6 +175,7 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `left_constant` | Direct 2D MatMul with constant left-hand matrix. |
| `matrix_vector` | Matrix-vector multiplication producing a 1D output. |
| `vector_matrix` | Vector-matrix multiplication producing a 1D output. |
| `yolo_attention` | YOLO11n rank-4 dynamic MatMul-scale-transpose-MatMul attention chain. |
### Mul (5)
+63
View File
@@ -242,6 +242,50 @@ def conv_yolo11n_stem():
save_model(model, "conv/yolo11n_stem", "conv_yolo11n_stem.onnx")
def conv_yolo11n_heavy():
"""Two largest YOLO11n standard Conv-SiLU blocks by MAC count."""
X = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1, 64, 80, 80])
Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, 64, 80, 80])
rng = np.random.default_rng(111)
W0 = numpy_helper.from_array(rng.uniform(-1, 1, (64, 64, 3, 3)).astype(np.float32), name="W0")
B0 = numpy_helper.from_array(rng.uniform(-1, 1, (64,)).astype(np.float32), name="B0")
W1 = numpy_helper.from_array(rng.uniform(-1, 1, (64, 64, 3, 3)).astype(np.float32), name="W1")
B1 = numpy_helper.from_array(rng.uniform(-1, 1, (64,)).astype(np.float32), name="B1")
nodes = [
helper.make_node("Conv", ["X", "W0", "B0"], ["C0"],
kernel_shape=[3, 3], strides=[1, 1], pads=[1, 1, 1, 1]),
helper.make_node("Sigmoid", ["C0"], ["S0"]),
helper.make_node("Mul", ["C0", "S0"], ["A0"]),
helper.make_node("Conv", ["A0", "W1", "B1"], ["C1"],
kernel_shape=[3, 3], strides=[1, 1], pads=[1, 1, 1, 1]),
helper.make_node("Sigmoid", ["C1"], ["S1"]),
helper.make_node("Mul", ["C1", "S1"], ["Y"]),
]
graph = helper.make_graph(nodes, "conv_yolo11n_heavy", [X], [Y], initializer=[W0, B0, W1, B1])
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)])
save_model(model, "conv/yolo11n_heavy", "conv_yolo11n_heavy.onnx")
def conv_yolo11n_depthwise_head():
"""YOLO11n pointwise-to-depthwise head boundary at its largest feature map."""
X = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1, 64, 80, 80])
Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, 64, 80, 80])
rng = np.random.default_rng(110)
W0 = numpy_helper.from_array(rng.uniform(-1, 1, (64, 64, 1, 1)).astype(np.float32), name="W0")
B0 = numpy_helper.from_array(rng.uniform(-1, 1, (64,)).astype(np.float32), name="B0")
W1 = numpy_helper.from_array(rng.uniform(-1, 1, (64, 1, 3, 3)).astype(np.float32), name="W1")
B1 = numpy_helper.from_array(rng.uniform(-1, 1, (64,)).astype(np.float32), name="B1")
nodes = [
helper.make_node("Conv", ["X", "W0", "B0"], ["P"],
kernel_shape=[1, 1], strides=[1, 1], pads=[0, 0, 0, 0]),
helper.make_node("Conv", ["P", "W1", "B1"], ["Y"],
kernel_shape=[3, 3], strides=[1, 1], pads=[1, 1, 1, 1], group=64),
]
graph = helper.make_graph(nodes, "conv_yolo11n_depthwise_head", [X], [Y], initializer=[W0, B0, W1, B1])
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)])
save_model(model, "conv/yolo11n_depthwise_head", "conv_yolo11n_depthwise_head.onnx")
def conv_pointwise_tiled_chain():
"""Chained pointwise Convs with a tiled intermediate."""
X = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1, 1024, 1, 1])
@@ -760,6 +804,22 @@ def matmul_batched_3d_dynamic():
save_model(model, "matmul/batched_3d_dynamic", "matmul_batched_3d_dynamic.onnx")
def matmul_yolo_attention():
"""YOLO11n attention chain with rank-4 dynamic matrices."""
Q = helper.make_tensor_value_info("Q", TensorProto.FLOAT, [1, 2, 400, 32])
K = helper.make_tensor_value_info("K", TensorProto.FLOAT, [1, 2, 32, 400])
V = helper.make_tensor_value_info("V", TensorProto.FLOAT, [1, 2, 64, 400])
Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, 2, 64, 400])
scale = numpy_helper.from_array(np.asarray([0.1767767], dtype=np.float32), name="scale")
nodes = [helper.make_node("MatMul", ["Q", "K"], ["scores"]),
helper.make_node("Mul", ["scores", "scale"], ["scaled"]),
helper.make_node("Transpose", ["scaled"], ["weights"], perm=[0, 1, 3, 2]),
helper.make_node("MatMul", ["V", "weights"], ["Y"])]
graph = helper.make_graph(nodes, "matmul_yolo_attention", [Q, K, V], [Y], initializer=[scale])
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)])
save_model(model, "matmul/yolo_attention", "matmul_yolo_attention.onnx")
def matmul_batched_left_constant():
"""Batched 3D MatMul with constant LHS and runtime RHS."""
rng = np.random.default_rng(70)
@@ -2057,6 +2117,8 @@ if __name__ == "__main__":
conv_huge_pointwise_1024()
conv_huge_pointwise_1024_dynamic()
conv_yolo11n_stem()
conv_yolo11n_heavy()
conv_yolo11n_depthwise_head()
conv_pointwise_tiled_chain()
conv_large_output_channels_1x1()
conv_large_input_channels_1x1()
@@ -2078,6 +2140,7 @@ if __name__ == "__main__":
matmul_dynamic()
matmul_batched_3d()
matmul_batched_3d_dynamic()
matmul_yolo_attention()
matmul_batched_left_constant()
matmul_batched_rhs_broadcast()
matmul_batched_lhs_broadcast()
+168 -165
View File
@@ -1,166 +1,169 @@
Operation,Result,Compile,Host mem,Cores mem,Cores,Xbars,Latency,Power,Energy
add/after_gemm,PASS,0.141 s,0.01 MiB,0.01 MiB,5,4,0.007784 ms,104.703618 mW,815012.960000 pJ
add/basic,PASS,0.114 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
add/broadcast_row,PASS,0.110 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
add/channel_broadcast_1024,PASS,0.105 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ
add/leading_dimension_broadcast,PASS,0.137 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
concat/channel_axis,PASS,0.127 s,0.00 MiB,0.00 MiB,1,0,0.000457 ms,78.157549 mW,35718.000000 pJ
concat/negative_axis,PASS,0.102 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.127 s,0.00 MiB,0.00 MiB,1,0,0.000644 ms,78.149068 mW,50328.000000 pJ
conv/batch_2,PASS,0.102 s,0.00 MiB,0.00 MiB,2,2,0.013694 ms,82.623885 mW,1131451.480000 pJ
conv/batch_4_pointwise,PASS,0.110 s,0.00 MiB,0.01 MiB,5,4,0.003932 ms,116.078576 mW,456420.960000 pJ
conv/depthwise_1024_channels,PASS,0.167 s,0.19 MiB,0.38 MiB,129,128,0.220751 ms,178.454307 mW,39393966.720000 pJ
conv/depthwise_grouped,PASS,0.109 s,0.01 MiB,0.00 MiB,5,4,0.006024 ms,108.326521 mW,652558.960000 pJ
conv/dilated_3x3,PASS,0.143 s,0.00 MiB,0.00 MiB,3,3,0.004045 ms,110.234541 mW,445898.720000 pJ
conv/dynamic,PASS,0.107 s,0.00 MiB,0.00 MiB,5,0,0.001835 ms,92.281199 mW,169336.000000 pJ
conv/explicit_padding,PASS,0.139 s,0.00 MiB,0.00 MiB,4,4,0.004327 ms,115.794768 mW,501043.960000 pJ
conv/grouped_many_groups,PASS,0.616 s,0.05 MiB,0.09 MiB,65,64,0.181845 ms,142.210104 mW,25860196.360000 pJ
conv/grouped_two_groups,PASS,0.145 s,0.00 MiB,0.00 MiB,3,2,0.005360 ms,101.459418 mW,543822.480000 pJ
conv/huge_pointwise_1024,PASS,0.749 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.097 s,8.04 MiB,12.61 MiB,168,0,2.627964 ms,169.518697 mW,445489032.000000 pJ
conv/kernel_3x3,PASS,0.155 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.156 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.160 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.141 s,0.00 MiB,0.01 MiB,1,8,0.004964 ms,117.628106 mW,583905.920000 pJ
conv/large_spatial,PASS,0.141 s,0.00 MiB,0.01 MiB,6,6,0.004096 ms,129.015000 mW,528445.440000 pJ
conv/multi_channel,PASS,0.143 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.127 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.085 s,0.00 MiB,0.00 MiB,3,3,0.005526 ms,105.464843 mW,582798.720000 pJ
conv/non_uniform_stride,PASS,0.120 s,0.00 MiB,0.00 MiB,4,4,0.005808 ms,110.081433 mW,639352.960000 pJ
conv/pointwise_1x1,PASS,0.139 s,0.00 MiB,0.00 MiB,4,4,0.004539 ms,114.835858 mW,521239.960000 pJ
conv/pointwise_tiled_chain,PASS,0.943 s,0.01 MiB,0.02 MiB,2,80,0.084437 ms,102.289307 mW,8637002.200000 pJ
conv/real_asymmetric_padding,PASS,0.115 s,0.00 MiB,0.00 MiB,4,4,0.005232 ms,111.870214 mW,585304.960000 pJ
conv/relu_conv_store,PASS,0.107 s,0.05 MiB,0.08 MiB,32,32,0.062978 ms,246.649827 mW,15533512.800000 pJ
conv/same_lower_3x3,PASS,0.094 s,0.00 MiB,0.00 MiB,5,5,0.004700 ms,119.232170 mW,560391.200000 pJ
conv/same_padding_3x3,PASS,0.126 s,0.00 MiB,0.00 MiB,5,5,0.004700 ms,119.232170 mW,560391.200000 pJ
conv/simple,PASS,0.125 s,0.00 MiB,0.00 MiB,2,2,0.003148 ms,94.665972 mW,298008.480000 pJ
conv/stride_2,PASS,0.146 s,0.00 MiB,0.00 MiB,2,2,0.002827 ms,96.393873 mW,272505.480000 pJ
conv/with_bias_3x3,PASS,0.121 s,0.00 MiB,0.00 MiB,3,3,0.004898 ms,107.176546 mW,524950.720000 pJ
conv/with_constant,PASS,0.106 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.105 s,0.00 MiB,0.00 MiB,3,3,0.003091 ms,115.862414 mW,358130.720000 pJ
conv/yolo11n_stem,PASS,3.091 s,29.20 MiB,31.38 MiB,168,488,9.799213 ms,361.075203 mW,3538252825.000010 pJ
div/after_gemm,PASS,0.130 s,0.01 MiB,0.01 MiB,5,4,0.007784 ms,104.703618 mW,815012.960000 pJ
div/basic,PASS,0.116 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
div/channel_broadcast_1024,PASS,0.121 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ
div/leading_dimension_broadcast,PASS,0.095 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
div/runtime_scalar_rhs,PASS,0.126 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ
div/scalar_constant,PASS,0.069 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
gather/3d_input_axis1,PASS,0.121 s,0.00 MiB,0.00 MiB,1,0,0.000589 ms,78.081494 mW,45990.000000 pJ
gather/axis0_matrix_indices,PASS,0.088 s,0.00 MiB,0.00 MiB,1,0,0.000697 ms,78.068867 mW,54414.000000 pJ
gather/axis1,PASS,0.123 s,0.00 MiB,0.00 MiB,1,0,0.000801 ms,78.059925 mW,62526.000000 pJ
gather/negative_axis,PASS,0.087 s,0.00 MiB,0.00 MiB,1,0,0.001437 ms,78.033403 mW,112134.000000 pJ
gather/negative_indices,PASS,0.115 s,0.00 MiB,0.00 MiB,1,0,0.000376 ms,78.127660 mW,29376.000000 pJ
gemm/alpha_beta,PASS,0.121 s,0.01 MiB,0.01 MiB,5,4,0.007456 ms,105.272125 mW,784908.960000 pJ
gemm/bias_rank2_broadcast,PASS,0.127 s,0.00 MiB,0.01 MiB,5,4,0.007072 ms,105.979208 mW,749484.960000 pJ
gemm/dynamic,PASS,0.084 s,0.00 MiB,0.00 MiB,5,0,0.002421 ms,91.480793 mW,221475.000000 pJ
gemm/dynamic_alpha,PASS,0.104 s,0.00 MiB,0.00 MiB,5,0,0.003262 ms,91.415696 mW,298198.000000 pJ
gemm/dynamic_beta,PASS,0.110 s,0.00 MiB,0.00 MiB,5,0,0.004365 ms,91.316151 mW,398595.000000 pJ
gemm/dynamic_bias,PASS,0.092 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.089 s,0.00 MiB,0.00 MiB,5,0,0.005629 ms,91.279268 mW,513811.000000 pJ
gemm/dynamic_transB,PASS,0.077 s,0.00 MiB,0.00 MiB,5,0,0.001301 ms,91.378171 mW,118883.000000 pJ
gemm/huge_1024,PASS,0.219 s,0.01 MiB,0.10 MiB,73,64,0.017522 ms,215.037402 mW,3767885.360000 pJ
gemm/large,PASS,0.097 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.106 s,0.01 MiB,0.01 MiB,9,8,0.004748 ms,133.481449 mW,633769.920000 pJ
gemm/non_square,PASS,0.102 s,0.00 MiB,0.01 MiB,5,4,0.003527 ms,118.958310 mW,419565.960000 pJ
gemm/scalar_bias,PASS,0.111 s,0.00 MiB,0.01 MiB,5,4,0.007072 ms,105.979208 mW,749484.960000 pJ
gemm/simple,PASS,0.135 s,0.03 MiB,0.08 MiB,42,40,0.021640 ms,151.774196 mW,3284393.600000 pJ
gemm/small,PASS,0.074 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.129 s,0.01 MiB,0.02 MiB,17,8,0.007962 ms,131.005014 mW,1043061.920000 pJ
gemm/transA,PASS,0.089 s,0.00 MiB,0.01 MiB,5,4,0.005762 ms,109.140743 mW,628868.960000 pJ
gemm/transA_transB,PASS,0.115 s,0.00 MiB,0.01 MiB,5,4,0.005762 ms,109.140743 mW,628868.960000 pJ
gemm/transB,PASS,0.068 s,0.00 MiB,0.01 MiB,5,4,0.003527 ms,118.958310 mW,419565.960000 pJ
gemm/transB_with_bias,PASS,0.069 s,0.01 MiB,0.01 MiB,5,4,0.005046 ms,110.546762 mW,557818.960000 pJ
gemm/with_bias,PASS,0.113 s,0.01 MiB,0.01 MiB,5,4,0.005562 ms,108.767882 mW,604966.960000 pJ
gemv/constant,PASS,0.112 s,0.00 MiB,0.00 MiB,0,0,0.000000 ms,2.000000 mW,0.000000 pJ
gemv/simple,PASS,0.136 s,0.00 MiB,0.01 MiB,6,4,0.005160 ms,111.150380 mW,573535.960000 pJ
gemv/with_heterogeneous_constant,PASS,0.138 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ
gemv/with_homogeneous_constant,PASS,0.140 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ
gemv/with_scalar_constant,PASS,0.124 s,0.00 MiB,0.01 MiB,6,4,0.005549 ms,109.816536 mW,609371.960000 pJ
matmul/basic,PASS,0.081 s,0.00 MiB,0.00 MiB,2,2,0.004420 ms,90.144000 mW,398436.480000 pJ
matmul/batched_3d,PASS,0.133 s,0.00 MiB,0.01 MiB,5,4,0.005958 ms,108.588949 mW,646972.960000 pJ
matmul/batched_3d_dynamic,PASS,0.105 s,0.00 MiB,0.00 MiB,9,0,0.003525 ms,92.330213 mW,325464.000000 pJ
matmul/batched_left_constant,PASS,0.136 s,0.00 MiB,0.02 MiB,9,8,0.008822 ms,114.385164 mW,1009105.920000 pJ
matmul/batched_lhs_broadcast,PASS,0.133 s,0.00 MiB,0.01 MiB,5,4,0.005681 ms,109.389361 mW,621440.960000 pJ
matmul/batched_rhs_broadcast,PASS,0.134 s,0.00 MiB,0.01 MiB,5,4,0.005958 ms,108.588949 mW,646972.960000 pJ
matmul/dynamic,PASS,0.093 s,0.00 MiB,0.00 MiB,5,0,0.001621 ms,91.421962 mW,148195.000000 pJ
matmul/huge_1024,PASS,0.277 s,0.01 MiB,0.10 MiB,73,64,0.017522 ms,215.037402 mW,3767885.360000 pJ
matmul/left_constant,PASS,0.123 s,0.00 MiB,0.01 MiB,5,4,0.005853 ms,108.861944 mW,637168.960000 pJ
matmul/matrix_vector,PASS,0.150 s,0.52 MiB,0.78 MiB,168,173,0.384660 ms,202.131271 mW,77751814.880000 pJ
matmul/vector_matrix,PASS,0.186 s,0.01 MiB,0.01 MiB,9,8,0.007409 ms,118.680243 mW,879301.920000 pJ
mul/after_conv,PASS,0.127 s,0.00 MiB,0.00 MiB,4,3,0.005453 ms,107.639046 mW,586955.720000 pJ
mul/basic,PASS,0.070 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
mul/channel_broadcast_1024,PASS,0.065 s,0.02 MiB,0.01 MiB,1,0,0.006913 ms,78.118038 mW,540030.000000 pJ
mul/leading_dimension_broadcast,PASS,0.108 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
mul/scalar_constant,PASS,0.082 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
pool/avg_basic,PASS,0.070 s,0.00 MiB,0.00 MiB,1,0,0.011939 ms,78.022112 mW,931506.000000 pJ
pool/avg_ceil_mode,PASS,0.107 s,0.00 MiB,0.00 MiB,1,0,0.004359 ms,78.033035 mW,340146.000000 pJ
pool/avg_explicit_padding,PASS,0.111 s,0.00 MiB,0.00 MiB,1,0,0.008822 ms,78.027205 mW,688356.000000 pJ
pool/avg_include_pad,PASS,0.075 s,0.00 MiB,0.00 MiB,1,0,0.008506 ms,78.016929 mW,663612.000000 pJ
pool/avg_large_channels,PASS,0.064 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.068 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.064 s,0.00 MiB,0.00 MiB,1,0,0.025206 ms,78.024756 mW,1966692.000000 pJ
pool/max_after_conv,PASS,0.070 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.065 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.089 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.103 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.088 s,0.00 MiB,0.00 MiB,4,0,0.003078 ms,93.124756 mW,286638.000000 pJ
pool/max_same_upper,PASS,0.065 s,0.00 MiB,0.00 MiB,3,0,0.003024 ms,92.095238 mW,278496.000000 pJ
pool/max_stride2_multichannel,PASS,0.086 s,0.00 MiB,0.00 MiB,3,0,0.004012 ms,92.269192 mW,370184.000000 pJ
reduce_mean/4d_spatial,PASS,0.061 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.067 s,0.00 MiB,0.00 MiB,4,0,0.000655 ms,94.352672 mW,61801.000000 pJ
reduce_mean/after_conv,PASS,0.115 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.103 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.086 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ
reduce_mean/basic,PASS,0.099 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.078 s,0.03 MiB,0.02 MiB,4,0,0.164926 ms,93.596631 mW,15436518.000000 pJ
reduce_mean/keepdims_0,PASS,0.073 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.115 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.074 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.111 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.086 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.115 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.090 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.110 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.107 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ
reduce_mean/negative_axis,PASS,0.120 s,0.00 MiB,0.00 MiB,6,0,0.000553 ms,93.520796 mW,51717.000000 pJ
relu/4d,PASS,0.116 s,0.00 MiB,0.00 MiB,1,0,0.000521 ms,78.184261 mW,40734.000000 pJ
relu/after_conv,PASS,0.078 s,0.00 MiB,0.00 MiB,3,3,0.004956 ms,106.998935 mW,530286.720000 pJ
relu/after_gemm,PASS,0.077 s,0.01 MiB,0.01 MiB,5,4,0.007513 ms,105.158653 mW,790056.960000 pJ
relu/basic,PASS,0.066 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.098 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.105 s,0.00 MiB,0.00 MiB,1,0,0.000162 ms,78.296296 mW,12684.000000 pJ
reshape/same_rank,PASS,0.057 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.090 s,0.00 MiB,0.00 MiB,1,0,0.000162 ms,78.296296 mW,12684.000000 pJ
resize/height_only,PASS,0.104 s,0.00 MiB,0.00 MiB,1,0,0.000795 ms,78.060377 mW,62058.000000 pJ
resize/nearest_2x,PASS,0.104 s,0.00 MiB,0.00 MiB,1,0,0.001422 ms,78.033755 mW,110964.000000 pJ
resize/nearest_downsample,PASS,0.059 s,0.00 MiB,0.00 MiB,1,0,0.000481 ms,78.099792 mW,37566.000000 pJ
resize/non_uniform,PASS,0.063 s,0.00 MiB,0.00 MiB,1,0,0.002108 ms,78.034156 mW,164496.000000 pJ
resize/width_only,PASS,0.068 s,0.00 MiB,0.00 MiB,1,0,0.000792 ms,78.060606 mW,61824.000000 pJ
resize/with_sizes,PASS,0.056 s,0.00 MiB,0.00 MiB,1,0,0.000951 ms,78.050473 mW,74226.000000 pJ
sigmoid/4d,PASS,0.104 s,0.00 MiB,0.00 MiB,1,0,0.000521 ms,78.184261 mW,40734.000000 pJ
sigmoid/after_gemm,PASS,0.118 s,0.01 MiB,0.01 MiB,5,4,0.007513 ms,105.158653 mW,790056.960000 pJ
sigmoid/basic,PASS,0.083 s,0.00 MiB,0.00 MiB,1,0,0.000221 ms,78.217195 mW,17286.000000 pJ
slice/2d_basic,PASS,0.067 s,0.00 MiB,0.00 MiB,1,0,0.000242 ms,78.297521 mW,18948.000000 pJ
slice/after_conv,PASS,0.111 s,0.00 MiB,0.01 MiB,7,6,0.011296 ms,118.190765 mW,1335082.880000 pJ
slice/default_axes,PASS,0.065 s,0.00 MiB,0.00 MiB,1,0,0.000242 ms,78.297521 mW,18948.000000 pJ
slice/large_channel_1024,PASS,0.060 s,0.01 MiB,0.00 MiB,1,0,0.002832 ms,78.144068 mW,221304.000000 pJ
slice/nchw_spatial_crop,PASS,0.067 s,0.00 MiB,0.00 MiB,1,0,0.001302 ms,78.239631 mW,101868.000000 pJ
slice/negative_axis,PASS,0.098 s,0.00 MiB,0.00 MiB,1,0,0.000562 ms,78.298932 mW,44004.000000 pJ
slice/negative_indices,PASS,0.106 s,0.00 MiB,0.00 MiB,1,0,0.000322 ms,78.298137 mW,25212.000000 pJ
slice/step2,PASS,0.068 s,0.00 MiB,0.00 MiB,1,0,0.002042 ms,78.293830 mW,159876.000000 pJ
softmax/3d_last_axis,PASS,0.072 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED
softmax/basic,PASS,0.088 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED
softmax/channel_axis,PASS,0.110 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED
softmax/large_dimension_1024,PASS,0.104 s,0.01 MiB,0.01 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED
softmax/negative_axis,PASS,0.118 s,0.00 MiB,0.00 MiB,1,0,UNSUPPORTED,UNSUPPORTED,UNSUPPORTED
split/basic,PASS,0.077 s,0.00 MiB,0.00 MiB,1,0,0.000403 ms,78.297767 mW,31554.000000 pJ
split/equal_three_way,PASS,0.107 s,0.00 MiB,0.00 MiB,1,0,0.000564 ms,78.297872 mW,44160.000000 pJ
split/negative_axis,PASS,0.111 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.061 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.063 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
sub/broadcast_row,PASS,0.075 s,0.00 MiB,0.00 MiB,1,0,0.000323 ms,78.222910 mW,25266.000000 pJ
sub/channel_broadcast_1024,PASS,0.052 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.053 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
1 Operation Result Compile Host mem Cores mem Cores Xbars Latency Power Energy
2 add/after_gemm PASS 0.141 s 0.062 s 0.01 MiB 0.01 MiB 5 4 0.007784 ms SKIP 104.703618 mW SKIP 815012.960000 pJ SKIP
3 add/basic PASS 0.114 s 0.053 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
4 add/broadcast_row PASS 0.110 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
5 add/channel_broadcast_1024 PASS 0.105 s 0.061 s 0.02 MiB 0.01 MiB 1 0 0.006913 ms SKIP 78.118038 mW SKIP 540030.000000 pJ SKIP
6 add/leading_dimension_broadcast PASS 0.137 s 0.057 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
7 concat/channel_axis PASS 0.127 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.000457 ms SKIP 78.157549 mW SKIP 35718.000000 pJ SKIP
8 concat/negative_axis PASS 0.102 s 0.061 s 0.00 MiB 0.00 MiB 1 0 0.001043 ms SKIP 78.092042 mW SKIP 81450.000000 pJ SKIP
9 concat/three_inputs_channel_axis PASS 0.127 s 0.053 s 0.00 MiB 0.00 MiB 1 0 0.000644 ms SKIP 78.149068 mW SKIP 50328.000000 pJ SKIP
10 conv/batch_2 PASS 0.102 s 0.056 s 0.00 MiB 0.00 MiB 2 2 0.013694 ms SKIP 82.623885 mW SKIP 1131451.480000 pJ SKIP
11 conv/batch_4_pointwise PASS 0.110 s 0.058 s 0.00 MiB 0.01 MiB 5 4 0.003932 ms SKIP 116.078576 mW SKIP 456420.960000 pJ SKIP
12 conv/depthwise_1024_channels PASS 0.167 s 0.085 s 0.19 MiB 0.38 MiB 129 128 0.220751 ms SKIP 178.454307 mW SKIP 39393966.720000 pJ SKIP
13 conv/depthwise_grouped PASS 0.109 s 0.074 s 0.01 MiB 0.00 MiB 5 4 0.006024 ms SKIP 108.326521 mW SKIP 652558.960000 pJ SKIP
14 conv/dilated_3x3 PASS 0.143 s 0.067 s 0.00 MiB 0.00 MiB 3 3 0.004045 ms SKIP 110.234541 mW SKIP 445898.720000 pJ SKIP
15 conv/dynamic PASS 0.107 s 0.059 s 0.00 MiB 0.00 MiB 5 0 0.001835 ms SKIP 92.281199 mW SKIP 169336.000000 pJ SKIP
16 conv/explicit_padding PASS 0.139 s 0.065 s 0.00 MiB 0.00 MiB 4 4 0.004327 ms SKIP 115.794768 mW SKIP 501043.960000 pJ SKIP
17 conv/grouped_many_groups PASS 0.616 s 0.500 s 0.05 MiB 0.09 MiB 65 64 0.181845 ms SKIP 142.210104 mW SKIP 25860196.360000 pJ SKIP
18 conv/grouped_two_groups PASS 0.145 s 0.070 s 0.00 MiB 0.00 MiB 3 2 0.005360 ms SKIP 101.459418 mW SKIP 543822.480000 pJ SKIP
19 conv/huge_pointwise_1024 PASS 0.749 s 0.652 s 0.01 MiB 0.01 MiB 1 64 0.028261 ms SKIP 133.488743 mW SKIP 3772525.360000 pJ SKIP
20 conv/huge_pointwise_1024_dynamic PASS 0.097 s 0.083 s 8.04 MiB 12.61 MiB 168 0 2.627964 ms SKIP 169.518697 mW SKIP 445489032.000000 pJ SKIP
21 conv/kernel_3x3 PASS 0.155 s 0.067 s 0.00 MiB 0.00 MiB 3 3 0.003091 ms SKIP 115.862414 mW SKIP 358130.720000 pJ SKIP
22 conv/kernel_equals_input_spatial PASS 0.156 s 0.072 s 0.00 MiB 0.00 MiB 1 2 0.008443 ms SKIP 83.863849 mW SKIP 708062.480000 pJ SKIP
23 conv/large_input_channels_1x1 PASS 0.160 s 0.098 s 0.01 MiB 0.01 MiB 1 8 0.017167 ms SKIP 89.416900 mW SKIP 1535019.920000 pJ SKIP
24 conv/large_output_channels_1x1 PASS 0.141 s 0.089 s 0.00 MiB 0.01 MiB 1 8 0.004964 ms SKIP 117.628106 mW SKIP 583905.920000 pJ SKIP
25 conv/large_spatial PASS 0.141 s 0.063 s 0.00 MiB 0.01 MiB 6 6 0.004096 ms SKIP 129.015000 mW SKIP 528445.440000 pJ SKIP
26 conv/multi_channel PASS 0.143 s 0.068 s 0.00 MiB 0.00 MiB 3 3 0.005148 ms SKIP 106.453520 mW SKIP 548022.720000 pJ SKIP
27 conv/non_square_kernel_1x3 PASS 0.127 s 0.080 s 0.00 MiB 0.00 MiB 5 5 0.004029 ms SKIP 123.600943 mW SKIP 497988.200000 pJ SKIP
28 conv/non_square_kernel_3x1 PASS 0.085 s 0.068 s 0.00 MiB 0.00 MiB 3 3 0.005526 ms SKIP 105.464843 mW SKIP 582798.720000 pJ SKIP
29 conv/non_uniform_stride PASS 0.120 s 0.068 s 0.00 MiB 0.00 MiB 4 4 0.005808 ms SKIP 110.081433 mW SKIP 639352.960000 pJ SKIP
30 conv/pointwise_1x1 PASS 0.139 s 0.071 s 0.00 MiB 0.00 MiB 4 4 0.004539 ms SKIP 114.835858 mW SKIP 521239.960000 pJ SKIP
31 conv/pointwise_tiled_chain PASS 0.943 s 0.911 s 0.01 MiB 0.02 MiB 2 80 0.084437 ms SKIP 102.289307 mW SKIP 8637002.200000 pJ SKIP
32 conv/real_asymmetric_padding PASS 0.115 s 0.057 s 0.00 MiB 0.00 MiB 4 4 0.005232 ms SKIP 111.870214 mW SKIP 585304.960000 pJ SKIP
33 conv/relu_conv_store PASS 0.107 s 0.070 s 0.05 MiB 0.02 MiB 0.08 MiB 0.10 MiB 32 32 0.062978 ms SKIP 246.649827 mW SKIP 15533512.800000 pJ SKIP
34 conv/same_lower_3x3 PASS 0.094 s 0.061 s 0.00 MiB 0.00 MiB 5 5 0.004700 ms SKIP 119.232170 mW SKIP 560391.200000 pJ SKIP
35 conv/same_padding_3x3 PASS 0.126 s 0.075 s 0.00 MiB 0.00 MiB 5 5 0.004700 ms SKIP 119.232170 mW SKIP 560391.200000 pJ SKIP
36 conv/simple PASS 0.125 s 0.055 s 0.00 MiB 0.00 MiB 2 2 0.003148 ms SKIP 94.665972 mW SKIP 298008.480000 pJ SKIP
37 conv/stride_2 PASS 0.146 s 0.057 s 0.00 MiB 0.00 MiB 2 2 0.002827 ms SKIP 96.393873 mW SKIP 272505.480000 pJ SKIP
38 conv/with_bias_3x3 PASS 0.121 s 0.067 s 0.00 MiB 0.00 MiB 3 3 0.004898 ms SKIP 107.176546 mW SKIP 524950.720000 pJ SKIP
39 conv/with_constant PASS 0.106 s 0.063 s 0.00 MiB 0.00 MiB 3 3 0.004273 ms SKIP 109.362677 mW SKIP 467306.720000 pJ SKIP
40 conv/without_kernel_shape_attr PASS 0.105 s 0.059 s 0.00 MiB 0.00 MiB 3 3 0.003091 ms SKIP 115.862414 mW SKIP 358130.720000 pJ SKIP
41 conv/yolo11n_stem conv/yolo11n_depthwise_head PASS 3.091 s 0.624 s 29.20 MiB 4.82 MiB 31.38 MiB 15.92 MiB 168 160 488 720 9.799213 ms SKIP 361.075203 mW SKIP 3538252825.000010 pJ SKIP
42 div/after_gemm conv/yolo11n_heavy PASS 0.130 s 0.496 s 0.01 MiB 4.82 MiB 0.01 MiB 20.66 MiB 5 160 4 800 0.007784 ms SKIP 104.703618 mW SKIP 815012.960000 pJ SKIP
43 div/basic conv/yolo11n_stem PASS 0.116 s 0.935 s 0.00 MiB 12.86 MiB 0.00 MiB 31.38 MiB 1 168 0 488 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
44 div/channel_broadcast_1024 div/after_gemm PASS 0.121 s 0.060 s 0.02 MiB 0.01 MiB 0.01 MiB 1 5 0 4 0.006913 ms SKIP 78.118038 mW SKIP 540030.000000 pJ SKIP
45 div/leading_dimension_broadcast div/basic PASS 0.095 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
46 div/runtime_scalar_rhs div/channel_broadcast_1024 PASS 0.126 s 0.049 s 0.02 MiB 0.01 MiB 1 0 0.006913 ms SKIP 78.118038 mW SKIP 540030.000000 pJ SKIP
47 div/scalar_constant div/leading_dimension_broadcast PASS 0.069 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
48 gather/3d_input_axis1 div/runtime_scalar_rhs PASS 0.121 s 0.053 s 0.00 MiB 0.02 MiB 0.00 MiB 0.01 MiB 1 0 0.000589 ms SKIP 78.081494 mW SKIP 45990.000000 pJ SKIP
49 gather/axis0_matrix_indices div/scalar_constant PASS 0.088 s 0.078 s 0.00 MiB 0.00 MiB 1 0 0.000697 ms SKIP 78.068867 mW SKIP 54414.000000 pJ SKIP
50 gather/axis1 gather/3d_input_axis1 PASS 0.123 s 0.059 s 0.00 MiB 0.00 MiB 1 0 0.000801 ms SKIP 78.059925 mW SKIP 62526.000000 pJ SKIP
51 gather/negative_axis gather/axis0_matrix_indices PASS 0.087 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.001437 ms SKIP 78.033403 mW SKIP 112134.000000 pJ SKIP
52 gather/negative_indices gather/axis1 PASS 0.115 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000376 ms SKIP 78.127660 mW SKIP 29376.000000 pJ SKIP
53 gemm/alpha_beta gather/negative_axis PASS 0.121 s 0.056 s 0.01 MiB 0.00 MiB 0.01 MiB 0.00 MiB 5 1 4 0 0.007456 ms SKIP 105.272125 mW SKIP 784908.960000 pJ SKIP
54 gemm/bias_rank2_broadcast gather/negative_indices PASS 0.127 s 0.051 s 0.00 MiB 0.01 MiB 0.00 MiB 5 1 4 0 0.007072 ms SKIP 105.979208 mW SKIP 749484.960000 pJ SKIP
55 gemm/dynamic gemm/alpha_beta PASS 0.084 s 0.078 s 0.00 MiB 0.01 MiB 0.00 MiB 0.01 MiB 5 0 4 0.002421 ms SKIP 91.480793 mW SKIP 221475.000000 pJ SKIP
56 gemm/dynamic_alpha gemm/bias_rank2_broadcast PASS 0.104 s 0.055 s 0.00 MiB 0.00 MiB 0.01 MiB 5 0 4 0.003262 ms SKIP 91.415696 mW SKIP 298198.000000 pJ SKIP
57 gemm/dynamic_beta gemm/dynamic PASS 0.110 s 0.058 s 0.00 MiB 0.00 MiB 5 0 0.004365 ms SKIP 91.316151 mW SKIP 398595.000000 pJ SKIP
58 gemm/dynamic_bias gemm/dynamic_alpha PASS 0.092 s 0.056 s 0.00 MiB 0.00 MiB 5 0 0.002665 ms SKIP 91.445779 mW SKIP 243703.000000 pJ SKIP
59 gemm/dynamic_bias_alpha_beta gemm/dynamic_beta PASS 0.089 s 0.053 s 0.00 MiB 0.00 MiB 5 0 0.005629 ms SKIP 91.279268 mW SKIP 513811.000000 pJ SKIP
60 gemm/dynamic_transB gemm/dynamic_bias PASS 0.077 s 0.062 s 0.00 MiB 0.00 MiB 5 0 0.001301 ms SKIP 91.378171 mW SKIP 118883.000000 pJ SKIP
61 gemm/huge_1024 gemm/dynamic_bias_alpha_beta PASS 0.219 s 0.114 s 0.01 MiB 0.00 MiB 0.10 MiB 0.00 MiB 73 5 64 0 0.017522 ms SKIP 215.037402 mW SKIP 3767885.360000 pJ SKIP
62 gemm/large gemm/dynamic_transB PASS 0.097 s 0.055 s 0.02 MiB 0.00 MiB 0.03 MiB 0.00 MiB 17 5 16 0 0.011229 ms SKIP 140.152181 mW SKIP 1573768.840000 pJ SKIP
63 gemm/large_k_small_n gemm/huge_1024 PASS 0.106 s 0.150 s 0.01 MiB 0.01 MiB 0.10 MiB 9 73 8 64 0.004748 ms SKIP 133.481449 mW SKIP 633769.920000 pJ SKIP
64 gemm/non_square gemm/large PASS 0.102 s 0.059 s 0.00 MiB 0.02 MiB 0.01 MiB 0.03 MiB 5 17 4 16 0.003527 ms SKIP 118.958310 mW SKIP 419565.960000 pJ SKIP
65 gemm/scalar_bias gemm/large_k_small_n PASS 0.111 s 0.087 s 0.00 MiB 0.01 MiB 0.01 MiB 5 9 4 8 0.007072 ms SKIP 105.979208 mW SKIP 749484.960000 pJ SKIP
66 gemm/simple gemm/non_square PASS 0.135 s 0.058 s 0.03 MiB 0.00 MiB 0.08 MiB 0.01 MiB 42 5 40 4 0.021640 ms SKIP 151.774196 mW SKIP 3284393.600000 pJ SKIP
67 gemm/small gemm/scalar_bias PASS 0.074 s 0.058 s 0.00 MiB 0.00 MiB 0.01 MiB 2 5 2 4 0.004420 ms SKIP 90.144000 mW SKIP 398436.480000 pJ SKIP
68 gemm/small_k_large_n gemm/simple PASS 0.129 s 0.075 s 0.01 MiB 0.03 MiB 0.02 MiB 0.08 MiB 17 42 8 40 0.007962 ms SKIP 131.005014 mW SKIP 1043061.920000 pJ SKIP
69 gemm/transA gemm/small PASS 0.089 s 0.056 s 0.00 MiB 0.01 MiB 0.00 MiB 5 2 4 2 0.005762 ms SKIP 109.140743 mW SKIP 628868.960000 pJ SKIP
70 gemm/transA_transB gemm/small_k_large_n PASS 0.115 s 0.091 s 0.00 MiB 0.01 MiB 0.01 MiB 0.02 MiB 5 17 4 8 0.005762 ms SKIP 109.140743 mW SKIP 628868.960000 pJ SKIP
71 gemm/transB gemm/transA PASS 0.068 s 0.057 s 0.00 MiB 0.01 MiB 5 4 0.003527 ms SKIP 118.958310 mW SKIP 419565.960000 pJ SKIP
72 gemm/transB_with_bias gemm/transA_transB PASS 0.069 s 0.079 s 0.01 MiB 0.00 MiB 0.01 MiB 5 4 0.005046 ms SKIP 110.546762 mW SKIP 557818.960000 pJ SKIP
73 gemm/with_bias gemm/transB PASS 0.113 s 0.065 s 0.01 MiB 0.00 MiB 0.01 MiB 5 4 0.005562 ms SKIP 108.767882 mW SKIP 604966.960000 pJ SKIP
74 gemv/constant gemm/transB_with_bias PASS 0.112 s 0.058 s 0.00 MiB 0.01 MiB 0.00 MiB 0.01 MiB 0 5 0 4 0.000000 ms SKIP 2.000000 mW SKIP 0.000000 pJ SKIP
75 gemv/simple gemm/with_bias PASS 0.136 s 0.054 s 0.00 MiB 0.01 MiB 0.01 MiB 6 5 4 0.005160 ms SKIP 111.150380 mW SKIP 573535.960000 pJ SKIP
76 gemv/with_heterogeneous_constant gemv/constant PASS 0.138 s 0.049 s 0.00 MiB 0.01 MiB 0.00 MiB 6 0 4 0 0.005549 ms SKIP 109.816536 mW SKIP 609371.960000 pJ SKIP
77 gemv/with_homogeneous_constant gemv/simple PASS 0.140 s 0.066 s 0.00 MiB 0.01 MiB 6 4 0.005549 ms SKIP 109.816536 mW SKIP 609371.960000 pJ SKIP
78 gemv/with_scalar_constant gemv/with_heterogeneous_constant PASS 0.124 s 0.069 s 0.00 MiB 0.01 MiB 6 4 0.005549 ms SKIP 109.816536 mW SKIP 609371.960000 pJ SKIP
79 matmul/basic gemv/with_homogeneous_constant PASS 0.081 s 0.065 s 0.00 MiB 0.00 MiB 0.01 MiB 2 6 2 4 0.004420 ms SKIP 90.144000 mW SKIP 398436.480000 pJ SKIP
80 matmul/batched_3d gemv/with_scalar_constant PASS 0.133 s 0.100 s 0.00 MiB 0.01 MiB 5 6 4 0.005958 ms SKIP 108.588949 mW SKIP 646972.960000 pJ SKIP
81 matmul/batched_3d_dynamic matmul/basic PASS 0.105 s 0.056 s 0.00 MiB 0.00 MiB 9 2 0 2 0.003525 ms SKIP 92.330213 mW SKIP 325464.000000 pJ SKIP
82 matmul/batched_left_constant matmul/batched_3d PASS 0.136 s 0.057 s 0.00 MiB 0.02 MiB 0.01 MiB 9 5 8 4 0.008822 ms SKIP 114.385164 mW SKIP 1009105.920000 pJ SKIP
83 matmul/batched_lhs_broadcast matmul/batched_3d_dynamic PASS 0.133 s 0.053 s 0.00 MiB 0.01 MiB 0.00 MiB 5 4 4 0 0.005681 ms SKIP 109.389361 mW SKIP 621440.960000 pJ SKIP
84 matmul/batched_rhs_broadcast matmul/batched_left_constant PASS 0.134 s 0.062 s 0.00 MiB 0.01 MiB 0.02 MiB 5 9 4 8 0.005958 ms SKIP 108.588949 mW SKIP 646972.960000 pJ SKIP
85 matmul/dynamic matmul/batched_lhs_broadcast PASS 0.093 s 0.059 s 0.00 MiB 0.00 MiB 0.01 MiB 5 0 4 0.001621 ms SKIP 91.421962 mW SKIP 148195.000000 pJ SKIP
86 matmul/huge_1024 matmul/batched_rhs_broadcast PASS 0.277 s 0.087 s 0.01 MiB 0.00 MiB 0.10 MiB 0.01 MiB 73 5 64 4 0.017522 ms SKIP 215.037402 mW SKIP 3767885.360000 pJ SKIP
87 matmul/left_constant matmul/dynamic PASS 0.123 s 0.058 s 0.00 MiB 0.01 MiB 0.00 MiB 5 4 0 0.005853 ms SKIP 108.861944 mW SKIP 637168.960000 pJ SKIP
88 matmul/matrix_vector matmul/huge_1024 PASS 0.150 s 0.164 s 0.52 MiB 0.01 MiB 0.78 MiB 0.10 MiB 168 73 173 64 0.384660 ms SKIP 202.131271 mW SKIP 77751814.880000 pJ SKIP
89 matmul/vector_matrix matmul/left_constant PASS 0.186 s 0.058 s 0.01 MiB 0.00 MiB 0.01 MiB 9 5 8 4 0.007409 ms SKIP 118.680243 mW SKIP 879301.920000 pJ SKIP
90 mul/after_conv matmul/matrix_vector PASS 0.127 s 0.100 s 0.00 MiB 0.52 MiB 0.00 MiB 0.78 MiB 4 168 3 173 0.005453 ms SKIP 107.639046 mW SKIP 586955.720000 pJ SKIP
91 mul/basic matmul/vector_matrix PASS 0.070 s 0.148 s 0.00 MiB 0.01 MiB 0.00 MiB 0.01 MiB 1 9 0 8 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
92 mul/channel_broadcast_1024 matmul/yolo_attention PASS 0.065 s 0.466 s 0.02 MiB 1.02 MiB 0.01 MiB 43.44 MiB 1 168 0 0.006913 ms SKIP 78.118038 mW SKIP 540030.000000 pJ SKIP
93 mul/leading_dimension_broadcast mul/after_conv PASS 0.108 s 0.059 s 0.00 MiB 0.00 MiB 1 4 0 3 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
94 mul/scalar_constant mul/basic PASS 0.082 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
95 pool/avg_basic mul/channel_broadcast_1024 PASS 0.070 s 0.051 s 0.00 MiB 0.02 MiB 0.00 MiB 0.01 MiB 1 0 0.011939 ms SKIP 78.022112 mW SKIP 931506.000000 pJ SKIP
96 pool/avg_ceil_mode mul/leading_dimension_broadcast PASS 0.107 s 0.071 s 0.00 MiB 0.00 MiB 1 0 0.004359 ms SKIP 78.033035 mW SKIP 340146.000000 pJ SKIP
97 pool/avg_explicit_padding mul/scalar_constant PASS 0.111 s 0.053 s 0.00 MiB 0.00 MiB 1 0 0.008822 ms SKIP 78.027205 mW SKIP 688356.000000 pJ SKIP
98 pool/avg_include_pad pool/avg_basic PASS 0.075 s 0.055 s 0.00 MiB 0.00 MiB 1 0 0.008506 ms SKIP 78.016929 mW SKIP 663612.000000 pJ SKIP
99 pool/avg_large_channels pool/avg_ceil_mode PASS 0.064 s 0.059 s 0.04 MiB 0.00 MiB 0.02 MiB 0.00 MiB 1 0 0.178249 ms SKIP 78.280327 mW SKIP 13953390.000000 pJ SKIP
100 pool/avg_non_uniform_stride pool/avg_explicit_padding PASS 0.068 s 0.059 s 0.00 MiB 0.00 MiB 1 0 0.014513 ms SKIP 78.016537 mW SKIP 1132254.000000 pJ SKIP
101 pool/avg_real_asymmetric_padding pool/avg_include_pad PASS 0.064 s 0.062 s 0.00 MiB 0.00 MiB 1 0 0.025206 ms SKIP 78.024756 mW SKIP 1966692.000000 pJ SKIP
102 pool/max_after_conv pool/avg_large_channels PASS 0.070 s 0.078 s 0.00 MiB 0.04 MiB 0.00 MiB 0.02 MiB 6 1 4 0 0.006452 ms SKIP 96.374606 mW SKIP 621808.960000 pJ SKIP
103 pool/max_basic pool/avg_non_uniform_stride PASS 0.063 s 0.057 s 0.00 MiB 0.00 MiB 3 1 0 0.001634 ms SKIP 92.132191 mW SKIP 150544.000000 pJ SKIP
104 pool/max_ceil_mode pool/avg_real_asymmetric_padding PASS 0.065 s 0.061 s 0.00 MiB 0.00 MiB 2 1 0 0.001297 ms SKIP 79.111025 mW SKIP 102607.000000 pJ SKIP
105 pool/max_global_style_kernel_equals_input pool/max_after_conv PASS 0.089 s 0.061 s 0.00 MiB 0.00 MiB 1 6 0 4 0.004366 ms SKIP 78.010994 mW SKIP 340596.000000 pJ SKIP
106 pool/max_non_square_kernel pool/max_basic PASS 0.103 s 0.060 s 0.00 MiB 0.00 MiB 4 3 0 0.003409 ms SKIP 93.253447 mW SKIP 317901.000000 pJ SKIP
107 pool/max_real_asymmetric_padding pool/max_ceil_mode PASS 0.088 s 0.055 s 0.00 MiB 0.00 MiB 4 2 0 0.003078 ms SKIP 93.124756 mW SKIP 286638.000000 pJ SKIP
108 pool/max_same_upper pool/max_global_style_kernel_equals_input PASS 0.065 s 0.100 s 0.00 MiB 0.00 MiB 3 1 0 0.003024 ms SKIP 92.095238 mW SKIP 278496.000000 pJ SKIP
109 pool/max_stride2_multichannel pool/max_non_square_kernel PASS 0.086 s 0.062 s 0.00 MiB 0.00 MiB 3 4 0 0.004012 ms SKIP 92.269192 mW SKIP 370184.000000 pJ SKIP
110 reduce_mean/4d_spatial pool/max_real_asymmetric_padding PASS 0.061 s 0.062 s 0.00 MiB 0.00 MiB 3 4 0 0.000321 ms SKIP 92.448598 mW SKIP 29676.000000 pJ SKIP
111 reduce_mean/4d_spatial_keepdims_0 pool/max_same_upper PASS 0.067 s 0.059 s 0.00 MiB 0.00 MiB 4 3 0 0.000655 ms SKIP 94.352672 mW SKIP 61801.000000 pJ SKIP
112 reduce_mean/after_conv pool/max_stride2_multichannel PASS 0.115 s 0.056 s 0.00 MiB 0.00 MiB 5 3 3 0 0.005342 ms SKIP 106.951089 mW SKIP 571332.720000 pJ SKIP
113 reduce_mean/all_axes_keepdims_0 reduce_mean/4d_spatial PASS 0.103 s 0.053 s 0.00 MiB 0.00 MiB 2 3 0 0.000391 ms SKIP 79.237852 mW SKIP 30982.000000 pJ SKIP
114 reduce_mean/all_axes_keepdims_1 reduce_mean/4d_spatial_keepdims_0 PASS 0.086 s 0.070 s 0.00 MiB 0.00 MiB 1 4 0 0.000221 ms SKIP 78.217195 mW SKIP 17286.000000 pJ SKIP
115 reduce_mean/basic reduce_mean/after_conv PASS 0.099 s 0.063 s 0.00 MiB 0.00 MiB 4 5 0 3 0.000373 ms SKIP 93.514745 mW SKIP 34881.000000 pJ SKIP
116 reduce_mean/channel_axis_nchw reduce_mean/all_axes_keepdims_0 PASS 0.078 s 0.053 s 0.03 MiB 0.00 MiB 0.02 MiB 0.00 MiB 4 2 0 0.164926 ms SKIP 93.596631 mW SKIP 15436518.000000 pJ SKIP
117 reduce_mean/keepdims_0 reduce_mean/all_axes_keepdims_1 PASS 0.073 s 0.051 s 0.00 MiB 0.00 MiB 5 1 0 0.000748 ms SKIP 91.401070 mW SKIP 68368.000000 pJ SKIP
118 reduce_mean/large_dimension_1024 reduce_mean/basic PASS 0.115 s 0.052 s 0.01 MiB 0.00 MiB 0.00 MiB 1 4 0 0.002785 ms SKIP 78.017235 mW SKIP 217278.000000 pJ SKIP
119 reduce_mean/legacy_axes_1_2_keepdims_1 reduce_mean/channel_axis_nchw PASS 0.074 s 0.053 s 0.00 MiB 0.03 MiB 0.00 MiB 0.02 MiB 2 4 0 0.000271 ms SKIP 79.354244 mW SKIP 21505.000000 pJ SKIP
120 reduce_mean/legacy_axis1_keepdims_0 reduce_mean/keepdims_0 PASS 0.111 s 0.053 s 0.00 MiB 0.00 MiB 9 5 0 0.001986 ms SKIP 92.501511 mW SKIP 183708.000000 pJ SKIP
121 reduce_mean/legacy_axis1_keepdims_1 reduce_mean/large_dimension_1024 PASS 0.086 s 0.061 s 0.00 MiB 0.01 MiB 0.00 MiB 8 1 0 0.001373 ms SKIP 94.559359 mW SKIP 129830.000000 pJ SKIP
122 reduce_mean/legacy_empty_axes_noop reduce_mean/legacy_axes_1_2_keepdims_1 PASS 0.115 s 0.052 s 0.00 MiB 0.00 MiB 1 2 0 0.000221 ms SKIP 78.217195 mW SKIP 17286.000000 pJ SKIP
123 reduce_mean/legacy_nchw_spatial reduce_mean/legacy_axis1_keepdims_0 PASS 0.090 s 0.052 s 0.00 MiB 0.00 MiB 3 9 0 0.000321 ms SKIP 92.448598 mW SKIP 29676.000000 pJ SKIP
124 reduce_mean/legacy_negative_axis reduce_mean/legacy_axis1_keepdims_1 PASS 0.110 s 0.058 s 0.00 MiB 0.00 MiB 6 8 0 0.000553 ms SKIP 93.520796 mW SKIP 51717.000000 pJ SKIP
125 reduce_mean/legacy_reduce_all_keepdims_1 reduce_mean/legacy_empty_axes_noop PASS 0.107 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000221 ms SKIP 78.217195 mW SKIP 17286.000000 pJ SKIP
126 reduce_mean/negative_axis reduce_mean/legacy_nchw_spatial PASS 0.120 s 0.058 s 0.00 MiB 0.00 MiB 6 3 0 0.000553 ms SKIP 93.520796 mW SKIP 51717.000000 pJ SKIP
127 relu/4d reduce_mean/legacy_negative_axis PASS 0.116 s 0.059 s 0.00 MiB 0.00 MiB 1 6 0 0.000521 ms SKIP 78.184261 mW SKIP 40734.000000 pJ SKIP
128 relu/after_conv reduce_mean/legacy_reduce_all_keepdims_1 PASS 0.078 s 0.068 s 0.00 MiB 0.00 MiB 3 1 3 0 0.004956 ms SKIP 106.998935 mW SKIP 530286.720000 pJ SKIP
129 relu/after_gemm reduce_mean/negative_axis PASS 0.077 s 0.065 s 0.01 MiB 0.00 MiB 0.01 MiB 0.00 MiB 5 6 4 0 0.007513 ms SKIP 105.158653 mW SKIP 790056.960000 pJ SKIP
130 relu/basic relu/4d PASS 0.066 s 0.00 MiB 0.00 MiB 1 0 0.000221 ms SKIP 78.217195 mW SKIP 17286.000000 pJ SKIP
131 reshape/4d_to_2d_flatten relu/after_conv PASS 0.098 s 0.075 s 0.00 MiB 0.00 MiB 1 3 0 3 0.000258 ms SKIP 78.279070 mW SKIP 20196.000000 pJ SKIP
132 reshape/infer_dim_minus_one relu/after_gemm PASS 0.105 s 0.079 s 0.00 MiB 0.01 MiB 0.00 MiB 0.01 MiB 1 5 0 4 0.000162 ms SKIP 78.296296 mW SKIP 12684.000000 pJ SKIP
133 reshape/same_rank relu/basic PASS 0.057 s 0.049 s 0.00 MiB 0.00 MiB 1 0 0.000162 ms SKIP 78.296296 mW SKIP 12684.000000 pJ SKIP
134 reshape/zero_copies_input_dim reshape/4d_to_2d_flatten PASS 0.090 s 0.049 s 0.00 MiB 0.00 MiB 1 0 0.000162 ms SKIP 78.296296 mW SKIP 12684.000000 pJ SKIP
135 resize/height_only reshape/infer_dim_minus_one PASS 0.104 s 0.061 s 0.00 MiB 0.00 MiB 1 0 0.000795 ms SKIP 78.060377 mW SKIP 62058.000000 pJ SKIP
136 resize/nearest_2x reshape/same_rank PASS 0.104 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.001422 ms SKIP 78.033755 mW SKIP 110964.000000 pJ SKIP
137 resize/nearest_downsample reshape/zero_copies_input_dim PASS 0.059 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.000481 ms SKIP 78.099792 mW SKIP 37566.000000 pJ SKIP
138 resize/non_uniform resize/height_only PASS 0.063 s 0.079 s 0.00 MiB 0.00 MiB 1 4 0 0.002108 ms SKIP 78.034156 mW SKIP 164496.000000 pJ SKIP
139 resize/width_only resize/nearest_2x PASS 0.068 s 0.053 s 0.00 MiB 0.00 MiB 1 4 0 0.000792 ms SKIP 78.060606 mW SKIP 61824.000000 pJ SKIP
140 resize/with_sizes resize/nearest_downsample PASS 0.056 s 0.055 s 0.00 MiB 0.00 MiB 1 2 0 0.000951 ms SKIP 78.050473 mW SKIP 74226.000000 pJ SKIP
141 sigmoid/4d resize/non_uniform PASS 0.104 s 0.053 s 0.00 MiB 0.00 MiB 1 6 0 0.000521 ms SKIP 78.184261 mW SKIP 40734.000000 pJ SKIP
142 sigmoid/after_gemm resize/width_only PASS 0.118 s 0.052 s 0.01 MiB 0.00 MiB 0.01 MiB 0.00 MiB 5 2 4 0 0.007513 ms SKIP 105.158653 mW SKIP 790056.960000 pJ SKIP
143 sigmoid/basic resize/with_sizes PASS 0.083 s 0.053 s 0.00 MiB 0.00 MiB 1 3 0 0.000221 ms SKIP 78.217195 mW SKIP 17286.000000 pJ SKIP
144 slice/2d_basic sigmoid/4d PASS 0.067 s 0.077 s 0.00 MiB 0.00 MiB 1 0 0.000242 ms SKIP 78.297521 mW SKIP 18948.000000 pJ SKIP
145 slice/after_conv sigmoid/after_gemm PASS 0.111 s 0.058 s 0.00 MiB 0.01 MiB 0.01 MiB 7 5 6 4 0.011296 ms SKIP 118.190765 mW SKIP 1335082.880000 pJ SKIP
146 slice/default_axes sigmoid/basic PASS 0.065 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.000242 ms SKIP 78.297521 mW SKIP 18948.000000 pJ SKIP
147 slice/large_channel_1024 slice/2d_basic PASS 0.060 s 0.055 s 0.01 MiB 0.00 MiB 0.00 MiB 1 0 0.002832 ms SKIP 78.144068 mW SKIP 221304.000000 pJ SKIP
148 slice/nchw_spatial_crop slice/after_conv PASS 0.067 s 0.061 s 0.00 MiB 0.00 MiB 0.01 MiB 1 7 0 6 0.001302 ms SKIP 78.239631 mW SKIP 101868.000000 pJ SKIP
149 slice/negative_axis slice/default_axes PASS 0.098 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000562 ms SKIP 78.298932 mW SKIP 44004.000000 pJ SKIP
150 slice/negative_indices slice/large_channel_1024 PASS 0.106 s 0.068 s 0.00 MiB 0.01 MiB 0.00 MiB 1 0 0.000322 ms SKIP 78.298137 mW SKIP 25212.000000 pJ SKIP
151 slice/step2 slice/nchw_spatial_crop PASS 0.068 s 0.052 s 0.00 MiB 0.00 MiB 1 0 0.002042 ms SKIP 78.293830 mW SKIP 159876.000000 pJ SKIP
152 softmax/3d_last_axis slice/negative_axis PASS 0.072 s 0.051 s 0.00 MiB 0.00 MiB 1 0 UNSUPPORTED SKIP UNSUPPORTED SKIP UNSUPPORTED SKIP
153 softmax/basic slice/negative_indices PASS 0.088 s 0.048 s 0.00 MiB 0.00 MiB 1 0 UNSUPPORTED SKIP UNSUPPORTED SKIP UNSUPPORTED SKIP
154 softmax/channel_axis slice/step2 PASS 0.110 s 0.051 s 0.00 MiB 0.00 MiB 1 0 UNSUPPORTED SKIP UNSUPPORTED SKIP UNSUPPORTED SKIP
155 softmax/large_dimension_1024 softmax/3d_last_axis PASS 0.104 s 0.056 s 0.01 MiB 0.00 MiB 0.01 MiB 0.00 MiB 1 0 UNSUPPORTED SKIP UNSUPPORTED SKIP UNSUPPORTED SKIP
156 softmax/negative_axis softmax/basic PASS 0.118 s 0.051 s 0.00 MiB 0.00 MiB 1 0 UNSUPPORTED SKIP UNSUPPORTED SKIP UNSUPPORTED SKIP
157 split/basic softmax/channel_axis PASS 0.077 s 0.056 s 0.00 MiB 0.00 MiB 1 0 0.000403 ms SKIP 78.297767 mW SKIP 31554.000000 pJ SKIP
158 split/equal_three_way softmax/large_dimension_1024 PASS 0.107 s 0.048 s 0.00 MiB 0.01 MiB 0.00 MiB 0.01 MiB 1 0 0.000564 ms SKIP 78.297872 mW SKIP 44160.000000 pJ SKIP
159 split/negative_axis softmax/negative_axis PASS 0.111 s 0.053 s 0.00 MiB 0.00 MiB 1 0 0.001083 ms SKIP 78.288089 mW SKIP 84786.000000 pJ SKIP
160 split/uneven_channel_axis_4d split/basic PASS 0.061 s 0.050 s 0.00 MiB 0.00 MiB 1 0 0.000242 ms SKIP 78.297521 mW SKIP 18948.000000 pJ SKIP
161 sub/after_gemm split/equal_three_way PASS 0.068 s 0.049 s 0.01 MiB 0.00 MiB 0.01 MiB 0.00 MiB 5 1 4 0 0.007784 ms SKIP 104.703618 mW SKIP 815012.960000 pJ SKIP
162 sub/basic split/negative_axis PASS 0.063 s 0.067 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
163 sub/broadcast_row split/uneven_channel_axis_4d PASS 0.075 s 0.053 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
164 sub/channel_broadcast_1024 sub/after_gemm PASS 0.052 s 0.059 s 0.02 MiB 0.01 MiB 0.01 MiB 1 5 0 4 0.006913 ms SKIP 78.118038 mW SKIP 540030.000000 pJ SKIP
165 sub/constant_lhs_broadcast sub/basic PASS 0.057 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000322 ms SKIP 78.223602 mW SKIP 25188.000000 pJ SKIP
166 sub/leading_dimension_broadcast sub/broadcast_row PASS 0.053 s 0.051 s 0.00 MiB 0.00 MiB 1 0 0.000323 ms SKIP 78.222910 mW SKIP 25266.000000 pJ SKIP
167 sub/channel_broadcast_1024 PASS 0.056 s 0.02 MiB 0.01 MiB 1 0 SKIP SKIP SKIP
168 sub/constant_lhs_broadcast PASS 0.052 s 0.00 MiB 0.00 MiB 1 0 SKIP SKIP SKIP
169 sub/leading_dimension_broadcast PASS 0.066 s 0.00 MiB 0.00 MiB 1 0 SKIP SKIP SKIP
@@ -0,0 +1,182 @@
#!/usr/bin/env python3
"""Save exact attention-tap arrays and quantify the MatMul error sources."""
import argparse
import hashlib
import json
import sys
from pathlib import Path
import numpy as np
import onnx
from onnx import numpy_helper
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT / "validation"))
from raptor_validation.onnx_utils import onnx_io # noqa: E402
from raptor_validation.validate_one import ( # noqa: E402
parse_pim_simulator_outputs,
sanitize_output_name,
)
TAP_NAMES = {
"v": "/model.10/m/m.0/attn/Split_output_2",
"raw": "/model.10/m/m.0/attn/MatMul_output_0",
"scaled": "/model.10/m/m.0/attn/Mul_output_0",
"rhs": "/model.10/m/m.0/attn/Transpose_1_output_0",
"c": "/model.10/m/m.0/attn/MatMul_1_output_0",
}
ABSOLUTE_TOLERANCE = 1e-3
RELATIVE_TOLERANCE = 1e-5
def sha256(path):
digest = hashlib.sha256()
with Path(path).open("rb") as stream:
for block in iter(lambda: stream.read(1 << 20), b""):
digest.update(block)
return digest.hexdigest()
def metric(actual, expected):
difference = np.abs(actual.astype(np.float64) - expected.astype(np.float64))
allowed = ABSOLUTE_TOLERANCE + RELATIVE_TOLERANCE * np.abs(expected.astype(np.float64))
return {
"max_abs": float(np.max(difference)),
"mean_abs": float(np.mean(difference)),
"rms": float(np.sqrt(np.mean(np.square(difference)))),
"elements_over_validator_limit": int(np.count_nonzero(difference > allowed)),
}
def f32_matmul(lhs, rhs):
return np.matmul(lhs.astype(np.float32), rhs.astype(np.float32)).astype(np.float32)
def f64_matmul(lhs, rhs):
return np.matmul(lhs.astype(np.float64), rhs.astype(np.float64)).astype(np.float64)
def load_constant(model, output_name):
for initializer in model.graph.initializer:
if initializer.name == output_name:
return float(numpy_helper.to_array(initializer).reshape(-1)[0])
for node in model.graph.node:
if output_name not in node.output:
continue
for attribute in node.attribute:
if attribute.name == "value" and attribute.HasField("t"):
return float(numpy_helper.to_array(attribute.t).reshape(-1)[0])
raise ValueError(f"could not find ONNX Constant producing {output_name}")
def load_arrays(workspace, model_path):
model = onnx.load(model_path)
descriptors = onnx_io(model_path)
output_descriptors = {name: (index, dtype, shape) for index, name, dtype, shape in descriptors[1]}
missing = sorted(set(TAP_NAMES.values()) - set(output_descriptors))
if missing:
raise ValueError("tap model is missing outputs: " + ", ".join(missing))
sim_arrays = parse_pim_simulator_outputs(
workspace / "simulation" / "out.bin", descriptors[1]
)
reference = {}
simulated = {}
input_files = {}
for key, name in TAP_NAMES.items():
index, _dtype, shape = output_descriptors[name]
csv_path = workspace / "outputs" / f"output{index}_{sanitize_output_name(name)}.csv"
reference[key] = np.loadtxt(csv_path, delimiter=",", dtype=np.float32).reshape(shape)
simulated[key] = np.asarray(sim_arrays[index], dtype=np.float32).reshape(shape)
input_files[key] = csv_path
return reference, simulated, input_files
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--workspace", type=Path, required=True,
help="validator workspace containing inputs, outputs, and simulation")
parser.add_argument("--model", type=Path, required=True, help="five-output ONNX tap model")
parser.add_argument("--output-dir", type=Path, required=True,
help="directory for arrays.npz, metadata.json, and decomposition.json")
args = parser.parse_args()
args.output_dir.mkdir(parents=True, exist_ok=True)
model = onnx.load(args.model)
reference, simulated, source_files = load_arrays(args.workspace, args.model)
scale = np.float32(load_constant(model, "/model.10/m/m.0/attn/Constant_1_output_0"))
ref_v, ref_raw, ref_scaled, ref_rhs, ref_c = (reference[key] for key in ("v", "raw", "scaled", "rhs", "c"))
sim_v, sim_raw, sim_scaled, sim_rhs, sim_c = (simulated[key] for key in ("v", "raw", "scaled", "rhs", "c"))
ref_score_transpose = np.swapaxes(ref_scaled, -1, -2)
sim_score_transpose = np.swapaxes(sim_scaled, -1, -2)
ref_ss_f32 = f32_matmul(ref_v, ref_rhs)
sim_ss_f32 = f32_matmul(sim_v, sim_rhs)
ref_ss_f64 = f64_matmul(ref_v, ref_rhs)
sim_ss_f64 = f64_matmul(sim_v, sim_rhs)
ref_split_f32 = (f32_matmul(ref_v, np.swapaxes(ref_raw, -1, -2)) * scale).astype(np.float32)
sim_split_f32 = (f32_matmul(sim_v, np.swapaxes(sim_raw, -1, -2)) * scale).astype(np.float32)
ref_split_f64 = f64_matmul(ref_v, np.swapaxes(ref_raw, -1, -2)) * np.float64(scale)
sim_split_f64 = f64_matmul(sim_v, np.swapaxes(sim_raw, -1, -2)) * np.float64(scale)
arrays = {
**{f"ref_{key}": value for key, value in reference.items()},
**{f"sim_{key}": value for key, value in simulated.items()},
"ref_ss_f32": ref_ss_f32,
"sim_ss_f32": sim_ss_f32,
"ref_ss_f64": ref_ss_f64,
"sim_ss_f64": sim_ss_f64,
"ref_split_f32": ref_split_f32,
"sim_split_f32": sim_split_f32,
"ref_split_f64": ref_split_f64,
"sim_split_f64": sim_split_f64,
}
arrays_path = args.output_dir / "arrays.npz"
np.savez_compressed(arrays_path, **arrays)
metrics = {
"validator_policy": {
"absolute_tolerance": ABSOLUTE_TOLERANCE,
"relative_tolerance": RELATIVE_TOLERANCE,
},
"scale": float(scale),
"shape": list(ref_c.shape),
"tap_differences": {key: metric(simulated[key], reference[key]) for key in TAP_NAMES},
"rhs_transpose_consistency": metric(ref_rhs, ref_score_transpose),
"sim_rhs_transpose_consistency": metric(sim_rhs, sim_score_transpose),
"c_sim_vs_ss_f32": metric(sim_c, ref_ss_f32),
"c_ref_vs_ss_f32": metric(ref_c, ref_ss_f32),
"c_sim_vs_simulated_inputs_ss_f32": metric(sim_c, sim_ss_f32),
"v_drift_only": metric(f32_matmul(sim_v, ref_rhs), ref_ss_f32),
"rhs_drift_only": metric(f32_matmul(ref_v, sim_rhs), ref_ss_f32),
"joint_input_drift": metric(sim_ss_f32, ref_ss_f32),
"scale_reassociation_reference": metric(ref_split_f32, ref_ss_f32),
"scale_reassociation_simulated": metric(sim_split_f32, sim_ss_f32),
"reference_accumulation_f32_vs_f64": metric(ref_ss_f32, ref_ss_f64),
"simulated_accumulation_f32_vs_f64": metric(sim_ss_f32, sim_ss_f64),
"split_accumulation_reference_f32_vs_f64": metric(ref_split_f32, ref_split_f64),
"split_accumulation_simulated_f32_vs_f64": metric(sim_split_f32, sim_split_f64),
}
decomposition_path = args.output_dir / "decomposition.json"
decomposition_path.write_text(json.dumps(metrics, indent=2) + "\n", encoding="utf-8")
metadata = {
"model": str(args.model),
"model_sha256": sha256(args.model),
"workspace": str(args.workspace),
"arrays_sha256": sha256(arrays_path),
"source_sha256": {key: sha256(path) for key, path in source_files.items()},
"simulator_output_sha256": sha256(args.workspace / "simulation" / "out.bin"),
"input_sha256": sha256(args.workspace / "inputs" / "in0.csv"),
"outputs": TAP_NAMES,
"arrays": {key: {"dtype": str(value.dtype), "shape": list(value.shape)} for key, value in arrays.items()},
}
(args.output_dir / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8")
print(json.dumps(metrics, indent=2))
if __name__ == "__main__":
main()