28#define DEBUG_TYPE "mir2vec"
31 "Number of lookups to MIR entities not present in the vocabulary");
33 "Number of register operands with no register class");
46 cl::desc(
"Weight for machine opcode embeddings"),
50 cl::desc(
"Weight for common operand embeddings"),
54 cl::desc(
"Weight for register operand embeddings"),
59 "Generate symbolic embeddings for MIR")),
64 "mir2vec-print-all-vocab-entries",
cl::init(
false),
65 cl::desc(
"Print all vocabulary entries including zero embeddings"),
76 VocabMap &&PhysicalRegisterMap,
77 VocabMap &&VirtualRegisterMap,
82 buildCanonicalOpcodeMapping();
83 unsigned CanonicalOpcodeCount = UniqueBaseOpcodeNames.size();
84 assert(CanonicalOpcodeCount > 0 &&
85 "No canonical opcodes found for target - invalid vocabulary");
87 buildRegisterOperandMapping();
90 Layout.OpcodeBase = 0;
91 Layout.CommonOperandBase = CanonicalOpcodeCount;
93 Layout.PhyRegBase = Layout.CommonOperandBase + std::size(CommonOperandNames);
94 Layout.VirtRegBase = Layout.PhyRegBase + RegisterOperandNames.size();
96 generateStorage(OpcodeMap, CommonOperandMap, PhysicalRegisterMap,
98 Layout.TotalEntries = Storage.size();
106 if (OpcodeMap.empty() || CommonOperandMap.empty() || PhyRegMap.empty() ||
109 "Empty vocabulary entries provided");
111 MIRVocabulary Vocab(std::move(OpcodeMap), std::move(CommonOperandMap),
112 std::move(PhyRegMap), std::move(
VirtRegMap), TII, TRI,
118 "Failed to create valid vocabulary storage");
120 return std::move(Vocab);
135 assert(!InstrName.
empty() &&
"Instruction name should not be empty");
138 static const Regex BaseOpcodeRegex(
"([a-zA-Z_]+)");
141 if (BaseOpcodeRegex.
match(InstrName, &Matches) && Matches.
size() > 1) {
144 while (!Match.
empty() && Match.
back() ==
'_')
150 return InstrName.
str();
154 assert(!UniqueBaseOpcodeNames.empty() &&
"Canonical mapping not built");
155 auto It = std::find(UniqueBaseOpcodeNames.begin(),
156 UniqueBaseOpcodeNames.end(), BaseName.
str());
157 assert(It != UniqueBaseOpcodeNames.end() &&
158 "Base name not found in unique opcodes");
159 return std::distance(UniqueBaseOpcodeNames.begin(), It);
162unsigned MIRVocabulary::getCanonicalOpcodeIndex(
unsigned Opcode)
const {
163 auto BaseOpcode = extractBaseOpcodeName(
TII.getName(Opcode));
164 return getCanonicalIndexForBaseName(BaseOpcode);
169 auto It = std::find(std::begin(CommonOperandNames),
170 std::end(CommonOperandNames), OperandName);
171 assert(It != std::end(CommonOperandNames) &&
172 "Operand name not found in common operands");
173 return Layout.CommonOperandBase +
174 std::distance(std::begin(CommonOperandNames), It);
179 bool IsPhysical)
const {
180 auto It = std::find(RegisterOperandNames.begin(), RegisterOperandNames.end(),
182 assert(It != RegisterOperandNames.end() &&
183 "Register name not found in register operands");
184 unsigned LocalIndex = std::distance(RegisterOperandNames.begin(), It);
185 return (IsPhysical ? Layout.PhyRegBase : Layout.VirtRegBase) + LocalIndex;
189 assert(Pos < Layout.TotalEntries &&
"Position out of bounds in vocabulary");
192 if (Pos < Layout.CommonOperandBase) {
194 auto It = UniqueBaseOpcodeNames.begin();
195 std::advance(It, Pos);
196 assert(It != UniqueBaseOpcodeNames.end() &&
197 "Canonical index out of bounds in opcode section");
201 auto getLocalIndex = [](
unsigned Pos,
size_t BaseOffset,
size_t Bound,
203 unsigned LocalIndex = Pos - BaseOffset;
209 if (Pos < Layout.PhyRegBase) {
210 unsigned LocalIndex = getLocalIndex(
211 Pos, Layout.CommonOperandBase, std::size(CommonOperandNames),
212 "Local index out of bounds in common operands");
213 return CommonOperandNames[LocalIndex].str();
217 if (Pos < Layout.VirtRegBase) {
218 unsigned LocalIndex =
219 getLocalIndex(Pos, Layout.PhyRegBase, RegisterOperandNames.size(),
220 "Local index out of bounds in physical registers");
221 return "PhyReg_" + RegisterOperandNames[LocalIndex];
225 unsigned LocalIndex =
226 getLocalIndex(Pos, Layout.VirtRegBase, RegisterOperandNames.size(),
227 "Local index out of bounds in virtual registers");
228 return "VirtReg_" + RegisterOperandNames[LocalIndex];
231void MIRVocabulary::generateStorage(
const VocabMap &OpcodeMap,
232 const VocabMap &CommonOperandsMap,
233 const VocabMap &PhyRegMap,
241 <<
"; using zero vector. This will result in an error "
243 ++MIRVocabMissCounter;
247 unsigned EmbeddingDim = OpcodeMap.begin()->second.size();
248 std::vector<Embedding> OpcodeEmbeddings(Layout.CommonOperandBase,
252 for (
auto COpcodeName : UniqueBaseOpcodeNames) {
253 if (
auto It = OpcodeMap.find(COpcodeName); It != OpcodeMap.end()) {
254 auto COpcodeIndex = getCanonicalIndexForBaseName(COpcodeName);
255 assert(COpcodeIndex < Layout.CommonOperandBase &&
256 "Canonical index out of bounds");
257 OpcodeEmbeddings[COpcodeIndex] = It->second;
259 handleMissingEntity(COpcodeName);
264 std::vector<Embedding> CommonOperandEmbeddings(std::size(CommonOperandNames),
266 unsigned OperandIndex = 0;
267 for (
const auto &CommonOperandName : CommonOperandNames) {
268 if (
auto It = CommonOperandsMap.find(CommonOperandName.str());
269 It != CommonOperandsMap.end()) {
270 CommonOperandEmbeddings[OperandIndex] = It->second;
272 handleMissingEntity(CommonOperandName);
278 auto createRegisterEmbeddings = [&](
const VocabMap &RegMap) {
279 std::vector<Embedding> RegEmbeddings(
TRI.getNumRegClasses(),
281 unsigned RegOperandIndex = 0;
282 for (
const auto &RegOperandName : RegisterOperandNames) {
283 if (
auto It = RegMap.find(RegOperandName); It != RegMap.end())
284 RegEmbeddings[RegOperandIndex] = It->second;
286 handleMissingEntity(RegOperandName);
289 return RegEmbeddings;
293 std::vector<Embedding> PhyRegEmbeddings = createRegisterEmbeddings(PhyRegMap);
294 std::vector<Embedding> VirtRegEmbeddings =
298 auto scaleVocabSection = [](std::vector<Embedding> &Embeddings,
303 scaleVocabSection(OpcodeEmbeddings,
OpcWeight);
308 std::vector<std::vector<Embedding>> Sections(
309 static_cast<unsigned>(Section::MaxSections));
310 Sections[
static_cast<unsigned>(Section::Opcodes)] =
311 std::move(OpcodeEmbeddings);
312 Sections[
static_cast<unsigned>(Section::CommonOperands)] =
313 std::move(CommonOperandEmbeddings);
314 Sections[
static_cast<unsigned>(Section::PhyRegisters)] =
315 std::move(PhyRegEmbeddings);
316 Sections[
static_cast<unsigned>(Section::VirtRegisters)] =
317 std::move(VirtRegEmbeddings);
322void MIRVocabulary::buildCanonicalOpcodeMapping() {
324 if (!UniqueBaseOpcodeNames.empty())
328 for (
unsigned Opcode = 0; Opcode <
TII.getNumOpcodes(); ++Opcode) {
329 std::string BaseOpcode = extractBaseOpcodeName(
TII.getName(Opcode));
330 UniqueBaseOpcodeNames.insert(BaseOpcode);
333 LLVM_DEBUG(
dbgs() <<
"MIR2Vec: Built canonical mapping for target with "
334 << UniqueBaseOpcodeNames.size()
335 <<
" unique base opcodes\n");
338void MIRVocabulary::buildRegisterOperandMapping() {
340 if (!RegisterOperandNames.empty())
343 for (
unsigned RC = 0; RC <
TRI.getNumRegClasses(); ++RC) {
350 RegisterOperandNames.push_back(ClassName.
str());
354unsigned MIRVocabulary::getCommonOperandIndex(
357 "Expected non-register operand type");
363std::optional<unsigned>
364MIRVocabulary::getRegisterOperandIndex(
Register Reg)
const {
365 assert(!RegisterOperandNames.empty() &&
"Register operand mapping not built");
368 "Expected a physical or virtual register");
376 RegClass =
TRI.getMinimalPhysRegClass(
Reg);
396 <<
"; using zero vector.\n");
397 ++MIRClasslessRegCounter;
401 return RegClass->
getID();
407 assert(Dim > 0 &&
"Dimension must be greater than zero");
409 float DummyVal = 0.1f;
411 VocabMap DummyOpcMap, DummyOperandMap, DummyPhyRegMap, DummyVirtRegMap;
414 for (
unsigned Opcode = 0; Opcode < TII.getNumOpcodes(); ++Opcode) {
416 if (DummyOpcMap.count(BaseOpcode) == 0) {
417 DummyOpcMap[BaseOpcode] =
Embedding(Dim, DummyVal);
423 for (
const auto &CommonOperandName : CommonOperandNames) {
424 DummyOperandMap[CommonOperandName.str()] =
Embedding(Dim, DummyVal);
429 for (
unsigned RC = 0; RC < TRI.getNumRegClasses(); ++RC) {
434 std::string ClassName = TRI.getRegClassName(RegClass);
435 DummyPhyRegMap[ClassName] =
Embedding(Dim, DummyVal);
436 DummyVirtRegMap[ClassName] =
Embedding(Dim, DummyVal);
442 std::move(DummyOpcMap), std::move(DummyOperandMap),
443 std::move(DummyPhyRegMap), std::move(DummyVirtRegMap), TII, TRI, MRI);
452 VocabMap OpcVocab, CommonOperandVocab, PhyRegVocabMap, VirtRegVocabMap;
454 if (
Error Err = readVocabulary(OpcVocab, CommonOperandVocab, PhyRegVocabMap,
456 return std::move(Err);
458 for (
const auto &
F : M) {
459 if (
F.isDeclaration())
462 if (
auto *MF = MMI.getMachineFunction(
F)) {
463 auto &Subtarget = MF->getSubtarget();
464 if (
const auto *
TII = Subtarget.getInstrInfo())
465 if (
const auto *
TRI = Subtarget.getRegisterInfo())
467 std::move(OpcVocab), std::move(CommonOperandVocab),
468 std::move(PhyRegVocabMap), std::move(VirtRegVocabMap), *
TII, *
TRI,
473 "No machine functions found in module");
476Error MIR2VecVocabProvider::readVocabulary(VocabMap &OpcodeVocab,
477 VocabMap &CommonOperandVocab,
478 VocabMap &PhyRegVocabMap,
479 VocabMap &VirtRegVocabMap) {
483 "MIR2Vec vocabulary file path not specified; set it "
484 "using --mir2vec-vocab-path");
490 auto Content = BufOrError.get()->getBuffer();
493 if (!ParsedVocabValue)
496 unsigned OpcodeDim = 0, CommonOperandDim = 0, PhyRegOperandDim = 0,
497 VirtRegOperandDim = 0;
499 "Opcodes", *ParsedVocabValue, OpcodeVocab, OpcodeDim))
503 "CommonOperands", *ParsedVocabValue, CommonOperandVocab,
508 "PhysicalRegisters", *ParsedVocabValue, PhyRegVocabMap,
513 "VirtualRegisters", *ParsedVocabValue, VirtRegVocabMap,
518 if (!(OpcodeDim == CommonOperandDim && CommonOperandDim == PhyRegOperandDim &&
519 PhyRegOperandDim == VirtRegOperandDim)) {
522 "MIR2Vec vocabulary sections have different dimensions");
530 "MIR2Vec Vocabulary Analysis",
false,
true)
536 return "MIR2Vec Vocabulary Analysis";
548 return std::make_unique<SymbolicMIREmbedder>(
MF,
Vocab);
563 const auto &Subtarget =
MF.getSubtarget();
564 const auto *
TII = Subtarget.getInstrInfo();
566 MF.getFunction().getContext().emitError(
567 "MIR2Vec: No TargetInstrInfo available; cannot compute embeddings");
572 for (
const auto &
MI :
MBB) {
574 if (
MI.isDebugInstr())
598std::unique_ptr<SymbolicMIREmbedder>
601 return std::make_unique<SymbolicMIREmbedder>(
MF,
Vocab);
606 if (
MI.isDebugInstr())
614 InstructionEmbedding +=
Vocab[MO];
616 return InstructionEmbedding;
625 "MIR2Vec Vocabulary Printer Pass",
false,
true)
637 auto MIR2VecVocabOrErr =
Analysis.getMIR2VecVocabulary(M);
639 if (!MIR2VecVocabOrErr) {
640 OS <<
"MIR2Vec Vocabulary Printer: Failed to get vocabulary - "
641 <<
toString(MIR2VecVocabOrErr.takeError()) <<
"\n";
645 auto &MIR2VecVocab = *MIR2VecVocabOrErr;
647 for (
const auto &Entry : MIR2VecVocab) {
651 OS <<
"Key: " << MIR2VecVocab.getStringKey(Pos) <<
": ";
667 "MIR2Vec Embedder Printer Pass",
false,
true)
671 "MIR2Vec Embedder Printer Pass",
false,
true)
676 Analysis.getMIR2VecVocabulary(*MF.getFunction().getParent());
677 assert(VocabOrErr &&
"Failed to get MIR2Vec vocabulary");
678 auto &MIRVocab = *VocabOrErr;
682 OS <<
"Error creating MIR2Vec embeddings for function " << MF.getName()
687 OS <<
"MIR2Vec embeddings for machine function " << MF.getName() <<
":\n";
688 OS <<
"Machine Function vector: ";
689 Emb->getMFunctionVector().print(OS);
691 OS <<
"Machine basic block vectors:\n";
693 OS <<
"Machine basic block: " <<
MBB.getFullName() <<
":\n";
694 Emb->getMBBVector(
MBB).print(OS);
697 OS <<
"Machine instruction vectors:\n";
702 if (
MI.isDebugInstr())
705 OS <<
"Machine instruction: ";
707 Emb->getMInstVector(
MI).print(OS);
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
block Block Frequency Analysis
#define clEnumValN(ENUMVAL, FLAGNAME, DESC)
This file builds on the ADT/GraphTraits.h file to build generic depth first graph iterator.
const HexagonInstrInfo * TII
Module.h This file contains the declarations for the Module class.
This file defines the MIR2Vec framework for generating Machine IR embeddings.
Register const TargetRegisterInfo * TRI
#define INITIALIZE_PASS_DEPENDENCY(depName)
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
SmallVector< MachineBasicBlock *, 4 > MBBVector
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
Lightweight error class with error context and mandatory checking.
static ErrorSuccess success()
Create a success value.
Tagged union holding either a T or a Error.
Error takeError()
Take ownership of the stored error.
unsigned getID() const
getID() - Return the register class ID number.
This pass prints the MIR2Vec embeddings for machine functions, basic blocks, and instructions.
MIR2VecPrinterLegacyPass(raw_ostream &OS)
bool runOnMachineFunction(MachineFunction &MF) override
runOnMachineFunction - This method must be overloaded to perform the desired machine code transformat...
Pass to analyze and populate MIR2Vec vocabulary from a module.
This pass prints the embeddings in the MIR2Vec vocabulary.
bool doFinalization(Module &M) override
doFinalization - Virtual method overriden by subclasses to do any necessary clean up after all passes...
bool runOnMachineFunction(MachineFunction &MF) override
runOnMachineFunction - This method must be overloaded to perform the desired machine code transformat...
MIR2VecVocabPrinterLegacyPass(raw_ostream &OS)
LLVM_ABI Expected< mir2vec::MIRVocabulary > getVocabulary(const Module &M)
MachineFunctionPass - This class adapts the FunctionPass interface to allow convenient creation of pa...
Representation of each machine instruction.
MachineOperand class - Representation of each machine instruction operand.
@ MO_Register
Register operand.
MachineRegisterInfo - Keep track of information for virtual and physical registers,...
const TargetRegisterClass * getRegClassOrNull(Register Reg) const
Return the register class of Reg, or null if Reg has not been assigned a register class yet.
static ErrorOr< std::unique_ptr< MemoryBuffer > > getFileOrSTDIN(const Twine &Filename, bool IsText=false, bool RequiresNullTerminator=true, std::optional< Align > Alignment=std::nullopt)
Open the specified file as a MemoryBuffer, or open stdin if the Filename is "-".
A Module instance is used to store all the information related to an LLVM module.
AnalysisType & getAnalysis() const
getAnalysis<AnalysisType>() - This function is used by subclasses to get to the analysis information ...
LLVM_ABI bool match(StringRef String, SmallVectorImpl< StringRef > *Matches=nullptr, std::string *Error=nullptr) const
matches - Match the regex against a given String.
Wrapper class representing virtual and physical registers.
constexpr bool isValid() const
constexpr bool isVirtual() const
Return true if the specified register number is in the virtual register namespace.
constexpr unsigned id() const
constexpr bool isPhysical() const
Return true if the specified register number is in the physical register namespace.
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
Represent a constant reference to a string, i.e.
std::string str() const
Get the contents as an std::string.
constexpr bool empty() const
Check if the string is empty.
char back() const
Get the last character in the string.
StringRef drop_back(size_t N=1) const
Return a StringRef equal to 'this' but with the last N elements dropped.
TargetInstrInfo - Interface to description of machine instruction set.
TargetRegisterInfo base class - We assume that the target defines a static array of TargetRegisterDes...
Generic storage class for section-based vocabularies.
static LLVM_ABI Error parseVocabSection(StringRef Key, const json::Value &ParsedVocabValue, VocabMap &TargetVocab, unsigned &Dim)
Parse a vocabulary section from JSON and populate the target vocabulary map.
unsigned getDimension() const
Get vocabulary dimension.
bool isValid() const
Check if vocabulary is valid (has data)
const unsigned Dimension
Dimension of the embeddings; Captured from the vocabulary.
const MIRVocabulary & Vocab
const float RegOperandWeight
const float CommonOperandWeight
LLVM_ABI Embedding computeEmbeddings() const
Function to compute embeddings.
const float OpcWeight
Weight for opcode embeddings.
const MachineFunction & MF
static LLVM_ABI std::unique_ptr< MIREmbedder > create(MIR2VecKind Mode, const MachineFunction &MF, const MIRVocabulary &Vocab)
Factory method to create an Embedder object of the specified kind Returns nullptr if the requested ki...
LLVM_ABI MIREmbedder(const MachineFunction &MF, const MIRVocabulary &Vocab)
Class for storing and accessing the MIR2Vec vocabulary.
LLVM_ABI unsigned getCanonicalIndexForOperandName(StringRef OperandName) const
LLVM_ABI unsigned getCanonicalIndexForRegisterClass(StringRef RegName, bool IsPhysical=true) const
static LLVM_ABI Expected< MIRVocabulary > create(VocabMap &&OpcMap, VocabMap &&CommonOperandsMap, VocabMap &&PhyRegMap, VocabMap &&VirtRegMap, const TargetInstrInfo &TII, const TargetRegisterInfo &TRI, const MachineRegisterInfo &MRI)
Factory method to create MIRVocabulary from vocabulary map.
static LLVM_ABI std::string extractBaseOpcodeName(StringRef InstrName)
Static method for extracting base opcode names (public for testing)
static LLVM_ABI Expected< MIRVocabulary > createDummyVocabForTest(const TargetInstrInfo &TII, const TargetRegisterInfo &TRI, const MachineRegisterInfo &MRI, unsigned Dim=1)
Create a dummy vocabulary for testing purposes.
LLVM_ABI std::string getStringKey(unsigned Pos) const
Get the string key for a vocabulary entry at the given position.
LLVM_ABI unsigned getCanonicalIndexForBaseName(StringRef BaseName) const
Get indices from opcode or operand names.
static std::unique_ptr< SymbolicMIREmbedder > create(const MachineFunction &MF, const MIRVocabulary &Vocab)
SymbolicMIREmbedder(const MachineFunction &F, const MIRVocabulary &Vocab)
This class implements an extremely fast bulk output stream that can only output to a stream.
OperandType
Operands are tagged with one of the values of this enum.
ValuesClass values(OptsTy... Options)
Helper to build a ValuesClass by forwarding a variable number of arguments as an initializer list to ...
initializer< Ty > init(const Ty &Val)
LLVM_ABI llvm::Expected< Value > parse(llvm::StringRef JSON)
Parses the provided JSON source, or returns a ParseError.
LLVM_ABI llvm::cl::OptionCategory MIR2VecCategory
static cl::opt< float > OpcWeight("mir2vec-opc-weight", cl::init(1.0), cl::desc("Weight for machine opcode embeddings"), cl::cat(MIR2VecCategory))
static cl::opt< float > RegOperandWeight("mir2vec-reg-operand-weight", cl::init(1.0), cl::desc("Weight for register operand embeddings"), cl::cat(MIR2VecCategory))
static cl::opt< bool > PrintAllVocabEntries("mir2vec-print-all-vocab-entries", cl::init(false), cl::desc("Print all vocabulary entries including zero embeddings"), cl::cat(MIR2VecCategory))
ir2vec::Embedding Embedding
cl::opt< MIR2VecKind > MIR2VecEmbeddingKind("mir2vec-kind", cl::values(clEnumValN(MIR2VecKind::Symbolic, "symbolic", "Generate symbolic embeddings for MIR")), cl::init(MIR2VecKind::Symbolic), cl::desc("MIR2Vec embedding kind"), cl::cat(MIR2VecCategory))
static cl::opt< std::string > VocabFile("mir2vec-vocab-path", cl::desc("Path to the vocabulary file for MIR2Vec"), cl::init(""), cl::cat(MIR2VecCategory))
static cl::opt< float > CommonOperandWeight("mir2vec-common-operand-weight", cl::init(1.0), cl::desc("Weight for common operand embeddings"), cl::cat(MIR2VecCategory))
This is an optimization pass for GlobalISel generic memory operations.
Error createFileError(const Twine &F, Error E)
Concatenate a source file path and/or name with an Error.
Error createStringError(std::error_code EC, char const *Fmt, const Ts &... Vals)
Create formatted StringError object.
LLVM_ABI MachineFunctionPass * createMIR2VecPrinterLegacyPass(raw_ostream &OS)
Create a machine pass that prints MIR2Vec embeddings.
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
LLVM_ATTRIBUTE_VISIBILITY_DEFAULT AnalysisKey InnerAnalysisManagerProxy< AnalysisManagerT, IRUnitT, ExtraArgTs... >::Key
LLVM_ABI raw_fd_ostream & errs()
This returns a reference to a raw_ostream for standard error.
LLVM_ABI MachineFunctionPass * createMIR2VecVocabPrinterLegacyPass(raw_ostream &OS)
MIR2VecVocabPrinter pass - This pass prints out the MIR2Vec vocabulary contents to the given stream a...
std::string toString(const APInt &I, unsigned Radix, bool Signed, bool formatAsCLiteral=false, bool UpperCase=true, bool InsertSeparators=false)
iterator_range< df_iterator< T > > depth_first(const T &G)
MCRegisterClass TargetRegisterClass