blazingly fasterer
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-07-21 15:43:35 +02:00
parent a893d23a74
commit 3e468b58c8
37 changed files with 1119 additions and 933 deletions
+67 -26
View File
@@ -3,6 +3,7 @@
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "llvm/ADT/DenseSet.h"
#include "llvm/ADT/STLExtras.h"
#include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp"
@@ -12,8 +13,10 @@ namespace onnx_mlir {
bool isCoreStaticAddressOp(mlir::Operation* op) {
if (mlir::isa<mlir::affine::AffineApplyOp,
mlir::arith::ConstantOp, mlir::arith::AddIOp,
mlir::arith::SubIOp, mlir::arith::MulIOp,
mlir::arith::ConstantOp,
mlir::arith::AddIOp,
mlir::arith::SubIOp,
mlir::arith::MulIOp,
mlir::arith::DivUIOp,
mlir::arith::DivSIOp,
mlir::arith::MinUIOp,
@@ -33,13 +36,13 @@ bool isCoreStaticAddressOp(mlir::Operation* op) {
namespace {
enum class CoreWalkMode { ExecuteAllIterations, StructuralExtremes };
using CoreWalkCallback =
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)>;
enum class CoreWalkMode {
ExecuteCommunication,
StructuralExtremes
};
using CoreWalkCallback = llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)>;
static void propagateRegionResults(mlir::ValueRange results,
mlir::Region& region,
StaticValueKnowledge& knowledge) {
static void propagateRegionResults(mlir::ValueRange results, mlir::Region& region, StaticValueKnowledge& knowledge) {
if (region.empty())
return;
auto yield = mlir::cast<mlir::scf::YieldOp>(region.front().getTerminator());
@@ -50,23 +53,38 @@ static void propagateRegionResults(mlir::ValueRange results,
static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
const StaticValueKnowledge& initialKnowledge,
CoreWalkMode mode,
const PimCoreCommunicationPlan* communicationPlan,
CoreWalkCallback callback) {
bool hasFailure = false;
StaticValueKnowledge knowledge = initialKnowledge;
llvm::StringRef purpose = mode == CoreWalkMode::ExecuteAllIterations ? "codegen" : "verification";
for (mlir::Operation& op : block) {
llvm::StringRef purpose = mode == CoreWalkMode::ExecuteCommunication ? "communication verification" : "verification";
llvm::SmallVector<mlir::Operation*, 0> structuralOperations;
llvm::ArrayRef<mlir::Operation*> operations;
if (communicationPlan) {
auto it = communicationPlan->find(&block);
if (it != communicationPlan->end())
operations = it->second;
}
else {
structuralOperations.reserve(block.getOperations().size());
for (mlir::Operation& op : block)
structuralOperations.push_back(&op);
operations = structuralOperations;
}
for (mlir::Operation* operation : operations) {
mlir::Operation& op = *operation;
if (mlir::isa<pim::PimHaltOp, mlir::scf::YieldOp>(op) || isCoreStaticAddressOp(&op))
continue;
if (auto loadOp = mlir::dyn_cast<mlir::memref::LoadOp>(op);
loadOp && succeeded(resolveIndexValue(loadOp.getResult(), knowledge)))
continue;
if (auto forOp = mlir::dyn_cast<mlir::scf::ForOp>(op)) {
auto lower = resolveIndexValue(forOp.getLowerBound(), knowledge);
auto upper = resolveIndexValue(forOp.getUpperBound(), knowledge);
auto step = resolveIndexValue(forOp.getStep(), knowledge);
if (failed(lower) || failed(upper) || failed(step)
|| (mode == CoreWalkMode::ExecuteAllIterations && *step <= 0)) {
|| (mode == CoreWalkMode::ExecuteCommunication && *step <= 0)) {
forOp.emitOpError() << "requires statically evaluable scf.for bounds for PIM " << purpose;
hasFailure = true;
continue;
@@ -84,13 +102,13 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
loopKnowledge.indexValues[forOp.getInductionVar()] = induction;
for (auto [index, iterArg] : llvm::enumerate(forOp.getRegionIterArgs()))
loopKnowledge.aliases[iterArg] = carryValues ? iterValues[index] : forOp.getInitArgs()[index];
hasFailure |= failed(walkPimCoreBlockImpl(body, loopKnowledge, mode, callback));
hasFailure |= failed(walkPimCoreBlockImpl(body, loopKnowledge, mode, communicationPlan, callback));
auto yield = mlir::cast<mlir::scf::YieldOp>(body.getTerminator());
for (auto [index, yielded] : llvm::enumerate(yield.getOperands()))
iterValues[index] = resolveLoopCarriedAlias(yielded, loopKnowledge);
};
if (mode == CoreWalkMode::ExecuteAllIterations) {
if (mode == CoreWalkMode::ExecuteCommunication) {
for (int64_t induction = *lower; induction < *upper; induction += *step)
visitIteration(induction, true);
}
@@ -113,12 +131,14 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
continue;
}
mlir::Region& selected = *condition != 0 ? ifOp.getThenRegion() : ifOp.getElseRegion();
if (mode == CoreWalkMode::ExecuteAllIterations) {
hasFailure |= !selected.empty() && failed(walkPimCoreBlockImpl(selected.front(), knowledge, mode, callback));
if (mode == CoreWalkMode::ExecuteCommunication) {
hasFailure |= !selected.empty()
&& failed(walkPimCoreBlockImpl(selected.front(), knowledge, mode, communicationPlan, callback));
}
else {
for (mlir::Region* region : {&ifOp.getThenRegion(), &ifOp.getElseRegion()})
hasFailure |= !region->empty() && failed(walkPimCoreBlockImpl(region->front(), knowledge, mode, callback));
hasFailure |= !region->empty()
&& failed(walkPimCoreBlockImpl(region->front(), knowledge, mode, communicationPlan, callback));
}
propagateRegionResults(ifOp.getResults(), selected, knowledge);
continue;
@@ -138,33 +158,54 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
break;
}
for (mlir::Region& region : switchOp->getRegions()) {
if (mode == CoreWalkMode::ExecuteAllIterations && &region != selected)
if (mode == CoreWalkMode::ExecuteCommunication && &region != selected)
continue;
hasFailure |= failed(walkPimCoreBlockImpl(region.front(), knowledge, mode, callback));
hasFailure |= failed(walkPimCoreBlockImpl(region.front(), knowledge, mode, communicationPlan, callback));
}
propagateRegionResults(switchOp.getResults(), *selected, knowledge);
continue;
}
hasFailure |= failed(callback(op, knowledge));
if (mode != CoreWalkMode::ExecuteCommunication || mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op))
hasFailure |= failed(callback(op, knowledge));
}
return mlir::success(!hasFailure);
}
} // namespace
mlir::LogicalResult
walkPimCoreBlock(mlir::Block& block,
const StaticValueKnowledge& knowledge,
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::ExecuteAllIterations, callback);
PimCoreCommunicationPlan buildPimCoreCommunicationPlan(mlir::Block& block) {
llvm::DenseSet<mlir::Operation*> communicationAncestors;
block.walk([&](mlir::Operation* op) {
if (!mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op))
return;
for (mlir::Operation* parent = op->getParentOp(); parent; parent = parent->getParentOp())
communicationAncestors.insert(parent);
});
PimCoreCommunicationPlan plan;
block.walk([&](mlir::Operation* op) {
bool isCommunication = mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op);
bool isControlFlow = mlir::isa<mlir::scf::ForOp, mlir::scf::IfOp, mlir::scf::IndexSwitchOp>(op);
if (isCommunication || (isControlFlow && (op->getNumResults() != 0 || communicationAncestors.contains(op))))
plan[op->getBlock()].push_back(op);
});
return plan;
}
mlir::LogicalResult walkPimCoreCommunicationBlock(
mlir::Block& block,
const PimCoreCommunicationPlan& plan,
const StaticValueKnowledge& knowledge,
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::ExecuteCommunication, &plan, callback);
}
mlir::LogicalResult walkPimCoreBlockStructurally(
mlir::Block& block,
const StaticValueKnowledge& knowledge,
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::StructuralExtremes, callback);
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::StructuralExtremes, nullptr, callback);
}
} // namespace onnx_mlir
+14 -7
View File
@@ -3,23 +3,30 @@
#include "mlir/IR/Block.h"
#include "mlir/Support/LogicalResult.h"
#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/STLFunctionalExtras.h"
#include "llvm/ADT/SmallVector.h"
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
namespace onnx_mlir {
using PimCoreCommunicationPlan = llvm::DenseMap<mlir::Block*, llvm::SmallVector<mlir::Operation*, 8>>;
/// Returns true for ops in a `pim.core` body that only participate in static
/// address or index computation and therefore do not emit PIM instructions.
bool isCoreStaticAddressOp(mlir::Operation* op);
/// Walks a `pim.core` body, statically unrolling nested `scf.for` loops when
/// their bounds are known and invoking `callback` only on instruction-emitting
/// operations.
mlir::LogicalResult
walkPimCoreBlock(mlir::Block& block,
const StaticValueKnowledge& knowledge,
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback);
/// Walks a `pim.core` body's communication stream, statically unrolling
/// control flow that contains send/receive operations and invoking `callback`
/// on those operations in execution order.
PimCoreCommunicationPlan buildPimCoreCommunicationPlan(mlir::Block& block);
mlir::LogicalResult walkPimCoreCommunicationBlock(
mlir::Block& block,
const PimCoreCommunicationPlan& plan,
const StaticValueKnowledge& knowledge,
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback);
/// Walks a `pim.core`-like body structurally for verification without
/// enumerating full loop trip counts. Loop bounds must still be statically
+10
View File
@@ -22,9 +22,19 @@
#include "src/Accelerators/PIM/Common/Support/FileSystemUtils.hpp"
#include "src/Compiler/CompilerOptions.hpp"
#include <cstdint>
#include <limits>
namespace onnx_mlir {
inline constexpr llvm::StringLiteral kCoreIdAttrName = "coreId";
inline constexpr llvm::StringLiteral kCoreIdsAttrName = "coreIds";
inline constexpr llvm::StringLiteral kLocalMemoryAddressAttrName = "pim.local_memory_address";
inline constexpr llvm::StringLiteral kLocalMemorySlotAttrName = "pim.local_memory_slot";
inline constexpr llvm::StringLiteral kLocalMemorySlotSizeAttrName = "pim.local_memory_slot_size";
inline constexpr llvm::StringLiteral kLocalMemoryFallbackCountAttrName = "pim.local_memory_fallback_count";
inline constexpr llvm::StringLiteral kLocalMemoryNestedSingleUseCountAttrName =
"pim.local_memory_nested_single_use_count";
inline constexpr size_t kPimLocalMemoryAddressLimit = static_cast<size_t>(std::numeric_limits<int32_t>::max());
} // namespace onnx_mlir