[llvm] [llvm-ir2vec] Decoupling Vocab loading from initEmbedding (PR #190507)

Nishant Sachdeva via llvm-commits llvm-commits at lists.llvm.org
Tue Apr 7 21:57:43 PDT 2026


https://github.com/nishant-sachdeva updated https://github.com/llvm/llvm-project/pull/190507

>From 4a483069863f4a7e3e6891150994385506b54ed6 Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Sun, 5 Apr 2026 13:03:28 +0530
Subject: [PATCH 1/6] Decoupling Vocab loading from initEmbedding in order to
 save time during entire dataset processing

---
 .../bindings/ir2vec-getBBEmbMap.py            |   3 +-
 .../llvm-ir2vec/bindings/ir2vec-getFuncEmb.py |   3 +-
 .../bindings/ir2vec-getFuncEmbMap.py          |   3 +-
 .../bindings/ir2vec-getFuncNames.py           |   3 +-
 .../bindings/ir2vec-getInstEmbMap.py          |   3 +-
 .../bindings/ir2vec-initEmbedding.py          | 136 +++++++++++++-----
 llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp  |  57 +++++---
 llvm/tools/llvm-ir2vec/lib/Utils.cpp          |  12 +-
 llvm/tools/llvm-ir2vec/lib/Utils.h            |  12 +-
 llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp        |   6 +-
 10 files changed, 171 insertions(+), 67 deletions(-)

diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getBBEmbMap.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getBBEmbMap.py
index 57fe1b032c6fd..963e0d0adeca5 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getBBEmbMap.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getBBEmbMap.py
@@ -6,8 +6,9 @@
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
+vocab = ir2vec.loadVocab(vocab_path)
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 
 # Success case
diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmb.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmb.py
index bb58b59e54825..b7ac30f689da6 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmb.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmb.py
@@ -6,8 +6,9 @@
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
+vocab = ir2vec.loadVocab(vocab_path)
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 
 # Success case
diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmbMap.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmbMap.py
index 03c7baab4d349..6bc3adbca80ac 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmbMap.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncEmbMap.py
@@ -6,8 +6,9 @@
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
+vocab = ir2vec.loadVocab(vocab_path)
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 
 # Success case
diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncNames.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncNames.py
index 4bb19e9bd8115..f4420bae6caac 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncNames.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getFuncNames.py
@@ -6,8 +6,9 @@
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
+vocab = ir2vec.loadVocab(vocab_path)
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 
 # Success case
diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getInstEmbMap.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getInstEmbMap.py
index b04222ce4943b..a76d23ed84146 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getInstEmbMap.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-getInstEmbMap.py
@@ -6,8 +6,9 @@
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
+vocab = ir2vec.loadVocab(vocab_path)
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 
 # Success case
diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
index dcdcd90da0847..f7e37e7e1bdd5 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
@@ -2,37 +2,90 @@
 
 import sys
 import ir2vec
+import tempfile
+import os
 
 ll_file = sys.argv[1]
 vocab_path = sys.argv[2]
 
-# Success case
+# ============================================================
+# loadVocab tests
+# ============================================================
+
+# Success: Load a valid vocabulary
+vocab = ir2vec.loadVocab(vocab_path)
+print(f"VOCAB: {type(vocab).__name__}")
+# CHECK: VOCAB: Vocab
+
+# Error: Empty vocab path
+try:
+    ir2vec.loadVocab("")
+except ValueError:
+    print("ERROR: Empty vocab path")
+# CHECK: ERROR: Empty vocab path
+
+# Error: Non-existent vocab file
+try:
+    ir2vec.loadVocab("/nonexistent/path/bad.json")
+except ValueError:
+    print("ERROR: Invalid vocab path")
+# CHECK: ERROR: Invalid vocab path
+
+# Error: Malformed JSON vocab file
+with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
+    f.write("{ this is not valid json }")
+    bad_vocab = f.name
+try:
+    ir2vec.loadVocab(bad_vocab)
+except ValueError:
+    print("ERROR: Malformed vocab file")
+finally:
+    os.unlink(bad_vocab)
+# CHECK: ERROR: Malformed vocab file
+
+# Error: Wrong type for vocab path (not a string)
+try:
+    ir2vec.loadVocab(42)
+except TypeError:
+    print("ERROR: Invalid vocab path type")
+# CHECK: ERROR: Invalid vocab path type
+
+# ============================================================
+# initEmbedding tests
+# ============================================================
+
+# Success: Create embedding tool with valid inputs
 tool = ir2vec.initEmbedding(
-    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+    filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab
 )
 print(f"SUCCESS: {type(tool).__name__}")
 # CHECK: SUCCESS: IR2VecTool
 
-# Error: Invalid mode
+# Success: Default mode (Symbolic) when mode is omitted
+tool_default = ir2vec.initEmbedding(filename=ll_file, vocab=vocab)
+print(f"DEFAULT MODE: {type(tool_default).__name__}")
+# CHECK: DEFAULT MODE: IR2VecTool
+
+# Error: Invalid mode (string instead of IR2VecKind enum)
 try:
-    ir2vec.initEmbedding(filename=ll_file, mode="invalid", vocabPath=vocab_path)
+    ir2vec.initEmbedding(filename=ll_file, mode="invalid", vocab=vocab)
 except TypeError:
-    print("ERROR: Invalid mode")
-# CHECK: ERROR: Invalid mode
+    print("ERROR: Invalid mode type")
+# CHECK: ERROR: Invalid mode type
 
-# Error: Empty vocab path
+# Error: Invalid mode (integer instead of IR2VecKind enum)
 try:
-    ir2vec.initEmbedding(
-        filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=""
-    )
-except ValueError:
-    print("ERROR: Empty vocab path")
-# CHECK: ERROR: Empty vocab path
+    ir2vec.initEmbedding(filename=ll_file, mode=99, vocab=vocab)
+except TypeError:
+    print("ERROR: Invalid mode int")
+# CHECK: ERROR: Invalid mode int
 
-# Error: Invalid file
+# Error: Non-existent IR file
 try:
     ir2vec.initEmbedding(
-        filename="/bad.ll", mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
+        filename="/nonexistent/bad.ll",
+        mode=ir2vec.IR2VecKind.Symbolic,
+        vocab=vocab,
     )
 except ValueError:
     print("ERROR: Invalid file")
@@ -40,35 +93,44 @@
 
 # Error: Empty filename
 try:
-    ir2vec.initEmbedding(
-        filename="", mode=ir2vec.IR2VecKind.Symbolic, vocabPath=vocab_path
-    )
+    ir2vec.initEmbedding(filename="", mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab)
 except ValueError:
     print("ERROR: Empty filename")
 # CHECK: ERROR: Empty filename
 
-# Error: Invalid vocab file
+# Error: Wrong type for vocab (string instead of Vocab object)
 try:
     ir2vec.initEmbedding(
-        filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath="/bad.json"
+        filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab="vocab.json"
     )
-except ValueError:
-    print("ERROR: Invalid vocab")
-# CHECK: ERROR: Invalid vocab
+except TypeError:
+    print("ERROR: Vocab is string")
+# CHECK: ERROR: Vocab is string
 
-# Error: Malformed JSON vocab
-import tempfile
-import os
+# Error: Wrong type for vocab (None)
+try:
+    ir2vec.initEmbedding(filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=None)
+except TypeError:
+    print("ERROR: Vocab is None")
+# CHECK: ERROR: Vocab is None
 
-with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
-    f.write("{ this is not valid json }")
-    bad_vocab = f.name
+# Error: Wrong type for vocab (integer)
 try:
-    ir2vec.initEmbedding(
-        filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocabPath=bad_vocab
-    )
-except ValueError:
-    print("ERROR: Invalid vocab file")
-finally:
-    os.unlink(bad_vocab)
-# CHECK: ERROR: Invalid vocab file
+    ir2vec.initEmbedding(filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic, vocab=123)
+except TypeError:
+    print("ERROR: Vocab is int")
+# CHECK: ERROR: Vocab is int
+
+# Error: Missing vocab argument entirely
+try:
+    ir2vec.initEmbedding(filename=ll_file, mode=ir2vec.IR2VecKind.Symbolic)
+except TypeError:
+    print("ERROR: Vocab missing")
+# CHECK: ERROR: Vocab missing
+
+# Error: Wrong type for filename (not a string)
+try:
+    ir2vec.initEmbedding(filename=42, mode=ir2vec.IR2VecKind.Symbolic, vocab=vocab)
+except TypeError:
+    print("ERROR: Filename is int")
+# CHECK: ERROR: Filename is int
diff --git a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
index 4c77e933b3ac2..d95fdca508a5c 100644
--- a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
+++ b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
@@ -38,6 +38,25 @@ std::unique_ptr<Module> getLLVMIR(const std::string &Filename,
   return M;
 }
 
+class PyVocab {
+private:
+  std::shared_ptr<Vocabulary> Vocab;
+
+public:
+  PyVocab(const std::string &VocabPath) {
+    if (VocabPath.empty())
+      throw nb::value_error("Empty vocabulary path not allowed");
+    auto VocabOrErr = ir2vec::loadVocabulary(VocabPath);
+    if (!VocabOrErr)
+      throw nb::value_error(
+          ("Failed to load vocabulary: " + toString(VocabOrErr.takeError()))
+              .c_str());
+    Vocab = std::move(*VocabOrErr);
+  }
+
+  std::shared_ptr<Vocabulary> getVocab() const { return Vocab; }
+};
+
 class PyIR2VecTool {
 private:
   std::unique_ptr<LLVMContext> Ctx;
@@ -46,22 +65,19 @@ class PyIR2VecTool {
   IR2VecKind OutputEmbeddingMode;
 
 public:
-  PyIR2VecTool(const std::string &Filename, IR2VecKind Mode,
-               const std::string &VocabPath) {
-    OutputEmbeddingMode = Mode;
+  PyIR2VecTool(const std::string &Filename, IR2VecKind Mode, PyVocab &Vocab) {
+    if (!Vocab.getVocab())
+      throw nb::value_error("Vocabulary object is not initialized");
 
-    if (VocabPath.empty())
-      throw nb::value_error("Empty Vocab Path not allowed");
+    if (Filename.empty())
+      throw nb::value_error("Empty filename not allowed");
+
+    OutputEmbeddingMode = Mode;
 
     Ctx = std::make_unique<LLVMContext>();
     M = getLLVMIR(Filename, *Ctx);
     Tool = std::make_unique<IR2VecTool>(*M);
-
-    if (auto Err = Tool->initializeVocabulary(VocabPath)) {
-      throw nb::value_error(("Failed to initialize IR2Vec vocabulary: " +
-                             toString(std::move(Err)))
-                                .c_str());
-    }
+    Tool->setVocabulary(Vocab.getVocab());
   }
 
   nb::list getFuncNames() {
@@ -201,9 +217,19 @@ NB_MODULE(ir2vec, m) {
              "Flow-aware encodings (includes data/control flow)")
       .export_values();
 
+  nb::class_<PyVocab>(m, "Vocab");
+
+  m.def(
+      "loadVocab",
+      [](const std::string &vocabPath) {
+        return std::make_unique<PyVocab>(vocabPath);
+      },
+      nb::arg("vocabPath"), "Load an IR2Vec vocabulary from a JSON file",
+      nb::rv_policy::take_ownership);
+
   nb::class_<PyIR2VecTool>(m, "IR2VecTool")
-      .def(nb::init<const std::string &, IR2VecKind, const std::string &>(),
-           nb::arg("filename"), nb::arg("mode"), nb::arg("vocabPath"))
+      .def(nb::init<const std::string &, IR2VecKind, PyVocab &>(),
+           nb::arg("filename"), nb::arg("mode"), nb::arg("vocab"))
       .def("getFuncNames", &PyIR2VecTool::getFuncNames,
            "Get list of all defined functions in the module\n"
            "Returns: list[str] - Function names")
@@ -228,9 +254,8 @@ NB_MODULE(ir2vec, m) {
 
   m.def(
       "initEmbedding",
-      [](const std::string &filename, IR2VecKind mode,
-         const std::string &vocabPath) {
-        return std::make_unique<PyIR2VecTool>(filename, mode, vocabPath);
+      [](const std::string &filename, IR2VecKind mode, PyVocab &vocab) {
+        return std::make_unique<PyIR2VecTool>(filename, mode, vocab);
       },
       nb::arg("filename"), nb::arg("mode"), nb::arg("vocabPath"),
       nb::rv_policy::take_ownership);
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.cpp b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
index 1334aa1783583..0d78d93ef71b1 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.cpp
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
@@ -43,17 +43,21 @@ namespace llvm {
 
 namespace ir2vec {
 
-Error IR2VecTool::initializeVocabulary(StringRef VocabPath) {
+Expected<std::shared_ptr<Vocabulary>> loadVocabulary(StringRef VocabPath) {
   auto VocabOrErr = Vocabulary::fromFile(VocabPath);
   if (!VocabOrErr)
     return VocabOrErr.takeError();
 
-  Vocab = std::make_unique<Vocabulary>(std::move(*VocabOrErr));
+  auto V = std::make_shared<Vocabulary>(std::move(*VocabOrErr));
 
-  if (!Vocab->isValid())
+  if (!V->isValid())
     return createStringError(errc::invalid_argument,
                              "Failed to initialize IR2Vec vocabulary");
-  return Error::success();
+  return V;
+}
+
+void IR2VecTool::setVocabulary(std::shared_ptr<Vocabulary> V) {
+  Vocab = std::move(V);
 }
 
 TripletResult IR2VecTool::generateTriplets(const Function &F) const {
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.h b/llvm/tools/llvm-ir2vec/lib/Utils.h
index ae2f931a90cf9..be0f319297069 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.h
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.h
@@ -84,12 +84,15 @@ enum RelationType {
   ArgRelation = 2   ///< Instruction to operand relationship (ArgRelation + N)
 };
 
+/// Load an IR2Vec vocabulary from a JSON file on disk.
+Expected<std::shared_ptr<Vocabulary>> loadVocabulary(StringRef VocabPath);
+
 /// Helper class for collecting IR triplets and generating embeddings
 class IR2VecTool {
 private:
   Module &M;
   ModuleAnalysisManager MAM;
-  std::unique_ptr<Vocabulary> Vocab;
+  std::shared_ptr<Vocabulary> Vocab;
 
 public:
   explicit IR2VecTool(Module &M) : M(M) {}
@@ -98,8 +101,11 @@ class IR2VecTool {
   Expected<std::unique_ptr<Embedder>>
   createIR2VecEmbedder(const Function &F, IR2VecKind Kind) const;
 
-  /// Initialize the IR2Vec vocabulary from the specified file path.
-  Error initializeVocabulary(StringRef VocabPath);
+  /// Load vocabulary from a shared pointer. This allows sharing the same
+  /// vocabulary instance across multiple IR2VecTool instances, which is useful
+  /// for generating embeddings for multiple functions without needing to reload
+  /// the vocabulary each time.
+  void setVocabulary(std::shared_ptr<Vocabulary> V);
 
   /// Generate triplets for a single function
   /// Returns a TripletResult with:
diff --git a/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp b/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
index cf290ba931023..78a2e3f657705 100644
--- a/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
+++ b/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
@@ -161,8 +161,10 @@ static Error processModule(Module &M, raw_ostream &OS) {
           "You may need to set it using --ir2vec-vocab-path");
     }
 
-    if (Error Err = Tool.initializeVocabulary(VocabFile))
-      return Err;
+    auto VocabOrErr = ir2vec::loadVocabulary(VocabFile);
+    if (!VocabOrErr)
+      return VocabOrErr.takeError();
+    Tool.setVocabulary(std::move(*VocabOrErr));
 
     if (!FunctionName.empty()) {
       // Process single function

>From a0c43d64953ffabcf92988a1a1835faff6822819 Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Mon, 6 Apr 2026 12:37:22 +0530
Subject: [PATCH 2/6] Nit commit - adding more protection around valid vocab,
 and more error handling

---
 .../tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py   | 5 -----
 llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp             | 5 +++--
 llvm/tools/llvm-ir2vec/lib/Utils.cpp                     | 9 ++++++++-
 llvm/tools/llvm-ir2vec/lib/Utils.h                       | 2 +-
 llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp                   | 9 +++++----
 5 files changed, 17 insertions(+), 13 deletions(-)

diff --git a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
index f7e37e7e1bdd5..071a9dc202dd0 100644
--- a/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
+++ b/llvm/test/tools/llvm-ir2vec/bindings/ir2vec-initEmbedding.py
@@ -61,11 +61,6 @@
 print(f"SUCCESS: {type(tool).__name__}")
 # CHECK: SUCCESS: IR2VecTool
 
-# Success: Default mode (Symbolic) when mode is omitted
-tool_default = ir2vec.initEmbedding(filename=ll_file, vocab=vocab)
-print(f"DEFAULT MODE: {type(tool_default).__name__}")
-# CHECK: DEFAULT MODE: IR2VecTool
-
 # Error: Invalid mode (string instead of IR2VecKind enum)
 try:
     ir2vec.initEmbedding(filename=ll_file, mode="invalid", vocab=vocab)
diff --git a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
index d95fdca508a5c..00db465390cef 100644
--- a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
+++ b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
@@ -77,7 +77,8 @@ class PyIR2VecTool {
     Ctx = std::make_unique<LLVMContext>();
     M = getLLVMIR(Filename, *Ctx);
     Tool = std::make_unique<IR2VecTool>(*M);
-    Tool->setVocabulary(Vocab.getVocab());
+    if (auto Err = Tool->setVocabulary(Vocab.getVocab()))
+      throw nb::value_error(toString(std::move(Err)).c_str());
   }
 
   nb::list getFuncNames() {
@@ -257,6 +258,6 @@ NB_MODULE(ir2vec, m) {
       [](const std::string &filename, IR2VecKind mode, PyVocab &vocab) {
         return std::make_unique<PyIR2VecTool>(filename, mode, vocab);
       },
-      nb::arg("filename"), nb::arg("mode"), nb::arg("vocabPath"),
+      nb::arg("filename"), nb::arg("mode"), nb::arg("vocab"),
       nb::rv_policy::take_ownership);
 }
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.cpp b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
index 0d78d93ef71b1..cdf3e0196d92c 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.cpp
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
@@ -56,8 +56,15 @@ Expected<std::shared_ptr<Vocabulary>> loadVocabulary(StringRef VocabPath) {
   return V;
 }
 
-void IR2VecTool::setVocabulary(std::shared_ptr<Vocabulary> V) {
+Error IR2VecTool::setVocabulary(std::shared_ptr<Vocabulary> V) {
+  if (!V)
+    return createStringError(errc::invalid_argument,
+                             "Null pointer provided for vocabulary. Will not set IR2VecTool vocabulary.");
+  if (!V->isValid())
+    return createStringError(errc::invalid_argument,
+                             "Vocabulary is not valid. Will not set IR2VecTool vocabulary.");
   Vocab = std::move(V);
+  return Error::success();
 }
 
 TripletResult IR2VecTool::generateTriplets(const Function &F) const {
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.h b/llvm/tools/llvm-ir2vec/lib/Utils.h
index be0f319297069..f3ece02997f23 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.h
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.h
@@ -105,7 +105,7 @@ class IR2VecTool {
   /// vocabulary instance across multiple IR2VecTool instances, which is useful
   /// for generating embeddings for multiple functions without needing to reload
   /// the vocabulary each time.
-  void setVocabulary(std::shared_ptr<Vocabulary> V);
+  Error setVocabulary(std::shared_ptr<Vocabulary> V);
 
   /// Generate triplets for a single function
   /// Returns a TripletResult with:
diff --git a/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp b/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
index 78a2e3f657705..372f1d06b6b6e 100644
--- a/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
+++ b/llvm/tools/llvm-ir2vec/llvm-ir2vec.cpp
@@ -161,10 +161,11 @@ static Error processModule(Module &M, raw_ostream &OS) {
           "You may need to set it using --ir2vec-vocab-path");
     }
 
-    auto VocabOrErr = ir2vec::loadVocabulary(VocabFile);
-    if (!VocabOrErr)
-      return VocabOrErr.takeError();
-    Tool.setVocabulary(std::move(*VocabOrErr));
+    std::shared_ptr<Vocabulary> Vocab;
+    if (auto Err = ir2vec::loadVocabulary(VocabFile).moveInto(Vocab))
+      return Err;
+    if (auto Err = Tool.setVocabulary(std::move(Vocab)))
+      return Err;
 
     if (!FunctionName.empty()) {
       // Process single function

>From 41523c4a680487ce14f1aa917aece8410887f324 Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Mon, 6 Apr 2026 12:42:04 +0530
Subject: [PATCH 3/6] Nit commit - formatting fixup

---
 llvm/tools/llvm-ir2vec/lib/Utils.cpp | 8 +++++---
 1 file changed, 5 insertions(+), 3 deletions(-)

diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.cpp b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
index cdf3e0196d92c..d37fd5e24de8e 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.cpp
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.cpp
@@ -59,10 +59,12 @@ Expected<std::shared_ptr<Vocabulary>> loadVocabulary(StringRef VocabPath) {
 Error IR2VecTool::setVocabulary(std::shared_ptr<Vocabulary> V) {
   if (!V)
     return createStringError(errc::invalid_argument,
-                             "Null pointer provided for vocabulary. Will not set IR2VecTool vocabulary.");
+                             "Null pointer provided for vocabulary. Will not "
+                             "set IR2VecTool vocabulary.");
   if (!V->isValid())
-    return createStringError(errc::invalid_argument,
-                             "Vocabulary is not valid. Will not set IR2VecTool vocabulary.");
+    return createStringError(
+        errc::invalid_argument,
+        "Vocabulary is not valid. Will not set IR2VecTool vocabulary.");
   Vocab = std::move(V);
   return Error::success();
 }

>From ec83fd690a20765fcc8664e6a1f5ca2f19f25ad3 Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Mon, 6 Apr 2026 16:30:17 +0530
Subject: [PATCH 4/6] Nit commit - Adding a note about thread safety around the
 vocab API

---
 llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp |  6 ++++++
 llvm/tools/llvm-ir2vec/lib/Utils.h           | 13 +++++++++----
 2 files changed, 15 insertions(+), 4 deletions(-)

diff --git a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
index 00db465390cef..16b86ab3fad0f 100644
--- a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
+++ b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
@@ -65,6 +65,11 @@ class PyIR2VecTool {
   IR2VecKind OutputEmbeddingMode;
 
 public:
+
+  /// \note
+  /// In the currently exposed API, the vocabulary is set once at construction and there is no
+  /// public interface to call this again. Callers should treat
+  /// the vocabulary as immutable for the lifetime of the tool instance.
   PyIR2VecTool(const std::string &Filename, IR2VecKind Mode, PyVocab &Vocab) {
     if (!Vocab.getVocab())
       throw nb::value_error("Vocabulary object is not initialized");
@@ -77,6 +82,7 @@ class PyIR2VecTool {
     Ctx = std::make_unique<LLVMContext>();
     M = getLLVMIR(Filename, *Ctx);
     Tool = std::make_unique<IR2VecTool>(*M);
+
     if (auto Err = Tool->setVocabulary(Vocab.getVocab()))
       throw nb::value_error(toString(std::move(Err)).c_str());
   }
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.h b/llvm/tools/llvm-ir2vec/lib/Utils.h
index f3ece02997f23..c6d1682942d94 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.h
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.h
@@ -92,6 +92,11 @@ class IR2VecTool {
 private:
   Module &M;
   ModuleAnalysisManager MAM;
+
+  /// \note The API around vocab object is not thread-safe.
+  /// Specifically, calling setVocabulary() on an instance while
+  /// another thread reading the Vocab object with the same instance
+  /// can cause a data race on this internal shared_ptr<Vocabulary> member.
   std::shared_ptr<Vocabulary> Vocab;
 
 public:
@@ -101,10 +106,10 @@ class IR2VecTool {
   Expected<std::unique_ptr<Embedder>>
   createIR2VecEmbedder(const Function &F, IR2VecKind Kind) const;
 
-  /// Load vocabulary from a shared pointer. This allows sharing the same
-  /// vocabulary instance across multiple IR2VecTool instances, which is useful
-  /// for generating embeddings for multiple functions without needing to reload
-  /// the vocabulary each time.
+  /// Sets the vocabulary for this tool instance.
+  /// This allows sharing the same vocabulary instance across multiple IR2VecTool
+  /// instances, which is useful for generating embeddings for multiple functions
+  /// without needing to reload the vocabulary each time.
   Error setVocabulary(std::shared_ptr<Vocabulary> V);
 
   /// Generate triplets for a single function

>From 7177c5eb807a77ecaf411dbc95891e5ad6aded9c Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Mon, 6 Apr 2026 16:37:51 +0530
Subject: [PATCH 5/6] Nit commit - formatting fixups

---
 llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp | 5 ++---
 llvm/tools/llvm-ir2vec/lib/Utils.h           | 6 +++---
 2 files changed, 5 insertions(+), 6 deletions(-)

diff --git a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
index 16b86ab3fad0f..ba05cdfdf02b0 100644
--- a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
+++ b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
@@ -65,10 +65,9 @@ class PyIR2VecTool {
   IR2VecKind OutputEmbeddingMode;
 
 public:
-
   /// \note
-  /// In the currently exposed API, the vocabulary is set once at construction and there is no
-  /// public interface to call this again. Callers should treat
+  /// In the currently exposed API, the vocabulary is set once at construction
+  /// and there is no public interface to call this again. Callers should treat
   /// the vocabulary as immutable for the lifetime of the tool instance.
   PyIR2VecTool(const std::string &Filename, IR2VecKind Mode, PyVocab &Vocab) {
     if (!Vocab.getVocab())
diff --git a/llvm/tools/llvm-ir2vec/lib/Utils.h b/llvm/tools/llvm-ir2vec/lib/Utils.h
index c6d1682942d94..9856ae6f9c168 100644
--- a/llvm/tools/llvm-ir2vec/lib/Utils.h
+++ b/llvm/tools/llvm-ir2vec/lib/Utils.h
@@ -107,9 +107,9 @@ class IR2VecTool {
   createIR2VecEmbedder(const Function &F, IR2VecKind Kind) const;
 
   /// Sets the vocabulary for this tool instance.
-  /// This allows sharing the same vocabulary instance across multiple IR2VecTool
-  /// instances, which is useful for generating embeddings for multiple functions
-  /// without needing to reload the vocabulary each time.
+  /// This allows sharing the same vocabulary instance across multiple
+  /// IR2VecTool instances, which is useful for generating embeddings for
+  /// multiple functions without needing to reload the vocabulary each time.
   Error setVocabulary(std::shared_ptr<Vocabulary> V);
 
   /// Generate triplets for a single function

>From 9aec63b6db3f501f74f191b1372665911a2f3fe1 Mon Sep 17 00:00:00 2001
From: nishant-sachdeva <nishant.sachdeva at research.iiit.ac.in>
Date: Wed, 8 Apr 2026 10:26:27 +0530
Subject: [PATCH 6/6] Nit commit - adding explicit keyword to the PyVocab
 constructor

---
 llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
index ba05cdfdf02b0..fa61c478e5c6d 100644
--- a/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
+++ b/llvm/tools/llvm-ir2vec/Bindings/PyIR2Vec.cpp
@@ -43,7 +43,7 @@ class PyVocab {
   std::shared_ptr<Vocabulary> Vocab;
 
 public:
-  PyVocab(const std::string &VocabPath) {
+  explicit PyVocab(const std::string &VocabPath) {
     if (VocabPath.empty())
       throw nb::value_error("Empty vocabulary path not allowed");
     auto VocabOrErr = ir2vec::loadVocabulary(VocabPath);



More information about the llvm-commits mailing list