finally fast googlenet with correct latency artifacts for fair comparison
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-07-29 18:20:44 +02:00
parent 060a21172e
commit 1b4f070bef
74 changed files with 2773 additions and 1311 deletions
+94 -109
View File
@@ -92,11 +92,12 @@ static MemoryReportKind classifyMemoryReportKind(mlir::Value value) {
return MemoryReportKind::None;
}
static int32_t getVectorByteSizeOrCrash(ShapedType type) {
auto byteSize = pim::getCheckedShapedTypeSizeInBytes(type, UnknownLoc::get(type.getContext()), "vector byte size");
if (failed(byteSize))
llvm_unreachable("Failed to compute checked vector byte size");
return pim::checkedI32OrCrash(*byteSize, "vector byte size");
static int32_t getVectorElementCountOrCrash(ShapedType type) {
return pim::checkedI32OrCrash(type.getNumElements(), "vector element count");
}
static int32_t getVectorElementBitwidthOrCrash(ShapedType type) {
return pim::checkedI32OrCrash(static_cast<uint64_t>(type.getElementTypeBitWidth()), "vector element bitwidth");
}
static Operation* getDiagnosticAnchor(mlir::Value value) {
@@ -165,8 +166,7 @@ size_t PimMemory::allocateAddress(size_t size, const MemoryValueKey& key) {
checkedAlignedEnd = checkedAlignTo(*checkedEnd, minAlignment, anchor, "local memory alignment");
if (address > kPimLocalMemoryAddressLimit || failed(checkedEnd) || *checkedEnd > kPimLocalMemoryAddressLimit
|| failed(checkedAlignedEnd) || *checkedAlignedEnd > kPimLocalMemoryAddressLimit) {
printMemoryOverflowDiagnostic(
key,
printMemoryOverflowDiagnostic(key,
size,
firstAvailableAddress,
succeeded(checkedAlignedEnd) ? *checkedAlignedEnd : kPimLocalMemoryAddressLimit);
@@ -208,9 +208,7 @@ void PimMemory::allocateMemoryForValue(const MemoryValueKey& key, MemEntry& memE
switch (reportKind) {
case MemoryReportKind::Alloca:
case MemoryReportKind::Global:
case MemoryReportKind::Input:
++reportRow.hostObjectCount;
break;
case MemoryReportKind::Input: ++reportRow.hostObjectCount; break;
case MemoryReportKind::None: break;
}
}
@@ -259,8 +257,7 @@ void PimMemory::allocateCore(const CompiledCoreMemoryPlan& plan, std::optional<u
reportRow.logicalLocalAllocationCount = plan.logicalAllocationCount;
reportRow.logicalLocalBytes = plan.logicalBytes;
}
else if (*localArenaSize != plan.arenaSize
|| reportRow.logicalLocalAllocationCount != plan.logicalAllocationCount
else if (*localArenaSize != plan.arenaSize || reportRow.logicalLocalAllocationCount != plan.logicalAllocationCount
|| reportRow.logicalLocalBytes != plan.logicalBytes)
llvm_unreachable("inconsistent PIM local-memory plan across core-batch lanes");
for (const CompiledLocalMemoryEntry& entry : plan.entries) {
@@ -358,8 +355,8 @@ llvm::FailureOr<int64_t> PimAcceleratorMemory::getIndexValue(mlir::Value value,
PimAcceleratorMemory::PimAcceleratorMemory()
: hostMem(memEntriesMap), fileReport(openMemoryReport(pimMemoryReport == PimMemoryReportSummary)) {}
PimAcceleratorMemory::PimAcceleratorMemory(
const llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& initialMemEntries, bool enableReport)
PimAcceleratorMemory::PimAcceleratorMemory(const llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& initialMemEntries,
bool enableReport)
: memEntriesMap(initialMemEntries),
hostMem(memEntriesMap),
fileReport(enableReport ? openMemoryReport(true) : std::fstream()) {}
@@ -367,10 +364,8 @@ PimAcceleratorMemory::PimAcceleratorMemory(
void PimAcceleratorMemory::reportHost() { hostReportRow = hostMem.getReportRow(); }
void PimAcceleratorMemory::recordCoreReport(size_t coreId, const MemoryReportRow& row) {
reportEntries.push_back({MemoryReportEntry::Kind::Core,
coreId,
{pim::checkedI32OrCrash(coreId, "memory report core id")},
row});
reportEntries.push_back(
{MemoryReportEntry::Kind::Core, coreId, {pim::checkedI32OrCrash(coreId, "memory report core id")}, row});
}
void PimAcceleratorMemory::recordBatchReport(uint64_t batchId,
@@ -441,8 +436,7 @@ void PimAcceleratorMemory::flushReport() {
os << " Weights memory: " << formatReportMemory(totalWeightBytes) << "\n";
os << " Local memory before reuse: " << formatReportMemory(logicalBytes) << "\n";
os << " Local memory after reuse: " << formatReportMemory(physicalBytes) << "\n";
os << " Saved local memory: " << formatReportMemory(savedBytes) << " ("
<< formatv("{0:F1}%", savedPercent) << ")\n";
os << " Saved local memory: " << formatReportMemory(savedBytes) << " (" << formatv("{0:F1}%", savedPercent) << ")\n";
os << " Largest core local memory: " << formatReportMemory(largest) << "\n";
if (!groups.empty()) {
os << " ";
@@ -456,20 +450,15 @@ void PimAcceleratorMemory::flushReport() {
printLabel(group);
os << "\n";
uint64_t groupSaved = group.row.logicalLocalBytes - group.row.physicalLocalBytes;
double groupPercent = group.row.logicalLocalBytes == 0
? 0.0
: 100.0 * groupSaved / group.row.logicalLocalBytes;
double groupPercent = group.row.logicalLocalBytes == 0 ? 0.0 : 100.0 * groupSaved / group.row.logicalLocalBytes;
if (group.coreIds.size() == 1) {
os << " Local memory: " << formatReportMemory(group.row.logicalLocalBytes) << " → "
<< formatReportMemory(group.row.physicalLocalBytes) << " (" << formatv("{0:F1}% saved", groupPercent)
<< ")\n";
<< formatReportMemory(group.row.physicalLocalBytes) << " (" << formatv("{0:F1}% saved", groupPercent) << ")\n";
}
else {
os << " Per core: " << formatReportMemory(group.row.logicalLocalBytes) << " → "
<< formatReportMemory(group.row.physicalLocalBytes) << " (" << formatv("{0:F1}% saved", groupPercent)
<< ")\n";
os << " Total after reuse: " << formatReportMemory(group.row.physicalLocalBytes * group.coreIds.size())
<< "\n";
<< formatReportMemory(group.row.physicalLocalBytes) << " (" << formatv("{0:F1}% saved", groupPercent) << ")\n";
os << " Total after reuse: " << formatReportMemory(group.row.physicalLocalBytes * group.coreIds.size()) << "\n";
}
}
if (groups.size() > kGroupLimit)
@@ -479,12 +468,6 @@ void PimAcceleratorMemory::flushReport() {
fileReport.close();
}
size_t PimCodeGen::remapCoreId(size_t coreId) const {
auto it = emittedCoreIds.find(coreId);
assert(it != emittedCoreIds.end() && "Missing emitted core id remapping");
return it->second;
}
void PimCodeGen::emitInstruction(const pim_binary::InstructionRecord& instruction) const {
if (failed(instructionWriter.append(instruction)))
return;
@@ -493,6 +476,19 @@ void PimCodeGen::emitInstruction(const pim_binary::InstructionRecord& instructio
updateScalarRegisterCache(instruction);
}
void PimCodeGen::ensureVectorBitwidth(int32_t inputBitwidth, int32_t outputBitwidth) const {
std::array<int32_t, 2> requested = {inputBitwidth, outputBitwidth};
if (vectorBitwidths == requested)
return;
pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::setbw;
instruction.generic1 = inputBitwidth;
instruction.generic2 = outputBitwidth;
emitInstruction(instruction);
vectorBitwidths = requested;
}
void PimCodeGen::updateScalarRegisterCache(const pim_binary::InstructionRecord& instruction) const {
switch (instruction.opcode) {
case pim_binary::Opcode::sldi: scalarRegisterValues[instruction.rd] = instruction.r2OrImm; break;
@@ -563,7 +559,7 @@ void PimCodeGen::emitCommunicationOp(pim_binary::Opcode opcode, size_t bufferAdd
pim_binary::InstructionRecord instruction;
instruction.opcode = opcode;
instruction.rd = 0;
instruction.r2OrImm = pim::checkedI32OrCrash(remapCoreId(coreId), "communication core id");
instruction.r2OrImm = pim::checkedI32OrCrash(coreId, "physical communication core id");
instruction.generic1 = 0;
instruction.generic2 = 0;
instruction.generic3 = pim::checkedI32OrCrash(size, "communication byte size");
@@ -679,6 +675,8 @@ void PimCodeGen::codeGenMVMLikeOp(size_t mvmId,
MVMTy mvmLikeOp,
bool transposeMatrix,
const StaticValueKnowledge& knowledge) {
ensureVectorBitwidth(getVectorElementBitwidthOrCrash(cast<ShapedType>(mvmLikeOp.getInput().getType())),
getVectorElementBitwidthOrCrash(cast<ShapedType>(mvmLikeOp.getOutputBuffer().getType())));
emitMvmOp(mvmId, addressOf(mvmLikeOp.getOutputBuffer(), knowledge), 0, addressOf(mvmLikeOp.getInput(), knowledge), 0);
// TODO: save weights somewhere (if transposeMatrix=true, transpose the weight matrix)
@@ -688,25 +686,29 @@ void PimCodeGen::emitBinaryVectorOp(pim_binary::Opcode opcode,
mlir::Value output,
mlir::Value lhs,
mlir::Value rhs,
size_t byteSize,
const StaticValueKnowledge& knowledge) const {
auto inputType = cast<ShapedType>(lhs.getType());
ensureVectorBitwidth(getVectorElementBitwidthOrCrash(inputType),
getVectorElementBitwidthOrCrash(cast<ShapedType>(output.getType())));
setupRdRs1Rs2(addressOf(output, knowledge), 0, addressOf(lhs, knowledge), 0, addressOf(rhs, knowledge), 0);
pim_binary::InstructionRecord instruction;
instruction.opcode = opcode;
instruction.rd = 0;
instruction.r1 = 1;
instruction.r2OrImm = 2;
instruction.generic3 = pim::checkedI32OrCrash(byteSize, "vector byte size");
instruction.generic3 = getVectorElementCountOrCrash(inputType);
emitInstruction(instruction);
}
void PimCodeGen::emitUnaryVectorOp(pim_binary::Opcode opcode,
mlir::Value output,
mlir::Value input,
size_t byteSize,
const StaticValueKnowledge& knowledge,
int32_t r2OrImm,
int32_t generic1) const {
auto inputType = cast<ShapedType>(input.getType());
ensureVectorBitwidth(getVectorElementBitwidthOrCrash(inputType),
getVectorElementBitwidthOrCrash(cast<ShapedType>(output.getType())));
setupRdRs1(addressOf(output, knowledge), 0, addressOf(input, knowledge), 0);
pim_binary::InstructionRecord instruction;
instruction.opcode = opcode;
@@ -714,7 +716,7 @@ void PimCodeGen::emitUnaryVectorOp(pim_binary::Opcode opcode,
instruction.r1 = 1;
instruction.r2OrImm = r2OrImm;
instruction.generic1 = generic1;
instruction.generic3 = pim::checkedI32OrCrash(byteSize, "vector byte size");
instruction.generic3 = getVectorElementCountOrCrash(inputType);
emitInstruction(instruction);
}
@@ -812,8 +814,7 @@ static SmallVector<Operation*> collectTopLevelCoreLikeOps(func::FuncOp funcOp) {
static FailureOr<CompiledCoreMemoryPlan> compileCoreMemoryPlan(Operation* coreLikeOp) {
CompiledCoreMemoryPlan plan;
auto arenaAttr = coreLikeOp->getAttrOfType<IntegerAttr>(kLocalMemorySizeAttrName);
if (!arenaAttr || arenaAttr.getInt() < 0
|| static_cast<uint64_t>(arenaAttr.getInt()) > kPimLocalMemoryAddressLimit) {
if (!arenaAttr || arenaAttr.getInt() < 0 || static_cast<uint64_t>(arenaAttr.getInt()) > kPimLocalMemoryAddressLimit) {
coreLikeOp->emitError("requires a valid pim.local_memory_size attribute before codegen");
return failure();
}
@@ -841,7 +842,9 @@ static FailureOr<CompiledCoreMemoryPlan> compileCoreMemoryPlan(Operation* coreLi
hasFailure = true;
return;
}
plan.entries.push_back({allocOp.getResult(), {address, static_cast<size_t>(*checkedSize)}});
plan.entries.push_back({
allocOp.getResult(), {address, static_cast<size_t>(*checkedSize)}
});
auto logicalBytes = pim::checkedAdd(
static_cast<size_t>(plan.logicalBytes), static_cast<size_t>(*checkedSize), allocOp, "logical local bytes");
if (failed(logicalBytes)) {
@@ -995,21 +998,10 @@ static LogicalResult executeCompiledCorePlan(
}
auto emitBinary = [&](auto op, pim_binary::Opcode opcode) {
coreCodeGen.emitBinaryVectorOp(opcode,
op.getOutputBuffer(),
op.getLhs(),
op.getRhs(),
getVectorByteSizeOrCrash(cast<ShapedType>(op.getLhs().getType())),
knowledge);
coreCodeGen.emitBinaryVectorOp(opcode, op.getOutputBuffer(), op.getLhs(), op.getRhs(), knowledge);
};
auto emitUnary = [&](auto op, pim_binary::Opcode opcode, int32_t r2OrImm, int32_t generic1) {
coreCodeGen.emitUnaryVectorOp(opcode,
op.getOutputBuffer(),
op.getInput(),
getVectorByteSizeOrCrash(cast<ShapedType>(op.getInput().getType())),
knowledge,
r2OrImm,
generic1);
coreCodeGen.emitUnaryVectorOp(opcode, op.getOutputBuffer(), op.getInput(), knowledge, r2OrImm, generic1);
};
switch (node.opKind) {
@@ -1092,8 +1084,8 @@ static void aliasMaterializedHostGlobals(CoreLikeOpTy coreLikeOp,
});
}
static OnnxMlirCompilerErrorCodes emitEmptyCoreArtifacts(StringRef outputDirPath, size_t emittedCoreId) {
std::string outputCorePath = (outputDirPath + "/core_" + std::to_string(emittedCoreId) + ".pim").str();
static OnnxMlirCompilerErrorCodes emitEmptyCoreArtifacts(StringRef outputDirPath, size_t physicalCoreId) {
std::string outputCorePath = (outputDirPath + "/core_" + std::to_string(physicalCoreId) + ".pim").str();
std::error_code errorCode;
raw_fd_ostream coreBinaryStream(outputCorePath, errorCode, sys::fs::OF_None);
if (errorCode) {
@@ -1117,7 +1109,7 @@ static OnnxMlirCompilerErrorCodes emitEmptyCoreArtifacts(StringRef outputDirPath
if (!pimEmitJson.getValue())
return CompilerSuccess;
std::string outputCoreJsonPath = (outputDirPath + "/core_" + std::to_string(emittedCoreId) + ".json").str();
std::string outputCoreJsonPath = (outputDirPath + "/core_" + std::to_string(physicalCoreId) + ".json").str();
errorCode = std::error_code();
raw_fd_ostream coreJsonStream(outputCoreJsonPath, errorCode);
if (errorCode) {
@@ -1179,36 +1171,15 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
return it->second.get();
};
llvm::DenseMap<size_t, size_t> emittedCoreIds;
size_t nextEmittedCoreId = 0;
for (Operation* op : coreLikeOps) {
if (auto coreOp = dyn_cast<pim::PimCoreOp>(op)) {
size_t originalCoreId = static_cast<size_t>(coreOp.getCoreId());
if (!emittedCoreIds.contains(originalCoreId))
emittedCoreIds[originalCoreId] = nextEmittedCoreId++;
continue;
}
auto coreBatchOp = cast<pim::PimCoreBatchOp>(op);
auto batchCoreIds = getBatchCoreIds(coreBatchOp);
for (unsigned lane = 0; lane < static_cast<unsigned>(coreBatchOp.getLaneCount()); ++lane) {
size_t originalCoreId = static_cast<size_t>(batchCoreIds[lane]);
if (!emittedCoreIds.contains(originalCoreId))
emittedCoreIds[originalCoreId] = nextEmittedCoreId++;
}
}
SmallVector<CoreEmissionJob> jobs;
SmallVector<SmallVector<size_t>> batchJobIndices;
for (Operation* op : coreLikeOps) {
if (auto coreOp = dyn_cast<pim::PimCoreOp>(op)) {
size_t originalCoreId = static_cast<size_t>(coreOp.getCoreId());
CoreEmissionJob job;
job.coreLikeOp = coreOp;
job.program = getCompiledProgram(op);
job.memoryPlan = getMemoryPlan(op);
job.emittedCoreId = emittedCoreIds.lookup(originalCoreId);
job.physicalCoreId = static_cast<size_t>(coreOp.getCoreId());
jobs.push_back(std::move(job));
continue;
}
@@ -1220,16 +1191,15 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
lanesByCoreId[static_cast<size_t>(batchCoreIds[lane])].push_back(lane);
SmallVector<size_t> jobIndices;
SmallVector<size_t> orderedOriginalCoreIds = llvm::to_vector(lanesByCoreId.keys());
llvm::sort(orderedOriginalCoreIds,
[&](size_t lhs, size_t rhs) { return emittedCoreIds.lookup(lhs) < emittedCoreIds.lookup(rhs); });
for (size_t originalCoreId : orderedOriginalCoreIds) {
SmallVector<size_t> physicalCoreIds = llvm::to_vector(lanesByCoreId.keys());
llvm::sort(physicalCoreIds);
for (size_t physicalCoreId : physicalCoreIds) {
CoreEmissionJob job;
job.coreLikeOp = coreBatchOp;
job.program = getCompiledProgram(op);
job.memoryPlan = getMemoryPlan(op);
job.emittedCoreId = emittedCoreIds.lookup(originalCoreId);
job.lanes = lanesByCoreId.lookup(originalCoreId);
job.physicalCoreId = physicalCoreId;
job.lanes = lanesByCoreId.lookup(physicalCoreId);
job.batchReportId = nextBatchReportId;
jobIndices.push_back(jobs.size());
jobs.push_back(std::move(job));
@@ -1238,8 +1208,11 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
++nextBatchReportId;
}
auto linkCoreWeights =
[&](size_t coreId, ArrayRef<std::string> weightFiles, json::Array& xbarsPerGroup) -> OnnxMlirCompilerErrorCodes {
auto linkCoreWeights = [&](size_t coreId,
ArrayRef<std::string> weightFiles,
ArrayRef<ResolvedWeightView> weights,
json::Array& xbarsPerGroup) -> OnnxMlirCompilerErrorCodes {
assert(weightFiles.size() == weights.size() && "weight files must match resolved weight views");
auto coreWeightsDirPath = outputDirPath + "/core_" + std::to_string(coreId);
if (auto error = sys::fs::create_directory(coreWeightsDirPath); error && error != std::errc::file_exists) {
errs() << "Error creating core directory: " << coreWeightsDirPath << ": " << error.message() << '\n';
@@ -1247,7 +1220,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
}
for (auto [slot, fileName] : llvm::enumerate(weightFiles)) {
xbarsPerGroup.push_back(1);
xbarsPerGroup.push_back(weights[slot].shape[1] / static_cast<int64_t>(crossbarSize));
std::string sourcePath = outputDirPath + "/weights/" + fileName;
std::string targetPath = coreWeightsDirPath + "/crossbar_" + std::to_string(slot) + ".bin";
sys::fs::remove(targetPath);
@@ -1295,7 +1268,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
};
std::error_code errorCode;
auto outputCorePath = outputDirPath + "/core_" + std::to_string(job.emittedCoreId) + ".pim";
auto outputCorePath = outputDirPath + "/core_" + std::to_string(job.physicalCoreId) + ".pim";
raw_fd_ostream coreBinaryStream(outputCorePath, errorCode, sys::fs::OF_None);
if (errorCode) {
errs() << "Error while opening core file `" << outputCorePath << "`: " << errorCode.message() << '\n';
@@ -1305,7 +1278,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
std::unique_ptr<raw_fd_ostream> coreJsonStream;
if (pimEmitJson.getValue()) {
std::string outputCoreJsonPath = outputDirPath + "/core_" + std::to_string(job.emittedCoreId) + ".json";
std::string outputCoreJsonPath = outputDirPath + "/core_" + std::to_string(job.physicalCoreId) + ".json";
errorCode = std::error_code();
coreJsonStream = std::make_unique<raw_fd_ostream>(outputCoreJsonPath, errorCode);
if (errorCode) {
@@ -1317,7 +1290,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
}
PimInstructionWriter instructionWriter(coreBinaryStream);
PimCodeGen coreCodeGen(jobMemory, instructionWriter, coreJsonStream.get(), emittedCoreIds);
PimCodeGen coreCodeGen(jobMemory, instructionWriter, coreJsonStream.get());
auto finalizeInstructions = [&]() {
bool succeeded = mlir::succeeded(instructionWriter.finalize());
coreBinaryStream.close();
@@ -1330,7 +1303,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
if (auto coreOp = dyn_cast<pim::PimCoreOp>(job.coreLikeOp)) {
aliasMaterializedHostGlobals(coreOp, moduleOp, materializedHostGlobals, jobMemory);
auto& deviceMemory = jobMemory.getOrCreateDeviceMem(job.emittedCoreId);
auto& deviceMemory = jobMemory.getOrCreateDeviceMem(job.physicalCoreId);
deviceMemory.allocateCore(*job.memoryPlan);
StaticValueKnowledge knowledge = seedCoreCodegenKnowledge(coreOp);
@@ -1345,7 +1318,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
else {
auto coreBatchOp = cast<pim::PimCoreBatchOp>(job.coreLikeOp);
aliasMaterializedHostGlobals(coreBatchOp, moduleOp, materializedHostGlobals, jobMemory);
auto& deviceMemory = jobMemory.getOrCreateDeviceMem(job.emittedCoreId);
auto& deviceMemory = jobMemory.getOrCreateDeviceMem(job.physicalCoreId);
for (unsigned lane : job.lanes) {
StaticValueKnowledge knowledge = seedCoreBatchCodegenKnowledge(coreBatchOp, lane);
@@ -1398,18 +1371,29 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
if (jobResults[jobIndex].status != CompilerSuccess)
return jobResults[jobIndex].status;
if (jobs.empty()) {
if (auto err = emitEmptyCoreArtifacts(outputDirPath, 0))
size_t maxPhysicalCoreId = 0;
for (const CoreEmissionJob& job : jobs)
maxPhysicalCoreId = std::max(maxPhysicalCoreId, job.physicalCoreId);
std::vector<bool> activePhysicalCores(maxPhysicalCoreId + 1);
for (const CoreEmissionJob& job : jobs)
activePhysicalCores[job.physicalCoreId] = true;
for (size_t physicalCoreId = 0;
physicalCoreId < activePhysicalCores.size(); ++physicalCoreId) {
if (activePhysicalCores[physicalCoreId])
continue;
if (auto err =
emitEmptyCoreArtifacts(outputDirPath, physicalCoreId))
return err;
xbarsPerArrayGroup["core0"] = json::Array {};
memory.recordCoreReport(0, MemoryReportRow {});
xbarsPerArrayGroup["core" + std::to_string(physicalCoreId)] =
json::Array {};
memory.recordCoreReport(physicalCoreId, MemoryReportRow {});
}
llvm::SmallVector<WeightFileRequest, 8> weightRequests;
weightRequests.reserve(jobs.size());
for (size_t jobIndex = 0; jobIndex < jobs.size(); ++jobIndex) {
WeightFileRequest request;
request.coreId = jobs[jobIndex].emittedCoreId;
request.coreId = jobs[jobIndex].physicalCoreId;
request.weights = jobResults[jobIndex].usedWeights;
weightRequests.push_back(std::move(request));
}
@@ -1422,10 +1406,11 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
json::Array xbarsPerGroup;
if (auto coreOp = dyn_cast<pim::PimCoreOp>(job.coreLikeOp)) {
if (auto err = linkCoreWeights(job.emittedCoreId, mapCoreWeightToFileName[job.emittedCoreId], xbarsPerGroup))
if (auto err = linkCoreWeights(
job.physicalCoreId, mapCoreWeightToFileName[job.physicalCoreId], result.usedWeights, xbarsPerGroup))
return err;
xbarsPerArrayGroup["core" + std::to_string(job.emittedCoreId)] = std::move(xbarsPerGroup);
memory.recordCoreReport(job.emittedCoreId, result.reportRow);
xbarsPerArrayGroup["core" + std::to_string(job.physicalCoreId)] = std::move(xbarsPerGroup);
memory.recordCoreReport(job.physicalCoreId, result.reportRow);
continue;
}
}
@@ -1438,10 +1423,11 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
const CoreEmissionJob& job = jobs[jobIndex];
const CoreEmissionResult& result = jobResults[jobIndex];
json::Array xbarsPerGroup;
if (auto err = linkCoreWeights(job.emittedCoreId, mapCoreWeightToFileName[job.emittedCoreId], xbarsPerGroup))
if (auto err = linkCoreWeights(
job.physicalCoreId, mapCoreWeightToFileName[job.physicalCoreId], result.usedWeights, xbarsPerGroup))
return err;
xbarsPerArrayGroup["core" + std::to_string(job.emittedCoreId)] = std::move(xbarsPerGroup);
reportedCoreIds.push_back(pim::checkedI32OrCrash(job.emittedCoreId, "batch report core id"));
xbarsPerArrayGroup["core" + std::to_string(job.physicalCoreId)] = std::move(xbarsPerGroup);
reportedCoreIds.push_back(pim::checkedI32OrCrash(job.physicalCoreId, "batch report physical core id"));
if (!batchPerCoreRow)
batchPerCoreRow = result.reportRow;
else if (!(*batchPerCoreRow == result.reportRow))
@@ -1449,11 +1435,10 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
}
uint64_t batchReportId = jobs[group.front()].batchReportId.value_or(0);
memory.recordBatchReport(
batchReportId, reportedCoreIds, batchPerCoreRow.value_or(MemoryReportRow {}));
memory.recordBatchReport(batchReportId, reportedCoreIds, batchPerCoreRow.value_or(MemoryReportRow {}));
}
maxCoreId = nextEmittedCoreId == 0 ? 0 : nextEmittedCoreId - 1;
maxCoreId = maxPhysicalCoreId;
memory.flushReport();
return writeConfigJson(funcOp, memory, maxCoreId, std::move(xbarsPerArrayGroup), outputDirPath);
+7 -15
View File
@@ -134,8 +134,7 @@ private:
public:
PimAcceleratorMemory();
PimAcceleratorMemory(
const llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& initialMemEntries, bool enableReport);
PimAcceleratorMemory(const llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& initialMemEntries, bool enableReport);
PimMemory& getOrCreateDeviceMem(size_t id);
@@ -145,8 +144,7 @@ public:
llvm::FailureOr<int64_t> getIndexValue(mlir::Value value, const StaticValueKnowledge& knowledge = {}) const;
void reportHost();
void recordCoreReport(size_t coreId, const MemoryReportRow& row);
void recordBatchReport(
uint64_t batchId, llvm::ArrayRef<int32_t> coreIds, const MemoryReportRow& perCoreRow);
void recordBatchReport(uint64_t batchId, llvm::ArrayRef<int32_t> coreIds, const MemoryReportRow& perCoreRow);
void setTotalWeightBytes(uint64_t bytes) { totalWeightBytes = bytes; }
void flushReport();
};
@@ -155,7 +153,7 @@ struct CoreEmissionJob {
mlir::Operation* coreLikeOp = nullptr;
const CompiledCoreProgram* program = nullptr;
const CompiledCoreMemoryPlan* memoryPlan = nullptr;
size_t emittedCoreId = 0;
size_t physicalCoreId = 0;
llvm::SmallVector<unsigned, 4> lanes;
std::optional<uint64_t> batchReportId;
};
@@ -164,17 +162,16 @@ class PimCodeGen {
PimAcceleratorMemory& memory;
PimInstructionWriter& instructionWriter;
llvm::raw_fd_ostream* coreJsonStream;
const llvm::DenseMap<size_t, size_t>& emittedCoreIds;
std::optional<unsigned> batchLane;
mutable std::array<std::optional<int32_t>, 256> scalarRegisterValues = {};
mutable std::optional<std::array<int32_t, 2>> vectorBitwidths;
size_t addressOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
return memory.getValueAddress(value, knowledge, batchLane);
}
size_t remapCoreId(size_t coreId) const;
void emitInstruction(const pim_binary::InstructionRecord& instruction) const;
void updateScalarRegisterCache(const pim_binary::InstructionRecord& instruction) const;
void ensureVectorBitwidth(int32_t inputBitwidth, int32_t outputBitwidth) const;
void genSetRegisterImmediate(uint8_t registerNumber, int32_t immediate) const;
void genSetRegisterImmediateUnsigned(size_t registerNumber, size_t immediate) const;
@@ -198,21 +195,16 @@ public:
mlir::Value output,
mlir::Value lhs,
mlir::Value rhs,
size_t byteSize,
const StaticValueKnowledge& knowledge) const;
void emitUnaryVectorOp(pim_binary::Opcode opcode,
mlir::Value output,
mlir::Value input,
size_t byteSize,
const StaticValueKnowledge& knowledge,
int32_t r2OrImm = 0,
int32_t generic1 = 0) const;
PimCodeGen(PimAcceleratorMemory& memory,
PimInstructionWriter& instructionWriter,
llvm::raw_fd_ostream* coreJson,
const llvm::DenseMap<size_t, size_t>& emittedCoreIds)
: memory(memory), instructionWriter(instructionWriter), coreJsonStream(coreJson), emittedCoreIds(emittedCoreIds) {}
PimCodeGen(PimAcceleratorMemory& memory, PimInstructionWriter& instructionWriter, llvm::raw_fd_ostream* coreJson)
: memory(memory), instructionWriter(instructionWriter), coreJsonStream(coreJson) {}
void setBatchLane(std::optional<unsigned> lane) { batchLane = lane; }
llvm::FailureOr<int64_t> indexOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
+6
View File
@@ -125,6 +125,12 @@ llvm::cl::opt<long> coresCount("core-count",
llvm::cl::desc("Number of cores in the chip. Required for PIM compilation."),
llvm::cl::init(-1));
llvm::cl::opt<std::string> pimTargetConfig(
"pim-target-config",
llvm::cl::desc("PIM target configuration used to construct the Spatial scheduling cost model"),
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"),
+3
View File
@@ -2,6 +2,8 @@
#include "llvm/Support/CommandLine.h"
#include <string>
#define INSTRUMENTSTAGE_ENUM_PIM
#define INSTRUMENTSTAGE_CL_ENUM_PIM
@@ -63,6 +65,7 @@ extern llvm::cl::opt<bool> pimTraceCommunicationMaterialization;
extern llvm::cl::opt<size_t> crossbarSize;
extern llvm::cl::opt<size_t> crossbarCountInCore;
extern llvm::cl::opt<long> coresCount;
extern llvm::cl::opt<std::string> pimTargetConfig;
extern llvm::cl::opt<uint64_t> pimConvIm2colMaxElements;
extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions;
+264 -2
View File
@@ -1,9 +1,21 @@
#include "mlir/Conversion/AffineToStandard/AffineToStandard.h"
#include "mlir/Transforms/Passes.h"
#include "llvm/Support/Error.h"
#include "llvm/Support/ErrorHandling.h"
#include "llvm/Support/JSON.h"
#include "llvm/Support/MemoryBuffer.h"
#include "llvm/Support/Path.h"
#include "llvm/ADT/SmallString.h"
#include <cmath>
#include <limits>
#include <tuple>
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerUtils.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/Compiler/CompilerPasses.hpp"
@@ -14,6 +26,254 @@ using namespace onnx_mlir;
namespace onnx_mlir {
namespace {
void setDefaultPimInterProcessorLatencies(
spatial::SchedulingTarget& target) {
size_t rows = static_cast<size_t>(
std::sqrt(static_cast<long double>(target.processorCount)));
while (rows > 1 && target.processorCount % rows != 0)
--rows;
size_t columns = (target.processorCount + rows - 1) / rows;
target.interProcessorLatencyNs.assign(
target.processorCount * target.processorCount, 0);
Cost latencySum = 0;
size_t pairCount = 0;
for (size_t source = 0; source < target.processorCount; ++source) {
for (size_t destination = 0;
destination < target.processorCount; ++destination) {
if (source == destination)
continue;
size_t sourceRow = source / columns;
size_t sourceColumn = source % columns;
size_t destinationRow = destination / columns;
size_t destinationColumn = destination % columns;
size_t rowDistance = sourceRow > destinationRow
? sourceRow - destinationRow
: destinationRow - sourceRow;
size_t columnDistance = sourceColumn > destinationColumn
? sourceColumn - destinationColumn
: destinationColumn - sourceColumn;
Cost latency = static_cast<Cost>(2 + rowDistance + columnDistance);
target.interProcessorLatencyNs[
source * target.processorCount + destination] = latency;
latencySum = checkedAdd(latencySum, latency);
++pairCount;
}
}
target.averageInterProcessorLatencyNs =
pairCount == 0
? 0
: (latencySum + static_cast<Cost>(pairCount) - 1)
/ static_cast<Cost>(pairCount);
}
spatial::SchedulingTarget getDefaultPimSchedulingTarget() {
spatial::SchedulingTarget target;
target.processorCount = static_cast<size_t>(coresCount.getValue());
target.residentWeightCapacity = crossbarCountInCore.getValue();
target.matrixRows = crossbarSize.getValue();
target.matrixColumns = crossbarSize.getValue();
setDefaultPimInterProcessorLatencies(target);
return target;
}
const llvm::json::Object& requireObject(const llvm::json::Object& object,
llvm::StringRef key,
llvm::StringRef path) {
const llvm::json::Object* nested = object.getObject(key);
if (!nested)
llvm::report_fatal_error("PIM target config is missing object '" + path + "." + key + "'");
return *nested;
}
Cost getConfigCost(const llvm::json::Object& object,
llvm::StringRef key,
Cost fallback,
bool allowZero = false) {
std::optional<double> number = object.getNumber(key);
if (!number)
return fallback;
if (!std::isfinite(*number) || *number < 0.0 || (!allowZero && *number == 0.0)
|| *number > static_cast<double>(std::numeric_limits<Cost>::max()))
llvm::report_fatal_error("PIM target config field '" + key + "' must be a valid positive number");
return static_cast<Cost>(std::ceil(*number));
}
std::pair<size_t, size_t> getConfigPair(const llvm::json::Object& object,
llvm::StringRef key) {
const llvm::json::Array* values = object.getArray(key);
if (!values || values->size() != 2)
llvm::report_fatal_error("PIM target config field '" + key + "' must contain two integers");
std::optional<int64_t> first = (*values)[0].getAsInteger();
std::optional<int64_t> second = (*values)[1].getAsInteger();
if (!first || !second || *first <= 0 || *second <= 0)
llvm::report_fatal_error("PIM target config field '" + key + "' must contain two positive integers");
return {static_cast<size_t>(*first), static_cast<size_t>(*second)};
}
void loadPimInterProcessorLatencies(
spatial::SchedulingTarget& target,
const llvm::json::Object& network) {
std::optional<llvm::StringRef> filename =
network.getString("net_config_file_path");
if (!filename)
llvm::report_fatal_error(
"PIM target config is missing network latency file path");
llvm::SmallString<256> networkPath(*filename);
if (!llvm::sys::path::is_absolute(networkPath)) {
llvm::SmallString<256> configDirectory(pimTargetConfig.getValue());
llvm::sys::path::remove_filename(configDirectory);
llvm::sys::path::append(configDirectory, networkPath);
networkPath = configDirectory;
}
auto buffer = llvm::MemoryBuffer::getFile(networkPath);
if (!buffer)
llvm::report_fatal_error(
llvm::Twine("failed to read PIM network config '")
+ networkPath + "': " + buffer.getError().message());
auto parsed = llvm::json::parse(buffer.get()->getBuffer());
if (!parsed)
llvm::report_fatal_error(
llvm::Twine("failed to parse PIM network config '")
+ networkPath + "': " + llvm::toString(parsed.takeError()));
const llvm::json::Object* root = parsed->getAsObject();
const llvm::json::Object* latencies =
root ? root->getObject("latency") : nullptr;
if (!latencies)
llvm::report_fatal_error(
"PIM network config is missing its latency matrix");
target.interProcessorLatencyNs.assign(
target.processorCount * target.processorCount, 0);
Cost latencySum = 0;
size_t pairCount = 0;
for (size_t source = 0; source < target.processorCount; ++source) {
std::string sourceKey = std::to_string(source);
const llvm::json::Object* row = latencies->getObject(sourceKey);
if (!row)
llvm::report_fatal_error(
llvm::Twine("PIM network config is missing latency row ")
+ sourceKey);
for (size_t destination = 0;
destination < target.processorCount; ++destination) {
if (source == destination)
continue;
std::string destinationKey = std::to_string(destination);
std::optional<double> latency = row->getNumber(destinationKey);
if (!latency || !std::isfinite(*latency) || *latency <= 0.0)
llvm::report_fatal_error(
llvm::Twine("PIM network config is missing latency ")
+ sourceKey + " -> " + destinationKey);
Cost roundedLatency = static_cast<Cost>(std::ceil(*latency));
target.interProcessorLatencyNs[
source * target.processorCount + destination] = roundedLatency;
latencySum = checkedAdd(latencySum, roundedLatency);
++pairCount;
}
}
target.averageInterProcessorLatencyNs =
pairCount == 0
? 0
: (latencySum + static_cast<Cost>(pairCount) - 1)
/ static_cast<Cost>(pairCount);
}
spatial::SchedulingTarget getPimSchedulingTarget() {
spatial::SchedulingTarget target = getDefaultPimSchedulingTarget();
if (pimTargetConfig.empty())
return target;
auto buffer = llvm::MemoryBuffer::getFile(pimTargetConfig);
if (!buffer)
llvm::report_fatal_error(
llvm::Twine("failed to read PIM target config '")
+ pimTargetConfig.getValue() + "': " + buffer.getError().message());
auto parsed = llvm::json::parse(buffer.get()->getBuffer());
if (!parsed)
llvm::report_fatal_error(
llvm::Twine("failed to parse PIM target config '")
+ pimTargetConfig.getValue() + "': "
+ llvm::toString(parsed.takeError()));
const llvm::json::Object* root = parsed->getAsObject();
if (!root)
llvm::report_fatal_error("PIM target config must contain a JSON object");
const llvm::json::Object& chip = requireObject(*root, "chip_config", "root");
const llvm::json::Object& core = requireObject(chip, "core_config", "chip_config");
const llvm::json::Object& matrix =
requireObject(core, "matrix_config", "chip_config.core_config");
const llvm::json::Object& localMemory =
requireObject(core, "local_memory_config", "chip_config.core_config");
const llvm::json::Object& network =
requireObject(chip, "network_config", "chip_config");
std::optional<int64_t> coreCount = chip.getInteger("core_cnt");
if (!coreCount || *coreCount <= 0)
llvm::report_fatal_error("PIM target config field 'core_cnt' must be a positive integer");
target.processorCount = static_cast<size_t>(*coreCount);
target.residentWeightCapacity =
getConfigCost(matrix, "xbar_array_count", target.residentWeightCapacity);
std::tie(target.matrixRows, target.matrixColumns) =
getConfigPair(matrix, "xbar_size");
if (target.processorCount != static_cast<size_t>(coresCount.getValue())
|| target.residentWeightCapacity != crossbarCountInCore.getValue()
|| target.matrixRows != crossbarSize.getValue()
|| target.matrixColumns != crossbarSize.getValue())
llvm::report_fatal_error("PIM target config resources do not match --core-count, "
"--crossbar-count, and --crossbar-size");
loadPimInterProcessorLatencies(target, network);
target.processorPeriodNs =
getConfigCost(core, "period", target.processorPeriodNs);
target.localMemoryWidthBytes =
getConfigCost(localMemory, "data_width", target.localMemoryWidthBytes);
target.localMemoryReadLatencyCycles =
getConfigCost(localMemory, "read_latency_cycle", target.localMemoryReadLatencyCycles);
target.localMemoryWriteLatencyCycles =
getConfigCost(localMemory, "write_latency_cycle", target.localMemoryWriteLatencyCycles);
target.transferWidthBytes =
getConfigCost(network, "bus_width", target.transferWidthBytes);
target.vectorWidth = getConfigCost(core, "vector_width", target.vectorWidth);
target.vectorLatencyCycles =
getConfigCost(core, "vector_latency_cycle", target.vectorLatencyCycles);
target.matrixPeriodNs =
getConfigCost(matrix, "period", target.matrixPeriodNs);
target.matrixInputResolutionBits =
getConfigCost(matrix, "dac_resolution", target.matrixInputResolutionBits);
target.matrixInputLatencyCycles =
getConfigCost(matrix, "dac_latency_cycle", target.matrixInputLatencyCycles);
target.matrixInputParallelism =
getConfigCost(matrix, "dac_count", target.matrixInputParallelism);
target.matrixReadLatencyNs =
getConfigCost(matrix, "xbar_latency", target.matrixReadLatencyNs);
target.matrixSampleLatencyCycles =
getConfigCost(matrix, "sample_hold_latency_cycle", target.matrixSampleLatencyCycles);
target.matrixOutputLatencyCycles =
getConfigCost(matrix, "adc_latency_cycle", target.matrixOutputLatencyCycles);
target.matrixOutputParallelism =
getConfigCost(matrix, "adc_count", target.matrixOutputParallelism);
target.matrixShiftLatencyCycles =
getConfigCost(matrix, "shift_adder_latency_cycle", target.matrixShiftLatencyCycles);
target.matrixBufferLatencyCycles =
getConfigCost(matrix, "output_buffer_latency_cycle", target.matrixBufferLatencyCycles);
target.matrixInputBufferLatencyCycles =
getConfigCost(matrix,
"input_buffer_latency_cycle",
target.matrixInputBufferLatencyCycles,
/*allowZero=*/true);
target.matrixPipeline = matrix.getBoolean("pipeline_mode").value_or(target.matrixPipeline);
return target;
}
} // namespace
void addPassesPim(OwningOpRef<ModuleOp>& module,
PassManager& pm,
EmissionTargetType& emissionTarget,
@@ -31,11 +291,13 @@ 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());
pm.addPass(createTrivialGraphComputeMergePass());
pm.addPass(createMergeComputeNodesPass());
pm.addPass(createTrivialGraphComputeMergePass(
schedulingTarget.residentWeightCapacity));
pm.addPass(createMergeComputeNodesPass(schedulingTarget));
pm.addPass(createMessagePass("Onnx lowered to Spatial"));
}
+5 -3
View File
@@ -42,7 +42,9 @@ WeightEmissionResult createAndPopulateWeightFolder(ArrayRef<WeightFileRequest> r
assert(isMatrixShape(shape) && "Weight matrix must be 2-dimensional");
int64_t numRows = shape[0];
int64_t numCols = shape[1];
assert(numRows <= xbarSize && numCols <= xbarSize && "Weight dimensions must not exceed crossbar size");
assert(numRows <= xbarSize && numCols % xbarSize == 0
&& numCols / xbarSize <= static_cast<int64_t>(crossbarCountInCore)
&& "Weight dimensions must fit in one array group");
size_t elementByteWidth = getElementTypeSizeInBytes(denseAttr.getElementType());
@@ -57,7 +59,7 @@ WeightEmissionResult createAndPopulateWeightFolder(ArrayRef<WeightFileRequest> r
uint64_t zero = 0;
for (int64_t row = 0; row < xbarSize; row++) {
for (int64_t col = 0; col < xbarSize; col++) {
for (int64_t col = 0; col < numCols; col++) {
if (row < numRows && col < numCols) {
int64_t elementIndex = weightView.offset + row * weightView.strides[0] + col * weightView.strides[1];
APInt bits = denseAttr.getValues<APFloat>()[elementIndex].bitcastToAPInt();
@@ -73,7 +75,7 @@ WeightEmissionResult createAndPopulateWeightFolder(ArrayRef<WeightFileRequest> r
weightFileStream.close();
materializedWeights.push_back({weightView, newFileName});
uint64_t weightBytes = pim::checkedMulOrCrash(
pim::checkedMulOrCrash(static_cast<size_t>(xbarSize), static_cast<size_t>(xbarSize), "weight element count"),
pim::checkedMulOrCrash(static_cast<size_t>(xbarSize), static_cast<size_t>(numCols), "weight element count"),
elementByteWidth,
"weight byte size");
result.totalWeightBytes = pim::checkedAddOrCrash(result.totalWeightBytes, weightBytes, "total weight bytes");