This commit is contained in:
@@ -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 && ®ion != selected)
|
||||
if (mode == CoreWalkMode::ExecuteCommunication && ®ion != 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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user