Raptor sync wait

This commit is contained in:
ilgeco
2026-08-06 14:32:46 +02:00
parent a963009855
commit a39fdba366
48 changed files with 3357 additions and 96 deletions
+2 -2
View File
@@ -162,8 +162,8 @@ inline constexpr std::array<InstructionJsonFormat, kOpcodeCount> kInstructionJso
{true, true, true, "", "", "", "len" }, // lmv
{true, false, true, "core", "", "", "size"}, // send
{true, false, true, "core", "", "", "size"}, // recv
{false, false, false, "", "", "", "" }, // wait
{false, false, false, "", "", "", "" }, // sync
{false, false, false, "", "event_register", "wait_value", ""}, // wait
{false, false, false, "core", "event_register", "", ""}, // sync
}};
static_assert(kInstructionJsonFormats.size() == kOpcodeCount);
+30
View File
@@ -692,6 +692,34 @@ void PimCodeGen::codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge
pim_binary::Opcode::send, addressOf(sendOp.getInput(), knowledge), *targetCoreId, sendOp.getSize());
}
void PimCodeGen::codeGenWaitOp(
pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const {
auto eventRegister = indexOf(waitOp.getEventRegister(), knowledge);
assert(succeeded(eventRegister)
&& "pim.wait event register must be statically resolvable during codegen");
pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::wait;
instruction.generic1 = pim::checkedI32OrCrash(
*eventRegister, "wait event register");
instruction.generic2 = waitOp.getWaitValue();
emitInstruction(instruction);
}
void PimCodeGen::codeGenSyncOp(
pim::PimSyncOp syncOp, const StaticValueKnowledge& knowledge) const {
auto targetCoreId = indexOf(syncOp.getTargetCoreId(), knowledge);
auto eventRegister = indexOf(syncOp.getEventRegister(), knowledge);
assert(succeeded(targetCoreId) && succeeded(eventRegister)
&& "pim.sync operands must be statically resolvable during codegen");
pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::sync;
instruction.r2OrImm = pim::checkedI32OrCrash(
*targetCoreId, "sync target core id");
instruction.generic1 = pim::checkedI32OrCrash(
*eventRegister, "sync event register");
emitInstruction(instruction);
}
void PimCodeGen::codeGenConcatOp(pim::PimConcatOp concatOp, const StaticValueKnowledge& knowledge) const {
auto outputType = cast<ShapedType>(concatOp.getOutputBuffer().getType());
assert(outputType.hasStaticShape() && "concat codegen requires static output shape");
@@ -991,6 +1019,8 @@ static LogicalResult executeCompiledCorePlan(
case CompiledCoreOpKind::VMV: coreCodeGen.codeGenVMVOp(cast<pim::PimVMVOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Receive: coreCodeGen.codeGenReceiveOp(cast<pim::PimReceiveOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Send: coreCodeGen.codeGenSendOp(cast<pim::PimSendOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Wait: coreCodeGen.codeGenWaitOp(cast<pim::PimWaitOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Sync: coreCodeGen.codeGenSyncOp(cast<pim::PimSyncOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Concat: coreCodeGen.codeGenConcatOp(cast<pim::PimConcatOp>(node.op), knowledge); break;
case CompiledCoreOpKind::Vmm:
if (auto weightSlot = resolveWeightSlot(cast<pim::PimVMMOp>(node.op), knowledge); succeeded(weightSlot))
+2
View File
@@ -217,6 +217,8 @@ public:
void codeGenReceiveOp(pim::PimReceiveOp receiveOp, const StaticValueKnowledge& knowledge) const;
void codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge& knowledge) const;
void codeGenWaitOp(pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const;
void codeGenSyncOp(pim::PimSyncOp syncOp, const StaticValueKnowledge& knowledge) const;
void codeGenConcatOp(pim::PimConcatOp concatOp, const StaticValueKnowledge& knowledge) const;
template <typename MVMTy>
+18
View File
@@ -2,6 +2,8 @@
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include <limits>
#define DEBUG_TYPE "PimCompilerOptions"
namespace onnx_mlir {
@@ -110,6 +112,12 @@ llvm::cl::opt<size_t>
llvm::cl::opt<size_t>
crossbarCountInCore("crossbar-count", llvm::cl::desc("Number of crossbars in each core"), llvm::cl::init(64));
llvm::cl::opt<size_t> pipelineStages(
"pipeline",
llvm::cl::desc("Number of throughput pipeline stages (1 preserves latency scheduling)"),
llvm::cl::init(1),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<long> coresCount("core-count",
llvm::cl::desc("Number of cores in the chip. Required for PIM compilation."),
llvm::cl::init(-1));
@@ -129,4 +137,14 @@ void verifyExplicitPimCoreCount() {
llvm::report_fatal_error("PIM compilation requires --core-count to be a positive integer");
}
void verifyPimPipelineStages() {
if (pipelineStages.getValue() == 0)
llvm::report_fatal_error("PIM compilation requires --pipeline to be positive");
if (static_cast<size_t>(coresCount.getValue()) % pipelineStages.getValue() != 0)
llvm::report_fatal_error("PIM compilation requires --core-count to be divisible by --pipeline");
if (crossbarCountInCore.getValue()
> std::numeric_limits<size_t>::max() / pipelineStages.getValue())
llvm::report_fatal_error("PIM compilation --crossbar-count * --pipeline overflows");
}
} // namespace onnx_mlir
+2
View File
@@ -62,6 +62,7 @@ extern llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom;
extern llvm::cl::opt<size_t> crossbarSize;
extern llvm::cl::opt<size_t> crossbarCountInCore;
extern llvm::cl::opt<size_t> pipelineStages;
extern llvm::cl::opt<long> coresCount;
extern llvm::cl::opt<std::string> pimTargetConfig;
extern llvm::cl::opt<uint64_t> pimConvIm2colMaxElements;
@@ -69,5 +70,6 @@ extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions;
bool hasExplicitPimCoreCount();
void verifyExplicitPimCoreCount();
void verifyPimPipelineStages();
} // namespace onnx_mlir
+2 -1
View File
@@ -330,6 +330,7 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
EmissionTargetType& emissionTarget,
std::string outputNameNoExt) {
verifyExplicitPimCoreCount();
verifyPimPipelineStages();
spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget();
spatial::SpatialTargetResources targetResources = getPimSpatialTargetResources(schedulingTarget);
@@ -354,7 +355,7 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
pm.addPass(createTrivialGraphComputeMergePass(
schedulingTarget.residentWeightCapacity, exportStage));
pm.addPass(spatial::createScheduleAndRealizeSpatialPass(
schedulingTarget, exportStage));
schedulingTarget, exportStage, pipelineStages.getValue()));
pm.addPass(createMessagePass("Onnx lowered to Spatial"));
}
+2
View File
@@ -17,6 +17,8 @@ static FailureOr<CompiledCoreOpKind> classifyCompiledCoreOpKind(Operation& op) {
if (isa<pim::PimVMVOp>(op)) return CompiledCoreOpKind::VMV;
if (isa<pim::PimReceiveOp>(op)) return CompiledCoreOpKind::Receive;
if (isa<pim::PimSendOp>(op)) return CompiledCoreOpKind::Send;
if (isa<pim::PimWaitOp>(op)) return CompiledCoreOpKind::Wait;
if (isa<pim::PimSyncOp>(op)) return CompiledCoreOpKind::Sync;
if (isa<pim::PimConcatOp>(op)) return CompiledCoreOpKind::Concat;
if (isa<pim::PimVMMOp>(op)) return CompiledCoreOpKind::Vmm;
if (isa<pim::PimVVAddOp>(op)) return CompiledCoreOpKind::VVAdd;
+2
View File
@@ -17,6 +17,8 @@ enum class CompiledCoreOpKind : uint8_t {
VMV,
Receive,
Send,
Wait,
Sync,
Concat,
Vmm,
VVAdd,