big refactor
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-08-04 11:28:05 +02:00
parent f4a3b012cc
commit 10b6ee6c32
150 changed files with 6737 additions and 4816 deletions
+3 -19
View File
@@ -70,11 +70,6 @@ llvm::cl::opt<bool>
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> useExperimentalConvImpl("use-experimental-conv-impl",
llvm::cl::desc("Use experimental implementation for convolution"),
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<uint64_t> pimConvIm2colMaxElements(
"pim-conv-im2col-max-elements",
llvm::cl::desc("Maximum number of im2col elements to materialize globally for one Conv before streaming/chunking"),
@@ -103,15 +98,9 @@ llvm::cl::opt<bool> pimDetectCommunicationDeadlock(
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimMaterializeScalarFanoutGlobalOrder(
"pim-materialize-scalar-fanout-global-order",
llvm::cl::desc("Experimental expensive materializer mode: emit scalar-source fanout as globally ordered communication events instead of all-send fanout loops"),
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimTraceCommunicationMaterialization(
"pim-trace-communication-materialization",
llvm::cl::desc("Emit verbose materializer-time diagnostics and provenance attributes for every Spatial communication op"),
llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom(
"pim-verify-bufferization-copy-freedom",
llvm::cl::desc("Run the expensive official PIM tensor-copy freedom proof before bufferization"),
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
@@ -131,11 +120,6 @@ llvm::cl::opt<std::string> pimTargetConfig(
llvm::cl::init(""),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool>
ignoreConcatError("ignore-concat-error",
llvm::cl::desc("Ignore ConcatOp corner case: do not assert and do a simplification"),
llvm::cl::init(false));
bool hasExplicitPimCoreCount() { return coresCount.getNumOccurrences() != 0; }
void verifyExplicitPimCoreCount() {
+1 -11
View File
@@ -55,12 +55,10 @@ extern llvm::cl::opt<PimConvLoweringType> pimConvLowering;
extern llvm::cl::opt<PimSpatialDataflowExportType> pimExportSpatialDataflow;
extern llvm::cl::opt<bool> pimOnlyCodegen;
extern llvm::cl::opt<bool> useExperimentalConvImpl;
extern llvm::cl::opt<bool> pimEmitJson;
extern llvm::cl::opt<bool> pimReportConvLowering;
extern llvm::cl::opt<bool> pimDetectCommunicationDeadlock;
extern llvm::cl::opt<bool> pimMaterializeScalarFanoutGlobalOrder;
extern llvm::cl::opt<bool> pimTraceCommunicationMaterialization;
extern llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom;
extern llvm::cl::opt<size_t> crossbarSize;
extern llvm::cl::opt<size_t> crossbarCountInCore;
@@ -72,12 +70,4 @@ extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions;
bool hasExplicitPimCoreCount();
void verifyExplicitPimCoreCount();
// This option, by default set to false, will ignore an error when resolving a
// specific tiles of the operands of a concat. This specific case is when the
// wanted tile is generated by two separate operands of the concat. If this is
// set to false, this corner case will assert an error. If this is set to true,
// a simplification is performed and only the tile from the first operand is
// taken.
extern llvm::cl::opt<bool> ignoreConcatError;
} // namespace onnx_mlir
+71 -12
View File
@@ -14,9 +14,12 @@
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialOptions.hpp"
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/Scheduling/SchedulingTarget.hpp"
#include "src/Accelerators/PIM/Pass/PIMPasses.h"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetResources.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/SpatialDataflowCsvExporter.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/SchedulingTarget.hpp"
#include "src/Accelerators/PIM/Passes/PIMPasses.h"
#include "src/Compiler/CompilerPasses.hpp"
#define DEBUG_TYPE "PimCompilerUtils"
@@ -80,6 +83,54 @@ spatial::SchedulingTarget getDefaultPimSchedulingTarget() {
return target;
}
spatial::ConvLoweringStrategy getSpatialConvLoweringStrategy(PimConvLoweringType strategy) {
switch (strategy) {
case PimConvLoweringAuto: return spatial::ConvLoweringStrategy::Auto;
case PimConvLoweringLegacy: return spatial::ConvLoweringStrategy::Legacy;
case PimConvLoweringDepthwise: return spatial::ConvLoweringStrategy::Depthwise;
case PimConvLoweringPackedIm2Col: return spatial::ConvLoweringStrategy::PackedIm2Col;
case PimConvLoweringStreamedPatch: return spatial::ConvLoweringStrategy::StreamedPatch;
case PimConvLoweringStreamedPacked: return spatial::ConvLoweringStrategy::StreamedPacked;
case PimConvLoweringOutputChannelTiled: return spatial::ConvLoweringStrategy::OutputChannelTiled;
case PimConvLoweringInputKTiled: return spatial::ConvLoweringStrategy::InputKTiled;
case PimConvLoweringTiled2D: return spatial::ConvLoweringStrategy::Tiled2D;
}
llvm_unreachable("unknown PIM Conv lowering strategy");
}
spatial::SpatialDataflowExportStage getPimSpatialDataflowExportStage(
PimSpatialDataflowExportType stage) {
switch (stage) {
case SpatialDataflowExportNone: return spatial::SpatialDataflowExportStage::None;
case SpatialDataflowExportSpatial1: return spatial::SpatialDataflowExportStage::Spatial1;
case SpatialDataflowExportSpatial2: return spatial::SpatialDataflowExportStage::Spatial2;
case SpatialDataflowExportSpatial3: return spatial::SpatialDataflowExportStage::Spatial3;
case SpatialDataflowExportSpatial4: return spatial::SpatialDataflowExportStage::Spatial4;
case SpatialDataflowExportAll: return spatial::SpatialDataflowExportStage::All;
}
llvm_unreachable("unknown PIM Spatial dataflow export stage");
}
spatial::SpatialTargetResources getPimSpatialTargetResources(const spatial::SchedulingTarget& target) {
spatial::SpatialTargetResources resources;
resources.matrixShape = {target.matrixRows, target.matrixColumns};
resources.matrixUnitsPerProcessor = target.residentWeightCapacity;
resources.processorCount = target.processorCount;
resources.vectorWidth = target.vectorWidth;
if (failed(resources.verify()))
llvm::report_fatal_error("PIM target resources are incomplete");
return resources;
}
ONNXToSpatialPlanningOptions getPimONNXToSpatialPlanningOptions() {
ONNXToSpatialPlanningOptions options;
options.convIm2colMaxElements = pimConvIm2colMaxElements.getValue();
options.convStreamChunkPositions = pimConvStreamChunkPositions.getValue();
options.forcedConvStrategy = getSpatialConvLoweringStrategy(pimConvLowering.getValue());
options.reportConvLowering = pimReportConvLowering.getValue();
return options;
}
const llvm::json::Object& requireObject(const llvm::json::Object& object,
llvm::StringRef key,
llvm::StringRef path) {
@@ -279,11 +330,13 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
EmissionTargetType& emissionTarget,
std::string outputNameNoExt) {
verifyExplicitPimCoreCount();
spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget();
spatial::SpatialTargetResources targetResources = getPimSpatialTargetResources(schedulingTarget);
if (pimOnlyCodegen) {
pm.addPass(createPimInstructionSelectionPass());
pm.addPass(createPimLocalMemoryPlanningPass());
pm.addPass(createPimVerificationPass());
pm.addPass(createPimVerificationPass(targetResources, pimDetectCommunicationDeadlock.getValue()));
pm.addPass(createEmitPimCodePass());
return;
}
@@ -292,23 +345,29 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
addONNXToMLIRPasses(pm, /*target CPU*/ false);
if (pimEmissionTarget >= EmitSpatial) {
spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget();
pm.addPass(createONNXToSpatialPass());
pm.addPass(createSpatialLayoutPlanningPass());
pm.addPass(createLowerSpatialPlansPass());
ONNXToSpatialPlanningOptions planningOptions = getPimONNXToSpatialPlanningOptions();
spatial::SpatialDataflowExportStage exportStage =
getPimSpatialDataflowExportStage(pimExportSpatialDataflow.getValue());
pm.addPass(createONNXToSpatialPass(targetResources, planningOptions));
pm.addPass(createSpatialLayoutPlanningPass(targetResources));
pm.addPass(createLowerSpatialPlansPass(targetResources, planningOptions, exportStage));
pm.addPass(createTrivialGraphComputeMergePass(
schedulingTarget.residentWeightCapacity));
pm.addPass(createMergeComputeNodesPass(schedulingTarget));
schedulingTarget.residentWeightCapacity, exportStage));
pm.addPass(spatial::createScheduleAndRealizeSpatialPass(
schedulingTarget, exportStage));
pm.addPass(createMessagePass("Onnx lowered to Spatial"));
}
if (pimEmissionTarget >= EmitPim) {
pm.addPass(createSpatialToPimPass());
pm.addPass(createSpatialToPimPass(targetResources));
pm.addPass(createMessagePass("Spatial lowered to Pim"));
}
if (pimEmissionTarget >= EmitPimBufferized) {
pm.addPass(createPimBufferizationPass());
pm.addPass(createPimBufferizationPreparationPass(pimVerifyBufferizationCopyFreedom.getValue()));
pm.addPass(createPimOneShotBufferizationPass());
pm.addPass(createPimMemoryNormalizationPass());
pm.addPass(createPimBufferizationVerificationPass());
pm.addPass(createMessagePass("Pim bufferized"));
}
@@ -320,7 +379,7 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
pm.addPass(createMessagePass("Pim instructions selected"));
pm.addPass(createPimLocalMemoryPlanningPass());
pm.addPass(createMessagePass("Pim local memory planned"));
pm.addPass(createPimVerificationPass());
pm.addPass(createPimVerificationPass(targetResources, pimDetectCommunicationDeadlock.getValue()));
pm.addPass(createMessagePass("Pim verified"));
pm.addPass(createEmitPimCodePass());
pm.addPass(createMessagePass("Pim code emitted"));