diff --git a/.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md b/.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md index 82a86ba..8c731e2 100644 --- a/.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md +++ b/.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md @@ -44,6 +44,19 @@ input size and target linear or sublinear time and space. Avoid repeated full-IR walks, nested scans, and per-operation recomputation when indexing, caching, or a single traversal can express the same behavior. +Bufferization cost scales with both the number of operations and the number of +MLIR values it receives. Upstream lowering and scheduling must therefore keep +repeated work in compact structured operations, such as statically evaluable +loops and batches, until after bufferization. Do not compensate for avoidable +pre-bufferization expansion by weakening, partitioning, or special-casing the +bufferization analysis; measure both operation and value counts at its input. +Compactness must also be preserved after bufferization so that liveness, +verification, memory planning, and other downstream passes do not repeat work +per logical lane or iteration. Keep structured loops and batches intact until +PIM ISA code generation is forced to scalarize them into concrete per-core +instructions; earlier expansion requires an explicit semantic necessity and +before/after operation and value counts. + When linear-or-better complexity is not possible, use the lowest justified complexity and report the actual time and space Big-O, the input variable, and why a lower bound is not practical. Include that cost in the final report; do diff --git a/src/PIM/CMakeLists.txt b/src/PIM/CMakeLists.txt index e4af4ab..cd94643 100644 --- a/src/PIM/CMakeLists.txt +++ b/src/PIM/CMakeLists.txt @@ -121,6 +121,7 @@ add_pim_library(OMPIMAccel OMPimCommon OMPimBufferization OMPimHostConstantFolding + OMPimInstructionSelection OMPimVerification MLIRTensorInferTypeOpInterfaceImpl ) diff --git a/src/PIM/Compiler/CMakeLists.txt b/src/PIM/Compiler/CMakeLists.txt index 95e2015..7cf2f2d 100644 --- a/src/PIM/Compiler/CMakeLists.txt +++ b/src/PIM/Compiler/CMakeLists.txt @@ -31,6 +31,7 @@ add_pim_library(OMPimCompilerUtils OMPimCommon OMPimBufferization OMPimHostConstantFolding + OMPimInstructionSelection OMPimLocalMemoryPlanning OMPimVerification OMPimPasses diff --git a/src/PIM/Compiler/PimCodeGen.cpp b/src/PIM/Compiler/PimCodeGen.cpp index fa84349..e6980f1 100644 --- a/src/PIM/Compiler/PimCodeGen.cpp +++ b/src/PIM/Compiler/PimCodeGen.cpp @@ -108,6 +108,23 @@ static Operation* getDiagnosticAnchor(mlir::Value value) { return nullptr; } +static bool isZeroSplatGlobal(mlir::Value value) { + auto getGlobalOp = value.getDefiningOp(); + auto moduleOp = getGlobalOp ? getGlobalOp->getParentOfType() : ModuleOp(); + auto globalOp = moduleOp ? lookupGlobalForGetGlobal(moduleOp, getGlobalOp) : memref::GlobalOp(); + if (!globalOp || !globalOp.getConstant() || !globalOp.getInitialValue()) + return false; + auto denseAttr = dyn_cast(*globalOp.getInitialValue()); + if (!denseAttr || !denseAttr.isSplat()) + return false; + Attribute valueAttr = denseAttr.getSplatValue(); + if (auto floatAttr = dyn_cast(valueAttr)) + return floatAttr.getValue().isZero() && !floatAttr.getValue().isNegative(); + if (auto integerAttr = dyn_cast(valueAttr)) + return integerAttr.getValue().isZero(); + return false; +} + // PIM instruction immediates are serialized as signed int32_t fields today // (`sldi` goes through checkedI32OrCrash), so local addresses must stay within // the non-negative int32_t range. @@ -541,14 +558,23 @@ void PimCodeGen::emitMemCopyOp(pim_binary::Opcode opcode, size_t rs1Offset, size_t size, StringRef sizeFieldName) const { - setupRdRs1(rdAddr, rdOffset, rs1Addr, rs1Offset); - pim_binary::InstructionRecord instruction; instruction.opcode = opcode; instruction.rd = 0; instruction.r1 = 1; - instruction.generic1 = 0; - instruction.generic2 = 0; + if (rdOffset == rs1Offset) { + setupRdRs1(rdAddr, 0, rs1Addr, 0); + instruction.generic1 = rdOffset == 0 ? 0 : 3; + instruction.generic2 = pim::checkedI32OrCrash(rdOffset, "shared address offset"); + } else if (rdOffset != 0) { + setupRdRs1(rdAddr, 0, rs1Addr, rs1Offset); + instruction.generic1 = 1; + instruction.generic2 = pim::checkedI32OrCrash(rdOffset, "rd address offset"); + } else { + setupRdRs1(rdAddr, 0, rs1Addr, 0); + instruction.generic1 = rs1Offset == 0 ? 0 : 2; + instruction.generic2 = pim::checkedI32OrCrash(rs1Offset, "rs1 address offset"); + } instruction.generic3 = pim::checkedI32OrCrash(size, sizeFieldName); emitInstruction(instruction); } @@ -584,6 +610,16 @@ void PimCodeGen::codeGenLoadOp(pim::PimMemCopyHostToDevOp loadOp, const StaticVa auto hostSourceOffset = indexOf(loadOp.getHostSourceOffset(), knowledge); assert(succeeded(deviceTargetOffset) && succeeded(hostSourceOffset) && "pim.memcp_hd offsets must be statically resolvable during codegen"); + if (isZeroSplatGlobal(loadOp.getHostSource())) { + setupRd(addressOf(loadOp.getDeviceTarget(), knowledge), *deviceTargetOffset); + pim_binary::InstructionRecord instruction; + instruction.opcode = pim_binary::Opcode::lldi; + instruction.rd = 0; + instruction.r2OrImm = 0; + instruction.generic3 = loadOp.getSize(); + emitInstruction(instruction); + return; + } emitMemCopyOp(pim_binary::Opcode::ld, addressOf(loadOp.getDeviceTarget(), knowledge), *deviceTargetOffset, @@ -619,6 +655,26 @@ void PimCodeGen::codeGenLmvOp(pim::PimMemCopyOp lmvOp, const StaticValueKnowledg "len"); } +void PimCodeGen::codeGenVMVOp(pim::PimVMVOp vmvOp, const StaticValueKnowledge& knowledge) const { + auto targetOffset = indexOf(vmvOp.getTargetOffset(), knowledge); + auto sourceOffset = indexOf(vmvOp.getSourceOffset(), knowledge); + auto sourceStride = indexOf(vmvOp.getSourceStride(), knowledge); + assert(succeeded(targetOffset) && succeeded(sourceOffset) && succeeded(sourceStride) + && "pim.vmv operands must be statically resolvable during codegen"); + auto sourceType = cast(vmvOp.getSource().getType()); + int32_t bitwidth = getVectorElementBitwidthOrCrash(sourceType); + ensureVectorBitwidth(bitwidth, bitwidth); + setupRdRs1Rs2(addressOf(vmvOp.getTarget(), knowledge), *targetOffset, + addressOf(vmvOp.getSource(), knowledge), *sourceOffset, 0, *sourceStride); + pim_binary::InstructionRecord instruction; + instruction.opcode = pim_binary::Opcode::vmv; + instruction.rd = 0; + instruction.r1 = 1; + instruction.r2OrImm = 2; + instruction.generic3 = vmvOp.getLength(); + emitInstruction(instruction); +} + void PimCodeGen::codeGenReceiveOp(pim::PimReceiveOp receiveOp, const StaticValueKnowledge& knowledge) const { auto outputOffset = indexOf(receiveOp.getOutputOffset(), knowledge); auto sourceCoreId = indexOf(receiveOp.getSourceCoreId(), knowledge); @@ -723,89 +779,6 @@ void PimCodeGen::emitUnaryVectorOp(pim_binary::Opcode opcode, emitInstruction(instruction); } -void PimCodeGen::codeGenTransposeOp(const CompiledTransposePlan& plan, const StaticValueKnowledge& knowledge) const { - auto srcAddr = addressOf(plan.source, knowledge); - auto dstAddr = addressOf(plan.destination, knowledge); - - size_t maxElementOffset = plan.totalBytes == 0 ? 0 : plan.totalBytes - plan.elementBytes; - int32_t maxSourceAddress = pim::checkedI32OrCrash( - pim::checkedAddOrCrash(srcAddr, maxElementOffset, "transpose source address"), "transpose source address"); - int32_t maxDestinationAddress = - pim::checkedI32OrCrash(pim::checkedAddOrCrash(dstAddr, maxElementOffset, "transpose destination address"), - "transpose destination address"); - (void) maxSourceAddress; - (void) maxDestinationAddress; - - pim_binary::InstructionRecord copyInstruction; - copyInstruction.opcode = pim_binary::Opcode::lmv; - copyInstruction.rd = 0; - copyInstruction.r1 = 1; - - size_t maxRunElements = static_cast(std::numeric_limits::max()) / plan.elementBytes; - auto emitRun = [&](size_t sourceStart, size_t destinationStart, size_t runLength) { - while (runLength != 0) { - size_t chunkElements = std::min(runLength, maxRunElements); - // totalBytes was checked when the plan was compiled, so these bounded - // products cannot overflow. - size_t sourceOffset = sourceStart * plan.elementBytes; - size_t destinationOffset = destinationStart * plan.elementBytes; - size_t byteSize = chunkElements * plan.elementBytes; - assert(sourceOffset <= plan.totalBytes - byteSize && destinationOffset <= plan.totalBytes - byteSize); - genSetRegisterImmediate(0, static_cast(dstAddr + destinationOffset)); - genSetRegisterImmediate(1, static_cast(srcAddr + sourceOffset)); - copyInstruction.generic3 = static_cast(byteSize); - emitInstruction(copyInstruction); - sourceStart += chunkElements; - destinationStart += chunkElements; - runLength -= chunkElements; - } - }; - - if (plan.storagePreserving) { - emitRun(0, 0, plan.totalElements); - return; - } - - size_t rank = plan.sourceShape.size(); - SmallVector sourceIndices(rank, 0); - size_t destinationFlat = 0; - size_t runSourceStart = 0; - size_t runDestinationStart = 0; - size_t runLength = 0; - - for (size_t sourceFlat = 0; sourceFlat < plan.totalElements; ++sourceFlat) { - if (runLength != 0 && destinationFlat != runDestinationStart + runLength) { - emitRun(runSourceStart, runDestinationStart, runLength); - runSourceStart = sourceFlat; - runDestinationStart = destinationFlat; - runLength = 0; - } - if (runLength == 0) { - runSourceStart = sourceFlat; - runDestinationStart = destinationFlat; - } - ++runLength; - - if (runLength == maxRunElements) { - emitRun(runSourceStart, runDestinationStart, runLength); - runLength = 0; - } - - if (sourceFlat + 1 == plan.totalElements) - break; - for (size_t sourceDim = rank; sourceDim-- > 0;) { - destinationFlat += plan.destinationStrides[plan.destinationDimensionForSource[sourceDim]]; - if (++sourceIndices[sourceDim] < plan.sourceShape[sourceDim]) - break; - sourceIndices[sourceDim] = 0; - destinationFlat -= plan.destinationRewinds[sourceDim]; - } - } - - if (runLength != 0) - emitRun(runSourceStart, runDestinationStart, runLength); -} - static SmallVector collectTopLevelCoreLikeOps(func::FuncOp funcOp) { SmallVector coreLikeOps; for (Operation& op : funcOp.getBody().front()) @@ -1015,6 +988,7 @@ static LogicalResult executeCompiledCorePlan( coreCodeGen.codeGenStoreOp(cast(node.op), knowledge); break; case CompiledCoreOpKind::Lmv: coreCodeGen.codeGenLmvOp(cast(node.op), knowledge); break; + case CompiledCoreOpKind::VMV: coreCodeGen.codeGenVMVOp(cast(node.op), knowledge); break; case CompiledCoreOpKind::Receive: coreCodeGen.codeGenReceiveOp(cast(node.op), knowledge); break; case CompiledCoreOpKind::Send: coreCodeGen.codeGenSendOp(cast(node.op), knowledge); break; case CompiledCoreOpKind::Concat: coreCodeGen.codeGenConcatOp(cast(node.op), knowledge); break; @@ -1024,7 +998,6 @@ static LogicalResult executeCompiledCorePlan( else return failure(); break; - case CompiledCoreOpKind::Transpose: coreCodeGen.codeGenTransposeOp(*node.transposePlan, knowledge); break; case CompiledCoreOpKind::VVAdd: emitBinary(cast(node.op), pim_binary::Opcode::vvadd); break; case CompiledCoreOpKind::VVSub: emitBinary(cast(node.op), pim_binary::Opcode::vvsub); break; case CompiledCoreOpKind::VVMul: emitBinary(cast(node.op), pim_binary::Opcode::vvmul); break; @@ -1224,7 +1197,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std:: for (auto [slot, fileName] : llvm::enumerate(weightFiles)) { xbarsPerGroup.push_back(weights[slot].shape[1] / static_cast(crossbarSize)); - std::string sourcePath = outputDirPath + "/weights/" + fileName; + std::string sourcePath = "../weights/" + fileName; std::string targetPath = coreWeightsDirPath + "/crossbar_" + std::to_string(slot) + ".bin"; sys::fs::remove(targetPath); if (auto error = sys::fs::create_link(sourcePath, targetPath)) { diff --git a/src/PIM/Compiler/PimCodeGen.hpp b/src/PIM/Compiler/PimCodeGen.hpp index ab6c2fb..7c045c2 100644 --- a/src/PIM/Compiler/PimCodeGen.hpp +++ b/src/PIM/Compiler/PimCodeGen.hpp @@ -25,7 +25,6 @@ namespace onnx_mlir { struct CompiledCoreProgram; -struct CompiledTransposePlan; class PimInstructionWriter; struct MemEntry { @@ -214,6 +213,7 @@ public: void codeGenLoadOp(pim::PimMemCopyHostToDevOp loadOp, const StaticValueKnowledge& knowledge) const; void codeGenStoreOp(pim::PimMemCopyDevToHostOp storeOp, const StaticValueKnowledge& knowledge) const; void codeGenLmvOp(pim::PimMemCopyOp lmvOp, const StaticValueKnowledge& knowledge) const; + void codeGenVMVOp(pim::PimVMVOp vmvOp, const StaticValueKnowledge& knowledge) const; void codeGenReceiveOp(pim::PimReceiveOp receiveOp, const StaticValueKnowledge& knowledge) const; void codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge& knowledge) const; @@ -222,7 +222,6 @@ public: template void codeGenMVMLikeOp(size_t mvmId, MVMTy mvmLikeOp, bool transposeMatrix, const StaticValueKnowledge& knowledge); - void codeGenTransposeOp(const CompiledTransposePlan& plan, const StaticValueKnowledge& knowledge) const; }; OnnxMlirCompilerErrorCodes compileToPimCode(mlir::ModuleOp& moduleOpRef, std::string& outputDirName); diff --git a/src/PIM/Compiler/PimCompilerUtils.cpp b/src/PIM/Compiler/PimCompilerUtils.cpp index 9fc8580..6e64791 100644 --- a/src/PIM/Compiler/PimCompilerUtils.cpp +++ b/src/PIM/Compiler/PimCompilerUtils.cpp @@ -281,6 +281,7 @@ void addPassesPim(OwningOpRef& module, verifyExplicitPimCoreCount(); if (pimOnlyCodegen) { + pm.addPass(createPimInstructionSelectionPass()); pm.addPass(createPimLocalMemoryPlanningPass()); pm.addPass(createPimVerificationPass()); pm.addPass(createEmitPimCodePass()); @@ -315,6 +316,8 @@ void addPassesPim(OwningOpRef& module, pm.addPass(mlir::createLowerAffinePass()); pm.addPass(createPimHostConstantFoldingPass()); pm.addPass(createMessagePass("Pim host constants folded")); + pm.addPass(createPimInstructionSelectionPass()); + pm.addPass(createMessagePass("Pim instructions selected")); pm.addPass(createPimLocalMemoryPlanningPass()); pm.addPass(createMessagePass("Pim local memory planned")); pm.addPass(createPimVerificationPass()); diff --git a/src/PIM/Compiler/PimCoreProgram.cpp b/src/PIM/Compiler/PimCoreProgram.cpp index 612533a..98ce549 100644 --- a/src/PIM/Compiler/PimCoreProgram.cpp +++ b/src/PIM/Compiler/PimCoreProgram.cpp @@ -1,12 +1,6 @@ #include "mlir/Dialect/MemRef/IR/MemRef.h" #include "mlir/Dialect/SCF/IR/SCF.h" -#include "mlir/IR/BuiltinTypes.h" - -#include "llvm/ADT/STLExtras.h" - #include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp" -#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" -#include "src/Accelerators/PIM/Common/Support/CheckedArithmetic.hpp" #include "src/Accelerators/PIM/Compiler/PimCoreProgram.hpp" #include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" @@ -20,11 +14,11 @@ static FailureOr classifyCompiledCoreOpKind(Operation& op) { if (isa(op)) return CompiledCoreOpKind::Load; if (isa(op)) return CompiledCoreOpKind::Store; if (isa(op)) return CompiledCoreOpKind::Lmv; + if (isa(op)) return CompiledCoreOpKind::VMV; if (isa(op)) return CompiledCoreOpKind::Receive; if (isa(op)) return CompiledCoreOpKind::Send; if (isa(op)) return CompiledCoreOpKind::Concat; if (isa(op)) return CompiledCoreOpKind::Vmm; - if (isa(op)) return CompiledCoreOpKind::Transpose; if (isa(op)) return CompiledCoreOpKind::VVAdd; if (isa(op)) return CompiledCoreOpKind::VVSub; if (isa(op)) return CompiledCoreOpKind::VVMul; @@ -38,77 +32,6 @@ static FailureOr classifyCompiledCoreOpKind(Operation& op) { return failure(); } -static bool isStoragePreservingTranspose(ArrayRef sourceShape, ArrayRef permutation) { - SmallVector sourceNonUnitDims; - SmallVector destinationSourceNonUnitDims; - for (auto [dim, size] : llvm::enumerate(sourceShape)) - if (size != 1) - sourceNonUnitDims.push_back(dim); - for (int64_t sourceDim : permutation) - if (sourceShape[sourceDim] != 1) - destinationSourceNonUnitDims.push_back(static_cast(sourceDim)); - return sourceNonUnitDims == destinationSourceNonUnitDims; -} - -static FailureOr compileTransposePlan(pim::PimTransposeOp transposeOp) { - auto sourceType = cast(transposeOp.getInput().getType()); - ArrayRef sourceShape = sourceType.getShape(); - size_t rank = sourceShape.size(); - CompiledTransposePlan plan; - plan.source = transposeOp.getInput(); - plan.destination = transposeOp.getOutputBuffer(); - plan.elementBytes = getElementTypeSizeInBytes(sourceType.getElementType()); - auto totalElements = pim::checkedSize(sourceType.getNumElements(), transposeOp, "transpose elements"); - if (failed(totalElements)) return failure(); - plan.totalElements = *totalElements; - auto totalBytes = pim::checkedMul(plan.totalElements, plan.elementBytes, transposeOp, "transpose byte size"); - if (failed(totalBytes)) return failure(); - plan.totalBytes = *totalBytes; - - SmallVector permutation = map_to_vector(transposeOp.getPermutation().getAsRange(), - [](IntegerAttr attr) { return attr.getInt(); }); - if (permutation.size() != rank) { - transposeOp.emitOpError("requires permutation rank to match source rank for PIM codegen"); - return failure(); - } - - SmallVector destinationShape(rank); - plan.destinationStrides.assign(rank, 1); - plan.destinationDimensionForSource.assign(rank, 0); - plan.destinationRewinds.assign(rank, 0); - SmallVector seenSourceDimensions(rank, false); - for (size_t dim = 0; dim < rank; ++dim) { - auto size = pim::checkedSize(sourceShape[dim], transposeOp, "transpose source dimension"); - if (failed(size)) return failure(); - plan.sourceShape.push_back(*size); - } - for (auto [destinationDim, sourceDim] : llvm::enumerate(permutation)) { - if (sourceDim < 0 || static_cast(sourceDim) >= rank || seenSourceDimensions[sourceDim]) { - transposeOp.emitOpError("requires a valid permutation containing each source dimension exactly once"); - return failure(); - } - seenSourceDimensions[sourceDim] = true; - destinationShape[destinationDim] = plan.sourceShape[sourceDim]; - plan.destinationDimensionForSource[sourceDim] = destinationDim; - } - for (size_t dim = rank; dim > 1; --dim) { - auto stride = pim::checkedMul( - plan.destinationStrides[dim - 1], destinationShape[dim - 1], transposeOp, "transpose destination stride"); - if (failed(stride)) return failure(); - plan.destinationStrides[dim - 2] = *stride; - } - for (size_t sourceDim = 0; sourceDim < rank; ++sourceDim) { - auto rewind = pim::checkedMul(plan.sourceShape[sourceDim], - plan.destinationStrides[plan.destinationDimensionForSource[sourceDim]], - transposeOp, - "transpose destination rewind"); - if (failed(rewind)) return failure(); - plan.destinationRewinds[sourceDim] = *rewind; - } - plan.storagePreserving = isStoragePreservingTranspose(plan.sourceShape, permutation); - return plan; -} - static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl& plan) { for (Operation& op : block) { if (isa(op) || isCoreStaticAddressOp(&op)) @@ -188,11 +111,6 @@ static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl(op)); - if (failed(transposePlan)) return failure(); - node.transposePlan = *transposePlan; - } plan.push_back(std::move(node)); } return success(); diff --git a/src/PIM/Compiler/PimCoreProgram.hpp b/src/PIM/Compiler/PimCoreProgram.hpp index 312714b..8afd4b4 100644 --- a/src/PIM/Compiler/PimCoreProgram.hpp +++ b/src/PIM/Compiler/PimCoreProgram.hpp @@ -6,34 +6,19 @@ #include "llvm/ADT/SmallVector.h" #include -#include - #include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp" namespace onnx_mlir { -struct CompiledTransposePlan { - mlir::Value source; - mlir::Value destination; - size_t elementBytes = 0; - size_t totalElements = 0; - size_t totalBytes = 0; - llvm::SmallVector sourceShape; - llvm::SmallVector destinationStrides; - llvm::SmallVector destinationDimensionForSource; - llvm::SmallVector destinationRewinds; - bool storagePreserving = false; -}; - enum class CompiledCoreOpKind : uint8_t { Load, Store, Lmv, + VMV, Receive, Send, Concat, Vmm, - Transpose, VVAdd, VVSub, VVMul, @@ -62,7 +47,6 @@ struct CompiledCoreNode { llvm::SmallVector caseValues; llvm::SmallVector>> caseBodies; std::unique_ptr> defaultBody; - std::optional transposePlan; }; struct CompiledCoreProgram { diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp index 99b9762..413911c 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.cpp @@ -180,7 +180,11 @@ FailureOr createRowStripAssemblyBlueprint(const RowStripPhysicalValue& va kRowStripIndexMap, rewriter, loc); } -FailureOr applyRowStripRelu(const RowStripPhysicalValue& value, PatternRewriter& rewriter, Location loc) { +template +static FailureOr applyRowStripActivation(const RowStripPhysicalValue& value, + PatternRewriter& rewriter, + Location loc, + BuildActivation buildActivation) { auto storageType = cast(value.storage.getType()); const int64_t laneCount = storageType.getDimSize(0); auto batchOp = createSpatComputeBatch(rewriter, @@ -193,10 +197,9 @@ FailureOr applyRowStripRelu(const RowStripPhysicalValue& value, PatternRe FailureOr fragment = extractGraphBatchPhysicalFragment( rewriter, loc, args.inputs.front(), args.lane, value.fragmentType); if (failed(fragment)) return failure(); - Value relu = spatial::SpatReluOp::create( - rewriter, loc, value.fragmentType, *fragment).getResult(); + Value result = buildActivation(*fragment); publishGraphBatchPhysicalFragment( - rewriter, loc, relu, args.outputs.front(), args.lane); + rewriter, loc, result, args.outputs.front(), args.lane); return success(); }); if (failed(batchOp)) @@ -204,6 +207,19 @@ FailureOr applyRowStripRelu(const RowStripPhysicalValue& value, PatternRe return batchOp->getResult(0); } +FailureOr applyRowStripRelu(const RowStripPhysicalValue& value, PatternRewriter& rewriter, Location loc) { + return applyRowStripActivation(value, rewriter, loc, [&](Value fragment) { + return spatial::SpatReluOp::create(rewriter, loc, value.fragmentType, fragment).getResult(); + }); +} + +FailureOr applyRowStripSilu(const RowStripPhysicalValue& value, PatternRewriter& rewriter, Location loc) { + return applyRowStripActivation(value, rewriter, loc, [&](Value fragment) { + Value sigmoid = spatial::SpatSigmoidOp::create(rewriter, loc, value.fragmentType, fragment).getResult(); + return spatial::SpatVMulOp::create(rewriter, loc, value.fragmentType, fragment, sigmoid).getResult(); + }); +} + FailureOr applyRowStripBiasAdd(const RowStripPhysicalValue& value, Value bias, PatternRewriter& rewriter, diff --git a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp index d66cbc8..a7f75e9 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp @@ -61,6 +61,10 @@ mlir::FailureOr applyRowStripRelu(const RowStripPhysicalValue& valu mlir::PatternRewriter& rewriter, mlir::Location loc); +mlir::FailureOr applyRowStripSilu(const RowStripPhysicalValue& value, + mlir::PatternRewriter& rewriter, + mlir::Location loc); + mlir::FailureOr applyRowStripBiasAdd(const RowStripPhysicalValue& value, mlir::Value bias, mlir::PatternRewriter& rewriter, diff --git a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp index 5c8795b..cdd6ff0 100644 --- a/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/LowerSpatialPlansPass.cpp @@ -58,6 +58,11 @@ lowerRowStripRelu(const RowStripPhysicalValue& input, spatial::SpatReluPlanOp pl return applyRowStripRelu(input, rewriter, planOp.getLoc()); } +static FailureOr +lowerRowStripSilu(const RowStripPhysicalValue& input, spatial::SpatSiluPlanOp planOp, PatternRewriter& rewriter) { + return applyRowStripSilu(input, rewriter, planOp.getLoc()); +} + static FailureOr lowerRowStripBiasAdd(const RowStripPhysicalValue& input, spatial::SpatBiasAddPlanOp planOp, PatternRewriter& rewriter) { @@ -349,6 +354,51 @@ struct LowerSpatialPlansPass final : PassWrapper(&op)) { + if (succeeded(getRowStripValue(rowStripValues, planOp.getInput()))) { + auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { + auto blueprint = dyn_cast(user); + return blueprint && blueprint.getPhysicalLayout() == kRowStripLayout; + }); + if (outputBlueprint == planOp.getResult().getUsers().end()) { + planOp.emitOpError("row-strip SiLU plan requires a row-strip blueprint result"); + signalPassFailure(); + return; + } + + FailureOr input = getRowStripValue(rowStripValues, planOp.getInput()); + rewriter.setInsertionPoint(planOp); + FailureOr lowered = lowerRowStripSilu(*input, planOp, rewriter); + if (failed(lowered)) { + planOp.emitOpError("failed to lower selected row-strip Spatial SiLU plan"); + signalPassFailure(); + return; + } + auto blueprint = cast(*outputBlueprint); + FailureOr output = buildRowStripValue(blueprint, *lowered); + if (failed(output)) { + signalPassFailure(); + return; + } + rowStripValues[blueprint.getResult()] = *output; + eraseAfterLowering.insert(planOp); + eraseAfterLowering.insert(blueprint); + continue; + } + + rewriter.setInsertionPoint(planOp); + auto computeOp = createSpatCompute<1>( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) { + Value sigmoid = spatial::SpatSigmoidOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), x).getResult(); + Value silu = spatial::SpatVMulOp::create( + rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, sigmoid).getResult(); + spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), silu); + }); + rewriter.replaceOp(planOp, computeOp.getResults()); + continue; + } if (auto planOp = dyn_cast(&op)) { auto outputBlueprint = llvm::find_if(planOp.getResult().getUsers(), [](Operation* user) { auto blueprint = dyn_cast(user); @@ -650,6 +700,7 @@ struct LowerSpatialPlansPass final : PassWrapper(op) diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp index 7bc8445..1e25d72 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialPass.cpp @@ -8,6 +8,7 @@ #include "mlir/Pass/Pass.h" #include "mlir/Pass/PassManager.h" #include "mlir/Transforms/Passes.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" #include "llvm/ADT/SmallVector.h" @@ -50,13 +51,14 @@ static void populateEmptyFunction(func::FuncOp funcOp) { SmallVector addPlans(funcOp.getOps()); SmallVector concatPlans(funcOp.getOps()); SmallVector reluPlans(funcOp.getOps()); + SmallVector siluPlans(funcOp.getOps()); SmallVector maxPoolPlans(funcOp.getOps()); SmallVector globalAveragePoolPlans( funcOp.getOps()); SmallVector blueprints(funcOp.getOps()); SmallVector materializers(funcOp.getOps()); if (!computes.empty() || !computeBatches.empty() || !convPlans.empty() || !biasAddPlans.empty() || !addPlans.empty() - || !concatPlans.empty() || !reluPlans.empty() || !maxPoolPlans.empty() || !blueprints.empty() + || !concatPlans.empty() || !reluPlans.empty() || !siluPlans.empty() || !maxPoolPlans.empty() || !blueprints.empty() || !globalAveragePoolPlans.empty() || !materializers.empty()) { return; } @@ -121,6 +123,14 @@ void ONNXToSpatialPass::runOnOperation() { return; } + RewritePatternSet fusionPatterns(ctx); + populateElementwiseFusionPatterns(fusionPatterns, ctx); + if (failed(applyPatternsGreedily(moduleOp, std::move(fusionPatterns)))) { + moduleOp.emitError("failed to fuse layout-aware ONNX elementwise patterns"); + signalPassFailure(); + return; + } + auto entryFunc = getPimEntryFunc(moduleOp); if (failed(entryFunc)) { moduleOp.emitError("failed to locate the PIM entry function during ONNX-to-Spatial lowering"); diff --git a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp index 2f60fd8..9e92dcf 100644 --- a/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/ONNXToSpatialVerifier.cpp @@ -149,6 +149,7 @@ void verifyLogicalTopLevelOps(func::FuncOp funcOp, pim::CappedDiagnosticReporter spatial::SpatAddPlanOp, spatial::SpatConcatPlanOp, spatial::SpatReluPlanOp, + spatial::SpatSiluPlanOp, spatial::SpatMaxPool2DPlanOp, spatial::SpatGlobalAveragePoolPlanOp, spatial::SpatBlueprintOp, diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp b/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp index 0a0e6aa..e3cec73 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns.hpp @@ -17,6 +17,7 @@ void populateWeightPromotionPatterns(mlir::RewritePatternSet& patterns, mlir::ML void populateConvPatterns(mlir::RewritePatternSet& patterns, mlir::MLIRContext* ctx); 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); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp index 847288d..2450228 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Conv.cpp @@ -1256,7 +1256,7 @@ static Value buildPackedWeights(DenseElementsAttr wDenseAttr, for (int64_t kernelIndex = 0; kernelIndex < tiling.kernelElements; ++kernelIndex) { const int64_t kernelH = kernelIndex / wType.getDimSize(3); const int64_t kernelW = kernelIndex % wType.getDimSize(3); - const int64_t targetRow = localChannel * tiling.kernelElements + kernelIndex; + const int64_t targetRow = kernelIndex * tiling.channelsPerTile + localChannel; for (int64_t multiplierIndex = 0; multiplierIndex < tiling.outputMultiplier; ++multiplierIndex) { const int64_t globalOutChannel = globalChannel * tiling.outputMultiplier + multiplierIndex; const int64_t sourceFlatIndex = @@ -1326,16 +1326,49 @@ static Value createInputTile(Value input, Value channelOffset = tiling.channelsPerTile == 1 ? channelTileIndex : affineMulConst(rewriter, loc, channelTileIndex, tiling.channelsPerTile, anchorOp); - Value tile4D = createConvInputPatch(input, - inputTileType, - batchIndex, - channelOffset, - inputHeightOffset, - inputWidthOffset, - dilationHeight, - dilationWidth, - rewriter, - loc); + Value tile4D; + if (dilationHeight == 1 && dilationWidth == 1) { + SmallVector offsets {batchIndex, inputHeightOffset, inputWidthOffset, channelOffset}; + SmallVector sizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(inputTileType.getDimSize(1)), + rewriter.getIndexAttr(inputTileType.getDimSize(2)), + rewriter.getIndexAttr(tiling.channelsPerTile)}; + tile4D = tensor::ExtractSliceOp::create( + rewriter, loc, inputTileType, input, offsets, sizes, getUnitStrides(rewriter, 4)); + } + else { + auto pixelType = RankedTensorType::get( + {1, 1, 1, tiling.channelsPerTile}, inputTileType.getElementType(), inputTileType.getEncoding()); + tile4D = tensor::EmptyOp::create(rewriter, loc, inputTileType.getShape(), inputTileType.getElementType()); + for (int64_t kernelH = 0; kernelH < inputTileType.getDimSize(1); ++kernelH) + for (int64_t kernelW = 0; kernelW < inputTileType.getDimSize(2); ++kernelW) { + Value sourceHeight = affineAddConst(rewriter, loc, inputHeightOffset, kernelH * dilationHeight, anchorOp); + Value sourceWidth = affineAddConst(rewriter, loc, inputWidthOffset, kernelW * dilationWidth, anchorOp); + SmallVector sizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(tiling.channelsPerTile)}; + Value pixel = tensor::ExtractSliceOp::create( + rewriter, + loc, + pixelType, + input, + SmallVector {batchIndex, sourceHeight, sourceWidth, channelOffset}, + sizes, + getUnitStrides(rewriter, 4)); + tile4D = tensor::InsertSliceOp::create( + rewriter, + loc, + pixel, + tile4D, + SmallVector {rewriter.getIndexAttr(0), + rewriter.getIndexAttr(kernelH), + rewriter.getIndexAttr(kernelW), + rewriter.getIndexAttr(0)}, + sizes, + getUnitStrides(rewriter, 4)); + } + } auto collapsedType = RankedTensorType::get({1, tiling.tileInputRows}, inputTileType.getElementType()); return tensor::CollapseShapeOp::create(rewriter, loc, @@ -1531,10 +1564,18 @@ rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& state.padWidthEnd, rewriter, loc); + auto paddedInputType = cast(paddedInput.getType()); + auto channelLastInputType = RankedTensorType::get({paddedInputType.getDimSize(0), + paddedInputType.getDimSize(2), + paddedInputType.getDimSize(3), + paddedInputType.getDimSize(1)}, + paddedInputType.getElementType()); + Value channelLastInput = ONNXTransposeOp::create( + rewriter, loc, channelLastInputType, paddedInput, rewriter.getI64ArrayAttr({0, 2, 3, 1})); Value packedWeights = buildPackedWeights(wDenseAttr, state.wType, *tiling, rewriter, loc); Value expandedBias; - SmallVector batchInputs {paddedInput}; + SmallVector batchInputs {channelLastInput}; if (state.hasBias) { expandedBias = expandBiasIfNeeded(state.b, rewriter, loc); auto biasType = dyn_cast(expandedBias.getType()); @@ -1553,9 +1594,8 @@ rewriteConv(Operation* convOp, const ConvLoweringState& state, PatternRewriter& {1, static_cast(crossbarSize.getValue())}, state.outType.getElementType()); auto piecesType = spatial::getGraphBatchPhysicalResultType( tiling->totalPatches * tiling->numChannelTiles, rowTileType); - auto paddedInputType = cast(paddedInput.getType()); auto inputTileType = - RankedTensorType::get({1, tiling->channelsPerTile, state.wType.getDimSize(2), state.wType.getDimSize(3)}, + RankedTensorType::get({1, state.wType.getDimSize(2), state.wType.getDimSize(3), tiling->channelsPerTile}, paddedInputType.getElementType()); SmallVector batchWeights; if (tiling->numChannelTiles == 1) { @@ -1766,7 +1806,7 @@ static Value unpackRowsFromParallelGemm(Value packedRows, } static Value createWeightMatrix( - Value weights, const ConvGemmPlan& plan, PatternRewriter& rewriter, Location loc) { + Value weights, const ConvGemmPlan& plan, bool transpose, PatternRewriter& rewriter, Location loc) { auto buildWeightMatrix = [&](Value weight) -> Value { Value flattened = tensor::CollapseShapeOp::create(rewriter, loc, @@ -1776,6 +1816,8 @@ static Value createWeightMatrix( {0}, {1, 2, 3} }); + if (!transpose) + return flattened; return ONNXTransposeOp::create(rewriter, loc, plan.wTransType, flattened, rewriter.getI64ArrayAttr({1, 0})) .getResult(); }; @@ -1783,8 +1825,9 @@ static Value createWeightMatrix( if (isCompileTimeComputable(weights)) return buildWeightMatrix(weights); + RankedTensorType resultType = transpose ? plan.wTransType : plan.wFlatType; auto computeOp = - createSpatCompute<1>(rewriter, loc, TypeRange {plan.wTransType}, {}, ValueRange {weights}, [&](Value weight) { + createSpatCompute<1>(rewriter, loc, TypeRange {resultType}, {}, ValueRange {weights}, [&](Value weight) { spatial::SpatYieldOp::create(rewriter, loc, buildWeightMatrix(weight)); }); return computeOp.getResult(0); @@ -1852,21 +1895,26 @@ static Value createPaddedPixelMajorWeightConstant(DenseElementsAttr sourceAttr, const ConvLoweringState& state, int64_t paddedK, int64_t paddedC, + int64_t packFactor, PatternRewriter& rewriter) { auto paddedType = RankedTensorType::get({paddedK, paddedC}, state.wType.getElementType()); SmallVector sourceValues(sourceAttr.getValues()); SmallVector paddedValues( paddedType.getNumElements(), cast(rewriter.getZeroAttr(paddedType.getElementType()))); - for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel) - for (int64_t kernelH = 0; kernelH < state.wHeight; ++kernelH) - for (int64_t kernelW = 0; kernelW < state.wWidth; ++kernelW) - for (int64_t inChannel = 0; inChannel < state.numChannelsIn; ++inChannel) { - const int64_t sourceFlatIndex = - (((outChannel * state.numChannelsIn) + inChannel) * state.wHeight + kernelH) * state.wWidth + kernelW; - const int64_t patchIndex = - ((kernelH * state.wWidth) + kernelW) * state.numChannelsIn + inChannel; - paddedValues[patchIndex * paddedC + outChannel] = sourceValues[sourceFlatIndex]; - } + const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + for (int64_t copy = 0; copy < packFactor; ++copy) + for (int64_t outChannel = 0; outChannel < state.numChannelsOut; ++outChannel) + for (int64_t kernelH = 0; kernelH < state.wHeight; ++kernelH) + for (int64_t kernelW = 0; kernelW < state.wWidth; ++kernelW) + for (int64_t inChannel = 0; inChannel < state.numChannelsIn; ++inChannel) { + const int64_t sourceFlatIndex = + (((outChannel * state.numChannelsIn) + inChannel) * state.wHeight + kernelH) * state.wWidth + kernelW; + const int64_t patchIndex = + ((kernelH * state.wWidth) + kernelW) * state.numChannelsIn + inChannel; + const int64_t packedRow = copy * patchSize + patchIndex; + const int64_t packedColumn = copy * state.numChannelsOut + outChannel; + paddedValues[packedRow * paddedC + packedColumn] = sourceValues[sourceFlatIndex]; + } return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), DenseElementsAttr::get(paddedType, paddedValues), paddedType); } @@ -2381,7 +2429,7 @@ static Value createStreamedConvRows(const ConvLoweringState& state, Value gemmBias = state.hasBias ? state.b : createZeroGemmBias(plan.gemmOutputRowsType, rewriter); Value packedBias = buildPackedBias(gemmBias, biasMatrix, biasDenseAttr, state, plan, rewriter, loc); Value gemmRows = ONNXGemmOp::create(rewriter, loc, plan.gemmOutputRowsType, inputRows, - packedWeights, packedBias, APFloat(1.0f), APFloat(1.0f), 0, 0).getY(); + packedWeights, packedBias, APFloat(1.0f), APFloat(1.0f), 0, !wDenseAttr).getY(); return maybeUnpackChunkRows(gemmRows, plan, rewriter, loc); } @@ -2401,9 +2449,10 @@ static Value rewritePackedIm2ColConv(const ConvLoweringState& state, ConvGemmPlan plan = buildConvGemmPlan(state, static_cast(wDenseAttr), !state.hasBias || static_cast(biasDenseAttr), 0, state.batchSize * state.outHeight * state.outWidth); - // Prepare weight matrix W for crossbar storage: - // W: [Cout, Cin, KH, KW] -> [Cout, patchSize] -> [patchSize, Cout] - Value weightMatrix = createWeightMatrix(state.w, plan, rewriter, loc); + // Static weights use the crossbar [patchSize, Cout] layout. Runtime weights + // stay in ONNX's contiguous [Cout, patchSize] layout and Gemm consumes them + // through transB without materializing a transpose. + Value weightMatrix = createWeightMatrix(state.w, plan, static_cast(wDenseAttr), rewriter, loc); Value gemmInputRows = createIm2colRows(state, preparedInput, plan, rewriter, loc); Value gemmB = buildPackedWeights(wDenseAttr, weightMatrix, state, plan, rewriter, loc); Value gemmBias = createZeroGemmBias(plan.gemmOutputRowsType, rewriter); @@ -2420,7 +2469,7 @@ static Value rewritePackedIm2ColConv(const ConvLoweringState& state, APFloat(1.0f), APFloat(1.0f), /*transA=*/0, - /*transB=*/0) + /*transB=*/!wDenseAttr) .getY(); return createCollectedConvOutput(ValueRange {gemmRows}, @@ -2452,7 +2501,7 @@ static Value rewriteStreamedConv(const ConvLoweringState& state, ConvGemmPlan seedPlan = buildConvGemmPlan( state, static_cast(wDenseAttr), !state.hasBias || static_cast(biasDenseAttr), 0, 1, forcedPackFactor); - Value weightMatrix = createWeightMatrix(state.w, seedPlan, rewriter, loc); + Value weightMatrix = createWeightMatrix(state.w, seedPlan, static_cast(wDenseAttr), rewriter, loc); Value collectedRows = createStreamedConvRows(state, preparedInput, weightMatrix, @@ -2516,6 +2565,20 @@ static bool rowStripOutputChannelTileFitsOneCore(const ConvGeometry& geometry) { <= static_cast(crossbarCountInCore.getValue()); } +static int64_t chooseRowStripPixelPackFactor(const ConvLoweringState& state, int64_t xbarDim) { + const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + const int64_t baseWeightGroups = ceilIntegerDivide(patchSize, xbarDim) + * ceilIntegerDivide(state.numChannelsOut, xbarDim); + int64_t factor = std::min(state.outWidth, xbarDim / state.numChannelsOut); + while (factor > 1 + && (state.outWidth % factor != 0 + || ceilIntegerDivide(factor * patchSize, xbarDim) + * ceilIntegerDivide(factor * state.numChannelsOut, xbarDim) + > baseWeightGroups)) + --factor; + return std::max(factor, 1); +} + static bool canConsumePixelMajorRowStripFragments(const ConvLoweringState& state, StringRef& failureReason) { if (state.batchSize != 1) { failureReason = "batch_not_one"; @@ -2558,6 +2621,7 @@ static Value createZeroTensorConstant(RankedTensorType type, PatternRewriter& re } static FailureOr createBiasRowConstant(const ConvLoweringState& state, + int64_t packFactor, PatternRewriter& rewriter) { DenseElementsAttr denseAttr; if (!isSupportedBiasAddValue(state.b, state.outType, &denseAttr)) @@ -2566,10 +2630,14 @@ static FailureOr createBiasRowConstant(const ConvLoweringState& state, if (failed(channelValues)) return failure(); - auto biasType = RankedTensorType::get({1, state.numChannelsOut}, state.outType.getElementType()); + SmallVector packedValues; + packedValues.reserve(packFactor * state.numChannelsOut); + for (int64_t copy = 0; copy < packFactor; ++copy) + packedValues.append(channelValues->begin(), channelValues->end()); + auto biasType = RankedTensorType::get({1, packFactor * state.numChannelsOut}, state.outType.getElementType()); return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), - DenseElementsAttr::get(biasType, *channelValues), + DenseElementsAttr::get(biasType, packedValues), biasType); } @@ -2825,39 +2893,47 @@ static FailureOr createConvInputWindow(Value input, Value initWindow = createZeroTensorConstant(paddedWindowType, rewriter); Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); - Value window = initWindow; - for (int64_t kernelRowIndex = 0; kernelRowIndex < state.wHeight; ++kernelRowIndex) { - Value kernelRow = getOrCreateIndexConstant(rewriter, anchorOp, kernelRowIndex); + Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); + Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); + Value cKernelRows = getOrCreateIndexConstant(rewriter, anchorOp, state.wHeight); + auto loop = buildNormalizedScfFor( + rewriter, + loc, + c0, + cKernelRows, + c1, + ValueRange {initWindow}, + [&](OpBuilder&, Location rowLoc, Value kernelRow, ValueRange iterArgs, + SmallVectorImpl& yielded) -> LogicalResult { FailureOr sourceRow = denseInput ? FailureOr( - extractDenseConvWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, loc)) - : extractProjectedRowStripWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, loc); + extractDenseConvWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, rowLoc)) + : extractProjectedRowStripWindowRow(input, sourceIndexTable, state, outputHeight, kernelRow, rewriter, rowLoc); if (failed(sourceRow)) return failure(); Value semanticRow = *sourceRow; if (state.padHeightBegin != 0 || state.padHeightEnd != 0) { - Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, loc); - semanticRow = spatial::SpatVMulOp::create(rewriter, loc, fragmentType, semanticRow, mask).getResult(); + Value mask = extractProjectedRowStripWindowMask(*maskTable, state, outputHeight, kernelRow, rewriter, rowLoc); + semanticRow = spatial::SpatVMulOp::create(rewriter, rowLoc, fragmentType, semanticRow, mask).getResult(); } - Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, loc); - window = tensor::InsertSliceOp::create(rewriter, - loc, - paddedRow, - window, - SmallVector {rewriter.getIndexAttr(0), - rewriter.getIndexAttr(kernelRowIndex), - rewriter.getIndexAttr(0), - rewriter.getIndexAttr(0)}, - SmallVector {rewriter.getIndexAttr(1), - rewriter.getIndexAttr(1), - rewriter.getIndexAttr( - state.xWidth + state.padWidthBegin - + state.padWidthEnd), - rewriter.getIndexAttr(state.numChannelsIn)}, - getUnitStrides(rewriter, 4)); - } - return window; + Value paddedRow = createHorizontallyPaddedRowStripFragment(semanticRow, state, rewriter, rowLoc); + yielded.push_back(tensor::InsertSliceOp::create( + rewriter, + rowLoc, + paddedRow, + iterArgs.front(), + SmallVector {rewriter.getIndexAttr(0), kernelRow, + rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)}, + SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.xWidth + state.padWidthBegin + + state.padWidthEnd), + rewriter.getIndexAttr(state.numChannelsIn)}, + getUnitStrides(rewriter, 4))); + return success(); + }); + return failed(loop) ? FailureOr(failure()) + : FailureOr(loop->results.front()); } static FailureOr createPixelMajorConvPatchRow(Value paddedWindow, @@ -2878,20 +2954,92 @@ static FailureOr createPixelMajorConvPatchRow(Value paddedWindow, rewriter.getIndexAttr(state.wHeight), rewriter.getIndexAttr(state.wWidth), rewriter.getIndexAttr(state.numChannelsIn)}; - SmallVector strides {rewriter.getIndexAttr(1), - rewriter.getIndexAttr(1), - rewriter.getIndexAttr(state.dilationWidth), - rewriter.getIndexAttr(1)}; - Value patch = tensor::ExtractSliceOp::create( - rewriter, loc, patchType, paddedWindow, offsets, sizes, strides); + Value patch; + if (state.dilationWidth == 1) + patch = tensor::ExtractSliceOp::create( + rewriter, loc, patchType, paddedWindow, offsets, sizes, getUnitStrides(rewriter, 4)); + else { + auto columnType = RankedTensorType::get({1, state.wHeight, 1, state.numChannelsIn}, + state.xType.getElementType(), state.xType.getEncoding()); + patch = tensor::EmptyOp::create(rewriter, loc, patchType.getShape(), patchType.getElementType()); + for (int64_t kernelColumn = 0; kernelColumn < state.wWidth; ++kernelColumn) { + Value sourceWidth = + affineAddConst(rewriter, loc, inputWidthOffset, kernelColumn * state.dilationWidth, anchorOp); + SmallVector columnSizes {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.wHeight), + rewriter.getIndexAttr(1), + rewriter.getIndexAttr(state.numChannelsIn)}; + Value column = tensor::ExtractSliceOp::create( + rewriter, loc, columnType, paddedWindow, + SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), sourceWidth, + rewriter.getIndexAttr(0)}, + columnSizes, getUnitStrides(rewriter, 4)); + patch = tensor::InsertSliceOp::create( + rewriter, loc, column, patch, + SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), + rewriter.getIndexAttr(kernelColumn), rewriter.getIndexAttr(0)}, + columnSizes, getUnitStrides(rewriter, 4)); + } + } return tensor::CollapseShapeOp::create( rewriter, loc, rowType, patch, SmallVector {{0}, {1, 2, 3}}) .getResult(); } +static FailureOr createPackedPixelMajorConvPatchRow(Value paddedWindow, + const ConvLoweringState& state, + Value outputGroup, + int64_t packFactor, + PatternRewriter& rewriter, + Location loc) { + if (packFactor == 1) + return createPixelMajorConvPatchRow(paddedWindow, state, outputGroup, rewriter, loc); + + Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); + const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + auto packedType = RankedTensorType::get( + {1, packFactor * patchSize}, state.xType.getElementType(), state.xType.getEncoding()); + Value packed = tensor::EmptyOp::create(rewriter, loc, packedType.getShape(), packedType.getElementType()); + Value outputStart = affineMulConst(rewriter, loc, outputGroup, packFactor, anchorOp); + Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); + Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); + Value cPackFactor = getOrCreateIndexConstant(rewriter, anchorOp, packFactor); + auto loop = buildNormalizedScfFor( + rewriter, + loc, + c0, + cPackFactor, + c1, + ValueRange {packed}, + [&](OpBuilder&, Location copyLoc, Value copy, ValueRange iterArgs, + SmallVectorImpl& yielded) -> LogicalResult { + Value outputWidth = createOrFoldAffineApply( + rewriter, copyLoc, rewriter.getAffineDimExpr(0) + rewriter.getAffineDimExpr(1), + ValueRange {outputStart, copy}, anchorOp); + FailureOr patch = + createPixelMajorConvPatchRow(paddedWindow, state, outputWidth, rewriter, copyLoc); + if (failed(patch)) + return failure(); + Value packedOffset = affineMulConst(rewriter, copyLoc, copy, patchSize, anchorOp); + yielded.push_back(tensor::InsertSliceOp::create( + rewriter, + copyLoc, + *patch, + iterArgs.front(), + SmallVector {rewriter.getIndexAttr(0), packedOffset}, + SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(patchSize)}, + getUnitStrides(rewriter, 2))); + return success(); + }); + if (failed(loop)) + return failure(); + return loop->results.front(); +} + static FailureOr> createConvInputTiles(Value paddedWindow, const ConvLoweringState& state, Value outputWidth, + int64_t packFactor, Value& partialInputScratch, int64_t patchSize, int64_t numKSlices, @@ -2903,7 +3051,7 @@ static FailureOr> createConvInputTiles(Value paddedWindow, SmallVector inputTiles; inputTiles.reserve(numKSlices); - if (state.numChannelsIn % xbarDim == 0) { + if (packFactor == 1 && state.numChannelsIn % xbarDim == 0) { Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); auto inputTileType = RankedTensorType::get( {1, 1, 1, xbarDim}, elementType, state.xType.getEncoding()); @@ -2938,8 +3086,8 @@ static FailureOr> createConvInputTiles(Value paddedWindow, return inputTiles; } - FailureOr patchRow = - createPixelMajorConvPatchRow(paddedWindow, state, outputWidth, rewriter, loc); + FailureOr patchRow = createPackedPixelMajorConvPatchRow( + paddedWindow, state, outputWidth, packFactor, rewriter, loc); if (failed(patchRow)) return failure(); @@ -3081,17 +3229,18 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, Value input, Value paddedWeights, Value bias, + int64_t packFactor, int64_t paddedK, int64_t numKSlices, int64_t xbarDim, PatternRewriter& rewriter, Location loc) { const int64_t laneCount = state.outHeight; - const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; + const int64_t patchSize = packFactor * state.numChannelsIn * state.wHeight * state.wWidth; const bool hasPartialInputTile = patchSize % xbarDim != 0; auto elementType = state.outType.getElementType(); auto partialInputScratchType = RankedTensorType::get({1, xbarDim}, elementType); - auto outputPixelType = RankedTensorType::get({1, 1, 1, state.numChannelsOut}, elementType); + auto outputPixelType = RankedTensorType::get({1, 1, packFactor, state.numChannelsOut}, elementType); auto fragmentType = getRowStripFragmentType(state.outType); auto storageType = getRowStripStorageType(state.outType); @@ -3106,7 +3255,7 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); Value c0 = getOrCreateIndexConstant(rewriter, anchorOp, 0); Value c1 = getOrCreateIndexConstant(rewriter, anchorOp, 1); - Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth); + Value cOutWidth = getOrCreateIndexConstant(rewriter, anchorOp, state.outWidth / packFactor); FailureOr inputWindow = createConvInputWindow(args.inputs.front(), state, args.lane, rewriter, loc); if (failed(inputWindow)) @@ -3131,6 +3280,7 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, FailureOr> inputTiles = createConvInputTiles(*inputWindow, state, localColumn, + packFactor, partialInputScratch, patchSize, numKSlices, @@ -3141,7 +3291,7 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, return failure(); FailureOr output = createConvOutputRow(*inputTiles, paddedK, - state.numChannelsOut, + packFactor * state.numChannelsOut, args.weights.front(), bias ? args.inputs[1] : Value(), xbarDim, @@ -3150,7 +3300,7 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, if (failed(output)) return failure(); Value outputPixel = tensor::ExpandShapeOp::create( - rewriter, pixelLoc, outputPixelType, *output, SmallVector {{0, 1, 2}, {3}}); + rewriter, pixelLoc, outputPixelType, *output, SmallVector {{0, 1}, {2, 3}}); Value next = tensor::InsertSliceOp::create( rewriter, pixelLoc, @@ -3158,11 +3308,11 @@ static FailureOr createRowStripConvOutput(const ConvLoweringState& state, iterArgs.front(), SmallVector {rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), - localColumn, + affineMulConst(rewriter, pixelLoc, localColumn, packFactor, anchorOp), rewriter.getIndexAttr(0)}, SmallVector {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), - rewriter.getIndexAttr(1), + rewriter.getIndexAttr(packFactor), rewriter.getIndexAttr(state.numChannelsOut)}, getUnitStrides(rewriter, 4)); yielded.push_back(next); @@ -3255,6 +3405,7 @@ static FailureOr createOutputChannelTiledRowStripConvOutput(const ConvLow FailureOr> inputTiles = createConvInputTiles(*inputWindow, state, widthIndex, + /*packFactor=*/1, partialInputScratch, patchSize, numKSlices, @@ -3313,29 +3464,34 @@ static FailureOr return failure(); const int64_t xbarDim = geometry.xbarSize; - const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; - const int64_t numKSlices = ceilIntegerDivide(patchSize, xbarDim); - const int64_t paddedK = numKSlices * xbarDim; + const int64_t basePatchSize = state.numChannelsIn * state.wHeight * state.wWidth; + const int64_t baseNumKSlices = ceilIntegerDivide(basePatchSize, xbarDim); + const int64_t basePaddedK = baseNumKSlices * xbarDim; if (!rowStripOutputTileFitsOneCore(geometry)) { Value tiledWeights = - standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter); + standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, basePaddedK, xbarDim, rewriter); return createOutputChannelTiledRowStripConvOutput( - state, state.x, tiledWeights, paddedK, numKSlices, xbarDim, rewriter, loc); + state, state.x, tiledWeights, basePaddedK, baseNumKSlices, xbarDim, rewriter, loc); } - const int64_t paddedOutputChannels = ceilIntegerDivide(state.numChannelsOut, xbarDim) * xbarDim; - Value paddedWeights = - standard::createPaddedPixelMajorWeightConstant(weightDenseAttr, state, paddedK, paddedOutputChannels, rewriter); + const int64_t packFactor = chooseRowStripPixelPackFactor(state, xbarDim); + const int64_t packedPatchSize = packFactor * basePatchSize; + const int64_t numKSlices = ceilIntegerDivide(packedPatchSize, xbarDim); + const int64_t paddedK = numKSlices * xbarDim; + const int64_t packedOutputChannels = packFactor * state.numChannelsOut; + const int64_t paddedOutputChannels = ceilIntegerDivide(packedOutputChannels, xbarDim) * xbarDim; + Value paddedWeights = standard::createPaddedPixelMajorWeightConstant( + weightDenseAttr, state, paddedK, paddedOutputChannels, packFactor, rewriter); FailureOr bias = failure(); if (state.hasBias) - bias = createBiasRowConstant(state, rewriter); + bias = createBiasRowConstant(state, packFactor, rewriter); if (state.hasBias && failed(bias)) return failure(); return createRowStripConvOutput( state, state.x, paddedWeights, state.hasBias ? *bias : Value(), - paddedK, numKSlices, xbarDim, rewriter, loc); + packFactor, paddedK, numKSlices, xbarDim, rewriter, loc); } static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value rowStripStorage, @@ -3351,32 +3507,36 @@ static FailureOr createConvOutputFromPixelMajorRowStripFragments(Value ro ConvGeometry geometry = buildConvGeometry(state); const int64_t xbarDim = geometry.xbarSize; - const int64_t patchSize = state.numChannelsIn * state.wHeight * state.wWidth; - const int64_t numKSlices = ceilIntegerDivide(patchSize, xbarDim); - const int64_t paddedK = numKSlices * xbarDim; + const int64_t basePatchSize = state.numChannelsIn * state.wHeight * state.wWidth; + const int64_t baseNumKSlices = ceilIntegerDivide(basePatchSize, xbarDim); + const int64_t basePaddedK = baseNumKSlices * xbarDim; auto weightDenseAttr = getHostConstDenseElementsAttr(state.w); if (!weightDenseAttr) return failure(); if (!rowStripOutputTileFitsOneCore(geometry)) { Value tiledWeights = - standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, paddedK, xbarDim, rewriter); + standard::createPaddedOutputChannelTiledWeightConstant(weightDenseAttr, state, basePaddedK, xbarDim, rewriter); return createOutputChannelTiledRowStripConvOutput( - state, rowStripStorage, tiledWeights, paddedK, numKSlices, xbarDim, rewriter, loc); + state, rowStripStorage, tiledWeights, basePaddedK, baseNumKSlices, xbarDim, rewriter, loc); } - const int64_t paddedOutputChannels = - ceilIntegerDivide(state.numChannelsOut, xbarDim) * xbarDim; + const int64_t packFactor = chooseRowStripPixelPackFactor(state, xbarDim); + const int64_t packedPatchSize = packFactor * basePatchSize; + const int64_t numKSlices = ceilIntegerDivide(packedPatchSize, xbarDim); + const int64_t paddedK = numKSlices * xbarDim; + const int64_t packedOutputChannels = packFactor * state.numChannelsOut; + const int64_t paddedOutputChannels = ceilIntegerDivide(packedOutputChannels, xbarDim) * xbarDim; Value paddedWeights = standard::createPaddedPixelMajorWeightConstant( - weightDenseAttr, state, paddedK, paddedOutputChannels, rewriter); + weightDenseAttr, state, paddedK, paddedOutputChannels, packFactor, rewriter); FailureOr bias = failure(); if (state.hasBias) - bias = createBiasRowConstant(state, rewriter); + bias = createBiasRowConstant(state, packFactor, rewriter); if (state.hasBias && failed(bias)) return failure(); return createRowStripConvOutput( state, rowStripStorage, paddedWeights, state.hasBias ? *bias : Value(), - paddedK, numKSlices, xbarDim, rewriter, loc); + packFactor, paddedK, numKSlices, xbarDim, rewriter, loc); } static FailureOr createPointwiseOutputFromRowStripFragments(Value rowStripStorage, diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp index 1172902..a71808d 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Elementwise.cpp @@ -16,6 +16,28 @@ using namespace mlir; namespace onnx_mlir { namespace { +struct SiluToSpatialPlan : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(ONNXMulOp mulOp, PatternRewriter& rewriter) const override { + ONNXSigmoidOp sigmoidOp = mulOp.getA().getDefiningOp(); + Value input = mulOp.getB(); + if (!sigmoidOp) { + sigmoidOp = mulOp.getB().getDefiningOp(); + input = mulOp.getA(); + } + if (!sigmoidOp || sigmoidOp.getX() != input || !sigmoidOp->hasOneUse() + || sigmoidOp.getResult().getType() != input.getType() || mulOp.getResult().getType() != input.getType()) + return failure(); + + auto plan = spatial::SpatSiluPlanOp::create( + rewriter, mulOp.getLoc(), mulOp.getResult().getType(), input, rewriter.getStringAttr("nchw")); + rewriter.replaceOp(mulOp, plan.getResult()); + rewriter.eraseOp(sigmoidOp); + return success(); + } +}; + static DenseElementsAttr getDenseConstantAttr(Value value) { if (auto constantOp = value.getDefiningOp()) return dyn_cast(constantOp.getValue()); @@ -219,6 +241,10 @@ struct AddToSpatialCompute : OpConversionPattern { } // namespace +void populateElementwiseFusionPatterns(RewritePatternSet& patterns, MLIRContext* ctx) { + patterns.add(ctx); +} + void populateElementwisePatterns(RewritePatternSet& patterns, MLIRContext* ctx) { patterns.add(ctx); patterns.add>(ctx); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp index c47f6a4..082ef04 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/Gemm.cpp @@ -259,17 +259,6 @@ static FailureOr createVmmBatch(Value a, return *batchOp; } -static Value -createDynamicGemmBatchRow(Value lane, int64_t numOutCols, ConversionPatternRewriter& rewriter, Location loc) { - if (numOutCols == 1) - return lane; - - MLIRContext* context = rewriter.getContext(); - AffineExpr d0 = getAffineDimExpr(0, context); - return createOrFoldAffineApply( - rewriter, loc, d0.floorDiv(numOutCols), ValueRange {lane}, rewriter.getInsertionBlock()->getParentOp()); -} - static Value extractDynamicGemmBColumn( Value matrix, Value column, RankedTensorType vectorType, ConversionPatternRewriter& rewriter, Location loc) { SmallVector offsets {rewriter.getIndexAttr(0), column}; @@ -373,33 +362,56 @@ static FailureOr createVvdmulBatch(Value a, Value b, RankedTensorType aType, RankedTensorType bType, - RankedTensorType scalarPiecesType, + RankedTensorType columnPiecesType, RankedTensorType outType, + bool transposeB, ConversionPatternRewriter& rewriter, Location loc) { const int64_t numOutRows = outType.getDimSize(0); const int64_t numOutCols = outType.getDimSize(1); const int64_t reductionSize = aType.getDimSize(1); - const int64_t laneCount = numOutRows * numOutCols; + auto vectorType = RankedTensorType::get({1, reductionSize}, aType.getElementType()); + auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType()); + auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType()); auto batchOp = createSpatComputeBatch( rewriter, loc, - TypeRange {scalarPiecesType}, - laneCount, + TypeRange {columnPiecesType}, + numOutCols, ValueRange {}, ValueRange {a, b}, [&](detail::SpatComputeBatchBodyArgs args) { - Value row = createDynamicGemmBatchRow(args.lane, numOutCols, rewriter, loc); - Value column = - onnx_mlir::affineModConst(rewriter, loc, args.lane, numOutCols, rewriter.getInsertionBlock()->getParentOp()); - - auto vectorType = RankedTensorType::get({1, reductionSize}, aType.getElementType()); - auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType()); - Value aVector = extractDynamicGemmRowVector(args.inputs[0], row, vectorType, rewriter, loc); - Value bVector = extractDynamicGemmBColumn(args.inputs[1], column, vectorType, rewriter, loc); - Value scalar = spatial::SpatVVDMulOp::create(rewriter, loc, scalarType, aVector, bVector).getResult(); - - publishGraphBatchPhysicalFragment(rewriter, loc, scalar, args.outputs.front(), args.lane); + Value bVector = transposeB + ? extractDynamicGemmRowVector(args.inputs[1], args.lane, vectorType, rewriter, loc) + : extractDynamicGemmBColumn(args.inputs[1], args.lane, vectorType, rewriter, loc); + Value columnInit = tensor::EmptyOp::create(rewriter, loc, columnType.getShape(), columnType.getElementType()); + 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( + rewriter, + loc, + c0, + cNumOutRows, + c1, + ValueRange {columnInit}, + [&](OpBuilder&, Location nestedLoc, Value row, ValueRange iterArgs, SmallVectorImpl& yielded) { + Value aVector = extractDynamicGemmRowVector(args.inputs[0], 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 {row, rewriter.getIndexAttr(0)}, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 2)); + yielded.push_back(next); + return success(); + }); + assert(succeeded(loop) && "dynamic Gemm row loop construction must succeed"); + publishGraphBatchPhysicalFragment(rewriter, loc, loop->results.front(), args.outputs.front(), args.lane); }); if (failed(batchOp)) return failure(); @@ -415,7 +427,7 @@ static FailureOr createDynamicGemmOutputCompute(Value scal float beta, ConversionPatternRewriter& rewriter, Location loc) { - const int64_t laneCount = scalarPiecesType.getDimSize(0); + const int64_t numOutRows = outType.getDimSize(0); const int64_t numOutCols = outType.getDimSize(1); SmallVector inputs {scalarPieces}; if (bias) @@ -428,43 +440,62 @@ static FailureOr createDynamicGemmOutputCompute(Value scal 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); + Value cNumOutCols = getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), numOutCols); + auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType()); auto loop = buildNormalizedScfFor( rewriter, loc, c0, - cLaneCount, + cNumOutCols, c1, ValueRange {outputInit}, - [&](OpBuilder&, Location nestedLoc, Value lane, ValueRange iterArgs, SmallVectorImpl& yielded) { + [&](OpBuilder&, Location nestedLoc, Value column, ValueRange iterArgs, SmallVectorImpl& yielded) { Value outputAcc = iterArgs.front(); - Value row = createDynamicGemmBatchRow(lane, numOutCols, rewriter, nestedLoc); - Value column = - onnx_mlir::affineModConst(rewriter, nestedLoc, lane, numOutCols, rewriter.getInsertionBlock()->getParentOp()); - SmallVector scalarSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; - SmallVector unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; - FailureOr scalar = extractGraphBatchPhysicalFragment(rewriter, nestedLoc, pieces, lane, scalarType); - if (failed(scalar)) + FailureOr columnPiece = + extractGraphBatchPhysicalFragment(rewriter, nestedLoc, pieces, column, columnType); + if (failed(columnPiece)) return failure(); - if (alpha != 1.0f) { - Value alphaTensor = createScalarTensorConstant(scalarType, alpha, rewriter, nestedLoc); - *scalar = spatial::SpatVMulOp::create(rewriter, nestedLoc, scalarType, *scalar, alphaTensor).getResult(); - } - if (biasArg) { - Value biasScalar = - createBroadcastedBiasScalar(biasArg, biasType, row, column, scalarType, rewriter, nestedLoc); - if (beta != 1.0f) { - Value betaTensor = createScalarTensorConstant(scalarType, beta, rewriter, nestedLoc); - biasScalar = - spatial::SpatVMulOp::create(rewriter, nestedLoc, scalarType, biasScalar, betaTensor).getResult(); - } - *scalar = spatial::SpatVAddOp::create(rewriter, nestedLoc, scalarType, *scalar, biasScalar).getResult(); - } - SmallVector outputOffsets {row, column}; - Value outputNext = - tensor::InsertSliceOp::create(rewriter, nestedLoc, *scalar, outputAcc, outputOffsets, scalarSizes, unitStrides) - .getResult(); - yielded.push_back(outputNext); + Value cNumOutRows = + getOrCreateIndexConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), numOutRows); + auto rowLoop = buildNormalizedScfFor( + rewriter, + nestedLoc, + c0, + cNumOutRows, + c1, + ValueRange {outputAcc}, + [&](OpBuilder&, Location rowLoc, Value row, ValueRange rowIterArgs, SmallVectorImpl& rowYielded) { + SmallVector scalarOffsets {row, rewriter.getIndexAttr(0)}; + SmallVector scalarSizes {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; + SmallVector unitStrides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; + Value scalar = tensor::ExtractSliceOp::create( + rewriter, rowLoc, scalarType, *columnPiece, scalarOffsets, scalarSizes, unitStrides); + if (alpha != 1.0f) { + Value alphaTensor = createScalarTensorConstant(scalarType, alpha, rewriter, rowLoc); + scalar = spatial::SpatVMulOp::create(rewriter, rowLoc, scalarType, scalar, alphaTensor).getResult(); + } + if (biasArg) { + Value biasScalar = createBroadcastedBiasScalar(biasArg, biasType, row, column, scalarType, rewriter, rowLoc); + if (beta != 1.0f) { + Value betaTensor = createScalarTensorConstant(scalarType, beta, rewriter, rowLoc); + biasScalar = + spatial::SpatVMulOp::create(rewriter, rowLoc, scalarType, biasScalar, betaTensor).getResult(); + } + scalar = spatial::SpatVAddOp::create(rewriter, rowLoc, scalarType, scalar, biasScalar).getResult(); + } + Value next = tensor::InsertSliceOp::create(rewriter, + rowLoc, + scalar, + rowIterArgs.front(), + SmallVector {row, column}, + scalarSizes, + unitStrides); + rowYielded.push_back(next); + return success(); + }); + if (failed(rowLoop)) + return failure(); + yielded.push_back(rowLoop->results.front()); return success(); }); if (failed(loop)) @@ -660,16 +691,10 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, aType = transposedType; } - if (gemmOpAdaptor.getTransB()) { - 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(); - bType = 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(); if (!isCompileTimeComputable(b)) { bool hasC = hasGemmBias(c); @@ -690,8 +715,9 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, biasType = *verifiedBiasType; } - if (aType.getDimSize(0) != numOutRows || bType.getDimSize(0) != reductionSize - || bType.getDimSize(1) != numOutCols) { + 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"); return failure(); } @@ -702,8 +728,9 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, return failure(); } - auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType(laneCount64, RankedTensorType::get({1, 1}, outType.getElementType())); - auto batchOp = createVvdmulBatch(a, b, aType, bType, scalarPiecesType, outType, rewriter, loc); + auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType()); + auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType(numOutCols, columnType); + auto batchOp = createVvdmulBatch(a, b, aType, bType, scalarPiecesType, outType, transposeB, rewriter, loc); if (failed(batchOp)) return failure(); auto outputCompute = createDynamicGemmOutputCompute( @@ -714,6 +741,13 @@ LogicalResult GemmToSpatialComputes::matchAndRewrite(ONNXGemmOp gemmOp, return success(); } + 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(); + bType = transposedType; + } + auto scaledB = materializeScaledConstantTensor(b, gemmOpAdaptor.getAlpha().convertToFloat(), rewriter, loc); if (failed(scaledB)) { gemmOp.emitOpError("requires constant Gemm input B when alpha is not 1.0"); diff --git a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp index 875f53f..fd26d78 100644 --- a/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/Patterns/Math/MatMul.cpp @@ -389,41 +389,6 @@ static FailureOr createBatchedVmmBatch(Value a, return *batchOp; } -static Value extractDynamicBatchedBColumn(Value matrix, - ArrayRef sourceBatchShape, - ArrayRef outputBatchShape, - Value outputBatchIndex, - Value column, - RankedTensorType vectorType, - PatternRewriter& rewriter, - Location loc) { - auto columnSliceType = RankedTensorType::get({1, vectorType.getDimSize(1), 1}, vectorType.getElementType()); - Value sourceBatchIndex = - mapOutputBatchIndexToSourceBatchIndex(outputBatchIndex, sourceBatchShape, outputBatchShape, rewriter, loc); - SmallVector offsets {OpFoldResult(sourceBatchIndex), rewriter.getIndexAttr(0), column}; - SmallVector sizes { - rewriter.getIndexAttr(1), rewriter.getIndexAttr(vectorType.getDimSize(1)), rewriter.getIndexAttr(1)}; - SmallVector strides {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; - Value columnSlice = tensor::ExtractSliceOp::create(rewriter, loc, columnSliceType, matrix, offsets, sizes, strides); - auto collapsedType = RankedTensorType::get({vectorType.getDimSize(1)}, vectorType.getElementType()); - Value collapsed = tensor::CollapseShapeOp::create(rewriter, - loc, - collapsedType, - columnSlice, - SmallVector { - {0, 1, 2} - }) - .getResult(); - return tensor::ExpandShapeOp::create(rewriter, - loc, - vectorType, - collapsed, - SmallVector { - {0, 1} - }) - .getResult(); -} - static Value extractDynamicBatchedRowVector(Value matrix, ArrayRef sourceBatchShape, ArrayRef outputBatchShape, @@ -449,7 +414,7 @@ static FailureOr createBatchedVvdmulBatch(Value a, ArrayRef outputBatchShape, RankedTensorType aType, RankedTensorType bType, - RankedTensorType scalarPiecesType, + RankedTensorType columnPiecesType, RankedTensorType outType, PatternRewriter& rewriter, Location loc) { @@ -457,29 +422,52 @@ static FailureOr createBatchedVvdmulBatch(Value a, 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 * numOutRows * numOutCols; + const int64_t laneCount = numBatches * numOutCols; + auto vectorType = RankedTensorType::get({1, reductionSize}, aType.getElementType()); + auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType()); + auto columnType = RankedTensorType::get({numOutRows, 1}, outType.getElementType()); auto batchOp = createSpatComputeBatch( rewriter, loc, - TypeRange {scalarPiecesType}, + TypeRange {columnPiecesType}, laneCount, ValueRange {}, ValueRange {a, b}, [&](detail::SpatComputeBatchBodyArgs args) { Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); - Value batch = affineFloorDivConst(rewriter, loc, args.lane, numOutRows * numOutCols, anchorOp); - Value batchLane = affineModConst(rewriter, loc, args.lane, numOutRows * numOutCols, anchorOp); - Value row = affineFloorDivConst(rewriter, loc, batchLane, numOutCols, anchorOp); - Value column = affineModConst(rewriter, loc, batchLane, numOutCols, anchorOp); - - auto vectorType = RankedTensorType::get({1, reductionSize}, aType.getElementType()); - auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType()); - Value aVector = extractDynamicBatchedRowVector( - args.inputs[0], aBatchShape, outputBatchShape, batch, row, vectorType, rewriter, loc); - Value bVector = extractDynamicBatchedBColumn( + 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 scalar = spatial::SpatVVDMulOp::create(rewriter, loc, scalarType, aVector, bVector).getResult(); - publishGraphBatchPhysicalFragment(rewriter, loc, scalar, args.outputs.front(), args.lane); + Value columnInit = tensor::EmptyOp::create(rewriter, loc, columnType.getShape(), columnType.getElementType()); + 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( + rewriter, + loc, + c0, + cNumOutRows, + c1, + ValueRange {columnInit}, + [&](OpBuilder&, Location nestedLoc, Value row, ValueRange iterArgs, SmallVectorImpl& yielded) { + 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 {row, rewriter.getIndexAttr(0)}, + SmallVector {rewriter.getIndexAttr(1), + rewriter.getIndexAttr(1)}, + getUnitStrides(rewriter, 2)); + yielded.push_back(next); + return success(); + }); + assert(succeeded(loop) && "dynamic MatMul row loop construction must succeed"); + publishGraphBatchPhysicalFragment(rewriter, loc, loop->results.front(), args.outputs.front(), args.lane); }); if (failed(batchOp)) return failure(); @@ -492,9 +480,8 @@ static FailureOr createBatchedDynamicOutputCompute(Value scalarPieces, PatternRewriter& rewriter, Location loc) { const int64_t laneCount = scalarPiecesType.getDimSize(0); - const int64_t numOutRows = outType.getDimSize(1); const int64_t numOutCols = outType.getDimSize(2); - auto scalarType = RankedTensorType::get({1, 1}, outType.getElementType()); + auto columnType = RankedTensorType::get({outType.getDimSize(1), 1}, outType.getElementType()); auto computeOp = createSpatCompute<1>( rewriter, loc, TypeRange {outType}, {}, ValueRange {scalarPieces}, [&](Value pieces) -> LogicalResult { @@ -513,19 +500,18 @@ static FailureOr createBatchedDynamicOutputCompute(Value scalarPieces, [&](OpBuilder&, Location nestedLoc, Value lane, ValueRange iterArgs, SmallVectorImpl& yielded) { Value outputAcc = iterArgs.front(); Operation* anchorOp = rewriter.getInsertionBlock()->getParentOp(); - Value batch = affineFloorDivConst(rewriter, nestedLoc, lane, numOutRows * numOutCols, anchorOp); - Value batchLane = affineModConst(rewriter, nestedLoc, lane, numOutRows * numOutCols, anchorOp); - Value row = affineFloorDivConst(rewriter, nestedLoc, batchLane, numOutCols, anchorOp); - Value column = affineModConst(rewriter, nestedLoc, batchLane, numOutCols, anchorOp); - FailureOr scalar = extractGraphBatchPhysicalFragment(rewriter, nestedLoc, pieces, lane, scalarType); - if (failed(scalar)) + Value batch = affineFloorDivConst(rewriter, nestedLoc, lane, numOutCols, anchorOp); + Value column = affineModConst(rewriter, nestedLoc, lane, numOutCols, anchorOp); + FailureOr columnPiece = + extractGraphBatchPhysicalFragment(rewriter, nestedLoc, pieces, lane, columnType); + if (failed(columnPiece)) return failure(); - SmallVector outputOffsets {batch, row, column}; + SmallVector outputOffsets {batch, rewriter.getIndexAttr(0), column}; SmallVector outputSizes = { - rewriter.getIndexAttr(1), rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; + rewriter.getIndexAttr(1), rewriter.getIndexAttr(outType.getDimSize(1)), rewriter.getIndexAttr(1)}; Value next = tensor::InsertSliceOp::create( - rewriter, nestedLoc, *scalar, outputAcc, outputOffsets, outputSizes, getUnitStrides(rewriter, 3)) + rewriter, nestedLoc, *columnPiece, outputAcc, outputOffsets, outputSizes, getUnitStrides(rewriter, 3)) .getResult(); yielded.push_back(next); return success(); @@ -1011,12 +997,14 @@ struct MatMulBatchedToSpatialComputes : OpRewritePattern { return success(); } } - const int64_t laneCount = plan.batch * plan.m * plan.n; + const int64_t laneCount = plan.batch * plan.n; + auto columnType = RankedTensorType::get({plan.m, 1}, shapeInfo->outType.getElementType()); auto scalarPiecesType = spatial::getGraphBatchPhysicalResultType( - laneCount, RankedTensorType::get({1, 1}, shapeInfo->outType.getElementType())); + laneCount, columnType); + Value transposedRhs = transposeLastTwoDims(plan.rhs, rewriter, loc); auto batchOp = createBatchedVvdmulBatch(plan.lhs, plan.lhsBatchShape, - plan.rhs, + transposedRhs, plan.rhsBatchShape, plan.outputBatchShape, plan.lhsType, diff --git a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp index f132fae..df69231 100644 --- a/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp +++ b/src/PIM/Conversion/ONNXToSpatial/SpatialLayoutPlanningPass.cpp @@ -34,6 +34,8 @@ static SelectedLayout getSelectedLayout(llvm::DenseMap& l static bool usesSelectedRowStrip(Operation* user, llvm::DenseMap& layouts) { if (auto reluPlan = dyn_cast(user)) return getSelectedLayout(layouts, reluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; + if (auto siluPlan = dyn_cast(user)) + return getSelectedLayout(layouts, siluPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; if (auto biasAddPlan = dyn_cast(user)) return getSelectedLayout(layouts, biasAddPlan.getResult()) == SelectedLayout::PixelMajorRowStrip; if (auto addPlan = dyn_cast(user)) @@ -62,7 +64,7 @@ static bool allUsersCanHandleRowStrip(Value value, llvm::DenseMap(user)) + if (isa(user)) return true; if (auto biasAddPlan = dyn_cast(user)) { auto resultType = dyn_cast(biasAddPlan.getOutput().getType()); @@ -105,13 +107,12 @@ static SelectedLayout chooseConvLayout(spatial::SpatConv2DPlanOp convPlan, return SelectedLayout::PixelMajorRowStrip; } -static SelectedLayout chooseReluLayout(spatial::SpatReluPlanOp reluPlan, - llvm::DenseMap& layouts) { - if (getSelectedLayout(layouts, reluPlan.getInput()) != SelectedLayout::PixelMajorRowStrip) +static SelectedLayout chooseActivationLayout(Value input, + Value result, + llvm::DenseMap& layouts) { + if (getSelectedLayout(layouts, input) != SelectedLayout::PixelMajorRowStrip) return SelectedLayout::DenseNchw; - if (!hasRowStripConsumer(reluPlan.getResult())) - return SelectedLayout::DenseNchw; - if (!allUsersCanHandleRowStrip(reluPlan.getResult(), layouts)) + if (!allUsersCanHandleRowStrip(result, layouts)) return SelectedLayout::DenseNchw; return SelectedLayout::PixelMajorRowStrip; } @@ -239,13 +240,21 @@ struct SpatialLayoutPlanningPass final : PassWrapper(&op)) { - SelectedLayout selected = chooseReluLayout(reluPlan, layouts); + SelectedLayout selected = chooseActivationLayout(reluPlan.getInput(), reluPlan.getResult(), layouts); if (layouts[reluPlan.getResult()] != selected) { layouts[reluPlan.getResult()] = selected; changed = true; } continue; } + if (auto siluPlan = dyn_cast(&op)) { + SelectedLayout selected = chooseActivationLayout(siluPlan.getInput(), siluPlan.getResult(), layouts); + if (layouts[siluPlan.getResult()] != selected) { + layouts[siluPlan.getResult()] = selected; + changed = true; + } + continue; + } if (auto biasAddPlan = dyn_cast(&op)) { SelectedLayout selected = chooseBiasAddLayout(biasAddPlan, layouts); if (layouts[biasAddPlan.getResult()] != selected) { @@ -301,6 +310,8 @@ struct SpatialLayoutPlanningPass final : PassWrapper(&op)) producedValue = reluPlan.getResult(); + else if (auto siluPlan = dyn_cast(&op)) + producedValue = siluPlan.getResult(); else if (auto maxPoolPlan = dyn_cast(&op)) producedValue = maxPoolPlan.getResult(); else if (auto averagePoolPlan = dyn_cast(&op)) diff --git a/src/PIM/Dialect/Pim/Analysis/LocalMemoryLifetimeAnalysis.cpp b/src/PIM/Dialect/Pim/Analysis/LocalMemoryLifetimeAnalysis.cpp index ea8013c..4bd7b94 100644 --- a/src/PIM/Dialect/Pim/Analysis/LocalMemoryLifetimeAnalysis.cpp +++ b/src/PIM/Dialect/Pim/Analysis/LocalMemoryLifetimeAnalysis.cpp @@ -47,11 +47,11 @@ static bool isRuntimeMemoryTouch(Operation* op) { return isa { }]; } +def PimVMVOp : PimOp<"vmv", [DestinationStyleOpInterface]> { + let summary = "Gather a strided vector within device memory"; + + let arguments = (ins + Index:$targetOffset, + Index:$sourceOffset, + Index:$sourceStride, + PimTensor:$target, + PimTensor:$source, + I32Attr:$length + ); + + let results = (outs + PimTensor:$output + ); + + let extraClassDeclaration = [{ + mlir::MutableOperandRange getDpsInitsMutable() { + return getTargetMutable(); + } + }]; + + let hasVerifier = 1; + let assemblyFormat = [{ + `[` $targetOffset `,` $sourceOffset `,` $sourceStride `]` + `(` $target `,` $source `)` attr-dict + `:` type($target) `,` type($source) `->` type($output) + }]; +} + def PimConcatOp : PimOp<"concat", [DestinationStyleOpInterface]> { let summary = "Concatenate tensors"; diff --git a/src/PIM/Dialect/Pim/PimOpsVerify.cpp b/src/PIM/Dialect/Pim/PimOpsVerify.cpp index 30ca606..a16942c 100644 --- a/src/PIM/Dialect/Pim/PimOpsVerify.cpp +++ b/src/PIM/Dialect/Pim/PimOpsVerify.cpp @@ -118,6 +118,20 @@ LogicalResult PimCoreBatchOp::verify() { return verifyOnlyConstantExternalValues(getOperation(), getBody(), "pim.core_batch"); } +LogicalResult PimVMVOp::verify() { + if (getLength() <= 0) + return emitError("length must be positive"); + if (failed(verifyCompatibleShapedTypes( + getOperation(), getTarget().getType(), getOutput().getType(), "target and output types must match"))) + return failure(); + auto targetType = cast(getTarget().getType()); + auto sourceType = dyn_cast(getSource().getType()); + if (!sourceType || !haveSameShapedContainerKind(getTarget().getType(), getSource().getType()) + || targetType.getElementType() != sourceType.getElementType()) + return emitError("target and source must use the same shaped container kind and element type"); + return success(); +} + LogicalResult PimVMMOp::verify() { if (failed(verifyCompatibleShapedTypes( getOperation(), getOutputBuffer().getType(), getOutput().getType(), "output buffer and output must match"))) diff --git a/src/PIM/Dialect/Pim/Transforms/InstructionSelection/CMakeLists.txt b/src/PIM/Dialect/Pim/Transforms/InstructionSelection/CMakeLists.txt new file mode 100644 index 0000000..540ae69 --- /dev/null +++ b/src/PIM/Dialect/Pim/Transforms/InstructionSelection/CMakeLists.txt @@ -0,0 +1,12 @@ +add_pim_library(OMPimInstructionSelection + InstructionSelectionPass.cpp + + EXCLUDE_FROM_OM_LIBS + + LINK_LIBS PUBLIC + MLIRArithDialect + MLIRSCFDialect + OMPimCommon + OMPimBufferization + PimOps +) diff --git a/src/PIM/Dialect/Pim/Transforms/InstructionSelection/InstructionSelectionPass.cpp b/src/PIM/Dialect/Pim/Transforms/InstructionSelection/InstructionSelectionPass.cpp new file mode 100644 index 0000000..cf5fbbb --- /dev/null +++ b/src/PIM/Dialect/Pim/Transforms/InstructionSelection/InstructionSelectionPass.cpp @@ -0,0 +1,167 @@ +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/SCF/IR/SCF.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +#include "llvm/ADT/STLExtras.h" + +#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" +#include "src/Accelerators/PIM/Common/Support/CheckedArithmetic.hpp" +#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" +#include "src/Accelerators/PIM/Dialect/Pim/Transforms/Bufferization/ContiguityPatterns.hpp" +#include "src/Accelerators/PIM/Pass/PIMPasses.h" + +using namespace llvm; +using namespace mlir; + +namespace onnx_mlir { +namespace { + +static bool isStoragePreservingTranspose(ArrayRef sourceShape, ArrayRef permutation) { + SmallVector sourceNonUnitDims; + SmallVector destinationSourceNonUnitDims; + for (auto [dim, size] : llvm::enumerate(sourceShape)) + if (size != 1) + sourceNonUnitDims.push_back(dim); + for (int64_t sourceDim : permutation) + if (sourceShape[sourceDim] != 1) + destinationSourceNonUnitDims.push_back(sourceDim); + return sourceNonUnitDims == destinationSourceNonUnitDims; +} + +static Value indexConstant(PatternRewriter& rewriter, Location loc, int64_t value) { + return arith::ConstantIndexOp::create(rewriter, loc, value); +} + +struct LowerTransposePattern final : OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(pim::PimTransposeOp op, PatternRewriter& rewriter) const override { + auto sourceType = dyn_cast(op.getInput().getType()); + auto targetType = dyn_cast(op.getOutputBuffer().getType()); + if (!sourceType || !targetType || !sourceType.hasStaticShape() || !targetType.hasStaticShape()) + return op.emitOpError("requires static memref operands before PIM instruction selection"); + + ArrayRef sourceShape = sourceType.getShape(); + size_t rank = sourceShape.size(); + SmallVector permutation = map_to_vector( + op.getPermutation().getAsRange(), [](IntegerAttr attr) { return attr.getInt(); }); + if (permutation.size() != rank) + return op.emitOpError("requires permutation rank to match source rank"); + SmallVector seen(rank, false); + for (int64_t sourceDim : permutation) { + if (sourceDim < 0 || static_cast(sourceDim) >= rank || seen[sourceDim]) + return op.emitOpError("requires a permutation containing each source dimension exactly once"); + seen[sourceDim] = true; + } + + auto totalElements = pim::checkedSize(sourceType.getNumElements(), op, "transpose elements"); + size_t elementBytes = getElementTypeSizeInBytes(sourceType.getElementType()); + auto totalBytes = failed(totalElements) + ? FailureOr(failure()) + : pim::checkedMul(*totalElements, elementBytes, op, "transpose byte size"); + if (failed(totalBytes)) + return failure(); + auto byteSizeAttr = pim::getCheckedI32Attr(rewriter, op, *totalBytes, "transpose byte size"); + if (failed(byteSizeAttr)) + return failure(); + + Location loc = op.getLoc(); + Value zero = indexConstant(rewriter, loc, 0); + if (isStoragePreservingTranspose(sourceShape, permutation)) { + auto copy = pim::PimMemCopyOp::create(rewriter, loc, targetType, zero, zero, + op.getOutputBuffer(), op.getInput(), *byteSizeAttr); + rewriter.replaceOp(op, copy.getOutput()); + return success(); + } + + SmallVector sourceStrides(rank, 1); + for (size_t dim = rank; dim > 1; --dim) + sourceStrides[dim - 2] = sourceStrides[dim - 1] * sourceShape[dim - 1]; + + SmallVector destinationShape; + destinationShape.reserve(rank); + for (int64_t sourceDim : permutation) + destinationShape.push_back(sourceShape[sourceDim]); + + size_t runStart = rank; + int64_t runLength = 1; + int64_t sourceStride = 1; + for (size_t destinationDim = rank; destinationDim-- > 0;) { + if (destinationShape[destinationDim] == 1) + continue; + int64_t candidateStride = sourceStrides[permutation[destinationDim]]; + if (runStart != rank && candidateStride != runLength * sourceStride) + break; + runStart = destinationDim; + if (runLength == 1) + sourceStride = candidateStride; + runLength *= destinationShape[destinationDim]; + } + assert(runStart != rank && "non-storage-preserving transpose must have a non-unit dimension"); + + int64_t outerCount = sourceType.getNumElements() / runLength; + auto lengthAttr = pim::getCheckedI32Attr(rewriter, op, runLength, "vmv length"); + auto runBytesAttr = pim::getCheckedI32Attr(rewriter, op, runLength * elementBytes, "transpose run byte size"); + if (failed(lengthAttr) || failed(runBytesAttr)) + return failure(); + + Value upper = indexConstant(rewriter, loc, outerCount); + Value one = indexConstant(rewriter, loc, 1); + auto loop = scf::ForOp::create(rewriter, loc, zero, upper, one, ValueRange{op.getOutputBuffer()}); + rewriter.setInsertionPointToStart(loop.getBody()); + Value outer = loop.getInductionVar(); + Value remaining = outer; + Value sourceElementOffset = zero; + for (size_t destinationDim = runStart; destinationDim-- > 0;) { + int64_t dimension = destinationShape[destinationDim]; + if (dimension == 1) + continue; + Value dimensionValue = indexConstant(rewriter, loc, dimension); + Value coordinate = arith::RemUIOp::create(rewriter, loc, remaining, dimensionValue); + remaining = arith::DivUIOp::create(rewriter, loc, remaining, dimensionValue); + Value stride = indexConstant(rewriter, loc, sourceStrides[permutation[destinationDim]]); + Value contribution = arith::MulIOp::create(rewriter, loc, coordinate, stride); + sourceElementOffset = arith::AddIOp::create(rewriter, loc, sourceElementOffset, contribution); + } + Value bytes = indexConstant(rewriter, loc, elementBytes); + Value sourceOffset = arith::MulIOp::create(rewriter, loc, sourceElementOffset, bytes); + Value runByteCount = indexConstant(rewriter, loc, runLength * elementBytes); + Value targetOffset = arith::MulIOp::create(rewriter, loc, outer, runByteCount); + Value target = loop.getRegionIterArg(0); + Value output; + if (sourceStride == 1) + output = pim::PimMemCopyOp::create( + rewriter, loc, targetType, targetOffset, sourceOffset, target, op.getInput(), *runBytesAttr).getOutput(); + else + output = pim::PimVMVOp::create(rewriter, loc, targetType, targetOffset, sourceOffset, + indexConstant(rewriter, loc, sourceStride), target, op.getInput(), *lengthAttr) + .getOutput(); + scf::YieldOp::create(rewriter, loc, output); + rewriter.replaceOp(op, loop.getResult(0)); + return success(); + } +}; + +struct InstructionSelectionPass : PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(InstructionSelectionPass) + + StringRef getArgument() const override { return "pim-instruction-selection"; } + StringRef getDescription() const override { return "Select explicit PIM ISA operations"; } + + void runOnOperation() override { + RewritePatternSet patterns(&getContext()); + patterns.add(&getContext()); + pim::populatePimContiguityNormalizationPatterns(patterns); + if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace + +std::unique_ptr createPimInstructionSelectionPass() { + return std::make_unique(); +} + +} // namespace onnx_mlir diff --git a/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp b/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp index d032c0c..d63ae7a 100644 --- a/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp @@ -238,11 +238,11 @@ static bool isSupportedCoreInstructionOp(Operation* op) { return isa { let hasVerifier = 1; } +def SpatSiluPlanOp : SpatOp<"silu_plan", []> { + let summary = "Layout-aware SiLU planning op"; + + let arguments = (ins + SpatTensor:$input, + StrAttr:$logicalLayout + ); + + let results = (outs + SpatTensor:$output + ); + + let hasVerifier = 1; +} + def SpatMaxPool2DPlanOp : SpatOp<"max_pool2d_plan", []> { let summary = "Layout-aware 2D NCHW MaxPool planning op"; diff --git a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp index 5b8b4a7..09f7eed 100644 --- a/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp +++ b/src/PIM/Dialect/Spatial/SpatialOpsVerify.cpp @@ -472,6 +472,14 @@ LogicalResult SpatReluPlanOp::verify() { return success(); } +LogicalResult SpatSiluPlanOp::verify() { + if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.silu_plan"))) + return failure(); + if (!isKnownLogicalLayout(getLogicalLayout())) + return emitError("requires a known logical layout"); + return success(); +} + LogicalResult SpatMaxPool2DPlanOp::verify() { if (failed(verifyPlanTensorTypes(getOperation(), getInput(), getOutput(), "spat.max_pool2d_plan"))) return failure(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp index 1089ee9..eca577b 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp @@ -265,6 +265,46 @@ static FailureOr emitReceiveValue(ArrayRef slices return receive.getOutput(); } +static FailureOr> +buildRectangularReceiveMetadata(const EmitReceiveAssemblyRun &run, + unsigned laneCount) { + if (run.lanes.size() != laneCount + || !llvm::all_of(run.entryLanes, + [&](const LaneSet &lanes) { return lanes.size() == laneCount; })) + return failure(); + SmallVector result; + for (size_t entry = 0; entry < run.positions.size(); ++entry) { + ArrayRef slices = ArrayRef(run.slices).slice( + run.entryOffsets[entry], + run.entryOffsets[entry + 1] - run.entryOffsets[entry]); + SmallVector ordered; + for (const ScheduledTransferSlice &slice : slices) + ordered.push_back(&slice); + llvm::sort(ordered, + [](const ScheduledTransferSlice *left, + const ScheduledTransferSlice *right) { + LaneInterval leftFamily = left->family->targetLanes.intervals().front(); + LaneInterval rightFamily = right->family->targetLanes.intervals().front(); + return leftFamily.begin + left->familyOffset + < rightFamily.begin + right->familyOffset; + }); + unsigned covered = 0; + LogicalTransferMetadataView metadata; + for (const ScheduledTransferSlice *slice : ordered) { + LaneInterval family = slice->family->targetLanes.intervals().front(); + unsigned begin = family.begin + slice->familyOffset; + if (begin != covered) + return failure(); + appendMetadata(*slice, metadata); + covered += slice->transferCount; + } + if (covered != laneCount) + return failure(); + result.push_back(std::move(metadata)); + } + return result; +} + template static FailureOr emitReceiveAssembly(const EmitReceiveAssemblyRun &run, Value lane, unsigned laneCount, Value initial, DeferredEmissionContext &context, Insert insert) { @@ -275,39 +315,11 @@ static FailureOr emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V SmallVector counts(laneCount); std::optional transferGrids; std::optional positions; - SmallVector metadataByEntry; - bool rectangular = run.lanes.size() == laneCount && llvm::all_of(run.entryLanes, [&](const LaneSet &lanes) { return lanes.size() == laneCount; }); - for (size_t entry = 0; rectangular && entry < run.positions.size(); ++entry) { - ArrayRef slices = - ArrayRef(run.slices).slice(run.entryOffsets[entry], run.entryOffsets[entry + 1] - run.entryOffsets[entry]); - SmallVector ordered; - for (const ScheduledTransferSlice &slice : slices) - ordered.push_back(&slice); - llvm::sort(ordered, [](const ScheduledTransferSlice *left, const ScheduledTransferSlice *right) { - LaneInterval leftFamily = left->family->targetLanes.intervals().front(); - LaneInterval rightFamily = right->family->targetLanes.intervals().front(); - return leftFamily.begin + left->familyOffset < rightFamily.begin + right->familyOffset; - }); - unsigned covered = 0; - LogicalTransferMetadataView metadata; - for (const ScheduledTransferSlice *slice : ordered) { - LaneInterval family = slice->family->targetLanes.intervals().front(); - unsigned begin = family.begin + slice->familyOffset; - if (begin != covered) { - rectangular = false; - break; - } - appendMetadata(*slice, metadata); - covered += slice->transferCount; - } - rectangular &= covered == laneCount; - if (rectangular) - metadataByEntry.push_back(std::move(metadata)); - } - if (rectangular) { + auto metadataByEntry = buildRectangularReceiveMetadata(run, laneCount); + if (succeeded(metadataByEntry)) { auto buildRows = [&](auto member) { SmallVector rows; - for (const LogicalTransferMetadataView &metadata : metadataByEntry) + for (const LogicalTransferMetadataView &metadata : *metadataByEntry) rows.push_back((metadata.*member).canonicalize()); return StaticIntGrid::fromRows(rows); }; @@ -400,6 +412,62 @@ static FailureOr emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V return loop->results.front(); } +static FailureOr transformProjectionFragment( + Value fragment, Value geometryRow, Value runtimeLane, + const DeferredProjectionLeafTemplate &leaf, const GridGeometry &geometry, + DeferredExchangePlan &exchange, DeferredEmissionContext &context) { + if (leaf.form != DeferredLeafForm::GraphBatchProjection) + return fragment; + Value shaped = extractMixedSliceOrIdentity( + context.rewriter, exchange.deferred.getLoc(), fragment, + getDeferredProjectedFragmentType(leaf), + lookupGeometry(geometry, geometryRow, runtimeLane, exchange.deferred, + context, exchange.deferred.getLoc())); + return shaped ? FailureOr(shaped) + : FailureOr(failure()); +} + +static bool canStreamLoopedReceiveResult( + const EmitReceiveAssemblyRun &run, unsigned laneCount, + const DeferredResultPlan &plan, DeferredEmissionContext &context) { + if (!run.collection + || run.collection->key.kind != FragmentCollectionKind::Leaf + || run.collection->key.leafIndex != 0 + || context.fragmentCollections.count(run.collection->key) + || run.lanes.size() != laneCount) + return false; + if (!canRealizeDeferredLoopResult(plan)) + return false; + if (run.collection->positionCount + != plan.exchange->program.leaves.front().loopIterationCount) + return false; + return llvm::all_of(run.positions, [&](unsigned position) { + return position < run.collection->positionCount; + }); +} + +static FailureOr emitLoopedReceiveResult( + const EmitReceiveAssemblyRun &run, Value lane, unsigned laneCount, + const DeferredResultPlan &plan, DeferredEmissionContext &context) { + return realizeDeferredLoopResult( + plan, lane, + [&](Value initial, + DeferredLoopIterationEmitter emitIteration) -> FailureOr { + return emitReceiveAssembly( + run, lane, laneCount, initial, context, + [&](Value fragment, Value position, Value, Value runtimeLane, + Value current) -> FailureOr { + auto shaped = transformProjectionFragment( + fragment, position, runtimeLane, plan.exchange->program.leaves.front(), + plan.innerGeometry.front(), *plan.exchange, context); + return failed(shaped) + ? FailureOr(failure()) + : emitIteration(*shaped, position, current); + }); + }, + context); +} + template static LogicalResult emitCollectionUpdate(const LaneSet &lanes, Value lane, unsigned laneCount, FragmentCollectionKey key, Value current, Operation *anchor, DeferredEmissionContext &context, Emit emit, bool local = false) { @@ -424,26 +492,59 @@ static LogicalResult emitCollectionUpdate(const LaneSet &lanes, Value lane, unsi return success(); } -static FailureOr insertProjectionFragment(Value fragment, Value specialization, Value position, Value geometryRow, Value runtimeLane, +static FailureOr insertProjectionFragment(Value fragment, + Value specialization, + Value position, + unsigned fragmentPositionCount, + Value runtimeLane, Value assembled, const DeferredProjectionLeafTemplate &leaf, const GridGeometry &geometry, DeferredExchangePlan &exchange, bool grouped, DeferredEmissionContext &context) { - Value shaped = fragment; - if (leaf.form == DeferredLeafForm::GraphBatchProjection) { - RankedTensorType projectedType = getDeferredProjectedFragmentType(leaf); - shaped = extractMixedSliceOrIdentity(context.rewriter, exchange.deferred.getLoc(), shaped, - projectedType, - lookupGeometry(geometry, geometryRow, runtimeLane, exchange.deferred, context, exchange.deferred.getLoc())); - if (!shaped) return failure(); + Value iteration = position; + Value fragmentPosition = position; + if (leaf.enclosingLoop && fragmentPositionCount > 1) { + Value divisor = context.constants.getIndex(fragmentPositionCount); + iteration = arith::DivUIOp::create( + context.rewriter, exchange.deferred.getLoc(), position, divisor); + fragmentPosition = arith::RemUIOp::create( + context.rewriter, exchange.deferred.getLoc(), position, divisor); } + Value geometryRow = leaf.enclosingLoop + ? iteration + : context.constants.getIndex(0); + if (grouped) { + Value base = arith::MulIOp::create( + context.rewriter, exchange.deferred.getLoc(), specialization, + context.constants.getIndex(leaf.loopIterationCount)); + geometryRow = arith::AddIOp::create( + context.rewriter, exchange.deferred.getLoc(), base, geometryRow); + } + auto transformed = transformProjectionFragment( + fragment, geometryRow, runtimeLane, leaf, geometry, exchange, context); + if (failed(transformed)) + return failure(); + Value shaped = *transformed; auto sourceType = dyn_cast(shaped.getType()); auto assembledType = dyn_cast(assembled.getType()); int64_t rankDifference = sourceType && assembledType ? assembledType.getRank() - sourceType.getRank() : 0; - if (rankDifference < 0 || rankDifference > 2 || (grouped && rankDifference == 0)) return failure(); - if (rankDifference == 0 && sourceType != assembledType) return failure(); + bool leading = sourceType && leaf.reconstructedType + && leaf.reconstructedType.getRank() == sourceType.getRank() + 1 + && leaf.reconstructedType.getShape().drop_front() + == sourceType.getShape(); + int64_t expectedRankDifference = static_cast(grouped) + + static_cast(!!leaf.enclosingLoop) + + static_cast(leading); + if (!sourceType || !assembledType + || rankDifference != expectedRankDifference) + return failure(); MixedSliceGeometry slice; slice.offsets.assign(assembledType.getRank(), context.rewriter.getIndexAttr(0)); - if (rankDifference) slice.offsets.front() = grouped ? specialization : position; - if (rankDifference == 2) slice.offsets[1] = position; + unsigned prefix = 0; + if (grouped) + slice.offsets[prefix++] = specialization; + if (leaf.enclosingLoop) + slice.offsets[prefix++] = iteration; + if (leading) + slice.offsets[prefix++] = fragmentPosition; slice.sizes.assign(rankDifference, context.rewriter.getIndexAttr(1)); for (int64_t dimension : sourceType.getShape()) slice.sizes.push_back(context.rewriter.getIndexAttr(dimension)); slice.strides.assign(assembledType.getRank(), context.rewriter.getIndexAttr(1)); @@ -488,7 +589,10 @@ static LogicalResult emitLeafCollectionUpdate(const EmitReceiveAssemblyRun &run, specialization = arith::DivUIOp::create(context.rewriter, exchange.deferred.getLoc(), position, divisor); leafPosition = arith::RemUIOp::create(context.rewriter, exchange.deferred.getLoc(), position, divisor); } - return insertProjectionFragment(fragment, specialization, leafPosition, grouped ? specialization : context.constants.getIndex(0), runtimeLane, + unsigned fragmentPositionCount = + collection.positionCount / leaf.loopIterationCount; + return insertProjectionFragment(fragment, specialization, leafPosition, + fragmentPositionCount, runtimeLane, assembled, leaf, geometry, exchange, grouped, context); }); }; @@ -642,10 +746,11 @@ static LogicalResult emitLoopedLocalCollectionUpdate( unsigned leafIndex = collection.key.leafIndex; const DeferredProjectionLeafTemplate &leaf = exchange.program.leaves[leafIndex]; + unsigned fragmentPositionCount = + collection.positionCount / leaf.loopIterationCount; auto inserted = insertProjectionFragment( *source, specialization, leafPosition, - grouped ? specialization : context.constants.getIndex(0), - runtimeLane, iterArgs.front(), leaf, + fragmentPositionCount, runtimeLane, iterArgs.front(), leaf, resultPlan.innerGeometry[leafIndex], exchange, grouped, context); if (failed(inserted)) return failure(); @@ -818,8 +923,10 @@ static LogicalResult emitLocalCollectionUpdate(const EmitLocalCollectionRun &upd auto emit = [&](Value assembled) -> FailureOr { auto fragment = materialize(); if (failed(fragment)) return failure(); + unsigned fragmentPositionCount = + collection.positionCount / leaf.loopIterationCount; return insertProjectionFragment(*fragment, context.constants.getIndex(specialization), - context.constants.getIndex(leafPosition), context.constants.getIndex(specialization), + context.constants.getIndex(leafPosition), fragmentPositionCount, lane ? lane : context.constants.getIndex(0), assembled, leaf, resultPlan.innerGeometry[leafIndex], exchange, grouped, context); }; return emitCollectionUpdate(update.lanes, lane, laneCount, collection.key, current, exchange.deferred, context, emit, true); @@ -828,7 +935,29 @@ static LogicalResult emitLocalCollectionUpdate(const EmitLocalCollectionRun &upd static FailureOr> emitInstructions(ArrayRef instructions, Value lane, unsigned laneCount, ArrayRef results, DeferredEmissionContext &context) { SmallVector produced; - for (const BoundaryInstruction &instruction : instructions) { + for (size_t instructionIndex = 0; + instructionIndex < instructions.size(); ++instructionIndex) { + const BoundaryInstruction &instruction = instructions[instructionIndex]; + if (auto assembly = std::get_if(&instruction)) { + if (instructionIndex + 1 < instructions.size()) + if (auto result = std::get_if( + &instructions[instructionIndex + 1])) { + const DeferredResultPlan *plan = findResultPlan(results, result->exchange); + if (plan && assembly->collection + && assembly->collection->key.exchange == result->exchange + && canStreamLoopedReceiveResult( + *assembly, laneCount, *plan, context)) { + auto value = emitLoopedReceiveResult( + *assembly, lane, laneCount, *plan, context); + if (failed(value)) + return result->exchange->deferred.emitOpError( + "failed to stream loop-indexed received result"), failure(); + produced.push_back(*value); + ++instructionIndex; + continue; + } + } + } if (auto send = std::get_if(&instruction)) { if (failed(emitConditionalSendRun(*send, lane, laneCount, context))) return failure(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp index a4a82b0..4c734a8 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp @@ -1,6 +1,7 @@ #pragma once #include "mlir/Dialect/Tensor/IR/Tensor.h" +#include "mlir/Dialect/SCF/IR/SCF.h" #include "mlir/IR/Operation.h" #include "llvm/ADT/ArrayRef.h" @@ -146,6 +147,8 @@ struct DeferredProjectionLeafTemplate { DeferredSliceTemplate innerGeometry; mlir::RankedTensorType reconstructedType; bool leadingRankReduced = false; + mlir::scf::ForOp enclosingLoop; + unsigned loopIterationCount = 1; }; struct DeferredInsertAssemblyEntryTemplate { diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp index b1c3f8f..e6846b2 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationPlanning.cpp @@ -171,6 +171,13 @@ static bool isTopLevelDeferredOperation(Operation *op, Block &body, && (isDeferredPayloadCandidateOp(op) || isa(op)); } +static Operation *getTopLevelDeferredOperation( + Operation *op, Block &body, const DeferredInputPlan &plan) { + while (op && op->getBlock() != &body) + op = op->getParentOp(); + return op && isTopLevelDeferredOperation(op, body, plan) ? op : nullptr; +} + static bool isEligible(Value value, Block &body, const DeferredInputPlan &plan, llvm::SmallPtrSetImpl &seen) { if (value == plan.graphInput || value == plan.graphLane || value == plan.scheduledLane) @@ -183,6 +190,27 @@ static bool isEligible(Value value, Block &body, const DeferredInputPlan &plan, return true; if (!op || !isTopLevelDeferredOperation(op, body, plan) || !seen.insert(op).second) return op && seen.contains(op); + if (auto loop = dyn_cast(op)) { + bool eligible = true; + loop.getRegion().walk([&](Operation *nested) { + if (isa(nested) && nested != loop) + eligible = false; + for (Value operand : nested->getOperands()) { + Operation *definition = operand.getDefiningOp(); + auto argument = dyn_cast(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)) + eligible = false; + } + }); + if (!eligible) + return false; + } return llvm::all_of(op->getOperands(), [&](Value operand) { return isEligible(operand, body, plan, seen); }); } @@ -198,8 +226,16 @@ static FailureOr clonePayloadRoot(Value root, Block &body, const Deferred argument && argument.getOwner() == &transfer.getBody().front() && argument.getArgNumber() >= transfer.getSources().size()) return value; - if (isa(value)) - return transfer.emitOpError("phase 1 payload shaping captures an unsupported block argument"), failure(); + if (auto argument = dyn_cast(value)) { + Operation *owner = argument.getOwner()->getParentOp(); + if (isa_and_nonnull(owner) + && owner->isProperAncestor(transfer)) + return value; + return transfer.emitOpError( + "phase 1 scheduled graph-lane expression captures an unsupported block argument: ") + << value << " owned by " << owner->getName(), + failure(); + } Operation *op = value.getDefiningOp(); if (!op || (!isDeferredPayloadCandidateOp(op) && !op->hasTrait())) return transfer.emitOpError("phase 1 cannot clone the scheduled graph-lane expression"), failure(); @@ -216,11 +252,39 @@ static FailureOr clonePayloadRoot(Value root, Block &body, const Deferred mapping.map(value, *mappedLane); return *mappedLane; } - if (isa(value)) - return transfer.emitOpError("phase 1 payload shaping captures an unsupported block argument"), failure(); + if (auto argument = dyn_cast(value)) { + Operation *owner = argument.getOwner()->getParentOp(); + if (isa_and_nonnull(owner) + && owner->isProperAncestor(transfer)) + return value; + return transfer.emitOpError( + "phase 1 payload shaping captures an unsupported block argument: ") + << value << " owned by " << owner->getName(), failure(); + } Operation *op = value.getDefiningOp(); if (!op || (!isTopLevelDeferredOperation(op, body, plan) && !op->hasTrait())) return transfer.emitOpError("phase 1 payload shaping contains an unsupported operation"), failure(); + if (auto loop = dyn_cast(op)) { + SmallVector captures; + loop.getRegion().walk([&](Operation *nested) { + for (Value operand : nested->getOperands()) { + Operation *definition = operand.getDefiningOp(); + auto argument = dyn_cast(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)) + captures.push_back(operand); + } + }); + for (Value capture : captures) + if (!mapping.contains(capture) && failed(clone(capture))) + return failure(); + } for (Value operand : op->getOperands()) if (failed(clone(operand))) return failure(); Operation *copy = builder.clone(*op, mapping); for (auto pair : llvm::zip(op->getResults(), copy->getResults())) mapping.map(std::get<0>(pair), std::get<1>(pair)); @@ -352,10 +416,19 @@ LogicalResult materializeDeferredPayloadDemands(OpBuilder &builder, Location loc Value value = worklist.pop_back_val(); if (!seen.insert(value).second) continue; for (OpOperand &use : value.getUses()) { - Operation *user = use.getOwner(); - if (!isTopLevelDeferredOperation(user, body, plan)) { needsIdentity = true; continue; } + Operation *user = getTopLevelDeferredOperation( + use.getOwner(), body, plan); + if (!user) { + if (value == plan.graphInput) needsIdentity = true; + else roots.push_back(value); + continue; + } llvm::SmallPtrSet eligibility; - if (!isEligible(user->getResult(0), body, plan, eligibility)) { needsIdentity = true; continue; } + if (!isEligible(user->getResult(0), body, plan, eligibility)) { + if (value == plan.graphInput) needsIdentity = true; + else roots.push_back(value); + continue; + } for (Value result : user->getResults()) { bool hasShapingUse = llvm::any_of(result.getUses(), [&](OpOperand &next) { return isTopLevelDeferredOperation(next.getOwner(), body, plan); }); bool hasOtherUse = llvm::any_of(result.getUses(), [&](OpOperand &next) { return !isTopLevelDeferredOperation(next.getOwner(), body, plan); }); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp index d29e072..42ae3a1 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp @@ -8,6 +8,8 @@ #include "llvm/ADT/SmallPtrSet.h" #include "llvm/ADT/SmallSet.h" +#include + #include "src/Accelerators/PIM/Common/IR/ShapingUtils.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp" @@ -187,6 +189,35 @@ static bool isInsideDeferredLoop(Operation *op, return false; } +static FailureOr getLoopIterationCount( + scf::ForOp loop, Value scheduledLane, Value specialization, + int64_t laneCount, int64_t specializationCount) { + std::optional commonCount; + for (int64_t specializationIndex = 0; + specializationIndex < specializationCount; ++specializationIndex) + for (int64_t lane = 0; lane < laneCount; ++lane) { + StaticIndexEnvironment environment; + if (scheduledLane) + environment.bindings[scheduledLane] = lane; + if (specialization) + environment.bindings[specialization] = specializationIndex; + auto lower = evaluateDeferredIndex(loop.getLowerBound(), environment); + auto upper = evaluateDeferredIndex(loop.getUpperBound(), environment); + auto step = evaluateDeferredIndex(loop.getStep(), environment); + if (failed(lower) || failed(upper) || failed(step) || *step <= 0) + return failure(); + int64_t distance = std::max(0, *upper - *lower); + uint64_t count = (static_cast(distance) + + static_cast(*step) - 1) + / static_cast(*step); + if (count > std::numeric_limits::max() + || (commonCount && *commonCount != count)) + return failure(); + commonCount = count; + } + return commonCount.value_or(0); +} + static FailureOr> getPossibleDeferredSourceOperandIndices( Value sourceRoot, SpatDeferredCommunicationOp deferred) { @@ -220,7 +251,7 @@ static LogicalResult validateDeferredProgram( op->emitOpError(message); return WalkResult::interrupt(); }; - if (isa(op) + if (isa(op) || (isa(op) && isa(op->getParentOp()))) return WalkResult::advance(); @@ -250,15 +281,39 @@ static LogicalResult validateDeferredProgram( return reject("has a non-positive or non-static deferred loop step"); return WalkResult::advance(); } - if (op->getNumRegions() != 0 && !isa(op)) + bool structuredLoopShaping = isInsideDeferredLoop(op, deferred) + && (isShapingOnlyOp(op) + || isa(op)); + if (op->getNumRegions() != 0 && !isa(op) + && !structuredLoopShaping) return reject("contains an unsupported region operation"); + bool structuredLoopResidual = isInsideDeferredLoop(op, deferred) + && (op->getNumRegions() == 0 + || structuredLoopShaping); if (!isShapingOnlyOp(op) && !isCompileTimeOp(op) - && !isPureIndexComputationOp(op)) + && !isPureIndexComputationOp(op) && !structuredLoopResidual) return reject("contains an unsupported deferred operation"); - for (Value operand : op->getOperands()) - if (isInsideDeferredLoop(op, deferred) - && originatesFromDeferredSource(operand, deferred)) - return reject("projects a deferred source inside a residual loop"); + if (isInsideDeferredLoop(op, deferred) + && llvm::any_of(op->getOperands(), [&](Value operand) { + return originatesFromDeferredSource(operand, deferred); + })) { + auto loop = op->getParentOfType(); + scf::ForOp outerLoop; + for (Operation *parent = loop ? loop->getParentOp() : nullptr; + parent && parent != deferred; parent = parent->getParentOp()) + if (auto candidate = dyn_cast(parent)) + outerLoop = candidate; + if (!loop || outerLoop) { + op->emitOpError("projects a deferred source inside an unsupported loop nest") + << " (inner " << (loop ? loop.getLoc() : op->getLoc()) + << ", outer " << (outerLoop ? outerLoop.getLoc() : op->getLoc()) + << ")"; + return WalkResult::interrupt(); + } + if (getConstantIntValue(loop.getLowerBound()) != 0 + || getConstantIntValue(loop.getStep()) != 1) + return reject("projects a deferred source inside a non-normalized loop"); + } return WalkResult::advance(); }); return failure(result.wasInterrupted()); @@ -376,9 +431,10 @@ FailureOr evaluateDeferredIndex( DeferredLaneValueEvaluator::DeferredLaneValueEvaluator( const DeferredProgramTemplate &program, unsigned laneCount, - unsigned specializationIndex) + unsigned specializationIndex, scf::ForOp loop, unsigned loopIteration) : program(program), laneCount(laneCount), - specializationIndex(specializationIndex) {} + specializationIndex(specializationIndex), loop(loop), + loopIteration(loopIteration) {} FailureOr DeferredLaneValueEvaluator::evaluate(Value value) { if (auto it = values.find(value); it != values.end()) @@ -397,6 +453,8 @@ FailureOr DeferredLaneValueEvaluator::evaluate(Value value) { if (program.specializationArgument) environment.bindings[program.specializationArgument] = specializationIndex; + if (loop) + environment.bindings[loop.getInductionVar()] = loopIteration; auto result = evaluateDeferredIndex(value, environment); if (failed(result)) return failure(); @@ -435,6 +493,8 @@ DeferredLaneValueEvaluator::resolveSourceOperandIndices(Value sourceRoot) { if (program.specializationArgument) environment.bindings[program.specializationArgument] = specializationIndex; + if (loop) + environment.bindings[loop.getInductionVar()] = loopIteration; auto index = sourceArgument(sourceRoot, program.deferred, environment); if (failed(index) || !*index) return failure(); @@ -518,6 +578,7 @@ FailureOr analyzeDeferredProgramTemplate( SmallVector( ArrayRef(slice.getMixedStrides()).drop_front())}; leaf.reconstructedType = cast(value.getType()); + leaf.enclosingLoop = slice->getParentOfType(); if (graphProjection && slice.getSourceType().getRank() == leaf.reconstructedType.getRank() + 1 @@ -540,30 +601,59 @@ FailureOr analyzeDeferredProgramTemplate( if (!type) return deferred.emitOpError( "deferred source is not a ranked tensor"); - program.leaves.push_back({DeferredLeafForm::DirectSource, value, value, - {}, {}, type}); + DeferredProjectionLeafTemplate leaf; + leaf.form = DeferredLeafForm::DirectSource; + leaf.sourceRoot = value; + leaf.replacementRoot = value; + leaf.reconstructedType = type; + program.leaves.push_back(std::move(leaf)); return success(); } if (value.getType().isIndex() || isa(value.getType())) return success(); + if (auto argument = dyn_cast(value)) { + auto loop = dyn_cast_or_null( + argument.getOwner()->getParentOp()); + if (loop && argument.getArgNumber() > 0) + return visit(loop.getInitArgs()[argument.getArgNumber() - 1]); + } Operation *op = value.getDefiningOp(); - if (!op || (op->getBlock() != &body && !isa(op))) + if (!op || (op->getBlock() != &body + && !op->getParentOfType())) return deferred.emitOpError( - "deferred residual escapes its verified body"); + "deferred residual escapes its verified body: ") << value; if (auto loop = dyn_cast(op)) { for (Value init : loop.getInitArgs()) if (failed(visit(init))) return failure(); + auto loopYield = cast(loop.getBody()->getTerminator()); + for (Value yielded : loopYield.getResults()) + if (failed(visit(yielded))) + return failure(); } else { for (Value operand : op->getOperands()) if (failed(visit(operand))) return failure(); } - program.residualOps.push_back(op); + if (op->getBlock() == &body) + program.residualOps.push_back(op); return success(); }; if (failed(visit(program.yieldedValue))) return failure(); + for (DeferredProjectionLeafTemplate &leaf : program.leaves) { + if (!leaf.enclosingLoop) + continue; + auto count = getLoopIterationCount( + leaf.enclosingLoop, program.scheduledLane, + program.specializationArgument, laneCount, + program.specializationCount); + if (failed(count) || *count == 0) + return deferred.emitOpError( + "deferred source projection loop must have a common non-zero trip count"), + failure(); + leaf.loopIterationCount = *count; + } auto assembly = analyzeInsertAssembly(program); if (failed(assembly)) return failure(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.hpp index aba4830..f28965e 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.hpp @@ -25,7 +25,9 @@ class DeferredLaneValueEvaluator { public: DeferredLaneValueEvaluator(const DeferredProgramTemplate &program, unsigned laneCount, - unsigned specializationIndex = 0); + unsigned specializationIndex = 0, + mlir::scf::ForOp loop = {}, + unsigned loopIteration = 0); mlir::FailureOr evaluate(mlir::Value value); mlir::FailureOr evaluate(mlir::OpFoldResult value); @@ -36,6 +38,8 @@ private: const DeferredProgramTemplate &program; unsigned laneCount; unsigned specializationIndex; + mlir::scf::ForOp loop; + unsigned loopIteration; llvm::DenseMap values; llvm::DenseMap sourceOperands; }; diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp index ac036d6..6cdb525 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp @@ -73,15 +73,19 @@ static LogicalResult buildLeafCollections(DeferredExchangePlan &exchange, RankedTensorType normalized = fragmentType; if (leaf.form == DeferredLeafForm::GraphBatchProjection) normalized = getDeferredProjectedFragmentType(leaf); - bool direct = positionCount == 1 && normalized == leaf.reconstructedType; + if (positionCount % leaf.loopIterationCount != 0) + return failure(); + unsigned fragmentPositionCount = positionCount / leaf.loopIterationCount; + bool direct = fragmentPositionCount == 1 + && normalized == leaf.reconstructedType; bool leading = leaf.reconstructedType.getRank() == normalized.getRank() + 1 - && leaf.reconstructedType.getDimSize(0) == positionCount + && leaf.reconstructedType.getDimSize(0) == fragmentPositionCount && leaf.reconstructedType.getShape().drop_front() == normalized.getShape(); if (!direct && !leading) return exchange.deferred.emitOpError( "cannot form fragment collection for leaf ") - << leafIndex << ": " << positionCount << " fragments of " + << leafIndex << ": " << fragmentPositionCount << " fragments of " << normalized << " do not reconstruct " << leaf.reconstructedType; FragmentCollectionKind kind = specializationCount == 1 ? FragmentCollectionKind::Leaf @@ -89,6 +93,8 @@ static LogicalResult buildLeafCollections(DeferredExchangePlan &exchange, SmallVector shape; if (specializationCount > 1) shape.push_back(specializationCount); + if (leaf.enclosingLoop) + shape.push_back(leaf.loopIterationCount); llvm::append_range(shape, leaf.reconstructedType.getShape()); FragmentCollectionPlan collection; collection.key = {&exchange, kind, static_cast(leafIndex)}; @@ -196,11 +202,14 @@ static constexpr std::array templateGeometryMembers{ template static FailureOr buildGrid( const DeferredProgramTemplate &program, unsigned laneCount, - unsigned rowCount, bool specializeRows, GetValue getValue) { + unsigned rowCount, bool specializeRows, scf::ForOp loop, + unsigned loopIterationCount, GetValue getValue) { SmallVector rows; for (unsigned row = 0; row < rowCount; ++row) { + unsigned specialization = specializeRows ? row / loopIterationCount : 0; + unsigned iteration = loop ? row % loopIterationCount : 0; DeferredLaneValueEvaluator evaluator( - program, laneCount, specializeRows ? row : 0); + program, laneCount, specialization, loop, iteration); auto sequence = evaluator.evaluate(getValue(row)); if (failed(sequence)) return failure(); rows.push_back(std::move(*sequence)); @@ -211,7 +220,8 @@ static FailureOr buildGrid( template static FailureOr buildGeometryGrids( const DeferredProgramTemplate &program, unsigned laneCount, - unsigned rowCount, bool specializeRows, GetGeometry getGeometry) { + unsigned rowCount, bool specializeRows, scf::ForOp loop, + unsigned loopIterationCount, GetGeometry getGeometry) { DeferredGridSliceGeometry result; const DeferredSliceTemplate &first = getGeometry(0); for (auto [group, sourceMember] : llvm::enumerate(templateGeometryMembers)) { @@ -219,6 +229,7 @@ static FailureOr buildGeometryGrids( for (unsigned dimension = 0; dimension < (first.*member).size(); ++dimension) { auto grid = buildGrid(program, laneCount, rowCount, specializeRows, + loop, loopIterationCount, [&](unsigned row) { return (getGeometry(row).*member)[dimension]; }); if (failed(grid)) return failure(); result[group].push_back(std::move(*grid)); @@ -226,8 +237,92 @@ static FailureOr buildGeometryGrids( } return result; } -static Value cloneResidual(DeferredExchangePlan &exchange, IRMapping &mapping, DeferredEmissionContext &context) { + +static LogicalResult cloneLoopIteration( + DeferredExchangePlan &exchange, scf::ForOp oldLoop, + const IRMapping &mapping, Value induction, ValueRange iterArgs, + llvm::function_ref(unsigned)> getLeaf, + DeferredEmissionContext &context, SmallVectorImpl &yielded) { + IRMapping nested(mapping); + nested.map(oldLoop.getInductionVar(), induction); + for (auto [oldArg, newArg] : + llvm::zip(oldLoop.getRegionIterArgs(), iterArgs)) + nested.map(oldArg, newArg); + for (Operation &nestedOp : oldLoop.getBody()->without_terminator()) { + bool replaced = false; + for (auto [leafIndex, leaf] : llvm::enumerate(exchange.program.leaves)) { + if (leaf.enclosingLoop != oldLoop + || leaf.replacementRoot.getDefiningOp() != &nestedOp) + continue; + auto selected = getLeaf(leafIndex); + if (failed(selected)) + return failure(); + nested.map(leaf.replacementRoot, *selected); + replaced = true; + } + if (replaced) + continue; + Operation *copy = context.rewriter.clone(nestedOp, nested); + for (auto [oldValue, newValue] : + llvm::zip(nestedOp.getResults(), copy->getResults())) + nested.map(oldValue, newValue); + } + auto oldYield = cast(oldLoop.getBody()->getTerminator()); + for (Value value : oldYield.getResults()) + yielded.push_back(nested.lookupOrDefault(value)); + return success(); +} + +static Value cloneResidual( + DeferredExchangePlan &exchange, IRMapping &mapping, + ArrayRef leafCollections, + DeferredEmissionContext &context) { + for (Operation &op : + exchange.deferred.getBody().front().without_terminator()) { + if (!op.hasTrait() + || llvm::all_of(op.getResults(), + [&](Value result) { return mapping.contains(result); })) + continue; + Operation *copy = context.rewriter.clone(op, mapping); + for (auto [oldValue, newValue] : + llvm::zip(op.getResults(), copy->getResults())) + mapping.map(oldValue, newValue); + } for (Operation *op : exchange.program.residualOps) { + if (op->hasTrait()) + continue; + if (auto oldLoop = dyn_cast(op)) { + SmallVector initArgs; + for (Value init : oldLoop.getInitArgs()) + initArgs.push_back(mapping.lookupOrDefault(init)); + auto loop = buildNormalizedScfFor( + context.rewriter, oldLoop.getLoc(), + mapping.lookupOrDefault(oldLoop.getLowerBound()), + mapping.lookupOrDefault(oldLoop.getUpperBound()), + mapping.lookupOrDefault(oldLoop.getStep()), initArgs, + [&](OpBuilder &, Location loc, Value induction, ValueRange iterArgs, + SmallVectorImpl &yielded) -> LogicalResult { + return cloneLoopIteration( + exchange, oldLoop, mapping, induction, iterArgs, + [&](unsigned leafIndex) -> FailureOr { + const auto &leaf = exchange.program.leaves[leafIndex]; + Value selected = extractMixedSliceOrIdentity( + context.rewriter, loc, leafCollections[leafIndex], + leaf.reconstructedType, + leadingSlice(context.rewriter, leaf.reconstructedType, + induction)); + return selected ? FailureOr(selected) + : FailureOr(failure()); + }, + context, yielded); + }); + if (failed(loop)) + return {}; + for (auto [oldValue, newValue] : + llvm::zip(oldLoop.getResults(), loop->results)) + mapping.map(oldValue, newValue); + continue; + } Operation *copy = context.rewriter.clone(*op, mapping); for (auto [oldValue, newValue] : llvm::zip(op->getResults(), copy->getResults())) mapping.map(oldValue, newValue); } @@ -252,16 +347,26 @@ static FailureOr realizeOne(const DeferredResultPlan &plan, Value lane, D if (!value.getType().isIndex()) selected = arith::IndexCastOp::create(context.rewriter, exchange.deferred.getLoc(), value.getType(), selected); mapping.map(value, selected); } + SmallVector leafCollections; for (unsigned index = 0; index < exchange.program.leaves.size(); ++index) { Value value = context.fragmentCollections.lookup( {&exchange, FragmentCollectionKind::Leaf, index}); - if (!value || value.getType() != exchange.program.leaves[index].reconstructedType) { + const auto &leaf = exchange.program.leaves[index]; + SmallVector expectedShape; + if (leaf.enclosingLoop) + expectedShape.push_back(leaf.loopIterationCount); + llvm::append_range(expectedShape, leaf.reconstructedType.getShape()); + Type expectedType = RankedTensorType::get( + expectedShape, leaf.reconstructedType.getElementType()); + if (!value || value.getType() != expectedType) { exchange.deferred.emitOpError("failed to reconstruct deferred result leaf ") << index; return failure(); } - mapping.map(exchange.program.leaves[index].replacementRoot, value); + leafCollections.push_back(value); + if (!exchange.program.leaves[index].enclosingLoop) + mapping.map(exchange.program.leaves[index].replacementRoot, value); } - Value result = cloneResidual(exchange, mapping, context); + Value result = cloneResidual(exchange, mapping, leafCollections, context); Type expected = exchange.program.specializationCount > 1 ? Type(exchange.program.specializationFragmentType) : exchange.deferred.getOutput().getType(); return result && result.getType() == expected ? FailureOr(result) : FailureOr(failure()); } @@ -274,8 +379,10 @@ FailureOr buildDeferredResultPlan(DeferredExchangePlan &exch : buildLeafCollections(exchange, plan))) return failure(); unsigned specializations = exchange.program.specializationCount; for (const auto &leaf : exchange.program.leaves) { + unsigned geometryRows = specializations * leaf.loopIterationCount; auto geometry = buildGeometryGrids(exchange.program, - exchange.targetLaneCount, specializations, true, + exchange.targetLaneCount, geometryRows, true, leaf.enclosingLoop, + leaf.loopIterationCount, [&](unsigned) -> const DeferredSliceTemplate & { return leaf.innerGeometry; }); @@ -285,7 +392,7 @@ FailureOr buildDeferredResultPlan(DeferredExchangePlan &exch if (exchange.program.insertAssembly) { const auto &entries = exchange.program.insertAssembly->entries; auto geometry = buildGeometryGrids(exchange.program, - exchange.targetLaneCount, entries.size(), false, + exchange.targetLaneCount, entries.size(), false, {}, 1, [&](unsigned row) -> const DeferredSliceTemplate & { return entries[row].targetGeometry; }); @@ -301,7 +408,7 @@ FailureOr buildDeferredResultPlan(DeferredExchangePlan &exch }); for (Value value : residualValues) { auto grid = buildGrid(exchange.program, exchange.targetLaneCount, - specializations, true, + specializations, true, {}, 1, [&](unsigned) { return OpFoldResult(value); }); if (succeeded(grid)) plan.residualValues.try_emplace(value, std::move(*grid)); @@ -320,6 +427,8 @@ FailureOr realizeDeferredResult(const DeferredResultPlan &plan, Value lan FragmentCollectionKey key{&exchange, FragmentCollectionKind::GroupedLeaf, static_cast(index)}; Value collection = context.fragmentCollections.lookup(key); SmallVector shape{exchange.program.specializationCount}; + if (leaf.enclosingLoop) + shape.push_back(leaf.loopIterationCount); llvm::append_range(shape, leaf.reconstructedType.getShape()); if (!collection || collection.getType() != RankedTensorType::get(shape, leaf.reconstructedType.getElementType())) return failure(); leafStacks.push_back(collection); @@ -337,12 +446,102 @@ FailureOr realizeDeferredResult(const DeferredResultPlan &plan, Value lan if (!value.getType().isIndex()) selected = arith::IndexCastOp::create(context.rewriter, loc, value.getType(), selected); mapping.map(value, selected); } - for (auto [index, leaf] : llvm::enumerate(exchange.program.leaves)) { Value selected = extractMixedSliceOrIdentity(context.rewriter, loc, leafStacks[index], leaf.reconstructedType, leadingSlice(context.rewriter, leaf.reconstructedType, specialization)); if (!selected) return failure(); mapping.map(leaf.replacementRoot, selected); } - Value fragment = cloneResidual(exchange, mapping, context); + SmallVector selectedLeaves; + for (auto [index, leaf] : llvm::enumerate(exchange.program.leaves)) { + RankedTensorType selectedType = + leaf.reconstructedType; + if (leaf.enclosingLoop) { + SmallVector shape { + static_cast(leaf.loopIterationCount)}; + llvm::append_range(shape, leaf.reconstructedType.getShape()); + selectedType = RankedTensorType::get( + shape, leaf.reconstructedType.getElementType()); + } + Value selected = extractMixedSliceOrIdentity(context.rewriter, loc, leafStacks[index], selectedType, leadingSlice(context.rewriter, selectedType, specialization)); + if (!selected) return failure(); + selectedLeaves.push_back(selected); + if (!leaf.enclosingLoop) + mapping.map(leaf.replacementRoot, selected); + } + Value fragment = cloneResidual( + exchange, mapping, selectedLeaves, + context); if (!fragment || fragment.getType() != fragmentType) return failure(); yielded.push_back(insertMixedSlice(context.rewriter, loc, fragment, iterArgs.front(), leadingSlice(context.rewriter, fragmentType, specialization))); return success(); }); return failed(loop) || loop->results.size() != 1 ? FailureOr(failure()) : FailureOr(loop->results.front()); } + +bool canRealizeDeferredLoopResult(const DeferredResultPlan &plan) { + DeferredExchangePlan &exchange = *plan.exchange; + if (exchange.program.specializationCount != 1 + || exchange.program.leaves.size() != 1) + return false; + const DeferredProjectionLeafTemplate &leaf = exchange.program.leaves.front(); + scf::ForOp loop = leaf.enclosingLoop; + return loop + && llvm::count_if(exchange.program.residualOps, [](Operation *op) { + return !op->hasTrait(); + }) == 1 + && llvm::is_contained(exchange.program.residualOps, loop.getOperation()) + && loop.getInitArgs().size() == 1 && loop.getNumResults() == 1 + && exchange.program.yieldedValue == loop.getResult(0) + && (leaf.form != DeferredLeafForm::GraphBatchProjection + || getDeferredProjectedFragmentType(leaf) == leaf.reconstructedType); +} + +FailureOr realizeDeferredLoopResult( + const DeferredResultPlan &plan, Value lane, + DeferredLoopEmitter emitLoop, + DeferredEmissionContext &context) { + if (!canRealizeDeferredLoopResult(plan)) + return failure(); + DeferredExchangePlan &exchange = *plan.exchange; + scf::ForOp oldLoop = exchange.program.leaves.front().enclosingLoop; + IRMapping mapping; + if (exchange.program.scheduledLane) + mapping.map(exchange.program.scheduledLane, lane); + for (auto &[value, grid] : plan.residualValues) { + Value selected = grid.emitLookup( + context.constants.getIndex(0), + lane ? lane : context.constants.getIndex(0), exchange.deferred, + context.constants, context.rewriter, exchange.deferred.getLoc()); + if (!value.getType().isIndex()) + selected = arith::IndexCastOp::create( + context.rewriter, exchange.deferred.getLoc(), value.getType(), selected); + mapping.map(value, selected); + } + for (Operation &op : + exchange.deferred.getBody().front().without_terminator()) { + if (!op.hasTrait()) + continue; + Operation *copy = context.rewriter.clone(op, mapping); + for (auto [oldValue, newValue] : + llvm::zip(op.getResults(), copy->getResults())) + mapping.map(oldValue, newValue); + } + Value initial = mapping.lookupOrDefault(oldLoop.getInitArgs().front()); + auto result = emitLoop( + initial, + [&](Value fragment, Value iteration, Value current) -> FailureOr { + if (fragment.getType() + != exchange.program.leaves.front().reconstructedType) + return failure(); + SmallVector yielded; + if (failed(cloneLoopIteration( + exchange, oldLoop, mapping, iteration, ValueRange {current}, + [&](unsigned leafIndex) -> FailureOr { + return leafIndex == 0 ? FailureOr(fragment) + : FailureOr(failure()); + }, + context, yielded)) + || yielded.size() != 1) + return failure(); + return yielded.front(); + }); + return succeeded(result) + && result->getType() == exchange.deferred.getOutput().getType() + ? result : FailureOr(failure()); +} } // namespace onnx_mlir::spatial diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp index 37c2018..c7acebd 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp @@ -2,6 +2,7 @@ #include "DeferredCommunicationModel.hpp" #include "src/Accelerators/PIM/Common/IR/StaticIntGrid.hpp" +#include "llvm/ADT/STLFunctionalExtras.h" #include namespace onnx_mlir::spatial { @@ -54,4 +55,16 @@ mlir::FailureOr realizeDeferredResult(const DeferredResultPlan& pla mlir::Value lane, DeferredEmissionContext& context); +using DeferredLoopIterationEmitter = llvm::function_ref< + mlir::FailureOr(mlir::Value, mlir::Value, mlir::Value)>; +using DeferredLoopEmitter = llvm::function_ref( + mlir::Value, DeferredLoopIterationEmitter)>; + +bool canRealizeDeferredLoopResult(const DeferredResultPlan &plan); + +mlir::FailureOr realizeDeferredLoopResult( + const DeferredResultPlan &plan, mlir::Value lane, + DeferredLoopEmitter emitLoop, + DeferredEmissionContext &context); + } // namespace onnx_mlir::spatial diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp index 0a76450..f6e554b 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp @@ -133,11 +133,14 @@ static LogicalResult buildRequirementFamilies(DeferredTransferPlan& plan, for (unsigned specialization = 0; specialization < exchange.program.specializationCount; ++specialization) { - DeferredLaneValueEvaluator evaluator( - exchange.program, exchange.targetLaneCount, specialization); for (auto leafItem : llvm::enumerate(exchange.program.leaves)) { unsigned leafIndex = leafItem.index(); const DeferredProjectionLeafTemplate &leaf = leafItem.value(); + for (unsigned iteration = 0; + iteration < leaf.loopIterationCount; ++iteration) { + DeferredLaneValueEvaluator evaluator( + exchange.program, exchange.targetLaneCount, specialization, + leaf.enclosingLoop, iteration); auto sourceIndices = evaluator.resolveSourceOperandIndices(leaf.sourceRoot); if (failed(sourceIndices)) return failure(); @@ -273,7 +276,9 @@ static LogicalResult buildRequirementFamilies(DeferredTransferPlan& plan, ++end; RequirementFamily family; family.exchange = &exchange; - family.coordinate = {specialization, leafIndex, position}; + family.coordinate = { + specialization, leafIndex, + iteration * positionCount + position}; family.targetLanes = LaneSet::range(begin, end); family.producer = producers[begin]; family.publicationFragmentType = fragmentTypes[begin]; @@ -297,6 +302,7 @@ static LogicalResult buildRequirementFamilies(DeferredTransferPlan& plan, begin = end; } } + } } } return success(); diff --git a/src/PIM/Pass/PIMPasses.h b/src/PIM/Pass/PIMPasses.h index ad5fa09..cc5d513 100644 --- a/src/PIM/Pass/PIMPasses.h +++ b/src/PIM/Pass/PIMPasses.h @@ -28,6 +28,8 @@ std::unique_ptr createTrivialGraphComputeMergePass( std::unique_ptr createPimHostConstantFoldingPass(); +std::unique_ptr createPimInstructionSelectionPass(); + std::unique_ptr createPimLocalMemoryPlanningPass(); std::unique_ptr createPimVerificationPass(); diff --git a/src/PIM/PimAccelerator.cpp b/src/PIM/PimAccelerator.cpp index 4fbe760..2c2ae5e 100644 --- a/src/PIM/PimAccelerator.cpp +++ b/src/PIM/PimAccelerator.cpp @@ -79,6 +79,7 @@ void PimAccelerator::registerPasses(int optLevel) const { mlir::registerPass([] { return createTrivialGraphComputeMergePass(); }); mlir::registerPass([] { return createMergeComputeNodesPass(); }); registerPass(createPimHostConstantFoldingPass); + registerPass(createPimInstructionSelectionPass); registerPass(createPimLocalMemoryPlanningPass); registerPass(createPimVerificationPass); registerPass(createEmitPimCodePass); diff --git a/validation/README.md b/validation/README.md index 74bcb47..34a61ab 100644 --- a/validation/README.md +++ b/validation/README.md @@ -58,7 +58,8 @@ Validate a network or network slice: ``` `--operations-dir` may point to any directory tree containing `.onnx` files. -The script discovers them recursively. +The script discovers them recursively and writes `validation_results.csv` in +that directory while retaining the terminal table. ## Raptor vs PIMCOMP comparison diff --git a/validation/networks/pimcomp_models/README.md b/validation/networks/pimcomp_models/README.md index c4f52fe..88401a5 100644 --- a/validation/networks/pimcomp_models/README.md +++ b/validation/networks/pimcomp_models/README.md @@ -1,26 +1,85 @@ -# PIMCOMP paper models +# PIMCOMP comparison models This directory contains the four networks evaluated in [PIMCOMP: An End-to-End DNN Compiler for Processing-In-Memory Accelerators](https://arxiv.org/pdf/2411.09159): -VGG-8, ResNet-18, ResNet-34, and GoogLeNet. +VGG-8, ResNet-18, ResNet-34, and GoogLeNet. It also contains YOLO11n as an +additional compiler comparison model. See the runner-generated [results.csv](results.csv) for the current latency and energy results. ## Models and provenance -| Directory | Model | Input | Provenance | -|--------------|----------------------|---------------|-------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| -| `resnet18/` | ResNet-18 v1 | `1x3x224x224` | Symlink to the complete [ONNX Model Zoo `resnet18-v1-7`](https://huggingface.co/onnxmodelzoo/resnet18-v1-7) model already present at `../resnet18/depth_68/resnet18_depth_68.onnx`. | -| `resnet34/` | ResNet-34 v1 | `1x3x224x224` | [ONNX Model Zoo `resnet34-v1-7`](https://huggingface.co/onnxmodelzoo/resnet34-v1-7), with its symbolic batch fixed to 1 as PIMCOMP's frontend does. | -| `googlenet/` | GoogLeNet | `1x3x224x224` | Unmodified [ONNX Model Zoo `googlenet-12`](https://huggingface.co/onnxmodelzoo/googlenet-12). | -| `vgg8/` | VGG-8 reconstruction | `1x1x28x28` | Reconstruction of the [PIMCOMP VGG-8 benchmark](https://arxiv.org/html/2411.09159#S8.SS1), with six convolution and two fully connected layers. | +| Directory | Model | Input | Provenance | +|--------------|----------------------|---------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| `resnet18/` | ResNet-18 v1 | `1x3x224x224` | Symlink to the complete [ONNX Model Zoo `resnet18-v1-7`](https://huggingface.co/onnxmodelzoo/resnet18-v1-7) model already present at `../resnet18/depth_68/resnet18_depth_68.onnx`. | +| `resnet34/` | ResNet-34 v1 | `1x3x224x224` | [ONNX Model Zoo `resnet34-v1-7`](https://huggingface.co/onnxmodelzoo/resnet34-v1-7), with its symbolic batch fixed to 1 as PIMCOMP's frontend does. | +| `googlenet/` | GoogLeNet | `1x3x224x224` | Unmodified [ONNX Model Zoo `googlenet-12`](https://huggingface.co/onnxmodelzoo/googlenet-12). | +| `vgg8/` | VGG-8 reconstruction | `1x1x28x28` | Reconstruction of the [PIMCOMP VGG-8 benchmark](https://arxiv.org/html/2411.09159#S8.SS1), with six convolution and two fully connected layers. | +| `yolo11n/` | YOLO11n detection | `1x3x640x640` | Derived from the canonical local model at `../yolo11n/depth_51/yolo11n_depth_51.onnx`, exported from [Ultralytics YOLO11n](https://github.com/ultralytics/ultralytics/blob/main/docs/en/models/yolo11.md). | -`googlenet/googlenet-12-latency.onnx` is the explicit common latency model. -It exposes the original model's final FC logits (`loss3/classifier_1`) and -removes its two LRN nodes and terminal Softmax so the comparison covers only -operations scheduled by PIMCOMP. Keep `googlenet-12.onnx` as the unmodified -source model. +`googlenet/googlenet-12-latency.onnx` is the explicit pimsim-nn-ready GoogLeNet model. +It removes the two LRN nodes and terminal Softmax from the original model, +so that the comparison covers only operations scheduled by PIMCOMP and supported by pimsim-nn. + +`yolo11n/yolo11n-latency.onnx` is the explicit pimsim-nn-ready YOLO11n model. +It removes the Softmax nodes from the original model, +so that the compiled artifact can be simulated in pimsim-nn. + +## Unsupported and ignored operations + +PIMCOMP's frontend accepts exactly these ONNX operations: + +```text +Add, AveragePool, BatchNormalization, Clip, Concat, Conv, Dropout, Flatten, +Gather, Gemm, GlobalAveragePool, LRN, MatMul, MaxPool, Mul, Pad, Relu, Reshape, +Shape, Sigmoid, Softmax, Squeeze, Sub, Sum, Tanh, Transpose, Unsqueeze +``` + +`Constant` is consumed as frontend metadata rather than emitted as a PIMCOMP +node. Every other ONNX operation is unsupported: the frontend prints +`operation: not considered` and stops at the first occurrence. Thus the +complete unsupported set is the complement of the allowlist above for the +model's ONNX opset. In particular, YOLO11n contains unsupported `Split` and +`Resize` nodes. + +PIMCOMP's low-latency scheduler, hierarchy mapper, and genetic algorithm use +this complete explicit no-consider set: + +```text +Input, BatchNormalization, Clip, Dropout, Flatten, LRN, MatMul, Reshape, +Softmax, Squeeze, Transpose +``` + +Those nodes do not receive scheduled latency instructions when they remain in +the backend graph. Before that point the frontend may fuse BatchNormalization +and activation nodes into Conv/Gemm, convert a +Reshape-Transpose-Reshape channel-shuffle pattern to `OP_SHUFFLE`, remove a +specific Shape-Gather-Unsqueeze-Concat shape chain, and merge Pad into its +consumer. These transformations do not make an otherwise standalone ignored +operation timed. + +`pimsim-nn` consumes PIM ISA instructions. It supports every named opcode +in the shared serialized range except `vsoftmax` (opcode 21), +which is rejected explicitly in both JSON and binary input. It +silently ignores no opcode; unknown names and numbers are errors. + +These boundaries explain the dedicated artifacts: + +- GoogLeNet's two LRN nodes and terminal Softmax perform real computation but + are ignored by PIMCOMP, so the common latency artifact removes them. Its + inference Dropout and shape-only Reshape can remain without adding compute. +- YOLO11n's latency artifact bypasses exactly its two Softmax nodes so it can + run in `pimsim-nn`. Every other node, including MatMul, Transpose, and the + final detection-decoding tail, remains present and timed by Raptor. No + PIMCOMP latency is reported because its frontend stops at `Split` and also + lacks `Resize`; compiling that prefix would not represent YOLO11n. + +The authoritative lists are in +[`frontend.py`](../../../third_party/PIMCOMP-NN/frontend/frontend.py), +[`ElementPipelineSchedule.cpp`](../../../third_party/PIMCOMP-NN/backend/ElementPipelineSchedule.cpp), +[`ISA.h`](../../../backend-simulators/pim/pimsim-nn/src/isa/ISA.h), and +[`Instruction.cpp`](../../../backend-simulators/pim/pimsim-nn/src/isa/Instruction.cpp). The PIMCOMP authors did not publish the ONNX checkpoints used by the paper. Running PIMCOMP's frontend on the three Model Zoo files above produces JSON @@ -41,6 +100,7 @@ c3231061d081bdd47884137b02134f85142752a39e87263c529cd14ed242b096 resnet34/resne c99c507058eaf41de8723408fdda7db8325cb57f0a89f2ee07a716d6e963e14e googlenet/googlenet-12.onnx a26f9e33901c573e60c34a3f0abbb4744fff83e4e0f21b18fc66e20395e72982 googlenet/googlenet-12-latency.onnx 396cdea21e5e7d02c3f26f14d22ef20975171702493f5c5e79b8e0d896e541ef vgg8/vgg8-mnist-reconstructed.onnx +229f3975af8933d39aee8d9031d969bff074b69c33e304a78abb35ff0c5f445f yolo11n/yolo11n-latency.onnx ``` ## Paper hardware profiles @@ -190,6 +250,13 @@ RAPTOR_ROOT=$PWD --pimcomp-pipeline element ``` +Use the same command with +`yolo11n/yolo11n-latency.onnx` to probe YOLO11n. Released PIMCOMP-NN cannot +compile it: the frontend stops at `/model.2/Split`, and it also has no mapping +for YOLO11n's two nearest-neighbor `Resize` nodes. Treating the emitted prefix +as YOLO11n would produce a misleading latency, so no PIMCOMP number is +reported for this model. + For Arch-A high throughput, use `--pimsim-mode throughput --pimcomp-pipeline batch`. For Arch-B, use 138 cores, 128 crossbars, a `6x23` mesh, and diff --git a/validation/networks/pimcomp_models/results.csv b/validation/networks/pimcomp_models/results.csv index 878a2a2..49345f5 100644 --- a/validation/networks/pimcomp_models/results.csv +++ b/validation/networks/pimcomp_models/results.csv @@ -1,5 +1,5 @@ model,raptor_latency_ms,pimcomp_latency_ms,raptor_energy_pj,pimcomp_energy_pj,faster_compiler,speedup -vgg8,2.463417,7.985074,649524564.120001,1597904071.120000,raptor,3.24 -resnet18,31.449811,58.853733,9335506856.119984,13983168748.119974,raptor,1.87 -resnet34,58.071522,91.607980,17027195415.679953,22722922369.680016,raptor,1.58 -googlenet,24.247113,62.923463,8019605936.239990,14547526780.240000,raptor,2.60 +vgg8,1.465778,7.985074,477298145.040001,1597904071.120000,raptor,5.45 +resnet18,28.099952,58.853733,8781611766.119984,13983168748.119974,raptor,2.09 +resnet34,45.781486,91.607980,14962940227.679951,22722922369.680016,raptor,2.00 +googlenet,13.371204,62.923463,6117835798.919991,14547526780.240000,raptor,4.71 diff --git a/validation/networks/pimcomp_models/yolo11n/yolo11n-latency.onnx b/validation/networks/pimcomp_models/yolo11n/yolo11n-latency.onnx new file mode 100644 index 0000000..a043ea6 Binary files /dev/null and b/validation/networks/pimcomp_models/yolo11n/yolo11n-latency.onnx differ diff --git a/validation/operations/README.md b/validation/operations/README.md index 22c368d..861a54e 100644 --- a/validation/operations/README.md +++ b/validation/operations/README.md @@ -38,11 +38,12 @@ Run the complete suite with deadlock detection: Use `--compile-only` for compiler and deadlock checks, then `--run-only` to reuse those artifacts for reference execution, simulation, and comparison. -The validator prints the complete operation results table before its summary. +The validator prints the complete operation results table before its summary +and writes the same rows to `validation_results.csv`. ## Complete inventory -The suite contains 164 models. Tensor shapes, attributes, and constants are +The suite contains 165 models. Tensor shapes, attributes, and constants are defined in `gen_tests.py` and in the checked-in ONNX models. ### Add (5) @@ -63,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 (31) +### Conv (32) | Case | Description | |---|---| @@ -98,6 +99,7 @@ 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_stem` | First two YOLO11n `Conv-SiLU` blocks at `640x640`, including the distributed activation boundary. | ### Div (6) diff --git a/validation/operations/conv/yolo11n_stem/conv_yolo11n_stem.onnx b/validation/operations/conv/yolo11n_stem/conv_yolo11n_stem.onnx new file mode 100644 index 0000000..d54ff19 Binary files /dev/null and b/validation/operations/conv/yolo11n_stem/conv_yolo11n_stem.onnx differ diff --git a/validation/operations/gen_tests.py b/validation/operations/gen_tests.py index d77a417..9ca077f 100644 --- a/validation/operations/gen_tests.py +++ b/validation/operations/gen_tests.py @@ -218,6 +218,30 @@ def conv_huge_pointwise_1024_dynamic(): save_model(model, "conv/huge_pointwise_1024_dynamic", "conv_huge_pointwise_1024_dynamic.onnx") +def conv_yolo11n_stem(): + """First two YOLO11n Conv-SiLU blocks, preserving the distributed intermediate.""" + X = helper.make_tensor_value_info("X", TensorProto.FLOAT, [1, 3, 640, 640]) + Y = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1, 32, 160, 160]) + rng = np.random.default_rng(109) + W0 = numpy_helper.from_array(rng.uniform(-1, 1, (16, 3, 3, 3)).astype(np.float32), name="W0") + B0 = numpy_helper.from_array(rng.uniform(-1, 1, (16,)).astype(np.float32), name="B0") + W1 = numpy_helper.from_array(rng.uniform(-1, 1, (32, 16, 3, 3)).astype(np.float32), name="W1") + B1 = numpy_helper.from_array(rng.uniform(-1, 1, (32,)).astype(np.float32), name="B1") + nodes = [ + helper.make_node("Conv", ["X", "W0", "B0"], ["C0"], + kernel_shape=[3, 3], strides=[2, 2], 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=[2, 2], 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_stem", [X], [Y], initializer=[W0, B0, W1, B1]) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 13)]) + save_model(model, "conv/yolo11n_stem", "conv_yolo11n_stem.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]) @@ -2032,6 +2056,7 @@ if __name__ == "__main__": conv_dynamic() conv_huge_pointwise_1024() conv_huge_pointwise_1024_dynamic() + conv_yolo11n_stem() conv_pointwise_tiled_chain() conv_large_output_channels_1x1() conv_large_input_channels_1x1() diff --git a/validation/operations/validation_results.csv b/validation/operations/validation_results.csv new file mode 100644 index 0000000..e00c427 --- /dev/null +++ b/validation/operations/validation_results.csv @@ -0,0 +1,166 @@ +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 diff --git a/validation/raptor_validation/pimsim_nn.py b/validation/raptor_validation/pimsim_nn.py index e1491a7..36a1a2f 100644 --- a/validation/raptor_validation/pimsim_nn.py +++ b/validation/raptor_validation/pimsim_nn.py @@ -1,8 +1,11 @@ import json import re import shutil +import struct from pathlib import Path +import numpy as np + _METRIC_PATTERNS = { "output_count": r"output count:\s+([0-9]+)\s+samples", @@ -14,6 +17,14 @@ _METRIC_PATTERNS = { } +def read_raptor_instruction_count(path): + with open(path, "rb") as f: + magic, version, count = struct.unpack("<4sII", f.read(12)) + if magic != b"PIMB" or version != 1: + raise ValueError(f"Invalid Raptor binary instruction file: {path}") + return count + + def parse_pimsim_nn_metrics(output): metrics = {"raw_output": output} for name, pattern in _METRIC_PATTERNS.items(): @@ -45,6 +56,46 @@ def export_raptor_latency_artifact(pim_dir, output_dir): json.dump(config, f, separators=(",", ":")) f.write("\n") + binary_sources = sorted(pim_dir.glob("core_*.pim"), key=lambda path: int(path.stem.split("_")[1])) + if binary_sources: + record_dtype = np.dtype([ + ("opcode", "u1"), ("rd", "u1"), ("r1", "u1"), ("flags", "u1"), + ("r2_or_imm", "= 0]]] = True + for start in range(0, count, 1_000_000): + chunk = records[start:start + 1_000_000] + setbw = chunk["opcode"] == 8 + chunk["generic1"][setbw] = 8 + chunk["generic2"][setbw] = 8 + address_sldi = (chunk["opcode"] == 1) & ~scalar_sldi[start:start + len(chunk)] + if np.any(chunk["r2_or_imm"][address_sldi] % 4): + raise ValueError(f"Raptor address is not aligned to its fp32 element width: {source}") + chunk["r2_or_imm"][address_sldi] //= 4 + transfer = (chunk["opcode"] >= 25) & (chunk["opcode"] <= 30) + if np.any(chunk["generic2"][transfer] % 4) or np.any(chunk["generic3"][transfer] % 4): + raise ValueError(f"Raptor transfer field is not aligned to its fp32 element width: {source}") + chunk["generic2"][transfer] //= 4 + chunk["generic3"][transfer] //= 4 + records.flush() + del records + return output_dir + byte_size_fields = { "ld": "size", "st": "size", @@ -56,12 +107,19 @@ def export_raptor_latency_artifact(pim_dir, output_dir): for source in sorted(pim_dir.glob("core_*.json"), key=lambda path: int(path.stem.split("_")[1])): with open(source, encoding="utf-8") as f: instructions = json.load(f) - for instruction in instructions: + last_sldi = {} + scalar_sldi = set() + for index, instruction in enumerate(instructions): + if instruction["op"] == "sldi": + last_sldi[instruction["rd"]] = index + elif instruction["op"] == "vmv" and instruction["rs2"] in last_sldi: + scalar_sldi.add(last_sldi[instruction["rs2"]]) + for index, instruction in enumerate(instructions): op = instruction["op"] if op == "setbw": instruction["ibiw"] = 8 instruction["obiw"] = 8 - elif op == "sldi": + elif op == "sldi" and index not in scalar_sldi: instruction["imm"] = int8_bytes(instruction["imm"], "address") if field := byte_size_fields.get(op): instruction[field] = int8_bytes(instruction[field], f"{op} {field}") diff --git a/validation/raptor_validation/validate_one.py b/validation/raptor_validation/validate_one.py index 89985b4..c432ffc 100644 --- a/validation/raptor_validation/validate_one.py +++ b/validation/raptor_validation/validate_one.py @@ -1,8 +1,10 @@ import json import os +import re import shutil import subprocess import sys +import time import numpy as np from dataclasses import dataclass, field from pathlib import Path @@ -10,7 +12,7 @@ from colorama import Style, Fore from .gen_network_runner import gen_network_runner from .onnx_utils import gen_random_inputs, save_inputs_to_files, onnx_io, write_inputs_to_memory_bin, _ONNX_TO_NP from .raptor import compile_with_raptor -from .pimsim_nn import export_raptor_latency_artifact, parse_pimsim_nn_metrics +from .pimsim_nn import export_raptor_latency_artifact, parse_pimsim_nn_metrics, read_raptor_instruction_count from .subprocess_utils import run_command_with_reporter STAGE_TITLES = ( @@ -80,6 +82,35 @@ class ValidationResult: pimsim_power_mw: float | None = None pimsim_energy_pj: float | None = None pimsim_status: str = PIMSIM_SKIPPED + compile_time_s: float | None = None + host_memory_bytes: int | None = None + cores_memory_bytes: int | None = None + used_core_count: int | None = None + used_crossbar_count: int | None = None + + +_MEMORY_UNITS = {"B": 1, "KB": 1 << 10, "MB": 1 << 20, "GB": 1 << 30} + + +def collect_pim_resource_metrics(pim_dir): + pim_dir = Path(pim_dir) + report_path = pim_dir.parent / "reports" / "memory_report.txt" + report = report_path.read_text(encoding="utf-8") if report_path.exists() else "" + + def memory_bytes(label): + match = re.search(rf"^\s*{re.escape(label)}:\s+([0-9.]+)\s+(B|KB|MB|GB)$", report, re.MULTILINE) + return round(float(match.group(1)) * _MEMORY_UNITS[match.group(2)]) if match else None + + with open(pim_dir / "config.json", encoding="utf-8") as f: + config = json.load(f) + used_cores = sum(read_raptor_instruction_count(path) > 0 for path in pim_dir.glob("core_*.pim")) + used_crossbars = sum(sum(groups) for groups in config.get("array_group_map", {}).values()) + return { + "host_memory_bytes": memory_bytes("Host memory"), + "cores_memory_bytes": memory_bytes("Local memory after reuse"), + "used_core_count": used_cores, + "used_crossbar_count": used_crossbars, + } class ProgressReporter: @@ -416,8 +447,6 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, pimsim_nn_build_dir = Path(pimsim_nn_build_dir).resolve() pimsim_config_path = Path(pimsim_config_path).resolve() compile_extra_args = list(raptor_extra_args or []) - if pimsim_enabled and "--pim-emit-json" not in compile_extra_args: - compile_extra_args.append("--pim-emit-json") owns_reporter = reporter is None reporter = reporter or ProgressReporter(model_total, stages_per_model=len(MODE_STAGE_TITLES[mode]), verbose=verbose) @@ -435,6 +464,8 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, f" {Style.BRIGHT}Validating {network_onnx_path.name}{Style.RESET_ALL}") failed_with_exception = False pim_pass_timings = {} + compile_time_s = None + resource_metrics = {} try: stem = network_onnx_path.stem @@ -443,6 +474,19 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, runner_path = runner_build_dir / "runner" pim_output_base = raptor_dir / stem + def compile_pim(): + nonlocal compile_time_s, resource_metrics + started = time.perf_counter() + timings = compile_with_raptor( + network_onnx_path, raptor_path, pim_output_base, crossbar_size, + crossbar_count, core_count=core_count, + raptor_extra_args=compile_extra_args, cwd=raptor_dir, + verbose=verbose, reporter=reporter, + timeout_sec=command_timeout_seconds) + compile_time_s = time.perf_counter() - started + resource_metrics = collect_pim_resource_metrics(raptor_dir / "pim") + return timings + if mode != MODE_RUN_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile ONNX") network_so_path, network_mlir_path = compile_onnx_network( @@ -468,16 +512,14 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, if mode == MODE_COMPILE_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM") - pim_pass_timings = compile_with_raptor( - network_onnx_path, raptor_path, pim_output_base, crossbar_size, crossbar_count, - core_count=core_count, - raptor_extra_args=compile_extra_args, - cwd=raptor_dir, verbose=verbose, reporter=reporter, timeout_sec=command_timeout_seconds) + pim_pass_timings = compile_pim() print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}") reporter.advance() reporter.record_result(True) reporter.log(Style.BRIGHT + f"Result: {Fore.GREEN}PASS{Style.RESET_ALL}" + Style.RESET_ALL) - return ValidationResult(passed=True, pim_pass_timings=pim_pass_timings) + return ValidationResult( + passed=True, pim_pass_timings=pim_pass_timings, + compile_time_s=compile_time_s, **resource_metrics) if mode == MODE_RUN_ONLY: required_paths = [ @@ -489,6 +531,7 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, missing = [f"{description} at {path}" for path, description in required_paths if not path.exists()] if missing: raise FileNotFoundError("run-only mode requires existing artifacts:\n " + "\n ".join(missing)) + resource_metrics = collect_pim_resource_metrics(raptor_dir / "pim") print_stage(reporter, model_index, model_total, network_onnx_path.name, "Generate Inputs") inputs_descriptor, outputs_descriptor = onnx_io(network_onnx_path) @@ -508,11 +551,7 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, if mode != MODE_RUN_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM") - pim_pass_timings = compile_with_raptor( - network_onnx_path, raptor_path, pim_output_base, crossbar_size, crossbar_count, - core_count=core_count, - raptor_extra_args=compile_extra_args, - cwd=raptor_dir, verbose=verbose, reporter=reporter, timeout_sec=command_timeout_seconds) + pim_pass_timings = compile_pim() print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}") reporter.advance() @@ -587,6 +626,8 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir, pimsim_power_mw=pimsim_power_mw, pimsim_energy_pj=pimsim_energy_pj, pimsim_status=pimsim_status, + compile_time_s=compile_time_s, + **resource_metrics, ) except Exception: failed_with_exception = True diff --git a/validation/validate.py b/validation/validate.py index 45be381..1984ae0 100755 --- a/validation/validate.py +++ b/validation/validate.py @@ -1,6 +1,7 @@ #!/usr/bin/env python3 import argparse +import csv import os import signal import subprocess @@ -140,6 +141,17 @@ def format_pimsim_metric(result, value, unit): return result.pimsim_status +def format_memory(byte_count): + if byte_count is None: + return "-" + return f"{byte_count / (1 << 20):.2f} MiB" + + +def operation_label(relative_path): + path = Path(relative_path) + return str(path.parent) if path.parent != Path(".") else path.stem + + def main(): script_dir = Path(__file__).parent.resolve() pimcomp_configs_dir = script_dir / "pimsim_configs" / "pimcomp" @@ -216,6 +228,7 @@ def main(): removed_count = 0 for onnx_path in onnx_files: removed_count += len(clean_workspace_artifacts(onnx_path.parent, onnx_path.stem)) + (operations_dir / "validation_results.csv").unlink(missing_ok=True) print(Style.BRIGHT + f"Removed {removed_count} generated artifact path(s)." + Style.RESET_ALL) sys.exit(0) @@ -343,38 +356,44 @@ def main(): # Summary n_passed = sum(1 for result in results.values() if result.passed) n_total = len(results) - status_width = len("Result") - path_width = max(len("Operation"), *(len(rel) for rel in results)) - formatted_metrics = { - rel: ( + headers = ("Operation", "Result", "Compile", "Host mem", "Cores mem", + "Cores", "Xbars", "Latency", "Power", "Energy") + rows = [] + for rel, result in results.items(): + rows.append(( + operation_label(rel), "PASS" if result.passed else "FAIL", + f"{result.compile_time_s:.3f} s" if result.compile_time_s is not None else "-", + format_memory(result.host_memory_bytes), + format_memory(result.cores_memory_bytes), + str(result.used_core_count) if result.used_core_count is not None else "-", + str(result.used_crossbar_count) if result.used_crossbar_count is not None else "-", format_pimsim_metric(result, result.pimsim_latency_ms, "ms"), format_pimsim_metric(result, result.pimsim_power_mw, "mW"), format_pimsim_metric(result, result.pimsim_energy_pj, "pJ"), - ) - for rel, result in results.items() - } - latency_width = max(len("Latency"), *(len(metrics[0]) for metrics in formatted_metrics.values())) - power_width = max(len("Power"), *(len(metrics[1]) for metrics in formatted_metrics.values())) - energy_width = max(len("Energy"), *(len(metrics[2]) for metrics in formatted_metrics.values())) - separator = ( - f"+-{'-' * path_width}-+-{'-' * status_width}-+-{'-' * latency_width}" - f"-+-{'-' * power_width}-+-{'-' * energy_width}-+") + )) + widths = [max(len(header), *(len(row[index]) for row in rows)) + for index, header in enumerate(headers)] + separator = "+-" + "-+-".join("-" * width for width in widths) + "-+" + + def table_line(row): + return "| " + " | ".join( + value.ljust(widths[index]) if index < 2 else value.rjust(widths[index]) + for index, value in enumerate(row)) + " |" + print(separator) - print( - f"| {'Operation'.ljust(path_width)} | {'Result'.ljust(status_width)} | " - f"{'Latency'.rjust(latency_width)} | {'Power'.rjust(power_width)} | " - f"{'Energy'.rjust(energy_width)} |" - ) + print(table_line(headers)) print(separator) - for rel, result in results.items(): - plain_status = "PASS" if result.passed else "FAIL" - status = Fore.GREEN + plain_status.ljust(status_width) + Style.RESET_ALL if result.passed else \ - Fore.RED + plain_status.ljust(status_width) + Style.RESET_ALL - latency, power, energy = formatted_metrics[rel] - print( - f"| {rel.ljust(path_width)} | {status} | {latency.rjust(latency_width)} | " - f"{power.rjust(power_width)} | {energy.rjust(energy_width)} |") + for row in rows: + line = table_line(row) + color = Fore.GREEN if row[1] == "PASS" else Fore.RED + line = line.replace(row[1].ljust(widths[1]), + color + row[1].ljust(widths[1]) + Style.RESET_ALL, 1) + print(line) print(separator) + with (operations_dir / "validation_results.csv").open( + "w", encoding="utf-8", newline="" + ) as results_file: + csv.writer(results_file).writerows((headers, *rows)) print("\n" + Style.BRIGHT + Fore.CYAN + "Summary" + Style.RESET_ALL) print(Style.BRIGHT + f"Passed: {n_passed}" + Style.RESET_ALL) print(Style.BRIGHT + f"Failed: {n_total - n_passed}" + Style.RESET_ALL)