Merge branch 'TestRottoConDeadLock' of chef.heaplab.deib.polimi.it:nnicolosi/Raptor into TestRottoConDeadLock
Validate Operations / validate-operations (push) Has been cancelled
Validate Operations / validate-operations (push) Has been cancelled
This commit is contained in:
@@ -17,6 +17,7 @@ add_pim_library(OMPimCompilerUtils
|
||||
PimCompilerUtils.cpp
|
||||
PimArtifactWriter.cpp
|
||||
PimCodeGen.cpp
|
||||
PimCoreProgram.cpp
|
||||
PimMemoryLiveness.cpp
|
||||
PimWeightEmitter.cpp
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstring>
|
||||
#include <vector>
|
||||
#include <limits>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/WeightUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimArtifactWriter.hpp"
|
||||
@@ -20,6 +20,65 @@ using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
PimInstructionWriter::PimInstructionWriter(raw_pwrite_stream& stream)
|
||||
: stream(stream), buffer(kBufferSize) {
|
||||
pim_binary::writeHeader(stream);
|
||||
}
|
||||
|
||||
PimInstructionWriter::PimInstructionWriter(raw_fd_ostream& stream)
|
||||
: PimInstructionWriter(static_cast<raw_pwrite_stream&>(stream)) {
|
||||
fileStream = &stream;
|
||||
}
|
||||
|
||||
bool PimInstructionWriter::consumeStreamError() {
|
||||
if (!fileStream || !fileStream->has_error())
|
||||
return false;
|
||||
fileStream->clear_error();
|
||||
hasFailure = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::flushBuffer() {
|
||||
if (hasFailure)
|
||||
return failure();
|
||||
if (bufferedBytes != 0) {
|
||||
stream.write(buffer.data(), bufferedBytes);
|
||||
bufferedBytes = 0;
|
||||
}
|
||||
return consumeStreamError() ? failure() : success();
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::append(const pim_binary::InstructionRecord& record) {
|
||||
if (hasFailure || finalized || count == std::numeric_limits<uint32_t>::max()) {
|
||||
hasFailure = true;
|
||||
return failure();
|
||||
}
|
||||
|
||||
if (bufferedBytes + pim_binary::kRecordSize > buffer.size() && failed(flushBuffer()))
|
||||
return failure();
|
||||
|
||||
pim_binary::EncodedInstruction encoded = pim_binary::encodeInstructionRecord(record);
|
||||
std::memcpy(buffer.data() + bufferedBytes, encoded.data(), encoded.size());
|
||||
bufferedBytes += encoded.size();
|
||||
++count;
|
||||
return success();
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::finalize() {
|
||||
if (finalized)
|
||||
return failure();
|
||||
finalized = true;
|
||||
if (hasFailure || failed(flushBuffer()))
|
||||
return failure();
|
||||
|
||||
stream.flush();
|
||||
if (consumeStreamError())
|
||||
return failure();
|
||||
pim_binary::patchInstructionCount(stream, count);
|
||||
stream.flush();
|
||||
return consumeStreamError() ? failure() : success();
|
||||
}
|
||||
|
||||
OnnxMlirCompilerErrorCodes
|
||||
writeMemoryBinary(ModuleOp moduleOp, func::FuncOp funcOp, PimAcceleratorMemory& memory, StringRef outputDirPath) {
|
||||
auto memoryFilePath = (outputDirPath + "/memory.bin").str();
|
||||
|
||||
@@ -2,16 +2,45 @@
|
||||
|
||||
#include "mlir/Dialect/Func/IR/FuncOps.h"
|
||||
#include "mlir/IR/BuiltinOps.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <vector>
|
||||
|
||||
#include "onnx-mlir/Compiler/OMCompilerTypes.h"
|
||||
#include "src/Accelerators/PIM/Compiler/PimBinaryFormat.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
class PimAcceleratorMemory;
|
||||
|
||||
class PimInstructionWriter {
|
||||
public:
|
||||
explicit PimInstructionWriter(llvm::raw_pwrite_stream& stream);
|
||||
explicit PimInstructionWriter(llvm::raw_fd_ostream& stream);
|
||||
|
||||
mlir::LogicalResult append(const pim_binary::InstructionRecord& record);
|
||||
mlir::LogicalResult finalize();
|
||||
|
||||
private:
|
||||
static constexpr size_t kBufferSize = 256 * 1024;
|
||||
mlir::LogicalResult flushBuffer();
|
||||
bool consumeStreamError();
|
||||
|
||||
llvm::raw_pwrite_stream& stream;
|
||||
llvm::raw_fd_ostream* fileStream = nullptr;
|
||||
std::vector<char> buffer;
|
||||
size_t bufferedBytes = 0;
|
||||
uint32_t count = 0;
|
||||
bool finalized = false;
|
||||
bool hasFailure = false;
|
||||
};
|
||||
|
||||
OnnxMlirCompilerErrorCodes writeMemoryBinary(mlir::ModuleOp moduleOp,
|
||||
mlir::func::FuncOp funcOp,
|
||||
PimAcceleratorMemory& memory,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
#pragma once
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
#include "llvm/Support/Endian.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
@@ -53,6 +54,14 @@ enum class Opcode : uint32_t {
|
||||
sync = 32,
|
||||
};
|
||||
|
||||
inline constexpr size_t kOpcodeCount = static_cast<size_t>(Opcode::sync) + 1;
|
||||
inline constexpr std::array<llvm::StringLiteral, kOpcodeCount> kOpcodeNames = {
|
||||
"nop", "sldi", "sld", "sadd", "ssub", "smul", "saddi", "smuli", "setbw", "mvmul", "vvadd",
|
||||
"vvsub", "vvmul", "vvdmul", "vvmax", "vvsll", "vvsra", "vavg", "vrelu", "vtanh", "vsigm", "vsoftmax",
|
||||
"vmv", "vrsu", "vrsl", "ld", "st", "lldi", "lmv", "send", "recv", "wait", "sync",
|
||||
};
|
||||
static_assert(kOpcodeNames.size() == kOpcodeCount);
|
||||
|
||||
struct InstructionRecord {
|
||||
Opcode opcode = Opcode::nop;
|
||||
uint8_t rd = 0;
|
||||
@@ -64,14 +73,29 @@ struct InstructionRecord {
|
||||
uint8_t flags = 0;
|
||||
};
|
||||
|
||||
using EncodedInstruction = std::array<char, kRecordSize>;
|
||||
static_assert(kRecordSize == 20);
|
||||
static_assert(sizeof(EncodedInstruction) == kRecordSize);
|
||||
|
||||
inline EncodedInstruction encodeInstructionRecord(const InstructionRecord& record) {
|
||||
EncodedInstruction encoded = {};
|
||||
encoded[0] = static_cast<char>(static_cast<uint8_t>(record.opcode));
|
||||
encoded[1] = static_cast<char>(record.rd);
|
||||
encoded[2] = static_cast<char>(record.r1);
|
||||
encoded[3] = static_cast<char>(record.flags);
|
||||
llvm::support::endian::write32le(encoded.data() + 4, static_cast<uint32_t>(record.r2OrImm));
|
||||
llvm::support::endian::write32le(encoded.data() + 8, static_cast<uint32_t>(record.generic1));
|
||||
llvm::support::endian::write32le(encoded.data() + 12, static_cast<uint32_t>(record.generic2));
|
||||
llvm::support::endian::write32le(encoded.data() + 16, static_cast<uint32_t>(record.generic3));
|
||||
return encoded;
|
||||
}
|
||||
|
||||
inline void writeUint32LE(llvm::raw_ostream& os, uint32_t value) {
|
||||
std::array<char, sizeof(uint32_t)> bytes;
|
||||
llvm::support::endian::write32le(bytes.data(), value);
|
||||
os.write(bytes.data(), bytes.size());
|
||||
}
|
||||
|
||||
inline void writeInt32LE(llvm::raw_ostream& os, int32_t value) { writeUint32LE(os, static_cast<uint32_t>(value)); }
|
||||
|
||||
inline void writeHeader(llvm::raw_ostream& os) {
|
||||
os.write(kMagic, sizeof(kMagic));
|
||||
writeUint32LE(os, kVersion);
|
||||
@@ -84,17 +108,6 @@ inline void patchInstructionCount(llvm::raw_pwrite_stream& os, uint32_t instruct
|
||||
os.pwrite(bytes.data(), bytes.size(), kCountOffset);
|
||||
}
|
||||
|
||||
inline void writeInstructionRecord(llvm::raw_ostream& os, const InstructionRecord& record) {
|
||||
os << static_cast<char>(static_cast<uint8_t>(record.opcode));
|
||||
os << static_cast<char>(record.rd);
|
||||
os << static_cast<char>(record.r1);
|
||||
os << static_cast<char>(record.flags);
|
||||
writeInt32LE(os, record.r2OrImm);
|
||||
writeInt32LE(os, record.generic1);
|
||||
writeInt32LE(os, record.generic2);
|
||||
writeInt32LE(os, record.generic3);
|
||||
}
|
||||
|
||||
inline int32_t toI32(int64_t value) { return onnx_mlir::pim::checkedI32OrCrash(value, "binary field"); }
|
||||
|
||||
inline uint8_t toU8(int64_t value) {
|
||||
@@ -107,113 +120,64 @@ inline int32_t getOptionalInt(const llvm::json::Object& object, llvm::StringRef
|
||||
return defaultValue;
|
||||
}
|
||||
|
||||
struct InstructionJsonFormat {
|
||||
bool rd;
|
||||
bool r1;
|
||||
bool offset;
|
||||
llvm::StringLiteral r2;
|
||||
llvm::StringLiteral generic1;
|
||||
llvm::StringLiteral generic2;
|
||||
llvm::StringLiteral generic3;
|
||||
};
|
||||
|
||||
inline constexpr std::array<InstructionJsonFormat, kOpcodeCount> kInstructionJsonFormats = {{
|
||||
{false, false, false, "", "", "", "" }, // nop
|
||||
{true, false, false, "imm", "", "", "" }, // sldi
|
||||
{true, true, true, "", "", "", "" }, // sld
|
||||
{true, true, false, "rs2", "", "", "" }, // sadd
|
||||
{true, true, false, "rs2", "", "", "" }, // ssub
|
||||
{true, true, false, "rs2", "", "", "" }, // smul
|
||||
{true, true, false, "imm", "", "", "" }, // saddi
|
||||
{true, true, false, "imm", "", "", "" }, // smuli
|
||||
{false, false, false, "", "ibiw", "obiw", "" }, // setbw
|
||||
{true, true, false, "mbiw", "relu", "group", "" }, // mvmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvadd
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsub
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvdmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvmax
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsll
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsra
|
||||
{true, true, true, "rs2", "", "", "len" }, // vavg
|
||||
{true, true, true, "", "", "", "len" }, // vrelu
|
||||
{true, true, true, "", "", "", "len" }, // vtanh
|
||||
{true, true, true, "", "", "", "len" }, // vsigm
|
||||
{true, true, true, "", "", "", "len" }, // vsoftmax
|
||||
{true, true, true, "rs2", "", "", "len" }, // vmv
|
||||
{true, true, true, "rs2", "", "", "len" }, // vrsu
|
||||
{true, true, true, "rs2", "", "", "len" }, // vrsl
|
||||
{true, true, true, "", "", "", "size"}, // ld
|
||||
{true, true, true, "", "", "", "size"}, // st
|
||||
{true, false, true, "imm", "", "", "len" }, // lldi
|
||||
{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
|
||||
}};
|
||||
static_assert(kInstructionJsonFormats.size() == kOpcodeCount);
|
||||
|
||||
inline Opcode opcodeFromString(llvm::StringRef opName) {
|
||||
if (opName == "nop")
|
||||
return Opcode::nop;
|
||||
if (opName == "sldi")
|
||||
return Opcode::sldi;
|
||||
if (opName == "sld")
|
||||
return Opcode::sld;
|
||||
if (opName == "sadd")
|
||||
return Opcode::sadd;
|
||||
if (opName == "ssub")
|
||||
return Opcode::ssub;
|
||||
if (opName == "smul")
|
||||
return Opcode::smul;
|
||||
if (opName == "saddi")
|
||||
return Opcode::saddi;
|
||||
if (opName == "smuli")
|
||||
return Opcode::smuli;
|
||||
if (opName == "setbw")
|
||||
return Opcode::setbw;
|
||||
if (opName == "mvmul")
|
||||
return Opcode::mvmul;
|
||||
if (opName == "vvadd")
|
||||
return Opcode::vvadd;
|
||||
if (opName == "vvsub")
|
||||
return Opcode::vvsub;
|
||||
if (opName == "vvmul")
|
||||
return Opcode::vvmul;
|
||||
if (opName == "vvdmul")
|
||||
return Opcode::vvdmul;
|
||||
if (opName == "vvmax")
|
||||
return Opcode::vvmax;
|
||||
if (opName == "vvsll")
|
||||
return Opcode::vvsll;
|
||||
if (opName == "vvsra")
|
||||
return Opcode::vvsra;
|
||||
if (opName == "vavg")
|
||||
return Opcode::vavg;
|
||||
if (opName == "vrelu")
|
||||
return Opcode::vrelu;
|
||||
if (opName == "vtanh")
|
||||
return Opcode::vtanh;
|
||||
if (opName == "vsigm")
|
||||
return Opcode::vsigm;
|
||||
if (opName == "vsoftmax")
|
||||
return Opcode::vsoftmax;
|
||||
if (opName == "vmv")
|
||||
return Opcode::vmv;
|
||||
if (opName == "vrsu")
|
||||
return Opcode::vrsu;
|
||||
if (opName == "vrsl")
|
||||
return Opcode::vrsl;
|
||||
if (opName == "ld")
|
||||
return Opcode::ld;
|
||||
if (opName == "st")
|
||||
return Opcode::st;
|
||||
if (opName == "lldi")
|
||||
return Opcode::lldi;
|
||||
if (opName == "lmv")
|
||||
return Opcode::lmv;
|
||||
if (opName == "send")
|
||||
return Opcode::send;
|
||||
if (opName == "recv")
|
||||
return Opcode::recv;
|
||||
if (opName == "wait")
|
||||
return Opcode::wait;
|
||||
if (opName == "sync")
|
||||
return Opcode::sync;
|
||||
for (auto [index, name] : llvm::enumerate(kOpcodeNames))
|
||||
if (opName == name)
|
||||
return static_cast<Opcode>(index);
|
||||
llvm_unreachable("Unsupported PIM binary opcode");
|
||||
}
|
||||
|
||||
inline llvm::StringRef opcodeToString(Opcode opcode) {
|
||||
switch (opcode) {
|
||||
case Opcode::nop: return "nop";
|
||||
case Opcode::sldi: return "sldi";
|
||||
case Opcode::sld: return "sld";
|
||||
case Opcode::sadd: return "sadd";
|
||||
case Opcode::ssub: return "ssub";
|
||||
case Opcode::smul: return "smul";
|
||||
case Opcode::saddi: return "saddi";
|
||||
case Opcode::smuli: return "smuli";
|
||||
case Opcode::setbw: return "setbw";
|
||||
case Opcode::mvmul: return "mvmul";
|
||||
case Opcode::vvadd: return "vvadd";
|
||||
case Opcode::vvsub: return "vvsub";
|
||||
case Opcode::vvmul: return "vvmul";
|
||||
case Opcode::vvdmul: return "vvdmul";
|
||||
case Opcode::vvmax: return "vvmax";
|
||||
case Opcode::vvsll: return "vvsll";
|
||||
case Opcode::vvsra: return "vvsra";
|
||||
case Opcode::vavg: return "vavg";
|
||||
case Opcode::vrelu: return "vrelu";
|
||||
case Opcode::vtanh: return "vtanh";
|
||||
case Opcode::vsigm: return "vsigm";
|
||||
case Opcode::vsoftmax: return "vsoftmax";
|
||||
case Opcode::vmv: return "vmv";
|
||||
case Opcode::vrsu: return "vrsu";
|
||||
case Opcode::vrsl: return "vrsl";
|
||||
case Opcode::ld: return "ld";
|
||||
case Opcode::st: return "st";
|
||||
case Opcode::lldi: return "lldi";
|
||||
case Opcode::lmv: return "lmv";
|
||||
case Opcode::send: return "send";
|
||||
case Opcode::recv: return "recv";
|
||||
case Opcode::wait: return "wait";
|
||||
case Opcode::sync: return "sync";
|
||||
}
|
||||
llvm_unreachable("Unsupported PIM binary opcode");
|
||||
size_t index = static_cast<size_t>(opcode);
|
||||
assert(index < kOpcodeNames.size() && "Unsupported PIM binary opcode");
|
||||
return kOpcodeNames[index];
|
||||
}
|
||||
|
||||
inline InstructionRecord makeInstructionRecord(const llvm::json::Object& instruction) {
|
||||
@@ -221,43 +185,25 @@ inline InstructionRecord makeInstructionRecord(const llvm::json::Object& instruc
|
||||
std::optional<llvm::StringRef> opName = instruction.getString("op");
|
||||
assert(opName && "Missing op field in PIM instruction");
|
||||
record.opcode = opcodeFromString(*opName);
|
||||
record.rd = toU8(getOptionalInt(instruction, "rd"));
|
||||
record.r1 = toU8(getOptionalInt(instruction, "rs1"));
|
||||
|
||||
switch (record.opcode) {
|
||||
case Opcode::sldi:
|
||||
case Opcode::saddi:
|
||||
case Opcode::smuli:
|
||||
case Opcode::lldi: record.r2OrImm = getOptionalInt(instruction, "imm"); break;
|
||||
case Opcode::mvmul:
|
||||
record.r2OrImm = getOptionalInt(instruction, "mbiw");
|
||||
record.generic1 = getOptionalInt(instruction, "relu");
|
||||
record.generic2 = getOptionalInt(instruction, "group");
|
||||
break;
|
||||
case Opcode::setbw:
|
||||
record.generic1 = getOptionalInt(instruction, "ibiw");
|
||||
record.generic2 = getOptionalInt(instruction, "obiw");
|
||||
break;
|
||||
case Opcode::send:
|
||||
case Opcode::recv:
|
||||
record.r2OrImm = getOptionalInt(instruction, "core");
|
||||
record.generic3 = getOptionalInt(instruction, "size");
|
||||
break;
|
||||
default: record.r2OrImm = getOptionalInt(instruction, "rs2"); break;
|
||||
}
|
||||
|
||||
if (record.opcode != Opcode::mvmul && record.opcode != Opcode::setbw) {
|
||||
const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)];
|
||||
if (format.rd)
|
||||
record.rd = toU8(getOptionalInt(instruction, "rd"));
|
||||
if (format.r1)
|
||||
record.r1 = toU8(getOptionalInt(instruction, "rs1"));
|
||||
if (!format.r2.empty())
|
||||
record.r2OrImm = getOptionalInt(instruction, format.r2);
|
||||
if (!format.generic1.empty())
|
||||
record.generic1 = getOptionalInt(instruction, format.generic1);
|
||||
if (!format.generic2.empty())
|
||||
record.generic2 = getOptionalInt(instruction, format.generic2);
|
||||
if (format.offset) {
|
||||
if (auto* offsetValue = instruction.getObject("offset")) {
|
||||
record.generic1 = getOptionalInt(*offsetValue, "offset_select");
|
||||
record.generic2 = getOptionalInt(*offsetValue, "offset_value");
|
||||
}
|
||||
}
|
||||
|
||||
if (instruction.get("len"))
|
||||
record.generic3 = getOptionalInt(instruction, "len");
|
||||
else if (instruction.get("size") && record.opcode != Opcode::send && record.opcode != Opcode::recv)
|
||||
record.generic3 = getOptionalInt(instruction, "size");
|
||||
|
||||
if (!format.generic3.empty())
|
||||
record.generic3 = getOptionalInt(instruction, format.generic3);
|
||||
return record;
|
||||
}
|
||||
|
||||
@@ -271,98 +217,21 @@ inline llvm::json::Object makeInstructionJson(const InstructionRecord& record) {
|
||||
offset["offset_value"] = offsetValue;
|
||||
instruction["offset"] = std::move(offset);
|
||||
};
|
||||
|
||||
switch (record.opcode) {
|
||||
case Opcode::sldi:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["imm"] = record.r2OrImm;
|
||||
break;
|
||||
case Opcode::sld:
|
||||
const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)];
|
||||
if (format.rd)
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
if (format.r1)
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
if (!format.r2.empty())
|
||||
instruction[format.r2] = record.r2OrImm;
|
||||
if (!format.generic1.empty())
|
||||
instruction[format.generic1] = record.generic1;
|
||||
if (!format.generic2.empty())
|
||||
instruction[format.generic2] = record.generic2;
|
||||
if (format.offset)
|
||||
addOffset(record.generic1, record.generic2);
|
||||
break;
|
||||
case Opcode::sadd:
|
||||
case Opcode::ssub:
|
||||
case Opcode::smul:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
instruction["rs2"] = record.r2OrImm;
|
||||
break;
|
||||
case Opcode::saddi:
|
||||
case Opcode::smuli:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
instruction["imm"] = record.r2OrImm;
|
||||
break;
|
||||
case Opcode::setbw:
|
||||
instruction["ibiw"] = record.generic1;
|
||||
instruction["obiw"] = record.generic2;
|
||||
break;
|
||||
case Opcode::mvmul:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
instruction["mbiw"] = record.r2OrImm;
|
||||
instruction["relu"] = record.generic1;
|
||||
instruction["group"] = record.generic2;
|
||||
break;
|
||||
case Opcode::vvadd:
|
||||
case Opcode::vvsub:
|
||||
case Opcode::vvmul:
|
||||
case Opcode::vvdmul:
|
||||
case Opcode::vvmax:
|
||||
case Opcode::vvsll:
|
||||
case Opcode::vvsra:
|
||||
case Opcode::vavg:
|
||||
case Opcode::vmv:
|
||||
case Opcode::vrsu:
|
||||
case Opcode::vrsl:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
instruction["rs2"] = record.r2OrImm;
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["len"] = record.generic3;
|
||||
break;
|
||||
case Opcode::vrelu:
|
||||
case Opcode::vtanh:
|
||||
case Opcode::vsigm:
|
||||
case Opcode::vsoftmax:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["len"] = record.generic3;
|
||||
break;
|
||||
case Opcode::ld:
|
||||
case Opcode::st:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["size"] = record.generic3;
|
||||
break;
|
||||
case Opcode::lldi:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["imm"] = record.r2OrImm;
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["len"] = record.generic3;
|
||||
break;
|
||||
case Opcode::lmv:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["len"] = record.generic3;
|
||||
break;
|
||||
case Opcode::send:
|
||||
case Opcode::recv:
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
instruction["core"] = record.r2OrImm;
|
||||
addOffset(record.generic1, record.generic2);
|
||||
instruction["size"] = record.generic3;
|
||||
break;
|
||||
case Opcode::wait:
|
||||
case Opcode::sync:
|
||||
case Opcode::nop: break;
|
||||
}
|
||||
|
||||
if (!format.generic3.empty())
|
||||
instruction[format.generic3] = record.generic3;
|
||||
return instruction;
|
||||
}
|
||||
|
||||
|
||||
+220
-559
File diff suppressed because it is too large
Load Diff
@@ -24,6 +24,10 @@
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct CompiledCoreProgram;
|
||||
struct CompiledTransposePlan;
|
||||
class PimInstructionWriter;
|
||||
|
||||
struct MemEntry {
|
||||
size_t address;
|
||||
size_t size;
|
||||
@@ -35,9 +39,7 @@ struct PhysicalSlotInfo {
|
||||
size_t size = 0;
|
||||
};
|
||||
|
||||
struct MemoryPlanArtifacts {
|
||||
std::string textReport;
|
||||
};
|
||||
using MemoryPlanArtifacts = std::string;
|
||||
|
||||
struct MemoryValueKey {
|
||||
mlir::Value value;
|
||||
@@ -98,6 +100,7 @@ class PimMemory {
|
||||
size_t nextPhysicalSlotId = 0;
|
||||
|
||||
MemEntry* gatherMemEntry(mlir::Value value, std::optional<unsigned> lane = std::nullopt);
|
||||
size_t allocateAddress(size_t size, const MemoryValueKey& key);
|
||||
void allocateGatheredMemory();
|
||||
void allocateMemoryForValue(const MemoryValueKey& key, MemEntry& memEntry, MemoryReportKind reportKind);
|
||||
PhysicalSlotInfo allocatePhysicalSlot(size_t slotSize, const MemoryValueKey& key);
|
||||
@@ -110,7 +113,6 @@ public:
|
||||
void allocateCore(mlir::Operation* op, std::optional<unsigned> lane = std::nullopt);
|
||||
MemoryReportRow getReportRow() const;
|
||||
const MemoryPlanArtifacts& getLivenessArtifacts() const { return livenessArtifacts; }
|
||||
void remove(mlir::Value val);
|
||||
|
||||
size_t getFirstAvailableAddress() const { return firstAvailableAddress; }
|
||||
MemEntry getMemEntry(const MemoryValueKey& key) const;
|
||||
@@ -153,12 +155,11 @@ public:
|
||||
uint64_t totalAllocaBytes);
|
||||
void setTotalWeightBytes(uint64_t bytes) { totalWeightBytes = bytes; }
|
||||
void flushReport();
|
||||
void clean(mlir::Operation* op);
|
||||
};
|
||||
|
||||
struct CoreEmissionJob {
|
||||
mlir::Operation* coreLikeOp = nullptr;
|
||||
size_t originalCoreId = 0;
|
||||
const CompiledCoreProgram* program = nullptr;
|
||||
size_t emittedCoreId = 0;
|
||||
llvm::SmallVector<unsigned, 4> lanes;
|
||||
std::optional<uint64_t> batchReportId;
|
||||
@@ -166,11 +167,10 @@ struct CoreEmissionJob {
|
||||
|
||||
class PimCodeGen {
|
||||
PimAcceleratorMemory& memory;
|
||||
llvm::raw_fd_ostream& coreBinaryStream;
|
||||
PimInstructionWriter& instructionWriter;
|
||||
llvm::raw_fd_ostream* coreJsonStream;
|
||||
const llvm::DenseMap<size_t, size_t>& emittedCoreIds;
|
||||
std::optional<unsigned> batchLane;
|
||||
mutable uint32_t emittedInstructionCount = 0;
|
||||
mutable std::array<std::optional<int32_t>, 256> scalarRegisterValues = {};
|
||||
|
||||
size_t addressOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
|
||||
@@ -181,30 +181,44 @@ class PimCodeGen {
|
||||
void emitInstruction(const pim_binary::InstructionRecord& instruction) const;
|
||||
void updateScalarRegisterCache(const pim_binary::InstructionRecord& instruction) const;
|
||||
|
||||
void genSetRegisterImmediate(uint8_t registerNumber, int32_t immediate) const;
|
||||
void genSetRegisterImmediateUnsigned(size_t registerNumber, size_t immediate) const;
|
||||
void setupRd(size_t rdAddress, size_t rdOffset) const;
|
||||
void setupRdRs1(size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset) const;
|
||||
void setupRdRs1Rs2(
|
||||
size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const;
|
||||
|
||||
void emitMemCopyOp(mlir::StringRef opName,
|
||||
void emitMemCopyOp(pim_binary::Opcode opcode,
|
||||
size_t rdAddr,
|
||||
size_t rdOffset,
|
||||
size_t rs1Addr,
|
||||
size_t rs1Offset,
|
||||
size_t size,
|
||||
mlir::StringRef sizeFieldName = "size") const;
|
||||
void emitCommunicationOp(mlir::StringRef opName, size_t bufferAddr, size_t coreId, size_t size) const;
|
||||
void emitCommunicationOp(pim_binary::Opcode opcode, size_t bufferAddr, size_t coreId, size_t size) const;
|
||||
void emitMvmOp(size_t groupId, size_t rdAddr, size_t rdOffset, size_t rs1Addr, size_t rs1Offset) const;
|
||||
|
||||
public:
|
||||
void emitBinaryVectorOp(pim_binary::Opcode opcode,
|
||||
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,
|
||||
llvm::raw_fd_ostream& coreBinary,
|
||||
PimInstructionWriter& instructionWriter,
|
||||
llvm::raw_fd_ostream* coreJson,
|
||||
const llvm::DenseMap<size_t, size_t>& emittedCoreIds)
|
||||
: memory(memory), coreBinaryStream(coreBinary), coreJsonStream(coreJson), emittedCoreIds(emittedCoreIds) {}
|
||||
: memory(memory), instructionWriter(instructionWriter), coreJsonStream(coreJson), emittedCoreIds(emittedCoreIds) {}
|
||||
|
||||
uint32_t getEmittedInstructionCount() const { return emittedInstructionCount; }
|
||||
void setBatchLane(std::optional<unsigned> lane) { batchLane = lane; }
|
||||
llvm::FailureOr<int64_t> indexOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
|
||||
return memory.getIndexValue(value, knowledge);
|
||||
@@ -221,18 +235,7 @@ public:
|
||||
template <typename MVMTy>
|
||||
void codeGenMVMLikeOp(size_t mvmId, MVMTy mvmLikeOp, bool transposeMatrix, const StaticValueKnowledge& knowledge);
|
||||
|
||||
void codeGenVVAddOp(pim::PimVVAddOp vvaddOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVSubOp(pim::PimVVSubOp vvsubOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVMulOp(pim::PimVVMulOp vvmulOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVMaxOp(pim::PimVVMaxOp vvmaxOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVDMulOp(pim::PimVVDMulOp vvdmulOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVAvgOp(pim::PimVAvgOp vavgOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVReluOp(pim::PimVReluOp vreluOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVTanhOp(pim::PimVTanhOp vtanhOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVSigmOp(pim::PimVSigmOp vsigmOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVSoftmaxOp(pim::PimVSoftmaxOp vsoftmaxOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGetGlobalOp(mlir::memref::GetGlobalOp getGlobalOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenTransposeOp(pim::PimTransposeOp transposeOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenTransposeOp(const CompiledTransposePlan& plan, const StaticValueKnowledge& knowledge) const;
|
||||
};
|
||||
|
||||
OnnxMlirCompilerErrorCodes compileToPimCode(mlir::ModuleOp& moduleOpRef, std::string& outputDirName);
|
||||
|
||||
@@ -21,7 +21,7 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
|
||||
verifyExplicitPimCoreCount();
|
||||
|
||||
if (pimOnlyCodegen) {
|
||||
// Skip all the lowering passes and directly generate code for PIM.
|
||||
pm.addPass(createEmitPimCodePass());
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/CheckedArithmetic.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCoreProgram.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
using namespace llvm;
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace {
|
||||
|
||||
static FailureOr<CompiledCoreOpKind> classifyCompiledCoreOpKind(Operation& op) {
|
||||
if (isa<pim::PimMemCopyHostToDevOp>(op)) return CompiledCoreOpKind::Load;
|
||||
if (isa<pim::PimMemCopyDevToHostOp>(op)) return CompiledCoreOpKind::Store;
|
||||
if (isa<pim::PimMemCopyOp>(op)) return CompiledCoreOpKind::Lmv;
|
||||
if (isa<pim::PimReceiveOp>(op)) return CompiledCoreOpKind::Receive;
|
||||
if (isa<pim::PimSendOp>(op)) return CompiledCoreOpKind::Send;
|
||||
if (isa<pim::PimConcatOp>(op)) return CompiledCoreOpKind::Concat;
|
||||
if (isa<pim::PimVMMOp>(op)) return CompiledCoreOpKind::Vmm;
|
||||
if (isa<pim::PimTransposeOp>(op)) return CompiledCoreOpKind::Transpose;
|
||||
if (isa<pim::PimVVAddOp>(op)) return CompiledCoreOpKind::VVAdd;
|
||||
if (isa<pim::PimVVSubOp>(op)) return CompiledCoreOpKind::VVSub;
|
||||
if (isa<pim::PimVVMulOp>(op)) return CompiledCoreOpKind::VVMul;
|
||||
if (isa<pim::PimVVMaxOp>(op)) return CompiledCoreOpKind::VVMax;
|
||||
if (isa<pim::PimVVDMulOp>(op)) return CompiledCoreOpKind::VVDMul;
|
||||
if (isa<pim::PimVAvgOp>(op)) return CompiledCoreOpKind::VAvg;
|
||||
if (isa<pim::PimVReluOp>(op)) return CompiledCoreOpKind::VRelu;
|
||||
if (isa<pim::PimVTanhOp>(op)) return CompiledCoreOpKind::VTanh;
|
||||
if (isa<pim::PimVSigmOp>(op)) return CompiledCoreOpKind::VSigm;
|
||||
if (isa<pim::PimVSoftmaxOp>(op)) return CompiledCoreOpKind::VSoftmax;
|
||||
return failure();
|
||||
}
|
||||
|
||||
static bool isStoragePreservingTranspose(ArrayRef<size_t> sourceShape, ArrayRef<int64_t> permutation) {
|
||||
SmallVector<unsigned> sourceNonUnitDims;
|
||||
SmallVector<unsigned> destinationSourceNonUnitDims;
|
||||
for (auto [dim, size] : llvm::enumerate(sourceShape))
|
||||
if (size != 1)
|
||||
sourceNonUnitDims.push_back(dim);
|
||||
for (int64_t sourceDim : permutation)
|
||||
if (sourceShape[sourceDim] != 1)
|
||||
destinationSourceNonUnitDims.push_back(static_cast<unsigned>(sourceDim));
|
||||
return sourceNonUnitDims == destinationSourceNonUnitDims;
|
||||
}
|
||||
|
||||
static FailureOr<CompiledTransposePlan> compileTransposePlan(pim::PimTransposeOp transposeOp) {
|
||||
auto sourceType = cast<ShapedType>(transposeOp.getInput().getType());
|
||||
ArrayRef<int64_t> sourceShape = sourceType.getShape();
|
||||
size_t rank = sourceShape.size();
|
||||
CompiledTransposePlan plan;
|
||||
plan.source = transposeOp.getInput();
|
||||
plan.destination = transposeOp.getOutputBuffer();
|
||||
plan.elementBytes = getElementTypeSizeInBytes(sourceType.getElementType());
|
||||
auto totalElements = pim::checkedSize(sourceType.getNumElements(), transposeOp, "transpose elements");
|
||||
if (failed(totalElements)) return failure();
|
||||
plan.totalElements = *totalElements;
|
||||
auto totalBytes = pim::checkedMul(plan.totalElements, plan.elementBytes, transposeOp, "transpose byte size");
|
||||
if (failed(totalBytes)) return failure();
|
||||
plan.totalBytes = *totalBytes;
|
||||
|
||||
SmallVector<int64_t> permutation = map_to_vector(transposeOp.getPermutation().getAsRange<IntegerAttr>(),
|
||||
[](IntegerAttr attr) { return attr.getInt(); });
|
||||
if (permutation.size() != rank) {
|
||||
transposeOp.emitOpError("requires permutation rank to match source rank for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
|
||||
SmallVector<size_t> destinationShape(rank);
|
||||
plan.destinationStrides.assign(rank, 1);
|
||||
plan.destinationDimensionForSource.assign(rank, 0);
|
||||
plan.destinationRewinds.assign(rank, 0);
|
||||
SmallVector<bool> seenSourceDimensions(rank, false);
|
||||
for (size_t dim = 0; dim < rank; ++dim) {
|
||||
auto size = pim::checkedSize(sourceShape[dim], transposeOp, "transpose source dimension");
|
||||
if (failed(size)) return failure();
|
||||
plan.sourceShape.push_back(*size);
|
||||
}
|
||||
for (auto [destinationDim, sourceDim] : llvm::enumerate(permutation)) {
|
||||
if (sourceDim < 0 || static_cast<size_t>(sourceDim) >= rank || seenSourceDimensions[sourceDim]) {
|
||||
transposeOp.emitOpError("requires a valid permutation containing each source dimension exactly once");
|
||||
return failure();
|
||||
}
|
||||
seenSourceDimensions[sourceDim] = true;
|
||||
destinationShape[destinationDim] = plan.sourceShape[sourceDim];
|
||||
plan.destinationDimensionForSource[sourceDim] = destinationDim;
|
||||
}
|
||||
for (size_t dim = rank; dim > 1; --dim) {
|
||||
auto stride = pim::checkedMul(
|
||||
plan.destinationStrides[dim - 1], destinationShape[dim - 1], transposeOp, "transpose destination stride");
|
||||
if (failed(stride)) return failure();
|
||||
plan.destinationStrides[dim - 2] = *stride;
|
||||
}
|
||||
for (size_t sourceDim = 0; sourceDim < rank; ++sourceDim) {
|
||||
auto rewind = pim::checkedMul(plan.sourceShape[sourceDim],
|
||||
plan.destinationStrides[plan.destinationDimensionForSource[sourceDim]],
|
||||
transposeOp,
|
||||
"transpose destination rewind");
|
||||
if (failed(rewind)) return failure();
|
||||
plan.destinationRewinds[sourceDim] = *rewind;
|
||||
}
|
||||
plan.storagePreserving = isStoragePreservingTranspose(plan.sourceShape, permutation);
|
||||
return plan;
|
||||
}
|
||||
|
||||
static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl<CompiledCoreNode>& plan) {
|
||||
for (Operation& op : block) {
|
||||
if (isa<pim::PimHaltOp, scf::YieldOp, memref::GetGlobalOp>(op) || isCoreStaticAddressOp(&op))
|
||||
continue;
|
||||
if (auto loadOp = dyn_cast<memref::LoadOp>(op); loadOp && succeeded(compileIndexExpr(loadOp.getResult())))
|
||||
continue;
|
||||
|
||||
if (auto forOp = dyn_cast<scf::ForOp>(op)) {
|
||||
auto lower = compileIndexExpr(forOp.getLowerBound());
|
||||
auto upper = compileIndexExpr(forOp.getUpperBound());
|
||||
auto step = compileIndexExpr(forOp.getStep());
|
||||
if (failed(lower) || failed(upper) || failed(step)) {
|
||||
forOp.emitOpError("requires statically evaluable scf.for bounds for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.kind = CompiledCoreNode::Kind::Loop;
|
||||
node.op = forOp;
|
||||
node.lowerBound = *lower;
|
||||
node.upperBound = *upper;
|
||||
node.step = *step;
|
||||
node.loopBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(forOp.getRegion().front(), *node.loopBody))) return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
if (auto ifOp = dyn_cast<scf::IfOp>(op)) {
|
||||
auto condition = compileIndexExpr(ifOp.getCondition());
|
||||
if (failed(condition)) {
|
||||
ifOp.emitOpError("requires statically evaluable scf.if condition for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.kind = CompiledCoreNode::Kind::If;
|
||||
node.op = ifOp;
|
||||
node.condition = *condition;
|
||||
node.thenBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
node.elseBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(ifOp.getThenRegion().front(), *node.thenBody))) return failure();
|
||||
if (!ifOp.getElseRegion().empty()
|
||||
&& failed(compileCoreEmissionPlan(ifOp.getElseRegion().front(), *node.elseBody)))
|
||||
return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
if (auto switchOp = dyn_cast<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 node;
|
||||
node.kind = CompiledCoreNode::Kind::IndexSwitch;
|
||||
node.op = switchOp;
|
||||
node.condition = *selector;
|
||||
llvm::append_range(node.caseValues, switchOp.getCases());
|
||||
for (Region& region : switchOp.getCaseRegions()) {
|
||||
auto body = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(region.front(), *body))) return failure();
|
||||
node.caseBodies.push_back(std::move(body));
|
||||
}
|
||||
node.defaultBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(switchOp.getDefaultRegion().front(), *node.defaultBody))) return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
|
||||
auto opKind = classifyCompiledCoreOpKind(op);
|
||||
if (failed(opKind)) {
|
||||
InFlightDiagnostic diagnostic = op.emitError() << "unsupported codegen for op '" << op.getName() << "'";
|
||||
if (auto coreOp = op.getParentOfType<pim::PimCoreOp>())
|
||||
diagnostic << " inside pim.core " << coreOp.getCoreId();
|
||||
else if (auto batchOp = op.getParentOfType<pim::PimCoreBatchOp>())
|
||||
diagnostic << " inside pim.core_batch with laneCount " << batchOp.getLaneCount();
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.op = &op;
|
||||
node.opKind = *opKind;
|
||||
if (*opKind == CompiledCoreOpKind::Transpose) {
|
||||
auto transposePlan = compileTransposePlan(cast<pim::PimTransposeOp>(op));
|
||||
if (failed(transposePlan)) return failure();
|
||||
node.transposePlan = *transposePlan;
|
||||
}
|
||||
plan.push_back(std::move(node));
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
LogicalResult compileCoreProgram(Operation* coreLikeOp, CompiledCoreProgram& program) {
|
||||
Block& block = isa<pim::PimCoreOp>(coreLikeOp) ? cast<pim::PimCoreOp>(coreLikeOp).getBody().front()
|
||||
: cast<pim::PimCoreBatchOp>(coreLikeOp).getBody().front();
|
||||
return compileCoreEmissionPlan(block, program.nodes);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,74 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Operation.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct CompiledTransposePlan {
|
||||
mlir::Value source;
|
||||
mlir::Value destination;
|
||||
size_t elementBytes = 0;
|
||||
size_t totalElements = 0;
|
||||
size_t totalBytes = 0;
|
||||
llvm::SmallVector<size_t> sourceShape;
|
||||
llvm::SmallVector<size_t> destinationStrides;
|
||||
llvm::SmallVector<unsigned> destinationDimensionForSource;
|
||||
llvm::SmallVector<size_t> destinationRewinds;
|
||||
bool storagePreserving = false;
|
||||
};
|
||||
|
||||
enum class CompiledCoreOpKind : uint8_t {
|
||||
Load,
|
||||
Store,
|
||||
Lmv,
|
||||
Receive,
|
||||
Send,
|
||||
Concat,
|
||||
Vmm,
|
||||
Transpose,
|
||||
VVAdd,
|
||||
VVSub,
|
||||
VVMul,
|
||||
VVMax,
|
||||
VVDMul,
|
||||
VAvg,
|
||||
VRelu,
|
||||
VTanh,
|
||||
VSigm,
|
||||
VSoftmax
|
||||
};
|
||||
|
||||
struct CompiledCoreNode {
|
||||
enum class Kind : uint8_t { Op, Loop, If, IndexSwitch };
|
||||
|
||||
Kind kind = Kind::Op;
|
||||
mlir::Operation* op = nullptr;
|
||||
CompiledCoreOpKind opKind = CompiledCoreOpKind::Load;
|
||||
CompiledIndexExpr lowerBound;
|
||||
CompiledIndexExpr upperBound;
|
||||
CompiledIndexExpr step;
|
||||
CompiledIndexExpr condition;
|
||||
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;
|
||||
std::optional<CompiledTransposePlan> transposePlan;
|
||||
};
|
||||
|
||||
struct CompiledCoreProgram {
|
||||
llvm::SmallVector<CompiledCoreNode, 32> nodes;
|
||||
};
|
||||
|
||||
mlir::LogicalResult compileCoreProgram(mlir::Operation* coreLikeOp, CompiledCoreProgram& program);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -42,8 +42,6 @@ static MemoryValueKey getMemoryValueKey(mlir::Value value, std::optional<unsigne
|
||||
struct MemoryTouchInterval {
|
||||
uint64_t start = 0;
|
||||
uint64_t end = 0;
|
||||
Operation* startOp = nullptr;
|
||||
Operation* endOp = nullptr;
|
||||
Operation* firstTouchOp = nullptr;
|
||||
Operation* lastTouchOp = nullptr;
|
||||
uint64_t firstTouchPosition = 0;
|
||||
@@ -218,22 +216,18 @@ static void appendAliasDescription(llvm::SmallVectorImpl<std::string>& aliases,
|
||||
struct OrderedTouchRange {
|
||||
uint64_t start = 0;
|
||||
uint64_t end = 0;
|
||||
Operation* startOp = nullptr;
|
||||
Operation* endOp = nullptr;
|
||||
bool escapedLoop = false;
|
||||
};
|
||||
|
||||
static OrderedTouchRange
|
||||
getEffectiveTouchRange(mlir::Value definingValue, Operation* user, const OperationOrdering& ordering) {
|
||||
OrderedTouchRange range {ordering.position.lookup(user), ordering.position.lookup(user), user, user, false};
|
||||
OrderedTouchRange range {ordering.position.lookup(user), ordering.position.lookup(user), false};
|
||||
for (Operation* current = user; current; current = current->getParentOp()) {
|
||||
auto forOp = dyn_cast<scf::ForOp>(current);
|
||||
if (!forOp || isWithin(definingValue, &forOp.getRegion()))
|
||||
continue;
|
||||
range.start = std::min(range.start, ordering.position.lookup(forOp));
|
||||
range.end = std::max(range.end, ordering.subtreeEnd.lookup(forOp));
|
||||
range.startOp = forOp;
|
||||
range.endOp = forOp;
|
||||
range.escapedLoop = true;
|
||||
}
|
||||
return range;
|
||||
@@ -247,8 +241,6 @@ computeMemoryTouchInterval(memref::AllocOp allocOp,
|
||||
MemoryTouchInterval interval;
|
||||
interval.start = ordering.position.lookup(allocOp);
|
||||
interval.end = interval.start;
|
||||
interval.startOp = allocOp;
|
||||
interval.endOp = allocOp;
|
||||
auto recordAlias = [&](mlir::Value value) {
|
||||
if (includeAliasDescriptions)
|
||||
appendAliasDescription(interval.aliasesFollowed, value);
|
||||
@@ -350,18 +342,14 @@ computeMemoryTouchInterval(memref::AllocOp allocOp,
|
||||
if (!interval.hasRuntimeUse) {
|
||||
interval.start = range.start;
|
||||
interval.end = range.end;
|
||||
interval.startOp = range.startOp;
|
||||
interval.endOp = range.endOp;
|
||||
interval.hasRuntimeUse = true;
|
||||
}
|
||||
else {
|
||||
if (range.start < interval.start) {
|
||||
interval.start = range.start;
|
||||
interval.startOp = range.startOp;
|
||||
}
|
||||
if (range.end > interval.end) {
|
||||
interval.end = range.end;
|
||||
interval.endOp = range.endOp;
|
||||
}
|
||||
}
|
||||
continue;
|
||||
@@ -380,8 +368,6 @@ computeMemoryTouchInterval(memref::AllocOp allocOp,
|
||||
interval.endUsedFallback = true;
|
||||
interval.start = ordering.position.lookup(allocOp);
|
||||
interval.end = fallbackEnd;
|
||||
interval.startOp = allocOp;
|
||||
interval.endOp = allocOp->getParentOp();
|
||||
interval.firstTouchPosition = interval.start;
|
||||
interval.lastTouchPosition = interval.end;
|
||||
addFallbackReason(interval.fallbackReason, "no runtime memory touch");
|
||||
@@ -390,7 +376,6 @@ computeMemoryTouchInterval(memref::AllocOp allocOp,
|
||||
|
||||
if (interval.endUsedFallback) {
|
||||
interval.end = std::max(interval.end, fallbackEnd);
|
||||
interval.endOp = allocOp->getParentOp();
|
||||
}
|
||||
|
||||
return interval;
|
||||
@@ -447,8 +432,6 @@ SmallVector<LocalAllocInterval, 0> onnx_mlir::buildLocalAllocIntervals(Operation
|
||||
interval.start = touchInterval.start;
|
||||
interval.end = touchInterval.end;
|
||||
interval.size = *checkedSize;
|
||||
interval.startOp = touchInterval.startOp;
|
||||
interval.endOp = touchInterval.endOp;
|
||||
interval.firstTouchOp = touchInterval.firstTouchOp;
|
||||
interval.lastTouchOp = touchInterval.lastTouchOp;
|
||||
interval.firstTouchPosition = touchInterval.firstTouchPosition;
|
||||
@@ -574,7 +557,7 @@ MemoryPlanArtifacts onnx_mlir::buildMemoryPlanArtifacts(Operation* coreLikeOp,
|
||||
double savedPercent =
|
||||
totalLogicalBytes == 0 ? 0.0 : 100.0 * static_cast<double>(savedBytes) / static_cast<double>(totalLogicalBytes);
|
||||
|
||||
raw_string_ostream os(artifacts.textReport);
|
||||
raw_string_ostream os(artifacts);
|
||||
os << "=== PIM Memory Liveness Report ===\n";
|
||||
os << "Op: " << coreLikeOp->getName() << "\n";
|
||||
if (lane)
|
||||
|
||||
@@ -21,8 +21,6 @@ struct LocalAllocInterval {
|
||||
uint64_t start = 0;
|
||||
uint64_t end = 0;
|
||||
size_t size = 0;
|
||||
mlir::Operation* startOp = nullptr;
|
||||
mlir::Operation* endOp = nullptr;
|
||||
mlir::Operation* firstTouchOp = nullptr;
|
||||
mlir::Operation* lastTouchOp = nullptr;
|
||||
uint64_t firstTouchPosition = 0;
|
||||
|
||||
Reference in New Issue
Block a user