This commit is contained in:
@@ -1119,7 +1119,8 @@ struct CompiledCoreNode {
|
||||
enum class Kind : uint8_t {
|
||||
Op,
|
||||
Loop,
|
||||
If
|
||||
If,
|
||||
IndexSwitch
|
||||
};
|
||||
|
||||
Kind kind = Kind::Op;
|
||||
@@ -1132,6 +1133,9 @@ struct CompiledCoreNode {
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> loopBody;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> thenBody;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> elseBody;
|
||||
llvm::SmallVector<int64_t> caseValues;
|
||||
llvm::SmallVector<std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>>> caseBodies;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> defaultBody;
|
||||
};
|
||||
|
||||
static FailureOr<CompiledCoreOpKind> classifyCompiledCoreOpKind(Operation& op) {
|
||||
@@ -1231,6 +1235,31 @@ compileCoreEmissionPlan(Block& block, Operation* weightOwner, llvm::SmallVectorI
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto switchOp = dyn_cast<mlir::scf::IndexSwitchOp>(op)) {
|
||||
auto selector = compileIndexExpr(switchOp.getArg());
|
||||
if (failed(selector)) {
|
||||
switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode switchNode;
|
||||
switchNode.kind = CompiledCoreNode::Kind::IndexSwitch;
|
||||
switchNode.op = switchOp.getOperation();
|
||||
switchNode.condition = *selector;
|
||||
llvm::append_range(switchNode.caseValues, switchOp.getCases());
|
||||
for (mlir::Region& region : switchOp.getCaseRegions()) {
|
||||
auto body = std::make_unique<llvm::SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(region.front(), weightOwner, *body)))
|
||||
return failure();
|
||||
switchNode.caseBodies.push_back(std::move(body));
|
||||
}
|
||||
switchNode.defaultBody = std::make_unique<llvm::SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(
|
||||
switchOp.getDefaultRegion().front(), weightOwner, *switchNode.defaultBody)))
|
||||
return failure();
|
||||
plan.push_back(std::move(switchNode));
|
||||
continue;
|
||||
}
|
||||
|
||||
auto opKind = classifyCompiledCoreOpKind(op);
|
||||
if (failed(opKind)) {
|
||||
InFlightDiagnostic diag = op.emitError() << "unsupported codegen for op '" << op.getName().getStringRef() << "'";
|
||||
@@ -1313,6 +1342,31 @@ static LogicalResult executeCompiledCorePlan(
|
||||
continue;
|
||||
}
|
||||
|
||||
if (node.kind == CompiledCoreNode::Kind::IndexSwitch) {
|
||||
auto selector = node.condition.evaluate(knowledge);
|
||||
auto switchOp = cast<mlir::scf::IndexSwitchOp>(node.op);
|
||||
if (failed(selector)) {
|
||||
switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
const llvm::SmallVectorImpl<CompiledCoreNode>* selectedBody = node.defaultBody.get();
|
||||
mlir::Region* selectedRegion = &switchOp.getDefaultRegion();
|
||||
for (auto [index, caseValue] : llvm::enumerate(node.caseValues))
|
||||
if (caseValue == *selector) {
|
||||
selectedBody = node.caseBodies[index].get();
|
||||
selectedRegion = &switchOp.getCaseRegions()[index];
|
||||
break;
|
||||
}
|
||||
if (failed(executeCompiledCorePlan(*selectedBody, coreCodeGen, knowledge,
|
||||
resolveWeightSlot, processedOperations,
|
||||
batchLane, batchLaneCount)))
|
||||
return failure();
|
||||
auto yield = cast<mlir::scf::YieldOp>(selectedRegion->front().getTerminator());
|
||||
for (auto [result, yielded] : llvm::zip(switchOp.getResults(), yield.getOperands()))
|
||||
knowledge.aliases[result] = resolveLoopCarriedAlias(yielded, knowledge);
|
||||
continue;
|
||||
}
|
||||
|
||||
switch (node.opKind) {
|
||||
case CompiledCoreOpKind::Load:
|
||||
coreCodeGen.codeGenLoadOp(cast<pim::PimMemCopyHostToDevOp>(node.op), knowledge);
|
||||
@@ -1413,6 +1467,36 @@ static int64_t codeGenCoreOps(
|
||||
return failed(result) ? -1 : static_cast<int64_t>(processedOperations);
|
||||
}
|
||||
|
||||
static OnnxMlirCompilerErrorCodes emitEmptyCoreArtifacts(StringRef outputDirPath, size_t emittedCoreId) {
|
||||
std::string outputCorePath =
|
||||
(outputDirPath + "/core_" + std::to_string(emittedCoreId) + ".pim").str();
|
||||
std::error_code errorCode;
|
||||
raw_fd_ostream coreBinaryStream(outputCorePath, errorCode, sys::fs::OF_None);
|
||||
if (errorCode) {
|
||||
errs() << "Error while opening core file `" << outputCorePath << "`: " << errorCode.message() << '\n';
|
||||
return InvalidOutputFileAccess;
|
||||
}
|
||||
|
||||
pim_binary::writeHeader(coreBinaryStream);
|
||||
pim_binary::patchInstructionCount(coreBinaryStream, 0);
|
||||
coreBinaryStream.close();
|
||||
|
||||
if (!pimEmitJson.getValue())
|
||||
return CompilerSuccess;
|
||||
|
||||
std::string outputCoreJsonPath =
|
||||
(outputDirPath + "/core_" + std::to_string(emittedCoreId) + ".json").str();
|
||||
errorCode = std::error_code();
|
||||
raw_fd_ostream coreJsonStream(outputCoreJsonPath, errorCode);
|
||||
if (errorCode) {
|
||||
errs() << "Error while opening core json file `" << outputCoreJsonPath << "`: " << errorCode.message() << '\n';
|
||||
return InvalidOutputFileAccess;
|
||||
}
|
||||
coreJsonStream << "[]";
|
||||
coreJsonStream.close();
|
||||
return CompilerSuccess;
|
||||
}
|
||||
|
||||
OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::string& outputDirPath) {
|
||||
if (!outputDirPath.empty()) {
|
||||
if (auto error = sys::fs::create_directory(outputDirPath)) {
|
||||
@@ -1657,6 +1741,13 @@ 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))
|
||||
return err;
|
||||
xbarsPerArrayGroup["core0"] = json::Array {};
|
||||
memory.recordCoreReport(0, MemoryReportRow {});
|
||||
}
|
||||
|
||||
llvm::SmallVector<WeightFileRequest, 8> weightRequests;
|
||||
weightRequests.reserve(jobs.size());
|
||||
for (size_t jobIndex = 0; jobIndex < jobs.size(); ++jobIndex) {
|
||||
|
||||
Reference in New Issue
Block a user