1//===- MIR2Vec.cpp - Implementation of MIR2Vec ---------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM
4// Exceptions. See the LICENSE file for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8///
9/// \file
10/// This file implements the MIR2Vec algorithm for Machine IR embeddings.
11///
12//===----------------------------------------------------------------------===//
13
14#include "llvm/CodeGen/MIR2Vec.h"
15#include "llvm/ADT/DepthFirstIterator.h"
16#include "llvm/ADT/Statistic.h"
17#include "llvm/CodeGen/TargetInstrInfo.h"
18#include "llvm/IR/Module.h"
19#include "llvm/InitializePasses.h"
20#include "llvm/Pass.h"
21#include "llvm/Support/Errc.h"
22#include "llvm/Support/MemoryBuffer.h"
23#include "llvm/Support/Regex.h"
24
25using namespace llvm;
26using namespace mir2vec;
27
28#define DEBUG_TYPE "mir2vec"
29
30STATISTIC(MIRVocabMissCounter,
31 "Number of lookups to MIR entities not present in the vocabulary");
32STATISTIC(MIRClasslessRegCounter,
33 "Number of register operands with no register class");
34
35namespace llvm {
36namespace mir2vec {
37cl::OptionCategory MIR2VecCategory("MIR2Vec Options");
38
39// FIXME: Use a default vocab when not specified
40static cl::opt<std::string>
41 VocabFile("mir2vec-vocab-path",
42 cl::desc("Path to the vocabulary file for MIR2Vec"), cl::init(Val: ""),
43 cl::cat(MIR2VecCategory));
44static cl::opt<float>
45 OpcWeight("mir2vec-opc-weight", cl::init(Val: 1.0),
46 cl::desc("Weight for machine opcode embeddings"),
47 cl::cat(MIR2VecCategory));
48static cl::opt<float>
49 CommonOperandWeight("mir2vec-common-operand-weight", cl::init(Val: 1.0),
50 cl::desc("Weight for common operand embeddings"),
51 cl::cat(MIR2VecCategory));
52static cl::opt<float>
53 RegOperandWeight("mir2vec-reg-operand-weight", cl::init(Val: 1.0),
54 cl::desc("Weight for register operand embeddings"),
55 cl::cat(MIR2VecCategory));
56cl::opt<MIR2VecKind> MIR2VecEmbeddingKind(
57 "mir2vec-kind",
58 cl::values(clEnumValN(MIR2VecKind::Symbolic, "symbolic",
59 "Generate symbolic embeddings for MIR")),
60 cl::init(Val: MIR2VecKind::Symbolic), cl::desc("MIR2Vec embedding kind"),
61 cl::cat(MIR2VecCategory));
62
63static cl::opt<bool> PrintAllVocabEntries(
64 "mir2vec-print-all-vocab-entries", cl::init(Val: false),
65 cl::desc("Print all vocabulary entries including zero embeddings"),
66 cl::cat(MIR2VecCategory));
67
68} // namespace mir2vec
69} // namespace llvm
70
71//===----------------------------------------------------------------------===//
72// Vocabulary
73//===----------------------------------------------------------------------===//
74
75MIRVocabulary::MIRVocabulary(VocabMap &&OpcodeMap, VocabMap &&CommonOperandMap,
76 VocabMap &&PhysicalRegisterMap,
77 VocabMap &&VirtualRegisterMap,
78 const TargetInstrInfo &TII,
79 const TargetRegisterInfo &TRI,
80 const MachineRegisterInfo &MRI)
81 : TII(TII), TRI(TRI), MRI(MRI) {
82 buildCanonicalOpcodeMapping();
83 unsigned CanonicalOpcodeCount = UniqueBaseOpcodeNames.size();
84 assert(CanonicalOpcodeCount > 0 &&
85 "No canonical opcodes found for target - invalid vocabulary");
86
87 buildRegisterOperandMapping();
88
89 // Define layout of vocabulary sections
90 Layout.OpcodeBase = 0;
91 Layout.CommonOperandBase = CanonicalOpcodeCount;
92 // We expect same classes for physical and virtual registers
93 Layout.PhyRegBase = Layout.CommonOperandBase + std::size(CommonOperandNames);
94 Layout.VirtRegBase = Layout.PhyRegBase + RegisterOperandNames.size();
95
96 generateStorage(OpcodeMap, CommonOperandMap, PhyRegMap: PhysicalRegisterMap,
97 VirtRegMap: VirtualRegisterMap);
98 Layout.TotalEntries = Storage.size();
99}
100
101Expected<MIRVocabulary>
102MIRVocabulary::create(VocabMap &&OpcodeMap, VocabMap &&CommonOperandMap,
103 VocabMap &&PhyRegMap, VocabMap &&VirtRegMap,
104 const TargetInstrInfo &TII, const TargetRegisterInfo &TRI,
105 const MachineRegisterInfo &MRI) {
106 if (OpcodeMap.empty() || CommonOperandMap.empty() || PhyRegMap.empty() ||
107 VirtRegMap.empty())
108 return createStringError(EC: errc::invalid_argument,
109 S: "Empty vocabulary entries provided");
110
111 MIRVocabulary Vocab(std::move(OpcodeMap), std::move(CommonOperandMap),
112 std::move(PhyRegMap), std::move(VirtRegMap), TII, TRI,
113 MRI);
114
115 // Validate Storage after construction
116 if (!Vocab.Storage.isValid())
117 return createStringError(EC: errc::invalid_argument,
118 S: "Failed to create valid vocabulary storage");
119 Vocab.ZeroEmbedding = Embedding(Vocab.Storage.getDimension(), 0.0);
120 return std::move(Vocab);
121}
122
123std::string MIRVocabulary::extractBaseOpcodeName(StringRef InstrName) {
124 // Extract base instruction name using regex to capture letters and
125 // underscores Examples: "ADD32rr" -> "ADD", "ARITH_FENCE" -> "ARITH_FENCE"
126 //
127 // TODO: Consider more sophisticated extraction:
128 // - Handle complex prefixes like "AVX1_SETALLONES" correctly (Currently, it
129 // would naively map to "AVX")
130 // - Extract width suffixes (8,16,32,64) as separate features
131 // - Capture addressing mode suffixes (r,i,m,ri,etc.) for better analysis
132 // (Currently, instances like "MOV32mi" map to "MOV", but "ADDPDrr" would map
133 // to "ADDPDrr")
134
135 assert(!InstrName.empty() && "Instruction name should not be empty");
136
137 // Use regex to extract initial sequence of letters and underscores
138 static const Regex BaseOpcodeRegex("([a-zA-Z_]+)");
139 SmallVector<StringRef, 2> Matches;
140
141 if (BaseOpcodeRegex.match(String: InstrName, Matches: &Matches) && Matches.size() > 1) {
142 StringRef Match = Matches[1];
143 // Trim trailing underscores
144 while (!Match.empty() && Match.back() == '_')
145 Match = Match.drop_back();
146 return Match.str();
147 }
148
149 // Fallback to original name if no pattern matches
150 return InstrName.str();
151}
152
153unsigned MIRVocabulary::getCanonicalIndexForBaseName(StringRef BaseName) const {
154 assert(!UniqueBaseOpcodeNames.empty() && "Canonical mapping not built");
155 auto It = std::find(first: UniqueBaseOpcodeNames.begin(),
156 last: UniqueBaseOpcodeNames.end(), val: BaseName.str());
157 assert(It != UniqueBaseOpcodeNames.end() &&
158 "Base name not found in unique opcodes");
159 return std::distance(first: UniqueBaseOpcodeNames.begin(), last: It);
160}
161
162unsigned MIRVocabulary::getCanonicalOpcodeIndex(unsigned Opcode) const {
163 auto BaseOpcode = extractBaseOpcodeName(InstrName: TII.getName(Opcode));
164 return getCanonicalIndexForBaseName(BaseName: BaseOpcode);
165}
166
167unsigned
168MIRVocabulary::getCanonicalIndexForOperandName(StringRef OperandName) const {
169 auto It = std::find(first: std::begin(arr: CommonOperandNames),
170 last: std::end(arr: CommonOperandNames), val: OperandName);
171 assert(It != std::end(CommonOperandNames) &&
172 "Operand name not found in common operands");
173 return Layout.CommonOperandBase +
174 std::distance(first: std::begin(arr: CommonOperandNames), last: It);
175}
176
177unsigned
178MIRVocabulary::getCanonicalIndexForRegisterClass(StringRef RegName,
179 bool IsPhysical) const {
180 auto It = std::find(first: RegisterOperandNames.begin(), last: RegisterOperandNames.end(),
181 val: RegName);
182 assert(It != RegisterOperandNames.end() &&
183 "Register name not found in register operands");
184 unsigned LocalIndex = std::distance(first: RegisterOperandNames.begin(), last: It);
185 return (IsPhysical ? Layout.PhyRegBase : Layout.VirtRegBase) + LocalIndex;
186}
187
188std::string MIRVocabulary::getStringKey(unsigned Pos) const {
189 assert(Pos < Layout.TotalEntries && "Position out of bounds in vocabulary");
190
191 // Handle opcodes section
192 if (Pos < Layout.CommonOperandBase) {
193 // Convert canonical index back to base opcode name
194 auto It = UniqueBaseOpcodeNames.begin();
195 std::advance(i&: It, n: Pos);
196 assert(It != UniqueBaseOpcodeNames.end() &&
197 "Canonical index out of bounds in opcode section");
198 return *It;
199 }
200
201 auto getLocalIndex = [](unsigned Pos, size_t BaseOffset, size_t Bound,
202 const char *Msg) {
203 unsigned LocalIndex = Pos - BaseOffset;
204 assert(LocalIndex < Bound && Msg);
205 return LocalIndex;
206 };
207
208 // Handle common operands section
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();
214 }
215
216 // Handle physical registers section
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];
222 }
223
224 // Handle virtual registers section
225 unsigned LocalIndex =
226 getLocalIndex(Pos, Layout.VirtRegBase, RegisterOperandNames.size(),
227 "Local index out of bounds in virtual registers");
228 return "VirtReg_" + RegisterOperandNames[LocalIndex];
229}
230
231void MIRVocabulary::generateStorage(const VocabMap &OpcodeMap,
232 const VocabMap &CommonOperandsMap,
233 const VocabMap &PhyRegMap,
234 const VocabMap &VirtRegMap) {
235
236 // Helper for handling missing entities in the vocabulary.
237 // Currently, we use a zero vector. In the future, we will throw an error to
238 // ensure that *all* known entities are present in the vocabulary.
239 auto handleMissingEntity = [](StringRef Key) {
240 LLVM_DEBUG(errs() << "MIR2Vec: Missing vocabulary entry for " << Key
241 << "; using zero vector. This will result in an error "
242 "in the future.\n");
243 ++MIRVocabMissCounter;
244 };
245
246 // Initialize opcode embeddings section
247 unsigned EmbeddingDim = OpcodeMap.begin()->second.size();
248 std::vector<Embedding> OpcodeEmbeddings(Layout.CommonOperandBase,
249 Embedding(EmbeddingDim));
250
251 // Populate opcode embeddings using canonical mapping
252 for (auto COpcodeName : UniqueBaseOpcodeNames) {
253 if (auto It = OpcodeMap.find(x: COpcodeName); It != OpcodeMap.end()) {
254 auto COpcodeIndex = getCanonicalIndexForBaseName(BaseName: COpcodeName);
255 assert(COpcodeIndex < Layout.CommonOperandBase &&
256 "Canonical index out of bounds");
257 OpcodeEmbeddings[COpcodeIndex] = It->second;
258 } else {
259 handleMissingEntity(COpcodeName);
260 }
261 }
262
263 // Initialize common operand embeddings section
264 std::vector<Embedding> CommonOperandEmbeddings(std::size(CommonOperandNames),
265 Embedding(EmbeddingDim));
266 unsigned OperandIndex = 0;
267 for (const auto &CommonOperandName : CommonOperandNames) {
268 if (auto It = CommonOperandsMap.find(x: CommonOperandName.str());
269 It != CommonOperandsMap.end()) {
270 CommonOperandEmbeddings[OperandIndex] = It->second;
271 } else {
272 handleMissingEntity(CommonOperandName);
273 }
274 ++OperandIndex;
275 }
276
277 // Helper lambda for creating register operand embeddings
278 auto createRegisterEmbeddings = [&](const VocabMap &RegMap) {
279 std::vector<Embedding> RegEmbeddings(TRI.getNumRegClasses(),
280 Embedding(EmbeddingDim));
281 unsigned RegOperandIndex = 0;
282 for (const auto &RegOperandName : RegisterOperandNames) {
283 if (auto It = RegMap.find(x: RegOperandName); It != RegMap.end())
284 RegEmbeddings[RegOperandIndex] = It->second;
285 else
286 handleMissingEntity(RegOperandName);
287 ++RegOperandIndex;
288 }
289 return RegEmbeddings;
290 };
291
292 // Initialize register operand embeddings sections
293 std::vector<Embedding> PhyRegEmbeddings = createRegisterEmbeddings(PhyRegMap);
294 std::vector<Embedding> VirtRegEmbeddings =
295 createRegisterEmbeddings(VirtRegMap);
296
297 // Scale the vocabulary sections based on the provided weights
298 auto scaleVocabSection = [](std::vector<Embedding> &Embeddings,
299 double Weight) {
300 for (auto &Embedding : Embeddings)
301 Embedding *= Weight;
302 };
303 scaleVocabSection(OpcodeEmbeddings, OpcWeight);
304 scaleVocabSection(CommonOperandEmbeddings, CommonOperandWeight);
305 scaleVocabSection(PhyRegEmbeddings, RegOperandWeight);
306 scaleVocabSection(VirtRegEmbeddings, RegOperandWeight);
307
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);
318
319 Storage = ir2vec::VocabStorage(std::move(Sections));
320}
321
322void MIRVocabulary::buildCanonicalOpcodeMapping() {
323 // Check if already built
324 if (!UniqueBaseOpcodeNames.empty())
325 return;
326
327 // Build mapping from opcodes to canonical base opcode indices
328 for (unsigned Opcode = 0; Opcode < TII.getNumOpcodes(); ++Opcode) {
329 std::string BaseOpcode = extractBaseOpcodeName(InstrName: TII.getName(Opcode));
330 UniqueBaseOpcodeNames.insert(x: BaseOpcode);
331 }
332
333 LLVM_DEBUG(dbgs() << "MIR2Vec: Built canonical mapping for target with "
334 << UniqueBaseOpcodeNames.size()
335 << " unique base opcodes\n");
336}
337
338void MIRVocabulary::buildRegisterOperandMapping() {
339 // Check if already built
340 if (!RegisterOperandNames.empty())
341 return;
342
343 for (unsigned RC = 0; RC < TRI.getNumRegClasses(); ++RC) {
344 const TargetRegisterClass *RegClass = TRI.getRegClass(i: RC);
345 if (!RegClass)
346 continue;
347
348 // Get the register class name
349 StringRef ClassName = TRI.getRegClassName(Class: RegClass);
350 RegisterOperandNames.push_back(Elt: ClassName.str());
351 }
352}
353
354unsigned MIRVocabulary::getCommonOperandIndex(
355 MachineOperand::MachineOperandType OperandType) const {
356 assert(OperandType != MachineOperand::MO_Register &&
357 "Expected non-register operand type");
358 assert(OperandType > MachineOperand::MO_Register &&
359 OperandType < MachineOperand::MO_Last && "Operand type out of bounds");
360 return static_cast<unsigned>(OperandType) - 1;
361}
362
363std::optional<unsigned>
364MIRVocabulary::getRegisterOperandIndex(Register Reg) const {
365 assert(!RegisterOperandNames.empty() && "Register operand mapping not built");
366 assert(Reg.isValid() && "Invalid register; not expected here");
367 assert((Reg.isPhysical() || Reg.isVirtual()) &&
368 "Expected a physical or virtual register");
369
370 const TargetRegisterClass *RegClass = nullptr;
371
372 // For physical registers, use TRI to get minimal register class as a
373 // physical register can belong to multiple classes. For virtual
374 // registers, use MRI to uniquely identify the assigned register class.
375 if (Reg.isPhysical())
376 RegClass = TRI.getMinimalPhysRegClass(Reg);
377 else
378 RegClass = MRI.getRegClassOrNull(Reg);
379
380 // Not every register belongs to a register class. This can happen for
381 // physical registers, e.g. X86's $mxcsr and $fpcw or AMDGPU's $mode, for
382 // which getMinimalPhysRegClass() returns nullptr. It can also happen for
383 // generic virtual registers that have not yet been through (or completed)
384 // GlobalISel's register bank selection, and thus carry an LLT or a
385 // RegisterBank instead of a TargetRegisterClass, for which
386 // getRegClassOrNull() returns nullptr.
387 // TODO: Avoid special-casing these registers at every use site. Classless
388 // registers currently fall back to a zero embedding in operator[] and to
389 // VirtRegBase in getEntityIDForRegister(), which is the same ad-hoc handling
390 // the invalid/stack-slot cases already get. Give them a real vocabulary
391 // representation instead -- e.g. an explicit "no register class" entry, or
392 // keying generic vregs on their LLT/RegisterBank -- so that the lookup is
393 // total and the callers need no fallbacks.
394 if (!RegClass) {
395 LLVM_DEBUG(errs() << "MIR2Vec: No register class for register " << Reg.id()
396 << "; using zero vector.\n");
397 ++MIRClasslessRegCounter;
398 return std::nullopt;
399 }
400
401 return RegClass->getID();
402}
403
404Expected<MIRVocabulary> MIRVocabulary::createDummyVocabForTest(
405 const TargetInstrInfo &TII, const TargetRegisterInfo &TRI,
406 const MachineRegisterInfo &MRI, unsigned Dim) {
407 assert(Dim > 0 && "Dimension must be greater than zero");
408
409 float DummyVal = 0.1f;
410
411 VocabMap DummyOpcMap, DummyOperandMap, DummyPhyRegMap, DummyVirtRegMap;
412
413 // Process opcodes directly without creating temporary vocabulary
414 for (unsigned Opcode = 0; Opcode < TII.getNumOpcodes(); ++Opcode) {
415 std::string BaseOpcode = extractBaseOpcodeName(InstrName: TII.getName(Opcode));
416 if (DummyOpcMap.count(x: BaseOpcode) == 0) { // Only add if not already present
417 DummyOpcMap[BaseOpcode] = Embedding(Dim, DummyVal);
418 DummyVal += 0.1f;
419 }
420 }
421
422 // Add common operands
423 for (const auto &CommonOperandName : CommonOperandNames) {
424 DummyOperandMap[CommonOperandName.str()] = Embedding(Dim, DummyVal);
425 DummyVal += 0.1f;
426 }
427
428 // Process register classes directly
429 for (unsigned RC = 0; RC < TRI.getNumRegClasses(); ++RC) {
430 const TargetRegisterClass *RegClass = TRI.getRegClass(i: RC);
431 if (!RegClass)
432 continue;
433
434 std::string ClassName = TRI.getRegClassName(Class: RegClass);
435 DummyPhyRegMap[ClassName] = Embedding(Dim, DummyVal);
436 DummyVirtRegMap[ClassName] = Embedding(Dim, DummyVal);
437 DummyVal += 0.1f;
438 }
439
440 // Create vocabulary directly without temporary instance
441 return MIRVocabulary::create(
442 OpcodeMap: std::move(DummyOpcMap), CommonOperandMap: std::move(DummyOperandMap),
443 PhyRegMap: std::move(DummyPhyRegMap), VirtRegMap: std::move(DummyVirtRegMap), TII, TRI, MRI);
444}
445
446//===----------------------------------------------------------------------===//
447// MIR2VecVocabProvider and MIR2VecVocabLegacyAnalysis
448//===----------------------------------------------------------------------===//
449
450Expected<mir2vec::MIRVocabulary>
451MIR2VecVocabProvider::getVocabulary(const Module &M) {
452 VocabMap OpcVocab, CommonOperandVocab, PhyRegVocabMap, VirtRegVocabMap;
453
454 if (Error Err = readVocabulary(OpcVocab, CommonOperandVocab, PhyRegVocabMap,
455 VirtRegVocabMap))
456 return std::move(Err);
457
458 for (const auto &F : M) {
459 if (F.isDeclaration())
460 continue;
461
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())
466 return mir2vec::MIRVocabulary::create(
467 OpcodeMap: std::move(OpcVocab), CommonOperandMap: std::move(CommonOperandVocab),
468 PhyRegMap: std::move(PhyRegVocabMap), VirtRegMap: std::move(VirtRegVocabMap), TII: *TII, TRI: *TRI,
469 MRI: MF->getRegInfo());
470 }
471 }
472 return createStringError(EC: errc::invalid_argument,
473 S: "No machine functions found in module");
474}
475
476Error MIR2VecVocabProvider::readVocabulary(VocabMap &OpcodeVocab,
477 VocabMap &CommonOperandVocab,
478 VocabMap &PhyRegVocabMap,
479 VocabMap &VirtRegVocabMap) {
480 if (VocabFile.empty())
481 return createStringError(
482 EC: errc::invalid_argument,
483 S: "MIR2Vec vocabulary file path not specified; set it "
484 "using --mir2vec-vocab-path");
485
486 auto BufOrError = MemoryBuffer::getFileOrSTDIN(Filename: VocabFile, /*IsText=*/true);
487 if (!BufOrError)
488 return createFileError(F: VocabFile, EC: BufOrError.getError());
489
490 auto Content = BufOrError.get()->getBuffer();
491
492 Expected<json::Value> ParsedVocabValue = json::parse(JSON: Content);
493 if (!ParsedVocabValue)
494 return ParsedVocabValue.takeError();
495
496 unsigned OpcodeDim = 0, CommonOperandDim = 0, PhyRegOperandDim = 0,
497 VirtRegOperandDim = 0;
498 if (auto Err = ir2vec::VocabStorage::parseVocabSection(
499 Key: "Opcodes", ParsedVocabValue: *ParsedVocabValue, TargetVocab&: OpcodeVocab, Dim&: OpcodeDim))
500 return Err;
501
502 if (auto Err = ir2vec::VocabStorage::parseVocabSection(
503 Key: "CommonOperands", ParsedVocabValue: *ParsedVocabValue, TargetVocab&: CommonOperandVocab,
504 Dim&: CommonOperandDim))
505 return Err;
506
507 if (auto Err = ir2vec::VocabStorage::parseVocabSection(
508 Key: "PhysicalRegisters", ParsedVocabValue: *ParsedVocabValue, TargetVocab&: PhyRegVocabMap,
509 Dim&: PhyRegOperandDim))
510 return Err;
511
512 if (auto Err = ir2vec::VocabStorage::parseVocabSection(
513 Key: "VirtualRegisters", ParsedVocabValue: *ParsedVocabValue, TargetVocab&: VirtRegVocabMap,
514 Dim&: VirtRegOperandDim))
515 return Err;
516
517 // All sections must have the same embedding dimension
518 if (!(OpcodeDim == CommonOperandDim && CommonOperandDim == PhyRegOperandDim &&
519 PhyRegOperandDim == VirtRegOperandDim)) {
520 return createStringError(
521 EC: errc::illegal_byte_sequence,
522 S: "MIR2Vec vocabulary sections have different dimensions");
523 }
524
525 return Error::success();
526}
527
528char MIR2VecVocabLegacyAnalysis::ID = 0;
529INITIALIZE_PASS_BEGIN(MIR2VecVocabLegacyAnalysis, "mir2vec-vocab-analysis",
530 "MIR2Vec Vocabulary Analysis", false, true)
531INITIALIZE_PASS_DEPENDENCY(MachineModuleInfoWrapperPass)
532INITIALIZE_PASS_END(MIR2VecVocabLegacyAnalysis, "mir2vec-vocab-analysis",
533 "MIR2Vec Vocabulary Analysis", false, true)
534
535StringRef MIR2VecVocabLegacyAnalysis::getPassName() const {
536 return "MIR2Vec Vocabulary Analysis";
537}
538
539//===----------------------------------------------------------------------===//
540// MIREmbedder and its subclasses
541//===----------------------------------------------------------------------===//
542
543std::unique_ptr<MIREmbedder> MIREmbedder::create(MIR2VecKind Mode,
544 const MachineFunction &MF,
545 const MIRVocabulary &Vocab) {
546 switch (Mode) {
547 case MIR2VecKind::Symbolic:
548 return std::make_unique<SymbolicMIREmbedder>(args: MF, args: Vocab);
549 }
550 return nullptr;
551}
552
553MIREmbedder::MIREmbedder(const MachineFunction &MF, const MIRVocabulary &Vocab)
554 : MF(MF), Vocab(Vocab), Dimension(Vocab.getDimension()),
555 OpcWeight(mir2vec::OpcWeight),
556 CommonOperandWeight(mir2vec::CommonOperandWeight),
557 RegOperandWeight(mir2vec::RegOperandWeight) {}
558
559Embedding MIREmbedder::computeEmbeddings(const MachineBasicBlock &MBB) const {
560 Embedding MBBVector(Dimension, 0);
561
562 // Get instruction info for opcode name resolution
563 const auto &Subtarget = MF.getSubtarget();
564 const auto *TII = Subtarget.getInstrInfo();
565 if (!TII) {
566 MF.getFunction().getContext().emitError(
567 ErrorStr: "MIR2Vec: No TargetInstrInfo available; cannot compute embeddings");
568 return MBBVector;
569 }
570
571 // Process each machine instruction in the basic block
572 for (const auto &MI : MBB) {
573 // Skip debug instructions and other metadata
574 if (MI.isDebugInstr())
575 continue;
576 MBBVector += computeEmbeddings(MI);
577 }
578
579 return MBBVector;
580}
581
582Embedding MIREmbedder::computeEmbeddings() const {
583 Embedding MFuncVector(Dimension, 0);
584
585 if (MF.empty())
586 return MFuncVector;
587
588 // Consider all reachable machine basic blocks in the function
589 for (const auto *MBB : depth_first(G: &MF))
590 MFuncVector += computeEmbeddings(MBB: *MBB);
591 return MFuncVector;
592}
593
594SymbolicMIREmbedder::SymbolicMIREmbedder(const MachineFunction &MF,
595 const MIRVocabulary &Vocab)
596 : MIREmbedder(MF, Vocab) {}
597
598std::unique_ptr<SymbolicMIREmbedder>
599SymbolicMIREmbedder::create(const MachineFunction &MF,
600 const MIRVocabulary &Vocab) {
601 return std::make_unique<SymbolicMIREmbedder>(args: MF, args: Vocab);
602}
603
604Embedding SymbolicMIREmbedder::computeEmbeddings(const MachineInstr &MI) const {
605 // Skip debug instructions and other metadata
606 if (MI.isDebugInstr())
607 return Embedding(Dimension, 0);
608
609 // Opcode embedding
610 Embedding InstructionEmbedding = Vocab[MI.getOpcode()];
611
612 // Add operand contributions
613 for (const MachineOperand &MO : MI.operands())
614 InstructionEmbedding += Vocab[MO];
615
616 return InstructionEmbedding;
617}
618
619//===----------------------------------------------------------------------===//
620// Printer Passes
621//===----------------------------------------------------------------------===//
622
623char MIR2VecVocabPrinterLegacyPass::ID = 0;
624INITIALIZE_PASS_BEGIN(MIR2VecVocabPrinterLegacyPass, "print-mir2vec-vocab",
625 "MIR2Vec Vocabulary Printer Pass", false, true)
626INITIALIZE_PASS_DEPENDENCY(MIR2VecVocabLegacyAnalysis)
627INITIALIZE_PASS_DEPENDENCY(MachineModuleInfoWrapperPass)
628INITIALIZE_PASS_END(MIR2VecVocabPrinterLegacyPass, "print-mir2vec-vocab",
629 "MIR2Vec Vocabulary Printer Pass", false, true)
630
631bool MIR2VecVocabPrinterLegacyPass::runOnMachineFunction(MachineFunction &MF) {
632 return false;
633}
634
635bool MIR2VecVocabPrinterLegacyPass::doFinalization(Module &M) {
636 auto &Analysis = getAnalysis<MIR2VecVocabLegacyAnalysis>();
637 auto MIR2VecVocabOrErr = Analysis.getMIR2VecVocabulary(M);
638
639 if (!MIR2VecVocabOrErr) {
640 OS << "MIR2Vec Vocabulary Printer: Failed to get vocabulary - "
641 << toString(E: MIR2VecVocabOrErr.takeError()) << "\n";
642 return false;
643 }
644
645 auto &MIR2VecVocab = *MIR2VecVocabOrErr;
646 unsigned Pos = 0;
647 for (const auto &Entry : MIR2VecVocab) {
648 // Skip zero embeddings to avoid printing entries not in the vocabulary.
649 // This makes the output stable across changes to the opcode list.
650 if (PrintAllVocabEntries || !Entry.isZero()) {
651 OS << "Key: " << MIR2VecVocab.getStringKey(Pos) << ": ";
652 Entry.print(OS);
653 }
654 ++Pos;
655 }
656
657 return false;
658}
659
660MachineFunctionPass *
661llvm::createMIR2VecVocabPrinterLegacyPass(raw_ostream &OS) {
662 return new MIR2VecVocabPrinterLegacyPass(OS);
663}
664
665char MIR2VecPrinterLegacyPass::ID = 0;
666INITIALIZE_PASS_BEGIN(MIR2VecPrinterLegacyPass, "print-mir2vec",
667 "MIR2Vec Embedder Printer Pass", false, true)
668INITIALIZE_PASS_DEPENDENCY(MIR2VecVocabLegacyAnalysis)
669INITIALIZE_PASS_DEPENDENCY(MachineModuleInfoWrapperPass)
670INITIALIZE_PASS_END(MIR2VecPrinterLegacyPass, "print-mir2vec",
671 "MIR2Vec Embedder Printer Pass", false, true)
672
673bool MIR2VecPrinterLegacyPass::runOnMachineFunction(MachineFunction &MF) {
674 auto &Analysis = getAnalysis<MIR2VecVocabLegacyAnalysis>();
675 auto VocabOrErr =
676 Analysis.getMIR2VecVocabulary(M: *MF.getFunction().getParent());
677 assert(VocabOrErr && "Failed to get MIR2Vec vocabulary");
678 auto &MIRVocab = *VocabOrErr;
679
680 auto Emb = mir2vec::MIREmbedder::create(Mode: MIR2VecEmbeddingKind, MF, Vocab: MIRVocab);
681 if (!Emb) {
682 OS << "Error creating MIR2Vec embeddings for function " << MF.getName()
683 << "\n";
684 return false;
685 }
686
687 OS << "MIR2Vec embeddings for machine function " << MF.getName() << ":\n";
688 OS << "Machine Function vector: ";
689 Emb->getMFunctionVector().print(OS);
690
691 OS << "Machine basic block vectors:\n";
692 for (const MachineBasicBlock &MBB : MF) {
693 OS << "Machine basic block: " << MBB.getFullName() << ":\n";
694 Emb->getMBBVector(MBB).print(OS);
695 }
696
697 OS << "Machine instruction vectors:\n";
698 for (const MachineBasicBlock &MBB : MF) {
699 for (const MachineInstr &MI : MBB) {
700 // Skip debug instructions as they are not
701 // embedded
702 if (MI.isDebugInstr())
703 continue;
704
705 OS << "Machine instruction: ";
706 MI.print(OS);
707 Emb->getMInstVector(MI).print(OS);
708 }
709 }
710
711 return false;
712}
713
714MachineFunctionPass *llvm::createMIR2VecPrinterLegacyPass(raw_ostream &OS) {
715 return new MIR2VecPrinterLegacyPass(OS);
716}
717